diff --git a/s16code/core/live_graph/store.py b/s16code/core/live_graph/store.py index d482e0e..259149c 100644 --- a/s16code/core/live_graph/store.py +++ b/s16code/core/live_graph/store.py @@ -147,6 +147,16 @@ def context(self, run_id: str) -> dict[str, Any]: with self._lock: return dict(self._load(run_id)["graph"].graph.get("context", {})) + def update_context(self, run_id: str, patch: dict[str, Any]) -> dict[str, Any]: + """Merge durable run context (budget, principal, spend) across a wait.""" + with self._lock: + state = self._load(run_id) + context = dict(state["graph"].graph.get("context") or {}) + context.update(patch) + state["graph"].graph["context"] = context + self._save(state) + return dict(context) + def resume(self, run_id: str) -> None: with self._lock: state = self._load(run_id) diff --git a/s16code/routes.py b/s16code/routes.py index ee09f81..9279ca1 100644 --- a/s16code/routes.py +++ b/s16code/routes.py @@ -298,7 +298,8 @@ async def resume(run_id: str, request: Request): try: return await runtime.run(prompt=None, scope=None, llm=lambda prompt, system: gateway_text_llm(request.app, prompt, system), - source_uri=None, source_author=None, run_id=run_id, resume=True) + source_uri=None, source_author=None, run_id=run_id, resume=True, + transport=request.app.state.gateway) except KeyError: raise HTTPException(404, "run not found") from None except RuntimeError as error: diff --git a/s16code/runtime.py b/s16code/runtime.py index b8fcef7..c87183d 100644 --- a/s16code/runtime.py +++ b/s16code/runtime.py @@ -155,6 +155,7 @@ async def run(self, *, prompt: str | None, scope: MemoryScope | None, llm: TextL run behaves exactly as it did before economics existed. """ run_id = run_id or f"run-{uuid.uuid4().hex[:12]}" + prior_spent = 0.0 if resume: context = self.graph.context(run_id) prompt = str(context["prompt"]) @@ -163,6 +164,14 @@ async def run(self, *, prompt: str | None, scope: MemoryScope | None, llm: TextL allowed_side_effects = set(context.get("allowed_side_effects", [])) source_uri, source_author, inbound_id = context["source_uri"], context["source_author"], context.get("inbound_id") initial_evidence = context.get("initial_evidence") or {} + # A wait is not a new run. The ceiling and spend that held before + # the park must come back with the graph, even if the caller (HTTP + # resume, channel approval, job callback) passes budget=None. + if budget is None and context.get("budget") is not None: + budget = float(context["budget"]) + if principal is None and context.get("principal"): + principal = str(context["principal"]) + prior_spent = float(context.get("budget_spent") or 0.0) else: if prompt is None or scope is None or source_uri is None or source_author is None: raise ValueError("new runs require prompt, scope, and source identity") @@ -170,14 +179,23 @@ async def run(self, *, prompt: str | None, scope: MemoryScope | None, llm: TextL inbound = self.memory.write(MemoryRecord(MemoryKind.EPISODE, scope, prompt, [user_source], Principal("gateway", "gateway"), metadata={"run_id": run_id})) inbound_id = inbound.id - self.graph.start(run_id, context={"prompt": prompt, "scope": {"tenant_id": scope.tenant_id, - "project_id": scope.project_id, "user_id": scope.user_id, - "agent_id": scope.agent_id, "run_id": scope.run_id}, - "source_uri": source_uri, "source_author": source_author, - "inbound_id": inbound_id, "respond_as": respond_as, - "allowed_side_effects": sorted(allowed_side_effects or ()), - "initial_evidence": initial_evidence or {}}) assert prompt is not None and scope is not None and source_uri is not None and source_author is not None + who = principal or principal_for(scope) + if not resume: + start_context: dict[str, Any] = { + "prompt": prompt, + "scope": {"tenant_id": scope.tenant_id, "project_id": scope.project_id, + "user_id": scope.user_id, "agent_id": scope.agent_id, "run_id": scope.run_id}, + "source_uri": source_uri, "source_author": source_author, + "inbound_id": inbound_id, "respond_as": respond_as, + "allowed_side_effects": sorted(allowed_side_effects or ()), + "initial_evidence": initial_evidence or {}, + "principal": who, + } + if budget is not None: + start_context["budget"] = float(budget) + start_context["budget_spent"] = 0.0 + self.graph.start(run_id, context=start_context) user_source = SourceRef(source_uri, source_author, excerpt=prompt) runtime = self @@ -189,12 +207,15 @@ async def run(self, *, prompt: str | None, scope: MemoryScope | None, llm: TextL economics_config: EconomicsConfig | None = None run_budget: RunBudget | None = None controller: BudgetedGateway | None = None - who = principal or principal_for(scope) if budget is not None: if transport is None: raise ValueError("a budgeted run needs a gateway transport to meter") economics_config = economics or EconomicsConfig.load() run_budget = economics_config.budget(principal=who, amount=float(budget), run_id=run_id) + if prior_spent > 0: + # Restoring only the original ceiling with spend reset to zero + # would give the run a second full allowance after every wait. + run_budget.spent = prior_spent controller = BudgetedGateway( transport, budget=run_budget, policy=economics_config.policy(), pricing=economics_config.pricing, ladder=economics_config.ladder, @@ -942,6 +963,12 @@ async def execute_once(task: TaskSpec) -> dict[str, Any] | Deferred: for name, worker in skills.items()}, max_workers=int(os.getenv("S16_MAX_WORKERS", "4")), ).run(run_id, resume=resume) + if run_budget is not None: + self.graph.update_context(run_id, { + "budget": run_budget.total, + "principal": who, + "budget_spent": run_budget.spent, + }) snapshot = self.graph.snapshot(run_id) terminal_skills = registry.terminal_skills(respond_as) terminal_nodes = [node for node in snapshot.nodes.values() if node["skill"] in terminal_skills] diff --git a/s16code/ui/routes.py b/s16code/ui/routes.py index 564f2f0..daaacbf 100644 --- a/s16code/ui/routes.py +++ b/s16code/ui/routes.py @@ -179,6 +179,7 @@ async def action(body: ActionBody, request: Request): prompt=None, scope=None, llm=lambda prompt, system: request.app.state.gateway.complete(prompt, system), source_uri=None, source_author=None, run_id=body.run_id, resume=True, + transport=request.app.state.gateway, ) return {"resumed": True, "node_id": body.node_id, "reason": decision.reason, "run": result} diff --git a/tests/test_budget_runtime.py b/tests/test_budget_runtime.py index 6a3d232..29d9a70 100644 --- a/tests/test_budget_runtime.py +++ b/tests/test_budget_runtime.py @@ -271,3 +271,94 @@ def test_the_trace_route_serves_the_span_tree(budgeted_client): def test_the_trace_route_404s_for_an_unknown_run(budgeted_client): assert budgeted_client.get("/v1/agent/runs/no-such-run/trace").status_code == 404 + + +# --------------------------------------------------------------------------- # +# a wait is not a new run: the ceiling and spend survive resume +# --------------------------------------------------------------------------- # + +class ApprovalGateway(FakeGateway): + """Planner parks on request_approval, then answers after the wait completes.""" + + async def chat(self, *, prompt, system, request=None): + request = dict(request or {}) + self.calls.append(request) + ceiling = int(request.get("max_tokens") or self.wanted_output) + if "evidence-readiness critic" in system: + text = json.dumps({"ready": True, "missing": [], "reason": "complete"}) + used_output = min(80, ceiling) + elif "decision core of a live-graph agent" in system: + context = json.loads(prompt) + nodes = context["graph"]["nodes"] + if not nodes: + text = json.dumps({ + "add": [{"id": "approve", "capability": "request_approval", + "arguments": {"question": "Send this outbound message?", + "choices": ["yes", "no"]}, + "depends_on": []}], + "cancel": [], "finish": False, + "reason": "outbound send needs a human", + }) + else: + text = json.dumps({ + "add": [{"id": "answer", "capability": "answer_with_evidence", + "arguments": {"query": context["goal"]}, + "depends_on": ["approve"]}], + "cancel": [], "finish": False, + "reason": "approval arrived", + }) + used_output = min(120, ceiling) + else: + text = "approved; the same run finished" + used_output = min(self.wanted_output, ceiling) + return {"text": text, "provider": "prov_1", + "model": request.get("model", "unknown"), + "input_tokens": len(prompt or "") // 4 + len(system or "") // 4, + "output_tokens": used_output, + "cache_read_input_tokens": 0, "cache_creation_input_tokens": 0, "latency_ms": 12} + + +def test_resume_keeps_the_run_budget_and_prior_spend(budgeted_client): + """HTTP resume used to pass budget=None, so BudgetedGateway was never rebuilt. + + A run that entered waiting under a ceiling could continue unmetered. The + control has to hold across the wait: same total, spend that only goes up. + """ + budgeted_client.app.state.gateway = ApprovalGateway() + first = budgeted_client.post("/v1/agent/runs", json={ + **SCOPE, "prompt": "Draft an outbound note and ask before sending.", + "budget": 0.05, "allowed_side_effects": ["request_approval"], + }) + assert first.status_code == 200, first.text + parked = first.json() + assert parked["status"] == "waiting" + assert parked["budget"] is not None, "the parked run must be under a ceiling" + assert parked["budget"]["total"] == 0.05 + spent_before = parked["budget"]["spent"] + assert spent_before > 0 + run_id = parked["run_id"] + + waiting = [node for node in parked["graph"]["nodes"].values() if node["state"] == "waiting"] + assert len(waiting) == 1 + wait = waiting[0]["wait"] + completion = budgeted_client.app.state.runtime.graph.complete_waiting( + wait["handle"], wait["event_type"], {"response": "yes"}, + ) + assert completion is not None + + # The resume route does not accept a budget body. The ceiling has to come + # back from the durable graph context, the way a killed-and-restarted + # process would see it. + resumed = budgeted_client.post(f"/v1/agent/runs/{run_id}/resume") + assert resumed.status_code == 200, resumed.text + body = resumed.json() + assert body["run_id"] == run_id + assert body["status"] == "completed" + assert body["budget"] is not None, "resume must not drop the run ceiling" + assert body["budget"]["total"] == 0.05 + assert body["budget"]["spent"] >= spent_before + assert body["budget"]["spent"] <= body["budget"]["total"] + assert body["budget"]["calls"] >= parked["budget"]["calls"] + persisted = budgeted_client.app.state.runtime.graph.context(run_id) + assert persisted["budget"] == 0.05 + assert float(persisted["budget_spent"]) == pytest.approx(body["budget"]["spent"])