feat(core,backends,frontends,storage): conversations, inject queue, session pool, gateway mcp tools, api frontend

This commit is contained in:
hh
2026-08-28 03:08:30 +02:00
parent 50b7057fa4
commit 33ccc78fec
28 changed files with 3543 additions and 345 deletions
+1
View File
@@ -19,6 +19,7 @@ dependencies = [
"greenlet>=3.5.0", "greenlet>=3.5.0",
"itsdangerous>=2.2.0", "itsdangerous>=2.2.0",
"jinja2>=3.1.6", "jinja2>=3.1.6",
"psutil>=7.2.2",
"psycopg[binary]>=3.3.4", "psycopg[binary]>=3.3.4",
"pydantic>=2.13.4", "pydantic>=2.13.4",
"pydantic-settings>=2.14.1", "pydantic-settings>=2.14.1",
+16 -4
View File
@@ -42,15 +42,27 @@ class ClaudeOptions(BaseModel):
"""Extra inherited variable names to let through the env whitelist.""" """Extra inherited variable names to let through the env whitelist."""
max_turns: int | None = None max_turns: int | None = None
idle_session_ttl: float = 1800.0 session_store_flush: str = "eager"
"""Seconds a live session may sit unused before it is closed.""" """``eager`` mirrors every transcript frame as it lands, so a gateway
killed mid-turn loses at most the frame in flight."""
session_store_flush: str = "batched"
class ClaudeAgent(BaseAgent): class ClaudeAgent(BaseAgent):
cwd: Path cwd: Path
system_prompt: str = "" system_prompt: str = ""
prompt_sources: tuple[PromptSource, ...] = () 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)."""
skill_sets: tuple[Path, ...] = () skill_sets: tuple[Path, ...] = ()
gateway_tools: tuple[str, ...] = ()
"""Gateway tools exposed in-process (``read_conversation``, ``spawn``,
``say``, ``schedule``, ``inject``); empty = no gateway MCP server."""
options: ClaudeOptions = Field(default_factory=ClaudeOptions) options: ClaudeOptions = Field(default_factory=ClaudeOptions)
def prompt_for(self, kind: str) -> tuple[PromptSource, ...]:
return self.prompt_sources_by_kind.get(kind, self.prompt_sources)
+193 -108
View File
@@ -2,11 +2,19 @@
One :class:`ClaudeSdkBackend` per :class:`ClaudeAgent`. A live session is One :class:`ClaudeSdkBackend` per :class:`ClaudeAgent`. A live session is
one ``ClaudeSDKClient`` (one claude subprocess) and runs one turn at a one ``ClaudeSDKClient`` (one claude subprocess) and runs one turn at a
time. Sessions are keyed by ``conversation_id`` when the frontend passes time; the sessions of every agent live in one shared
one (markdown chats) or by a text-only fingerprint of ``messages[:-1]`` :class:`~beaver_gateway.core.sessions.SessionPool` that owns TTL and
memory-pressure eviction. Sessions are keyed by ``conversation_id`` when
the caller passes one or by a text-only fingerprint of ``messages[:-1]``
for stateless callers (``/v1/messages``). Without a live session the for stateless callers (``/v1/messages``). Without a live session the
adapter resumes ``session_id`` from the session store, or seeds the adapter resumes ``session_id`` from the session store (after closing any
incoming history into the store via ``core/transcript`` and resumes that. ``tool_use`` left open by a crash), or seeds the incoming history into the
store via ``core/transcript`` and resumes that.
Per-turn ``**options`` beyond the protocol's: ``kind`` (conversation kind,
picks the prompt assembly and the pool TTL), ``pinned`` (never evicted),
``tools=False`` (no MCP at all - forks and jobs), ``observer`` (callback
receiving every raw SDK message, subagent ones included), ``turn_id``.
Events on the wire are the Anthropic ``MessageStreamEvent`` family: one Events on the wire are the Anthropic ``MessageStreamEvent`` family: one
``message_start``/``message_stop`` envelope per turn, block indices ``message_start``/``message_stop`` envelope per turn, block indices
@@ -36,13 +44,14 @@ import uuid
from collections.abc import Mapping from collections.abc import Mapping
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, Protocol, Self, cast from typing import TYPE_CHECKING, Any, Self, cast
import claude_agent_sdk import claude_agent_sdk
from claude_agent_sdk import ( from claude_agent_sdk import (
AssistantMessage, AssistantMessage,
ClaudeAgentOptions, ClaudeAgentOptions,
ClaudeSDKClient, ClaudeSDKClient,
MirrorErrorMessage,
ResultMessage, ResultMessage,
StreamEvent, StreamEvent,
TextBlock, TextBlock,
@@ -68,14 +77,15 @@ from beaver_gateway.core.events import (
build_thinking_delta, build_thinking_delta,
build_tool_use_block_start, build_tool_use_block_start,
) )
from beaver_gateway.core.transcript import build_entries from beaver_gateway.core.sessions import Session, SessionClient, SessionPool
from beaver_gateway.core.transcript import build_entries, close_open_tool_uses
from beaver_gateway.core.turn_capture import TurnCapture, TurnUsage from beaver_gateway.core.turn_capture import TurnCapture, TurnUsage
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator, Awaitable, Callable, Iterable, Sequence from collections.abc import AsyncIterator, Awaitable, Callable, Iterable, Sequence
from anthropic.types import MessageParam from anthropic.types import MessageParam
from claude_agent_sdk import SessionStore from claude_agent_sdk import McpSdkServerConfig, SessionStore
from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.agents.claude import ClaudeAgent from beaver_gateway.agents.claude import ClaudeAgent
@@ -88,6 +98,7 @@ __all__ = [
"ClaudeSdkBackend", "ClaudeSdkBackend",
"RunnerConfig", "RunnerConfig",
"SessionClient", "SessionClient",
"ToolServerFactory",
"UsageSink", "UsageSink",
"fingerprint", "fingerprint",
] ]
@@ -120,7 +131,6 @@ ENV_KEEP: tuple[str, ...] = (
) )
ENV_KEEP_PREFIXES: tuple[str, ...] = ("CLAUDE_", "ANTHROPIC_", "DISABLE_") ENV_KEEP_PREFIXES: tuple[str, ...] = ("CLAUDE_", "ANTHROPIC_", "DISABLE_")
_REAP_INTERVAL = 60.0
_STOP_REASONS: dict[str, StopReason] = { _STOP_REASONS: dict[str, StopReason] = {
"end_turn": "end_turn", "end_turn": "end_turn",
"tool_use": "tool_use", "tool_use": "tool_use",
@@ -131,15 +141,10 @@ _STOP_REASONS: dict[str, StopReason] = {
} }
class SessionClient(Protocol):
async def connect(self) -> None: ...
async def query(self, prompt: str) -> None: ...
def receive_response(self) -> AsyncIterator[Any]: ...
async def disconnect(self) -> None: ...
ClientFactory = "Callable[[ClaudeAgentOptions], SessionClient]" ClientFactory = "Callable[[ClaudeAgentOptions], SessionClient]"
UsageSink = "Callable[[UsageEvent], Awaitable[None]]" UsageSink = "Callable[[UsageEvent], Awaitable[None]]"
ToolServerFactory = "Callable[[str, str], McpSdkServerConfig | None]"
"""``(conversation_key, kind) -> in-process MCP server config`` or ``None``."""
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -158,16 +163,6 @@ class UsageEvent:
usage: TurnUsage usage: TurnUsage
@dataclass
class _Live:
client: SessionClient
session_id: str | None
resumed: bool
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
last_used: float = field(default_factory=time.monotonic)
turns: int = 0
class _RunnerClient(ClaudeSDKClient): class _RunnerClient(ClaudeSDKClient):
"""``ClaudeSDKClient`` that hands the materialized resume dir to the runner uid.""" """``ClaudeSDKClient`` that hands the materialized resume dir to the runner uid."""
@@ -194,6 +189,8 @@ class ClaudeSdkBackend:
usage_sink: Callable[[UsageEvent], Awaitable[None]] | None = None, usage_sink: Callable[[UsageEvent], Awaitable[None]] | None = None,
client_factory: Callable[[ClaudeAgentOptions], SessionClient] | None = None, client_factory: Callable[[ClaudeAgentOptions], SessionClient] | None = None,
work_dir: Path | None = None, work_dir: Path | None = None,
pool: SessionPool | None = None,
tool_server: Callable[[str, str], McpSdkServerConfig | None] | None = None,
) -> None: ) -> None:
self._agent = agent self._agent = agent
self._store = session_store self._store = session_store
@@ -203,8 +200,8 @@ class ClaudeSdkBackend:
self._work_dir = work_dir or Path(tempfile.gettempdir()) / "beaver-claude" self._work_dir = work_dir or Path(tempfile.gettempdir()) / "beaver-claude"
self._servers = _mcp_servers(agent, mcp_internal_urls) self._servers = _mcp_servers(agent, mcp_internal_urls)
self._mcp_disallowed = _mcp_disallowed(agent, mcp_tool_names or {}) self._mcp_disallowed = _mcp_disallowed(agent, mcp_tool_names or {})
self._sessions: dict[str, _Live] = {} self._pool = pool if pool is not None else SessionPool()
self._reaper: asyncio.Task[None] | None = None self._tool_server = tool_server
self._uid, self._gid = _resolve_ids(self._runner.user) self._uid, self._gid = _resolve_ids(self._runner.user)
self._wrapper: Path | None = None self._wrapper: Path | None = None
@@ -212,34 +209,56 @@ class ClaudeSdkBackend:
def agent(self) -> ClaudeAgent: def agent(self) -> ClaudeAgent:
return self._agent return self._agent
@property
def pool(self) -> SessionPool:
return self._pool
@property @property
def sessions(self) -> dict[str, dict[str, Any]]: def sessions(self) -> dict[str, dict[str, Any]]:
now = time.monotonic()
return { return {
key: { row["key"]: row
"session_id": live.session_id, for row in self._pool.snapshot()
"idle_seconds": now - live.last_used, if row["agent"] == self._agent.name
"turns": live.turns,
"busy": live.lock.locked(),
}
for key, live in self._sessions.items()
} }
async def __aenter__(self) -> Self: async def __aenter__(self) -> Self:
self._reaper = asyncio.create_task(self._reap_loop())
return self return self
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None: async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
await self.aclose() await self.aclose()
async def aclose(self) -> None: async def aclose(self) -> None:
if self._reaper is not None: await self._pool.close_all(agent=self._agent.name)
self._reaper.cancel()
with contextlib.suppress(BaseException): async def close(self, key: str) -> None:
await self._reaper await self._pool.close(key)
self._reaper = None
for key in list(self._sessions): async def interrupt(self, key: str) -> bool:
await self._close(key) live = self._pool.get(key)
if live is None or not live.busy:
return False
live.interrupt_requested = True
await live.client.interrupt()
return True
def live(self, key: str) -> Session | None:
return self._pool.get(key)
async def repair_session(self, session_id: str) -> int:
"""Close ``tool_use`` blocks a crash left without a result; count added."""
key = self._store_key(session_id)
entries = await self._store.load(cast("Any", key))
if not entries:
return 0
fixes = close_open_tool_uses(cast("list[Mapping[str, Any]]", entries))
if fixes:
await self._store.append(cast("Any", key), cast("Any", fixes))
_log.warning(
"session %s: closed %d open tool_use with synthetic results",
session_id,
len(fixes),
)
return len(fixes)
async def complete( async def complete(
self, self,
@@ -250,6 +269,11 @@ class ClaudeSdkBackend:
conversation_id: str | None = None, conversation_id: str | None = None,
session_id: str | None = None, session_id: str | None = None,
capture: TurnCapture | None = None, capture: TurnCapture | None = None,
kind: str = "deep",
pinned: bool = False,
tools: bool = True,
observer: Callable[[Any], None] | None = None,
turn_id: str | None = None,
**options: Any, # noqa: ARG002 - per-request knobs are not supported **options: Any, # noqa: ARG002 - per-request knobs are not supported
) -> AsyncIterator[MessageStreamEvent]: ) -> AsyncIterator[MessageStreamEvent]:
if agent.name != self._agent.name: if agent.name != self._agent.name:
@@ -262,37 +286,70 @@ class ClaudeSdkBackend:
prompt = _prompt_text(history[-1].get("content")) prompt = _prompt_text(history[-1].get("content"))
prior = history[:-1] prior = history[:-1]
key = conversation_id or fingerprint(prior) key = conversation_id or fingerprint(prior)
live = await self._acquire(key, session_id=session_id, history=prior) spec = _SessionSpec(kind=kind, pinned=pinned, tools=tools)
live = await self._acquire(key, session_id=session_id, history=prior, spec=spec)
message_id = f"msg_{uuid.uuid4().hex}" message_id = f"msg_{uuid.uuid4().hex}"
yield build_message_start(message_id=message_id, model=self._agent.model) yield build_message_start(message_id=message_id, model=self._agent.model)
async with live.lock: async with live.lock:
live.running_turn = turn_id or message_id
live.last_used = time.monotonic() live.last_used = time.monotonic()
try: try:
turn = await self._run_turn(live, prompt) turn = await self._run_turn(live, prompt, observer=observer)
except Exception: except Exception:
if live.resumed and live.turns == 0: live.running_turn = None
_log.exception( await self._pool.close(key)
"resume of %s failed, reseeding from history", live.session_id if not (live.resumed and live.turns == 0):
)
await self._close(key)
live = await self._acquire(key, session_id=None, history=prior)
async with live.lock:
turn = await self._run_turn(live, prompt)
else:
raise raise
_log.exception(
"resume of %s failed, reseeding from history", live.session_id
)
live = await self._acquire(
key, session_id=None, history=prior, spec=spec
)
async with live.lock:
live.running_turn = turn_id or message_id
turn = await self._run_turn(live, prompt, observer=observer)
for event in turn.events: for event in turn.events:
yield event yield event
live.turns += 1 live.turns += 1
live.last_used = time.monotonic() live.last_used = time.monotonic()
live.running_turn = None
usage, interrupted = await self._after_turn(
live,
turn,
conversation_id=conversation_id,
history=history,
capture=capture,
)
if turn.result is not None and turn.result.is_error and not interrupted:
msg = f"claude: {turn.result.result or turn.result.subtype}"
raise RuntimeError(msg)
yield build_message_delta(
stop_reason=turn.stop_reason, usage=_wire_usage(usage)
)
yield build_message_stop()
async def _after_turn(
self,
live: Session,
turn: _Turn,
*,
conversation_id: str | None,
history: list[dict[str, Any]],
capture: TurnCapture | None,
) -> tuple[TurnUsage, bool]:
if turn.result is not None and turn.result.session_id: if turn.result is not None and turn.result.session_id:
live.session_id = turn.result.session_id live.session_id = turn.result.session_id
if conversation_id is None: if conversation_id is None:
self._rekey(key, fingerprint([*history, *turn.synthesized])) self._rekey(live.key, fingerprint([*history, *turn.synthesized]))
usage = _usage_of(turn.result) usage = _usage_of(turn.result)
interrupted = live.interrupt_requested
live.interrupt_requested = False
if capture is not None: if capture is not None:
capture.synthesized_messages = turn.synthesized capture.synthesized_messages = turn.synthesized
capture.session_id = live.session_id capture.session_id = live.session_id
capture.usage = usage capture.usage = usage
capture.interrupted = interrupted
if self._usage_sink is not None: if self._usage_sink is not None:
await self._usage_sink( await self._usage_sink(
UsageEvent( UsageEvent(
@@ -304,15 +361,15 @@ class ClaudeSdkBackend:
usage=usage, usage=usage,
) )
) )
if turn.result is not None and turn.result.is_error: return usage, interrupted
msg = f"claude: {turn.result.result or turn.result.subtype}"
raise RuntimeError(msg)
yield build_message_delta(
stop_reason=turn.stop_reason, usage=_wire_usage(usage)
)
yield build_message_stop()
async def _run_turn(self, live: _Live, prompt: str) -> _Turn: async def _run_turn(
self,
live: Session,
prompt: str,
*,
observer: Callable[[Any], None] | None = None,
) -> _Turn:
streaming = self._agent.options.include_partial_messages streaming = self._agent.options.include_partial_messages
turn = _Turn() turn = _Turn()
raw: list[Any] = [] raw: list[Any] = []
@@ -320,6 +377,16 @@ class ClaudeSdkBackend:
offset = 0 offset = 0
await live.client.query(prompt) await live.client.query(prompt)
async for message in live.client.receive_response(): async for message in live.client.receive_response():
if observer is not None:
observer(message)
if isinstance(message, MirrorErrorMessage):
live.dirty = True
_log.error(
"session %s: mirror error, marked dirty: %s",
live.session_id,
message.error,
)
continue
if getattr(message, "parent_tool_use_id", None) is not None: if getattr(message, "parent_tool_use_id", None) is not None:
continue continue
if isinstance(message, StreamEvent): if isinstance(message, StreamEvent):
@@ -357,17 +424,30 @@ class ClaudeSdkBackend:
return turn return turn
async def _acquire( async def _acquire(
self, key: str, *, session_id: str | None, history: list[dict[str, Any]] self,
) -> _Live: key: str,
live = self._sessions.get(key) *,
session_id: str | None,
history: list[dict[str, Any]],
spec: _SessionSpec,
) -> Session:
live = self._pool.get(key)
if live is not None: if live is not None:
return live return live
resume = session_id resume = session_id
if resume is None and history: if resume is not None:
await self.repair_session(resume)
elif history:
resume = await self._seed(history) resume = await self._seed(history)
live = await self._spawn(resume) await self._pool.make_room()
self._sessions[key] = live live = await self._spawn(resume, key=key, spec=spec)
return live return self._pool.add(live)
def _store_key(self, session_id: str) -> dict[str, str]:
return {
"project_key": project_key_for_directory(str(self._agent.cwd)),
"session_id": session_id,
}
async def _seed(self, history: list[dict[str, Any]]) -> str: async def _seed(self, history: list[dict[str, Any]]) -> str:
session_id = str(uuid.uuid4()) session_id = str(uuid.uuid4())
@@ -378,11 +458,9 @@ class ClaudeSdkBackend:
model=self._agent.model, model=self._agent.model,
permission_mode=self._agent.options.permission_mode, permission_mode=self._agent.options.permission_mode,
) )
key = { await self._store.append(
"project_key": project_key_for_directory(str(self._agent.cwd)), cast("Any", self._store_key(session_id)), cast("Any", entries)
"session_id": session_id, )
}
await self._store.append(cast("Any", key), cast("Any", entries))
_log.info( _log.info(
"seeded session %s with %d entries from %d messages", "seeded session %s with %d entries from %d messages",
session_id, session_id,
@@ -391,22 +469,36 @@ class ClaudeSdkBackend:
) )
return session_id return session_id
async def _spawn(self, resume: str | None) -> _Live: async def _spawn(
options = self._build_options(resume) self, resume: str | None, *, key: str, spec: _SessionSpec
) -> Session:
options = self._build_options(resume, key=key, spec=spec)
client = self._factory(options) client = self._factory(options)
await client.connect() await client.connect()
_log.info( _log.info(
"spawned claude: agent=%s resume=%s user=%s", "spawned claude: agent=%s kind=%s resume=%s tools=%s user=%s",
self._agent.name, self._agent.name,
spec.kind,
resume, resume,
spec.tools,
self._runner.user, self._runner.user,
) )
return _Live(client=client, session_id=resume, resumed=resume is not None) return Session(
key=key,
agent=self._agent.name,
kind=spec.kind,
client=client,
session_id=resume,
resumed=resume is not None,
pinned=spec.pinned,
)
def _default_factory(self, options: ClaudeAgentOptions) -> SessionClient: def _default_factory(self, options: ClaudeAgentOptions) -> SessionClient:
return _RunnerClient(options, uid=self._uid) return _RunnerClient(options, uid=self._uid)
def _build_options(self, resume: str | None) -> ClaudeAgentOptions: def _build_options(
self, resume: str | None, *, key: str, spec: _SessionSpec
) -> ClaudeAgentOptions:
agent = self._agent agent = self._agent
opt = agent.options opt = agent.options
env = dict(opt.env) env = dict(opt.env)
@@ -414,18 +506,25 @@ class ClaudeSdkBackend:
env["HOME"] = str(self._runner.home) env["HOME"] = str(self._runner.home)
env.setdefault("CLAUDE_CONFIG_DIR", str(self._runner.home / ".claude")) env.setdefault("CLAUDE_CONFIG_DIR", str(self._runner.home / ".claude"))
plugins = self._plugins() plugins = self._plugins()
sources = agent.prompt_for(spec.kind)
system_prompt = ( system_prompt = (
prompt_assembly.assemble(agent.prompt_sources) prompt_assembly.assemble(sources) if sources else agent.system_prompt
if agent.prompt_sources
else agent.system_prompt
) )
servers: dict[str, Any] = dict(self._servers) if spec.tools else {}
gateway = (
self._tool_server(key, spec.kind)
if spec.tools and self._tool_server is not None and agent.gateway_tools
else None
)
if gateway is not None:
servers[str(gateway["name"])] = gateway
return ClaudeAgentOptions( return ClaudeAgentOptions(
model=agent.model or None, model=agent.model or None,
effort=cast("Any", opt.effort), effort=cast("Any", opt.effort),
system_prompt=system_prompt, system_prompt=system_prompt,
setting_sources=[], setting_sources=[],
strict_mcp_config=True, strict_mcp_config=True,
mcp_servers=cast("Any", self._servers), mcp_servers=cast("Any", servers),
permission_mode=cast("Any", opt.permission_mode), permission_mode=cast("Any", opt.permission_mode),
tools=list(opt.tools) if opt.tools is not None else None, tools=list(opt.tools) if opt.tools is not None else None,
disallowed_tools=[*opt.disallowed_tools, *self._mcp_disallowed], disallowed_tools=[*opt.disallowed_tools, *self._mcp_disallowed],
@@ -473,8 +572,8 @@ class ClaudeSdkBackend:
target=json.dumps(target), target=json.dumps(target),
keep=json.dumps(keep), keep=json.dumps(keep),
prefixes=json.dumps(list(ENV_KEEP_PREFIXES)), prefixes=json.dumps(list(ENV_KEEP_PREFIXES)),
uid=json.dumps(self._uid), uid=repr(self._uid),
gid=json.dumps(self._gid), gid=repr(self._gid),
) )
digest = hashlib.sha256(script.encode("utf-8")).hexdigest()[:12] digest = hashlib.sha256(script.encode("utf-8")).hexdigest()[:12]
path = self._work_dir / f"claude-exec-{digest}.py" path = self._work_dir / f"claude-exec-{digest}.py"
@@ -487,30 +586,16 @@ class ClaudeSdkBackend:
return path return path
def _rekey(self, old: str, new: str) -> None: def _rekey(self, old: str, new: str) -> None:
live = self._sessions.pop(old, None) stale = self._pool.rekey(old, new)
if live is None: if stale is not None:
return
stale = self._sessions.pop(new, None)
self._sessions[new] = live
if stale is not None and stale is not live:
asyncio.get_running_loop().create_task(_disconnect(stale)) asyncio.get_running_loop().create_task(_disconnect(stale))
async def _close(self, key: str) -> None:
live = self._sessions.pop(key, None)
if live is not None:
await _disconnect(live)
async def _reap_loop(self) -> None: @dataclass(frozen=True, slots=True)
ttl = self._agent.options.idle_session_ttl class _SessionSpec:
while True: kind: str
await asyncio.sleep(_REAP_INTERVAL) pinned: bool
if ttl <= 0: tools: bool
continue
now = time.monotonic()
for key, live in list(self._sessions.items()):
if not live.lock.locked() and now - live.last_used > ttl:
_log.info("closing idle session %s (%s)", live.session_id, key)
await self._close(key)
@dataclass @dataclass
@@ -541,7 +626,7 @@ os.execve(TARGET, [TARGET, *sys.argv[1:]], env)
""" """
async def _disconnect(live: _Live) -> None: async def _disconnect(live: Session) -> None:
try: try:
await live.client.disconnect() await live.client.disconnect()
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
+47 -1
View File
@@ -22,6 +22,7 @@ from __future__ import annotations
import asyncio import asyncio
import contextlib import contextlib
import functools
import logging import logging
import signal import signal
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
@@ -43,13 +44,18 @@ from beaver_gateway.backends.claude_sdk import (
) )
from beaver_gateway.backends.raycast import RaycastBackend from beaver_gateway.backends.raycast import RaycastBackend
from beaver_gateway.core.auth import TokenStore from beaver_gateway.core.auth import TokenStore
from beaver_gateway.core.bus import EventBus
from beaver_gateway.core.conversations import Conversations
from beaver_gateway.core.gateway_tools import build_tool_server
from beaver_gateway.core.registry import AgentRegistry, McpRegistry from beaver_gateway.core.registry import AgentRegistry, McpRegistry
from beaver_gateway.core.sessions import SessionPool
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.mcp.internal_app import build_internal_app from beaver_gateway.mcp.internal_app import build_internal_app
from beaver_gateway.settings import Settings from beaver_gateway.settings import Settings
from beaver_gateway.storage import Database, PostgresSessionStore, Usage, append_usage from beaver_gateway.storage import Database, PostgresSessionStore, Usage, append_usage
if TYPE_CHECKING: if TYPE_CHECKING:
from claude_agent_sdk import McpSdkServerConfig
from fastmcp import FastMCP from fastmcp import FastMCP
from fastmcp.tools.base import Tool as FastMCPTool from fastmcp.tools.base import Tool as FastMCPTool
from starlette.applications import Starlette from starlette.applications import Starlette
@@ -140,15 +146,32 @@ async def _async_main() -> None:
# surfaces at startup instead of mid-conversation. # surfaces at startup instead of mid-conversation.
mcp_tools = await _prefetch_mcp_tools(mcp_servers) mcp_tools = await _prefetch_mcp_tools(mcp_servers)
pool = SessionPool()
bus = EventBus()
late = _LateConversations()
session_store = PostgresSessionStore(db)
backends: dict[str, Backend] = await _build_backends( backends: dict[str, Backend] = await _build_backends(
settings=settings, settings=settings,
agents=agents, agents=agents,
stack=stack, stack=stack,
db=db, db=db,
session_store=session_store,
mcp_internal_urls=internal_urls, mcp_internal_urls=internal_urls,
mcp_servers=mcp_servers, mcp_servers=mcp_servers,
mcp_tools=mcp_tools, mcp_tools=mcp_tools,
pool=pool,
late=late,
) )
conversations = Conversations(
db=db,
agents=agents,
backends=backends,
bus=bus,
pool=pool,
store=session_store,
texts=gateway.texts,
)
late.conversations = conversations
runtime = GatewayRuntime( runtime = GatewayRuntime(
agents=agents, agents=agents,
@@ -161,6 +184,9 @@ async def _async_main() -> None:
admin_pass=settings.admin_pass, admin_pass=settings.admin_pass,
session_secret=settings.session_secret, session_secret=settings.session_secret,
frontends=tuple(gateway.frontends), frontends=tuple(gateway.frontends),
conversations=conversations,
bus=bus,
pool=pool,
) )
for fe in gateway.frontends: for fe in gateway.frontends:
@@ -186,7 +212,10 @@ async def _async_main() -> None:
# it in this path and exit cleanly (Phase 0 DoD). # it in this path and exit cleanly (Phase 0 DoD).
return return
await conversations.start()
stack.push_async_callback(conversations.stop)
async with asyncio.TaskGroup() as tg: async with asyncio.TaskGroup() as tg:
tg.create_task(pool.reap_loop())
if internal_app is not None: if internal_app is not None:
tg.create_task(_serve_internal_mcp(internal_app, settings=settings)) tg.create_task(_serve_internal_mcp(internal_app, settings=settings))
for fe in gateway.frontends: for fe in gateway.frontends:
@@ -252,15 +281,31 @@ async def _serve_internal_mcp(app: Starlette, *, settings: Settings) -> None:
await server.serve() await server.serve()
class _LateConversations:
"""Backends need a tool-server factory before the service that backs it exists."""
conversations: Conversations | None = None
def server(
self, key: str, _kind: str, names: tuple[str, ...]
) -> McpSdkServerConfig | None:
if self.conversations is None or not names:
return None
return build_tool_server(self.conversations, conversation_key=key, names=names)
async def _build_backends( async def _build_backends(
*, *,
settings: Settings, settings: Settings,
agents: AgentRegistry, agents: AgentRegistry,
stack: AsyncExitStack, stack: AsyncExitStack,
db: Database, db: Database,
session_store: PostgresSessionStore,
mcp_internal_urls: dict[str, str], mcp_internal_urls: dict[str, str],
mcp_servers: dict[str, FastMCP], mcp_servers: dict[str, FastMCP],
mcp_tools: dict[str, list[FastMCPTool]], mcp_tools: dict[str, list[FastMCPTool]],
pool: SessionPool,
late: _LateConversations,
) -> dict[str, Backend]: ) -> dict[str, Backend]:
"""Construct one backend per agent name. """Construct one backend per agent name.
@@ -285,7 +330,6 @@ async def _build_backends(
for a in raycast_agents: for a in raycast_agents:
backends[a.name] = raycast_backend backends[a.name] = raycast_backend
session_store = PostgresSessionStore(db)
runner = RunnerConfig(user=settings.claude_runner_user, home=settings.claude_home) runner = RunnerConfig(user=settings.claude_runner_user, home=settings.claude_home)
mcp_tool_names = { mcp_tool_names = {
name: [t.name for t in tools] for name, tools in mcp_tools.items() name: [t.name for t in tools] for name, tools in mcp_tools.items()
@@ -321,6 +365,8 @@ async def _build_backends(
mcp_tool_names=mcp_tool_names, mcp_tool_names=mcp_tool_names,
runner=runner, runner=runner,
usage_sink=record_usage, usage_sink=record_usage,
pool=pool,
tool_server=functools.partial(late.server, names=a.gateway_tools),
) )
await stack.enter_async_context(adapter) await stack.enter_async_context(adapter)
backends[a.name] = adapter backends[a.name] = adapter
+2
View File
@@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Any
from beaver_gateway.agents.base import BaseAgent, ExposedMcp from beaver_gateway.agents.base import BaseAgent, ExposedMcp
from beaver_gateway.agents.claude import ClaudeAgent from beaver_gateway.agents.claude import ClaudeAgent
from beaver_gateway.agents.raycast import RaycastAgent from beaver_gateway.agents.raycast import RaycastAgent
from beaver_gateway.core.conversations import ConversationTexts
from beaver_gateway.core.registry import Gateway from beaver_gateway.core.registry import Gateway
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.mcp.types import HttpMcp, McpServer, PythonToolMcp, StdioMcp from beaver_gateway.mcp.types import HttpMcp, McpServer, PythonToolMcp, StdioMcp
@@ -37,6 +38,7 @@ _PUBLIC_NAMES: dict[str, Any] = {
"McpServer": McpServer, "McpServer": McpServer,
"ExposedMcp": ExposedMcp, "ExposedMcp": ExposedMcp,
"Gateway": Gateway, "Gateway": Gateway,
"ConversationTexts": ConversationTexts,
} }
_McpInstance = StdioMcp | HttpMcp | PythonToolMcp _McpInstance = StdioMcp | HttpMcp | PythonToolMcp
+2 -1
View File
@@ -60,12 +60,13 @@ class TokenStoreError(ValueError):
"""Malformed ``BOOTSTRAP_TOKENS`` value or duplicate token.""" """Malformed ``BOOTSTRAP_TOKENS`` value or duplicate token."""
VALID_SCOPES: frozenset[str] = frozenset({"*", "messages", "mcp", "admin"}) VALID_SCOPES: frozenset[str] = frozenset({"*", "messages", "mcp", "admin", "api"})
"""The scopes a ``Token.scope`` may hold (Phase 4.3 admin UI enforces). """The scopes a ``Token.scope`` may hold (Phase 4.3 admin UI enforces).
* ``*`` — wildcard, may use any frontend * ``*`` — wildcard, may use any frontend
* ``messages`` — Anthropic Messages frontend only * ``messages`` — Anthropic Messages frontend only
* ``mcp`` — MCP server frontend only * ``mcp`` — MCP server frontend only
* ``api`` — conversations API (``/api``) and its SSE
* ``admin`` — reserved for programmatic admin access; the AdminFrontend * ``admin`` — reserved for programmatic admin access; the AdminFrontend
itself authenticates via session cookies, not bearer tokens, so this itself authenticates via session cookies, not bearer tokens, so this
scope is unused today and kept for forward compatibility. scope is unused today and kept for forward compatibility.
+69
View File
@@ -0,0 +1,69 @@
"""In-process event bus: what the gateway does, as a stream frontends can tap.
Events are plain dicts ``{"type", "seq", "ts", ...}``. Publishers never
block; a subscriber that falls behind loses its oldest events rather than
stalling the turn that produced them. ``/api/events`` serialises the
stream as SSE, the panel builds its subagent tree from ``stream`` events
carrying ``parent_tool_use_id``.
"""
from __future__ import annotations
import asyncio
import contextlib
import itertools
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from collections.abc import AsyncIterator
__all__ = ["Event", "EventBus"]
Event = dict[str, Any]
class EventBus:
def __init__(self, *, maxsize: int = 2000) -> None:
self._maxsize = maxsize
self._subscribers: set[asyncio.Queue[Event]] = set()
self._seq = itertools.count(1)
def publish(self, type_: str, **data: Any) -> Event:
event: Event = {
"type": type_,
"seq": next(self._seq),
"ts": datetime.now(UTC).isoformat(timespec="milliseconds"),
**data,
}
for queue in list(self._subscribers):
if queue.full():
with contextlib.suppress(asyncio.QueueEmpty):
queue.get_nowait()
queue.put_nowait(event)
return event
@contextlib.asynccontextmanager
async def subscribe(self) -> AsyncIterator[asyncio.Queue[Event]]:
queue: asyncio.Queue[Event] = asyncio.Queue(maxsize=self._maxsize)
self._subscribers.add(queue)
try:
yield queue
finally:
self._subscribers.discard(queue)
async def stream(
self, *, conversation_id: str | None = None
) -> AsyncIterator[Event]:
async with self.subscribe() as queue:
while True:
event = await queue.get()
if (
conversation_id is None
or event.get("conversation_id") == conversation_id
):
yield event
@property
def subscribers(self) -> int:
return len(self._subscribers)
@@ -496,12 +496,14 @@ def _walk_prefix(
s.name for s in inc.structure if isinstance(s, ToolSegment) s.name for s in inc.structure if isinstance(s, ToolSegment)
) )
inc_text_count = sum(1 for s in inc.structure if isinstance(s, TextSegment)) inc_text_count = sum(1 for s in inc.structure if isinstance(s, TextSegment))
if inc_skeleton != st.skeleton: # Files rendered without tool callouts (§3.10) carry no skeleton:
# prose alone decides whether the turn matched.
if inc_skeleton and inc_skeleton != st.skeleton:
return spliced_groups, i return spliced_groups, i
if inc.text == st.spoken_text: if inc.text == st.spoken_text:
spliced_groups.append(list(st.messages)) spliced_groups.append(list(st.messages))
continue continue
if inc_text_count != st.text_segment_count: if inc_skeleton and inc_text_count != st.text_segment_count:
return spliced_groups, i return spliced_groups, i
spliced = _splice_assistant_group(stored_group=st, incoming=inc) spliced = _splice_assistant_group(stored_group=st, incoming=inc)
if spliced is None: if spliced is None:
File diff suppressed because it is too large Load Diff
+169
View File
@@ -0,0 +1,169 @@
"""In-process MCP server with the gateway's own tools (§3.1, §3.2).
One server per live session so every tool knows which conversation is
calling; ``alwaysLoad`` keeps the tools out of tool search. Which names a
session gets comes from ``ClaudeAgent.gateway_tools``.
"""
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Any, cast
from claude_agent_sdk import create_sdk_mcp_server, tool
if TYPE_CHECKING:
from collections.abc import Iterable
from claude_agent_sdk import McpSdkServerConfig, SdkMcpTool
from beaver_gateway.core.conversations import Conversations
__all__ = ["SERVER_NAME", "TOOL_NAMES", "build_tool_server"]
_log = logging.getLogger("beaver_gateway.core.gateway_tools")
SERVER_NAME = "gateway"
TOOL_NAMES = ("read_conversation", "spawn", "say", "schedule", "inject")
def build_tool_server(
conversations: Conversations, *, conversation_key: str, names: Iterable[str]
) -> McpSdkServerConfig | None:
wanted = set(names)
unknown = wanted - set(TOOL_NAMES)
if unknown:
msg = f"unknown gateway tools: {sorted(unknown)}"
raise ValueError(msg)
tools = [t for t in _tools(conversations, conversation_key) if t.name in wanted]
if not tools:
return None
server = create_sdk_mcp_server(SERVER_NAME, tools=tools)
return cast("McpSdkServerConfig", {**server, "alwaysLoad": True})
def _tools(conversations: Conversations, key: str) -> list[SdkMcpTool[Any]]:
async def current() -> Any:
conv = await conversations.get(key)
if conv is None:
msg = f"conversation {key} not found"
raise LookupError(msg)
return conv
@tool(
"read_conversation",
"Read another conversation (a branch, the master, a deep chat) as plain "
"text. `window` limits it to the last N user turns.",
{
"type": "object",
"properties": {
"id": {"type": "string", "description": "conversation id"},
"window": {"type": "integer", "minimum": 1},
},
"required": ["id"],
},
)
async def read_conversation(args: dict[str, Any]) -> dict[str, Any]:
conv = await conversations.get(str(args["id"]))
if conv is None:
return _error(f"conversation {args['id']} not found")
text = await conversations.read(conv, window=args.get("window"))
return _text(text or "(empty)")
@tool(
"spawn",
"Open a new conversation of the given kind (branch = your own thread, "
"deep = a long research chat, job = a headless task). `seed` is how it "
"starts: clean (nothing), morning (handout), copy (copy of this "
"conversation, last `window` turns), brief (your `text`). Returns the id.",
{
"type": "object",
"properties": {
"kind": {"type": "string", "enum": ["branch", "deep", "job"]},
"seed": {
"type": "string",
"enum": ["clean", "morning", "copy", "brief"],
"default": "clean",
},
"text": {"type": "string", "description": "brief for seed=brief"},
"title": {"type": "string"},
"window": {"type": "integer", "minimum": 1},
},
"required": ["kind"],
},
)
async def spawn(args: dict[str, Any]) -> dict[str, Any]:
parent = await current()
child = await conversations.spawn(
kind=str(args["kind"]),
agent=parent.agent_name,
seed=str(args.get("seed") or "clean"),
parent=parent,
text=args.get("text"),
title=args.get("title"),
window=args.get("window"),
origin="mcp",
)
return _text(f"spawned {child.kind} {child.external_id}")
@tool(
"say",
"Say something to the human in the frontend this conversation is bound "
"to. The only way an inject-started turn can speak; silence is simply "
"not calling it.",
{"text": str},
)
async def say(args: dict[str, Any]) -> dict[str, Any]:
conv = await current()
await conversations.say(conv, str(args["text"]))
return _text("ok")
@tool(
"schedule",
"Promise yourself an inject later: `at` is `+15m`, `+2h`, `+1d` or an "
"ISO datetime; `text` arrives in this conversation at that time.",
{"at": str, "text": str},
)
async def schedule(args: dict[str, Any]) -> dict[str, Any]:
conv = await current()
row = await conversations.schedule(conv, str(args["at"]), str(args["text"]))
return _text(f"scheduled #{row.id} at {row.execute_at.isoformat()}")
@tool(
"inject",
"Put a system-origin message into another conversation's queue.",
{
"type": "object",
"properties": {
"conversation": {"type": "string"},
"text": {"type": "string"},
"urgency": {
"type": "string",
"enum": ["normal", "urgent"],
"default": "normal",
},
},
"required": ["conversation", "text"],
},
)
async def inject(args: dict[str, Any]) -> dict[str, Any]:
target = await conversations.get(str(args["conversation"]))
if target is None:
return _error(f"conversation {args['conversation']} not found")
item = await conversations.inject(
target,
str(args["text"]),
urgency=cast("Any", args.get("urgency") or "normal"),
origin="агент",
)
return _text(f"queued #{item.id}")
return [read_conversation, spawn, say, schedule, inject]
def _text(text: str) -> dict[str, Any]:
return {"content": [{"type": "text", "text": text}]}
def _error(text: str) -> dict[str, Any]:
return {"content": [{"type": "text", "text": text}], "is_error": True}
+130
View File
@@ -0,0 +1,130 @@
"""Persisted per-conversation queue with priorities ``urgent > user > normal`` (§3.4).
One ``ClaudeSDKClient`` runs one turn at a time, so ordering has to happen
before the client: the rows here are the queue, ``core/conversations``
runs one worker per conversation over them. A row that is still
``running`` when the gateway starts was cut by a restart; it is flagged
``interrupted`` and never re-run.
"""
from __future__ import annotations
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Literal
from sqlmodel import col, select
from beaver_gateway.storage.models import InjectQueueItem
if TYPE_CHECKING:
from collections.abc import Iterable, Sequence
from beaver_gateway.storage.db import Database
__all__ = ["PRIORITY_RANK", "InjectQueue", "Priority", "inject_header"]
Priority = Literal["urgent", "user", "normal"]
PRIORITY_RANK: dict[str, int] = {"urgent": 0, "user": 1, "normal": 2}
def inject_header(origin: str) -> str:
return f"[инжект: {origin} - это не Бобёр, отвечать не нужно, голос не обязателен]"
class InjectQueue:
def __init__(self, db: Database) -> None:
self._db = db
async def push(
self, *, conversation_id: int, priority: Priority, origin: str, text: str
) -> InjectQueueItem:
if priority not in PRIORITY_RANK:
msg = f"unknown priority {priority!r}"
raise ValueError(msg)
row = InjectQueueItem(
conversation_id=conversation_id, priority=priority, origin=origin, text=text
)
async with self._db.session() as session:
session.add(row)
await session.commit()
await session.refresh(row)
return row
async def pending(self, conversation_id: int) -> list[InjectQueueItem]:
async with self._db.session() as session:
result = await session.exec(
select(InjectQueueItem)
.where(
InjectQueueItem.conversation_id == conversation_id,
InjectQueueItem.status == "queued",
)
.order_by(col(InjectQueueItem.created_at), col(InjectQueueItem.id))
)
rows = list(result.all())
rows.sort(key=lambda r: (PRIORITY_RANK.get(r.priority, 9), r.created_at))
return rows
async def conversations_with_pending(self) -> list[int]:
async with self._db.session() as session:
result = await session.exec(
select(InjectQueueItem.conversation_id)
.where(InjectQueueItem.status == "queued")
.distinct()
)
return list(result.all())
async def start(self, items: Iterable[InjectQueueItem], turn_id: str) -> None:
await self._mark(items, status="running", turn_id=turn_id, delivered=True)
async def finish(
self, items: Iterable[InjectQueueItem], status: str = "done"
) -> None:
await self._mark(items, status=status)
async def interrupted(self) -> Sequence[InjectQueueItem]:
async with self._db.session() as session:
result = await session.exec(
select(InjectQueueItem).where(InjectQueueItem.status == "running")
)
rows = list(result.all())
for row in rows:
row.status = "interrupted"
session.add(row)
await session.commit()
return rows
async def recent(
self, conversation_id: int, *, limit: int = 50
) -> list[InjectQueueItem]:
async with self._db.session() as session:
result = await session.exec(
select(InjectQueueItem)
.where(InjectQueueItem.conversation_id == conversation_id)
.order_by(col(InjectQueueItem.id).desc())
.limit(limit)
)
return list(result.all())
async def _mark(
self,
items: Iterable[InjectQueueItem],
*,
status: str,
turn_id: str | None = None,
delivered: bool = False,
) -> None:
ids = [i.id for i in items if i.id is not None]
if not ids:
return
async with self._db.session() as session:
result = await session.exec(
select(InjectQueueItem).where(col(InjectQueueItem.id).in_(ids))
)
for row in result.all():
row.status = status
if turn_id is not None:
row.turn_id = turn_id
if delivered:
row.delivered_at = datetime.now(UTC)
session.add(row)
await session.commit()
+3
View File
@@ -16,6 +16,7 @@ if TYPE_CHECKING:
from collections.abc import Iterable, Iterator from collections.abc import Iterable, Iterator
from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.core.conversations import ConversationTexts
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.mcp.types import McpServerT from beaver_gateway.mcp.types import McpServerT
@@ -81,3 +82,5 @@ class Gateway:
agents: list[BaseAgent] = field(default_factory=list) agents: list[BaseAgent] = field(default_factory=list)
mcps: list[McpServerT] = field(default_factory=list) mcps: list[McpServerT] = field(default_factory=list)
frontends: list[Frontend] = field(default_factory=list) frontends: list[Frontend] = field(default_factory=list)
texts: ConversationTexts | None = None
"""Merge prompt and seed bodies for ``core/conversations`` (§8.2-8.3)."""
+257
View File
@@ -0,0 +1,257 @@
"""Pool of live Agent SDK sessions across every Claude agent (§3.2).
One :class:`Session` is one ``ClaudeSDKClient`` (one claude subprocess).
The pool owns the two decisions the adapters used to make on their own:
when a session is closed for idleness (TTL by conversation kind) and
which one goes when memory runs out (measured RSS of the subprocess tree
against the cgroup limit, ``max_live`` where there is no limit). Eviction
only ever picks ``idle && !running_turn && !pending_question`` sessions
that are neither pinned (the master) nor ``dirty`` (mirror gap not yet
repaired); forks and jobs go first.
"""
from __future__ import annotations
import asyncio
import contextlib
import logging
import os
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import TYPE_CHECKING, Any, Protocol
import psutil
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterator, Mapping
__all__ = ["DEFAULT_TTL", "Session", "SessionClient", "SessionPool", "cgroup_limit"]
_log = logging.getLogger("beaver_gateway.core.sessions")
DEFAULT_TTL: Mapping[str, float | None] = {
"master": None,
"branch": 7200.0,
"deep": 1800.0,
"job": 0.0,
"fork": 0.0,
}
"""Idle seconds before a session is closed; ``None`` = never (pinned kinds)."""
_EVICT_ORDER = {"fork": 0, "job": 0, "deep": 1, "branch": 2, "master": 3}
_RSS_HEADROOM = 0.8
class SessionClient(Protocol):
async def connect(self) -> None: ...
async def query(self, prompt: str) -> None: ...
def receive_response(self) -> AsyncIterator[Any]: ...
async def interrupt(self) -> None: ...
async def disconnect(self) -> None: ...
@dataclass
class Session:
key: str
agent: str
kind: str
client: SessionClient
session_id: str | None
resumed: bool
pinned: bool = False
dirty: bool = False
running_turn: str | None = None
pending_question: bool = False
interrupt_requested: bool = False
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
last_used: float = field(default_factory=time.monotonic)
created_at: float = field(default_factory=time.monotonic)
turns: int = 0
@property
def busy(self) -> bool:
return self.lock.locked() or self.running_turn is not None
@property
def evictable(self) -> bool:
return not (self.pinned or self.dirty or self.busy or self.pending_question)
@property
def pid(self) -> int | None:
transport = getattr(self.client, "_transport", None)
process = getattr(transport, "_process", None)
pid = getattr(process, "pid", None)
return pid if isinstance(pid, int) else None
class SessionPool:
def __init__(
self,
*,
rss_limit: int | None = None,
max_live: int = 8,
ttl: Mapping[str, float | None] = DEFAULT_TTL,
reap_interval: float = 60.0,
) -> None:
self._sessions: dict[str, Session] = {}
self._rss_limit = rss_limit if rss_limit is not None else cgroup_limit()
self._max_live = max_live
self._ttl = dict(ttl)
self._reap_interval = reap_interval
def __len__(self) -> int:
return len(self._sessions)
def __iter__(self) -> Iterator[Session]:
return iter(list(self._sessions.values()))
def __contains__(self, key: object) -> bool:
return key in self._sessions
@property
def rss_limit(self) -> int | None:
return self._rss_limit
def get(self, key: str) -> Session | None:
return self._sessions.get(key)
def add(self, session: Session) -> Session:
self._sessions[session.key] = session
return session
def pop(self, key: str) -> Session | None:
return self._sessions.pop(key, None)
def rekey(self, old: str, new: str) -> Session | None:
session = self._sessions.pop(old, None)
if session is None:
return None
stale = self._sessions.pop(new, None)
session.key = new
self._sessions[new] = session
return stale if stale is not session else None
def ttl_for(self, kind: str) -> float | None:
return self._ttl.get(kind, self._ttl.get("deep"))
def rss(self) -> int:
try:
children = psutil.Process(os.getpid()).children(recursive=True)
except psutil.Error:
return 0
total = 0
for child in children:
with contextlib.suppress(psutil.Error):
total += child.memory_info().rss
return total
@staticmethod
def rss_of(session: Session) -> int | None:
pid = session.pid
if pid is None:
return None
try:
process = psutil.Process(pid)
return process.memory_info().rss + sum(
c.memory_info().rss for c in process.children(recursive=True)
)
except psutil.Error:
return None
def over_limit(self) -> bool:
if self._rss_limit is not None:
return self.rss() > self._rss_limit * _RSS_HEADROOM
return len(self._sessions) >= self._max_live
def victims(self) -> list[Session]:
candidates = [s for s in self._sessions.values() if s.evictable]
candidates.sort(key=lambda s: (_EVICT_ORDER.get(s.kind, 1), s.last_used))
return candidates
async def make_room(self) -> int:
closed = 0
while self.over_limit():
victims = self.victims()
if not victims:
_log.warning(
"session pool over limit (%d live, rss=%d) but nothing evictable",
len(self._sessions),
self.rss(),
)
break
await self.close(victims[0].key)
closed += 1
return closed
async def close(self, key: str) -> None:
session = self._sessions.pop(key, None)
if session is None:
return
_log.info("closing session %s (%s, %s)", session.session_id, session.kind, key)
try:
await session.client.disconnect()
except Exception: # noqa: BLE001
_log.exception("disconnect failed for session %s", session.session_id)
async def close_all(self, *, agent: str | None = None) -> None:
for session in list(self._sessions.values()):
if agent is None or session.agent == agent:
await self.close(session.key)
async def reap_once(self) -> int:
now = time.monotonic()
closed = 0
for session in list(self._sessions.values()):
ttl = self.ttl_for(session.kind)
if ttl is None or not session.evictable:
continue
if now - session.last_used > ttl:
await self.close(session.key)
closed += 1
return closed
async def reap_loop(self) -> None:
while True:
await asyncio.sleep(self._reap_interval)
try:
await self.reap_once()
await self.make_room()
except Exception: # noqa: BLE001
_log.exception("session reaper failed")
def snapshot(self) -> list[dict[str, Any]]:
now = time.monotonic()
return [
{
"key": s.key,
"agent": s.agent,
"kind": s.kind,
"session_id": s.session_id,
"pid": s.pid,
"rss": self.rss_of(s),
"idle_seconds": round(now - s.last_used, 1),
"age_seconds": round(now - s.created_at, 1),
"turns": s.turns,
"busy": s.busy,
"running_turn": s.running_turn,
"pending_question": s.pending_question,
"pinned": s.pinned,
"dirty": s.dirty,
}
for s in self._sessions.values()
]
def cgroup_limit() -> int | None:
for path in (
"/sys/fs/cgroup/memory.max",
"/sys/fs/cgroup/memory/memory.limit_in_bytes",
):
try:
raw = Path(path).read_text(encoding="ascii").strip()
except OSError:
continue
if raw.isdigit() and int(raw) < 1 << 60:
return int(raw)
return None
+229 -1
View File
@@ -27,7 +27,17 @@ except ImportError: # pragma: no cover
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Iterable, Mapping from collections.abc import Iterable, Mapping
__all__ = ["CLI_VERSION", "build_entries", "messages_from_entries"] __all__ = [
"CLI_VERSION",
"build_entries",
"close_open_tool_uses",
"messages_from_entries",
"open_tool_uses",
"prompt_count",
"render_messages",
"strip_tool_entries",
"window_entries",
]
CLI_VERSION = _cli_version CLI_VERSION = _cli_version
_ENTRYPOINT = "sdk-py" _ENTRYPOINT = "sdk-py"
@@ -275,3 +285,221 @@ def _zero_usage() -> dict[str, Any]:
"iterations": [], "iterations": [],
"speed": "standard", "speed": "standard",
} }
# ---- repair, windows, projections ---------------------------------------
_PROMPT_TYPES = ("user", "assistant")
_INTERRUPTED = "прервано"
def open_tool_uses(
entries: Iterable[Mapping[str, Any]],
) -> list[tuple[Mapping[str, Any], dict[str, Any]]]:
"""``(assistant entry, tool_use block)`` pairs that never got a ``tool_result``."""
closed: set[str] = set()
uses: list[tuple[Mapping[str, Any], dict[str, Any]]] = []
for entry in entries:
content = _entry_content(entry)
if entry.get("type") == "user":
closed.update(
str(b.get("tool_use_id", ""))
for b in content
if b.get("type") == "tool_result"
)
elif entry.get("type") == "assistant":
uses.extend((entry, b) for b in content if b.get("type") == "tool_use")
return [
(owner, block)
for owner, block in uses
if str(block.get("id", "")) not in closed
]
def close_open_tool_uses(
entries: list[Mapping[str, Any]], *, text: str = _INTERRUPTED
) -> list[dict[str, Any]]:
"""Synthetic ``tool_result`` entries on the leaf, one per open ``tool_use``.
Appending the result to the session store gives the next ``resume`` a
transcript the CLI accepts: an assistant message ending in ``tool_use``
without its result is rejected by the API on the next call.
"""
pending = open_tool_uses(entries)
if not pending:
return []
leaf = next(
(
e
for e in reversed(entries)
if e.get("type") in _PROMPT_TYPES and e.get("uuid")
),
None,
)
parent = str(leaf["uuid"]) if leaf is not None else None
stamp = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S.000Z")
out: list[dict[str, Any]] = []
for owner, block in pending:
uid = _new_uuid()
result = {
"type": "tool_result",
"tool_use_id": block.get("id", ""),
"content": text,
"is_error": True,
}
out.append(
{
"parentUuid": parent,
"promptId": owner.get("promptId"),
"type": "user",
"message": {"role": "user", "content": [result]},
"uuid": uid,
"timestamp": stamp,
"toolUseResult": text,
"sourceToolAssistantUUID": owner.get("uuid"),
**{k: owner[k] for k in _COMMON_KEYS if k in owner},
}
)
parent = uid
return out
_COMMON_KEYS = (
"isSidechain",
"userType",
"entrypoint",
"cwd",
"sessionId",
"version",
"gitBranch",
)
def window_entries(
entries: Iterable[Mapping[str, Any]], *, window: int | None = None
) -> list[dict[str, Any]]:
"""``user``/``assistant`` entries only, last ``window`` prompts, relinked."""
kept = [
dict(e)
for e in entries
if e.get("type") in _PROMPT_TYPES and isinstance(e.get("uuid"), str)
]
if window is not None and window > 0:
starts = [i for i, e in enumerate(kept) if _is_prompt_entry(e)]
if len(starts) > window:
kept = kept[starts[-window] :]
return _relink(kept)
def strip_tool_entries(entries: Iterable[Mapping[str, Any]]) -> list[dict[str, Any]]:
"""Drop ``tool_result`` entries and ``tool_use``/``thinking`` blocks; relink."""
out: list[dict[str, Any]] = []
for raw in entries:
entry = dict(raw)
message = entry.get("message")
if not isinstance(message, dict):
out.append(entry)
continue
content = message.get("content")
if isinstance(content, list):
blocks = [
b
for b in content
if isinstance(b, dict)
and b.get("type") not in ("tool_use", "tool_result", "thinking")
]
if not blocks:
continue
entry["message"] = {**message, "content": blocks}
out.append(entry)
return _relink(out)
def prompt_count(entries: Iterable[Mapping[str, Any]]) -> int:
return sum(1 for e in entries if _is_prompt_entry(e))
def render_messages(
messages: Iterable[Mapping[str, Any]], *, window: int | None = None
) -> str:
"""Plain-text projection for ``read_conversation``: ``user:``/``assistant:`` turns.
Tool calls collapse to a one-line summary per assistant turn; tool
results and thinking are dropped.
"""
turns: list[str] = []
tools: list[str] = []
current: list[str] = []
def flush() -> None:
if not current and not tools:
return
body = "\n\n".join(current).strip()
if tools:
body = (body + "\n" if body else "") + "(tools: " + ", ".join(tools) + ")"
turns.append("assistant:\n" + body)
current.clear()
tools.clear()
for message in messages:
role = message.get("role")
content = message.get("content")
if role == "user":
if _tool_results(content):
continue
flush()
turns.append("user:\n" + _text_of_content(content))
continue
for block in _assistant_blocks(content):
if block.get("type") == "text" and block.get("text"):
current.append(str(block["text"]))
elif block.get("type") == "tool_use":
tools.append(str(block.get("name", "")))
flush()
if window is not None and window > 0:
starts = [i for i, t in enumerate(turns) if t.startswith("user:")]
if len(starts) > window:
turns = turns[starts[-window] :]
return "\n\n".join(turns)
def _entry_content(entry: Mapping[str, Any]) -> list[dict[str, Any]]:
message = entry.get("message")
if not isinstance(message, dict):
return []
content = message.get("content")
if not isinstance(content, list):
return []
return [b for b in content if isinstance(b, dict)]
def _is_prompt_entry(entry: Mapping[str, Any]) -> bool:
if entry.get("type") != "user":
return False
message = entry.get("message")
if not isinstance(message, dict):
return False
content = message.get("content")
if isinstance(content, str):
return True
return not _tool_results(content)
def _relink(entries: list[dict[str, Any]]) -> list[dict[str, Any]]:
parent: str | None = None
for entry in entries:
entry["parentUuid"] = parent
parent = entry.get("uuid")
return entries
def _text_of_content(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
return "\n\n".join(
str(b.get("text", ""))
for b in content
if isinstance(b, dict) and b.get("type") == "text" and b.get("text")
)
return ""
+2
View File
@@ -35,3 +35,5 @@ class TurnCapture:
"""Backend session that ran the turn; persist and pass back as ``session_id``.""" """Backend session that ran the turn; persist and pass back as ``session_id``."""
usage: TurnUsage | None = None usage: TurnUsage | None = None
interrupted: bool = False
"""The turn was cut by ``interrupt()`` (urgent inject), not by an error."""
+63
View File
@@ -0,0 +1,63 @@
"""Server-sent events helpers shared by the markdown and API frontends."""
from __future__ import annotations
import asyncio
import contextlib
import json
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from collections.abc import AsyncIterator
__all__ = ["HEARTBEAT_INTERVAL", "SSE_HEADERS", "events_with_heartbeat", "sse_pack"]
HEARTBEAT_INTERVAL = 15.0
"""Seconds of backend silence before a comment frame keeps the socket warm."""
SSE_HEADERS = {
"Cache-Control": "no-cache, no-transform",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
}
KEEPALIVE = b": keepalive\n\n"
async def events_with_heartbeat(
events: AsyncIterator[Any], interval: float = HEARTBEAT_INTERVAL
) -> AsyncIterator[Any]:
"""Pass ``events`` through, yielding ``None`` after ``interval`` seconds of silence.
One in-flight ``__anext__`` task is reused across timeouts: a second
consumer on the same async generator raises ``RuntimeError``.
Cancellation of the outer scope cancels that task instead of leaving
it dangling.
"""
src = events.__aiter__()
next_task: asyncio.Task[Any] | None = None
try:
while True:
if next_task is None:
next_task = asyncio.ensure_future(src.__anext__())
done, _pending = await asyncio.wait({next_task}, timeout=interval)
if not done:
yield None
continue
task = next_task
next_task = None
try:
result = task.result()
except StopAsyncIteration:
return
yield result
finally:
if next_task is not None and not next_task.done():
next_task.cancel()
with contextlib.suppress(BaseException):
await next_task
def sse_pack(event: str, data: dict[str, Any]) -> bytes:
body = json.dumps(data, ensure_ascii=False)
return f"event: {event}\ndata: {body}\n\n".encode()
@@ -0,0 +1,5 @@
"""``ApiFrontend`` - the conversations API and event stream (§3.9)."""
from beaver_gateway.frontends.api.frontend import ApiFrontend
__all__ = ["ApiFrontend"]
@@ -0,0 +1,491 @@
"""``ApiFrontend`` - ``/api/conversations``, SSE events, sessions, usage (§3.9).
Bearer scope ``api``. Every write goes through ``core/conversations``; the
frontend only shapes JSON. ``/api/events`` and
``/api/conversations/{id}/events`` replay the gateway bus as SSE with the
same keepalive the markdown frontend uses, so a proxy never sees a
silent socket.
"""
from __future__ import annotations
import json
import logging
from datetime import UTC, datetime, timedelta
from typing import TYPE_CHECKING, Any, cast
from fastapi import FastAPI, HTTPException, Request, status
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, StreamingResponse
from sqlalchemy import func
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.frontends._auth import require_token
from beaver_gateway.frontends._sse import (
KEEPALIVE,
SSE_HEADERS,
events_with_heartbeat,
sse_pack,
)
from beaver_gateway.frontends.base import Frontend
from beaver_gateway.storage.models import Usage
if TYPE_CHECKING:
from collections.abc import AsyncIterator
from beaver_gateway.core.conversations import Conversations
from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.storage.models import Conversation
_log = logging.getLogger("beaver_gateway.frontends.api")
__all__ = ["ApiFrontend"]
SCOPE = "api"
class ApiFrontend(Frontend):
def __init__(
self,
*,
host: str = "0.0.0.0", # noqa: S104
port: int = 8004,
public_base_url: 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._app: FastAPI | None = None
def configure(self, runtime: GatewayRuntime) -> None:
if runtime.conversations is None or runtime.bus is None:
msg = "ApiFrontend needs runtime.conversations and runtime.bus"
raise RuntimeError(msg)
self._app = _build_app(runtime)
async def serve(self) -> None:
import uvicorn
if self._app is None:
msg = "configure() must be called before serve()"
raise RuntimeError(msg)
server = uvicorn.Server(
uvicorn.Config(self._app, host=self.host, port=self.port, log_level="info")
)
await server.serve()
def _build_app(runtime: GatewayRuntime) -> FastAPI: # noqa: PLR0915
app = FastAPI(title="beaver-gateway / API")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=False,
allow_methods=["*"],
allow_headers=["*"],
)
conversations = cast("Conversations", runtime.conversations)
async def body_of(request: Request) -> dict[str, Any]:
if not await request.body():
return {}
try:
data = await request.json()
except json.JSONDecodeError as exc:
raise HTTPException(
status.HTTP_400_BAD_REQUEST, f"invalid JSON: {exc}"
) from exc
if not isinstance(data, dict):
raise HTTPException(status.HTTP_400_BAD_REQUEST, "body must be an object")
return data
async def conv_of(public_id: str) -> Conversation:
conv = await conversations.get(public_id)
if conv is None:
raise HTTPException(
status.HTTP_404_NOT_FOUND, f"unknown conversation {public_id}"
)
return conv
def text_of(data: dict[str, Any], key: str = "text") -> str:
text = data.get(key)
if not isinstance(text, str) or not text.strip():
raise HTTPException(status.HTTP_400_BAD_REQUEST, f"missing `{key}`")
return text
def int_or_none(data: dict[str, Any], key: str) -> int | None:
value = data.get(key)
if value is None:
return None
if not isinstance(value, int) or value < 1:
raise HTTPException(
status.HTTP_400_BAD_REQUEST, f"`{key}` must be a positive int"
)
return value
@app.get("/healthz")
async def healthz() -> dict[str, str]:
return {"status": "ok"}
@app.get("/api/conversations")
async def list_conversations(request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
q = request.query_params
rows = await conversations.find(
status=q.get("status"), kind=q.get("kind"), limit=int(q.get("limit", "200"))
)
return {"conversations": [conversations.public(r) for r in rows]}
@app.post("/api/conversations", status_code=status.HTTP_201_CREATED)
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]}"
)
if not isinstance(agent, str) or agent not in runtime.agents:
raise HTTPException(
status.HTTP_400_BAD_REQUEST, "unknown or missing `agent`"
)
seed = str(data.get("seed") or "clean")
if seed not in SEEDS:
raise HTTPException(
status.HTTP_400_BAD_REQUEST, f"seed must be one of {SEEDS}"
)
parent = await conv_of(str(data["parent"])) if data.get("parent") else None
try:
conv = await conversations.spawn(
kind=kind,
agent=agent,
seed=seed,
parent=parent,
text=data.get("text"),
title=data.get("title"),
window=int_or_none(data, "window"),
origin="api",
)
except ValueError as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc
await audit.log(
runtime,
actor=f"token:{token}",
kind="api_spawn",
agent_name=agent,
conversation=conv.external_id,
seed=seed,
)
return await conversations.describe(conv)
@app.get("/api/conversations/{public_id}")
async def get_conversation(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
out = await conversations.describe(conv)
out["queue"] = [
{
"id": i.id,
"priority": i.priority,
"origin": i.origin,
"status": i.status,
"created_at": i.created_at.isoformat(),
"text": i.text[:200],
}
for i in await conversations.queue.recent(cast("int", conv.id), limit=20)
]
return out
@app.get("/api/conversations/{public_id}/messages")
async def get_messages(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
raw = request.query_params.get("window")
window = int(raw) if raw and raw.isdigit() else None
return {
"id": conv.external_id,
"text": await conversations.read(conv, window=window),
}
@app.post(
"/api/conversations/{public_id}/messages", status_code=status.HTTP_202_ACCEPTED
)
async def post_message(public_id: str, request: Request) -> dict[str, Any]:
token = await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
data = await body_of(request)
item = await conversations.post(
conv, text_of(data), origin=str(data.get("origin") or "user")
)
await audit.log(
runtime,
actor=f"token:{token}",
kind="api_message",
agent_name=conv.agent_name,
conversation=conv.external_id,
)
return {"id": conv.external_id, "item": item.id, "status": item.status}
@app.post(
"/api/conversations/{public_id}/inject", status_code=status.HTTP_202_ACCEPTED
)
async def post_inject(public_id: str, request: Request) -> dict[str, Any]:
token = await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
data = await body_of(request)
urgency = str(data.get("urgency") or "normal")
if urgency not in ("normal", "urgent"):
raise HTTPException(
status.HTTP_400_BAD_REQUEST, "urgency must be normal|urgent"
)
item = await conversations.inject(
conv,
text_of(data),
urgency=cast("Any", urgency),
origin=str(data.get("origin") or "api"),
)
await audit.log(
runtime,
actor=f"token:{token}",
kind="api_inject",
agent_name=conv.agent_name,
conversation=conv.external_id,
urgency=urgency,
)
return {"id": conv.external_id, "item": item.id, "priority": item.priority}
@app.post("/api/conversations/{public_id}/say")
async def post_say(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
return await conversations.say(conv, text_of(await body_of(request)))
@app.post(
"/api/conversations/{public_id}/branch", status_code=status.HTTP_201_CREATED
)
async def post_branch(public_id: str, request: Request) -> dict[str, Any]:
token = await require_token(request, runtime, scope=SCOPE)
parent = await conv_of(public_id)
data = await body_of(request)
seed = str(data.get("seed_mode") or data.get("seed") or "morning")
if seed not in SEEDS:
raise HTTPException(
status.HTTP_400_BAD_REQUEST, f"seed must be one of {SEEDS}"
)
agent = str(data.get("agent") or parent.agent_name)
try:
child = await conversations.spawn(
kind="branch",
agent=agent,
seed=seed,
parent=parent,
text=data.get("text"),
title=data.get("title"),
window=int_or_none(data, "window"),
origin="api",
)
except (ValueError, LookupError) as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc
await audit.log(
runtime,
actor=f"token:{token}",
kind="api_branch",
agent_name=agent,
conversation=child.external_id,
parent=parent.external_id,
seed=seed,
)
return await conversations.describe(child)
@app.post("/api/conversations/{public_id}/merge")
async def post_merge(public_id: str, request: Request) -> dict[str, Any]:
token = await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
try:
result = await conversations.merge(conv)
except (ValueError, LookupError) as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc
await audit.log(
runtime,
actor=f"token:{token}",
kind="api_merge",
agent_name=conv.agent_name,
conversation=conv.external_id,
)
return {
"id": conv.external_id,
"status": "merged",
"fork": result.conversation.external_id,
"text": result.text,
}
@app.post("/api/conversations/{public_id}/fork")
async def post_fork(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
data = await body_of(request)
try:
result = await conversations.fork(
conv,
text_of(data, "prompt"),
window=int_or_none(data, "window"),
strip_tools=bool(data.get("strip_tools", False)),
)
except (ValueError, RuntimeError) as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc
return {
"id": conv.external_id,
"fork": result.conversation.external_id,
"text": result.text,
}
@app.post("/api/conversations/{public_id}/bind")
async def post_bind(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
data = await body_of(request)
frontend = text_of(data, "frontend")
external_id = text_of(data, "external_id")
await conversations.bind(
conv,
frontend=frontend,
external_id=external_id,
visible=bool(data.get("visible", True)),
)
return await conversations.describe(conv)
@app.patch("/api/conversations/{public_id}/flags")
async def patch_flags(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
data = await body_of(request)
return conversations.public(await conversations.set_flags(conv, data))
@app.patch("/api/conversations/{public_id}")
async def patch_conversation(public_id: str, request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
data = await body_of(request)
try:
if isinstance(data.get("status"), str):
conv = await conversations.set_status(conv, data["status"])
if isinstance(data.get("title"), str):
conv = await conversations.set_title(conv, data["title"])
except ValueError as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc
return conversations.public(conv)
@app.get("/api/conversations/{public_id}/events")
async def conversation_events(public_id: str, request: Request) -> Any:
await require_token(request, runtime, scope=SCOPE)
conv = await conv_of(public_id)
return _sse(runtime, conversation_id=conv.external_id)
@app.get("/api/events")
async def all_events(request: Request) -> Any:
await require_token(request, runtime, scope=SCOPE)
return _sse(runtime, conversation_id=None)
@app.get("/api/sessions")
async def sessions(request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
pool = runtime.pool
return {
"rss": pool.rss() if pool is not None else None,
"rss_limit": pool.rss_limit if pool is not None else None,
"sessions": pool.snapshot() if pool is not None else [],
}
@app.get("/api/schedules")
async def schedules(request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
raw = request.query_params.get("conversation")
conv = await conv_of(raw) if raw else None
return {
"schedules": [
{
"id": s.id,
"conversation_row": s.conversation_id,
"execute_at": s.execute_at.isoformat(),
"text": s.text,
"delivered_at": s.delivered_at.isoformat()
if s.delivered_at
else None,
}
for s in await conversations.schedules(conv)
]
}
@app.get("/api/usage")
async def usage(request: Request) -> dict[str, Any]:
await require_token(request, runtime, scope=SCOPE)
hours = float(request.query_params.get("hours", "24"))
since = datetime.now(UTC) - timedelta(hours=hours)
columns = (
func.count(),
func.coalesce(func.sum(Usage.input_tokens), 0),
func.coalesce(func.sum(Usage.output_tokens), 0),
func.coalesce(func.sum(Usage.cache_read_tokens), 0),
func.coalesce(func.sum(Usage.cache_creation_tokens), 0),
func.coalesce(func.sum(Usage.cost_usd), 0.0),
)
async with runtime.db.session() as session:
by_agent = (
await session.execute( # ty: ignore[deprecated]
sa_select(col(Usage.agent_name), *columns)
.where(col(Usage.ts) >= since.replace(tzinfo=None))
.group_by(col(Usage.agent_name))
)
).all()
by_conversation = (
await session.execute( # ty: ignore[deprecated]
sa_select(col(Usage.conversation_id), *columns)
.where(col(Usage.ts) >= since.replace(tzinfo=None))
.group_by(col(Usage.conversation_id))
)
).all()
return {
"since": since.isoformat(timespec="seconds"),
"by_agent": [_usage_row("agent", r) for r in by_agent],
"by_conversation": [_usage_row("conversation", r) for r in by_conversation],
}
@app.exception_handler(HTTPException)
async def http_error(_request: Request, exc: HTTPException) -> JSONResponse:
return JSONResponse(
status_code=exc.status_code,
content={"error": exc.detail},
headers=exc.headers,
)
return app
def _usage_row(label: str, row: Any) -> dict[str, Any]:
key, turns, inp, out, cache_read, cache_creation, cost = row
return {
label: key,
"turns": turns,
"input": inp,
"output": out,
"cache_read": cache_read,
"cache_creation": cache_creation,
"cost_usd": round(float(cost or 0.0), 4),
}
def _sse(runtime: GatewayRuntime, *, conversation_id: str | None) -> StreamingResponse:
async def gen() -> AsyncIterator[bytes]:
stream = runtime.bus.stream(conversation_id=conversation_id)
yield sse_pack("hello", {"type": "hello", "conversation_id": conversation_id})
async for event in events_with_heartbeat(stream):
if event is None:
yield KEEPALIVE
continue
yield sse_pack(str(event["type"]), event)
return StreamingResponse(gen(), media_type="text/event-stream", headers=SSE_HEADERS)
+7 -1
View File
@@ -12,7 +12,7 @@ from __future__ import annotations
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING from typing import TYPE_CHECKING, Any
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Awaitable, Callable, Mapping, Sequence from collections.abc import Awaitable, Callable, Mapping, Sequence
@@ -77,6 +77,12 @@ class GatewayRuntime:
# to keep ``anthropic.types`` out of the runtime import graph for # to keep ``anthropic.types`` out of the runtime import graph for
# this base module. # this base module.
turn_log_handlers: list[TurnLogHandler] = field(default_factory=list) turn_log_handlers: list[TurnLogHandler] = field(default_factory=list)
# M1b: conversations service, event bus and the shared session pool.
# ``Any`` for the same import-graph reason as above; ``None`` only in
# tests that build a runtime without them.
conversations: Any = None
bus: Any = None
pool: Any = None
class Frontend(ABC): class Frontend(ABC):
+84 -194
View File
@@ -13,7 +13,13 @@ Concurrency model: an in-memory ``set[Path]`` of files currently in
flight. Two concurrent requests for the same file → the second gets flight. Two concurrent requests for the same file → the second gets
409. The set is single-process (one gateway instance) — that's by 409. The set is single-process (one gateway instance) — that's by
design; the markdown frontend is the only writer in its vault from design; the markdown frontend is the only writer in its vault from
the gateway side. the gateway side. The turn itself runs through ``core/conversations``
(one turn per conversation, ``running_turn`` in the DB), so a message
posted to the same conversation via ``/api`` waits its turn.
A chat file is a ``deep`` conversation bound as
``(markdown, <vault-relative path>)``; frontmatter carries only ``agent``
and ``conversation_id`` (§3.10), tool calls are never rendered.
Cross-frontend logging: when ``log_all_chats=True``, ``configure()`` Cross-frontend logging: when ``log_all_chats=True``, ``configure()``
registers a handler on ``runtime.turn_log_handlers`` so every other registers a handler on ``runtime.turn_log_handlers`` so every other
@@ -24,7 +30,6 @@ shape.
from __future__ import annotations from __future__ import annotations
import asyncio
import contextlib import contextlib
import json import json
import logging import logging
@@ -33,7 +38,7 @@ import tempfile
import time import time
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any, cast
import aiofile import aiofile
from anthropic.types import RawContentBlockStopEvent from anthropic.types import RawContentBlockStopEvent
@@ -44,29 +49,28 @@ from fastapi.responses import JSONResponse, StreamingResponse
from beaver_gateway.core import audit from beaver_gateway.core import audit
from beaver_gateway.core.conversation_store import ( from beaver_gateway.core.conversation_store import (
diff_and_fork, diff_and_fork,
load_conversation,
load_messages, load_messages,
mint_conversation,
rewrite_messages, rewrite_messages,
set_session_id,
) )
from beaver_gateway.core.turn_capture import TurnCapture from beaver_gateway.core.turn_capture import TurnCapture
from beaver_gateway.core.turn_record import TurnRecord from beaver_gateway.core.turn_record import TurnRecord
from beaver_gateway.frontends._accumulate import StreamAccumulator from beaver_gateway.frontends._accumulate import StreamAccumulator
from beaver_gateway.frontends._auth import require_token from beaver_gateway.frontends._auth import require_token
from beaver_gateway.frontends._sse import (
KEEPALIVE,
SSE_HEADERS,
events_with_heartbeat,
sse_pack,
)
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.frontends.markdown import parser, renderer from beaver_gateway.frontends.markdown import parser, renderer
from beaver_gateway.frontends.markdown.crossfront import ( from beaver_gateway.frontends.markdown.crossfront import CrossFrontendLogger
CrossFrontendLogger,
fingerprint_messages,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator, Callable from collections.abc import AsyncIterator, Callable
from anthropic.types import MessageParam
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.storage.models import Conversation
_log = logging.getLogger("beaver_gateway.frontends.markdown") _log = logging.getLogger("beaver_gateway.frontends.markdown")
@@ -89,12 +93,7 @@ _STREAM_FLUSH_DEBOUNCE = 0.4
# disk round-trip). # disk round-trip).
_SSE_FLUSH_DEBOUNCE = 0.1 _SSE_FLUSH_DEBOUNCE = 0.1
# Interval between SSE comment-frames sent when the backend is silent FRONTEND = "markdown"
# (e.g. claude is mid-thinking on a large context). The Obsidian plugin
# and any intermediate proxies will hold the connection open as long as
# bytes keep flowing; a comment-frame is the cheapest legal SSE keepalive.
# Set well under typical proxy/client idle timeouts (60s).
_SSE_HEARTBEAT_INTERVAL = 15.0
class MarkdownFrontend(Frontend): class MarkdownFrontend(Frontend):
@@ -134,6 +133,9 @@ class MarkdownFrontend(Frontend):
self._crossfront: CrossFrontendLogger | None = None self._crossfront: CrossFrontendLogger | None = None
def configure(self, runtime: GatewayRuntime) -> None: def configure(self, runtime: GatewayRuntime) -> None:
if runtime.conversations is None:
msg = "MarkdownFrontend needs runtime.conversations"
raise RuntimeError(msg)
self._runtime = runtime self._runtime = runtime
self.vault_path.mkdir(parents=True, exist_ok=True) self.vault_path.mkdir(parents=True, exist_ok=True)
if self.log_all_chats: if self.log_all_chats:
@@ -286,17 +288,7 @@ class MarkdownFrontend(Frontend):
self._busy.discard(file_path) self._busy.discard(file_path)
return StreamingResponse( return StreamingResponse(
gen(), gen(), media_type="text/event-stream", headers=SSE_HEADERS
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache, no-transform",
"Connection": "keep-alive",
# nginx default-buffers SSE bodies; this header tells
# both nginx and uvicorn-behind-proxy to flush as we
# write. Harmless if the deployment has no reverse
# proxy in front.
"X-Accel-Buffering": "no",
},
) )
return app return app
@@ -384,16 +376,19 @@ class MarkdownFrontend(Frontend):
# stored history, and feed the aligned messages to the backend # stored history, and feed the aligned messages to the backend
# - see ``core/conversation_store.py`` for the full rationale. # - see ``core/conversation_store.py`` for the full rationale.
conv, conv_external_id, stored_msgs = await self._resolve_conversation( conv, conv_external_id, stored_msgs = await self._resolve_conversation(
runtime=runtime, metadata=parsed.metadata, agent_name=agent.name runtime=runtime,
metadata=parsed.metadata,
agent_name=agent.name,
file_path=file_path,
) )
outcome = diff_and_fork(stored=stored_msgs, incoming=parsed.turns) outcome = diff_and_fork(stored=stored_msgs, incoming=parsed.turns)
capture = TurnCapture() capture = TurnCapture()
events = backend.complete( events = runtime.conversations.turn(
agent=agent, conv,
messages=outcome.messages, messages=outcome.messages,
system=None, origin="user",
capture=capture, capture=capture,
**_session_options(conv, outcome.divergence_index), use_session=outcome.divergence_index is None,
) )
try: try:
message = await self._stream_to_file( message = await self._stream_to_file(
@@ -422,7 +417,7 @@ class MarkdownFrontend(Frontend):
await self._persist_canonical_history( await self._persist_canonical_history(
runtime=runtime, runtime=runtime,
conversation_id=conv.id, conversation_id=cast("int", conv.id),
persist_messages=outcome.persist_messages, persist_messages=outcome.persist_messages,
new_user_text=parsed.turns[-1].text, new_user_text=parsed.turns[-1].text,
capture=capture, capture=capture,
@@ -487,7 +482,7 @@ class MarkdownFrontend(Frontend):
elif content_override is None: elif content_override is None:
file_text = await _read_or_empty(file_path) file_text = await _read_or_empty(file_path)
else: else:
yield _sse_pack( yield sse_pack(
"error", "error",
{ {
"status_code": status.HTTP_400_BAD_REQUEST, "status_code": status.HTTP_400_BAD_REQUEST,
@@ -503,7 +498,7 @@ class MarkdownFrontend(Frontend):
default=self.default_agent, default=self.default_agent,
) )
if not agent_name: if not agent_name:
yield _sse_pack( yield sse_pack(
"error", "error",
{ {
"status_code": status.HTTP_400_BAD_REQUEST, "status_code": status.HTTP_400_BAD_REQUEST,
@@ -516,7 +511,7 @@ class MarkdownFrontend(Frontend):
return return
if not parsed.messages: if not parsed.messages:
yield _sse_pack( yield sse_pack(
"done", "done",
{ {
"status": "nothing_to_do", "status": "nothing_to_do",
@@ -527,7 +522,7 @@ class MarkdownFrontend(Frontend):
return return
if parser.last_role(parsed.messages) == "assistant": if parser.last_role(parsed.messages) == "assistant":
yield _sse_pack( yield sse_pack(
"done", "done",
{ {
"status": "nothing_to_do", "status": "nothing_to_do",
@@ -539,7 +534,7 @@ class MarkdownFrontend(Frontend):
agent = runtime.agents.get(agent_name) agent = runtime.agents.get(agent_name)
if agent is None: if agent is None:
yield _sse_pack( yield sse_pack(
"error", "error",
{ {
"status_code": status.HTTP_404_NOT_FOUND, "status_code": status.HTTP_404_NOT_FOUND,
@@ -549,7 +544,7 @@ class MarkdownFrontend(Frontend):
return return
backend = runtime.backends.get(agent.name) backend = runtime.backends.get(agent.name)
if backend is None: if backend is None:
yield _sse_pack( yield sse_pack(
"error", "error",
{ {
"status_code": status.HTTP_503_SERVICE_UNAVAILABLE, "status_code": status.HTTP_503_SERVICE_UNAVAILABLE,
@@ -575,7 +570,10 @@ class MarkdownFrontend(Frontend):
) )
conv, conv_external_id, stored_msgs = await self._resolve_conversation( conv, conv_external_id, stored_msgs = await self._resolve_conversation(
runtime=runtime, metadata=parsed.metadata, agent_name=agent.name runtime=runtime,
metadata=parsed.metadata,
agent_name=agent.name,
file_path=file_path,
) )
_log.info( _log.info(
"chat/stream: file=%s conv_external_id=%s conv_id=%d " "chat/stream: file=%s conv_external_id=%s conv_id=%d "
@@ -602,12 +600,12 @@ class MarkdownFrontend(Frontend):
agent.name, agent.name,
conv.session_id, conv.session_id,
) )
events = backend.complete( events = runtime.conversations.turn(
agent=agent, conv,
messages=outcome.messages, messages=outcome.messages,
system=None, origin="user",
capture=capture, capture=capture,
**_session_options(conv, outcome.divergence_index), use_session=outcome.divergence_index is None,
) )
acc = StreamAccumulator() acc = StreamAccumulator()
@@ -624,13 +622,9 @@ class MarkdownFrontend(Frontend):
return _reattach_frontmatter(parsed.metadata, new_body) return _reattach_frontmatter(parsed.metadata, new_body)
try: try:
async for ev in _events_with_heartbeat(events): async for ev in events_with_heartbeat(events):
if ev is None: if ev is None:
# Backend is quiet (claude mid-thinking, MCP slow, yield KEEPALIVE
# whatever). SSE comment-frame keeps the TCP socket
# warm so the plugin / uvicorn / any reverse proxy
# doesn't time the request out before we finish.
yield b": keepalive\n\n"
continue continue
acc.feed(ev) acc.feed(ev)
now = time.monotonic() now = time.monotonic()
@@ -643,7 +637,7 @@ class MarkdownFrontend(Frontend):
# render to the same prefix as before they closed # render to the same prefix as before they closed
# (we don't surface the tool-call args in markdown). # (we don't surface the tool-call args in markdown).
if payload is not None and payload != last_payload: if payload is not None and payload != last_payload:
yield _sse_pack("delta", {"new_content": payload}) yield sse_pack("delta", {"new_content": payload})
last_payload = payload last_payload = payload
last_flush = now last_flush = now
except Exception as exc: # noqa: BLE001 — wire any backend failure as an SSE error frame except Exception as exc: # noqa: BLE001 — wire any backend failure as an SSE error frame
@@ -662,7 +656,7 @@ class MarkdownFrontend(Frontend):
await _write_atomic( await _write_atomic(
file_path, _reattach_frontmatter(parsed.metadata, new_body) file_path, _reattach_frontmatter(parsed.metadata, new_body)
) )
yield _sse_pack( yield sse_pack(
"error", "error",
{ {
"status_code": status.HTTP_500_INTERNAL_SERVER_ERROR, "status_code": status.HTTP_500_INTERNAL_SERVER_ERROR,
@@ -684,7 +678,7 @@ class MarkdownFrontend(Frontend):
await self._persist_canonical_history( await self._persist_canonical_history(
runtime=runtime, runtime=runtime,
conversation_id=conv.id, conversation_id=cast("int", conv.id),
persist_messages=outcome.persist_messages, persist_messages=outcome.persist_messages,
new_user_text=parsed.turns[-1].text, new_user_text=parsed.turns[-1].text,
capture=capture, capture=capture,
@@ -704,7 +698,7 @@ class MarkdownFrontend(Frontend):
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
_log.exception("turn_log_handler raised; continuing") _log.exception("turn_log_handler raised; continuing")
yield _sse_pack( yield sse_pack(
"done", "done",
{ {
"status": "ok", "status": "ok",
@@ -795,71 +789,53 @@ class MarkdownFrontend(Frontend):
rendered = renderer.render_assistant_message(message) rendered = renderer.render_assistant_message(message)
new_body = renderer.append_to_body(parsed.body, rendered) new_body = renderer.append_to_body(parsed.body, rendered)
new_body = renderer.append_to_body(new_body, renderer.USER_SCAFFOLD) new_body = renderer.append_to_body(new_body, renderer.USER_SCAFFOLD)
# Recompute fingerprint so a future cross-frontend hit on this
# same conversation can find it. Stored as hex string in
# frontmatter — only the markdown frontend reads it.
assistant_param: MessageParam = {
"role": "assistant",
"content": _flatten_assistant_text(message),
}
updated_messages: list[MessageParam] = [*parsed.messages, assistant_param]
updated_metadata = dict(parsed.metadata) updated_metadata = dict(parsed.metadata)
updated_metadata.pop("fingerprint", None)
updated_metadata["agent"] = agent_name updated_metadata["agent"] = agent_name
updated_metadata["conversation_id"] = conv_external_id updated_metadata["conversation_id"] = conv_external_id
updated_metadata["fingerprint"] = fingerprint_messages(updated_messages)
new_content = _reattach_frontmatter(updated_metadata, new_body) new_content = _reattach_frontmatter(updated_metadata, new_body)
if write_disk: if write_disk:
await _write_atomic(file_path, new_content) await _write_atomic(file_path, new_content)
return new_content return new_content
async def _resolve_conversation( async def _resolve_conversation(
self, *, runtime: GatewayRuntime, metadata: dict[str, Any], agent_name: str self,
) -> tuple[Any, str, list[dict[str, Any]]]: *,
"""Resolve the conversation row + stored messages for this request. runtime: GatewayRuntime,
metadata: dict[str, Any],
agent_name: str,
file_path: Path,
) -> tuple[Conversation, str, list[dict[str, Any]]]:
"""Resolve the ``deep`` conversation for this file + its stored messages.
Looks up by frontmatter ``conversation_id``, mints a new row if Frontmatter ``conversation_id`` wins; a file that lost it is found
missing, and returns ``(conv, external_id, stored_messages)``. by its visible ``(markdown, path)`` binding; otherwise a new
``conv.id`` is guaranteed non-None because both conversation is created. The binding follows the file: a moved
``load_conversation`` (after refresh on a committed row) and chat re-binds to its new path on the next turn.
``mint_conversation`` (post-commit refresh) populate it. We
coerce with a runtime check so the rest of the handler can
treat it as ``int``.
""" """
conversations = runtime.conversations
rel = file_path.relative_to(self.vault_path).as_posix()
raw = metadata.get("conversation_id") raw = metadata.get("conversation_id")
lookup_id = raw if isinstance(raw, str) and raw else None conv = await conversations.get(raw) if isinstance(raw, str) and raw else None
if conv is None:
conv = await conversations.find_bound(frontend=FRONTEND, external_id=rel)
if conv is None:
conv = await conversations.create(
kind="deep", agent=agent_name, origin=FRONTEND, title=file_path.stem
)
_log.info("minted conversation %s for %s", conv.external_id, rel)
bound = [
b
for b in await conversations.bindings(conv)
if b.frontend == FRONTEND and b.visible and b.external_id == rel
]
if not bound:
await conversations.bind(conv, frontend=FRONTEND, external_id=rel)
await conversations.touch_user(conv)
if conv.id is None:
msg = "conversation row missing primary key after commit"
raise RuntimeError(msg)
async with runtime.db.session() as session: async with runtime.db.session() as session:
conv = None
if lookup_id is not None:
conv = await load_conversation(
session, frontend="markdown", external_id=lookup_id
)
if conv is None:
_log.info(
"_resolve_conversation: frontmatter conv_id=%s "
"not found in DB, will mint new",
lookup_id,
)
else:
_log.info(
"_resolve_conversation: LOADED existing conv "
"id=%d external_id=%s",
conv.id or -1,
conv.external_id,
)
if conv is None:
conv = await mint_conversation(
session, frontend="markdown", agent_name=agent_name
)
_log.info(
"_resolve_conversation: MINTED new conv "
"id=%d external_id=%s agent=%s",
conv.id or -1,
conv.external_id,
agent_name,
)
if conv.id is None:
msg = "conversation row missing primary key after commit"
raise RuntimeError(msg)
stored = await load_messages(session, conversation_id=conv.id) stored = await load_messages(session, conversation_id=conv.id)
return conv, conv.external_id, stored return conv, conv.external_id, stored
@@ -894,12 +870,6 @@ class MarkdownFrontend(Frontend):
await rewrite_messages( await rewrite_messages(
session, conversation_id=conversation_id, messages=canonical session, conversation_id=conversation_id, messages=canonical
) )
if capture.session_id is not None:
await set_session_id(
session,
conversation_id=conversation_id,
session_id=capture.session_id,
)
_log.info( _log.info(
"_persist_canonical_history: conv_id=%d DB committed", conversation_id "_persist_canonical_history: conv_id=%d DB committed", conversation_id
) )
@@ -927,71 +897,6 @@ class MarkdownFrontend(Frontend):
# ---- module-level utilities ---------------------------------------------- # ---- module-level utilities ----------------------------------------------
def _session_options(conv: Any, divergence_index: int | None) -> dict[str, Any]:
"""Backend options that pin the turn to the conversation's live session.
A divergence means the file's history no longer matches what the
session saw, so the stored ``session_id`` is withheld and the backend
seeds a fresh one from the aligned messages.
"""
return {
"conversation_id": conv.external_id,
"session_id": conv.session_id if divergence_index is None else None,
}
async def _events_with_heartbeat(
events: AsyncIterator[Any], interval: float = _SSE_HEARTBEAT_INTERVAL
) -> AsyncIterator[Any]:
"""Wrap an async event stream with idle-time heartbeat markers.
Yields ``None`` every ``interval`` seconds during silence; real
events pass through unchanged. When the wrapped iterator is
exhausted, this generator returns. Cancellation propagates: if the
outer scope is cancelled we cancel the pending ``__anext__`` task
instead of leaving it dangling.
"""
src = events.__aiter__()
next_task: asyncio.Task[Any] | None = None
try:
while True:
# Reuse the in-flight task across timeouts. Spawning a fresh
# ``__anext__()`` while the previous one is still pending
# puts two consumers on the same async generator — that
# raises ``RuntimeError: anext(): asynchronous generator is
# already running``.
if next_task is None:
next_task = asyncio.ensure_future(src.__anext__())
done, _pending = await asyncio.wait({next_task}, timeout=interval)
if not done:
yield None
continue
task = next_task
next_task = None
try:
result = task.result()
except StopAsyncIteration:
return
yield result
finally:
if next_task is not None and not next_task.done():
next_task.cancel()
with contextlib.suppress(BaseException):
await next_task
def _sse_pack(event: str, data: dict[str, Any]) -> bytes:
r"""Format one Server-Sent Event frame.
Uses named events (``event: <name>``) so the plugin can dispatch on
type without parsing JSON discriminators. ``ensure_ascii=False`` so
multibyte content rides through verbatim instead of becoming
``\uXXXX`` blobs that bloat the wire.
"""
body = json.dumps(data, ensure_ascii=False)
return f"event: {event}\ndata: {body}\n\n".encode()
async def _read_or_empty(path: Path) -> str: async def _read_or_empty(path: Path) -> str:
"""Return file contents, or empty string if the file doesn't exist.""" """Return file contents, or empty string if the file doesn't exist."""
# ``path.exists()`` here is a metadata stat — microseconds — and # ``path.exists()`` here is a metadata stat — microseconds — and
@@ -1083,21 +988,6 @@ def _fallback_synthesized(message: Any) -> list[dict[str, Any]]:
return [{"role": "assistant", "content": content}] return [{"role": "assistant", "content": content}]
def _flatten_assistant_text(message: Any) -> str:
"""Pull all text blocks from an assistant ``Message`` and join them.
Used when we need the assistant content as a plain string for
fingerprinting / equality with a parser-shaped history (parser
already drops thinking + tool_use from assistant turns).
"""
chunks = [
getattr(block, "text", "") or ""
for block in getattr(message, "content", ())
if getattr(block, "type", None) == "text"
]
return "\n\n".join(c for c in chunks if c)
def _render_error_block(exc: BaseException) -> str: def _render_error_block(exc: BaseException) -> str:
"""Render a backend failure as an Assistant turn with a ``[!error]-`` callout.""" """Render a backend failure as an Assistant turn with a ``[!error]-`` callout."""
msg = str(exc) or exc.__class__.__name__ msg = str(exc) or exc.__class__.__name__
@@ -9,11 +9,10 @@ from other frontends).
from __future__ import annotations from __future__ import annotations
import json
import re import re
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
from anthropic.types import Message, TextBlock, ThinkingBlock, ToolUseBlock from anthropic.types import Message, TextBlock, ThinkingBlock
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Iterable from collections.abc import Iterable
@@ -61,9 +60,8 @@ def render_assistant_message(message: Message) -> str:
* ``ThinkingBlock`` → ``> [!thinking]-`` collapsed callout * ``ThinkingBlock`` → ``> [!thinking]-`` collapsed callout
* ``TextBlock`` → plain text (the spoken answer) * ``TextBlock`` → plain text (the spoken answer)
* ``ToolUseBlock`` → ``> [!tool]- <name>`` callout with the ``input`` * ``ToolUseBlock`` → nothing (§3.10: tool calls never reach the file;
JSON quoted inside. Tool *results* are not persisted — see "what the agent is doing" is the activity panel fed by SSE)
module docstring on ``parser.py`` for why.
Blank lines separate adjacent blocks; trailing newline guarantees Blank lines separate adjacent blocks; trailing newline guarantees
the next ``---`` / ``### User:`` marker lands on its own line. the next ``---`` / ``### User:`` marker lands on its own line.
@@ -147,10 +145,7 @@ def _render_block(block: object) -> Iterable[str]:
if isinstance(block, ThinkingBlock): if isinstance(block, ThinkingBlock):
yield from _render_thinking(block.thinking or "") yield from _render_thinking(block.thinking or "")
return return
if isinstance(block, ToolUseBlock): # Tool-use blocks and unknown block types never reach the file.
yield from _render_tool_use(block)
return
# Unknown block type — skip silently rather than corrupting the file.
def _render_thinking(text: str) -> Iterable[str]: def _render_thinking(text: str) -> Iterable[str]:
@@ -159,17 +154,6 @@ def _render_thinking(text: str) -> Iterable[str]:
yield f"> {line}" if line else ">" yield f"> {line}" if line else ">"
def _render_tool_use(block: ToolUseBlock) -> Iterable[str]:
title = summarize_tool_input(block.name, block.input)
yield f"> [!tool]- {title}"
yield "> **input:**"
yield "> ```json"
pretty = json.dumps(block.input, indent=2, ensure_ascii=False, sort_keys=True)
for line in pretty.splitlines():
yield f"> {line}" if line else ">"
yield "> ```"
def adaptive_fence(content: str) -> str: def adaptive_fence(content: str) -> str:
"""Return a backtick fence at least one longer than the longest run in ``content``. """Return a backtick fence at least one longer than the longest run in ``content``.
+14 -1
View File
@@ -16,13 +16,26 @@ from beaver_gateway.storage.db import (
revoke_token, revoke_token,
touch_token, touch_token,
) )
from beaver_gateway.storage.models import AuditLog, Token, TranscriptEntry, Usage from beaver_gateway.storage.models import (
AuditLog,
Conversation,
ConversationBinding,
InjectQueueItem,
Schedule,
Token,
TranscriptEntry,
Usage,
)
from beaver_gateway.storage.session_store import PostgresSessionStore from beaver_gateway.storage.session_store import PostgresSessionStore
__all__ = [ __all__ = [
"AuditLog", "AuditLog",
"Conversation",
"ConversationBinding",
"Database", "Database",
"InjectQueueItem",
"PostgresSessionStore", "PostgresSessionStore",
"Schedule",
"Token", "Token",
"TranscriptEntry", "TranscriptEntry",
"Usage", "Usage",
+23 -3
View File
@@ -105,9 +105,29 @@ def _add_missing_columns(conn: Connection) -> None:
if column.name in existing: if column.name in existing:
continue continue
kind = column.type.compile(conn.dialect) kind = column.type.compile(conn.dialect)
conn.execute( ddl = f"ALTER TABLE {table.name} ADD COLUMN {column.name} {kind}"
text(f"ALTER TABLE {table.name} ADD COLUMN {column.name} {kind}") default = _column_default(column)
) if default is not None:
ddl += f" DEFAULT {default}"
conn.execute(text(ddl))
def _column_default(column: Any) -> str | None:
"""Literal for ``ADD COLUMN ... DEFAULT`` so old rows get the model default."""
server_default = getattr(column, "server_default", None)
if server_default is not None:
return str(server_default.arg.text)
default = getattr(column, "default", None)
if default is None or default.is_callable:
return None
value = default.arg
if isinstance(value, bool):
return "TRUE" if value else "FALSE"
if isinstance(value, int | float):
return repr(value)
if isinstance(value, str):
return "'" + value.replace("'", "''") + "'"
return None
# ---- Token CRUD --------------------------------------------------------- # ---- Token CRUD ---------------------------------------------------------
+97 -9
View File
@@ -1,6 +1,6 @@
"""SQLModel tables. """SQLModel tables.
Four tables, all flat, no FK relationships modelled (``actor`` and Flat tables, no FK relationships modelled (``actor`` and
``agent_name`` are stored as strings — joining audit→token by name is ``agent_name`` are stored as strings — joining audit→token by name is
fine at this volume; we'll introduce FKs when the admin UI actually fine at this volume; we'll introduce FKs when the admin UI actually
demands them). demands them).
@@ -74,15 +74,16 @@ class AuditLog(SQLModel, table=True):
class Conversation(SQLModel, table=True): class Conversation(SQLModel, table=True):
"""One chat thread, scoped to a frontend. """One conversation (§3.1): a master thread, a branch, a deep chat or a job.
``external_id`` is the identifier the frontend uses to find this ``external_id`` is the public id (uuid) every frontend, the API and the
thread again on the next request — for the markdown frontend it's a usage table refer to; ``frontend`` names the frontend that created the
uuid we mint and persist into the file's frontmatter, for the row (``markdown``, ``api``, ``mcp``, ``system``). Where a conversation
anthropic frontend it'd be the same metadata.conversation_id the is *visible* lives in :class:`ConversationBinding`.
client passes. Unique per ``(frontend, external_id)`` because two
frontends sharing a uuid is fine; the same frontend reusing one is ``running_turn`` survives a restart: a non-null value at startup means
a bug. the gateway died mid-turn and the transcript needs its open ``tool_use``
closed before the session is resumed (``core/conversations``).
""" """
__tablename__ = "conversations" __tablename__ = "conversations"
@@ -95,8 +96,92 @@ class Conversation(SQLModel, table=True):
external_id: str = Field(index=True) external_id: str = Field(index=True)
agent_name: str = Field(index=True) agent_name: str = Field(index=True)
session_id: str | None = Field(default=None, index=True) session_id: str | None = Field(default=None, index=True)
kind: str = Field(default="deep", index=True)
parent_id: int | None = Field(default=None, index=True)
title: str | None = Field(default=None)
status: str = Field(default="open", index=True)
running_turn: str | None = Field(default=None)
pending_question: bool = Field(default=False)
flags: dict[str, Any] = Field(
default_factory=dict,
sa_column=Column(
JSON().with_variant(JSONB(), "postgresql"),
nullable=False,
server_default=text("'{}'"),
),
)
created_at: datetime = Field(default_factory=_utcnow) created_at: datetime = Field(default_factory=_utcnow)
updated_at: datetime = Field(default_factory=_utcnow) updated_at: datetime = Field(default_factory=_utcnow)
last_user_activity_at: datetime | None = Field(default=None)
last_activity_at: datetime | None = Field(default=None)
class ConversationBinding(SQLModel, table=True):
"""Where a conversation shows up: a Telegram topic id, a vault-relative path.
Invariant (§3.1): at most one *visible* binding per frontend per
conversation - enforced by the partial unique index. The same external
id may point at several conversations over time (a renamed topic gets a
new branch), only one of them visible.
"""
__tablename__ = "conversation_bindings"
__table_args__ = (
Index(
"uq_binding_visible",
"conversation_id",
"frontend",
unique=True,
postgresql_where=text("visible"),
sqlite_where=text("visible"),
),
Index("ix_binding_lookup", "frontend", "external_id"),
)
id: int | None = Field(default=None, primary_key=True)
conversation_id: int = Field(index=True)
frontend: str
external_id: str
visible: bool = Field(default=True)
created_at: datetime = Field(default_factory=_utcnow)
class InjectQueueItem(SQLModel, table=True):
"""Persisted per-conversation queue (§3.4), priorities ``urgent > user > normal``.
``status`` walks ``queued -> running -> done``; a row still ``running``
at startup was cut by a restart and becomes ``interrupted`` - it is
never re-run, the agent gets an inject saying so instead.
"""
__tablename__ = "inject_queue"
id: int | None = Field(default=None, primary_key=True)
conversation_id: int = Field(index=True)
priority: str = Field(index=True)
origin: str = Field(default="system")
text: str
status: str = Field(default="queued", index=True)
turn_id: str | None = Field(default=None)
created_at: datetime = Field(default_factory=_utcnow)
delivered_at: datetime | None = Field(default=None)
class Schedule(SQLModel, table=True):
"""Deferred inject written by the ``schedule`` tool (§3.6).
M1b only records the promise; the executor (pgqueuer, M3) will move
these into its own job table and this one goes away.
"""
__tablename__ = "schedules"
id: int | None = Field(default=None, primary_key=True)
conversation_id: int = Field(index=True)
execute_at: datetime = Field(index=True)
text: str
created_at: datetime = Field(default_factory=_utcnow)
delivered_at: datetime | None = Field(default=None)
class ConversationMessage(SQLModel, table=True): class ConversationMessage(SQLModel, table=True):
@@ -193,7 +278,10 @@ class Usage(SQLModel, table=True):
__all__ = [ __all__ = [
"AuditLog", "AuditLog",
"Conversation", "Conversation",
"ConversationBinding",
"ConversationMessage", "ConversationMessage",
"InjectQueueItem",
"Schedule",
"Token", "Token",
"TranscriptEntry", "TranscriptEntry",
"Usage", "Usage",
+3
View File
@@ -93,6 +93,9 @@ class FakeClient:
async def query(self, prompt: str) -> None: async def query(self, prompt: str) -> None:
self.prompts.append(prompt) self.prompts.append(prompt)
async def interrupt(self) -> None:
self.interrupted = True
async def receive_response(self): async def receive_response(self):
start = {"type": "message_start", "message": {}} start = {"type": "message_start", "message": {}}
yield StreamEvent(uuid="u", session_id="s", event=start) yield StreamEvent(uuid="u", session_id="s", event=start)
+546
View File
@@ -0,0 +1,546 @@
import asyncio
import tempfile
import uuid
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
import pytest
from claude_agent_sdk import (
AssistantMessage,
InMemorySessionStore,
ResultMessage,
StreamEvent,
TextBlock,
project_key_for_directory,
)
from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions
from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend
from beaver_gateway.core.bus import EventBus
from beaver_gateway.core.conversations import Conversations, ConversationTexts, parse_at
from beaver_gateway.core.registry import AgentRegistry
from beaver_gateway.core.sessions import SessionPool
from beaver_gateway.core.transcript import (
build_entries,
close_open_tool_uses,
open_tool_uses,
render_messages,
strip_tool_entries,
window_entries,
)
from beaver_gateway.storage import Database
from beaver_gateway.storage.models import InjectQueueItem
class ScriptedClient:
instances: list["ScriptedClient"] = []
hold: asyncio.Event | None = None
def __init__(self, options: Any) -> None:
self.options = options
self.prompts: list[str] = []
self.session_id = options.resume or str(uuid.uuid4())
self.interrupted = False
self.connected = False
ScriptedClient.instances.append(self)
async def connect(self) -> None:
self.connected = True
async def query(self, prompt: str) -> None:
self.prompts.append(prompt)
async def receive_response(self):
prompt = self.prompts[-1]
yield StreamEvent(
uuid="u", session_id="s", event={"type": "message_start", "message": {}}
)
yield StreamEvent(
uuid="u",
session_id="s",
event={
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
)
yield StreamEvent(
uuid="u",
session_id="s",
event={
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": f"ok:{prompt}"},
},
)
yield StreamEvent(
uuid="u", session_id="s", event={"type": "content_block_stop", "index": 0}
)
yield AssistantMessage(content=[TextBlock(text=f"ok:{prompt}")], model="m")
hold = ScriptedClient.hold
if hold is not None and not self.interrupted:
await hold.wait()
cut = self.interrupted
self.interrupted = False
yield ResultMessage(
subtype="error_during_execution" if cut else "success",
duration_ms=1,
duration_api_ms=1,
is_error=cut,
num_turns=1,
session_id=self.session_id,
stop_reason="end_turn",
total_cost_usd=0.0,
usage={"input_tokens": 1, "output_tokens": 1},
)
async def interrupt(self) -> None:
self.interrupted = True
if ScriptedClient.hold is not None:
ScriptedClient.hold.set()
async def disconnect(self) -> None:
self.connected = False
class World:
def __init__(self, root: Path) -> None:
ScriptedClient.instances.clear()
ScriptedClient.hold = None
self.root = root
self.db = Database(f"sqlite:///{root / 'w.db'}")
self.store = InMemorySessionStore()
self.agent = ClaudeAgent(
name="a",
model="m",
system_prompt="hi",
cwd=root,
gateway_tools=("say",),
options=ClaudeOptions(effort="low"),
)
self.pool = SessionPool(rss_limit=1 << 40, max_live=100)
self.backend = ClaudeSdkBackend(
agent=self.agent,
mcp_internal_urls={},
session_store=self.store,
client_factory=ScriptedClient,
work_dir=root / "work",
pool=self.pool,
)
self.bus = EventBus()
self.conversations = Conversations(
db=self.db,
agents=AgentRegistry([self.agent]),
backends={"a": self.backend},
bus=self.bus,
pool=self.pool,
store=self.store,
texts=ConversationTexts(),
idle_interval=3600,
)
async def setup(self) -> "World":
await self.db.create_all()
return self
def key(self, session_id: str) -> dict[str, str]:
return {
"project_key": project_key_for_directory(str(self.root)),
"session_id": session_id,
}
async def statuses(self, conv) -> list[tuple[str, str]]:
rows = await self.conversations.queue.recent(conv.id)
return [(r.priority, r.status) for r in reversed(rows)]
async def settle(self, conv, expected: int, timeout: float = 5.0) -> None:
deadline = asyncio.get_running_loop().time() + timeout
while asyncio.get_running_loop().time() < deadline:
rows = await self.conversations.queue.recent(conv.id)
if sum(1 for r in rows if r.status == "done") >= expected:
return
await asyncio.sleep(0.02)
msg = f"queue did not settle: {await self.statuses(conv)}"
raise AssertionError(msg)
@pytest.fixture
async def world() -> World:
root = Path(tempfile.mkdtemp(prefix="beaver-conv-"))
w = await World(root).setup()
yield w
await w.conversations.stop()
await w.pool.close_all()
await w.db.dispose()
async def test_two_messages_run_one_at_a_time_in_order(world: World) -> None:
conv = await world.conversations.create(kind="master", agent="a", origin="test")
ScriptedClient.hold = asyncio.Event()
await world.conversations.post(conv, "first")
await asyncio.sleep(0.2)
await world.conversations.post(conv, "second")
await asyncio.sleep(0.2)
assert await world.statuses(conv) == [("user", "running"), ("user", "queued")]
live = world.pool.get(conv.external_id)
assert live is not None and live.busy
assert (await world.conversations.get(conv.external_id)).running_turn is not None
ScriptedClient.hold.set()
await world.settle(conv, 2)
assert len(ScriptedClient.instances) == 1
assert ScriptedClient.instances[0].prompts == ["first", "second"]
row = await world.conversations.get(conv.external_id)
assert row.running_turn is None
assert row.session_id == ScriptedClient.instances[0].session_id
async def test_urgent_interrupts_and_goes_first(world: World) -> None:
conv = await world.conversations.create(kind="master", agent="a", origin="test")
ScriptedClient.hold = asyncio.Event()
await world.conversations.post(conv, "first")
await asyncio.sleep(0.2)
await world.conversations.post(conv, "second")
await world.conversations.inject(conv, "ALERT", urgency="urgent", origin="крон")
await world.settle(conv, 2)
client = ScriptedClient.instances[0]
assert await world.statuses(conv) == [
("user", "interrupted"),
("user", "done"),
("urgent", "done"),
]
assert await world.statuses(conv) == [
("user", "interrupted"),
("user", "done"),
("urgent", "done"),
]
assert client.prompts[0] == "first"
assert client.prompts[1].startswith("[инжект: крон")
assert client.prompts[1].endswith("ALERT")
assert client.prompts[2] == "second"
async def test_normal_injects_ride_with_the_next_user_message(world: World) -> None:
conv = await world.conversations.create(kind="master", agent="a", origin="test")
await world.conversations.inject(
conv, "vault changed", urgency="normal", origin="watch"
)
await asyncio.sleep(0.2)
assert await world.statuses(conv) == [("normal", "queued")]
await world.conversations.post(conv, "hello")
await world.settle(conv, 2)
prompts = ScriptedClient.instances[0].prompts
assert len(prompts) == 1
assert prompts[0].startswith("hello\n\n[инжекты")
assert "- [watch] vault changed" in prompts[0]
async def test_inject_turn_reply_is_not_routed(world: World) -> None:
conv = await world.conversations.create(kind="master", agent="a", origin="test")
seen: list[dict[str, Any]] = []
async def collect() -> None:
async for event in world.bus.stream(conversation_id=conv.external_id):
seen.append(event)
task = asyncio.create_task(collect())
world.conversations._normal_window = 0.05
await world.conversations.inject(conv, "tick", urgency="normal", origin="крон")
await world.settle(conv, 1)
await asyncio.sleep(0.05)
task.cancel()
types = [e["type"] for e in seen]
assert "turn.start" in types and "turn.end" in types
assert "reply" not in types
assert all(e.get("origin") == "inject" for e in seen if e["type"] == "turn.start")
async def test_spawn_seeds_first_user_message(world: World) -> None:
conv = await world.conversations.spawn(
kind="branch", agent="a", seed="brief", text="do X", title="t"
)
await world.settle(conv, 1)
prompt = ScriptedClient.instances[0].prompts[0]
assert prompt.startswith("[сид: brief] branch «t», ")
assert prompt.endswith("\n\ndo X")
assert ScriptedClient.instances[0].options.system_prompt == "hi"
async def test_fork_leaves_original_untouched(world: World) -> None:
sid = str(uuid.uuid4())
history = [
{"role": "user", "content": "one"},
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "t1", "name": "Read", "input": {}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "t1", "content": "x"}],
},
{"role": "assistant", "content": "done one"},
{"role": "user", "content": "two"},
{"role": "assistant", "content": "done two"},
{"role": "user", "content": "three"},
{"role": "assistant", "content": "done three"},
]
entries = build_entries(history, session_id=sid, cwd=str(world.root), model="m")
await world.store.append(world.key(sid), entries)
conv = await world.conversations.create(
kind="branch", agent="a", origin="test", session_id=sid
)
before = [dict(e) for e in await world.store.load(world.key(sid))]
result = await world.conversations.fork(
conv, "summarize", window=2, strip_tools=True
)
after = await world.store.load(world.key(sid))
assert after == before
assert result.text == "ok:summarize"
assert result.conversation.status == "closed"
assert result.conversation.parent_id == conv.id
fork_entries = await world.store.load(world.key(result.conversation.session_id))
assert not {e["uuid"] for e in fork_entries} & {e["uuid"] for e in before}
texts = [e["message"]["content"] for e in fork_entries]
assert texts[0] == "two"
assert len(fork_entries) == 4
assert fork_entries[0]["parentUuid"] is None
assert all(
b["type"] == "text"
for e in fork_entries[1:]
for b in e["message"]["content"]
if isinstance(e["message"]["content"], list)
)
client = ScriptedClient.instances[0]
assert client.options.resume == result.conversation.session_id
assert client.options.mcp_servers == {}
assert world.pool.get(result.conversation.external_id) is None
async def test_copy_seed_forks_parent_with_window(world: World) -> None:
sid = str(uuid.uuid4())
history = [
{"role": "user", "content": f"q{i}"}
if i % 2 == 0
else {"role": "assistant", "content": f"a{i}"}
for i in range(8)
]
await world.store.append(
world.key(sid),
build_entries(history, session_id=sid, cwd=str(world.root), model="m"),
)
parent = await world.conversations.create(
kind="master", agent="a", origin="test", session_id=sid
)
child = await world.conversations.spawn(
kind="branch", agent="a", seed="copy", parent=parent, window=1
)
assert child.session_id and child.session_id != sid
copied = await world.store.load(world.key(child.session_id))
assert [e["message"]["content"] for e in copied] == [
"q6",
[{"type": "text", "text": "a7"}],
]
await world.settle(child, 1)
assert ScriptedClient.instances[0].options.resume == child.session_id
async def test_merge_injects_summary_into_parent(world: World) -> None:
sid = str(uuid.uuid4())
await world.store.append(
world.key(sid),
build_entries(
[{"role": "user", "content": "q"}, {"role": "assistant", "content": "a"}],
session_id=sid,
cwd=str(world.root),
model="m",
),
)
master = await world.conversations.create(kind="master", agent="a", origin="test")
branch = await world.conversations.create(
kind="branch", agent="a", origin="test", parent=master, session_id=sid
)
result = await world.conversations.merge(branch)
assert result.text.startswith("ok:")
assert (await world.conversations.get(branch.external_id)).status == "merged"
assert await world.statuses(master) == [("normal", "queued")]
item = (await world.conversations.queue.recent(master.id))[0]
assert item.origin == "слив" and item.text == result.text
async def test_recover_closes_open_tool_use_and_injects_interrupted(
world: World,
) -> None:
sid = str(uuid.uuid4())
history = [
{"role": "user", "content": "go"},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "t9",
"name": "Bash",
"input": {"command": "sleep"},
}
],
},
]
await world.store.append(
world.key(sid),
build_entries(history, session_id=sid, cwd=str(world.root), model="m"),
)
conv = await world.conversations.create(
kind="master", agent="a", origin="test", session_id=sid
)
async with world.db.session() as session:
row = await session.get(type(conv), conv.id)
row.running_turn = "turn_dead"
session.add(row)
session.add(
InjectQueueItem(
conversation_id=conv.id,
priority="user",
origin="user",
text="go",
status="running",
turn_id="turn_dead",
)
)
await session.commit()
cut = await world.conversations.recover()
assert [c.external_id for c in cut] == [conv.external_id]
entries = await world.store.load(world.key(sid))
assert not open_tool_uses(entries)
tail = entries[-1]
assert tail["type"] == "user"
assert tail["message"]["content"][0] == {
"type": "tool_result",
"tool_use_id": "t9",
"content": "прервано",
"is_error": True,
}
assert tail["parentUuid"] == entries[-2]["uuid"]
assert (await world.conversations.get(conv.external_id)).running_turn is None
assert await world.statuses(conv) == [("user", "interrupted"), ("normal", "queued")]
note = (await world.conversations.queue.recent(conv.id))[0]
assert (
"turn_dead" in note.text
and "оборван" in note.text
and "1 незакрытых" in note.text
)
await asyncio.sleep(0.2)
assert ScriptedClient.instances == []
async def test_read_and_bindings(world: World) -> None:
sid = str(uuid.uuid4())
history = [
{"role": "user", "content": "q1"},
{
"role": "assistant",
"content": [
{"type": "text", "text": "a1"},
{"type": "tool_use", "id": "t", "name": "Read", "input": {}},
],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "t", "content": "x"}],
},
{"role": "assistant", "content": "a1b"},
{"role": "user", "content": "q2"},
{"role": "assistant", "content": "a2"},
]
await world.store.append(
world.key(sid),
build_entries(history, session_id=sid, cwd=str(world.root), model="m"),
)
conv = await world.conversations.create(
kind="deep", agent="a", origin="test", session_id=sid
)
assert (
await world.conversations.read(conv)
== "user:\nq1\n\nassistant:\na1\n\na1b\n(tools: Read)\n\nuser:\nq2\n\nassistant:\na2"
)
assert (
await world.conversations.read(conv, window=1) == "user:\nq2\n\nassistant:\na2"
)
await world.conversations.bind(conv, frontend="markdown", external_id="a.md")
await world.conversations.bind(conv, frontend="markdown", external_id="b.md")
bindings = await world.conversations.bindings(conv)
assert [(b.external_id, b.visible) for b in bindings] == [
("a.md", False),
("b.md", True),
]
found = await world.conversations.find_bound(
frontend="markdown", external_id="b.md"
)
assert found is not None and found.id == conv.id
assert (
await world.conversations.find_bound(frontend="markdown", external_id="a.md")
is None
)
async def test_schedule_rows_and_parse_at(world: World) -> None:
conv = await world.conversations.create(kind="master", agent="a", origin="test")
row = await world.conversations.schedule(conv, "+15m", "push X")
delta = (row.execute_at.replace(tzinfo=UTC) - datetime.now(UTC)).total_seconds()
assert 14 * 60 < delta <= 15 * 60
assert [s.text for s in await world.conversations.schedules(conv)] == ["push X"]
assert parse_at("2026-09-01T10:00:00+02:00") == datetime(
2026, 9, 1, 8, 0, tzinfo=UTC
)
with pytest.raises(ValueError, match="Invalid isoformat"):
parse_at("tomorrow")
def test_transcript_helpers() -> None:
history = [
{"role": "user", "content": "q1"},
{
"role": "assistant",
"content": [
{"type": "thinking", "thinking": "hm", "signature": "s"},
{"type": "tool_use", "id": "t", "name": "Read", "input": {}},
],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "t", "content": "x"}],
},
{"role": "assistant", "content": "a1"},
{"role": "user", "content": "q2"},
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "t2", "name": "Bash", "input": {}}],
},
]
entries = build_entries(history, session_id="s", cwd="/v", model="m")
assert [b["id"] for _, b in open_tool_uses(entries)] == ["t2"]
fixes = close_open_tool_uses(entries)
assert len(fixes) == 1 and fixes[0]["sessionId"] == "s" and fixes[0]["cwd"] == "/v"
assert not open_tool_uses([*entries, *fixes])
stripped = strip_tool_entries(entries)
assert [e["message"]["content"] for e in stripped] == [
"q1",
[{"type": "text", "text": "a1"}],
"q2",
]
assert (
stripped[0]["parentUuid"] is None
and stripped[1]["parentUuid"] == stripped[0]["uuid"]
)
windowed = window_entries(entries, window=1)
assert windowed[0]["message"]["content"] == "q2" and len(windowed) == 2
assert (
render_messages(
[{"role": "user", "content": "q"}, {"role": "assistant", "content": "a"}]
)
== "user:\nq\n\nassistant:\na"
)
+28
View File
@@ -31,3 +31,31 @@ async def test_create_all_adds_missing_columns() -> None:
session.add(conv) session.add(conv)
await session.commit() await session.commit()
await db.dispose() await db.dispose()
async def test_create_all_backfills_defaults_for_old_rows() -> None:
path = Path(tempfile.mkdtemp(prefix="beaver-migrate-")) / "old.db"
raw = sqlite3.connect(path)
raw.execute(
"CREATE TABLE conversations (id INTEGER PRIMARY KEY, frontend VARCHAR NOT NULL, "
"external_id VARCHAR NOT NULL, agent_name VARCHAR NOT NULL, session_id VARCHAR, "
"created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL)"
)
raw.execute(
"INSERT INTO conversations VALUES (1, 'markdown', 'x', 'a', NULL, '2026-01-01', '2026-01-01')"
)
raw.commit()
raw.close()
db = Database(f"sqlite:///{path}")
await db.create_all()
async with db.session() as session:
conv = (await session.exec(select(Conversation))).one()
assert (conv.kind, conv.status, conv.pending_question, conv.flags) == (
"deep",
"open",
False,
{},
)
assert conv.running_turn is None
await db.dispose()
Generated
+30
View File
@@ -273,6 +273,7 @@ dependencies = [
{ name = "greenlet" }, { name = "greenlet" },
{ name = "itsdangerous" }, { name = "itsdangerous" },
{ name = "jinja2" }, { name = "jinja2" },
{ name = "psutil" },
{ name = "psycopg", extra = ["binary"] }, { name = "psycopg", extra = ["binary"] },
{ name = "pydantic" }, { name = "pydantic" },
{ name = "pydantic-settings" }, { name = "pydantic-settings" },
@@ -313,6 +314,7 @@ requires-dist = [
{ name = "greenlet", specifier = ">=3.5.0" }, { name = "greenlet", specifier = ">=3.5.0" },
{ name = "itsdangerous", specifier = ">=2.2.0" }, { name = "itsdangerous", specifier = ">=2.2.0" },
{ name = "jinja2", specifier = ">=3.1.6" }, { name = "jinja2", specifier = ">=3.1.6" },
{ name = "psutil", specifier = ">=7.2.2" },
{ name = "psycopg", extras = ["binary"], specifier = ">=3.3.4" }, { name = "psycopg", extras = ["binary"], specifier = ">=3.3.4" },
{ name = "pydantic", specifier = ">=2.13.4" }, { name = "pydantic", specifier = ">=2.13.4" },
{ name = "pydantic-settings", specifier = ">=2.14.1" }, { name = "pydantic-settings", specifier = ">=2.14.1" },
@@ -1386,6 +1388,34 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/3a/ed/1cdcab6ba3d6ab7feca11fc14f0eeea80755bb53ef4e892079f31b10a25f/propcache-0.5.2-py3-none-any.whl", hash = "sha256:be1ddfcbb376e3de5d2e2db1d58d6d67463e6b4f9f040c000de8e300295465fe", size = 14036, upload-time = "2026-05-08T21:02:10.673Z" }, { url = "https://files.pythonhosted.org/packages/3a/ed/1cdcab6ba3d6ab7feca11fc14f0eeea80755bb53ef4e892079f31b10a25f/propcache-0.5.2-py3-none-any.whl", hash = "sha256:be1ddfcbb376e3de5d2e2db1d58d6d67463e6b4f9f040c000de8e300295465fe", size = 14036, upload-time = "2026-05-08T21:02:10.673Z" },
] ]
[[package]]
name = "psutil"
version = "7.2.2"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/aa/c6/d1ddf4abb55e93cebc4f2ed8b5d6dbad109ecb8d63748dd2b20ab5e57ebe/psutil-7.2.2.tar.gz", hash = "sha256:0746f5f8d406af344fd547f1c8daa5f5c33dbc293bb8d6a16d80b4bb88f59372", size = 493740, upload-time = "2026-01-28T18:14:54.428Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/51/08/510cbdb69c25a96f4ae523f733cdc963ae654904e8db864c07585ef99875/psutil-7.2.2-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:2edccc433cbfa046b980b0df0171cd25bcaeb3a68fe9022db0979e7aa74a826b", size = 130595, upload-time = "2026-01-28T18:14:57.293Z" },
{ url = "https://files.pythonhosted.org/packages/d6/f5/97baea3fe7a5a9af7436301f85490905379b1c6f2dd51fe3ecf24b4c5fbf/psutil-7.2.2-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:e78c8603dcd9a04c7364f1a3e670cea95d51ee865e4efb3556a3a63adef958ea", size = 131082, upload-time = "2026-01-28T18:14:59.732Z" },
{ url = "https://files.pythonhosted.org/packages/37/d6/246513fbf9fa174af531f28412297dd05241d97a75911ac8febefa1a53c6/psutil-7.2.2-cp313-cp313t-manylinux2010_x86_64.manylinux_2_12_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1a571f2330c966c62aeda00dd24620425d4b0cc86881c89861fbc04549e5dc63", size = 181476, upload-time = "2026-01-28T18:15:01.884Z" },
{ url = "https://files.pythonhosted.org/packages/b8/b5/9182c9af3836cca61696dabe4fd1304e17bc56cb62f17439e1154f225dd3/psutil-7.2.2-cp313-cp313t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:917e891983ca3c1887b4ef36447b1e0873e70c933afc831c6b6da078ba474312", size = 184062, upload-time = "2026-01-28T18:15:04.436Z" },
{ url = "https://files.pythonhosted.org/packages/16/ba/0756dca669f5a9300d0cbcbfae9a4c30e446dfc7440ffe43ded5724bfd93/psutil-7.2.2-cp313-cp313t-win_amd64.whl", hash = "sha256:ab486563df44c17f5173621c7b198955bd6b613fb87c71c161f827d3fb149a9b", size = 139893, upload-time = "2026-01-28T18:15:06.378Z" },
{ url = "https://files.pythonhosted.org/packages/1c/61/8fa0e26f33623b49949346de05ec1ddaad02ed8ba64af45f40a147dbfa97/psutil-7.2.2-cp313-cp313t-win_arm64.whl", hash = "sha256:ae0aefdd8796a7737eccea863f80f81e468a1e4cf14d926bd9b6f5f2d5f90ca9", size = 135589, upload-time = "2026-01-28T18:15:08.03Z" },
{ url = "https://files.pythonhosted.org/packages/81/69/ef179ab5ca24f32acc1dac0c247fd6a13b501fd5534dbae0e05a1c48b66d/psutil-7.2.2-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:eed63d3b4d62449571547b60578c5b2c4bcccc5387148db46e0c2313dad0ee00", size = 130664, upload-time = "2026-01-28T18:15:09.469Z" },
{ url = "https://files.pythonhosted.org/packages/7b/64/665248b557a236d3fa9efc378d60d95ef56dd0a490c2cd37dafc7660d4a9/psutil-7.2.2-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:7b6d09433a10592ce39b13d7be5a54fbac1d1228ed29abc880fb23df7cb694c9", size = 131087, upload-time = "2026-01-28T18:15:11.724Z" },
{ url = "https://files.pythonhosted.org/packages/d5/2e/e6782744700d6759ebce3043dcfa661fb61e2fb752b91cdeae9af12c2178/psutil-7.2.2-cp314-cp314t-manylinux2010_x86_64.manylinux_2_12_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1fa4ecf83bcdf6e6c8f4449aff98eefb5d0604bf88cb883d7da3d8d2d909546a", size = 182383, upload-time = "2026-01-28T18:15:13.445Z" },
{ url = "https://files.pythonhosted.org/packages/57/49/0a41cefd10cb7505cdc04dab3eacf24c0c2cb158a998b8c7b1d27ee2c1f5/psutil-7.2.2-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e452c464a02e7dc7822a05d25db4cde564444a67e58539a00f929c51eddda0cf", size = 185210, upload-time = "2026-01-28T18:15:16.002Z" },
{ url = "https://files.pythonhosted.org/packages/dd/2c/ff9bfb544f283ba5f83ba725a3c5fec6d6b10b8f27ac1dc641c473dc390d/psutil-7.2.2-cp314-cp314t-win_amd64.whl", hash = "sha256:c7663d4e37f13e884d13994247449e9f8f574bc4655d509c3b95e9ec9e2b9dc1", size = 141228, upload-time = "2026-01-28T18:15:18.385Z" },
{ url = "https://files.pythonhosted.org/packages/f2/fc/f8d9c31db14fcec13748d373e668bc3bed94d9077dbc17fb0eebc073233c/psutil-7.2.2-cp314-cp314t-win_arm64.whl", hash = "sha256:11fe5a4f613759764e79c65cf11ebdf26e33d6dd34336f8a337aa2996d71c841", size = 136284, upload-time = "2026-01-28T18:15:19.912Z" },
{ url = "https://files.pythonhosted.org/packages/e7/36/5ee6e05c9bd427237b11b3937ad82bb8ad2752d72c6969314590dd0c2f6e/psutil-7.2.2-cp36-abi3-macosx_10_9_x86_64.whl", hash = "sha256:ed0cace939114f62738d808fdcecd4c869222507e266e574799e9c0faa17d486", size = 129090, upload-time = "2026-01-28T18:15:22.168Z" },
{ url = "https://files.pythonhosted.org/packages/80/c4/f5af4c1ca8c1eeb2e92ccca14ce8effdeec651d5ab6053c589b074eda6e1/psutil-7.2.2-cp36-abi3-macosx_11_0_arm64.whl", hash = "sha256:1a7b04c10f32cc88ab39cbf606e117fd74721c831c98a27dc04578deb0c16979", size = 129859, upload-time = "2026-01-28T18:15:23.795Z" },
{ url = "https://files.pythonhosted.org/packages/b5/70/5d8df3b09e25bce090399cf48e452d25c935ab72dad19406c77f4e828045/psutil-7.2.2-cp36-abi3-manylinux2010_x86_64.manylinux_2_12_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:076a2d2f923fd4821644f5ba89f059523da90dc9014e85f8e45a5774ca5bc6f9", size = 155560, upload-time = "2026-01-28T18:15:25.976Z" },
{ url = "https://files.pythonhosted.org/packages/63/65/37648c0c158dc222aba51c089eb3bdfa238e621674dc42d48706e639204f/psutil-7.2.2-cp36-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:b0726cecd84f9474419d67252add4ac0cd9811b04d61123054b9fb6f57df6e9e", size = 156997, upload-time = "2026-01-28T18:15:27.794Z" },
{ url = "https://files.pythonhosted.org/packages/8e/13/125093eadae863ce03c6ffdbae9929430d116a246ef69866dad94da3bfbc/psutil-7.2.2-cp36-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:fd04ef36b4a6d599bbdb225dd1d3f51e00105f6d48a28f006da7f9822f2606d8", size = 148972, upload-time = "2026-01-28T18:15:29.342Z" },
{ url = "https://files.pythonhosted.org/packages/04/78/0acd37ca84ce3ddffaa92ef0f571e073faa6d8ff1f0559ab1272188ea2be/psutil-7.2.2-cp36-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:b58fabe35e80b264a4e3bb23e6b96f9e45a3df7fb7eed419ac0e5947c61e47cc", size = 148266, upload-time = "2026-01-28T18:15:31.597Z" },
{ url = "https://files.pythonhosted.org/packages/b4/90/e2159492b5426be0c1fef7acba807a03511f97c5f86b3caeda6ad92351a7/psutil-7.2.2-cp37-abi3-win_amd64.whl", hash = "sha256:eb7e81434c8d223ec4a219b5fc1c47d0417b12be7ea866e24fb5ad6e84b3d988", size = 137737, upload-time = "2026-01-28T18:15:33.849Z" },
{ url = "https://files.pythonhosted.org/packages/8c/c7/7bb2e321574b10df20cbde462a94e2b71d05f9bbda251ef27d104668306a/psutil-7.2.2-cp37-abi3-win_arm64.whl", hash = "sha256:8c233660f575a5a89e6d4cb65d9f938126312bca76d8fe087b947b3a1aaac9ee", size = 134617, upload-time = "2026-01-28T18:15:36.514Z" },
]
[[package]] [[package]]
name = "psycopg" name = "psycopg"
version = "3.3.4" version = "3.3.4"