rebase
This commit is contained in:
348
app/api/sessions.py
Normal file
348
app/api/sessions.py
Normal file
@@ -0,0 +1,348 @@
|
||||
"""Sessions API — state retrieval, orchestrator iterate stream, world_builder/editor streams, retry/rollback."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.api.deps import get_current_user, get_settings_dict
|
||||
from app.core.llm import LlmClient, MockLlmClient
|
||||
from app.core.logging import get_logger
|
||||
from app.db import get_db
|
||||
from app.engine.game_master import run_iteration
|
||||
from app.engine.sse import SseEmitter
|
||||
from app.engine.world_builder import run_world_builder
|
||||
from app.engine.world_editor import run_world_editor
|
||||
from app.models import Entity, Step, World, WorldPreset
|
||||
from app.schemas import AnswerRequest, IterateRequest
|
||||
|
||||
_logger = get_logger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/sessions", tags=["sessions"])
|
||||
|
||||
|
||||
def _llm_factory(settings: dict) -> LlmClient | MockLlmClient:
|
||||
api_url = settings.get("llm.api_url", "")
|
||||
if not api_url:
|
||||
return MockLlmClient()
|
||||
return LlmClient.from_settings(settings)
|
||||
|
||||
|
||||
async def _load_world(db: AsyncSession, world_id: uuid.UUID, user) -> World:
|
||||
world = (
|
||||
await db.execute(select(World).where(World.id == world_id))
|
||||
).scalar_one_or_none()
|
||||
if world is None:
|
||||
raise HTTPException(404, "not_found")
|
||||
if world.owner_id != user.id and not user.is_admin:
|
||||
raise HTTPException(403, "not_owner")
|
||||
return world
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# State retrieval
|
||||
# --------------------------------------------------------------------------- #
|
||||
@router.get("/worlds/{world_id}/state")
|
||||
async def get_state(
|
||||
world_id: uuid.UUID,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
) -> dict:
|
||||
"""Return current session state for the play page."""
|
||||
world = await _load_world(db, world_id, user)
|
||||
recent = (
|
||||
await db.execute(
|
||||
select(Step)
|
||||
.where(Step.world_id == world.id, Step.deleted_at.is_(None))
|
||||
.order_by(Step.sequence_number.desc())
|
||||
.limit(10)
|
||||
)
|
||||
).scalars().all()
|
||||
recent_steps = [
|
||||
{
|
||||
"id": str(s.id), "sequence_number": s.sequence_number,
|
||||
"player_action": s.player_action, "scene_text": s.scene_text,
|
||||
"suggested_actions": s.suggested_actions, "created_at": s.created_at.isoformat(),
|
||||
}
|
||||
for s in reversed(recent)
|
||||
]
|
||||
next_actions = recent_steps[-1]["suggested_actions"] if recent_steps else []
|
||||
if world.intro_scene and not recent_steps:
|
||||
next_actions = []
|
||||
return {
|
||||
"world": {
|
||||
"id": str(world.id), "name": world.name, "current_time": world.current_time,
|
||||
"language": world.language, "intro_scene": world.intro_scene,
|
||||
},
|
||||
"environment": world.environment,
|
||||
"recent_steps": recent_steps,
|
||||
"next_actions": next_actions,
|
||||
}
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# World builder stream (SSE)
|
||||
# --------------------------------------------------------------------------- #
|
||||
@router.get("/worlds/{world_id}/builder/stream")
|
||||
async def builder_stream(
|
||||
world_id: uuid.UUID,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
settings: dict = Depends(get_settings_dict),
|
||||
) -> StreamingResponse:
|
||||
world = await _load_world(db, world_id, user)
|
||||
preset: WorldPreset | None = None
|
||||
if world.preset_id:
|
||||
preset = (
|
||||
await db.execute(select(WorldPreset).where(WorldPreset.id == world.preset_id))
|
||||
).scalar_one_or_none()
|
||||
emitter = SseEmitter()
|
||||
player_name = (world.environment or {}).get("player", {}).get("name", "Hero")
|
||||
notes = world.description
|
||||
llm = _llm_factory(settings)
|
||||
|
||||
async def run_bg():
|
||||
async with _session_scope() as bg_db:
|
||||
# Reload world in this session
|
||||
bg_world = (
|
||||
await bg_db.execute(select(World).where(World.id == world.id))
|
||||
).scalar_one()
|
||||
await run_world_builder(
|
||||
db=bg_db, world=bg_world, player_name=player_name, notes=notes,
|
||||
llm=llm, sse=emitter, preset=preset,
|
||||
)
|
||||
|
||||
import asyncio
|
||||
|
||||
task = asyncio.create_task(run_bg())
|
||||
|
||||
async def gen():
|
||||
try:
|
||||
async for evt in emitter.stream():
|
||||
yield _format_sse(evt)
|
||||
finally:
|
||||
await task
|
||||
|
||||
return StreamingResponse(
|
||||
gen(), media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# World editor stream (SSE)
|
||||
# --------------------------------------------------------------------------- #
|
||||
@router.get("/worlds/{world_id}/editor/stream")
|
||||
async def editor_stream(
|
||||
world_id: uuid.UUID,
|
||||
instruction: str = Query(...),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
settings: dict = Depends(get_settings_dict),
|
||||
) -> StreamingResponse:
|
||||
world = await _load_world(db, world_id, user)
|
||||
emitter = SseEmitter()
|
||||
llm = _llm_factory(settings)
|
||||
|
||||
async def run_bg():
|
||||
async with _session_scope() as bg_db:
|
||||
bg_world = (
|
||||
await bg_db.execute(select(World).where(World.id == world.id))
|
||||
).scalar_one()
|
||||
await run_world_editor(
|
||||
db=bg_db, world=bg_world, instruction=instruction, llm=llm, sse=emitter,
|
||||
)
|
||||
|
||||
import asyncio
|
||||
|
||||
task = asyncio.create_task(run_bg())
|
||||
|
||||
async def gen():
|
||||
try:
|
||||
async for evt in emitter.stream():
|
||||
yield _format_sse(evt)
|
||||
finally:
|
||||
await task
|
||||
|
||||
return StreamingResponse(
|
||||
gen(), media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Orchestrator iterate
|
||||
# --------------------------------------------------------------------------- #
|
||||
@router.post("/worlds/{world_id}/iterate", response_model=dict, status_code=status.HTTP_202_ACCEPTED)
|
||||
async def iterate(
|
||||
world_id: uuid.UUID,
|
||||
body: IterateRequest,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
settings: dict = Depends(get_settings_dict),
|
||||
) -> dict:
|
||||
world = await _load_world(db, world_id, user)
|
||||
if world.status != "ready":
|
||||
raise HTTPException(422, "world_not_ready")
|
||||
# Compute next sequence number
|
||||
last_seq = (
|
||||
await db.execute(
|
||||
select(Step.sequence_number)
|
||||
.where(Step.world_id == world.id, Step.deleted_at.is_(None))
|
||||
.order_by(Step.sequence_number.desc())
|
||||
.limit(1)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
next_seq = (last_seq or 0) + 1
|
||||
step = Step(
|
||||
world_id=world.id,
|
||||
sequence_number=next_seq,
|
||||
player_action=body.action,
|
||||
status="pending",
|
||||
)
|
||||
db.add(step)
|
||||
await db.commit()
|
||||
await db.refresh(step)
|
||||
return {
|
||||
"stream_url": f"/api/sessions/worlds/{world.id}/iterate/stream?step_id={step.id}",
|
||||
"step_id": str(step.id),
|
||||
}
|
||||
|
||||
|
||||
@router.get("/worlds/{world_id}/iterate/stream")
|
||||
async def iterate_stream(
|
||||
world_id: uuid.UUID,
|
||||
step_id: uuid.UUID = Query(...),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
settings: dict = Depends(get_settings_dict),
|
||||
) -> StreamingResponse:
|
||||
world = await _load_world(db, world_id, user)
|
||||
step = (
|
||||
await db.execute(select(Step).where(Step.id == step_id, Step.world_id == world.id))
|
||||
).scalar_one_or_none()
|
||||
if step is None:
|
||||
raise HTTPException(404, "step not found")
|
||||
emitter = SseEmitter()
|
||||
llm = _llm_factory(settings)
|
||||
|
||||
async def run_bg():
|
||||
async with _session_scope() as bg_db:
|
||||
bg_world = (
|
||||
await bg_db.execute(select(World).where(World.id == world.id))
|
||||
).scalar_one()
|
||||
bg_step = (
|
||||
await bg_db.execute(select(Step).where(Step.id == step.id))
|
||||
).scalar_one()
|
||||
await run_iteration(db=bg_db, world=bg_world, step=bg_step, llm=llm, sse=emitter)
|
||||
|
||||
import asyncio
|
||||
|
||||
task = asyncio.create_task(run_bg())
|
||||
|
||||
async def gen():
|
||||
try:
|
||||
async for evt in emitter.stream():
|
||||
yield _format_sse(evt)
|
||||
finally:
|
||||
await task
|
||||
|
||||
return StreamingResponse(
|
||||
gen(), media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Retry / rollback
|
||||
# --------------------------------------------------------------------------- #
|
||||
@router.post("/worlds/{world_id}/retry", response_model=dict)
|
||||
async def retry_last(
|
||||
world_id: uuid.UUID,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
) -> dict:
|
||||
"""Soft-delete the last step and create a new one with the same action."""
|
||||
world = await _load_world(db, world_id, user)
|
||||
last = (
|
||||
await db.execute(
|
||||
select(Step)
|
||||
.where(Step.world_id == world.id, Step.deleted_at.is_(None))
|
||||
.order_by(Step.sequence_number.desc())
|
||||
.limit(1)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
if last is None:
|
||||
raise HTTPException(404, "no_step_to_retry")
|
||||
last.deleted_at = datetime.now(timezone.utc)
|
||||
new_step = Step(
|
||||
world_id=world.id,
|
||||
sequence_number=last.sequence_number + 1,
|
||||
player_action=last.player_action,
|
||||
status="pending",
|
||||
)
|
||||
db.add(new_step)
|
||||
await db.commit()
|
||||
await db.refresh(new_step)
|
||||
return {
|
||||
"step_id": str(new_step.id),
|
||||
"stream_url": f"/api/sessions/worlds/{world.id}/iterate/stream?step_id={new_step.id}",
|
||||
}
|
||||
|
||||
|
||||
@router.post("/worlds/{world_id}/rollback", status_code=status.HTTP_204_NO_CONTENT, response_class=Response)
|
||||
async def rollback_last(
|
||||
world_id: uuid.UUID,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
user=Depends(get_current_user),
|
||||
) -> Response:
|
||||
"""Soft-delete the last step."""
|
||||
world = await _load_world(db, world_id, user)
|
||||
last = (
|
||||
await db.execute(
|
||||
select(Step)
|
||||
.where(Step.world_id == world.id, Step.deleted_at.is_(None))
|
||||
.order_by(Step.sequence_number.desc())
|
||||
.limit(1)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
if last is None:
|
||||
raise HTTPException(404, "no_step_to_rollback")
|
||||
last.deleted_at = datetime.now(timezone.utc)
|
||||
await db.commit()
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Helpers
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _format_sse(evt: dict[str, str]) -> str:
|
||||
"""Format an SSE event dict into the wire format."""
|
||||
lines = []
|
||||
if "id" in evt:
|
||||
lines.append(f"id: {evt['id']}")
|
||||
if "event" in evt:
|
||||
lines.append(f"event: {evt['event']}")
|
||||
if "data" in evt:
|
||||
# Split multi-line data
|
||||
for chunk in evt["data"].split("\n"):
|
||||
lines.append(f"data: {chunk}")
|
||||
lines.append("")
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
async def _session_scope():
|
||||
"""Open a fresh DB session for the background task."""
|
||||
from app.db import get_sessionmaker
|
||||
|
||||
sm = get_sessionmaker()
|
||||
async with sm() as s:
|
||||
yield s
|
||||
Reference in New Issue
Block a user