181 lines
6.1 KiB
Python
181 lines
6.1 KiB
Python
|
|
"""Context manager — builds the LLM message list per stage.
|
||
|
|
|
||
|
|
Implements the compression strategy from §10.3 of the TDD:
|
||
|
|
- If history > threshold, prepend the latest summary as a system message.
|
||
|
|
- Truncate to last N guaranteed messages.
|
||
|
|
- Optionally include RAG results.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import json
|
||
|
|
from typing import Any
|
||
|
|
|
||
|
|
from sqlalchemy import select
|
||
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
|
|
||
|
|
from app.core.logging import get_logger
|
||
|
|
from app.core.settings_service import get_setting
|
||
|
|
from app.core.time_utils import summarize_schemas
|
||
|
|
from app.models import StoryEntry, Step, World
|
||
|
|
from app.prompts.registry import get_prompt
|
||
|
|
|
||
|
|
_logger = get_logger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
def _scene_text_truncate(text: str, max_tokens: int) -> str:
|
||
|
|
"""Crude truncation: ~4 chars per token."""
|
||
|
|
max_chars = max_tokens * 4
|
||
|
|
if len(text) <= max_chars:
|
||
|
|
return text
|
||
|
|
return text[:max_chars] + "…"
|
||
|
|
|
||
|
|
|
||
|
|
async def build_orchestrator_phase1_context(
|
||
|
|
*,
|
||
|
|
db: AsyncSession,
|
||
|
|
world: World,
|
||
|
|
player_action: str,
|
||
|
|
settings: dict[str, Any],
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
"""Build messages list for orchestrator Phase 1."""
|
||
|
|
guaranteed = int(settings.get("context.guaranteed_messages", 10))
|
||
|
|
threshold = int(settings.get("context.compression_threshold_messages", 20))
|
||
|
|
scene_trunc = int(settings.get("context.scene_text_truncate_tokens", 500))
|
||
|
|
|
||
|
|
# Fetch recent steps (most recent first)
|
||
|
|
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(max(threshold, guaranteed) + 1)
|
||
|
|
)
|
||
|
|
).scalars().all()
|
||
|
|
)
|
||
|
|
)
|
||
|
|
|
||
|
|
# Pull latest summary if available
|
||
|
|
summary_text: str | None = None
|
||
|
|
if len(recent_steps) > threshold:
|
||
|
|
latest_summary = (
|
||
|
|
await db.execute(
|
||
|
|
select(StoryEntry)
|
||
|
|
.where(
|
||
|
|
StoryEntry.world_id == world.id,
|
||
|
|
StoryEntry.entry_type == "event",
|
||
|
|
StoryEntry.metadata_["type"].as_string() == "summary",
|
||
|
|
)
|
||
|
|
.order_by(StoryEntry.created_at.desc())
|
||
|
|
.limit(1)
|
||
|
|
)
|
||
|
|
).scalar_one_or_none()
|
||
|
|
if latest_summary:
|
||
|
|
summary_text = latest_summary.content
|
||
|
|
|
||
|
|
# Build the message list
|
||
|
|
sys_prompt = get_prompt("orchestrator_phase1", "en").format(
|
||
|
|
world_name=world.name,
|
||
|
|
rules="\n".join(f"- {r}" for r in (world.rules or [])),
|
||
|
|
schemas_summary=summarize_schemas(world.schemas or []),
|
||
|
|
environment_json=json.dumps(world.environment or {}, ensure_ascii=False, indent=2),
|
||
|
|
plot_rails_json=json.dumps(world.plot_rails or {}, ensure_ascii=False, indent=2),
|
||
|
|
current_time=world.current_time,
|
||
|
|
recent_history=_format_recent_history(
|
||
|
|
recent_steps[-guaranteed:], scene_trunc
|
||
|
|
),
|
||
|
|
max_substeps=settings.get("game.max_substeps_per_iteration", 8),
|
||
|
|
language=world.language,
|
||
|
|
player_action=player_action,
|
||
|
|
)
|
||
|
|
messages: list[dict[str, Any]] = [{"role": "system", "content": sys_prompt}]
|
||
|
|
if summary_text:
|
||
|
|
messages.append({
|
||
|
|
"role": "system",
|
||
|
|
"content": f"Summary of earlier events:\n{summary_text}",
|
||
|
|
})
|
||
|
|
# Recent steps as user/assistant pairs
|
||
|
|
for s in recent_steps[-guaranteed:]:
|
||
|
|
messages.append({"role": "user", "content": s.player_action})
|
||
|
|
if s.scene_text:
|
||
|
|
messages.append({"role": "assistant", "content": s.scene_text})
|
||
|
|
# Current action
|
||
|
|
messages.append({"role": "user", "content": player_action})
|
||
|
|
return messages
|
||
|
|
|
||
|
|
|
||
|
|
def _format_recent_history(steps: list[Step], scene_trunc: int) -> str:
|
||
|
|
if not steps:
|
||
|
|
return "(no recent history)"
|
||
|
|
lines: list[str] = []
|
||
|
|
for s in steps[-5:]: # only show last 5 in the prompt
|
||
|
|
text = _scene_text_truncate(s.scene_text or "(no scene)", scene_trunc)
|
||
|
|
lines.append(f"[step {s.sequence_number}] {s.player_action}\n → {text}")
|
||
|
|
return "\n".join(lines)
|
||
|
|
|
||
|
|
|
||
|
|
async def build_orchestrator_phase2_context(
|
||
|
|
*,
|
||
|
|
db: AsyncSession,
|
||
|
|
world: World,
|
||
|
|
player_action: str,
|
||
|
|
plan: str,
|
||
|
|
summary: list[dict[str, Any]],
|
||
|
|
settings: dict[str, Any],
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
"""Build messages list for orchestrator Phase 2 (writer)."""
|
||
|
|
sys_prompt = get_prompt("orchestrator_phase2", "en").format(
|
||
|
|
world_name=world.name,
|
||
|
|
world_description=world.description or "",
|
||
|
|
language=world.language,
|
||
|
|
current_time=world.current_time,
|
||
|
|
player_action=player_action,
|
||
|
|
plan=plan,
|
||
|
|
summary_json=json.dumps(summary, ensure_ascii=False, indent=2),
|
||
|
|
environment_json=json.dumps(world.environment or {}, ensure_ascii=False, indent=2),
|
||
|
|
)
|
||
|
|
return [
|
||
|
|
{"role": "system", "content": sys_prompt},
|
||
|
|
{"role": "user", "content": "Write the scene and call submit_step."},
|
||
|
|
]
|
||
|
|
|
||
|
|
|
||
|
|
async def build_orchestrator_phase3_suggest_context(
|
||
|
|
*,
|
||
|
|
db: AsyncSession,
|
||
|
|
world: World,
|
||
|
|
scene_text: str,
|
||
|
|
settings: dict[str, Any],
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
sys_prompt = get_prompt("orchestrator_phase3_suggest", "en").format(
|
||
|
|
language=world.language,
|
||
|
|
scene_text=scene_text[:2000],
|
||
|
|
current_goals=", ".join((world.plot_rails or {}).get("current_goals", []) or ["(none)"]),
|
||
|
|
)
|
||
|
|
return [
|
||
|
|
{"role": "system", "content": sys_prompt},
|
||
|
|
{"role": "user", "content": "Suggest 1-3 next actions."},
|
||
|
|
]
|
||
|
|
|
||
|
|
|
||
|
|
async def build_summary_context(
|
||
|
|
*,
|
||
|
|
db: AsyncSession,
|
||
|
|
world: World,
|
||
|
|
old_steps: list[Step],
|
||
|
|
settings: dict[str, Any],
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
"""Build messages list for the summary LLM call."""
|
||
|
|
messages_json = json.dumps(
|
||
|
|
[{"action": s.player_action, "scene": s.scene_text} for s in old_steps],
|
||
|
|
ensure_ascii=False,
|
||
|
|
indent=2,
|
||
|
|
)
|
||
|
|
sys_prompt = get_prompt("summary", "en").format(messages_json=messages_json)
|
||
|
|
return [
|
||
|
|
{"role": "system", "content": sys_prompt},
|
||
|
|
{"role": "user", "content": "Summarize."},
|
||
|
|
]
|