mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
Merge pull request #201 from rohitg00/feat/phase-19-track-a-agent-harness-1
feat(phase-19): track A agent harness 20-24 deep capstones
This commit is contained in:
@@ -0,0 +1,343 @@
|
||||
"""Agent harness loop contract — deterministic state machine, hooks, pull points.
|
||||
|
||||
Conceptual references:
|
||||
- ./docs/en.md (this lesson)
|
||||
- Phase 14 lesson 01 (agent loop fundamentals)
|
||||
- Phase 13 lesson 02 (tool protocols overview)
|
||||
|
||||
Stdlib only. Run: python3 code/main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Iterable
|
||||
|
||||
|
||||
class State(str, Enum):
|
||||
IDLE = "idle"
|
||||
PLANNING = "planning"
|
||||
EXECUTING = "executing"
|
||||
AWAITING_TOOL = "awaiting_tool"
|
||||
REFLECTING = "reflecting"
|
||||
DONE = "done"
|
||||
|
||||
|
||||
HOOK_TOPICS = (
|
||||
"before_plan",
|
||||
"after_plan",
|
||||
"before_step",
|
||||
"after_step",
|
||||
"before_tool_call",
|
||||
"after_tool_call",
|
||||
"on_error",
|
||||
"on_pause",
|
||||
"on_budget_exceeded",
|
||||
"on_complete",
|
||||
)
|
||||
|
||||
EVENT_TYPES = (
|
||||
"session.start",
|
||||
"plan.draft",
|
||||
"plan.commit",
|
||||
"step.start",
|
||||
"step.end",
|
||||
"tool.call",
|
||||
"tool.result",
|
||||
"tool.error",
|
||||
"budget.warn",
|
||||
"session.pause",
|
||||
"session.complete",
|
||||
)
|
||||
|
||||
|
||||
class HookAbort(Exception):
|
||||
"""Raised by a hook to cancel the in-flight turn."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Event:
|
||||
type: str
|
||||
payload: dict
|
||||
ts: float
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {"type": self.type, "payload": self.payload, "ts": self.ts}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Budget:
|
||||
max_turns: int = 8
|
||||
max_tool_calls: int = 16
|
||||
max_wall_seconds: float = 30.0
|
||||
turns: int = 0
|
||||
tool_calls: int = 0
|
||||
started_at: float = field(default_factory=time.time)
|
||||
|
||||
def remaining_seconds(self) -> float:
|
||||
return max(0.0, self.max_wall_seconds - (time.time() - self.started_at))
|
||||
|
||||
def exceeded(self) -> str | None:
|
||||
if self.turns >= self.max_turns:
|
||||
return "turns"
|
||||
if self.tool_calls >= self.max_tool_calls:
|
||||
return "tool_calls"
|
||||
if self.remaining_seconds() <= 0.0:
|
||||
return "wall_clock"
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Step:
|
||||
id: int
|
||||
description: str
|
||||
requires_tool: bool
|
||||
tool_name: str | None = None
|
||||
tool_args: dict = field(default_factory=dict)
|
||||
result: Any = None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PullRequest:
|
||||
"""Returned from run()/resume() when the loop yields control."""
|
||||
reason: str
|
||||
state: State
|
||||
payload: dict
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionResult:
|
||||
state: State
|
||||
reason: str
|
||||
steps: list[Step]
|
||||
events: list[Event]
|
||||
|
||||
|
||||
class HookRegistry:
|
||||
def __init__(self) -> None:
|
||||
self._subs: dict[str, list[Callable[[dict], Any]]] = {t: [] for t in HOOK_TOPICS}
|
||||
|
||||
def on(self, topic: str, fn: Callable[[dict], Any]) -> None:
|
||||
if topic not in self._subs:
|
||||
raise ValueError(f"unknown hook topic: {topic}")
|
||||
self._subs[topic].append(fn)
|
||||
|
||||
def fire(self, topic: str, payload: dict) -> list[Any]:
|
||||
results = []
|
||||
for fn in self._subs[topic]:
|
||||
results.append(fn(payload))
|
||||
return results
|
||||
|
||||
|
||||
Planner = Callable[[str, list[Step]], list[Step]]
|
||||
|
||||
|
||||
def _default_planner(goal: str, history: list[Step]) -> list[Step]:
|
||||
"""Deterministic stand-in planner. Returns a fixed three-step plan."""
|
||||
if history:
|
||||
return []
|
||||
return [
|
||||
Step(id=1, description=f"interpret goal: {goal}", requires_tool=False),
|
||||
Step(id=2, description="fetch user record", requires_tool=True,
|
||||
tool_name="db.get_user", tool_args={"id": 42}),
|
||||
Step(id=3, description="summarize and respond", requires_tool=True,
|
||||
tool_name="format.summary", tool_args={"style": "short"}),
|
||||
]
|
||||
|
||||
|
||||
class HarnessLoop:
|
||||
"""Six-state deterministic loop with hook topics and event stream."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
planner: Planner | None = None,
|
||||
budget: Budget | None = None,
|
||||
) -> None:
|
||||
self.state: State = State.IDLE
|
||||
self.hooks = HookRegistry()
|
||||
self.budget = budget or Budget()
|
||||
self._planner: Planner = planner or _default_planner
|
||||
self._goal: str = ""
|
||||
self._plan: list[Step] = []
|
||||
self._cursor: int = 0
|
||||
self._events: list[Event] = []
|
||||
self._history: list[Step] = []
|
||||
self._reason: str = ""
|
||||
self._prev_state: State | None = None
|
||||
|
||||
@property
|
||||
def events(self) -> list[Event]:
|
||||
return list(self._events)
|
||||
|
||||
@property
|
||||
def plan(self) -> list[Step]:
|
||||
return list(self._plan)
|
||||
|
||||
def _emit(self, etype: str, payload: dict) -> None:
|
||||
if etype not in EVENT_TYPES:
|
||||
raise ValueError(f"unknown event type: {etype}")
|
||||
self._events.append(Event(type=etype, payload=payload, ts=time.time()))
|
||||
|
||||
def _transition(self, target: State) -> None:
|
||||
legal: dict[State, set[State]] = {
|
||||
State.IDLE: {State.PLANNING},
|
||||
State.PLANNING: {State.EXECUTING, State.IDLE, State.DONE},
|
||||
State.EXECUTING: {State.AWAITING_TOOL, State.REFLECTING, State.IDLE},
|
||||
State.AWAITING_TOOL: {State.REFLECTING, State.IDLE},
|
||||
State.REFLECTING: {State.PLANNING, State.EXECUTING, State.DONE, State.IDLE},
|
||||
State.DONE: set(),
|
||||
}
|
||||
if target not in legal[self.state]:
|
||||
raise RuntimeError(f"illegal transition {self.state.value} -> {target.value}")
|
||||
self.state = target
|
||||
|
||||
def _check_budget(self) -> PullRequest | None:
|
||||
which = self.budget.exceeded()
|
||||
if which is None:
|
||||
return None
|
||||
self._emit("budget.warn", {"limit": which})
|
||||
self.hooks.fire("on_budget_exceeded", {"limit": which, "budget": self.budget})
|
||||
self._reason = f"budget_exceeded:{which}"
|
||||
self._prev_state = self.state
|
||||
return self._pause(self._reason)
|
||||
|
||||
def _pause(self, reason: str) -> PullRequest:
|
||||
self._emit("session.pause", {"reason": reason})
|
||||
self.hooks.fire("on_pause", {"reason": reason})
|
||||
self._transition(State.IDLE)
|
||||
return PullRequest(reason=reason, state=self.state, payload={"reason": reason})
|
||||
|
||||
def run(self, goal: str) -> PullRequest | SessionResult:
|
||||
if self.state != State.IDLE:
|
||||
raise RuntimeError(f"run() requires IDLE, got {self.state.value}")
|
||||
self._goal = goal
|
||||
self.budget.started_at = time.time()
|
||||
self._emit("session.start", {"goal": goal})
|
||||
return self._step()
|
||||
|
||||
def resume(self, payload: dict | None = None) -> PullRequest | SessionResult:
|
||||
if self.state == State.IDLE and self._reason.startswith("budget_exceeded"):
|
||||
self.budget.turns = 0
|
||||
self.budget.tool_calls = 0
|
||||
self.budget.started_at = time.time()
|
||||
self._reason = ""
|
||||
prev = self._prev_state
|
||||
self._prev_state = None
|
||||
if not self._plan:
|
||||
return self._begin_plan()
|
||||
if prev == State.EXECUTING:
|
||||
self.state = State.EXECUTING
|
||||
else:
|
||||
self.state = State.REFLECTING
|
||||
return self._step()
|
||||
if self.state == State.AWAITING_TOOL:
|
||||
if payload is None:
|
||||
raise ValueError("resume from AWAITING_TOOL requires a payload")
|
||||
current = self._plan[self._cursor]
|
||||
if "error" in payload:
|
||||
current.error = str(payload["error"])
|
||||
self._emit("tool.error", {"step": current.id, "error": current.error})
|
||||
self.hooks.fire("on_error", {"step": current, "error": current.error})
|
||||
else:
|
||||
current.result = payload.get("result")
|
||||
self._emit("tool.result", {"step": current.id, "result": current.result})
|
||||
self.hooks.fire("after_tool_call", {"step": current})
|
||||
self._transition(State.REFLECTING)
|
||||
return self._step()
|
||||
raise RuntimeError(f"resume() unsupported from state {self.state.value}")
|
||||
|
||||
def _begin_plan(self) -> PullRequest | SessionResult:
|
||||
self._transition(State.PLANNING)
|
||||
self.hooks.fire("before_plan", {"goal": self._goal, "history": list(self._history)})
|
||||
draft = self._planner(self._goal, list(self._history))
|
||||
self._emit("plan.draft", {"steps": [s.description for s in draft]})
|
||||
self.hooks.fire("after_plan", {"steps": draft})
|
||||
self._plan = draft
|
||||
self._cursor = 0
|
||||
self._emit("plan.commit", {"count": len(draft)})
|
||||
if not draft:
|
||||
return self._complete("no_plan")
|
||||
self._transition(State.EXECUTING)
|
||||
return self._step()
|
||||
|
||||
def _step(self) -> PullRequest | SessionResult:
|
||||
if self.state == State.IDLE:
|
||||
return self._begin_plan()
|
||||
budget_hit = self._check_budget()
|
||||
if budget_hit is not None:
|
||||
return budget_hit
|
||||
if self.state == State.REFLECTING:
|
||||
self._cursor += 1
|
||||
self.budget.turns += 1
|
||||
if self._cursor >= len(self._plan):
|
||||
return self._complete("goal_met")
|
||||
self._transition(State.EXECUTING)
|
||||
return self._step()
|
||||
if self.state != State.EXECUTING:
|
||||
raise RuntimeError(f"_step requires EXECUTING/REFLECTING, got {self.state.value}")
|
||||
step = self._plan[self._cursor]
|
||||
self.hooks.fire("before_step", {"step": step})
|
||||
self._emit("step.start", {"step_id": step.id, "desc": step.description})
|
||||
if step.requires_tool:
|
||||
try:
|
||||
self.hooks.fire("before_tool_call", {"step": step})
|
||||
except HookAbort as exc:
|
||||
step.error = f"hook_abort:{exc}"
|
||||
self._emit("tool.error", {"step": step.id, "error": step.error})
|
||||
self.hooks.fire("on_error", {"step": step, "error": step.error})
|
||||
self._transition(State.REFLECTING)
|
||||
return self._step()
|
||||
self.budget.tool_calls += 1
|
||||
self._emit("tool.call", {"step": step.id, "tool": step.tool_name, "args": step.tool_args})
|
||||
self._transition(State.AWAITING_TOOL)
|
||||
self._emit("step.end", {"step_id": step.id, "outcome": "awaiting_tool"})
|
||||
self.hooks.fire("after_step", {"step": step, "outcome": "awaiting_tool"})
|
||||
return PullRequest(
|
||||
reason="tool_call",
|
||||
state=self.state,
|
||||
payload={"tool": step.tool_name, "args": step.tool_args, "step_id": step.id},
|
||||
)
|
||||
step.result = f"ok:{step.description}"
|
||||
self._emit("step.end", {"step_id": step.id, "outcome": "ok"})
|
||||
self.hooks.fire("after_step", {"step": step, "outcome": "ok"})
|
||||
self._transition(State.REFLECTING)
|
||||
return self._step()
|
||||
|
||||
def _complete(self, reason: str) -> SessionResult:
|
||||
self._emit("session.complete", {"reason": reason})
|
||||
self.hooks.fire("on_complete", {"reason": reason})
|
||||
self._transition(State.DONE)
|
||||
self._reason = reason
|
||||
return SessionResult(state=self.state, reason=reason, steps=list(self._plan), events=list(self._events))
|
||||
|
||||
|
||||
def _demo() -> None:
|
||||
loop = HarnessLoop()
|
||||
fired: list[str] = []
|
||||
for topic in HOOK_TOPICS:
|
||||
loop.hooks.on(topic, lambda payload, t=topic: fired.append(t))
|
||||
|
||||
out = loop.run("ship the release notes")
|
||||
assert isinstance(out, PullRequest) and out.reason == "tool_call"
|
||||
out = loop.resume({"result": {"id": 42, "name": "ada"}})
|
||||
assert isinstance(out, PullRequest) and out.reason == "tool_call"
|
||||
final = loop.resume({"result": "summary text"})
|
||||
assert isinstance(final, SessionResult)
|
||||
assert final.state == State.DONE
|
||||
assert final.reason == "goal_met"
|
||||
|
||||
report = {
|
||||
"events": [e.type for e in final.events],
|
||||
"hooks_fired": fired,
|
||||
"final_state": final.state.value,
|
||||
"final_reason": final.reason,
|
||||
}
|
||||
print(json.dumps(report, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
_demo()
|
||||
@@ -0,0 +1,208 @@
|
||||
"""Tests for HarnessLoop state machine, hooks, events, budget."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import unittest
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, os.path.dirname(HERE))
|
||||
|
||||
from main import ( # noqa: E402
|
||||
HOOK_TOPICS,
|
||||
Budget,
|
||||
HarnessLoop,
|
||||
HookAbort,
|
||||
PullRequest,
|
||||
SessionResult,
|
||||
State,
|
||||
Step,
|
||||
)
|
||||
|
||||
|
||||
def linear_planner(goal, history):
|
||||
if history:
|
||||
return []
|
||||
return [
|
||||
Step(id=1, description="step1", requires_tool=False),
|
||||
Step(id=2, description="step2", requires_tool=False),
|
||||
Step(id=3, description="step3", requires_tool=False),
|
||||
]
|
||||
|
||||
|
||||
def two_tool_planner(goal, history):
|
||||
if history:
|
||||
return []
|
||||
return [
|
||||
Step(id=1, description="prep", requires_tool=False),
|
||||
Step(id=2, description="fetch", requires_tool=True, tool_name="t.fetch", tool_args={}),
|
||||
Step(id=3, description="render", requires_tool=True, tool_name="t.render", tool_args={}),
|
||||
]
|
||||
|
||||
|
||||
class TestStateTransitions(unittest.TestCase):
|
||||
def test_idle_to_done_linear(self) -> None:
|
||||
loop = HarnessLoop(planner=linear_planner)
|
||||
result = loop.run("g")
|
||||
self.assertIsInstance(result, SessionResult)
|
||||
self.assertEqual(result.state, State.DONE)
|
||||
self.assertEqual(result.reason, "goal_met")
|
||||
|
||||
def test_run_twice_raises(self) -> None:
|
||||
loop = HarnessLoop(planner=linear_planner)
|
||||
loop.run("g")
|
||||
with self.assertRaises(RuntimeError):
|
||||
loop.run("again")
|
||||
|
||||
def test_tool_pull_point_then_resume(self) -> None:
|
||||
loop = HarnessLoop(planner=two_tool_planner)
|
||||
out = loop.run("g")
|
||||
self.assertIsInstance(out, PullRequest)
|
||||
self.assertEqual(out.reason, "tool_call")
|
||||
self.assertEqual(loop.state, State.AWAITING_TOOL)
|
||||
out2 = loop.resume({"result": 1})
|
||||
self.assertIsInstance(out2, PullRequest)
|
||||
final = loop.resume({"result": 2})
|
||||
self.assertIsInstance(final, SessionResult)
|
||||
self.assertEqual(final.state, State.DONE)
|
||||
|
||||
def test_resume_requires_payload(self) -> None:
|
||||
loop = HarnessLoop(planner=two_tool_planner)
|
||||
loop.run("g")
|
||||
with self.assertRaises(ValueError):
|
||||
loop.resume(None)
|
||||
|
||||
def test_illegal_transition_rejected(self) -> None:
|
||||
loop = HarnessLoop(planner=linear_planner)
|
||||
with self.assertRaises(RuntimeError):
|
||||
loop._transition(State.DONE)
|
||||
|
||||
def test_empty_plan_completes(self) -> None:
|
||||
def empty(goal, history):
|
||||
return []
|
||||
loop = HarnessLoop(planner=empty)
|
||||
result = loop.run("g")
|
||||
self.assertIsInstance(result, SessionResult)
|
||||
self.assertEqual(result.reason, "no_plan")
|
||||
|
||||
|
||||
class TestHooks(unittest.TestCase):
|
||||
def test_all_topics_register(self) -> None:
|
||||
loop = HarnessLoop()
|
||||
for t in HOOK_TOPICS:
|
||||
loop.hooks.on(t, lambda p: None)
|
||||
|
||||
def test_unknown_topic_rejected(self) -> None:
|
||||
loop = HarnessLoop()
|
||||
with self.assertRaises(ValueError):
|
||||
loop.hooks.on("not_a_topic", lambda p: None)
|
||||
|
||||
def test_hook_firing_order_linear(self) -> None:
|
||||
loop = HarnessLoop(planner=linear_planner)
|
||||
seen: list[str] = []
|
||||
for t in HOOK_TOPICS:
|
||||
loop.hooks.on(t, lambda p, t=t: seen.append(t))
|
||||
loop.run("g")
|
||||
self.assertEqual(seen[0], "before_plan")
|
||||
self.assertEqual(seen[1], "after_plan")
|
||||
self.assertEqual(seen[-1], "on_complete")
|
||||
self.assertEqual(seen.count("before_step"), 3)
|
||||
self.assertEqual(seen.count("after_step"), 3)
|
||||
self.assertNotIn("before_tool_call", seen)
|
||||
|
||||
def test_before_tool_call_fires_per_tool_step(self) -> None:
|
||||
loop = HarnessLoop(planner=two_tool_planner)
|
||||
before: list[int] = []
|
||||
after: list[int] = []
|
||||
loop.hooks.on("before_tool_call", lambda p: before.append(p["step"].id))
|
||||
loop.hooks.on("after_tool_call", lambda p: after.append(p["step"].id))
|
||||
loop.run("g")
|
||||
loop.resume({"result": "a"})
|
||||
loop.resume({"result": "b"})
|
||||
self.assertEqual(before, [2, 3])
|
||||
self.assertEqual(after, [2, 3])
|
||||
|
||||
def test_hook_abort_skips_tool_call(self) -> None:
|
||||
loop = HarnessLoop(planner=two_tool_planner)
|
||||
errors: list[str] = []
|
||||
loop.hooks.on("on_error", lambda p: errors.append(p["error"]))
|
||||
|
||||
def block(p):
|
||||
raise HookAbort("policy_denied")
|
||||
loop.hooks.on("before_tool_call", block)
|
||||
result = loop.run("g")
|
||||
self.assertIsInstance(result, SessionResult)
|
||||
self.assertEqual(len(errors), 2)
|
||||
self.assertTrue(errors[0].startswith("hook_abort"))
|
||||
|
||||
|
||||
class TestEvents(unittest.TestCase):
|
||||
def test_event_stream_shape(self) -> None:
|
||||
loop = HarnessLoop(planner=linear_planner)
|
||||
loop.run("g")
|
||||
types = [e.type for e in loop.events]
|
||||
self.assertEqual(types[0], "session.start")
|
||||
self.assertIn("plan.draft", types)
|
||||
self.assertIn("plan.commit", types)
|
||||
self.assertIn("step.start", types)
|
||||
self.assertIn("step.end", types)
|
||||
self.assertEqual(types[-1], "session.complete")
|
||||
|
||||
def test_tool_events_emitted(self) -> None:
|
||||
loop = HarnessLoop(planner=two_tool_planner)
|
||||
loop.run("g")
|
||||
loop.resume({"result": "x"})
|
||||
loop.resume({"result": "y"})
|
||||
types = [e.type for e in loop.events]
|
||||
self.assertEqual(types.count("tool.call"), 2)
|
||||
self.assertEqual(types.count("tool.result"), 2)
|
||||
self.assertNotIn("tool.error", types)
|
||||
|
||||
def test_tool_error_recorded(self) -> None:
|
||||
loop = HarnessLoop(planner=two_tool_planner)
|
||||
loop.run("g")
|
||||
loop.resume({"error": "boom"})
|
||||
types = [e.type for e in loop.events]
|
||||
self.assertIn("tool.error", types)
|
||||
|
||||
|
||||
class TestBudget(unittest.TestCase):
|
||||
def test_turn_limit_paused(self) -> None:
|
||||
budget = Budget(max_turns=1, max_tool_calls=10, max_wall_seconds=10.0)
|
||||
loop = HarnessLoop(planner=linear_planner, budget=budget)
|
||||
result = loop.run("g")
|
||||
self.assertIsInstance(result, PullRequest)
|
||||
self.assertTrue(result.reason.startswith("budget_exceeded"))
|
||||
|
||||
def test_tool_call_limit_paused(self) -> None:
|
||||
budget = Budget(max_turns=10, max_tool_calls=1, max_wall_seconds=10.0)
|
||||
loop = HarnessLoop(planner=two_tool_planner, budget=budget)
|
||||
out = loop.run("g")
|
||||
self.assertIsInstance(out, PullRequest)
|
||||
out2 = loop.resume({"result": "x"})
|
||||
self.assertIsInstance(out2, PullRequest)
|
||||
self.assertTrue(out2.reason.startswith("budget_exceeded"))
|
||||
|
||||
def test_wall_clock_check(self) -> None:
|
||||
budget = Budget(max_turns=10, max_tool_calls=10, max_wall_seconds=0.0)
|
||||
loop = HarnessLoop(planner=linear_planner, budget=budget)
|
||||
result = loop.run("g")
|
||||
self.assertIsInstance(result, PullRequest)
|
||||
self.assertEqual(result.reason, "budget_exceeded:wall_clock")
|
||||
|
||||
|
||||
class TestDeterminism(unittest.TestCase):
|
||||
def test_same_inputs_same_event_types(self) -> None:
|
||||
a = HarnessLoop(planner=linear_planner).run("g")
|
||||
b = HarnessLoop(planner=linear_planner).run("g")
|
||||
self.assertIsInstance(a, SessionResult)
|
||||
self.assertIsInstance(b, SessionResult)
|
||||
a_types = [e.type for e in a.events]
|
||||
b_types = [e.type for e in b.events]
|
||||
self.assertEqual(a_types, b_types)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,108 @@
|
||||
# Agent Harness Loop Contract
|
||||
|
||||
> The harness is the agent. The model is a coprocessor. This lesson freezes the loop contract you can wire any model into.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 13 lessons 01-07, Phase 14 lesson 01
|
||||
**Time:** ~90 minutes
|
||||
|
||||
## Learning Objectives
|
||||
- Specify an agent harness loop as a deterministic state machine with explicit transitions.
|
||||
- Implement ten lifecycle hook topics that operators wire policy, telemetry, and guardrails into.
|
||||
- Define two pull points where the loop yields control back to the caller and resumes on a fresh input.
|
||||
- Enforce per-session budgets (turns, tool calls, wall-clock) without leaking partial state on exceeding.
|
||||
- Emit a typed stream of eleven event types so downstream UIs and tracers can subscribe without inspecting the loop directly.
|
||||
|
||||
## The frame
|
||||
|
||||
A coding agent that runs unattended for forty turns is not a chat loop. It is a state machine whose nodes the operator can intercept and whose edges the operator can audit. Once you write the contract down, swapping models, tools, or policies stops being a refactor. It becomes a registration call.
|
||||
|
||||
This lesson builds that contract. We name six states, ten hook topics, two pull points, eleven event types, and a budget envelope. Everything else in the harness (tool registry, JSON-RPC transport, dispatcher, planner) plugs into this shape.
|
||||
|
||||
## The states
|
||||
|
||||
The loop has six states. Five are active. One is terminal.
|
||||
|
||||
```mermaid
|
||||
stateDiagram-v2
|
||||
[*] --> IDLE
|
||||
IDLE --> PLANNING: run(goal)
|
||||
PLANNING --> EXECUTING: plan committed
|
||||
EXECUTING --> AWAITING_TOOL: tool_call needed
|
||||
AWAITING_TOOL --> REFLECTING: result
|
||||
EXECUTING --> REFLECTING: no_tool step done
|
||||
REFLECTING --> EXECUTING: next step
|
||||
REFLECTING --> PLANNING: replan
|
||||
REFLECTING --> DONE: goal_met
|
||||
PLANNING --> DONE: no_plan
|
||||
DONE --> [*]
|
||||
```
|
||||
|
||||
`IDLE` is the only legal entry point. `DONE` is the only legal exit. `AWAITING_TOOL` is the only state that yields a pull point. Every other transition is internal.
|
||||
|
||||
The state machine is deterministic. Given the same event log, the harness re-enters the same state. That property is what lets you replay sessions for debugging without re-calling the model.
|
||||
|
||||
## The hook topics
|
||||
|
||||
Hooks are the operator's seam into the loop. The harness fires ten topics. Each topic accepts any number of subscribers. Subscribers fire in registration order. A subscriber may mutate the payload, raise to abort the turn, or return a sentinel to skip the next step.
|
||||
|
||||
```text
|
||||
before_plan after_plan
|
||||
before_tool_call after_tool_call
|
||||
before_step after_step
|
||||
on_error
|
||||
on_pause
|
||||
on_budget_exceeded
|
||||
on_complete
|
||||
```
|
||||
|
||||
The shape mirrors what Claude Code, Cursor, and OpenCode all converged on by mid-2025. The names are functional, not branded. A hook that blocks `rm -rf` lives in `before_tool_call`. A hook that ships an OpenTelemetry span lives in `after_step`. A hook that resumes on a paused session lives in `on_pause`.
|
||||
|
||||
## The pull points
|
||||
|
||||
The loop yields control twice. First on `AWAITING_TOOL` when it cannot make progress without a tool result. Second on `on_pause` when the budget is exhausted or a hook explicitly requests human review.
|
||||
|
||||
A pull point is not an exception. It is a return. The caller inspects the harness state, fetches whatever the harness asked for, and calls `resume(payload)`. The harness picks up where it stopped. This is the same shape as a Python generator. The transport over the pull point is your choice. In a TUI it is keypress. Over MCP it is `tools/call`. Over a queue it is a job poll.
|
||||
|
||||
## The event stream
|
||||
|
||||
The loop appends events to a typed stream at specific points in the contract. The stream is append-only and subscribers can replay from any offset. The eleven implemented event types are:
|
||||
|
||||
- `session.start` — emitted once when `run(goal)` is called
|
||||
- `plan.draft` — emitted when the planner returns a draft plan
|
||||
- `plan.commit` — emitted after the draft is committed as the active plan
|
||||
- `step.start` — emitted at the start of each executing step
|
||||
- `step.end` — emitted at the end of each executing step
|
||||
- `tool.call` — emitted when a tool-requiring step yields control to the caller
|
||||
- `tool.result` — emitted on resume with a tool result
|
||||
- `tool.error` — emitted on resume with an error or when a hook aborts the call
|
||||
- `budget.warn` — emitted when a budget limit is reached
|
||||
- `session.pause` — emitted when the loop yields on a pause (budget or hook)
|
||||
- `session.complete` — emitted once when the loop reaches `DONE`
|
||||
|
||||
The events do not duplicate hook payloads. Hooks are imperative (mutate, abort). Events are observational (record, ship). Treat them as orthogonal.
|
||||
|
||||
## The budget envelope
|
||||
|
||||
A session carries three limits. Turn count, tool call count, wall-clock seconds. Each turn increments turns by one. Each tool call increments tool calls by one. Wall-clock is checked on every state transition. When any limit is reached, the loop fires `on_budget_exceeded`, emits `budget.warn`, then transitions to `IDLE` with a budget-exceeded reason on the next pull point.
|
||||
|
||||
The budget is not a kill switch. It is a yield. The caller decides whether to extend the budget and resume, or to close the session.
|
||||
|
||||
## What this lesson does not do
|
||||
|
||||
It does not call a model. It does not register real tools. It does not implement a transport. Those are the next four lessons. This lesson nails the contract so the next four can plug into it without rewriting.
|
||||
|
||||
The deterministic planner in `main.py` is a stand-in. It returns a hardcoded plan of three steps, two of which require a tool result. The point is the loop, not the plan.
|
||||
|
||||
## How to read the code
|
||||
|
||||
`HarnessLoop` is the main class. It holds state, fires hooks, emits events. `Budget` tracks limits. `Event` is the typed envelope on the stream. `HookRegistry` is the dispatch table. `_transition` is the only function that changes state, so the state machine invariants live in one place.
|
||||
|
||||
Read `main.py` top to bottom. Then read `code/tests/test_loop.py`. The tests pin every transition and every hook firing order.
|
||||
|
||||
## Going further
|
||||
|
||||
The hardest part of building a harness in production is not the state machine. It is making the contract enforceable. The contract has to survive a hot reload of the planner. It has to survive a tool that returns malformed JSON. It has to survive a hook that raises in `before_tool_call` two-thirds of the way through a forty-turn session. The tests in this lesson exercise those failure modes. Run them. Break them. Add cases.
|
||||
|
||||
The next lesson adds the tool registry. After that, the JSON-RPC transport. After that, the dispatcher. By lesson twenty-four, the loop in this file will be running a real plan against real tools with real budgets enforced.
|
||||
@@ -0,0 +1,90 @@
|
||||
{
|
||||
"lesson": "20-agent-harness-loop-contract",
|
||||
"title": "Agent Harness Loop Contract",
|
||||
"questions": [
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Why model the harness loop as an explicit state machine rather than a while loop with flags?",
|
||||
"options": [
|
||||
"It lets the harness reject illegal transitions in one place and replay sessions from the event log",
|
||||
"It is faster on CPython",
|
||||
"It avoids the need for tests",
|
||||
"It is the only shape the model can call"
|
||||
],
|
||||
"correct": 0,
|
||||
"explanation": "A state machine puts the legality rules in one transition function. The same input sequence reproduces the same state, which is what makes session replay possible."
|
||||
},
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "What does a pull point in the loop represent?",
|
||||
"options": [
|
||||
"A point where the loop yields control and waits for the caller to provide an input",
|
||||
"A model call that returns a tool name",
|
||||
"A timer that fires every N seconds",
|
||||
"A hook subscriber that mutates payload"
|
||||
],
|
||||
"correct": 0,
|
||||
"explanation": "Pull points are returns, not exceptions. The caller fetches what the harness asked for and calls resume()."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Which state is the only legal pull point for a tool call?",
|
||||
"options": [
|
||||
"EXECUTING",
|
||||
"AWAITING_TOOL",
|
||||
"REFLECTING",
|
||||
"PLANNING"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Only AWAITING_TOOL yields a tool pull point. EXECUTING transitions into it when a step requires a tool."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "What happens when a hook raises HookAbort in before_tool_call?",
|
||||
"options": [
|
||||
"The session immediately transitions to DONE",
|
||||
"The tool call is skipped, on_error fires, and the loop moves to REFLECTING",
|
||||
"The hook is unregistered",
|
||||
"The budget is reset"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "HookAbort cancels the in-flight tool dispatch, fires on_error with the abort reason, and continues."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Why are hooks and events kept as separate concerns?",
|
||||
"options": [
|
||||
"Because event types are bytes and hooks are strings",
|
||||
"Hooks are imperative and can mutate or abort. Events are observational and append-only.",
|
||||
"Because the audit script forbids subscribers",
|
||||
"Because the event log cannot store dicts"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Hooks change behavior. Events describe behavior. Mixing them couples policy and telemetry, which is the bug we are avoiding."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "When the wall-clock budget is exceeded, what is the next state?",
|
||||
"options": [
|
||||
"DONE — the session is terminated",
|
||||
"EXECUTING — the loop retries",
|
||||
"IDLE — a pull point is returned with a budget_exceeded reason",
|
||||
"PLANNING — a fresh plan is drafted"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "Budget exhaustion is a yield. The harness transitions to IDLE, fires on_budget_exceeded, and returns a PullRequest so the caller can decide whether to extend or close."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "Why does this lesson stub out the planner instead of calling a model?",
|
||||
"options": [
|
||||
"Because models are slow to call locally",
|
||||
"To keep the loop contract testable and deterministic before any model is bound",
|
||||
"Because models cannot return tool calls",
|
||||
"To save tokens"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Pinning the loop contract first lets later lessons swap in real planners, registries, and dispatchers without renegotiating the state machine."
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,284 @@
|
||||
"""Tool registry with JSON Schema 2020-12 subset validation.
|
||||
|
||||
Conceptual references:
|
||||
- ./docs/en.md (this lesson)
|
||||
- IETF draft draft-bhutton-json-schema-2020-12 (subset: type, properties,
|
||||
required, enum, minLength, maxLength, pattern, items)
|
||||
- RFC 6901 (JSON Pointer for error paths)
|
||||
|
||||
Stdlib only. Run: python3 code/main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable
|
||||
|
||||
|
||||
PRIMITIVE_TYPE_MAP: dict[str, tuple[type, ...]] = {
|
||||
"string": (str,),
|
||||
"integer": (int,),
|
||||
"number": (int, float),
|
||||
"boolean": (bool,),
|
||||
"object": (dict,),
|
||||
"array": (list,),
|
||||
"null": (type(None),),
|
||||
}
|
||||
|
||||
ALLOWED_KEYWORDS = {
|
||||
"type", "properties", "required", "enum",
|
||||
"minLength", "maxLength", "pattern", "items", "description",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidationError:
|
||||
path: str
|
||||
keyword: str
|
||||
message: str
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {"path": self.path, "keyword": self.keyword, "message": self.message}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Ok:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolRecord:
|
||||
name: str
|
||||
description: str
|
||||
schema: dict
|
||||
handler: Callable[..., Any]
|
||||
idempotent: bool = False
|
||||
timeout_ms: int = 30_000
|
||||
|
||||
|
||||
class ToolRegistry:
|
||||
"""Name-keyed table of tool records with schema validation."""
|
||||
|
||||
_NAME_RE = re.compile(r"^[a-z][a-z0-9_]*(\.[a-z][a-z0-9_]*)*$")
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._records: dict[str, ToolRecord] = {}
|
||||
self._order: list[str] = []
|
||||
|
||||
def register(
|
||||
self,
|
||||
name: str,
|
||||
schema: dict,
|
||||
handler: Callable[..., Any],
|
||||
description: str = "",
|
||||
idempotent: bool = False,
|
||||
timeout_ms: int = 30_000,
|
||||
override: bool = False,
|
||||
) -> ToolRecord:
|
||||
if not self._NAME_RE.match(name):
|
||||
raise ValueError(f"tool name {name!r} must match {self._NAME_RE.pattern}")
|
||||
if name in self._records and not override:
|
||||
raise ValueError(f"tool {name!r} already registered; pass override=True to replace")
|
||||
validate_schema_shape(schema)
|
||||
rec = ToolRecord(
|
||||
name=name, description=description, schema=schema, handler=handler,
|
||||
idempotent=idempotent, timeout_ms=timeout_ms,
|
||||
)
|
||||
if name not in self._records:
|
||||
self._order.append(name)
|
||||
self._records[name] = rec
|
||||
return rec
|
||||
|
||||
def get(self, name: str) -> ToolRecord:
|
||||
if name not in self._records:
|
||||
raise KeyError(f"unknown tool {name!r}")
|
||||
return self._records[name]
|
||||
|
||||
def names(self) -> list[str]:
|
||||
return list(self._order)
|
||||
|
||||
def validate(self, name: str, args: Any) -> Ok | list[ValidationError]:
|
||||
rec = self.get(name)
|
||||
errors: list[ValidationError] = []
|
||||
_walk(rec.schema, args, "", errors)
|
||||
if errors:
|
||||
return errors
|
||||
return Ok()
|
||||
|
||||
|
||||
def validate_schema_shape(schema: dict) -> None:
|
||||
"""Reject schemas using keywords outside the supported subset."""
|
||||
if not isinstance(schema, dict):
|
||||
raise ValueError("schema must be a dict")
|
||||
unknown = set(schema.keys()) - ALLOWED_KEYWORDS
|
||||
if unknown:
|
||||
raise ValueError(f"unsupported schema keywords: {sorted(unknown)}")
|
||||
t = schema.get("type")
|
||||
if t is not None and t not in PRIMITIVE_TYPE_MAP:
|
||||
raise ValueError(f"unsupported type: {t!r}")
|
||||
enum_vals = schema.get("enum")
|
||||
if enum_vals is not None and not isinstance(enum_vals, list):
|
||||
raise ValueError("enum must be a list")
|
||||
min_len = schema.get("minLength")
|
||||
if min_len is not None:
|
||||
if isinstance(min_len, bool) or not isinstance(min_len, int) or min_len < 0:
|
||||
raise ValueError("minLength must be a non-negative integer")
|
||||
max_len = schema.get("maxLength")
|
||||
if max_len is not None:
|
||||
if isinstance(max_len, bool) or not isinstance(max_len, int) or max_len < 0:
|
||||
raise ValueError("maxLength must be a non-negative integer")
|
||||
if min_len is not None and max_len is not None and min_len > max_len:
|
||||
raise ValueError("minLength cannot be greater than maxLength")
|
||||
pattern = schema.get("pattern")
|
||||
if pattern is not None and not isinstance(pattern, str):
|
||||
raise ValueError("pattern must be a string")
|
||||
props = schema.get("properties")
|
||||
if props is not None:
|
||||
if not isinstance(props, dict):
|
||||
raise ValueError("properties must be a dict")
|
||||
for pname, psub in props.items():
|
||||
if not isinstance(pname, str):
|
||||
raise ValueError("property names must be strings")
|
||||
validate_schema_shape(psub)
|
||||
items = schema.get("items")
|
||||
if items is not None:
|
||||
validate_schema_shape(items)
|
||||
req = schema.get("required")
|
||||
if req is not None:
|
||||
if not isinstance(req, list) or not all(isinstance(x, str) for x in req):
|
||||
raise ValueError("required must be list[str]")
|
||||
|
||||
|
||||
def _path(prefix: str, segment: str | int) -> str:
|
||||
seg = str(segment).replace("~", "~0").replace("/", "~1")
|
||||
return f"{prefix}/{seg}"
|
||||
|
||||
|
||||
def _type_matches(value: Any, expected: str) -> bool:
|
||||
types = PRIMITIVE_TYPE_MAP[expected]
|
||||
if expected == "boolean":
|
||||
return isinstance(value, bool)
|
||||
if expected in ("integer", "number"):
|
||||
if isinstance(value, bool):
|
||||
return False
|
||||
return isinstance(value, types)
|
||||
return isinstance(value, types)
|
||||
|
||||
|
||||
def _walk(schema: dict, value: Any, path: str, errs: list[ValidationError]) -> None:
|
||||
t = schema.get("type")
|
||||
if t is not None and not _type_matches(value, t):
|
||||
errs.append(ValidationError(
|
||||
path=path or "/",
|
||||
keyword="type",
|
||||
message=f"expected {t}, got {type(value).__name__}",
|
||||
))
|
||||
return
|
||||
if "enum" in schema:
|
||||
if value not in schema["enum"]:
|
||||
errs.append(ValidationError(
|
||||
path=path or "/",
|
||||
keyword="enum",
|
||||
message=f"value {value!r} not in {schema['enum']!r}",
|
||||
))
|
||||
return
|
||||
if t == "string":
|
||||
_check_string(schema, value, path, errs)
|
||||
elif t == "object":
|
||||
_check_object(schema, value, path, errs)
|
||||
elif t == "array":
|
||||
_check_array(schema, value, path, errs)
|
||||
|
||||
|
||||
def _check_string(schema: dict, value: str, path: str, errs: list[ValidationError]) -> None:
|
||||
if "minLength" in schema and len(value) < schema["minLength"]:
|
||||
errs.append(ValidationError(
|
||||
path=path or "/", keyword="minLength",
|
||||
message=f"length {len(value)} < minLength {schema['minLength']}",
|
||||
))
|
||||
if "maxLength" in schema and len(value) > schema["maxLength"]:
|
||||
errs.append(ValidationError(
|
||||
path=path or "/", keyword="maxLength",
|
||||
message=f"length {len(value)} > maxLength {schema['maxLength']}",
|
||||
))
|
||||
if "pattern" in schema:
|
||||
try:
|
||||
if not re.search(schema["pattern"], value):
|
||||
errs.append(ValidationError(
|
||||
path=path or "/", keyword="pattern",
|
||||
message=f"value {value!r} does not match pattern {schema['pattern']!r}",
|
||||
))
|
||||
except re.error as exc:
|
||||
errs.append(ValidationError(
|
||||
path=path or "/", keyword="pattern",
|
||||
message=f"invalid regex: {exc}",
|
||||
))
|
||||
|
||||
|
||||
def _check_object(schema: dict, value: dict, path: str, errs: list[ValidationError]) -> None:
|
||||
required = schema.get("required", [])
|
||||
for req_name in required:
|
||||
if req_name not in value:
|
||||
errs.append(ValidationError(
|
||||
path=_path(path, req_name),
|
||||
keyword="required",
|
||||
message=f"missing required property {req_name!r}",
|
||||
))
|
||||
props = schema.get("properties", {})
|
||||
for prop_name, prop_value in value.items():
|
||||
if prop_name in props:
|
||||
_walk(props[prop_name], prop_value, _path(path, prop_name), errs)
|
||||
|
||||
|
||||
def _check_array(schema: dict, value: list, path: str, errs: list[ValidationError]) -> None:
|
||||
items_schema = schema.get("items")
|
||||
if items_schema is None:
|
||||
return
|
||||
for idx, item in enumerate(value):
|
||||
_walk(items_schema, item, _path(path, idx), errs)
|
||||
|
||||
|
||||
def _demo() -> None:
|
||||
registry = ToolRegistry()
|
||||
|
||||
def get_user(id: int) -> dict:
|
||||
return {"id": id, "name": "ada"}
|
||||
|
||||
registry.register(
|
||||
name="db.get_user",
|
||||
description="Fetch a user record by id.",
|
||||
schema={
|
||||
"type": "object",
|
||||
"required": ["id"],
|
||||
"properties": {
|
||||
"id": {"type": "integer"},
|
||||
"fields": {
|
||||
"type": "array",
|
||||
"items": {"type": "string", "enum": ["id", "name", "email"]},
|
||||
},
|
||||
},
|
||||
},
|
||||
handler=get_user,
|
||||
idempotent=True,
|
||||
)
|
||||
|
||||
cases = [
|
||||
{"id": 42, "fields": ["id", "name"]},
|
||||
{"id": "forty-two"},
|
||||
{"fields": ["id"]},
|
||||
{"id": 1, "fields": ["id", "phone"]},
|
||||
]
|
||||
report = []
|
||||
for c in cases:
|
||||
result = registry.validate("db.get_user", c)
|
||||
if isinstance(result, Ok):
|
||||
report.append({"args": c, "ok": True})
|
||||
else:
|
||||
report.append({"args": c, "ok": False, "errors": [e.to_dict() for e in result]})
|
||||
print(json.dumps({"tools": registry.names(), "cases": report}, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
_demo()
|
||||
+186
@@ -0,0 +1,186 @@
|
||||
"""Tests for ToolRegistry and JSON Schema subset validator."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, os.path.dirname(HERE))
|
||||
|
||||
from main import ( # noqa: E402
|
||||
Ok,
|
||||
ToolRecord,
|
||||
ToolRegistry,
|
||||
ValidationError,
|
||||
validate_schema_shape,
|
||||
)
|
||||
|
||||
|
||||
class TestRegistration(unittest.TestCase):
|
||||
def test_register_returns_record(self) -> None:
|
||||
r = ToolRegistry()
|
||||
rec = r.register(
|
||||
"fs.read",
|
||||
schema={"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]},
|
||||
handler=lambda path: open(path).read(),
|
||||
description="Read file",
|
||||
)
|
||||
self.assertIsInstance(rec, ToolRecord)
|
||||
self.assertEqual(rec.name, "fs.read")
|
||||
self.assertEqual(r.names(), ["fs.read"])
|
||||
|
||||
def test_duplicate_rejected_without_override(self) -> None:
|
||||
r = ToolRegistry()
|
||||
r.register("a", schema={"type": "string"}, handler=lambda x: x)
|
||||
with self.assertRaises(ValueError):
|
||||
r.register("a", schema={"type": "integer"}, handler=lambda x: x)
|
||||
|
||||
def test_override_replaces(self) -> None:
|
||||
r = ToolRegistry()
|
||||
r.register("a", schema={"type": "string"}, handler=lambda x: x)
|
||||
r.register("a", schema={"type": "integer"}, handler=lambda x: x, override=True)
|
||||
self.assertEqual(r.get("a").schema["type"], "integer")
|
||||
self.assertEqual(r.names(), ["a"])
|
||||
|
||||
def test_invalid_name_rejected(self) -> None:
|
||||
r = ToolRegistry()
|
||||
with self.assertRaises(ValueError):
|
||||
r.register("Bad-Name", schema={"type": "string"}, handler=lambda x: x)
|
||||
with self.assertRaises(ValueError):
|
||||
r.register("1starts-with-digit", schema={"type": "string"}, handler=lambda x: x)
|
||||
|
||||
def test_unknown_keyword_rejected(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
validate_schema_shape({"type": "object", "oneOf": []})
|
||||
|
||||
def test_unknown_type_rejected(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
validate_schema_shape({"type": "tuple"})
|
||||
|
||||
def test_get_unknown_raises(self) -> None:
|
||||
r = ToolRegistry()
|
||||
with self.assertRaises(KeyError):
|
||||
r.get("nope")
|
||||
|
||||
|
||||
class TestValidatorTypes(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.r = ToolRegistry()
|
||||
|
||||
def test_string_ok(self) -> None:
|
||||
self.r.register("s", schema={"type": "string"}, handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("s", "hi"), Ok)
|
||||
|
||||
def test_string_wrong_type(self) -> None:
|
||||
self.r.register("s", schema={"type": "string"}, handler=lambda x: x)
|
||||
errs = self.r.validate("s", 42)
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].keyword, "type")
|
||||
self.assertEqual(errs[0].path, "/")
|
||||
|
||||
def test_integer_vs_boolean(self) -> None:
|
||||
self.r.register("n", schema={"type": "integer"}, handler=lambda x: x)
|
||||
errs = self.r.validate("n", True)
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].keyword, "type")
|
||||
|
||||
def test_number_accepts_int_and_float(self) -> None:
|
||||
self.r.register("n", schema={"type": "number"}, handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("n", 1), Ok)
|
||||
self.assertIsInstance(self.r.validate("n", 1.5), Ok)
|
||||
|
||||
def test_null_type(self) -> None:
|
||||
self.r.register("z", schema={"type": "null"}, handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("z", None), Ok)
|
||||
errs = self.r.validate("z", 0)
|
||||
self.assertIsInstance(errs, list)
|
||||
|
||||
|
||||
class TestValidatorKeywords(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.r = ToolRegistry()
|
||||
|
||||
def test_min_max_length(self) -> None:
|
||||
self.r.register("s", schema={"type": "string", "minLength": 2, "maxLength": 4},
|
||||
handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("s", "abc"), Ok)
|
||||
e1 = self.r.validate("s", "a")
|
||||
self.assertIsInstance(e1, list)
|
||||
self.assertEqual(e1[0].keyword, "minLength")
|
||||
e2 = self.r.validate("s", "abcde")
|
||||
self.assertIsInstance(e2, list)
|
||||
self.assertEqual(e2[0].keyword, "maxLength")
|
||||
|
||||
def test_pattern(self) -> None:
|
||||
self.r.register("s", schema={"type": "string", "pattern": r"^[a-z]+$"},
|
||||
handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("s", "abc"), Ok)
|
||||
errs = self.r.validate("s", "abc1")
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].keyword, "pattern")
|
||||
|
||||
def test_enum(self) -> None:
|
||||
self.r.register("s", schema={"type": "string", "enum": ["a", "b"]},
|
||||
handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("s", "a"), Ok)
|
||||
errs = self.r.validate("s", "c")
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].keyword, "enum")
|
||||
|
||||
def test_required_missing(self) -> None:
|
||||
self.r.register("o", schema={
|
||||
"type": "object",
|
||||
"properties": {"id": {"type": "integer"}},
|
||||
"required": ["id"],
|
||||
}, handler=lambda **kw: kw)
|
||||
errs = self.r.validate("o", {})
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].keyword, "required")
|
||||
self.assertEqual(errs[0].path, "/id")
|
||||
|
||||
def test_nested_path(self) -> None:
|
||||
self.r.register("o", schema={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"user": {
|
||||
"type": "object",
|
||||
"properties": {"email": {"type": "string"}},
|
||||
"required": ["email"],
|
||||
},
|
||||
},
|
||||
"required": ["user"],
|
||||
}, handler=lambda **kw: kw)
|
||||
errs = self.r.validate("o", {"user": {"email": 0}})
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].path, "/user/email")
|
||||
|
||||
def test_array_items(self) -> None:
|
||||
self.r.register("a", schema={
|
||||
"type": "array",
|
||||
"items": {"type": "integer"},
|
||||
}, handler=lambda x: x)
|
||||
self.assertIsInstance(self.r.validate("a", [1, 2, 3]), Ok)
|
||||
errs = self.r.validate("a", [1, "x", 3])
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(errs[0].path, "/1")
|
||||
|
||||
def test_multiple_errors_collected(self) -> None:
|
||||
self.r.register("o", schema={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"a": {"type": "string"},
|
||||
"b": {"type": "integer"},
|
||||
},
|
||||
"required": ["a", "b"],
|
||||
}, handler=lambda **kw: kw)
|
||||
errs = self.r.validate("o", {})
|
||||
self.assertIsInstance(errs, list)
|
||||
self.assertEqual(len(errs), 2)
|
||||
paths = {e.path for e in errs}
|
||||
self.assertEqual(paths, {"/a", "/b"})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,99 @@
|
||||
# Tool Registry with Schema Validation
|
||||
|
||||
> A tool the agent cannot validate is a tool the agent cannot call. Build the registry and the schema checker before you build the tools.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 13 lessons 01-07, Phase 14 lesson 01
|
||||
**Time:** ~90 minutes
|
||||
|
||||
## Learning Objectives
|
||||
- Hold a typed registry of tool name → schema → handler that the dispatcher can ask once and trust afterwards.
|
||||
- Implement a JSON Schema 2020-12 subset that covers the keywords ninety percent of tool calls actually use.
|
||||
- Return precise, json-pointer-shaped error paths so the model can self-correct in one round trip.
|
||||
- Reject re-registration without explicit override, since silent overwrites are how production tool catalogs drift.
|
||||
- Keep the validator pure (no I/O, no time, no globals) so it can be re-run on a replay log.
|
||||
|
||||
## Why the registry comes before the tool
|
||||
|
||||
A coding agent in 2026 has more registered tools than the model can fit in a single context window. A non-trivial harness will register two hundred tools and surface ten to forty at any given turn. The registry is the source of truth for "what tools exist," "what shape do their arguments take," and "what handler do I call." Once those three answers are pinned, the rest of the harness can stop guessing.
|
||||
|
||||
The mistake we are avoiding is shipping handlers without schemas, or shipping schemas without validation. Both are common. Both turn the next layer (the dispatcher in lesson twenty-three) into a guessing game where the only failure mode is a stack trace from the handler.
|
||||
|
||||
## What a tool record looks like
|
||||
|
||||
```text
|
||||
ToolRecord
|
||||
name : str (unique, lowercase alphanumeric and underscore segments separated by dots, e.g., snake_case.segment.case)
|
||||
description : str (one line, shown to the model)
|
||||
schema : dict (JSON Schema 2020-12 subset)
|
||||
handler : Callable (async or sync, returns Any)
|
||||
idempotent : bool (dispatcher uses this for retry decisions)
|
||||
timeout_ms : int (override per-tool dispatcher default)
|
||||
```
|
||||
|
||||
The schema is the only field the validator touches. The handler is opaque to it. We separate them on purpose. The schema is data. The handler is code. Mixing them tempts you to put validation logic inside the handler, which is the bug we are stopping.
|
||||
|
||||
## The JSON Schema 2020-12 subset
|
||||
|
||||
The full 2020-12 spec is a paper. We need eight keywords.
|
||||
|
||||
```text
|
||||
type string / number / integer / boolean / object / array / null
|
||||
properties map of property name -> schema
|
||||
required list of property names
|
||||
enum list of allowed primitive values
|
||||
minLength integer, applies to strings
|
||||
maxLength integer, applies to strings
|
||||
pattern ECMA-262-compatible regex, applies to strings
|
||||
items schema applied to every array element
|
||||
```
|
||||
|
||||
That is enough to cover what a tool API actually needs. The keywords we are not adding (oneOf, anyOf, allOf, $ref, conditionals) are valid in production schemas but turn the validator into a tree walker with cycles. We are building a registry, not a JSON Schema engine.
|
||||
|
||||
## Json pointer error paths
|
||||
|
||||
When validation fails, the validator returns a list of errors. Each error carries a json-pointer path into the input. A pointer is a slash-prefixed sequence of property names and array indices.
|
||||
|
||||
```text
|
||||
{"a": {"b": [1, 2, "x"]}}
|
||||
^
|
||||
/a/b/2
|
||||
```
|
||||
|
||||
The model reads error paths better than it reads sentences. If a schema requires `args.user.email` and the model passed an integer, the error should be `/user/email` with `expected_type: string`. The model fixes that in the next call without a round of natural language.
|
||||
|
||||
## Registration and override
|
||||
|
||||
`register(name, schema, handler, **opts)` rejects re-registration by default. The caller has to pass `override=True` to replace. This is operational hygiene. Two parts of the codebase silently registering the same tool name is the kind of bug that takes a week to find in production.
|
||||
|
||||
The registry exposes three read methods. `get(name)` returns the record or raises. `validate(name, args)` returns an `Ok` or a list of errors. `names()` returns the tool names in registration order.
|
||||
|
||||
## What the validator is and is not
|
||||
|
||||
It is a single pass over the schema tree, recursive. It is pure. It does not call handlers. It does not coerce types (a string `"42"` does not pass a number schema). It does not silently truncate.
|
||||
|
||||
It is not a security boundary. A malicious handler can still misbehave after validation passes. The dispatcher in lesson twenty-three adds timeout and sandbox layers. The registry adds shape.
|
||||
|
||||
## Shape
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
code[your code]
|
||||
reg[ToolRegistry<br/>name<br/>schema<br/>handler<br/>timeout]
|
||||
out[Ok or list of errors]
|
||||
code -->|register name, schema, handler| reg
|
||||
reg -->|validate args| out
|
||||
```
|
||||
|
||||
## How to read the code
|
||||
|
||||
`code/main.py` defines `ToolRegistry`, `ToolRecord`, `ValidationError`, and the eight validator functions. The validator dispatches on `schema["type"]` (or treats a schema with `enum` as untyped enum check). Each type validator returns either an empty list or a list of `ValidationError`. The top-level walker concatenates errors and prepends path segments as it descends.
|
||||
|
||||
`code/tests/test_registry.py` covers registration, override, validation success, validation failure with paths, and every keyword in the subset.
|
||||
|
||||
## Going further
|
||||
|
||||
The two extensions you will want once this lesson lands are `$ref` resolution against a local definitions block, and `additionalProperties: false` for strict shape. Both are small. Both are common to add as the tool catalog grows past fifty tools. We left them out of the lesson to keep the file under one read.
|
||||
|
||||
The next lesson (twenty-two) builds the JSON-RPC stdio transport that surfaces this registry to a model client. The lesson after (twenty-three) wraps both behind a dispatcher with timeouts and retries.
|
||||
@@ -0,0 +1,78 @@
|
||||
{
|
||||
"lesson": "21-tool-registry-schema-validation",
|
||||
"title": "Tool Registry with Schema Validation",
|
||||
"questions": [
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Why register a tool's schema separately from its handler?",
|
||||
"options": [
|
||||
"Because schemas have to live in YAML",
|
||||
"So the validator stays pure data, and the handler stays opaque code",
|
||||
"Because handlers cannot return JSON",
|
||||
"Because schemas must be compiled"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Keeping the schema as data and the handler as code lets the validator be a pure function over the schema tree."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Which keywords are intentionally outside the subset this lesson implements?",
|
||||
"options": [
|
||||
"type, properties, required",
|
||||
"minLength, maxLength, pattern",
|
||||
"oneOf, anyOf, $ref, allOf",
|
||||
"items, enum"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "The omitted keywords add cycles and reference resolution. We keep the subset linear so the validator stays a single recursive walk."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Why does register() reject duplicates by default?",
|
||||
"options": [
|
||||
"Because the registry cannot store more than one record per name",
|
||||
"Silent overwrites are the cause of production tool catalog drift",
|
||||
"Because the schema cannot be re-validated",
|
||||
"Because handlers cannot be replaced"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Forcing override=True surfaces the case where two parts of the code intend to own the same tool name."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "What does the validator do when a schema specifies type integer but receives the value True?",
|
||||
"options": [
|
||||
"Treats it as 1 and passes",
|
||||
"Returns a type error because bool is not integer in this validator",
|
||||
"Raises an unhandled exception",
|
||||
"Coerces to integer silently"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Python bool is a subclass of int, but the validator filters bool out of the integer check. JSON booleans and integers are distinct."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "When the input object is {'user': {'email': 0}} and the schema requires user.email to be a string, what is the error path?",
|
||||
"options": [
|
||||
"/",
|
||||
"/user",
|
||||
"/user/email",
|
||||
"user.email"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "JSON pointer encodes the descent through user then email as /user/email."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "Why is the registry not a security boundary?",
|
||||
"options": [
|
||||
"Because handlers can still misbehave after validation passes",
|
||||
"Because handlers are async",
|
||||
"Because the schema is JSON",
|
||||
"Because the registry is in-process"
|
||||
],
|
||||
"correct": 0,
|
||||
"explanation": "The registry checks shape. Behavior, timeouts, and isolation belong to the dispatcher and the sandbox."
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,258 @@
|
||||
"""JSON-RPC 2.0 over newline-delimited stdio.
|
||||
|
||||
Conceptual references:
|
||||
- ./docs/en.md (this lesson)
|
||||
- JSON-RPC 2.0 specification (https://www.jsonrpc.org/specification)
|
||||
- RFC 8259 (JSON)
|
||||
|
||||
Stdlib only. Run: python3 code/main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import json
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, BinaryIO, Callable, Iterable
|
||||
|
||||
|
||||
ERR_PARSE = -32700
|
||||
ERR_INVALID_REQUEST = -32600
|
||||
ERR_METHOD_NOT_FOUND = -32601
|
||||
ERR_INVALID_PARAMS = -32602
|
||||
ERR_INTERNAL = -32603
|
||||
|
||||
|
||||
class JsonRpcError(Exception):
|
||||
code: int = ERR_INTERNAL
|
||||
|
||||
def __init__(self, message: str, data: Any | None = None) -> None:
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.data = data
|
||||
|
||||
|
||||
class MethodNotFound(JsonRpcError):
|
||||
code = ERR_METHOD_NOT_FOUND
|
||||
|
||||
|
||||
class InvalidParams(JsonRpcError):
|
||||
code = ERR_INVALID_PARAMS
|
||||
|
||||
|
||||
@dataclass
|
||||
class Request:
|
||||
method: str
|
||||
params: Any
|
||||
id: int | str | None
|
||||
is_notification: bool
|
||||
|
||||
|
||||
def _is_valid_envelope(msg: Any) -> bool:
|
||||
if not isinstance(msg, dict):
|
||||
return False
|
||||
if msg.get("jsonrpc") != "2.0":
|
||||
return False
|
||||
if not isinstance(msg.get("method"), str):
|
||||
return False
|
||||
if "params" in msg and not isinstance(msg["params"], (dict, list)):
|
||||
return False
|
||||
if "id" in msg:
|
||||
rid = msg["id"]
|
||||
if isinstance(rid, bool):
|
||||
return False
|
||||
if not isinstance(rid, (int, str, type(None))):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def parse_request(raw: str) -> tuple[Request | None, dict | None]:
|
||||
"""Parse one JSON line. Returns (Request, None) on success or (None, error_dict).
|
||||
|
||||
error_dict is a JSON-RPC error response ready to write out.
|
||||
"""
|
||||
try:
|
||||
msg = json.loads(raw)
|
||||
except json.JSONDecodeError as exc:
|
||||
return None, _err_envelope(None, ERR_PARSE, f"parse error: {exc}")
|
||||
if not _is_valid_envelope(msg):
|
||||
rid = msg.get("id") if isinstance(msg, dict) else None
|
||||
return None, _err_envelope(rid, ERR_INVALID_REQUEST, "invalid request envelope")
|
||||
is_notif = "id" not in msg
|
||||
return Request(
|
||||
method=msg["method"],
|
||||
params=msg.get("params"),
|
||||
id=msg.get("id"),
|
||||
is_notification=is_notif,
|
||||
), None
|
||||
|
||||
|
||||
def _err_envelope(rid: int | str | None, code: int, message: str, data: Any | None = None) -> dict:
|
||||
err: dict[str, Any] = {"code": code, "message": message}
|
||||
if data is not None:
|
||||
err["data"] = data
|
||||
return {"jsonrpc": "2.0", "id": rid, "error": err}
|
||||
|
||||
|
||||
def _ok_envelope(rid: int | str | None, result: Any) -> dict:
|
||||
return {"jsonrpc": "2.0", "id": rid, "result": result}
|
||||
|
||||
|
||||
def _notification_envelope(method: str, params: Any | None) -> dict:
|
||||
env: dict[str, Any] = {"jsonrpc": "2.0", "method": method}
|
||||
if params is not None:
|
||||
env["params"] = params
|
||||
return env
|
||||
|
||||
|
||||
class StdioTransport:
|
||||
"""Newline-delimited JSON-RPC 2.0 over a pair of byte streams."""
|
||||
|
||||
def __init__(self, stdin: BinaryIO, stdout: BinaryIO) -> None:
|
||||
self._in: BinaryIO = stdin
|
||||
self._out: BinaryIO = stdout
|
||||
|
||||
def read_line(self) -> bytes | None:
|
||||
line = self._in.readline()
|
||||
if not line:
|
||||
return None
|
||||
return line
|
||||
|
||||
def write_response(self, rid: int | str | None, result: Any) -> None:
|
||||
self._write(_ok_envelope(rid, result))
|
||||
|
||||
def write_error(self, rid: int | str | None, code: int, message: str, data: Any | None = None) -> None:
|
||||
self._write(_err_envelope(rid, code, message, data))
|
||||
|
||||
def write_notification(self, method: str, params: Any | None = None) -> None:
|
||||
self._write(_notification_envelope(method, params))
|
||||
|
||||
def _write(self, obj: dict) -> None:
|
||||
encoded = json.dumps(obj, separators=(",", ":"))
|
||||
self._out.write((encoded + "\n").encode("utf-8"))
|
||||
self._out.flush()
|
||||
|
||||
|
||||
Handler = Callable[[str, Any], Any]
|
||||
|
||||
|
||||
def _handle_one(handler: Handler, transport: StdioTransport, req: Request) -> dict | None:
|
||||
"""Dispatch a single Request. Returns the response envelope (or None for notification)."""
|
||||
try:
|
||||
result = handler(req.method, req.params)
|
||||
except MethodNotFound as exc:
|
||||
return None if req.is_notification else _err_envelope(req.id, exc.code, exc.message, exc.data)
|
||||
except InvalidParams as exc:
|
||||
return None if req.is_notification else _err_envelope(req.id, exc.code, exc.message, exc.data)
|
||||
except JsonRpcError as exc:
|
||||
return None if req.is_notification else _err_envelope(req.id, exc.code, exc.message, exc.data)
|
||||
except Exception as exc:
|
||||
return None if req.is_notification else _err_envelope(
|
||||
req.id, ERR_INTERNAL, "internal error",
|
||||
{"exception": type(exc).__name__, "detail": str(exc)},
|
||||
)
|
||||
if req.is_notification:
|
||||
return None
|
||||
return _ok_envelope(req.id, result)
|
||||
|
||||
|
||||
def _process_batch(handler: Handler, transport: StdioTransport, items: list) -> list | None:
|
||||
out: list = []
|
||||
for raw in items:
|
||||
if not isinstance(raw, dict) or not _is_valid_envelope(raw):
|
||||
rid = raw.get("id") if isinstance(raw, dict) else None
|
||||
out.append(_err_envelope(rid, ERR_INVALID_REQUEST, "invalid request envelope"))
|
||||
continue
|
||||
is_notif = "id" not in raw
|
||||
req = Request(
|
||||
method=raw["method"], params=raw.get("params"),
|
||||
id=raw.get("id"), is_notification=is_notif,
|
||||
)
|
||||
resp = _handle_one(handler, transport, req)
|
||||
if resp is not None:
|
||||
out.append(resp)
|
||||
if not out:
|
||||
return None
|
||||
return out
|
||||
|
||||
|
||||
def _write_raw(transport: StdioTransport, obj: Any) -> None:
|
||||
encoded = json.dumps(obj, separators=(",", ":"))
|
||||
transport._out.write((encoded + "\n").encode("utf-8"))
|
||||
transport._out.flush()
|
||||
|
||||
|
||||
def serve(handler: Handler, transport: StdioTransport) -> None:
|
||||
"""Read requests from transport until EOF. Dispatch each through handler."""
|
||||
while True:
|
||||
line = transport.read_line()
|
||||
if line is None:
|
||||
return
|
||||
text = line.decode("utf-8").rstrip("\n").strip()
|
||||
if not text:
|
||||
continue
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
except json.JSONDecodeError as exc:
|
||||
transport.write_error(None, ERR_PARSE, f"parse error: {exc}")
|
||||
continue
|
||||
if isinstance(parsed, list):
|
||||
if not parsed:
|
||||
transport.write_error(None, ERR_INVALID_REQUEST, "empty batch")
|
||||
continue
|
||||
batch_out = _process_batch(handler, transport, parsed)
|
||||
if batch_out is not None:
|
||||
_write_raw(transport, batch_out)
|
||||
continue
|
||||
req, err = parse_request(text)
|
||||
if err is not None:
|
||||
_write_raw(transport, err)
|
||||
continue
|
||||
resp = _handle_one(handler, transport, req)
|
||||
if resp is not None:
|
||||
_write_raw(transport, resp)
|
||||
|
||||
|
||||
def _demo() -> None:
|
||||
"""Self-terminating demo using io.BytesIO. No process spawn."""
|
||||
|
||||
def handler(method: str, params: Any) -> Any:
|
||||
if method == "math.add":
|
||||
if not isinstance(params, dict) or "a" not in params or "b" not in params:
|
||||
raise InvalidParams("a and b required")
|
||||
return params["a"] + params["b"]
|
||||
if method == "echo":
|
||||
return params
|
||||
if method == "boom":
|
||||
raise RuntimeError("intentional")
|
||||
raise MethodNotFound(f"method {method!r}")
|
||||
|
||||
requests = [
|
||||
{"jsonrpc": "2.0", "id": 1, "method": "math.add", "params": {"a": 2, "b": 3}},
|
||||
{"jsonrpc": "2.0", "id": 2, "method": "math.add", "params": {"a": 5}},
|
||||
{"jsonrpc": "2.0", "id": 3, "method": "missing"},
|
||||
{"jsonrpc": "2.0", "id": 4, "method": "boom"},
|
||||
{"jsonrpc": "2.0", "method": "log", "params": {"level": "info"}},
|
||||
[
|
||||
{"jsonrpc": "2.0", "id": 10, "method": "echo", "params": {"text": "hi"}},
|
||||
{"jsonrpc": "2.0", "method": "log", "params": {"msg": "skip-me"}},
|
||||
{"jsonrpc": "2.0", "id": 11, "method": "math.add", "params": {"a": 1, "b": 1}},
|
||||
],
|
||||
]
|
||||
|
||||
stdin = io.BytesIO()
|
||||
for r in requests:
|
||||
stdin.write((json.dumps(r) + "\n").encode("utf-8"))
|
||||
stdin.write(b"{not json\n")
|
||||
stdin.seek(0)
|
||||
stdout = io.BytesIO()
|
||||
transport = StdioTransport(stdin, stdout)
|
||||
serve(handler, transport)
|
||||
stdout.seek(0)
|
||||
lines = [json.loads(line) for line in stdout.read().decode("utf-8").splitlines() if line]
|
||||
print(json.dumps({"server_responses": lines}, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
_demo()
|
||||
@@ -0,0 +1,165 @@
|
||||
"""Tests for JSON-RPC 2.0 stdio transport: error codes, notifications, batches."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, os.path.dirname(HERE))
|
||||
|
||||
from main import ( # noqa: E402
|
||||
ERR_INTERNAL,
|
||||
ERR_INVALID_PARAMS,
|
||||
ERR_INVALID_REQUEST,
|
||||
ERR_METHOD_NOT_FOUND,
|
||||
ERR_PARSE,
|
||||
InvalidParams,
|
||||
MethodNotFound,
|
||||
StdioTransport,
|
||||
serve,
|
||||
)
|
||||
|
||||
|
||||
def _drive(requests, handler):
|
||||
"""Encode requests as newline-delimited JSON and run the server over them."""
|
||||
stdin = io.BytesIO()
|
||||
for r in requests:
|
||||
if isinstance(r, (bytes, bytearray)):
|
||||
stdin.write(r)
|
||||
else:
|
||||
stdin.write((json.dumps(r) + "\n").encode("utf-8"))
|
||||
stdin.seek(0)
|
||||
stdout = io.BytesIO()
|
||||
transport = StdioTransport(stdin, stdout)
|
||||
serve(handler, transport)
|
||||
stdout.seek(0)
|
||||
return [json.loads(line) for line in stdout.read().decode("utf-8").splitlines() if line.strip()]
|
||||
|
||||
|
||||
def echo_handler(method, params):
|
||||
if method == "echo":
|
||||
return params
|
||||
if method == "addone":
|
||||
if not isinstance(params, dict) or "n" not in params:
|
||||
raise InvalidParams("n required")
|
||||
return params["n"] + 1
|
||||
raise MethodNotFound(f"method {method!r}")
|
||||
|
||||
|
||||
class TestErrorCodes(unittest.TestCase):
|
||||
def test_parse_error(self) -> None:
|
||||
out = _drive([b"{ not json\n"], echo_handler)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_PARSE)
|
||||
self.assertIsNone(out[0]["id"])
|
||||
|
||||
def test_invalid_request_wrong_version(self) -> None:
|
||||
out = _drive([{"jsonrpc": "1.0", "id": 1, "method": "echo"}], echo_handler)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_INVALID_REQUEST)
|
||||
|
||||
def test_invalid_request_no_method(self) -> None:
|
||||
out = _drive([{"jsonrpc": "2.0", "id": 1}], echo_handler)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_INVALID_REQUEST)
|
||||
|
||||
def test_method_not_found(self) -> None:
|
||||
out = _drive([{"jsonrpc": "2.0", "id": 1, "method": "nope"}], echo_handler)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_METHOD_NOT_FOUND)
|
||||
self.assertEqual(out[0]["id"], 1)
|
||||
|
||||
def test_invalid_params(self) -> None:
|
||||
out = _drive([{"jsonrpc": "2.0", "id": 1, "method": "addone", "params": {}}], echo_handler)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_INVALID_PARAMS)
|
||||
|
||||
def test_internal_error_carries_exception_name(self) -> None:
|
||||
def bad(m, p):
|
||||
raise ValueError("kaboom")
|
||||
out = _drive([{"jsonrpc": "2.0", "id": 1, "method": "x"}], bad)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_INTERNAL)
|
||||
self.assertEqual(out[0]["error"]["data"]["exception"], "ValueError")
|
||||
|
||||
def test_boolean_id_rejected(self) -> None:
|
||||
out_true = _drive([{"jsonrpc": "2.0", "id": True, "method": "echo"}], echo_handler)
|
||||
self.assertEqual(out_true[0]["error"]["code"], ERR_INVALID_REQUEST)
|
||||
out_false = _drive([{"jsonrpc": "2.0", "id": False, "method": "echo"}], echo_handler)
|
||||
self.assertEqual(out_false[0]["error"]["code"], ERR_INVALID_REQUEST)
|
||||
|
||||
|
||||
class TestNotifications(unittest.TestCase):
|
||||
def test_notification_no_response(self) -> None:
|
||||
out = _drive([{"jsonrpc": "2.0", "method": "echo", "params": {"v": "x"}}], echo_handler)
|
||||
self.assertEqual(out, [])
|
||||
|
||||
def test_notification_handler_exception_silent(self) -> None:
|
||||
out = _drive([{"jsonrpc": "2.0", "method": "missing"}], echo_handler)
|
||||
self.assertEqual(out, [])
|
||||
|
||||
|
||||
class TestBatches(unittest.TestCase):
|
||||
def test_batch_mixed_returns_only_non_notifications(self) -> None:
|
||||
out = _drive([[
|
||||
{"jsonrpc": "2.0", "id": 1, "method": "echo", "params": {"v": "a"}},
|
||||
{"jsonrpc": "2.0", "method": "echo", "params": {"v": "b"}},
|
||||
{"jsonrpc": "2.0", "id": 3, "method": "echo", "params": {"v": "c"}},
|
||||
]], echo_handler)
|
||||
self.assertEqual(len(out), 1)
|
||||
self.assertIsInstance(out[0], list)
|
||||
self.assertEqual(len(out[0]), 2)
|
||||
ids = {r["id"] for r in out[0]}
|
||||
self.assertEqual(ids, {1, 3})
|
||||
|
||||
def test_batch_all_notifications_silent(self) -> None:
|
||||
out = _drive([[
|
||||
{"jsonrpc": "2.0", "method": "echo", "params": {"v": "a"}},
|
||||
{"jsonrpc": "2.0", "method": "echo", "params": {"v": "b"}},
|
||||
]], echo_handler)
|
||||
self.assertEqual(out, [])
|
||||
|
||||
def test_empty_batch_invalid_request(self) -> None:
|
||||
out = _drive([[]], echo_handler)
|
||||
self.assertEqual(len(out), 1)
|
||||
self.assertEqual(out[0]["error"]["code"], ERR_INVALID_REQUEST)
|
||||
|
||||
|
||||
class TestStreamDoesNotPoison(unittest.TestCase):
|
||||
def test_parse_error_then_continue(self) -> None:
|
||||
stdin = io.BytesIO()
|
||||
stdin.write(b"{ broken\n")
|
||||
stdin.write((json.dumps({"jsonrpc": "2.0", "id": 1, "method": "echo", "params": {"v": "ok"}}) + "\n").encode("utf-8"))
|
||||
stdin.seek(0)
|
||||
stdout = io.BytesIO()
|
||||
transport = StdioTransport(stdin, stdout)
|
||||
serve(echo_handler, transport)
|
||||
stdout.seek(0)
|
||||
lines = [json.loads(line) for line in stdout.read().decode("utf-8").splitlines() if line.strip()]
|
||||
self.assertEqual(lines[0]["error"]["code"], ERR_PARSE)
|
||||
self.assertEqual(lines[1]["result"], {"v": "ok"})
|
||||
|
||||
def test_empty_lines_skipped(self) -> None:
|
||||
stdin = io.BytesIO(b"\n\n" + json.dumps({"jsonrpc": "2.0", "id": 1, "method": "echo", "params": {"n": 5}}).encode("utf-8") + b"\n")
|
||||
stdout = io.BytesIO()
|
||||
transport = StdioTransport(stdin, stdout)
|
||||
serve(echo_handler, transport)
|
||||
stdout.seek(0)
|
||||
lines = [json.loads(line) for line in stdout.read().decode("utf-8").splitlines() if line.strip()]
|
||||
self.assertEqual(len(lines), 1)
|
||||
self.assertEqual(lines[0]["result"], {"n": 5})
|
||||
|
||||
|
||||
class TestNotificationHelper(unittest.TestCase):
|
||||
def test_write_notification_no_id(self) -> None:
|
||||
stdin = io.BytesIO()
|
||||
stdout = io.BytesIO()
|
||||
transport = StdioTransport(stdin, stdout)
|
||||
transport.write_notification("progress", {"pct": 50})
|
||||
stdout.seek(0)
|
||||
obj = json.loads(stdout.read().decode("utf-8").splitlines()[0])
|
||||
self.assertEqual(obj["method"], "progress")
|
||||
self.assertNotIn("id", obj)
|
||||
self.assertEqual(obj["jsonrpc"], "2.0")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,102 @@
|
||||
# JSON-RPC 2.0 Over Newline-Delimited Stdio
|
||||
|
||||
> The transport between a model client and a tool server is JSON-RPC over stdio. Hand-rolling it once teaches you what every framing layer is paying for.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 13 lessons 01-07, Phase 14 lesson 01
|
||||
**Time:** ~90 minutes
|
||||
|
||||
## Learning Objectives
|
||||
- Speak JSON-RPC 2.0 framed as newline-delimited JSON over stdin and stdout.
|
||||
- Map the five standard error codes (-32700, -32600, -32601, -32602, -32603) and surface them with the right semantics.
|
||||
- Distinguish requests, responses, notifications, and batches without inventing new envelope keys.
|
||||
- Handle one parse error per line without poisoning the rest of the stream.
|
||||
- Build a self-terminating demo using io.BytesIO so the lesson runs without spawning a child process.
|
||||
|
||||
## Why JSON-RPC stays the lingua franca
|
||||
|
||||
A coding agent in 2026 talks to maybe twelve tool servers in a single session. Each server is a separate process or a remote endpoint. The wire format has been the same since 2013. JSON-RPC 2.0 is two-page spec. It survives because the alternatives (gRPC, HTTP per call, custom binary) all impose a tradeoff JSON-RPC does not: they pick either streaming or batching or transport-coupling. JSON-RPC is symmetric across stdio, sockets, websockets, and HTTP, and a client can drive a server it has never seen if both honor the spec.
|
||||
|
||||
This lesson builds the stdio variant. Newline-delimited JSON. Each request is one line. Each response is one line. The transport boundary is `\n`.
|
||||
|
||||
## The wire shape
|
||||
|
||||
Four envelope shapes exist. Two are spoken by the client. Two are spoken by the server.
|
||||
|
||||
```mermaid
|
||||
sequenceDiagram
|
||||
participant Client
|
||||
participant Server
|
||||
Client->>Server: request {jsonrpc:"2.0", id:7, method:"foo", params:{...}}
|
||||
Server-->>Client: success {jsonrpc:"2.0", id:7, result:{...}}
|
||||
Client->>Server: notification {jsonrpc:"2.0", method:"bar", params:{...}} (no id)
|
||||
Note over Server: no response for notifications
|
||||
Client->>Server: request that fails
|
||||
Server-->>Client: error {jsonrpc:"2.0", id:7 or null, error:{code, message, data?}}
|
||||
```
|
||||
|
||||
A notification has no `id`. The server must not respond to it. If a server returns a response to a notification, the client has no way to attach it to a call site. That single rule keeps the framing math simple.
|
||||
|
||||
A batch is a JSON array of requests or notifications. The server replies with an array of responses, in any order, one per non-notification entry. If every entry in the batch is a notification, the server sends nothing back.
|
||||
|
||||
## The five error codes
|
||||
|
||||
```text
|
||||
-32700 Parse error JSON could not be parsed
|
||||
-32600 Invalid Request Envelope shape is wrong
|
||||
-32601 Method not found
|
||||
-32602 Invalid params
|
||||
-32603 Internal error
|
||||
```
|
||||
|
||||
The codes between -32000 and -32099 are reserved for server-defined errors. Everything else is application-defined. The lesson sticks to the five. If your handler raises, the transport wraps it as -32603 with the exception class name in `data.exception`.
|
||||
|
||||
A parse error has a special rule. The `id` in the response is `null`, because the request never parsed enough to extract an id.
|
||||
|
||||
## Newline framing and the BytesIO demo
|
||||
|
||||
The transport reads one line at a time. A line is bytes up to and including `\n`. If a line cannot be parsed, the transport writes a -32700 response with `id: null` and continues. The stream is not poisoned. The next line gets parsed fresh.
|
||||
|
||||
For the lesson we wrap an `io.BytesIO` pair as stdin and stdout. The server reads requests until EOF, writes responses for each, and returns. The client reads the responses back. No process spawn. No timeouts. The transport behavior is identical to a real subprocess pipe because Python's `io` interface presents the same `.readline()` and `.write()` contract.
|
||||
|
||||
## Method dispatch
|
||||
|
||||
The transport does not know which methods exist. It hands off to a callable `handler(method, params)` that the harness supplies. The handler returns a result or raises. Three exception classes surface specific codes.
|
||||
|
||||
```text
|
||||
MethodNotFound -> -32601
|
||||
InvalidParams -> -32602
|
||||
Anything else -> -32603 with exception name in data
|
||||
```
|
||||
|
||||
The transport never sees a tool registry. The registry sits behind the handler. This is the layering we want. The transport speaks JSON-RPC. The registry speaks tool shapes. The dispatcher (lesson twenty-three) stitches them together.
|
||||
|
||||
## Stream behavior on errors
|
||||
|
||||
```text
|
||||
client writes server reads server writes
|
||||
--------------- ----------- -------------
|
||||
{...valid request...} parses ok {...response, id matches...}
|
||||
{...broken json... parse fails {id:null, error: -32700}
|
||||
{...valid request...} parses ok {...response, id matches...}
|
||||
{...missing method...} invalid envelope {id:X, error: -32600}
|
||||
```
|
||||
|
||||
A broken JSON line does not stop the loop. A missing `method` field does not stop the loop. A handler exception does not stop the loop. The transport keeps reading until EOF.
|
||||
|
||||
## Notifications and asymmetric flows
|
||||
|
||||
A notification is fire-and-forget. The harness uses notifications for progress events, cancellation signals, and log lines. Notifications are how a long-running tool can stream status updates without round-tripping for each one.
|
||||
|
||||
The lesson implements one outbound notification helper, `write_notification`. The server uses it to emit progress while a request is in flight. The demo shows the pattern: a request comes in, the handler emits two progress notifications, then writes the final response.
|
||||
|
||||
## How to read the code
|
||||
|
||||
`code/main.py` defines `StdioTransport`, the parse helper (`parse_request`), the three write helpers (`write_response`, `write_error`, `write_notification`), and the dispatch loop `serve`. The error code constants live at module scope.
|
||||
|
||||
`code/tests/test_transport.py` covers the five error codes, notifications (no response written), batches (array in, array out, notifications skipped), broken JSON (parse error then continue), and the asymmetric flow where a handler writes a notification mid-call.
|
||||
|
||||
## Going further
|
||||
|
||||
This transport is enough for the lessons that follow. Production transports add three things. A correlation id field that survives forwarding (your `id` is already this, but in a mesh you need an outer trace id too). A cancellation channel (a notification like `$/cancelRequest` with the id of the in-flight call). And a content-type negotiation handshake so the same socket can speak JSON-RPC and Streamable HTTP. None of those change the wire. They add metadata.
|
||||
@@ -0,0 +1,90 @@
|
||||
{
|
||||
"lesson": "22-jsonrpc-stdio-transport",
|
||||
"title": "JSON-RPC 2.0 over Newline-Delimited Stdio",
|
||||
"questions": [
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "What distinguishes a JSON-RPC notification from a request on the wire?",
|
||||
"options": [
|
||||
"Different jsonrpc version field",
|
||||
"Notifications omit the id field; servers must not respond to them",
|
||||
"Notifications use a different method namespace",
|
||||
"Notifications are wrapped in an array"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Absence of id is the only marker. Servers do not respond, because there is no id to correlate against."
|
||||
},
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Why is newline-delimited JSON sufficient for stdio framing?",
|
||||
"options": [
|
||||
"Because JSON does not allow literal newlines inside strings without escaping, so each object can be one line",
|
||||
"Because pipes guarantee newline atomicity",
|
||||
"Because UTF-8 has no newline byte",
|
||||
"Because the spec mandates it"
|
||||
],
|
||||
"correct": 0,
|
||||
"explanation": "json.dumps emits a single line per object; literal newlines inside strings are escaped. That makes \\n a safe record separator."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Which error code does the transport return when JSON parsing fails?",
|
||||
"options": [
|
||||
"-32600",
|
||||
"-32700",
|
||||
"-32603",
|
||||
"-32601"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "-32700 is Parse error. The response id is null because no id could be extracted."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "If a handler raises a plain ValueError, what does the transport return?",
|
||||
"options": [
|
||||
"Method not found (-32601)",
|
||||
"Invalid params (-32602)",
|
||||
"Internal error (-32603) with the exception class name in data",
|
||||
"It does not return; the connection drops"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "Unmapped exceptions surface as -32603. data.exception preserves the class name for the client."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "What does the server return when a batch contains only notifications?",
|
||||
"options": [
|
||||
"An empty array",
|
||||
"Nothing — no bytes are written",
|
||||
"One null per notification",
|
||||
"A single null response"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Spec rule: if every entry is a notification, the server stays silent."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "A line with broken JSON arrives in the middle of a session. What happens to subsequent lines?",
|
||||
"options": [
|
||||
"The transport closes the stream",
|
||||
"They are buffered until the broken line is retried",
|
||||
"The transport writes a -32700 response and continues reading the next line as fresh",
|
||||
"They are all rejected with -32600"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "The framing layer must not poison the stream. One parse error means one error response and the loop continues."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "Why does the transport not consult the tool registry directly?",
|
||||
"options": [
|
||||
"Because the registry is async",
|
||||
"To keep the transport layer JSON-RPC-shaped only; the handler stitches transport to registry",
|
||||
"Because the registry holds secrets",
|
||||
"Because the transport runs in a different process"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "The transport speaks the wire. The dispatcher (next lesson) is where transport meets registry meets timeout."
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,352 @@
|
||||
"""Function call dispatcher with timeout, retry, idempotency, concurrency limit.
|
||||
|
||||
Conceptual references:
|
||||
- ./docs/en.md (this lesson)
|
||||
- JSON-RPC 2.0 specification (error envelope shape)
|
||||
- IETF draft draft-bhutton-json-schema-2020-12 (schema subset reused)
|
||||
|
||||
Stdlib only. Run: python3 code/main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Awaitable, Callable, Iterable
|
||||
|
||||
|
||||
ERR_METHOD_NOT_FOUND = -32601
|
||||
ERR_INVALID_PARAMS = -32602
|
||||
ERR_INTERNAL = -32603
|
||||
|
||||
|
||||
class TransientError(Exception):
|
||||
"""Raised by a handler to indicate the failure is worth retrying."""
|
||||
|
||||
|
||||
class _DispatchedError(Exception):
|
||||
"""Internal sentinel that wraps a DispatchError so dedup followers preserve kind."""
|
||||
|
||||
def __init__(self, error: "DispatchError") -> None:
|
||||
super().__init__(error.message)
|
||||
self.error = error
|
||||
|
||||
|
||||
@dataclass
|
||||
class DispatchError(Exception):
|
||||
kind: str
|
||||
message: str
|
||||
attempts: int
|
||||
jsonrpc_code: int = ERR_INTERNAL
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__init__(f"{self.kind}: {self.message}")
|
||||
|
||||
def to_envelope(self) -> dict:
|
||||
return {
|
||||
"code": self.jsonrpc_code,
|
||||
"message": self.message,
|
||||
"data": {"kind": self.kind, "attempts": self.attempts},
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class DispatchOk:
|
||||
result: Any
|
||||
attempts: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ToolRecord:
|
||||
name: str
|
||||
schema: dict
|
||||
handler: Callable[..., Any]
|
||||
idempotent: bool = False
|
||||
timeout_ms: int = 30_000
|
||||
|
||||
|
||||
class MiniRegistry:
|
||||
"""A trimmed registry: name, schema, handler, idempotent, timeout."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._recs: dict[str, _ToolRecord] = {}
|
||||
|
||||
def register(
|
||||
self, name: str, schema: dict, handler: Callable[..., Any],
|
||||
*, idempotent: bool = False, timeout_ms: int = 30_000,
|
||||
) -> None:
|
||||
self._recs[name] = _ToolRecord(
|
||||
name=name, schema=schema, handler=handler,
|
||||
idempotent=idempotent, timeout_ms=timeout_ms,
|
||||
)
|
||||
|
||||
def get(self, name: str) -> _ToolRecord:
|
||||
if name not in self._recs:
|
||||
raise KeyError(name)
|
||||
return self._recs[name]
|
||||
|
||||
def validate(self, name: str, args: Any) -> list[str]:
|
||||
rec = self.get(name)
|
||||
errs: list[str] = []
|
||||
_walk(rec.schema, args, "", errs)
|
||||
return errs
|
||||
|
||||
|
||||
def _walk(schema: dict, value: Any, path: str, errs: list[str]) -> None:
|
||||
t = schema.get("type")
|
||||
type_ok = True
|
||||
if t == "object" and not isinstance(value, dict):
|
||||
type_ok = False
|
||||
elif t == "integer" and (isinstance(value, bool) or not isinstance(value, int)):
|
||||
type_ok = False
|
||||
elif t == "string" and not isinstance(value, str):
|
||||
type_ok = False
|
||||
elif t == "array" and not isinstance(value, list):
|
||||
type_ok = False
|
||||
if not type_ok:
|
||||
errs.append(f"{path or '/'}: expected {t}, got {type(value).__name__}")
|
||||
return
|
||||
if t == "object":
|
||||
for req in schema.get("required", []):
|
||||
if req not in value:
|
||||
errs.append(f"{path}/{req}: required")
|
||||
for k, sub in schema.get("properties", {}).items():
|
||||
if k in value:
|
||||
_walk(sub, value[k], f"{path}/{k}", errs)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _InFlight:
|
||||
future: asyncio.Future
|
||||
started_at: float
|
||||
|
||||
|
||||
class Dispatcher:
|
||||
"""Per-call timeout, retry, idempotency, concurrency limit."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
registry: MiniRegistry,
|
||||
*,
|
||||
max_attempts: int = 3,
|
||||
concurrency: int = 8,
|
||||
cache_ttl_seconds: float = 60.0,
|
||||
sleep: Callable[[float], Awaitable[None]] | None = None,
|
||||
) -> None:
|
||||
if max_attempts <= 0:
|
||||
raise ValueError("max_attempts must be > 0")
|
||||
self.registry = registry
|
||||
self.max_attempts = max_attempts
|
||||
self._sem = asyncio.Semaphore(concurrency)
|
||||
self._inflight: dict[str, _InFlight] = {}
|
||||
self._cache: dict[str, tuple[Any, float]] = {}
|
||||
self._cache_ttl = cache_ttl_seconds
|
||||
self._sleep = sleep or asyncio.sleep
|
||||
|
||||
async def dispatch(
|
||||
self,
|
||||
name: str,
|
||||
args: dict,
|
||||
*,
|
||||
timeout_ms_override: int | None = None,
|
||||
idempotency_key: str | None = None,
|
||||
budget_tool_calls_remaining: int | None = None,
|
||||
) -> DispatchOk | DispatchError:
|
||||
try:
|
||||
rec = self.registry.get(name)
|
||||
except KeyError:
|
||||
return DispatchError(
|
||||
kind="not_found", message=f"tool {name!r}", attempts=0,
|
||||
jsonrpc_code=ERR_METHOD_NOT_FOUND,
|
||||
)
|
||||
|
||||
errs = self.registry.validate(name, args)
|
||||
if errs:
|
||||
return DispatchError(
|
||||
kind="schema", message="; ".join(errs), attempts=0,
|
||||
jsonrpc_code=ERR_INVALID_PARAMS,
|
||||
)
|
||||
|
||||
if budget_tool_calls_remaining is not None and budget_tool_calls_remaining <= 0:
|
||||
return DispatchError(
|
||||
kind="budget_exceeded", message="tool_calls remaining is 0", attempts=0,
|
||||
)
|
||||
|
||||
if idempotency_key is not None:
|
||||
now = time.monotonic()
|
||||
cached = self._cache.get(idempotency_key)
|
||||
if cached is not None and now - cached[1] < self._cache_ttl:
|
||||
return DispatchOk(result=cached[0], attempts=0)
|
||||
inflight = self._inflight.get(idempotency_key)
|
||||
if inflight is not None:
|
||||
try:
|
||||
res = await inflight.future
|
||||
return DispatchOk(result=res, attempts=0)
|
||||
except Exception as exc:
|
||||
return _map_exception(exc, attempts=0)
|
||||
|
||||
async with self._sem:
|
||||
return await self._run_with_retries(rec, args, timeout_ms_override, idempotency_key)
|
||||
|
||||
async def _run_with_retries(
|
||||
self,
|
||||
rec: _ToolRecord,
|
||||
args: dict,
|
||||
timeout_override: int | None,
|
||||
idempotency_key: str | None,
|
||||
) -> DispatchOk | DispatchError:
|
||||
timeout_ms = timeout_override if timeout_override is not None else rec.timeout_ms
|
||||
timeout_s = timeout_ms / 1000.0
|
||||
attempt = 0
|
||||
last_error: DispatchError | None = None
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
future: asyncio.Future | None = None
|
||||
if idempotency_key is not None:
|
||||
future = loop.create_future()
|
||||
self._inflight[idempotency_key] = _InFlight(future=future, started_at=time.monotonic())
|
||||
|
||||
try:
|
||||
while attempt < self.max_attempts:
|
||||
attempt += 1
|
||||
try:
|
||||
coro = _invoke(rec.handler, args)
|
||||
result = await asyncio.wait_for(coro, timeout=timeout_s)
|
||||
except asyncio.TimeoutError:
|
||||
last_error = DispatchError(
|
||||
kind="timeout",
|
||||
message=f"timeout after {timeout_ms}ms",
|
||||
attempts=attempt,
|
||||
)
|
||||
if not rec.idempotent:
|
||||
break
|
||||
if attempt >= self.max_attempts:
|
||||
break
|
||||
await self._sleep(_backoff(attempt))
|
||||
continue
|
||||
except TransientError as exc:
|
||||
last_error = DispatchError(
|
||||
kind="transient", message=str(exc), attempts=attempt,
|
||||
)
|
||||
if attempt >= self.max_attempts:
|
||||
break
|
||||
await self._sleep(_backoff(attempt))
|
||||
continue
|
||||
except Exception as exc:
|
||||
err = _map_exception(exc, attempts=attempt)
|
||||
if future is not None and not future.done():
|
||||
future.set_exception(_DispatchedError(err))
|
||||
return err
|
||||
if future is not None and not future.done():
|
||||
future.set_result(result)
|
||||
if idempotency_key is not None:
|
||||
self._cache[idempotency_key] = (result, time.monotonic())
|
||||
return DispatchOk(result=result, attempts=attempt)
|
||||
|
||||
assert last_error is not None
|
||||
if future is not None and not future.done():
|
||||
future.set_exception(_DispatchedError(last_error))
|
||||
return last_error
|
||||
finally:
|
||||
if idempotency_key is not None:
|
||||
self._inflight.pop(idempotency_key, None)
|
||||
|
||||
async def gather_bounded(self, calls: Iterable[tuple[str, dict]]) -> list[DispatchOk | DispatchError]:
|
||||
return await asyncio.gather(*(self.dispatch(n, a) for n, a in calls))
|
||||
|
||||
|
||||
async def _invoke(handler: Callable[..., Any], args: dict) -> Any:
|
||||
if inspect.iscoroutinefunction(handler):
|
||||
return await handler(**args)
|
||||
result = handler(**args)
|
||||
if inspect.isawaitable(result):
|
||||
return await result
|
||||
return result
|
||||
|
||||
|
||||
def _backoff(attempt: int) -> float:
|
||||
base = 0.1 * (4 ** (attempt - 1))
|
||||
return base * (1 + random.random() * 0.5)
|
||||
|
||||
|
||||
def _map_exception(exc: Exception, attempts: int) -> DispatchError:
|
||||
if isinstance(exc, _DispatchedError):
|
||||
original = exc.error
|
||||
return DispatchError(
|
||||
kind=original.kind,
|
||||
message=original.message,
|
||||
attempts=original.attempts,
|
||||
jsonrpc_code=original.jsonrpc_code,
|
||||
)
|
||||
return DispatchError(
|
||||
kind="internal",
|
||||
message=f"{type(exc).__name__}: {exc}",
|
||||
attempts=attempts,
|
||||
)
|
||||
|
||||
|
||||
async def _demo() -> None:
|
||||
reg = MiniRegistry()
|
||||
|
||||
counter = {"a": 0, "b": 0}
|
||||
|
||||
async def flaky_fetch(id: int) -> dict:
|
||||
counter["a"] += 1
|
||||
if counter["a"] < 2:
|
||||
raise TransientError("upstream not ready")
|
||||
return {"id": id, "name": "ada"}
|
||||
|
||||
async def slow(n: int) -> int:
|
||||
counter["b"] += 1
|
||||
await asyncio.sleep(0.05)
|
||||
return n
|
||||
|
||||
reg.register(
|
||||
"fetch_user",
|
||||
schema={"type": "object", "required": ["id"], "properties": {"id": {"type": "integer"}}},
|
||||
handler=flaky_fetch, idempotent=True, timeout_ms=200,
|
||||
)
|
||||
reg.register(
|
||||
"slow",
|
||||
schema={"type": "object", "required": ["n"], "properties": {"n": {"type": "integer"}}},
|
||||
handler=slow, idempotent=True, timeout_ms=10,
|
||||
)
|
||||
reg.register(
|
||||
"noop",
|
||||
schema={"type": "object", "properties": {}},
|
||||
handler=lambda: "ok",
|
||||
)
|
||||
|
||||
disp = Dispatcher(reg, max_attempts=3, concurrency=4)
|
||||
|
||||
out_retry = await disp.dispatch("fetch_user", {"id": 42})
|
||||
out_timeout = await disp.dispatch("slow", {"n": 1})
|
||||
out_schema = await disp.dispatch("fetch_user", {"id": "x"})
|
||||
out_missing = await disp.dispatch("does_not_exist", {})
|
||||
out_ok = await disp.dispatch("noop", {})
|
||||
|
||||
a, b = await asyncio.gather(
|
||||
disp.dispatch("noop", {}, idempotency_key="k1"),
|
||||
disp.dispatch("noop", {}, idempotency_key="k1"),
|
||||
)
|
||||
report = {
|
||||
"retry_then_success": {
|
||||
"attempts": getattr(out_retry, "attempts", None),
|
||||
"ok": isinstance(out_retry, DispatchOk),
|
||||
},
|
||||
"timeout": {"kind": getattr(out_timeout, "kind", None)},
|
||||
"schema": {"kind": getattr(out_schema, "kind", None)},
|
||||
"missing": {"kind": getattr(out_missing, "kind", None)},
|
||||
"happy": {"result": getattr(out_ok, "result", None)},
|
||||
"idempotency_pair": [isinstance(a, DispatchOk), isinstance(b, DispatchOk)],
|
||||
}
|
||||
print(json.dumps(report, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(_demo())
|
||||
+234
@@ -0,0 +1,234 @@
|
||||
"""Tests for the function call dispatcher: timeout, retry, idempotency, concurrency."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, os.path.dirname(HERE))
|
||||
|
||||
from main import ( # noqa: E402
|
||||
Dispatcher,
|
||||
DispatchError,
|
||||
DispatchOk,
|
||||
MiniRegistry,
|
||||
TransientError,
|
||||
)
|
||||
|
||||
|
||||
def run(coro):
|
||||
return asyncio.run(coro)
|
||||
|
||||
|
||||
def _registry_with(name: str, handler, **opts) -> MiniRegistry:
|
||||
r = MiniRegistry()
|
||||
r.register(name, schema={"type": "object", "properties": {}}, handler=handler, **opts)
|
||||
return r
|
||||
|
||||
|
||||
class TestSchemaPath(unittest.TestCase):
|
||||
def test_unknown_tool(self) -> None:
|
||||
r = MiniRegistry()
|
||||
d = Dispatcher(r)
|
||||
out = run(d.dispatch("absent", {}))
|
||||
self.assertIsInstance(out, DispatchError)
|
||||
self.assertEqual(out.kind, "not_found")
|
||||
self.assertEqual(out.jsonrpc_code, -32601)
|
||||
|
||||
def test_schema_failure_returns_invalid_params(self) -> None:
|
||||
r = MiniRegistry()
|
||||
r.register(
|
||||
"tool",
|
||||
schema={"type": "object", "required": ["id"], "properties": {"id": {"type": "integer"}}},
|
||||
handler=lambda id: id,
|
||||
)
|
||||
d = Dispatcher(r)
|
||||
out = run(d.dispatch("tool", {}))
|
||||
self.assertIsInstance(out, DispatchError)
|
||||
self.assertEqual(out.kind, "schema")
|
||||
self.assertEqual(out.jsonrpc_code, -32602)
|
||||
out2 = run(d.dispatch("tool", {"id": "nope"}))
|
||||
self.assertEqual(out2.kind, "schema")
|
||||
|
||||
def test_schema_error_does_not_retry(self) -> None:
|
||||
calls = []
|
||||
|
||||
def h(id):
|
||||
calls.append(id)
|
||||
return id
|
||||
r = MiniRegistry()
|
||||
r.register("tool",
|
||||
schema={"type": "object", "required": ["id"], "properties": {"id": {"type": "integer"}}},
|
||||
handler=h)
|
||||
d = Dispatcher(r)
|
||||
run(d.dispatch("tool", {"id": "x"}))
|
||||
self.assertEqual(calls, [])
|
||||
|
||||
|
||||
class TestTimeout(unittest.TestCase):
|
||||
def test_timeout_idempotent_retries(self) -> None:
|
||||
attempts = {"n": 0}
|
||||
|
||||
async def slow():
|
||||
attempts["n"] += 1
|
||||
await asyncio.sleep(0.05)
|
||||
return "done"
|
||||
|
||||
r = _registry_with("t", slow, idempotent=True, timeout_ms=5)
|
||||
|
||||
async def fake_sleep(_):
|
||||
return None
|
||||
|
||||
d = Dispatcher(r, max_attempts=3, sleep=fake_sleep)
|
||||
out = run(d.dispatch("t", {}))
|
||||
self.assertIsInstance(out, DispatchError)
|
||||
self.assertEqual(out.kind, "timeout")
|
||||
self.assertEqual(out.attempts, 3)
|
||||
self.assertEqual(attempts["n"], 3)
|
||||
|
||||
def test_timeout_non_idempotent_no_retry(self) -> None:
|
||||
attempts = {"n": 0}
|
||||
|
||||
async def slow():
|
||||
attempts["n"] += 1
|
||||
await asyncio.sleep(0.05)
|
||||
return "done"
|
||||
|
||||
r = _registry_with("t", slow, idempotent=False, timeout_ms=5)
|
||||
d = Dispatcher(r, max_attempts=3)
|
||||
out = run(d.dispatch("t", {}))
|
||||
self.assertEqual(out.kind, "timeout")
|
||||
self.assertEqual(attempts["n"], 1)
|
||||
|
||||
|
||||
class TestRetry(unittest.TestCase):
|
||||
def test_transient_retries_until_success(self) -> None:
|
||||
attempts = {"n": 0}
|
||||
|
||||
async def flaky():
|
||||
attempts["n"] += 1
|
||||
if attempts["n"] < 2:
|
||||
raise TransientError("not ready")
|
||||
return "ok"
|
||||
|
||||
r = _registry_with("t", flaky, idempotent=True, timeout_ms=200)
|
||||
|
||||
async def fake_sleep(_):
|
||||
return None
|
||||
|
||||
d = Dispatcher(r, max_attempts=3, sleep=fake_sleep)
|
||||
out = run(d.dispatch("t", {}))
|
||||
self.assertIsInstance(out, DispatchOk)
|
||||
self.assertEqual(out.attempts, 2)
|
||||
|
||||
def test_transient_exhausts_attempts(self) -> None:
|
||||
async def always_bad():
|
||||
raise TransientError("nope")
|
||||
|
||||
r = _registry_with("t", always_bad, idempotent=True, timeout_ms=200)
|
||||
|
||||
async def fake_sleep(_):
|
||||
return None
|
||||
|
||||
d = Dispatcher(r, max_attempts=2, sleep=fake_sleep)
|
||||
out = run(d.dispatch("t", {}))
|
||||
self.assertIsInstance(out, DispatchError)
|
||||
self.assertEqual(out.kind, "transient")
|
||||
self.assertEqual(out.attempts, 2)
|
||||
|
||||
def test_random_exception_no_retry(self) -> None:
|
||||
attempts = {"n": 0}
|
||||
|
||||
async def bad():
|
||||
attempts["n"] += 1
|
||||
raise RuntimeError("kaboom")
|
||||
|
||||
r = _registry_with("t", bad, idempotent=True, timeout_ms=200)
|
||||
d = Dispatcher(r, max_attempts=3)
|
||||
out = run(d.dispatch("t", {}))
|
||||
self.assertEqual(out.kind, "internal")
|
||||
self.assertEqual(attempts["n"], 1)
|
||||
|
||||
|
||||
class TestIdempotency(unittest.TestCase):
|
||||
def test_inflight_dedupe(self) -> None:
|
||||
counter = {"n": 0}
|
||||
ready = asyncio.Event()
|
||||
|
||||
async def slow_once():
|
||||
counter["n"] += 1
|
||||
await ready.wait()
|
||||
return counter["n"]
|
||||
|
||||
r = _registry_with("t", slow_once, idempotent=True, timeout_ms=2000)
|
||||
d = Dispatcher(r, max_attempts=1)
|
||||
|
||||
async def go():
|
||||
t1 = asyncio.create_task(d.dispatch("t", {}, idempotency_key="k"))
|
||||
await asyncio.sleep(0)
|
||||
t2 = asyncio.create_task(d.dispatch("t", {}, idempotency_key="k"))
|
||||
await asyncio.sleep(0)
|
||||
ready.set()
|
||||
return await asyncio.gather(t1, t2)
|
||||
|
||||
results = run(go())
|
||||
self.assertEqual(counter["n"], 1)
|
||||
self.assertTrue(all(isinstance(r, DispatchOk) for r in results))
|
||||
self.assertEqual(results[0].result, results[1].result)
|
||||
|
||||
def test_recent_cache_hit(self) -> None:
|
||||
counter = {"n": 0}
|
||||
|
||||
async def h():
|
||||
counter["n"] += 1
|
||||
return counter["n"]
|
||||
|
||||
r = _registry_with("t", h, idempotent=True, timeout_ms=200)
|
||||
d = Dispatcher(r, max_attempts=1, cache_ttl_seconds=60.0)
|
||||
|
||||
async def go():
|
||||
a = await d.dispatch("t", {}, idempotency_key="k")
|
||||
b = await d.dispatch("t", {}, idempotency_key="k")
|
||||
return a, b
|
||||
|
||||
a, b = run(go())
|
||||
self.assertEqual(counter["n"], 1)
|
||||
self.assertEqual(a.result, b.result)
|
||||
|
||||
|
||||
class TestConcurrency(unittest.TestCase):
|
||||
def test_semaphore_bounds_inflight(self) -> None:
|
||||
active = {"n": 0, "peak": 0}
|
||||
|
||||
async def h():
|
||||
active["n"] += 1
|
||||
active["peak"] = max(active["peak"], active["n"])
|
||||
await asyncio.sleep(0.01)
|
||||
active["n"] -= 1
|
||||
return 1
|
||||
|
||||
r = _registry_with("t", h, idempotent=True, timeout_ms=500)
|
||||
d = Dispatcher(r, max_attempts=1, concurrency=3)
|
||||
|
||||
async def go():
|
||||
tasks = [d.dispatch("t", {}) for _ in range(20)]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
run(go())
|
||||
self.assertLessEqual(active["peak"], 3)
|
||||
|
||||
|
||||
class TestBudget(unittest.TestCase):
|
||||
def test_zero_budget_fails_fast(self) -> None:
|
||||
r = _registry_with("t", lambda: 1)
|
||||
d = Dispatcher(r)
|
||||
out = run(d.dispatch("t", {}, budget_tool_calls_remaining=0))
|
||||
self.assertIsInstance(out, DispatchError)
|
||||
self.assertEqual(out.kind, "budget_exceeded")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,135 @@
|
||||
# Function Call Dispatcher
|
||||
|
||||
> The dispatcher is where the harness pays for every promise the schema made. Timeouts, retries, dedupe, error mapping. All on one seam.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 13 lessons 01-07, Phase 14 lesson 01
|
||||
**Time:** ~90 minutes
|
||||
|
||||
## Learning Objectives
|
||||
- Wrap a tool handler in a per-call timeout that returns a typed error instead of hanging the loop.
|
||||
- Apply exponential backoff retry with jitter and a maximum attempt count.
|
||||
- Deduplicate retries on an idempotency key so a retry that races with a slow original does not run twice.
|
||||
- Map handler exceptions and transport faults onto a single error envelope the harness loop already understands.
|
||||
- Bound parallel dispatch with a concurrency limit so a fan-out of forty tool calls does not exhaust the event loop.
|
||||
|
||||
## Where the dispatcher sits
|
||||
|
||||
Between the harness loop (lesson twenty) and the tool registry (lesson twenty-one). The transport (lesson twenty-two) feeds the loop. The loop hands a tool call to the dispatcher. The dispatcher calls the registry, runs the handler, and returns either a result or a JSON-RPC-shaped error envelope.
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
loop[harness loop]
|
||||
disp[dispatcher]
|
||||
reg[tool registry]
|
||||
handler[handler]
|
||||
loop --> disp
|
||||
disp -->|get name| reg
|
||||
disp -->|validate args| reg
|
||||
disp -->|asyncio.wait_for handler args timeout| handler
|
||||
handler -->|success| disp
|
||||
handler -->|TimeoutError -> retry or fail| disp
|
||||
handler -->|Exception -> map to error code| disp
|
||||
disp -->|Ok result or DispatchError| loop
|
||||
```
|
||||
|
||||
The dispatcher is the only layer that knows about timers, retries, and idempotency. The loop does not. The registry does not. The handler does not. That isolation is the point.
|
||||
|
||||
## Timeouts
|
||||
|
||||
Each tool has a default timeout. The registry record carries `timeout_ms`. The dispatcher overrides it from a per-call override when the harness passes one. We use `asyncio.wait_for`. On timeout, the handler task is cancelled and the dispatcher returns `DispatchError(kind="timeout")`.
|
||||
|
||||
A timeout is not a retryable error by default for non-idempotent tools. A `db.write` that timed out may or may not have committed. Retrying duplicates the write. The dispatcher honors the `idempotent` flag from the registry record. Idempotent tools retry. Non-idempotent tools do not.
|
||||
|
||||
## Retries with exponential backoff
|
||||
|
||||
The retry policy is three attempts maximum. Backoff is exponential with jitter.
|
||||
|
||||
```text
|
||||
attempt 1 -> delay 0
|
||||
attempt 2 -> delay 0.1s * (1 + random[0..0.5])
|
||||
attempt 3 -> delay 0.4s * (1 + random[0..0.5])
|
||||
```
|
||||
|
||||
Only `timeout` and `transient` errors retry. A `schema` error, a `not_found`, or an `internal` error does not retry. Schema errors are deterministic. Retrying does not change the outcome and burns the budget.
|
||||
|
||||
The retry loop respects the budget from the harness. If the caller's budget has zero remaining tool calls, the dispatcher fails fast on the first attempt and returns `kind="budget_exceeded"`.
|
||||
|
||||
## Idempotency key dedupe
|
||||
|
||||
A retry that fires while the original is still in flight is a real production bug. The first call hangs at four point nine seconds (just under the timeout). The retry fires at five seconds. Now two requests race against the same backend. If the tool is `payments.charge`, you charged twice.
|
||||
|
||||
The dispatcher accepts an optional `idempotency_key`. If the same key is in flight when a call arrives, the dispatcher waits on the in-flight future and returns its result. The cache holds keys for sixty seconds after completion to absorb late retries.
|
||||
|
||||
The key is the caller's responsibility. The harness derives it from the planner: `f"{step_id}:{tool_name}:{hash(args)}"`. The dispatcher does not invent keys, because deriving a key from arguments alone makes two semantically-different calls look the same.
|
||||
|
||||
## Error envelope
|
||||
|
||||
A failed dispatch returns a single shape.
|
||||
|
||||
```text
|
||||
DispatchError
|
||||
kind : "timeout" | "transient" | "schema" | "not_found" | "internal" | "budget_exceeded"
|
||||
message : str
|
||||
attempts : int
|
||||
jsonrpc_code: int (one of -32601, -32602, -32603)
|
||||
```
|
||||
|
||||
The harness loop maps `kind` to the next state. `schema` and `not_found` go to `on_error` and trigger a replan. `timeout` and `transient` go to `on_error` and may or may not replan depending on attempts. `budget_exceeded` triggers `on_budget_exceeded`.
|
||||
|
||||
## Concurrency limit on fan-out
|
||||
|
||||
`gather(*calls)` runs all coroutines simultaneously. With forty tool calls, that is forty open sockets or forty subprocess pipes. Most backends do not like forty parallel connections from one client.
|
||||
|
||||
The dispatcher wraps `gather` in a semaphore. Default concurrency limit is eight. Each call acquires the semaphore before dispatching and releases on completion. The caller sees `gather`-shaped output but the actual scheduling is bounded.
|
||||
|
||||
## Flow for one call
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
start([caller: dispatch name, args, opts])
|
||||
validate[registry.validate name, args]
|
||||
schema_err[DispatchError kind=schema]
|
||||
idem_check{idempotency cache?}
|
||||
in_flight[await existing future]
|
||||
cached[return cached result]
|
||||
attempt[asyncio.wait_for handler args, timeout]
|
||||
success[cache + return result]
|
||||
timeout_branch{TimeoutError + idempotent?}
|
||||
retry[retry with backoff]
|
||||
fail[DispatchError]
|
||||
transient_branch{TransientError?}
|
||||
other[map Exception to kind, no retry]
|
||||
exhausted[DispatchError]
|
||||
|
||||
start --> validate
|
||||
validate -->|errors| schema_err
|
||||
validate -->|ok| idem_check
|
||||
idem_check -->|hit in flight| in_flight
|
||||
idem_check -->|hit recent| cached
|
||||
idem_check -->|miss| attempt
|
||||
attempt --> success
|
||||
attempt --> timeout_branch
|
||||
timeout_branch -->|yes| retry
|
||||
timeout_branch -->|no| fail
|
||||
attempt --> transient_branch
|
||||
transient_branch -->|yes, attempts left| retry
|
||||
transient_branch -->|exhausted| exhausted
|
||||
attempt --> other
|
||||
retry --> attempt
|
||||
```
|
||||
|
||||
## How to read the code
|
||||
|
||||
`code/main.py` defines `Dispatcher`, `DispatchError`, and `TransientError`. The dispatcher takes a registry on construction. The async `dispatch(name, args, ...)` is the only entry point. Per-attempt timeouts are applied inline inside `_run_with_retries` using `asyncio.wait_for`. `gather_bounded(calls)` runs many dispatches with the concurrency limit.
|
||||
|
||||
`code/tests/test_dispatcher.py` covers timeout firing, retry on transient, no-retry on schema error, idempotency dedupe (two concurrent calls with the same key collapse to one handler invocation), and concurrency limiting (the semaphore in action).
|
||||
|
||||
The tests use `asyncio.sleep(0)` and deterministic `Counter`-based handlers, so they finish in milliseconds and do not depend on wall-clock timing.
|
||||
|
||||
## Going further
|
||||
|
||||
Two extensions production dispatchers add. First, structured logging at every transition (which the loop's event stream already gives you, but the dispatcher should also emit `dispatch.attempt` and `dispatch.retry` events). Second, circuit breakers: after N failures in a window, a tool gets a cool-down period where dispatches return immediately with `kind="circuit_open"` instead of attempting the handler. Both fit on top of this dispatcher without changing the contract.
|
||||
|
||||
Lesson twenty-four glues the dispatcher to a plan-and-execute agent so you see all four pieces in motion.
|
||||
@@ -0,0 +1,90 @@
|
||||
{
|
||||
"lesson": "23-function-call-dispatcher",
|
||||
"title": "Function Call Dispatcher",
|
||||
"questions": [
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Why does timeout on a non-idempotent tool default to no retry?",
|
||||
"options": [
|
||||
"Because the dispatcher cannot measure non-idempotent calls",
|
||||
"Because the tool may have partially committed; retry can duplicate the effect",
|
||||
"Because non-idempotent handlers run faster",
|
||||
"Because idempotency keys are slow"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Retrying a db.write that timed out at 4.9s can produce a second write. The dispatcher refuses unless the handler tells it the call is safe to repeat."
|
||||
},
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "What is the idempotency-key dedupe protecting against?",
|
||||
"options": [
|
||||
"Slow networks",
|
||||
"A retry that fires while the original call is still in flight",
|
||||
"Schema validation errors",
|
||||
"Process restarts"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Without dedupe, a retry at 5s races a slow original still pending. The cache collapses both into one handler invocation."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Which error kinds are eligible for retry in this dispatcher?",
|
||||
"options": [
|
||||
"schema and not_found",
|
||||
"internal and budget_exceeded",
|
||||
"timeout (if idempotent) and transient",
|
||||
"All of them"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "Schema and not_found are deterministic. Internal masks an unknown handler bug. Budget exceeded is a yield. Only timeout-on-idempotent and explicit TransientError retry."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "What does the semaphore around _run_with_retries protect?",
|
||||
"options": [
|
||||
"The handler from itself",
|
||||
"The backend from too many simultaneous calls fanned out by gather()",
|
||||
"The validator",
|
||||
"The retry counter"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "A gather of forty calls would open forty concurrent dispatches without the limit. The semaphore caps the wave."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Why is the idempotency key derived by the caller, not by the dispatcher?",
|
||||
"options": [
|
||||
"Because Python hashes are unstable",
|
||||
"Because deriving from args alone makes two semantically-different calls collide if they happen to share args",
|
||||
"Because the registry would not have the key",
|
||||
"Because the cache cannot store strings"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "The caller knows the step id, the planner intent, the user. Args alone are not enough context to safely dedupe."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "Which JSON-RPC error code maps to a schema failure on dispatch?",
|
||||
"options": [
|
||||
"-32601",
|
||||
"-32602",
|
||||
"-32603",
|
||||
"-32700"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "-32602 Invalid params. The dispatcher returns it before any handler runs."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "After max_attempts is reached with timeouts on an idempotent tool, what is returned?",
|
||||
"options": [
|
||||
"An exception is raised",
|
||||
"DispatchOk with no result",
|
||||
"DispatchError with kind=timeout and attempts==max_attempts",
|
||||
"DispatchError with kind=internal"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "The dispatcher returns a typed error envelope. The kind reflects the last failure mode, attempts reflects how many tries it made."
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
"""Plan-and-execute agent with replan on failure, plan diffs, and dual budgets.
|
||||
|
||||
Conceptual references:
|
||||
- ./docs/en.md (this lesson)
|
||||
- Phase 14 lesson 01 (agent loop fundamentals)
|
||||
- Phase 13 lesson 02 (tool protocols overview)
|
||||
|
||||
Stdlib only. Run: python3 code/main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable
|
||||
|
||||
|
||||
@dataclass
|
||||
class Step:
|
||||
id: int
|
||||
tool_name: str
|
||||
args: dict
|
||||
expected_outcome: str
|
||||
result: Any | None = None
|
||||
error: str | None = None
|
||||
|
||||
def signature(self) -> tuple:
|
||||
return (self.tool_name, json.dumps(self.args, sort_keys=True))
|
||||
|
||||
|
||||
@dataclass
|
||||
class PlanDiff:
|
||||
revision: int
|
||||
removed: list[int]
|
||||
added: list[int]
|
||||
revised: list[int]
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"revision": self.revision,
|
||||
"removed": list(self.removed),
|
||||
"added": list(self.added),
|
||||
"revised": list(self.revised),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Event:
|
||||
type: str
|
||||
payload: dict
|
||||
ts: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionResult:
|
||||
status: str
|
||||
reason: str
|
||||
history: list[Step]
|
||||
revisions: list[PlanDiff]
|
||||
events: list[Event]
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"status": self.status,
|
||||
"reason": self.reason,
|
||||
"history": [
|
||||
{"id": s.id, "tool": s.tool_name, "args": s.args,
|
||||
"result": s.result, "error": s.error}
|
||||
for s in self.history
|
||||
],
|
||||
"revisions": [r.to_dict() for r in self.revisions],
|
||||
"events": [{"type": e.type, "payload": e.payload, "ts": e.ts} for e in self.events],
|
||||
}
|
||||
|
||||
|
||||
Planner = Callable[[str, list[Step], str | None], list[Step]]
|
||||
ToolExecutor = Callable[[str, dict], Any]
|
||||
|
||||
|
||||
class ToolFailure(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _diff_plans(old: list[Step], new: list[Step], revision: int) -> PlanDiff:
|
||||
old_ids = {s.id for s in old}
|
||||
new_ids = {s.id for s in new}
|
||||
removed = sorted(old_ids - new_ids)
|
||||
added = sorted(new_ids - old_ids)
|
||||
revised: list[int] = []
|
||||
old_by_id = {s.id: s for s in old}
|
||||
for s in new:
|
||||
if s.id in old_ids and old_by_id[s.id].signature() != s.signature():
|
||||
revised.append(s.id)
|
||||
return PlanDiff(revision=revision, removed=removed, added=added, revised=revised)
|
||||
|
||||
|
||||
class PlanExecuteAgent:
|
||||
"""Sequential plan executor with replan on failure."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
planner: Planner,
|
||||
executor: ToolExecutor,
|
||||
*,
|
||||
max_steps: int = 12,
|
||||
max_replans: int = 5,
|
||||
) -> None:
|
||||
self._planner = planner
|
||||
self._executor = executor
|
||||
self.max_steps = max_steps
|
||||
self.max_replans = max_replans
|
||||
self._events: list[Event] = []
|
||||
|
||||
def _emit(self, etype: str, payload: dict) -> None:
|
||||
self._events.append(Event(type=etype, payload=payload))
|
||||
|
||||
def run(self, goal: str) -> SessionResult:
|
||||
self._events = []
|
||||
history: list[Step] = []
|
||||
revisions: list[PlanDiff] = []
|
||||
steps_taken = 0
|
||||
replans_used = 0
|
||||
last_error: str | None = None
|
||||
|
||||
plan = self._planner(goal, history, None)
|
||||
self._emit("plan.commit", {"revision": 0, "steps": _summarize(plan)})
|
||||
|
||||
if not plan:
|
||||
self._emit("session.complete", {"reason": "no_plan"})
|
||||
return SessionResult(
|
||||
status="failed", reason="no_plan",
|
||||
history=history, revisions=revisions, events=list(self._events),
|
||||
)
|
||||
|
||||
cursor = 0
|
||||
revision = 0
|
||||
|
||||
while cursor < len(plan):
|
||||
if steps_taken >= self.max_steps:
|
||||
self._emit("session.complete", {"reason": "step_budget"})
|
||||
return SessionResult(
|
||||
status="failed", reason="step_budget",
|
||||
history=history, revisions=revisions, events=list(self._events),
|
||||
)
|
||||
|
||||
step = plan[cursor]
|
||||
self._emit("step.start", {"step_id": step.id, "tool": step.tool_name})
|
||||
try:
|
||||
step.result = self._executor(step.tool_name, step.args)
|
||||
self._emit("step.end", {"step_id": step.id, "outcome": "ok"})
|
||||
history.append(step)
|
||||
cursor += 1
|
||||
steps_taken += 1
|
||||
continue
|
||||
except Exception as exc:
|
||||
step.error = f"{type(exc).__name__}: {exc}"
|
||||
self._emit("step.end", {"step_id": step.id, "outcome": "error", "error": step.error})
|
||||
history.append(step)
|
||||
steps_taken += 1
|
||||
last_error = step.error
|
||||
|
||||
if replans_used >= self.max_replans:
|
||||
self._emit("session.complete", {"reason": "replan_budget"})
|
||||
return SessionResult(
|
||||
status="failed", reason="replan_budget",
|
||||
history=history, revisions=revisions, events=list(self._events),
|
||||
)
|
||||
|
||||
replans_used += 1
|
||||
revision += 1
|
||||
new_plan = self._planner(goal, history, last_error)
|
||||
self._emit("plan.draft", {"revision": revision, "steps": _summarize(new_plan)})
|
||||
if not new_plan:
|
||||
self._emit("session.complete", {"reason": "no_plan"})
|
||||
return SessionResult(
|
||||
status="failed", reason="no_plan",
|
||||
history=history, revisions=revisions, events=list(self._events),
|
||||
)
|
||||
diff = _diff_plans(plan[cursor:], new_plan, revision)
|
||||
revisions.append(diff)
|
||||
self._emit("plan.diff", diff.to_dict())
|
||||
plan = new_plan
|
||||
cursor = 0
|
||||
self._emit("plan.commit", {"revision": revision, "steps": _summarize(plan)})
|
||||
|
||||
self._emit("session.complete", {"reason": "goal_met"})
|
||||
return SessionResult(
|
||||
status="completed", reason="goal_met",
|
||||
history=history, revisions=revisions, events=list(self._events),
|
||||
)
|
||||
|
||||
|
||||
def _summarize(plan: list[Step]) -> list[dict]:
|
||||
return [{"id": s.id, "tool": s.tool_name, "outcome": s.expected_outcome} for s in plan]
|
||||
|
||||
|
||||
def make_deterministic_planner(fail_step_id: int | None, recovery: str = "route_around") -> Planner:
|
||||
"""Planner used in the demo and tests.
|
||||
|
||||
When ``fail_step_id`` is given, the planner inserts a ``_force_fail`` marker
|
||||
into that step's args on the initial plan. Executors that honor the marker
|
||||
raise on that step, exercising the replan path. The marker is removed on the
|
||||
revised plan so the route-around can succeed.
|
||||
"""
|
||||
|
||||
def planner(goal: str, history: list[Step], last_error: str | None) -> list[Step]:
|
||||
if last_error is None:
|
||||
initial = [
|
||||
Step(1, "fetch", {"key": "input"}, "loaded user input"),
|
||||
Step(2, "transform", {"mode": "v1"}, "computed v1 form"),
|
||||
Step(3, "render", {}, "rendered output"),
|
||||
Step(4, "submit", {}, "submitted to backend"),
|
||||
]
|
||||
if fail_step_id is not None:
|
||||
for s in initial:
|
||||
if s.id == fail_step_id:
|
||||
s.args = {**s.args, "_force_fail": True}
|
||||
return initial
|
||||
if recovery == "route_around" and "transform" in last_error:
|
||||
return [
|
||||
Step(2, "transform", {"mode": "v2"}, "computed via fallback"),
|
||||
Step(3, "render", {}, "rendered output"),
|
||||
Step(4, "submit", {}, "submitted to backend"),
|
||||
]
|
||||
if recovery == "give_up":
|
||||
return [
|
||||
Step(98, "log_failure", {"why": last_error or ""}, "logged failure"),
|
||||
Step(99, "notify_user", {}, "told the user"),
|
||||
]
|
||||
return []
|
||||
|
||||
return planner
|
||||
|
||||
|
||||
def _demo() -> None:
|
||||
counters = {"transform_v1_calls": 0}
|
||||
|
||||
def executor(tool: str, args: dict) -> Any:
|
||||
if args.get("_force_fail"):
|
||||
counters["transform_v1_calls"] += 1
|
||||
raise ToolFailure(f"{tool} marker-forced failure")
|
||||
if tool == "fetch":
|
||||
return {"k": "v"}
|
||||
if tool == "transform":
|
||||
if args.get("mode") == "v1":
|
||||
counters["transform_v1_calls"] += 1
|
||||
raise ToolFailure("transform v1 backend down")
|
||||
return {"ok": True}
|
||||
if tool == "render":
|
||||
return "html"
|
||||
if tool == "submit":
|
||||
return {"id": 1}
|
||||
if tool in ("log_failure", "notify_user"):
|
||||
return "logged"
|
||||
raise ToolFailure(f"unknown tool {tool}")
|
||||
|
||||
agent = PlanExecuteAgent(
|
||||
planner=make_deterministic_planner(fail_step_id=2, recovery="route_around"),
|
||||
executor=executor,
|
||||
max_steps=12, max_replans=5,
|
||||
)
|
||||
res = agent.run("ship the report")
|
||||
print(json.dumps({
|
||||
"status": res.status,
|
||||
"reason": res.reason,
|
||||
"history": [(s.id, s.tool_name, bool(s.error)) for s in res.history],
|
||||
"revisions": [r.to_dict() for r in res.revisions],
|
||||
"events": [e.type for e in res.events],
|
||||
"transform_v1_calls": counters["transform_v1_calls"],
|
||||
}, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
_demo()
|
||||
@@ -0,0 +1,156 @@
|
||||
"""Tests for PlanExecuteAgent: linear, replan, replan exhaustion, step budget, diffs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, os.path.dirname(HERE))
|
||||
|
||||
from main import ( # noqa: E402
|
||||
PlanDiff,
|
||||
PlanExecuteAgent,
|
||||
SessionResult,
|
||||
Step,
|
||||
ToolFailure,
|
||||
make_deterministic_planner,
|
||||
)
|
||||
|
||||
|
||||
def perfect_executor(tool, args):
|
||||
return f"ok:{tool}"
|
||||
|
||||
|
||||
class TestLinear(unittest.TestCase):
|
||||
def test_linear_plan_completes(self) -> None:
|
||||
agent = PlanExecuteAgent(
|
||||
planner=make_deterministic_planner(fail_step_id=None),
|
||||
executor=perfect_executor,
|
||||
)
|
||||
res = agent.run("g")
|
||||
self.assertEqual(res.status, "completed")
|
||||
self.assertEqual(res.reason, "goal_met")
|
||||
self.assertEqual(len(res.history), 4)
|
||||
self.assertEqual([s.id for s in res.history], [1, 2, 3, 4])
|
||||
self.assertEqual(res.revisions, [])
|
||||
|
||||
def test_no_plan_initial(self) -> None:
|
||||
def empty(g, h, e):
|
||||
return []
|
||||
agent = PlanExecuteAgent(planner=empty, executor=perfect_executor)
|
||||
res = agent.run("g")
|
||||
self.assertEqual(res.status, "failed")
|
||||
self.assertEqual(res.reason, "no_plan")
|
||||
|
||||
|
||||
class TestReplan(unittest.TestCase):
|
||||
def test_replan_once_on_transform_failure(self) -> None:
|
||||
calls = {"transform_v1": 0, "transform_v2": 0}
|
||||
|
||||
def executor(tool, args):
|
||||
if tool == "transform":
|
||||
mode = args.get("mode")
|
||||
if mode == "v1":
|
||||
calls["transform_v1"] += 1
|
||||
raise ToolFailure("transform v1 down")
|
||||
if mode == "v2":
|
||||
calls["transform_v2"] += 1
|
||||
return "ok"
|
||||
raise ToolFailure("transform unknown mode")
|
||||
return f"ok:{tool}"
|
||||
|
||||
agent = PlanExecuteAgent(
|
||||
planner=make_deterministic_planner(fail_step_id=2, recovery="route_around"),
|
||||
executor=executor,
|
||||
)
|
||||
res = agent.run("g")
|
||||
self.assertEqual(res.status, "completed")
|
||||
self.assertEqual(res.reason, "goal_met")
|
||||
self.assertEqual(calls["transform_v1"], 1)
|
||||
self.assertEqual(calls["transform_v2"], 1)
|
||||
self.assertEqual(len(res.revisions), 1)
|
||||
|
||||
def test_replan_diff_event_emitted(self) -> None:
|
||||
def executor(tool, args):
|
||||
if tool == "transform" and args.get("mode") == "v1":
|
||||
raise ToolFailure("transform v1 boom")
|
||||
return "ok"
|
||||
|
||||
agent = PlanExecuteAgent(
|
||||
planner=make_deterministic_planner(None, recovery="route_around"),
|
||||
executor=executor,
|
||||
)
|
||||
res = agent.run("g")
|
||||
diff_events = [e for e in res.events if e.type == "plan.diff"]
|
||||
self.assertEqual(len(diff_events), 1)
|
||||
diff = diff_events[0].payload
|
||||
self.assertEqual(diff["revision"], 1)
|
||||
self.assertIn("removed", diff)
|
||||
self.assertIn("added", diff)
|
||||
self.assertIn("revised", diff)
|
||||
self.assertIn(2, diff["revised"])
|
||||
|
||||
def test_replan_exhaustion_returns_failed(self) -> None:
|
||||
def always_bad(tool, args):
|
||||
if tool == "transform":
|
||||
raise ToolFailure("nope")
|
||||
return "ok"
|
||||
|
||||
def planner(g, h, e):
|
||||
if e is None:
|
||||
return [
|
||||
Step(1, "fetch", {}, "fetch"),
|
||||
Step(2, "transform", {"mode": "v1"}, "transform"),
|
||||
]
|
||||
return [
|
||||
Step(2, "transform", {"mode": "v1"}, "transform again"),
|
||||
]
|
||||
|
||||
agent = PlanExecuteAgent(planner=planner, executor=always_bad, max_replans=2, max_steps=50)
|
||||
res = agent.run("g")
|
||||
self.assertEqual(res.status, "failed")
|
||||
self.assertEqual(res.reason, "replan_budget")
|
||||
|
||||
|
||||
class TestBudgets(unittest.TestCase):
|
||||
def test_step_budget_caps_execution(self) -> None:
|
||||
def planner(g, h, e):
|
||||
return [Step(i, "noop", {}, "noop") for i in range(20)]
|
||||
|
||||
agent = PlanExecuteAgent(planner=planner, executor=perfect_executor, max_steps=5)
|
||||
res = agent.run("g")
|
||||
self.assertEqual(res.status, "failed")
|
||||
self.assertEqual(res.reason, "step_budget")
|
||||
self.assertEqual(len(res.history), 5)
|
||||
|
||||
def test_max_replans_zero_returns_after_first_failure(self) -> None:
|
||||
def planner(g, h, e):
|
||||
return [Step(1, "bad", {}, "bad")]
|
||||
|
||||
def boom(tool, args):
|
||||
raise ToolFailure("nope")
|
||||
|
||||
agent = PlanExecuteAgent(planner=planner, executor=boom, max_replans=0)
|
||||
res = agent.run("g")
|
||||
self.assertEqual(res.status, "failed")
|
||||
self.assertEqual(res.reason, "replan_budget")
|
||||
|
||||
|
||||
class TestEvents(unittest.TestCase):
|
||||
def test_event_order(self) -> None:
|
||||
agent = PlanExecuteAgent(
|
||||
planner=make_deterministic_planner(None),
|
||||
executor=perfect_executor,
|
||||
)
|
||||
res = agent.run("g")
|
||||
types = [e.type for e in res.events]
|
||||
self.assertEqual(types[0], "plan.commit")
|
||||
self.assertEqual(types[-1], "session.complete")
|
||||
self.assertEqual(types.count("step.start"), 4)
|
||||
self.assertEqual(types.count("step.end"), 4)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,127 @@
|
||||
# Plan-Execute Control Flow
|
||||
|
||||
> A plan that cannot survive a failure is a script. A script that can replan is an agent. Build the replanner first.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 13 lessons 01-07, Phase 14 lesson 01
|
||||
**Time:** ~90 minutes
|
||||
|
||||
## Learning Objectives
|
||||
- Represent a plan as an ordered list of typed steps so the executor can reason about progress and outcome.
|
||||
- Execute steps sequentially with a controlled failure handoff back to the planner.
|
||||
- Replan from the current cursor with the prior error in the context so the next plan is informed.
|
||||
- Emit a plan diff on each revision so a downstream tracer or UI can show why the plan changed.
|
||||
- Enforce two budgets: a hard step ceiling and a hard replan ceiling.
|
||||
|
||||
## Plan and execute, not chain-of-thought
|
||||
|
||||
A chain-of-thought agent emits tokens and lets the loop guess where the tool call ends. A plan-and-execute agent emits a structured plan first, then executes each step deterministically. The plan is data the harness can introspect. The execution is the harness running that data through a dispatcher.
|
||||
|
||||
Two pieces. A planner that produces a plan. An executor that runs the plan. The interesting work is what happens when the executor hits a failure. Three options:
|
||||
|
||||
```text
|
||||
1. Abort (return failed, surface the error)
|
||||
2. Skip (mark step failed, continue with the rest)
|
||||
3. Replan (hand the error to the planner, get a new plan from the cursor)
|
||||
```
|
||||
|
||||
Replan is the one that turns a script into an agent.
|
||||
|
||||
## The Step shape
|
||||
|
||||
```text
|
||||
Step
|
||||
id : int (monotonic within a plan revision)
|
||||
tool_name : str
|
||||
args : dict
|
||||
expected_outcome: str (planner's stated success condition)
|
||||
result : Any | None
|
||||
error : str | None
|
||||
```
|
||||
|
||||
`expected_outcome` is a short sentence the planner emits alongside the step. It is not enforced by the executor. It is for two things: the replanner reads it when revising the plan; the event stream emits it so a tracer can show "this step was supposed to do X."
|
||||
|
||||
## The planner shape
|
||||
|
||||
```python
|
||||
def planner(goal: str, history: list[Step], last_error: str | None) -> list[Step]:
|
||||
...
|
||||
```
|
||||
|
||||
A pure function. `goal` is the user goal. `history` is the steps already executed (with results and errors filled in). `last_error` is None on the first call and the most recent failure message on every subsequent call. The planner returns the next plan starting from the cursor.
|
||||
|
||||
The planner does not know about the executor. It does not know about retries. It does not know about timeouts. It produces a plan. That is all.
|
||||
|
||||
## The executor
|
||||
|
||||
The executor is a small state machine. Each step runs through the dispatcher. The outcome is one of three things: success, failure-replannable, failure-fatal. Replannable failures hand back to the planner. Fatal failures (budget exceeded, replan ceiling hit) return a `FAILED` session result.
|
||||
|
||||
```mermaid
|
||||
stateDiagram-v2
|
||||
[*] --> EXEC
|
||||
EXEC --> NEXT: success
|
||||
NEXT --> EXEC: n+1 < len(plan)
|
||||
NEXT --> DONE: n+1 == len(plan)
|
||||
EXEC --> REPLAN: failure
|
||||
REPLAN --> EXEC: new plan, replans_used < max_replans
|
||||
REPLAN --> FAILED: replans_used >= max_replans
|
||||
FAILED --> [*]
|
||||
DONE --> [*]
|
||||
```
|
||||
|
||||
## Plan diffs on revision
|
||||
|
||||
When the planner returns a new plan after a failure, the executor emits a `plan.diff` event with three fields.
|
||||
|
||||
```text
|
||||
removed: list of step ids that were in the old plan and are not in the new
|
||||
added : list of step ids in the new plan that were not in the old
|
||||
revised: list of step ids whose tool_name or args changed
|
||||
```
|
||||
|
||||
A tracer or UI can render this as a strikethrough on the removed steps and a highlight on the added ones. The point is not the diff format. The point is that revision is a visible event, not a silent rewrite.
|
||||
|
||||
## Two budgets, both hard
|
||||
|
||||
`max_steps` caps total step executions across the whole session, including replans. Default is twelve. A linear five-step plan that replans twice and adds three steps each time hits sixteen executions and would exceed the budget. The executor will refuse the replan and return FAILED.
|
||||
|
||||
`max_replans` caps the number of times the planner is called after the first plan. Default is five. This is the more important limit. A planner that returns the same broken plan five times in a row would otherwise loop until the step budget catches it. Capping replans makes the failure faster and the reason clearer.
|
||||
|
||||
## The deterministic planner in this lesson
|
||||
|
||||
We do not call a model in this lesson. The lesson ships a deterministic planner that picks a plan based on `last_error`.
|
||||
|
||||
```text
|
||||
last_error is None -> emit a four-step plan
|
||||
last_error matches X -> emit a three-step plan that routes around X
|
||||
last_error matches Y -> emit a two-step plan that gives up gracefully
|
||||
otherwise -> return [] (signals nothing to replan)
|
||||
```
|
||||
|
||||
This is enough to test the executor's behavior on every transition path: success, replan-once, replan-twice, replan-exhaustion, and step-budget exhaustion.
|
||||
|
||||
## Result shape
|
||||
|
||||
```text
|
||||
SessionResult
|
||||
status : "completed" | "failed"
|
||||
reason : str ("goal_met" | "step_budget" | "replan_budget" | "no_plan")
|
||||
history : list[Step]
|
||||
revisions : list[PlanDiff]
|
||||
events : list[Event]
|
||||
```
|
||||
|
||||
The harness loop from lesson twenty can read this directly. The dispatcher from lesson twenty-three is what executes each step. The registry from lesson twenty-one validates each step's args. The transport from lesson twenty-two would surface this whole flow over JSON-RPC to a model client.
|
||||
|
||||
## How to read the code
|
||||
|
||||
`code/main.py` defines `PlanExecuteAgent`, `Step`, `PlanDiff`, `SessionResult`, and the deterministic planner. The executor is a single `run(goal)` method that returns a `SessionResult`. The plan diff is computed by comparing step ids and `(tool_name, args)` tuples.
|
||||
|
||||
`code/tests/test_agent.py` covers a linear success, a mid-plan failure that replans once, replan exhaustion that returns `failed:replan_budget`, step-budget exhaustion, and the plan-diff event format.
|
||||
|
||||
## Going further
|
||||
|
||||
Two extensions you will want once you wire this to a real model. First, partial-plan caching: when a plan succeeds for the first three of six steps and then fails, you do not want to re-run the first three. The executor already keeps history; the planner just needs to read it. Second, parallel branches: the current executor is strictly sequential. A planner that emits an independent branch (`gather_step` instead of `next_step`) can run two tool calls concurrently through the dispatcher.
|
||||
|
||||
Both add real complexity. Both are easier to add once the linear executor is pinned. That is what this lesson does.
|
||||
@@ -0,0 +1,90 @@
|
||||
{
|
||||
"lesson": "24-plan-execute-control-flow",
|
||||
"title": "Plan-Execute Control Flow",
|
||||
"questions": [
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "What is the planner's responsibility in a plan-execute agent?",
|
||||
"options": [
|
||||
"To call the dispatcher on each step",
|
||||
"To produce an ordered list of typed steps the executor will run",
|
||||
"To enforce timeouts",
|
||||
"To handle the JSON-RPC transport"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Planner emits structured plans. The executor runs them. Mixing the two is the chain-of-thought trap this lesson is moving away from."
|
||||
},
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Why is expected_outcome part of the Step shape even though the executor does not check it?",
|
||||
"options": [
|
||||
"Because the dispatcher uses it as a cache key",
|
||||
"Because it gives the replanner and tracers a stated success condition to read",
|
||||
"Because the registry requires it",
|
||||
"Because the model trains on it"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "The replanner reads it when revising. The tracer renders it on the timeline. It is a human-readable contract on the step."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "When a step fails and the replanner returns a new plan, what does the executor emit on the event stream?",
|
||||
"options": [
|
||||
"Nothing — replan is silent",
|
||||
"session.complete with status=failed",
|
||||
"plan.diff with removed/added/revised step ids",
|
||||
"tool.error followed by session.start"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "plan.diff makes the revision visible. A silent rewrite is the bug we are avoiding."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Why have both max_steps and max_replans as separate budgets?",
|
||||
"options": [
|
||||
"Because they measure different units. Steps cap total execution; replans cap how many times the planner is called.",
|
||||
"Because asyncio requires both",
|
||||
"Because the dispatcher rejects unbounded plans",
|
||||
"Because the registry counts replans"
|
||||
],
|
||||
"correct": 0,
|
||||
"explanation": "max_replans catches a planner that keeps returning the same broken plan. max_steps catches a runaway plan that keeps growing."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "Which reason is returned when the planner returns an empty list on the first call?",
|
||||
"options": [
|
||||
"step_budget",
|
||||
"replan_budget",
|
||||
"goal_met",
|
||||
"no_plan"
|
||||
],
|
||||
"correct": 3,
|
||||
"explanation": "No plan, no execution. The session ends with reason=no_plan and status=failed."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "What is in PlanDiff.revised?",
|
||||
"options": [
|
||||
"Step ids that were in the old plan and are not in the new",
|
||||
"Step ids in the new plan that were not in the old",
|
||||
"Step ids whose tool_name or args changed between revisions",
|
||||
"Step ids the planner asked to skip"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "revised is the set of ids whose (tool_name, args) signature changed across the revision."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "When the executor exhausts max_replans, what is the resulting SessionResult.reason?",
|
||||
"options": [
|
||||
"goal_met",
|
||||
"step_budget",
|
||||
"replan_budget",
|
||||
"no_plan"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "replan_budget. The status is failed. History carries every step that ran, including the failures that triggered replans."
|
||||
}
|
||||
]
|
||||
}
|
||||
Reference in New Issue
Block a user