Files
ai-rpg/backend/app/engine/orchestrator.py

616 lines
23 KiB
Python
Raw Normal View History

2026-06-19 16:31:45 +03:00
"""Game orchestrator: runs the multi-step LLM tool-calling loop and produces a narrative step.
Design (v2 tool-calling-first):
- The orchestrator LLM is given a set of game tools (dice_roll, update_state,
rag_query, rag_add, schedule_trigger, advance_time, run_subagent) PLUS a
`submit_plan` tool. The LLM calls game tools to execute its plan, then
calls `submit_plan` to terminate the loop with structured data.
- The step-writer LLM is given only a `submit_scene` tool. It calls this
to return the narrative + options; its text response is ignored.
- This replaces the old "return JSON in your text response" pattern which
conflicted with tool use and caused the model to dump raw JSON into chat.
"""
2026-06-19 11:28:04 +03:00
from __future__ import annotations
import json
import uuid
from datetime import datetime, timezone
from typing import Any, AsyncIterator, Dict, List, Optional
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.llm import LlmClient
from app.core.settings_service import cast_setting, get_all_settings
2026-06-19 16:31:45 +03:00
from app.core.triggers import advance_world_time, fire_due_triggers
2026-06-19 11:28:04 +03:00
from app.engine.context import (
build_orchestrator_messages,
build_step_writer_messages,
build_subagent_messages,
)
2026-06-19 16:31:45 +03:00
from app.engine.tools.tools import (
ALL_TOOL_SCHEMAS,
STEP_WRITER_TOOL_SCHEMAS,
ToolContext,
handle_tool_call,
)
2026-06-19 11:28:04 +03:00
from app.logging_setup import get_logger
2026-06-19 16:31:45 +03:00
from app.models import Message, Session, World
2026-06-19 11:28:04 +03:00
log = get_logger("orchestrator")
async def run_iteration(
db: AsyncSession,
user_id: uuid.UUID,
session_id: uuid.UUID,
action_text: str,
) -> AsyncIterator[Dict[str, Any]]:
"""Run one full iteration: plan -> tools -> step -> technical side-effects.
Yields SSE-ready event dicts:
{"type": "status", "data": {"message": "..."}}
{"type": "plan", "data": {...}} # orchestrator plan with tool calls
{"type": "tool_call", "data": {"name": ..., "args": ..., "result": ...}}
{"type": "narrative_chunk", "data": {"content": "..."}}
{"type": "step_complete", "data": {"message_id": ..., "options": [...], "state": ...}}
{"type": "error", "data": {"message": "..."}}
{"type": "done", "data": {}}
"""
# Load session + world
result = await db.execute(select(Session).where(Session.id == session_id))
session = result.scalars().first()
if not session:
yield {"type": "error", "data": {"message": "session_not_found"}}
return
result = await db.execute(select(World).where(World.id == session.world_id))
world = result.scalars().first()
if not world:
yield {"type": "error", "data": {"message": "world_not_found"}}
return
settings_map = await get_all_settings(db)
llm = LlmClient(settings_map)
2026-06-19 19:14:27 +03:00
# Save the player's action as a message — UNLESS this is a retry of the
# previous action (frontend re-sent the same action_text after an error).
# In that case we reuse the existing player_action row so the chat
# history doesn't fill up with duplicates.
last_msg_result = await db.execute(
select(Message)
.where(Message.session_id == session_id)
.order_by(Message.seq.desc())
.limit(1)
2026-06-19 11:28:04 +03:00
)
2026-06-19 19:14:27 +03:00
last_msg = last_msg_result.scalars().first()
is_retry = (
last_msg is not None
and last_msg.kind == "player_action"
and last_msg.content == action_text
)
if is_retry:
player_msg = last_msg
else:
next_seq = await _next_seq(db, session_id)
player_msg = Message(
session_id=session_id,
seq=next_seq,
role="user",
kind="player_action",
content=action_text,
payload={},
is_pinned=True,
hidden=False,
)
db.add(player_msg)
await db.commit()
await db.refresh(player_msg)
2026-06-19 11:28:04 +03:00
yield {"type": "status", "data": {"message": "planning"}}
# Subagent runner
async def _subagent(task: str, context: str) -> str:
sub_messages = await build_subagent_messages(world, task, context)
resp = await llm.chat(
messages=sub_messages,
temperature=0.7,
max_tokens=300,
purpose="subagent",
user_id=user_id,
session_id=session_id,
db=db,
)
return resp.text
2026-06-19 16:31:45 +03:00
ctx = ToolContext(
db=db,
world=world,
session_id=session_id,
user_id=user_id,
subagent_runner=_subagent,
settings_map=settings_map,
)
2026-06-19 11:28:04 +03:00
# === Phase 1: Orchestrator with tool calls (max 5 iterations) ===
orchestrator_messages, _ = await build_orchestrator_messages(db, world, session_id, action_text)
max_iters = 5
2026-06-19 16:31:45 +03:00
parsed: Dict[str, Any] = {}
plan_tool_calls_log: List[Dict[str, Any]] = []
orchestrator_text_log: str = ""
2026-06-19 11:28:04 +03:00
for i in range(max_iters):
yield {"type": "status", "data": {"message": f"orchestrator_turn_{i + 1}"}}
response = await llm.chat(
messages=orchestrator_messages,
tools=ALL_TOOL_SCHEMAS,
temperature=float(cast_setting("llm.temperature", settings_map.get("llm.temperature", 0.7))),
purpose="orchestrator",
user_id=user_id,
session_id=session_id,
db=db,
)
2026-06-19 16:31:45 +03:00
if not response.tool_calls:
# No tool calls — model gave up or errored. Treat its text as the
# outcome directly so the player still sees SOMETHING.
log.warning("orchestrator_no_tool_calls", iteration=i, text_len=len(response.text or ""))
orchestrator_text_log = response.text or ""
parsed = {
"assessment": "(no plan submitted)",
"outcome": response.text or "",
"narrative_prompt": "",
"next_options": [],
"state_patch": {},
"time_advance": None,
"rag_facts": [],
"rails_update": None,
}
break
# Append assistant message with tool_calls
orchestrator_messages.append({
"role": "assistant",
"content": response.text or "",
"tool_calls": response.tool_calls,
})
# Check for submit_plan — if present, extract plan and break
submit_plan_call = None
for tc in response.tool_calls:
if tc.get("function", {}).get("name") == "submit_plan":
submit_plan_call = tc
break
if submit_plan_call:
# Extract plan from the submit_plan tool call
args_str = submit_plan_call.get("function", {}).get("arguments", "{}")
try:
parsed = json.loads(args_str) if args_str else {}
except json.JSONDecodeError:
log.warning("submit_plan_invalid_json", args=args_str[:200])
parsed = {}
# Make sure required keys exist
parsed.setdefault("assessment", "")
parsed.setdefault("outcome", "")
parsed.setdefault("narrative_prompt", "")
parsed.setdefault("next_options", [])
parsed.setdefault("state_patch", {})
parsed.setdefault("time_advance", None)
parsed.setdefault("rag_facts", [])
parsed.setdefault("rails_update", None)
# Acknowledge the tool call so the model's history is consistent
2026-06-19 11:28:04 +03:00
orchestrator_messages.append({
2026-06-19 16:31:45 +03:00
"role": "tool",
"tool_call_id": submit_plan_call.get("id", ""),
"name": "submit_plan",
"content": json.dumps({"ok": True}),
2026-06-19 11:28:04 +03:00
})
2026-06-19 16:31:45 +03:00
# Log OTHER tool calls made this iteration (for debugging)
2026-06-19 11:28:04 +03:00
for tc in response.tool_calls:
fn = tc.get("function", {})
2026-06-19 16:31:45 +03:00
if fn.get("name") != "submit_plan":
plan_tool_calls_log.append({
"name": fn.get("name"),
"args": _safe_parse_json(fn.get("arguments", "{}")),
})
2026-06-19 11:28:04 +03:00
break
2026-06-19 16:31:45 +03:00
# Otherwise: execute all tool calls and continue
for tc in response.tool_calls:
fn = tc.get("function", {})
name = fn.get("name", "")
args_str = fn.get("arguments", "{}")
try:
args = json.loads(args_str) if args_str else {}
except json.JSONDecodeError:
args = {}
yield {"type": "tool_call", "data": {"name": name, "args": args}}
try:
result_dict = await handle_tool_call(name, args, ctx)
except Exception as e:
result_dict = {"error": f"{type(e).__name__}: {e}"}
log.error("tool_call_failed", name=name, error=str(e))
yield {"type": "tool_result", "data": {"name": name, "result": result_dict}}
plan_tool_calls_log.append({"name": name, "args": args, "result": result_dict})
# Append tool result message
orchestrator_messages.append({
"role": "tool",
"tool_call_id": tc.get("id", ""),
"name": name,
"content": json.dumps(result_dict, ensure_ascii=False, default=str)[:800],
})
await db.commit()
else:
# Ran out of iterations without submit_plan — use a minimal fallback.
log.warning("orchestrator_exhausted_iterations")
parsed = parsed or {
"assessment": "(iteration limit reached)",
"outcome": orchestrator_text_log or action_text,
"narrative_prompt": "",
"next_options": [],
"state_patch": {},
"time_advance": None,
"rag_facts": [],
"rails_update": None,
}
2026-06-19 11:28:04 +03:00
yield {"type": "status", "data": {"message": "writing_scene"}}
# Apply final state patch (if any)
if parsed.get("state_patch"):
from app.core.state_validator import apply_patch, validate_state
new_state = apply_patch(world.state, parsed["state_patch"])
schema = world.definition.get("world_schema", {})
ok, errors = validate_state(new_state, schema)
if ok:
world.state = new_state
else:
log.warning("state_patch_invalid", errors=errors)
# Advance time
time_advance = parsed.get("time_advance")
if time_advance and isinstance(time_advance, dict):
2026-06-19 16:31:45 +03:00
new_time, _total, _delta = advance_world_time(world.current_time, time_advance, world)
2026-06-19 11:28:04 +03:00
world.current_time = new_time
# Save orchestrator plan as hidden message
plan_seq = await _next_seq(db, session_id)
plan_msg = Message(
session_id=session_id,
seq=plan_seq,
role="assistant",
kind="orchestrator_plan",
2026-06-19 16:31:45 +03:00
content=(parsed.get("assessment", "") + " | " + parsed.get("outcome", ""))[:2000],
2026-06-19 11:28:04 +03:00
payload={
"assessment": parsed.get("assessment", ""),
"outcome": parsed.get("outcome", ""),
"state_patch": parsed.get("state_patch", {}),
"time_advance": time_advance,
2026-06-19 16:31:45 +03:00
"tool_calls_made": plan_tool_calls_log,
2026-06-19 11:28:04 +03:00
"scheduled_triggers": ctx.scheduled_triggers,
"rag_added": ctx.rag_added,
2026-06-19 16:31:45 +03:00
"narrative_prompt": parsed.get("narrative_prompt", ""),
"next_options": parsed.get("next_options", []),
2026-06-19 11:28:04 +03:00
},
is_pinned=False,
hidden=True,
)
db.add(plan_msg)
2026-06-19 16:31:45 +03:00
# === Phase 2: Step writer (narrative scene) — uses submit_scene tool ===
2026-06-19 11:28:04 +03:00
narrative_prompt_parts = [parsed.get("narrative_prompt", "")]
# Add RAG context if relevant
if parsed.get("outcome"):
try:
from app.core.rag import get_rag
rag = await get_rag(settings_map)
2026-06-19 16:31:45 +03:00
rag_results = await rag.search_glossary(
world.id,
parsed.get("outcome", ""),
limit=3,
settings_map=settings_map,
)
2026-06-19 11:28:04 +03:00
if rag_results:
2026-06-19 16:31:45 +03:00
rag_text = "\n".join(
f"- {r.get('name', '?')}: {r.get('description', '')[:120]}" for r in rag_results
)
2026-06-19 11:28:04 +03:00
narrative_prompt_parts.append(f"Relevant facts from glossary:\n{rag_text}")
except Exception as e:
log.warning("rag_lookup_failed", error=str(e))
step_messages = await build_step_writer_messages(
db=db,
world=world,
session_id=session_id,
outcome=parsed.get("outcome", action_text),
narrative_prompt="\n".join(p for p in narrative_prompt_parts if p),
)
step_resp = await llm.chat(
messages=step_messages,
2026-06-19 16:31:45 +03:00
tools=STEP_WRITER_TOOL_SCHEMAS,
2026-06-19 11:28:04 +03:00
temperature=float(cast_setting("llm.step_temperature", settings_map.get("llm.step_temperature", 0.85))),
2026-06-19 16:31:45 +03:00
max_tokens=1200,
2026-06-19 11:28:04 +03:00
purpose="step",
user_id=user_id,
session_id=session_id,
db=db,
)
2026-06-19 16:31:45 +03:00
# Extract scene from submit_scene tool call (if present); fall back to text.
step_text = step_resp.text or ""
2026-06-19 11:28:04 +03:00
step_options: List[str] = parsed.get("next_options", []) or []
2026-06-19 16:31:45 +03:00
for tc in (step_resp.tool_calls or []):
if tc.get("function", {}).get("name") == "submit_scene":
args_str = tc.get("function", {}).get("arguments", "{}")
try:
scene_data = json.loads(args_str) if args_str else {}
if scene_data.get("narrative"):
step_text = scene_data["narrative"]
if scene_data.get("options") and isinstance(scene_data["options"], list):
step_options = [str(o) for o in scene_data["options"]][:5]
except json.JSONDecodeError:
log.warning("submit_scene_invalid_json", args=args_str[:200])
break
else:
# No submit_scene call — try to extract JSON from text as a last resort.
import re as _re
json_match = _re.search(r"\{[\s\S]*\}", step_resp.text or "")
if json_match:
try:
step_data = json.loads(json_match.group(0))
if "narrative" in step_data:
step_text = step_data["narrative"]
if "options" in step_data and isinstance(step_data["options"], list):
step_options = [str(o) for o in step_data["options"]][:5]
except json.JSONDecodeError:
pass
# If still no narrative, use the orchestrator's outcome as fallback.
if not step_text.strip():
step_text = parsed.get("outcome", action_text)
2026-06-19 11:28:04 +03:00
# Save narrative step message
step_seq = await _next_seq(db, session_id)
step_msg = Message(
session_id=session_id,
seq=step_seq,
role="assistant",
kind="narrative_step",
content=step_text,
payload={
"options": step_options,
"outcome": parsed.get("outcome", ""),
"world_time": world.current_time,
"player_state": world.state.get("player", {}),
},
is_pinned=True,
hidden=False,
)
db.add(step_msg)
# === Phase 3: Update plot rails (if any) ===
rails_update = parsed.get("rails_update")
if rails_update and isinstance(rails_update, dict):
defn = dict(world.definition)
rails = dict(defn.get("plot_rails", {}))
if "main_goal" in rails_update:
rails["main_goal"] = rails_update["main_goal"]
if "new_subgoals" in rails_update:
existing = list(rails.get("subgoals", []))
existing.extend(rails_update["new_subgoals"])
rails["subgoals"] = existing
if "completed_subgoals" in rails_update:
completed = set(rails.get("completed_subgoals", []))
completed.update(rails_update["completed_subgoals"])
rails["completed_subgoals"] = list(completed)
rails["subgoals"] = [s for s in rails.get("subgoals", []) if s not in completed]
defn["plot_rails"] = rails
world.definition = defn
# Add RAG facts from orchestrator response
rag_facts = parsed.get("rag_facts", []) or []
if rag_facts:
from app.core.rag import get_rag
from app.models import GlossaryEntry
rag = await get_rag(settings_map)
for f in rag_facts:
if not isinstance(f, dict):
continue
entry = GlossaryEntry(
world_id=world.id,
session_id=session_id,
kind=f.get("kind", "lore"),
name=f.get("name", "unknown"),
description=f.get("description", ""),
payload={},
)
db.add(entry)
await db.flush()
await rag.upsert_glossary(
world_id=world.id,
entry_id=entry.id,
kind=entry.kind,
name=entry.name,
description=entry.description,
payload={},
settings_map=settings_map,
)
# Update session last_played_at
session.last_played_at = datetime.now(timezone.utc)
await db.commit()
await db.refresh(step_msg)
2026-06-19 16:31:45 +03:00
# Check for triggers that should fire now (fire_at <= current world time).
# Triggers fire on in-game time changes, not real-time polling — see
# app.core.triggers. We do this AFTER committing the narrative step so the
# player sees the main scene first, then any trigger consequences.
triggers_enabled = bool(cast_setting(
"triggers.enabled",
settings_map.get("triggers.enabled", True),
))
fired_now: List[Dict[str, Any]] = []
if triggers_enabled:
try:
await db.refresh(world)
fired_now = await fire_due_triggers(
db=db,
session_id=session_id,
world=world,
settings_map=settings_map,
user_id=user_id,
)
except Exception as e:
log.warning("trigger_fire_failed_in_iteration", error=f"{type(e).__name__}: {e}")
2026-06-19 11:28:04 +03:00
yield {
"type": "step_complete",
"data": {
"message_id": str(step_msg.id),
"seq": step_msg.seq,
"narrative": step_text,
"options": step_options,
"state": world.state,
"world_time": world.current_time,
"player_state": world.state.get("player", {}),
"fired_triggers": fired_now,
},
}
yield {"type": "done", "data": {}}
async def _next_seq(db: AsyncSession, session_id: uuid.UUID) -> int:
result = await db.execute(
select(Message.seq).where(Message.session_id == session_id).order_by(Message.seq.desc()).limit(1)
)
row = result.first()
return (row[0] + 1) if row else 1
2026-06-19 16:31:45 +03:00
def _safe_parse_json(s: str) -> Any:
try:
return json.loads(s) if s else {}
except Exception:
return s
2026-06-19 19:14:27 +03:00
async def generate_intro_scene(
db: AsyncSession,
user_id: uuid.UUID,
session_id: uuid.UUID,
) -> AsyncIterator[Dict[str, Any]]:
"""Generate the opening cinematic scene for a freshly-created session.
Yields the same SSE event stream shape as `run_iteration` so the
frontend can consume it identically. Saves a `narrative_step` message
of kind `intro_scene` (still kind=narrative_step for compatibility,
but with payload.kind=intro so the UI can style it differently if
desired).
"""
result = await db.execute(select(Session).where(Session.id == session_id))
session = result.scalars().first()
if not session:
yield {"type": "error", "data": {"message": "session_not_found"}}
return
result = await db.execute(select(World).where(World.id == session.world_id))
world = result.scalars().first()
if not world:
yield {"type": "error", "data": {"message": "world_not_found"}}
return
settings_map = await get_all_settings(db)
llm = LlmClient(settings_map)
yield {"type": "status", "data": {"message": "writing_scene"}}
import json as _json
defn = world.definition or {}
player_state = world.state.get("player", {}) if world.state else {}
system_prompt = get_prompt("intro_scene", world.language).format(
setting_description=defn.get("setting_description", "")[:1200],
current_time=world.current_time or "",
player_state=_json.dumps(player_state, ensure_ascii=False)[:800],
plot_rails=_json.dumps(defn.get("plot_rails", {}), ensure_ascii=False)[:600],
world_language=world.language or "en",
)
step_resp = await llm.chat(
messages=[{"role": "system", "content": system_prompt}],
tools=STEP_WRITER_TOOL_SCHEMAS,
temperature=float(cast_setting("llm.step_temperature", settings_map.get("llm.step_temperature", 0.85))),
max_tokens=1500,
purpose="intro_scene",
user_id=user_id,
session_id=session_id,
db=db,
)
step_text = step_resp.text or ""
step_options: List[str] = []
for tc in (step_resp.tool_calls or []):
if tc.get("function", {}).get("name") == "submit_scene":
args_str = tc.get("function", {}).get("arguments", "{}")
try:
scene_data = _json.loads(args_str) if args_str else {}
if scene_data.get("narrative"):
step_text = scene_data["narrative"]
if scene_data.get("options") and isinstance(scene_data["options"], list):
step_options = [str(o) for o in scene_data["options"]][:5]
except _json.JSONDecodeError:
log.warning("intro_scene_invalid_json", args=args_str[:200])
break
else:
# Fallback: extract JSON from text.
import re as _re
json_match = _re.search(r"\{[\s\S]*\}", step_resp.text or "")
if json_match:
try:
step_data = _json.loads(json_match.group(0))
if "narrative" in step_data:
step_text = step_data["narrative"]
if "options" in step_data and isinstance(step_data["options"], list):
step_options = [str(o) for o in step_data["options"]][:5]
except _json.JSONDecodeError:
pass
# Save as a narrative_step message flagged as intro in payload.
step_seq = await _next_seq(db, session_id)
step_msg = Message(
session_id=session_id,
seq=step_seq,
role="assistant",
kind="narrative_step",
content=step_text,
payload={
"kind": "intro",
"options": step_options,
"world_time": world.current_time,
"player_state": world.state.get("player", {}),
},
is_pinned=True,
hidden=False,
)
db.add(step_msg)
session.last_played_at = datetime.now(timezone.utc)
await db.commit()
await db.refresh(step_msg)
yield {
"type": "step_complete",
"data": {
"message_id": str(step_msg.id),
"seq": step_msg.seq,
"narrative": step_text,
"options": step_options,
"state": world.state,
"world_time": world.current_time,
"player_state": world.state.get("player", {}),
"fired_triggers": [],
"is_intro": True,
},
}
yield {"type": "done", "data": {}}