You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
548 lines
21 KiB
548 lines
21 KiB
from __future__ import annotations
|
|
import asyncio
|
|
import json
|
|
import uuid
|
|
import logging
|
|
from typing import AsyncGenerator
|
|
|
|
from fastapi import APIRouter, Request, HTTPException, Query
|
|
from pydantic import BaseModel
|
|
from sse_starlette.sse import EventSourceResponse
|
|
|
|
from app.models.simulation import SSEEvent, SimulationStatus, SpeedMode, ForkRequest
|
|
from app.models.world import WorldState, WorldMetrics
|
|
from app.models.demographics import DemographicProfile
|
|
from app.services.census import CensusService
|
|
from app.services.i18n import locale_from_request, translate
|
|
|
|
from enum import Enum
|
|
|
|
logger = logging.getLogger(__name__)
|
|
router = APIRouter(prefix="/api")
|
|
|
|
|
|
def _serialize_event_data(data: dict) -> str:
|
|
def _default(obj):
|
|
if isinstance(obj, Enum):
|
|
return obj.value
|
|
return str(obj)
|
|
return json.dumps(data, default=_default)
|
|
|
|
|
|
class SegmentInput(BaseModel):
|
|
name: str
|
|
description: str
|
|
count: int | None = None
|
|
|
|
|
|
class SimulateRequest(BaseModel):
|
|
# NOTE: We intentionally avoid field-level validators here because they run
|
|
# before the request middleware has access to the Accept-Language header,
|
|
# which means their messages would always be in English. Validation lives
|
|
# in the route handler below so it can localize errors.
|
|
rules: str = ""
|
|
population: int = 25
|
|
duration_days: int = 365
|
|
proposed_change: str | None = None
|
|
segments: list[SegmentInput] | None = None
|
|
city: str | None = None
|
|
# If true, the simulation is automatically marked is_public=1 as soon as
|
|
# it finishes successfully, so it shows up in /gallery without a manual
|
|
# publish click. Defaults to false so users opt in explicitly.
|
|
auto_publish: bool = False
|
|
|
|
|
|
class InjectRequest(BaseModel):
|
|
event: str
|
|
|
|
|
|
class SpeedRequest(BaseModel):
|
|
mode: str
|
|
|
|
|
|
@router.post("/simulate")
|
|
async def simulate(req: SimulateRequest, request: Request):
|
|
locale = locale_from_request(request)
|
|
# Sync narrator's locale so generated narratives and reports use this language.
|
|
narrator = getattr(request.app.state, "narrator", None)
|
|
if narrator is not None and hasattr(narrator, "set_locale"):
|
|
narrator.set_locale(locale)
|
|
# Same for the world generator — world name, description, locations should
|
|
# follow the user's language.
|
|
world_gen = getattr(request.app.state, "world_generator", None)
|
|
if world_gen is not None and hasattr(world_gen, "set_locale"):
|
|
world_gen.set_locale(locale)
|
|
# Citizen generator: agent names, roles, personality hooks should match.
|
|
citizen_gen = getattr(request.app.state, "citizen_generator", None)
|
|
if citizen_gen is not None and hasattr(citizen_gen, "set_locale"):
|
|
citizen_gen.set_locale(locale)
|
|
# Tension engine: hardcoded event descriptions ("Consumer patience wore
|
|
# thin: ...") should localize too.
|
|
engine_obj = getattr(request.app.state, "engine", None)
|
|
if engine_obj is not None and hasattr(engine_obj, "tension"):
|
|
engine_obj.tension.set_locale(locale)
|
|
|
|
# Re-validate here so error messages respect Accept-Language. Pydantic's
|
|
# own field validators run before locale is available; we ignore those
|
|
# errors and produce a localized one instead.
|
|
rules = (req.rules or "").strip()
|
|
if not rules:
|
|
raise HTTPException(400, translate("rules_cannot_be_empty", locale))
|
|
if req.population < 2 or req.population > 50:
|
|
raise HTTPException(400, translate("population_range", locale))
|
|
if req.duration_days < 1 or req.duration_days > 365:
|
|
raise HTTPException(400, translate("duration_range", locale))
|
|
|
|
sim_id = str(uuid.uuid4())[:12]
|
|
store = request.app.state.store
|
|
world_gen = request.app.state.world_generator
|
|
citizen_gen = request.app.state.citizen_generator
|
|
engine = request.app.state.engine
|
|
|
|
await store.create(sim_id, rules)
|
|
|
|
cancelled = request.app.state.cancelled_sims
|
|
|
|
async def run_pipeline():
|
|
event_queue = request.app.state.event_queues.get(sim_id)
|
|
if not event_queue:
|
|
return
|
|
|
|
try:
|
|
await event_queue.put(SSEEvent(type="status", data={"status": "generating_world"}))
|
|
await store.update_status(sim_id, SimulationStatus.GENERATING_WORLD.value)
|
|
|
|
blueprint = await world_gen.generate(
|
|
req.rules, req.population, req.duration_days,
|
|
proposed_change=req.proposed_change,
|
|
)
|
|
|
|
if sim_id in cancelled:
|
|
await event_queue.put(SSEEvent(type="cancelled", data={"message": translate("simulation_cancelled", locale)}))
|
|
return
|
|
|
|
await store.set_meta(sim_id, "blueprint", blueprint.model_dump())
|
|
if req.proposed_change:
|
|
await store.set_meta(sim_id, "proposed_change", req.proposed_change)
|
|
# Persist auto_publish so engine.run can flip is_public=1 when the
|
|
# sim finishes, without having to thread the flag through every
|
|
# method on the way down.
|
|
await store.set_meta(sim_id, "auto_publish", req.auto_publish)
|
|
await store.update_status(sim_id, SimulationStatus.GENERATING_WORLD.value, world_name=blueprint.name)
|
|
await event_queue.put(SSEEvent(type="world_ready", data=blueprint.model_dump()))
|
|
|
|
if sim_id in cancelled:
|
|
await event_queue.put(SSEEvent(type="cancelled", data={"message": translate("simulation_cancelled", locale)}))
|
|
return
|
|
|
|
demographics: DemographicProfile | None = None
|
|
if req.city:
|
|
try:
|
|
census = CensusService(request.app.state.llm)
|
|
if "," in req.city:
|
|
city_name, state = [p.strip() for p in req.city.split(",", 1)]
|
|
else:
|
|
city_name, state = req.city.strip(), None
|
|
demographics = await census.get_profile(city_name, state)
|
|
await event_queue.put(SSEEvent(type="demographics_loaded", data={
|
|
"city": demographics.city_name,
|
|
"state": demographics.state,
|
|
"population": demographics.population,
|
|
"median_income": demographics.median_household_income,
|
|
"poverty_rate": demographics.poverty_rate,
|
|
}))
|
|
except Exception as e:
|
|
logger.warning("Demographics fetch failed for %s: %s", req.city, e)
|
|
|
|
await event_queue.put(SSEEvent(type="status", data={"status": "generating_citizens"}))
|
|
await store.update_status(sim_id, SimulationStatus.GENERATING_CITIZENS.value)
|
|
|
|
async def on_citizen(agent):
|
|
if sim_id in cancelled:
|
|
raise asyncio.CancelledError()
|
|
await store.save_agent(sim_id, agent)
|
|
await event_queue.put(SSEEvent(type="citizen_generated", data=agent.model_dump()))
|
|
|
|
segments_dicts = [s.model_dump() for s in req.segments] if req.segments else None
|
|
agents = await citizen_gen.generate_fast(
|
|
blueprint, req.population, on_citizen=on_citizen,
|
|
proposed_change=req.proposed_change,
|
|
segments=segments_dicts,
|
|
demographics=demographics,
|
|
)
|
|
|
|
for agent in agents:
|
|
await store.save_agent(sim_id, agent)
|
|
await event_queue.put(SSEEvent(type="citizen_generated", data=agent.model_dump()))
|
|
|
|
if sim_id in cancelled:
|
|
await event_queue.put(SSEEvent(type="cancelled", data={"message": translate("simulation_cancelled", locale)}))
|
|
return
|
|
|
|
async def enrich_in_background():
|
|
try:
|
|
enriched = await citizen_gen.enrich_background(blueprint, agents)
|
|
engine.deliver_enriched_agents(sim_id, enriched)
|
|
for agent in enriched:
|
|
await store.save_agent(sim_id, agent)
|
|
logger.info("Background enrichment complete for %s", sim_id)
|
|
except Exception as e:
|
|
logger.warning("Background enrichment failed for %s: %s", sim_id, e)
|
|
|
|
asyncio.create_task(enrich_in_background())
|
|
|
|
world_state = WorldState(
|
|
blueprint=blueprint,
|
|
metrics=WorldMetrics(),
|
|
community_rules=[],
|
|
)
|
|
|
|
await store.update_status(sim_id, SimulationStatus.RUNNING.value)
|
|
|
|
async def emit(event: SSEEvent):
|
|
await event_queue.put(event)
|
|
|
|
await engine.run(sim_id, world_state, agents, emit)
|
|
|
|
except asyncio.CancelledError:
|
|
logger.info("Pipeline cancelled for %s", sim_id)
|
|
await event_queue.put(SSEEvent(type="cancelled", data={"message": translate("simulation_cancelled", locale)}))
|
|
except Exception as e:
|
|
logger.error("Pipeline failed for %s: %s", sim_id, e, exc_info=True)
|
|
await event_queue.put(SSEEvent(type="error", data={"message": str(e)}))
|
|
finally:
|
|
cancelled.discard(sim_id)
|
|
|
|
request.app.state.event_queues[sim_id] = asyncio.Queue()
|
|
asyncio.create_task(run_pipeline())
|
|
|
|
return {"simulation_id": sim_id}
|
|
|
|
|
|
@router.get("/simulation/{sim_id}/stream")
|
|
async def stream(sim_id: str, request: Request):
|
|
locale = locale_from_request(request)
|
|
queue = request.app.state.event_queues.get(sim_id)
|
|
if not queue:
|
|
# Distinguish "no live pipeline" (e.g. after a backend restart) from
|
|
# "never existed" so the client can give a useful message.
|
|
sim = await request.app.state.store.get_simulation(sim_id)
|
|
if not sim:
|
|
raise HTTPException(404, translate("simulation_not_found", locale))
|
|
raise HTTPException(
|
|
410,
|
|
translate("simulation_pipeline_gone", locale),
|
|
)
|
|
|
|
async def event_generator() -> AsyncGenerator:
|
|
try:
|
|
while True:
|
|
if await request.is_disconnected():
|
|
break
|
|
try:
|
|
event = await asyncio.wait_for(queue.get(), timeout=30)
|
|
yield {"event": event.type, "data": _serialize_event_data(event.data)}
|
|
if event.type in ("simulation_complete", "error", "cancelled"):
|
|
break
|
|
except asyncio.TimeoutError:
|
|
yield {"event": "keepalive", "data": "{}"}
|
|
finally:
|
|
request.app.state.event_queues.pop(sim_id, None)
|
|
|
|
return EventSourceResponse(event_generator())
|
|
|
|
|
|
@router.get("/simulation/{sim_id}/state")
|
|
async def get_state(sim_id: str, request: Request):
|
|
store = request.app.state.store
|
|
sim = await store.get_simulation(sim_id)
|
|
if not sim:
|
|
raise HTTPException(404, translate("simulation_not_found", locale_from_request(request)))
|
|
|
|
world_state = await store.get_world_state(sim_id)
|
|
agents = await store.get_all_agents(sim_id)
|
|
|
|
return {
|
|
"simulation": sim,
|
|
"world_state": world_state.model_dump() if world_state else None,
|
|
"agent_count": len(agents),
|
|
}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/inject")
|
|
async def inject(sim_id: str, req: InjectRequest, request: Request):
|
|
engine = request.app.state.engine
|
|
store = request.app.state.store
|
|
|
|
world_state = await store.get_world_state(sim_id)
|
|
agents = await store.get_all_agents(sim_id)
|
|
|
|
if not world_state:
|
|
raise HTTPException(404, translate("no_world_state", locale_from_request(request)))
|
|
|
|
desc = await engine.inject_event(sim_id, req.event, world_state, agents)
|
|
return {"description": desc}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/speed")
|
|
async def set_speed(sim_id: str, req: SpeedRequest, request: Request):
|
|
engine = request.app.state.engine
|
|
try:
|
|
mode = SpeedMode(req.mode)
|
|
except ValueError:
|
|
raise HTTPException(400, translate("invalid_speed_mode", locale_from_request(request), mode=req.mode))
|
|
await engine.set_speed(sim_id, mode)
|
|
return {"mode": mode.value}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/pause")
|
|
async def pause(sim_id: str, request: Request):
|
|
found = await request.app.state.engine.pause(sim_id)
|
|
if not found:
|
|
raise HTTPException(404, translate("simulation_not_running", locale_from_request(request)))
|
|
return {"status": "paused"}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/resume")
|
|
async def resume(sim_id: str, request: Request):
|
|
found = await request.app.state.engine.resume(sim_id)
|
|
if not found:
|
|
raise HTTPException(404, translate("simulation_not_running", locale_from_request(request)))
|
|
return {"status": "resumed"}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/stop")
|
|
async def stop(sim_id: str, request: Request):
|
|
found = await request.app.state.engine.stop(sim_id)
|
|
if not found:
|
|
raise HTTPException(404, translate("simulation_not_running", locale_from_request(request)))
|
|
return {"status": "stopped"}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/cancel")
|
|
async def cancel(sim_id: str, request: Request):
|
|
engine = request.app.state.engine
|
|
cancelled = request.app.state.cancelled_sims
|
|
|
|
stopped = await engine.stop(sim_id)
|
|
if stopped:
|
|
return {"status": "stopped"}
|
|
|
|
cancelled.add(sim_id)
|
|
return {"status": "cancelled"}
|
|
|
|
|
|
@router.get("/simulation/{sim_id}/report")
|
|
async def get_report(sim_id: str, request: Request):
|
|
locale = locale_from_request(request)
|
|
store = request.app.state.store
|
|
sim = await store.get_simulation(sim_id)
|
|
if not sim:
|
|
raise HTTPException(404, translate("simulation_not_found", locale))
|
|
|
|
# Cancelled: client should not keep polling, but partial data (if any) is
|
|
# still in the report table. Try to return it before bailing.
|
|
if sim.get("status") == "cancelled":
|
|
result = await store.get_report(sim_id)
|
|
if result.get("status") == "ready":
|
|
return result["report"]
|
|
return {"status": "cancelled", "message": translate("simulation_cancelled", locale)}
|
|
|
|
# Interrupted (typically because the backend restarted and the pipeline
|
|
# task died). If a partial report was already saved before the crash,
|
|
# surface it — losing 30 rounds of work to a server restart is bad UX.
|
|
# Only when there's no report do we return the interrupted terminal status.
|
|
if sim.get("status") == "interrupted":
|
|
result = await store.get_report(sim_id)
|
|
if result.get("status") == "ready":
|
|
return result["report"]
|
|
return {
|
|
"status": "interrupted",
|
|
"message": translate("simulation_interrupted_no_report", locale),
|
|
}
|
|
|
|
result = await store.get_report(sim_id)
|
|
if result.get("status") == "ready":
|
|
report = result["report"]
|
|
report["rules_text"] = sim.get("rules_text", "") or ""
|
|
return report
|
|
return result
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/publish")
|
|
async def publish(sim_id: str, request: Request):
|
|
locale = locale_from_request(request)
|
|
sim = await request.app.state.store.get_simulation(sim_id)
|
|
if not sim:
|
|
raise HTTPException(404, translate("simulation_not_found", locale))
|
|
await request.app.state.store.publish(sim_id)
|
|
return {"published": True, "is_public": True}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/unpublish")
|
|
async def unpublish(sim_id: str, request: Request):
|
|
locale = locale_from_request(request)
|
|
sim = await request.app.state.store.get_simulation(sim_id)
|
|
if not sim:
|
|
raise HTTPException(404, translate("simulation_not_found", locale))
|
|
await request.app.state.store.unpublish(sim_id)
|
|
return {"published": False, "is_public": False}
|
|
|
|
|
|
@router.get("/simulation/{sim_id}/compare/{fork_id}")
|
|
async def compare(sim_id: str, fork_id: str, request: Request):
|
|
store = request.app.state.store
|
|
narrator = request.app.state.narrator
|
|
|
|
source_state = await store.get_world_state(sim_id)
|
|
fork_state = await store.get_world_state(fork_id)
|
|
|
|
if not source_state or not fork_state:
|
|
raise HTTPException(404, translate("simulation_not_found", locale_from_request(request)))
|
|
|
|
source_metrics = await store.get_metrics_history(sim_id)
|
|
fork_metrics = await store.get_metrics_history(fork_id)
|
|
|
|
source_agents = await store.get_all_agents(sim_id)
|
|
fork_agents = await store.get_all_agents(fork_id)
|
|
|
|
async def _get_or_generate_report(sid):
|
|
cached = await store.get_report(sid)
|
|
if cached and cached.get("status") == "ready":
|
|
return cached["report"]
|
|
rpt = await narrator.generate_report(sid, store)
|
|
if "error" not in rpt:
|
|
await store.save_report(sid, rpt)
|
|
return rpt
|
|
|
|
source_report = await _get_or_generate_report(sim_id)
|
|
fork_report = await _get_or_generate_report(fork_id)
|
|
|
|
return {
|
|
"source": {
|
|
"id": sim_id,
|
|
"report": source_report,
|
|
"metrics_history": source_metrics,
|
|
"final_day": source_state.day,
|
|
},
|
|
"fork": {
|
|
"id": fork_id,
|
|
"report": fork_report,
|
|
"metrics_history": fork_metrics,
|
|
"final_day": fork_state.day,
|
|
},
|
|
}
|
|
|
|
|
|
@router.get("/simulation/{sim_id}/forecast")
|
|
async def get_forecast(
|
|
sim_id: str,
|
|
request: Request,
|
|
horizon: int = Query(default=30, ge=1, le=365),
|
|
):
|
|
store = request.app.state.store
|
|
forecast = getattr(request.app.state, "forecast", None)
|
|
|
|
if not forecast or not forecast.available:
|
|
return {"error": translate("forecast_not_available", locale_from_request(request))}
|
|
|
|
metrics_history = await store.get_metrics_history(sim_id)
|
|
if not metrics_history:
|
|
raise HTTPException(404, translate("no_metrics_data", locale_from_request(request)))
|
|
|
|
analysis = forecast.analyze(metrics_history, horizon=horizon)
|
|
causal_links = forecast.discover_causality(sim_id)
|
|
counterfactuals = forecast.compute_counterfactuals(sim_id)
|
|
|
|
return {
|
|
"analysis": analysis,
|
|
"causal_links": [
|
|
{"cause": l.cause, "effect": l.effect, "lag": l.lag,
|
|
"p_value": l.p_value, "strength": l.strength}
|
|
for l in causal_links
|
|
],
|
|
"counterfactuals": [
|
|
{"event_round": c.event_round, "event_description": c.event_description,
|
|
"metric_impacts": c.metric_impacts}
|
|
for c in counterfactuals
|
|
],
|
|
}
|
|
|
|
|
|
@router.post("/simulation/{sim_id}/fork")
|
|
async def fork(sim_id: str, req: ForkRequest, request: Request):
|
|
store = request.app.state.store
|
|
engine = request.app.state.engine
|
|
locale = locale_from_request(request)
|
|
|
|
source = await store.get_simulation(sim_id)
|
|
if not source:
|
|
raise HTTPException(404, translate("source_simulation_not_found", locale_from_request(request)))
|
|
|
|
fork_id = str(uuid.uuid4())[:12]
|
|
source_world = await store.get_world_state(sim_id)
|
|
rounds_per_day = source_world.blueprint.time_config.rounds_per_day if source_world else 3
|
|
fork_round = req.fork_at_day * rounds_per_day
|
|
|
|
await store.create(fork_id, source.get("rules_text", ""))
|
|
await store.update_status(fork_id, SimulationStatus.PAUSED.value,
|
|
forked_from=sim_id,
|
|
world_name=source.get("world_name", ""),
|
|
agent_count=source.get("agent_count", 0))
|
|
await store.copy_state_at_round(sim_id, fork_id, fork_round)
|
|
await store.increment_fork_count(sim_id)
|
|
|
|
if req.changes:
|
|
world_state = await store.get_world_state(fork_id)
|
|
agents = await store.get_all_agents(fork_id)
|
|
if world_state and agents:
|
|
await engine.inject_event(fork_id, req.changes, world_state, agents)
|
|
|
|
event_queue = asyncio.Queue()
|
|
request.app.state.event_queues[fork_id] = event_queue
|
|
|
|
async def run_fork():
|
|
try:
|
|
world_state = await store.get_world_state(fork_id)
|
|
agents = await store.get_all_agents(fork_id)
|
|
if not world_state or not agents:
|
|
await event_queue.put(SSEEvent(type="error", data={"message": translate("fork_state_not_found", locale)}))
|
|
return
|
|
|
|
await event_queue.put(SSEEvent(
|
|
type="status", data={"status": "forked_simulation"}
|
|
))
|
|
|
|
await event_queue.put(SSEEvent(
|
|
type="world_ready", data=world_state.blueprint.model_dump()
|
|
))
|
|
|
|
for agent in agents:
|
|
await event_queue.put(SSEEvent(
|
|
type="citizen_generated", data=agent.model_dump()
|
|
))
|
|
|
|
await store.update_status(
|
|
fork_id, SimulationStatus.RUNNING.value,
|
|
world_name=world_state.blueprint.name,
|
|
agent_count=len(agents),
|
|
)
|
|
|
|
async def emit(event: SSEEvent):
|
|
await event_queue.put(event)
|
|
|
|
# Keep tension engine's locale in sync with the fork request —
|
|
# otherwise the hardcoded "Consumer patience wore thin: ..." line
|
|
# will surface in English even for zh users.
|
|
if hasattr(engine, "tension"):
|
|
engine.tension.set_locale(locale)
|
|
|
|
await engine.run(fork_id, world_state, agents, emit, start_round=fork_round)
|
|
|
|
except Exception as e:
|
|
logger.error("Fork pipeline failed for %s: %s", fork_id, e, exc_info=True)
|
|
await event_queue.put(SSEEvent(type="error", data={"message": str(e)}))
|
|
|
|
asyncio.create_task(run_fork())
|
|
|
|
return {"simulation_id": fork_id, "forked_from": sim_id, "fork_at_day": req.fork_at_day}
|
|
|