diff --git a/src/beaver_gateway/agents/claude.py b/src/beaver_gateway/agents/claude.py index 5302c82..c6388dd 100644 --- a/src/beaver_gateway/agents/claude.py +++ b/src/beaver_gateway/agents/claude.py @@ -1,26 +1,24 @@ """Claude agent definition, backed by the Claude Agent SDK. -The system prompt is either ``system_prompt`` verbatim or, when -``prompt_sources`` is set, the concatenation of those files (or -``(tag, file)`` pairs) assembled at every session spawn (see -``core/prompt.py``). ``skill_sets`` are -directories of ``/SKILL.md`` folders; each becomes a local SDK -plugin. Nothing from disk is loaded otherwise: the adapter runs with -``setting_sources=[]``. +The system prompt is ``system_prompt`` verbatim or, per conversation kind, +the granules named in ``prompts`` assembled at every session spawn (see +``core/prompt.py``). ``skill_sets`` are directories of ``/SKILL.md`` +folders; each becomes a local SDK plugin. Nothing from disk is loaded +otherwise: the adapter runs with ``setting_sources=[]``. """ from __future__ import annotations from collections.abc import Mapping # noqa: TC003 - pydantic runtime from pathlib import Path # noqa: TC003 - pydantic runtime -from typing import Any from pydantic import BaseModel, ConfigDict, Field, model_validator from beaver_gateway.agents.base import BaseAgent +from beaver_gateway.core.kinds import KINDS, Kind from beaver_gateway.core.prompt import PromptSource # noqa: TC001 - pydantic runtime -__all__ = ["ClaudeAgent", "ClaudeOptions"] +__all__ = ["ClaudeAgent", "ClaudeOptions", "Prompts"] class ClaudeOptions(BaseModel): @@ -48,19 +46,39 @@ class ClaudeOptions(BaseModel): killed mid-turn loses at most the frame in flight.""" +class Prompts(BaseModel): + """Prompt assembly per conversation kind (§3.12): the granules, in order. + + A kind left ``None`` is not served by the agent; ``ClaudeAgent.kinds`` + follows from the kinds set here. + """ + + model_config = ConfigDict(frozen=True) + + master: tuple[PromptSource, ...] | None = None + branch: tuple[PromptSource, ...] | None = None + deep: tuple[PromptSource, ...] | None = None + job: tuple[PromptSource, ...] | None = None + fork: tuple[PromptSource, ...] | None = None + + def for_kind(self, kind: Kind) -> tuple[PromptSource, ...] | None: + return getattr(self, kind) + + @property + def kinds(self) -> tuple[Kind, ...]: + return tuple(k for k in KINDS if self.for_kind(k) is not None) + + class ClaudeAgent(BaseAgent): cwd: Path system_prompt: str = "" - prompt_sources: tuple[PromptSource, ...] = () - prompt_sources_by_kind: Mapping[str, tuple[PromptSource, ...]] = Field( - default_factory=dict - ) - """Per conversation kind (``master``/``branch``/``deep``/``job``/``fork``) - assembly; falls back to ``prompt_sources``. Constant per kind (§3.12).""" + """Verbatim prompt for agents without ``prompts`` (tests, one-offs).""" - kinds: tuple[str, ...] = () + prompts: Prompts = Field(default_factory=Prompts) + kinds: tuple[Kind, ...] = () """Conversation kinds this agent serves; ``create``/``spawn`` reject the - rest. Defaults to the keys of ``prompt_sources_by_kind`` or ``("deep",)``.""" + rest. Defaults to the kinds ``prompts`` covers, or ``("deep",)`` for a + verbatim ``system_prompt``.""" skill_sets: tuple[Path, ...] = () gateway_tools: tuple[str, ...] = () @@ -69,16 +87,20 @@ class ClaudeAgent(BaseAgent): options: ClaudeOptions = Field(default_factory=ClaudeOptions) - @model_validator(mode="before") - @classmethod - def _default_kinds(cls, data: Any) -> Any: - if isinstance(data, dict) and not data.get("kinds"): - by_kind = data.get("prompt_sources_by_kind") or {} - data = {**data, "kinds": tuple(by_kind) or ("deep",)} - return data + @model_validator(mode="after") + def _kinds_follow_prompts(self) -> ClaudeAgent: + covered = self.prompts.kinds + if not self.kinds: + object.__setattr__(self, "kinds", covered or ("deep",)) + elif covered: + missing = [k for k in self.kinds if k not in covered] + if missing: + msg = f"agent {self.name!r} serves {missing} without a prompt" + raise ValueError(msg) + return self - def prompt_for(self, kind: str) -> tuple[PromptSource, ...]: - return self.prompt_sources_by_kind.get(kind, self.prompt_sources) + def prompt_for(self, kind: Kind) -> tuple[PromptSource, ...] | None: + return self.prompts.for_kind(kind) def serves(self, kind: str) -> bool: return kind in self.kinds diff --git a/src/beaver_gateway/backends/claude_sdk.py b/src/beaver_gateway/backends/claude_sdk.py index a481d4f..e3bcfd5 100644 --- a/src/beaver_gateway/backends/claude_sdk.py +++ b/src/beaver_gateway/backends/claude_sdk.py @@ -77,6 +77,7 @@ from beaver_gateway.core.events import ( build_thinking_delta, build_tool_use_block_start, ) +from beaver_gateway.core.kinds import as_kind from beaver_gateway.core.sessions import Session, SessionClient, SessionPool from beaver_gateway.core.transcript import ( build_entries, @@ -511,7 +512,7 @@ class ClaudeSdkBackend: env["HOME"] = str(self._runner.home) env.setdefault("CLAUDE_CONFIG_DIR", str(self._runner.home / ".claude")) plugins = self._plugins() - sources = agent.prompt_for(spec.kind) + sources = agent.prompt_for(as_kind(spec.kind)) system_prompt = ( prompt_assembly.assemble(sources) if sources else agent.system_prompt ) diff --git a/src/beaver_gateway/core/conversations.py b/src/beaver_gateway/core/conversations.py index 89ed2f5..0ecf0dd 100644 --- a/src/beaver_gateway/core/conversations.py +++ b/src/beaver_gateway/core/conversations.py @@ -36,6 +36,7 @@ from claude_agent_sdk import ( from sqlmodel import col, select from beaver_gateway.core.injects import InjectQueue, inject_header +from beaver_gateway.core.kinds import KINDS, Kind from beaver_gateway.core.transcript import ( messages_from_entries, render_messages, @@ -77,7 +78,6 @@ __all__ = [ _log = logging.getLogger("beaver_gateway.core.conversations") -KINDS = ("master", "branch", "deep", "job", "fork") SEEDS = ("clean", "morning", "copy", "brief") _STATUSES = ("open", "merged", "closed", "archived") _DEFAULT_MERGE_PROMPT = ( @@ -91,7 +91,7 @@ _UNITS = {"s": 1, "m": 60, "h": 3600, "d": 86400} @dataclass(frozen=True, slots=True) class SeedContext: - kind: str + kind: Kind seed: str agent: str parent: Conversation | None @@ -177,7 +177,7 @@ class Conversations: async def create( self, *, - kind: str, + kind: Kind, agent: str, parent: Conversation | None = None, title: str | None = None, @@ -406,7 +406,7 @@ class Conversations: msg = f"unknown frontend {name!r}" raise ValueError(msg) - def default_agent(self, kind: str) -> str | None: + def default_agent(self, kind: Kind) -> str | None: for fe in self._frontends: if kind in fe.kinds and (agent := fe.agent_for(kind)): return agent @@ -426,7 +426,7 @@ class Conversations: async def spawn( self, *, - kind: str, + kind: Kind, agent: str | None = None, seed: str = "clean", parent: Conversation | None = None, diff --git a/src/beaver_gateway/core/gateway_tools.py b/src/beaver_gateway/core/gateway_tools.py index f83286d..382e58e 100644 --- a/src/beaver_gateway/core/gateway_tools.py +++ b/src/beaver_gateway/core/gateway_tools.py @@ -12,6 +12,8 @@ from typing import TYPE_CHECKING, Any, cast from claude_agent_sdk import create_sdk_mcp_server, tool +from beaver_gateway.core.kinds import as_kind + if TYPE_CHECKING: from collections.abc import Iterable @@ -99,7 +101,7 @@ def _tools(conversations: Conversations, key: str) -> list[SdkMcpTool[Any]]: parent = await current() try: child = await conversations.spawn( - kind=str(args["kind"]), + kind=as_kind(str(args["kind"])), agent=args.get("agent"), seed=str(args.get("seed") or "clean"), parent=parent, diff --git a/src/beaver_gateway/core/kinds.py b/src/beaver_gateway/core/kinds.py new file mode 100644 index 0000000..e5f5944 --- /dev/null +++ b/src/beaver_gateway/core/kinds.py @@ -0,0 +1,18 @@ +"""Conversation kinds (§3.1) as one closed type for agents, frontends, service.""" + +from __future__ import annotations + +from typing import Literal, get_args + +__all__ = ["KINDS", "Kind", "as_kind"] + +Kind = Literal["master", "branch", "deep", "job", "fork"] +KINDS: tuple[Kind, ...] = get_args(Kind) + + +def as_kind(value: str) -> Kind: + for kind in KINDS: + if value == kind: + return kind + msg = f"unknown conversation kind {value!r}" + raise ValueError(msg) diff --git a/src/beaver_gateway/frontends/api/frontend.py b/src/beaver_gateway/frontends/api/frontend.py index f1fe446..30fa872 100644 --- a/src/beaver_gateway/frontends/api/frontend.py +++ b/src/beaver_gateway/frontends/api/frontend.py @@ -22,7 +22,8 @@ from sqlalchemy import select as sa_select from sqlmodel import col from beaver_gateway.core import audit -from beaver_gateway.core.conversations import KINDS, SEEDS +from beaver_gateway.core.conversations import SEEDS +from beaver_gateway.core.kinds import Kind, as_kind from beaver_gateway.frontends._auth import require_token from beaver_gateway.frontends._sse import ( KEEPALIVE, @@ -34,7 +35,7 @@ from beaver_gateway.frontends.base import Frontend from beaver_gateway.storage.models import Usage if TYPE_CHECKING: - from collections.abc import AsyncIterator, Mapping + from collections.abc import AsyncIterator from beaver_gateway.core.conversations import Conversations from beaver_gateway.frontends.base import GatewayRuntime @@ -57,16 +58,27 @@ class ApiFrontend(Frontend): host: str = "0.0.0.0", # noqa: S104 port: int = 8004, public_base_url: str | None = None, - default_agents: Mapping[str, str] | None = None, + master_agent: str | None = None, + branch_agent: str | None = None, + deep_agent: str | None = None, + job_agent: str | None = None, ) -> None: self.host = host self.port = port self.public_base_url = public_base_url.rstrip("/") if public_base_url else None - self.default_agents = dict(default_agents or {}) + self.master_agent = master_agent + self.branch_agent = branch_agent + self.deep_agent = deep_agent + self.job_agent = job_agent self._app: FastAPI | None = None - def agent_for(self, kind: str) -> str | None: - return self.default_agents.get(kind) + def agent_for(self, kind: Kind) -> str | None: + return { + "master": self.master_agent, + "branch": self.branch_agent, + "deep": self.deep_agent, + "job": self.job_agent, + }.get(kind) def configure(self, runtime: GatewayRuntime) -> None: if runtime.conversations is None or runtime.bus is None: @@ -175,12 +187,13 @@ def _build_app(runtime: GatewayRuntime) -> FastAPI: # noqa: PLR0915 async def create_conversation(request: Request) -> dict[str, Any]: token = await require_token(request, runtime, scope=SCOPE) data = await body_of(request) - kind = str(data.get("kind") or "deep") agent = data.get("agent") - if kind not in KINDS or kind == "fork": - raise HTTPException( - status.HTTP_400_BAD_REQUEST, f"kind must be one of {KINDS[:-1]}" - ) + try: + kind = as_kind(str(data.get("kind") or "deep")) + except ValueError as exc: + raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc + if kind == "fork": + raise HTTPException(status.HTTP_400_BAD_REQUEST, "fork is internal") if agent is not None and ( not isinstance(agent, str) or agent not in runtime.agents ): diff --git a/src/beaver_gateway/frontends/base.py b/src/beaver_gateway/frontends/base.py index 04fb39c..e1ded4b 100644 --- a/src/beaver_gateway/frontends/base.py +++ b/src/beaver_gateway/frontends/base.py @@ -19,6 +19,7 @@ if TYPE_CHECKING: from beaver_gateway.backends.base import Backend from beaver_gateway.core.auth import TokenStore + from beaver_gateway.core.kinds import Kind from beaver_gateway.core.registry import AgentRegistry, McpRegistry from beaver_gateway.core.turn_record import TurnRecord from beaver_gateway.storage import Database @@ -101,7 +102,7 @@ class Frontend(ABC): """ name: str = "" - kinds: tuple[str, ...] = () + kinds: tuple[Kind, ...] = () @abstractmethod def configure(self, runtime: GatewayRuntime) -> None: ... @@ -109,7 +110,7 @@ class Frontend(ABC): @abstractmethod async def serve(self) -> None: ... - def agent_for(self, kind: str) -> str | None: # noqa: ARG002 + def agent_for(self, kind: Kind) -> str | None: # noqa: ARG002 return None async def materialize(self, conv: Conversation) -> ConversationBinding | None: # noqa: ARG002 diff --git a/src/beaver_gateway/frontends/markdown/frontend.py b/src/beaver_gateway/frontends/markdown/frontend.py index db05558..6c42f17 100644 --- a/src/beaver_gateway/frontends/markdown/frontend.py +++ b/src/beaver_gateway/frontends/markdown/frontend.py @@ -76,6 +76,7 @@ from beaver_gateway.frontends.markdown.mirror import FRONTEND, ChatMirror if TYPE_CHECKING: from collections.abc import AsyncIterator, Callable + from beaver_gateway.core.kinds import Kind from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.storage.models import Conversation, ConversationBinding @@ -173,7 +174,7 @@ class MarkdownFrontend(Frontend): raise RuntimeError(msg) return self._mirror - def agent_for(self, kind: str) -> str | None: + def agent_for(self, kind: Kind) -> str | None: return self.default_agent if kind == "deep" else None async def materialize(self, conv: Conversation) -> ConversationBinding | None: diff --git a/tests/test_claude_sdk_backend.py b/tests/test_claude_sdk_backend.py index 3e498c6..1e91f66 100644 --- a/tests/test_claude_sdk_backend.py +++ b/tests/test_claude_sdk_backend.py @@ -24,7 +24,7 @@ from claude_agent_sdk import ( ) from beaver_gateway.agents.base import ExposedMcp -from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions +from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions, Prompts from beaver_gateway.backends.claude_sdk import ( ClaudeSdkBackend, RunnerConfig, @@ -369,7 +369,7 @@ async def test_prompt_sources_are_assembled(cwd: Path) -> None: backend = _backend( cwd, InMemorySessionStore(), - prompt_sources=(("role", cwd / "a.md"), cwd / "b.md"), + prompts=Prompts(deep=(("role", cwd / "a.md"), cwd / "b.md")), ) await _drain( backend.complete( diff --git a/tests/test_conversations.py b/tests/test_conversations.py index a3b8d6f..3842db5 100644 --- a/tests/test_conversations.py +++ b/tests/test_conversations.py @@ -634,3 +634,24 @@ async def test_spawn_defaults_agent_and_materializes(world: World) -> None: with pytest.raises(ValueError, match="no default agent for kind 'job'"): await world.conversations.spawn(kind="job", seed="clean") await world.settle(deep, 1) + + +def test_agent_kinds_follow_prompts(tmp_path: Path) -> None: + from beaver_gateway.agents.claude import Prompts + + plain = ClaudeAgent(name="p", model="m", system_prompt="hi", cwd=tmp_path) + assert plain.kinds == ("deep",) + dispatcher = ClaudeAgent( + name="x", model="m", cwd=tmp_path, prompts=Prompts(master=("a.md",), fork=()) + ) + assert dispatcher.kinds == ("master", "fork") + assert dispatcher.prompt_for("master") == ("a.md",) + assert dispatcher.prompt_for("deep") is None + with pytest.raises(ValueError, match=r"serves \['deep'\] without a prompt"): + ClaudeAgent( + name="y", + model="m", + cwd=tmp_path, + kinds=("master", "deep"), + prompts=Prompts(master=()), + ) diff --git a/tests/test_routing.py b/tests/test_routing.py index 1cc867c..2c1fe43 100644 --- a/tests/test_routing.py +++ b/tests/test_routing.py @@ -23,7 +23,7 @@ class Stack: def __init__(self, world: World) -> None: self.world = world self.vault = world.root / "vault" - self.api = ApiFrontend(default_agents={"master": "a"}) + self.api = ApiFrontend(master_agent="a") self.markdown = MarkdownFrontend(vault_path=self.vault, default_agent="d") self.anthropic = AnthropicMessagesFrontend() frontends = [self.api, self.anthropic, self.markdown]