Files
ai-rpg/app/engine/game_master.py
2026-06-20 19:13:05 +03:00

315 lines
12 KiB
Python

"""Game Master (orchestrator) — three-phase iteration engine.
Phase 1: Planner + Executor (tool-calling loop until submit_plan)
Phase 2: Writer (single LLM call with submit_step tool)
Phase 3: Persist + Deferred triggers + Summary + Suggest actions
"""
from __future__ import annotations
import json
import uuid
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.llm import LlmClient, MockLlmClient
from app.core.logging import get_logger
from app.core.rag import rag_add
from app.core.settings_service import get_all_settings
from app.core.time_utils import advance_time, summarize_schemas
from app.engine.context import (
build_orchestrator_phase1_context,
build_orchestrator_phase2_context,
build_orchestrator_phase3_suggest_context,
build_summary_context,
)
from app.engine.sse import SseEmitter
from app.engine.tools.base import ToolContext, get_registry
from app.engine.world_builder import _run_tool_loop
from app.models import DeferredTrigger, Step, StoryEntry, World
_logger = get_logger(__name__)
async def run_iteration(
*,
db: AsyncSession,
world: World,
step: Step,
llm: LlmClient | MockLlmClient,
sse: SseEmitter,
) -> None:
"""Run the full three-phase orchestrator iteration for a single step."""
settings = await get_all_settings(db)
try:
# ============ Phase 1 ============
await sse.emit("phase_start", {"phase": 1, "name": "planner_executor"})
messages = await build_orchestrator_phase1_context(
db=db, world=world, player_action=step.player_action, settings=settings,
)
phase1_result = await _run_tool_loop(
db=db, world=world, llm=llm, sse=sse,
stage="orchestrator_phase1",
system_prompt=messages[0]["content"],
terminal_tool="submit_plan",
max_substeps=int(settings.get("game.max_substeps_per_iteration", 8)),
settings=settings,
)
await sse.emit("phase_end", {"phase": 1, "duration_ms": 0})
if not phase1_result or not phase1_result.get("ok"):
# Force-completion: synthesize a minimal plan
phase1_result = {
"ok": True,
"data": {
"plan": "The action was processed but no explicit plan was submitted.",
"summary": [],
"offscreen_events": [],
},
}
plan = phase1_result["data"].get("plan", "")
summary = phase1_result["data"].get("summary", [])
offscreen_events = phase1_result["data"].get("offscreen_events", [])
# Persist tool_calls_summary on the step
step.tool_calls_summary = summary
await db.commit()
# ============ Phase 2: Writer ============
await sse.emit("phase_start", {"phase": 2, "name": "writer"})
messages = await build_orchestrator_phase2_context(
db=db, world=world, player_action=step.player_action,
plan=plan, summary=summary, settings=settings,
)
registry = get_registry()
ctx = ToolContext(db=db, world=world, step_id=step.id, stage="orchestrator_phase2",
sse_emitter=sse.emit)
tools = registry.to_openai_format("orchestrator_phase2")
phase2_msg: dict[str, Any] = {}
for retry in range(3):
resp = await llm.complete(
stage="orchestrator_phase2",
messages=messages,
tools=tools,
temperature=float(settings.get("llm.temperature_writer", 0.85)),
max_tokens=int(settings.get("llm.max_tokens", 2048)),
world_id=world.id, step_id=step.id, session=db,
)
phase2_msg = resp.get("message", {})
tcs = phase2_msg.get("tool_calls") or []
if tcs:
# Execute submit_step
for tc in tcs:
fn = tc.get("function", {})
if fn.get("name") == "submit_step":
try:
args = json.loads(fn.get("arguments") or "{}")
except json.JSONDecodeError:
args = {}
result = await registry.execute("submit_step", args, ctx)
if result.ok:
scene_text = result.data.get("scene_text", "")
delta_time = result.data.get("delta_time", "hours_1")
step.scene_text = scene_text
step.scene_delta_time = delta_time
await sse.emit("scene_complete", {
"text": scene_text, "delta_time": delta_time,
})
break
if step.scene_text:
break
# Retry
messages.append(phase2_msg)
messages.append({
"role": "user",
"content": "You MUST call submit_step with scene_text and delta_time.",
})
else:
await sse.error("writer_no_submit", "Writer failed to call submit_step after 3 retries")
step.status = "failed"
await db.commit()
return
await sse.emit("phase_end", {"phase": 2, "duration_ms": 0})
# ============ Phase 3 ============
await sse.emit("phase_start", {"phase": 3, "name": "persist_triggers_summary_suggest"})
# 3.0 Persist
step.status = "completed"
world.last_played_at = datetime.now(timezone.utc)
world.current_time = advance_time(
world.current_time, step.scene_delta_time or "hours_1", world.time_schema
)
await db.commit()
# 3.1 Deferred triggers
if settings.get("game.deferred_triggers_enabled", True):
await _process_deferred_triggers(
db=db, world=world, step=step, llm=llm, sse=sse, settings=settings,
)
# 3.2 Summary (if history is too long)
await _maybe_generate_summary(
db=db, world=world, step=step, llm=llm, sse=sse, settings=settings,
)
# 3.3 Suggest actions
suggest_msgs = await build_orchestrator_phase3_suggest_context(
db=db, world=world, scene_text=step.scene_text or "", settings=settings,
)
suggest_tools = registry.to_openai_format("orchestrator_phase3_suggest")
for retry in range(2):
resp = await llm.complete(
stage="orchestrator_phase3_suggest",
messages=suggest_msgs,
tools=suggest_tools,
temperature=0.8,
max_tokens=512,
world_id=world.id, step_id=step.id, session=db,
)
msg = resp.get("message", {})
tcs = msg.get("tool_calls") or []
for tc in tcs:
fn = tc.get("function", {})
if fn.get("name") == "suggest_actions":
try:
args = json.loads(fn.get("arguments") or "{}")
except json.JSONDecodeError:
args = {}
result = await registry.execute("suggest_actions", args, ctx)
if result.ok:
step.suggested_actions = result.data.get("actions", [])
await sse.emit("suggested_actions", {"actions": step.suggested_actions})
break
if step.suggested_actions:
break
suggest_msgs.append(msg)
suggest_msgs.append({"role": "user", "content": "Call suggest_actions with 1-3 actions."})
await db.commit()
await sse.emit("iteration_complete", {
"step_id": str(step.id), "sequence_number": step.sequence_number,
})
await sse.done({"step_id": str(step.id), "status": "completed"})
except Exception as e: # noqa: BLE001
_logger.exception("orchestrator_failed", step_id=str(step.id), error=str(e))
step.status = "failed"
await db.commit()
await sse.error("internal_error", str(e))
async def _process_deferred_triggers(
*,
db: AsyncSession,
world: World,
step: Step,
llm: LlmClient | MockLlmClient,
sse: SseEmitter,
settings: dict[str, Any],
) -> None:
"""Fire all deferred triggers whose fire_at <= current_time."""
from app.core.time_utils import time_le
triggers = (
await db.execute(
select(DeferredTrigger).where(
DeferredTrigger.world_id == world.id,
DeferredTrigger.is_fired.is_(False),
)
)
).scalars().all()
fired = 0
for trig in triggers:
try:
if not time_le(trig.fire_at, world.current_time):
continue
except Exception: # noqa: BLE001
continue
# Simple firing: append a note to scene_text
summary = f"\n\n[Offscreen event: {trig.event_type} — payload: {json.dumps(trig.payload, ensure_ascii=False)}]"
if step.scene_text:
step.scene_text += summary
else:
step.scene_text = summary
trig.is_fired = True
trig.fired_at = datetime.now(timezone.utc)
await db.flush()
await sse.emit("trigger_fired", {
"trigger_id": str(trig.id), "event_type": trig.event_type,
"summary": summary.strip(),
})
fired += 1
# Persist the trigger event as a story entry
await rag_add(
db=db, world_id=world.id,
content=f"Deferred trigger fired: {trig.event_type} at {trig.fire_at}",
entry_type="event",
metadata={"trigger_id": str(trig.id), "step_id": str(step.id)},
step_id=step.id,
)
if fired:
await db.commit()
async def _maybe_generate_summary(
*,
db: AsyncSession,
world: World,
step: Step,
llm: LlmClient | MockLlmClient,
sse: SseEmitter,
settings: dict[str, Any],
) -> None:
"""Generate a summary if recent step count exceeds the threshold."""
threshold = int(settings.get("context.compression_threshold_messages", 20))
guaranteed = int(settings.get("context.guaranteed_messages", 10))
recent_steps = list(
reversed(
(
await db.execute(
select(Step)
.where(Step.world_id == world.id, Step.deleted_at.is_(None))
.order_by(Step.sequence_number.desc())
.limit(threshold + 1)
)
).scalars().all()
)
)
if len(recent_steps) <= threshold:
return
old_steps = recent_steps[:-guaranteed]
if not old_steps:
return
messages = await build_summary_context(
db=db, world=world, old_steps=old_steps, settings=settings,
)
resp = await llm.complete(
stage="orchestrator_phase3_summary",
messages=messages,
temperature=0.3,
max_tokens=1024,
world_id=world.id, step_id=step.id, session=db,
)
summary_text = resp.get("message", {}).get("content", "")
if not summary_text:
return
# Store as a story entry
se = StoryEntry(
world_id=world.id,
content=summary_text,
entry_type="event",
metadata_={
"type": "summary",
"step_range": [old_steps[0].sequence_number, old_steps[-1].sequence_number],
},
embedding_status="pending",
)
db.add(se)
await db.commit()
await sse.emit("summary_generated", {
"summary_id": str(se.id),
"message_range": [old_steps[0].sequence_number, old_steps[-1].sequence_number],
})