refactor: split flat core into capability packages, layer the conversations service, English defaults for every model-facing text

This commit is contained in:
hh
2026-09-02 00:13:20 +02:00
parent b96714338f
commit cae2ed4161
77 changed files with 2987 additions and 2944 deletions
+2 -2
View File
@@ -11,8 +11,8 @@ from pathlib import Path
from beaver_gateway.agents.base import ExposedMcp from beaver_gateway.agents.base import ExposedMcp
from beaver_gateway.agents.claude import ClaudeAgent from beaver_gateway.agents.claude import ClaudeAgent
from beaver_gateway.agents.raycast import RaycastAgent, RemoteTool, UserPreferences from beaver_gateway.agents.raycast import RaycastAgent, RemoteTool, UserPreferences
from beaver_gateway.core.registry import Gateway from beaver_gateway.app import Gateway
from beaver_gateway.core.turn_record import slugify from beaver_gateway.frontends.turn_record import slugify
from beaver_gateway.frontends.admin import AdminFrontend from beaver_gateway.frontends.admin import AdminFrontend
from beaver_gateway.frontends.anthropic import AnthropicMessagesFrontend from beaver_gateway.frontends.anthropic import AnthropicMessagesFrontend
from beaver_gateway.frontends.markdown import MarkdownFrontend from beaver_gateway.frontends.markdown import MarkdownFrontend
+3 -3
View File
@@ -17,9 +17,9 @@ from pathlib import Path # noqa: TC003 - pydantic runtime
from pydantic import BaseModel, ConfigDict, Field, model_validator from pydantic import BaseModel, ConfigDict, Field, model_validator
from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.core.kinds import KINDS, Kind from beaver_gateway.agents.policy import PolicyRule # noqa: TC001 - pydantic runtime
from beaver_gateway.core.policy import PolicyRule # noqa: TC001 - pydantic runtime from beaver_gateway.agents.prompts import PromptSource # noqa: TC001 - pydantic runtime
from beaver_gateway.core.prompt import PromptSource # noqa: TC001 - pydantic runtime from beaver_gateway.conversations.kinds import KINDS, Kind
__all__ = ["ClaudeAgent", "ClaudeOptions", "Prompts", "SkillSets"] __all__ = ["ClaudeAgent", "ClaudeOptions", "Prompts", "SkillSets"]
@@ -21,7 +21,7 @@ if TYPE_CHECKING:
__all__ = ["PromptSource", "assemble"] __all__ = ["PromptSource", "assemble"]
_log = logging.getLogger("beaver_gateway.core.prompt") _log = logging.getLogger("beaver_gateway.agents.prompts")
PromptSource = str | Path | tuple[str, str | Path] PromptSource = str | Path | tuple[str, str | Path]
+440
View File
@@ -0,0 +1,440 @@
"""Build the runtime from a ``Gateway`` and run every part of it until shutdown."""
from __future__ import annotations
import asyncio
import functools
import logging
from contextlib import AsyncExitStack
from typing import TYPE_CHECKING, Any
import psycopg
import uvicorn
from pgqueuer import PsycopgDriver
from raycast_api import Client as RaycastClient
from raycast_api.config import Config as RaycastConfig
from beaver_gateway.agents.claude import ClaudeAgent
from beaver_gateway.agents.raycast import RaycastAgent
from beaver_gateway.backends.claude_sdk import (
ClaudeSdkBackend,
RunnerConfig,
UsageEvent,
)
from beaver_gateway.backends.raycast import RaycastBackend
from beaver_gateway.backends.sessions import SessionPool
from beaver_gateway.conversations.envelope import Envelope
from beaver_gateway.conversations.rotation import Rotation, RotationPolicy
from beaver_gateway.conversations.service import Conversations
from beaver_gateway.conversations.tools import build_tool_server
from beaver_gateway.events.bus import EventBus
from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.frontends.bearer import require_token
from beaver_gateway.frontends.root import build_root_app
from beaver_gateway.jobs.scheduler import Scheduler
from beaver_gateway.mcp.internal_app import build_internal_app
from beaver_gateway.security.auth import TokenStore
from beaver_gateway.storage import (
Database,
PostgresSessionStore,
Usage,
append_audit,
append_usage,
)
if TYPE_CHECKING:
from collections.abc import Iterable, Iterator
from claude_agent_sdk import McpSdkServerConfig
from fastmcp import FastMCP
from fastmcp.tools.base import Tool as FastMCPTool
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.types import ASGIApp
from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.agents.policy import ToolAudit
from beaver_gateway.backends.base import Backend
from beaver_gateway.config import Gateway
from beaver_gateway.mcp.types import McpServerT
from beaver_gateway.settings import Settings
__all__ = ["AgentRegistry", "McpRegistry", "run"]
_log = logging.getLogger("beaver_gateway.app")
class AgentRegistry:
def __init__(self, agents: Iterable[BaseAgent]) -> None:
self._agents: dict[str, BaseAgent] = {}
for a in agents:
if a.name in self._agents:
msg = f"duplicate agent name: {a.name!r}"
raise ValueError(msg)
self._agents[a.name] = a
def __getitem__(self, name: str) -> BaseAgent:
return self._agents[name]
def get(self, name: str) -> BaseAgent | None:
return self._agents.get(name)
def __iter__(self) -> Iterator[BaseAgent]:
return iter(self._agents.values())
def __len__(self) -> int:
return len(self._agents)
def __contains__(self, name: object) -> bool:
return name in self._agents
class McpRegistry:
def __init__(self, mcps: Iterable[McpServerT]) -> None:
self._mcps: dict[str, McpServerT] = {}
for m in mcps:
if m.name in self._mcps:
msg = f"duplicate mcp name: {m.name!r}"
raise ValueError(msg)
self._mcps[m.name] = m
def __getitem__(self, name: str) -> McpServerT:
return self._mcps[name]
def get(self, name: str) -> McpServerT | None:
return self._mcps.get(name)
def __iter__(self) -> Iterator[McpServerT]:
return iter(self._mcps.values())
def __len__(self) -> int:
return len(self._mcps)
def __contains__(self, name: object) -> bool:
return name in self._mcps
async def run(gateway: Gateway, settings: Settings) -> None:
agents = AgentRegistry(gateway.agents)
mcps = McpRegistry(gateway.mcps)
db = Database(settings.database_url)
await db.create_all()
token_store = TokenStore(
db,
bootstrap=TokenStore.parse_bootstrap(settings.bootstrap_tokens),
bootstrap_scopes=TokenStore.parse_bootstrap_scopes(settings.bootstrap_tokens),
)
async with AsyncExitStack() as stack:
stack.push_async_callback(db.dispose)
await token_store.start()
stack.push_async_callback(token_store.stop)
internal_app, internal_urls, mcp_servers = _build_internal_mcp(
gateway.mcps, settings=settings
)
mcp_tools = await _prefetch_mcp_tools(mcp_servers)
pool = SessionPool()
bus = EventBus()
late = _LateConversations()
session_store = PostgresSessionStore(db)
backends = await _build_backends(
settings=settings,
agents=agents,
stack=stack,
db=db,
session_store=session_store,
mcp_internal_urls=internal_urls,
mcp_servers=mcp_servers,
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,
frontends=gateway.frontends,
envelope=Envelope(
watch=gateway.watch, tz=gateway.tz, recall=gateway.recall
),
distiller=gateway.distiller,
user_sink=gateway.user_sink,
)
late.conversations = conversations
scheduler = Scheduler(
conversations=conversations,
jobs=gateway.jobs,
driver=await _pgqueuer_driver(settings.database_url, stack),
budget=gateway.budget,
rotation=Rotation(
conversations, gateway.rotation or RotationPolicy(tz=gateway.tz)
),
tz=gateway.tz,
)
conversations.scheduler = scheduler
runtime = GatewayRuntime(
agents=agents,
mcps=mcps,
backends=backends,
token_store=token_store,
db=db,
mcp_internal_urls=internal_urls,
admin_user=settings.admin_user,
admin_pass=settings.admin_pass,
session_secret=settings.session_secret,
frontends=tuple(gateway.frontends),
conversations=conversations,
bus=bus,
pool=pool,
scheduler=scheduler,
public_url=gateway.public_url.rstrip("/") if gateway.public_url else None,
)
for fe in gateway.frontends:
fe.configure(runtime)
_log.info(
"beaver-gateway: loaded %d agents, %d mcps, %d frontends",
len(agents),
len(mcps),
len(gateway.frontends),
)
if not gateway.frontends:
return
await conversations.start()
stack.push_async_callback(conversations.stop)
await scheduler.start()
stack.push_async_callback(scheduler.stop)
hooks = scheduler.app(
functools.partial(_authorize_hook, runtime=runtime, scope="api")
)
async with asyncio.TaskGroup() as tg:
tg.create_task(pool.reap_loop())
if internal_app is not None:
tg.create_task(_serve_internal_mcp(internal_app, settings=settings))
tg.create_task(_serve_root(gateway, extra={"/hooks": hooks}))
if gateway.watch is not None:
tg.create_task(gateway.watch.run())
for fe in gateway.frontends:
tg.create_task(fe.serve())
async def _authorize_hook(
request: Request, *, runtime: GatewayRuntime, scope: str
) -> str:
return await require_token(request, runtime, scope=scope)
async def _pgqueuer_driver(url: str, stack: AsyncExitStack) -> PsycopgDriver | None:
plain = _plain_postgres_url(url)
if plain is None:
return None
conn = await psycopg.AsyncConnection.connect(plain, autocommit=True)
stack.push_async_callback(conn.close)
return PsycopgDriver(conn)
def _plain_postgres_url(url: str) -> str | None:
for prefix in ("postgresql+psycopg://", "postgresql://", "postgres://"):
if url.startswith(prefix):
return "postgresql://" + url[len(prefix) :]
return None
async def _serve_root(gateway: Gateway, *, extra: dict[str, ASGIApp]) -> None:
app = build_root_app(gateway.frontends, extra=extra)
config = uvicorn.Config(app, host=gateway.host, port=gateway.port, log_level="info")
_log.info(
"gateway on http://%s:%d - %s",
gateway.host,
gateway.port,
", ".join([*(fe.path for fe in gateway.frontends if fe.path), *extra])
or "no http frontends",
)
await uvicorn.Server(config).serve()
def _build_internal_mcp(
mcps: list[McpServerT], *, settings: Settings
) -> tuple[Starlette | None, dict[str, str], dict[str, FastMCP]]:
if not mcps:
return None, {}, {}
return build_internal_app(mcps, host="127.0.0.1", port=settings.internal_mcp_port)
async def _prefetch_mcp_tools(
servers: dict[str, FastMCP],
) -> dict[str, list[FastMCPTool]]:
out: dict[str, list[FastMCPTool]] = {}
for name, server in servers.items():
try:
out[name] = list(await server.list_tools())
except Exception: # noqa: BLE001
_log.exception("failed to list tools for MCP %r, skipping", name)
out[name] = []
return out
async def _serve_internal_mcp(app: Starlette, *, settings: Settings) -> None:
config = uvicorn.Config(
app,
host="127.0.0.1",
port=settings.internal_mcp_port,
log_level="warning",
loop="uvloop",
)
_log.info(
"internal MCP aggregator on http://127.0.0.1:%d/mcp/<name>",
settings.internal_mcp_port,
)
await uvicorn.Server(config).serve()
class _LateConversations:
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 ask(self, key: str, payload: dict[str, Any]) -> str:
if self.conversations is None:
msg = "conversations service is not up yet"
raise RuntimeError(msg)
answer = await self.conversations.ask(key, payload)
return self.conversations.answer_text(answer)
async def _build_backends(
*,
settings: Settings,
agents: AgentRegistry,
stack: AsyncExitStack,
db: Database,
session_store: PostgresSessionStore,
mcp_internal_urls: dict[str, str],
mcp_servers: dict[str, FastMCP],
mcp_tools: dict[str, list[FastMCPTool]],
pool: SessionPool,
late: _LateConversations,
) -> dict[str, Backend]:
backends: dict[str, Backend] = {}
raycast_agents = [a for a in agents if isinstance(a, RaycastAgent)]
if raycast_agents:
client = await _try_open_raycast_client(settings, stack)
if client is not None:
raycast_backend = RaycastBackend(
client, mcp_servers=mcp_servers, mcp_tools=mcp_tools
)
for a in raycast_agents:
backends[a.name] = raycast_backend
runner = RunnerConfig(user=settings.claude_runner_user, home=settings.claude_home)
mcp_tool_names = {
name: [t.name for t in tools] for name, tools in mcp_tools.items()
}
async def record_usage(event: UsageEvent) -> None:
row = Usage(
agent_name=event.agent_name,
conversation_id=event.conversation_id,
session_id=event.session_id,
model=event.model,
effort=event.effort,
input_tokens=event.usage.input_tokens,
output_tokens=event.usage.output_tokens,
cache_read_tokens=event.usage.cache_read_tokens,
cache_creation_tokens=event.usage.cache_creation_tokens,
context_tokens=event.usage.context_tokens,
cost_usd=event.usage.cost_usd,
duration_ms=event.usage.duration_ms,
num_turns=event.usage.num_turns,
model_usage=event.usage.model_usage,
)
try:
async with db.session() as session:
await append_usage(session, row)
except Exception: # noqa: BLE001
_log.exception("usage write failed for %s", event.agent_name)
async def record_tool(event: ToolAudit) -> None:
detail = {
"conversation": event.conversation,
"kind": event.kind,
"tool": event.tool,
"decision": event.decision,
"reason": event.reason,
"brief": event.brief,
}
try:
async with db.session() as session:
await append_audit(
session,
actor=f"agent:{event.agent}",
kind="tool_call",
agent_name=event.agent,
detail=detail,
)
except Exception: # noqa: BLE001
_log.exception("tool audit write failed for %s", event.agent)
for a in agents:
if isinstance(a, ClaudeAgent):
adapter = ClaudeSdkBackend(
agent=a,
mcp_internal_urls=mcp_internal_urls,
session_store=session_store,
mcp_tool_names=mcp_tool_names,
runner=runner,
usage_sink=record_usage,
pool=pool,
tool_server=functools.partial(late.server, names=a.gateway_tools),
asker=late.ask,
audit_sink=record_tool,
)
await stack.enter_async_context(adapter)
backends[a.name] = adapter
return backends
async def _try_open_raycast_client(
settings: Settings, stack: AsyncExitStack
) -> RaycastClient | None:
if not settings.raycast_bearer:
_log.warning(
"RaycastAgent present but RAYCAST_BEARER is unset; those agents 503"
)
return None
if not settings.raycast_device_id:
_log.warning(
"RaycastAgent present but RAYCAST_DEVICE_ID is unset, those agents 503 "
"(generate once with `python -c 'import secrets; "
"print(secrets.token_hex(32))'`)"
)
return None
if not settings.raycast_config_path.exists():
_log.warning(
"RaycastAgent present but %s is missing, those agents 503",
settings.raycast_config_path,
)
return None
config = RaycastConfig.load(settings.raycast_config_path)
client = RaycastClient(
config=config,
bearer_token=settings.raycast_bearer,
device_id=settings.raycast_device_id,
locale=settings.raycast_locale,
)
return await stack.enter_async_context(client)
+1 -1
View File
@@ -1,6 +1,6 @@
"""Backend adapters. """Backend adapters.
Each backend wraps a provider SDK (``raycast-api``, ``claude-agent-sdk``) Each backend wraps a provider SDK (``raycast-api``, ``claude-agent-sdk``)
and yields the unified :class:`~beaver_gateway.core.events.MessageStreamEvent` and yields the unified :class:`~beaver_gateway.events.stream.MessageStreamEvent`
family. The Anthropic-style frontend serialises events straight to SSE. family. The Anthropic-style frontend serialises events straight to SSE.
""" """
+3 -3
View File
@@ -1,7 +1,7 @@
"""Backend protocol. """Backend protocol.
A backend turns an Anthropic-style turn (``messages`` + agent definition) A backend turns an Anthropic-style turn (``messages`` + agent definition)
into a stream of :class:`~beaver_gateway.core.events.MessageStreamEvent` into a stream of :class:`~beaver_gateway.events.stream.MessageStreamEvent`
records. The frontend serializes whatever comes out straight to SSE, so records. The frontend serializes whatever comes out straight to SSE, so
backends are the only place where provider quirks are translated. backends are the only place where provider quirks are translated.
@@ -12,7 +12,7 @@ subclassing - to keep them swappable in tests with bare async generators.
ignored by backends that don't keep state: ``conversation_id`` (stable id ignored by backends that don't keep state: ``conversation_id`` (stable id
the backend may pin a live session to), ``session_id`` (backend session to the backend may pin a live session to), ``session_id`` (backend session to
resume when nothing is live), ``capture`` (a resume when nothing is live), ``capture`` (a
:class:`~beaver_gateway.core.turn_capture.TurnCapture` the backend fills :class:`~beaver_gateway.backends.capture.TurnCapture` the backend fills
after the stream closes). after the stream closes).
""" """
@@ -26,7 +26,7 @@ if TYPE_CHECKING:
from anthropic.types import MessageParam from anthropic.types import MessageParam
from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.core.events import MessageStreamEvent from beaver_gateway.events.stream import MessageStreamEvent
class Backend(Protocol): class Backend(Protocol):
+21 -17
View File
@@ -3,7 +3,7 @@
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; the sessions of every agent live in one shared time; the sessions of every agent live in one shared
:class:`~beaver_gateway.core.sessions.SessionPool` that owns TTL and :class:`~beaver_gateway.backends.sessions.SessionPool` that owns TTL and
memory-pressure eviction. Sessions are keyed by ``conversation_id`` when memory-pressure eviction. Sessions are keyed by ``conversation_id`` when
the caller passes one or by a text-only fingerprint of ``messages[:-1]`` 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
@@ -67,9 +67,18 @@ from claude_agent_sdk import (
project_key_for_directory, project_key_for_directory,
) )
from beaver_gateway.core import policy as policy_mod from beaver_gateway.agents import policy as policy_mod
from beaver_gateway.core import prompt as prompt_assembly from beaver_gateway.agents import prompts as prompt_assembly
from beaver_gateway.core.events import ( from beaver_gateway.backends.capture import TurnCapture, TurnUsage
from beaver_gateway.backends.sessions import Session, SessionClient, SessionPool
from beaver_gateway.backends.transcript import (
build_entries,
close_open_tool_uses,
fingerprint,
text_of,
)
from beaver_gateway.conversations.kinds import as_kind
from beaver_gateway.events.stream import (
StopReason, StopReason,
build_content_block_stop, build_content_block_stop,
build_input_json_delta, build_input_json_delta,
@@ -83,15 +92,6 @@ 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.kinds import as_kind
from beaver_gateway.core.sessions import Session, SessionClient, SessionPool
from beaver_gateway.core.transcript import (
build_entries,
close_open_tool_uses,
fingerprint,
text_of,
)
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
@@ -107,8 +107,8 @@ if TYPE_CHECKING:
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
from beaver_gateway.core.events import MessageStreamEvent from beaver_gateway.conversations.kinds import Kind
from beaver_gateway.core.kinds import Kind from beaver_gateway.events.stream import MessageStreamEvent
_log = logging.getLogger("beaver_gateway.backends.claude_sdk") _log = logging.getLogger("beaver_gateway.backends.claude_sdk")
@@ -281,13 +281,17 @@ class ClaudeSdkBackend:
def live(self, key: str) -> Session | None: def live(self, key: str) -> Session | None:
return self._pool.get(key) return self._pool.get(key)
async def repair_session(self, session_id: str) -> int: async def repair_session(
self, session_id: str, *, text: str = "interrupted"
) -> int:
"""Close ``tool_use`` blocks a crash left without a result; count added.""" """Close ``tool_use`` blocks a crash left without a result; count added."""
key = self._store_key(session_id) key = self._store_key(session_id)
entries = await self._store.load(cast("Any", key)) entries = await self._store.load(cast("Any", key))
if not entries: if not entries:
return 0 return 0
fixes = close_open_tool_uses(cast("list[Mapping[str, Any]]", entries)) fixes = close_open_tool_uses(
cast("list[Mapping[str, Any]]", entries), text=text
)
if fixes: if fixes:
await self._store.append(cast("Any", key), cast("Any", fixes)) await self._store.append(cast("Any", key), cast("Any", fixes))
_log.warning( _log.warning(
+2 -2
View File
@@ -42,7 +42,7 @@ from raycast_api import Message as RaycastMessage
from raycast_api import RemoteTool, Tool, ToolCall from raycast_api import RemoteTool, Tool, ToolCall
from beaver_gateway.agents.raycast import RaycastAgent from beaver_gateway.agents.raycast import RaycastAgent
from beaver_gateway.core.events import ( from beaver_gateway.events.stream import (
StopReason, StopReason,
build_content_block_stop, build_content_block_stop,
build_input_json_delta, build_input_json_delta,
@@ -65,7 +65,7 @@ if TYPE_CHECKING:
from raycast_api import ChatStreamChunk, Client from raycast_api import ChatStreamChunk, Client
from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.core.events import MessageStreamEvent from beaver_gateway.events.stream import MessageStreamEvent
else: else:
from collections.abc import Mapping from collections.abc import Mapping
@@ -28,7 +28,7 @@ if TYPE_CHECKING:
__all__ = ["DEFAULT_TTL", "Session", "SessionClient", "SessionPool", "cgroup_limit"] __all__ = ["DEFAULT_TTL", "Session", "SessionClient", "SessionPool", "cgroup_limit"]
_log = logging.getLogger("beaver_gateway.core.sessions") _log = logging.getLogger("beaver_gateway.backends.sessions")
DEFAULT_TTL: Mapping[str, float | None] = { DEFAULT_TTL: Mapping[str, float | None] = {
"master": None, "master": None,
@@ -296,7 +296,7 @@ def _zero_usage() -> dict[str, Any]:
# ---- repair, windows, projections --------------------------------------- # ---- repair, windows, projections ---------------------------------------
_PROMPT_TYPES = ("user", "assistant") _PROMPT_TYPES = ("user", "assistant")
_INTERRUPTED = "прервано" _INTERRUPTED = "interrupted"
def open_tool_uses( def open_tool_uses(
+13 -505
View File
@@ -1,533 +1,41 @@
"""Process entrypoint. """Process entrypoint: logging, signals, ``.env``, the config, then ``app.run``."""
Phase 1.4 — async ``main``: install uvloop, load the user config, build
registries + per-agent backends (only ``RaycastBackend`` so far), wire
each frontend with a ``GatewayRuntime``, and run all
``frontend.serve()`` coroutines concurrently. Without any frontends we
still print the Phase 0 DoD line and exit cleanly so the bare skeleton
keeps working.
Phase 2.1 — when the user declares any ``McpServer``, we additionally
build the internal MCP aggregator app and run it on
``127.0.0.1:INTERNAL_MCP_PORT`` as another task inside the same
TaskGroup. URLs are surfaced through ``GatewayRuntime.mcp_internal_urls``
so Phase 2.2's ClaudeCode adapter can find them.
Phase 3 — the same aggregator backs the external ``McpServerFrontend``;
``cli`` doesn't have to know that the frontend reverse-proxies into it,
it just keeps the aggregator running for anyone who needs it.
"""
from __future__ import annotations 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 typing import TYPE_CHECKING, Any
import psycopg
import uvicorn
import uvloop import uvloop
from dotenv import load_dotenv from dotenv import load_dotenv
from pgqueuer import PsycopgDriver
from raycast_api import Client as RaycastClient
from raycast_api.config import Config as RaycastConfig
from beaver_gateway import config_loader from beaver_gateway import app, config
from beaver_gateway.agents.claude import ClaudeAgent from beaver_gateway.security.redact import install as install_redaction
from beaver_gateway.agents.raycast import RaycastAgent from beaver_gateway.security.redact import load_secrets as load_secrets_to_mask
from beaver_gateway.backends.claude_sdk import (
ClaudeSdkBackend,
RunnerConfig,
UsageEvent,
)
from beaver_gateway.backends.raycast import RaycastBackend
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.envelope import Envelope
from beaver_gateway.core.gateway_tools import build_tool_server
from beaver_gateway.core.redact import install as install_redaction
from beaver_gateway.core.redact import load_secrets as load_secrets_to_mask
from beaver_gateway.core.registry import AgentRegistry, Gateway, McpRegistry
from beaver_gateway.core.rotation import Rotation, RotationPolicy
from beaver_gateway.core.scheduler import Scheduler
from beaver_gateway.core.sessions import SessionPool
from beaver_gateway.frontends._auth import require_token
from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.frontends.root import build_root_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_audit,
append_usage,
)
if TYPE_CHECKING:
from claude_agent_sdk import McpSdkServerConfig
from fastmcp import FastMCP
from fastmcp.tools.base import Tool as FastMCPTool
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.types import ASGIApp
from beaver_gateway.backends.base import Backend
from beaver_gateway.core.policy import ToolAudit
from beaver_gateway.mcp.types import McpServerT
_log = logging.getLogger("beaver_gateway.cli")
def main() -> None: def main() -> None:
"""Sync wrapper: uvloop loop factory + asyncio.run."""
logging.basicConfig( logging.basicConfig(
level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s" level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s"
) )
install_redaction() install_redaction()
_install_sigterm_handler() _sigterm_as_interrupt()
asyncio.run(_async_main(), loop_factory=uvloop.new_event_loop) asyncio.run(_run(), loop_factory=uvloop.new_event_loop)
def _install_sigterm_handler() -> None: async def _run() -> None:
"""Turn SIGTERM into a normal interpreter exit. load_dotenv(override=False)
load_secrets_to_mask()
settings = Settings() # ty: ignore[missing-argument]
gateway = config.load(settings.config_path)
await app.run(gateway, settings)
Python's default SIGTERM disposition kills the process outright, so
neither ``AsyncExitStack`` unwinding nor ``atexit`` hooks run. That
matters because every ``claude`` we spawn lives in its own session
(ptyprocess calls ``setsid``), which makes it immune to the signal
that took us down — a hard SIGTERM leaves one orphaned CLI per live
session, each holding hundreds of MB. Raising ``KeyboardInterrupt``
instead routes ``docker stop`` / ``systemctl stop`` through the same
shutdown path as Ctrl-C, which does reap them.
Only installed when we own the main thread's signal handlers; under
an embedding host that isn't ours to take.
"""
def _sigterm_as_interrupt() -> None:
def _raise_interrupt(_signum: int, _frame: object) -> None: def _raise_interrupt(_signum: int, _frame: object) -> None:
raise KeyboardInterrupt raise KeyboardInterrupt
with contextlib.suppress(ValueError, OSError): with contextlib.suppress(ValueError, OSError):
signal.signal(signal.SIGTERM, _raise_interrupt) signal.signal(signal.SIGTERM, _raise_interrupt)
async def _async_main() -> None:
# Populate ``os.environ`` from ``.env`` before anything else so the
# user's ``config.py`` can read its own secrets via ``os.environ[...]``
# (Firefly PAT, third-party MCP creds, etc.). ``Settings`` already
# reads ``.env`` independently via pydantic-settings, but that path
# populates Settings fields, not the process environment.
# ``override=False``: real env vars (Docker, systemd) win over .env.
load_dotenv(override=False)
# Only now does the process environment hold the credentials the
# redactor masks literally (in Docker they arrive via ``env_file``).
load_secrets_to_mask()
settings = Settings() # ty: ignore[missing-argument]
gateway = config_loader.load(settings.config_path)
agents = AgentRegistry(gateway.agents)
mcps = McpRegistry(gateway.mcps)
# Phase 4.1 — open the async DB and run create_all once. Engine
# pool is process-wide; ``dispose()`` after the TaskGroup unwinds.
db = Database(settings.database_url)
await db.create_all()
# Phase 4.2 — TokenStore now reads from the DB (in-memory cache
# primed at start, TTL-refreshed, last_used_at flushed by a
# background task). BOOTSTRAP_TOKENS layers on top so first-run /
# examples still work without DB writes.
token_store = TokenStore(
db,
bootstrap=TokenStore.parse_bootstrap(settings.bootstrap_tokens),
bootstrap_scopes=TokenStore.parse_bootstrap_scopes(settings.bootstrap_tokens),
)
async with AsyncExitStack() as stack:
stack.push_async_callback(db.dispose)
await token_store.start()
stack.push_async_callback(token_store.stop)
# Internal MCP URLs must exist before we construct any
# ClaudeSdkBackend - adapters bake the URLs into their
# ``mcp_servers`` at construction time. The
# ``mcp_servers`` map is used by the Raycast backend, which
# needs in-process ``list_tools`` / ``call_tool`` access (the
# Raycast wire has no native MCP concept).
internal_app, internal_urls, mcp_servers = _build_internal_mcp(
gateway.mcps, settings=settings
)
# Prefetch tool catalogs for every MCP so RaycastAgent requests
# don't pay a per-turn list_tools roundtrip and so a broken MCP
# surfaces at startup instead of mid-conversation.
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(
settings=settings,
agents=agents,
stack=stack,
db=db,
session_store=session_store,
mcp_internal_urls=internal_urls,
mcp_servers=mcp_servers,
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,
frontends=gateway.frontends,
envelope=Envelope(
watch=gateway.watch, tz=gateway.tz, recall=gateway.recall
),
distiller=gateway.distiller,
user_sink=gateway.user_sink,
)
late.conversations = conversations
scheduler = Scheduler(
conversations=conversations,
jobs=gateway.jobs,
driver=await _pgqueuer_driver(settings.database_url, stack),
budget=gateway.budget,
rotation=Rotation(
conversations, gateway.rotation or RotationPolicy(tz=gateway.tz)
),
tz=gateway.tz,
)
conversations.scheduler = scheduler
runtime = GatewayRuntime(
agents=agents,
mcps=mcps,
backends=backends,
token_store=token_store,
db=db,
mcp_internal_urls=internal_urls,
admin_user=settings.admin_user,
admin_pass=settings.admin_pass,
session_secret=settings.session_secret,
frontends=tuple(gateway.frontends),
conversations=conversations,
bus=bus,
pool=pool,
scheduler=scheduler,
public_url=gateway.public_url.rstrip("/") if gateway.public_url else None,
)
for fe in gateway.frontends:
fe.configure(runtime)
_log.info(
"beaver-gateway: loaded %d agents, %d mcps, %d frontends",
len(agents),
len(mcps),
len(gateway.frontends),
)
# Keep the Phase 0 DoD line on stdout for grep-friendly smoke
# tests, in addition to the structured log line above.
print(
f"beaver-gateway: loaded {len(agents)} agents, "
f"{len(mcps)} mcps, {len(gateway.frontends)} frontends"
)
if not gateway.frontends:
# No external listeners → nothing to serve. The internal
# MCP app has no consumer on its own, so we skip running
# it in this path and exit cleanly (Phase 0 DoD).
return
await conversations.start()
stack.push_async_callback(conversations.stop)
await scheduler.start()
stack.push_async_callback(scheduler.stop)
hooks = scheduler.app(
functools.partial(_authorize_hook, runtime=runtime, scope="api")
)
async with asyncio.TaskGroup() as tg:
tg.create_task(pool.reap_loop())
if internal_app is not None:
tg.create_task(_serve_internal_mcp(internal_app, settings=settings))
tg.create_task(_serve_root(gateway, extra={"/hooks": hooks}))
if gateway.watch is not None:
tg.create_task(gateway.watch.run())
for fe in gateway.frontends:
tg.create_task(fe.serve())
async def _authorize_hook(
request: Request, *, runtime: GatewayRuntime, scope: str
) -> str:
return await require_token(request, runtime, scope=scope)
async def _pgqueuer_driver(url: str, stack: AsyncExitStack) -> PsycopgDriver | None:
"""A dedicated autocommit connection for pgqueuer's LISTEN/NOTIFY."""
plain = _plain_postgres_url(url)
if plain is None:
return None
conn = await psycopg.AsyncConnection.connect(plain, autocommit=True)
stack.push_async_callback(conn.close)
return PsycopgDriver(conn)
def _plain_postgres_url(url: str) -> str | None:
for prefix in ("postgresql+psycopg://", "postgresql://", "postgres://"):
if url.startswith(prefix):
return "postgresql://" + url[len(prefix) :]
return None
async def _serve_root(gateway: Gateway, *, extra: dict[str, ASGIApp]) -> None:
app = build_root_app(gateway.frontends, extra=extra)
config = uvicorn.Config(app, host=gateway.host, port=gateway.port, log_level="info")
_log.info(
"gateway on http://%s:%d - %s",
gateway.host,
gateway.port,
", ".join([*(fe.path for fe in gateway.frontends if fe.path), *extra])
or "no http frontends",
)
await uvicorn.Server(config).serve()
def _build_internal_mcp(
mcps: list[McpServerT], *, settings: Settings
) -> tuple[Starlette | None, dict[str, str], dict[str, FastMCP]]:
"""Build the aggregator app + URL map + server map, or empty equivalents.
The URL map is always handed out (frontends may still introspect
``runtime.mcp_internal_urls`` even if nothing is configured); the
app is ``None`` when there are no MCPs to mount, so the caller
skips the uvicorn task entirely. The server map is the in-process
handle the Raycast backend needs to splice MCP tools into its
requests — empty when no MCPs are configured.
"""
if not mcps:
return None, {}, {}
return build_internal_app(mcps, host="127.0.0.1", port=settings.internal_mcp_port)
async def _prefetch_mcp_tools(
servers: dict[str, FastMCP],
) -> dict[str, list[FastMCPTool]]:
"""Eagerly enumerate tools per MCP so the Raycast loop has a static catalog.
Each underlying proxy is allowed to fail independently — a broken
MCP shouldn't take down the whole gateway. The result has one entry
per MCP that responded; agents that ``expose_mcps`` a missing entry
will simply expose no tools from it (logged once per request).
"""
out: dict[str, list[FastMCPTool]] = {}
for name, server in servers.items():
try:
out[name] = list(await server.list_tools())
except Exception: # noqa: BLE001 — proxy can raise any transport error; we degrade per-MCP rather than fail the whole gateway
_log.exception("failed to list tools for MCP %r — skipping", name)
out[name] = []
return out
async def _serve_internal_mcp(app: Starlette, *, settings: Settings) -> None:
"""Run the internal MCP aggregator on loopback.
Bound to ``127.0.0.1`` (never EXPOSE'd) — only the in-process
ClaudeCode subprocess reaches it. Logged at ``warning`` level so
we don't drown the gateway's own logs in per-request noise.
"""
config = uvicorn.Config(
app,
host="127.0.0.1",
port=settings.internal_mcp_port,
log_level="warning",
loop="uvloop",
)
server = uvicorn.Server(config)
_log.info(
"internal MCP aggregator on http://127.0.0.1:%d/mcp/<name>",
settings.internal_mcp_port,
)
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 ask(self, key: str, payload: dict[str, Any]) -> str:
if self.conversations is None:
msg = "conversations service is not up yet"
raise RuntimeError(msg)
answer = await self.conversations.ask(key, payload)
return self.conversations.answer_text(answer)
async def _build_backends(
*,
settings: Settings,
agents: AgentRegistry,
stack: AsyncExitStack,
db: Database,
session_store: PostgresSessionStore,
mcp_internal_urls: dict[str, str],
mcp_servers: dict[str, FastMCP],
mcp_tools: dict[str, list[FastMCPTool]],
pool: SessionPool,
late: _LateConversations,
) -> dict[str, Backend]:
"""Construct one backend per agent name.
The Raycast ``Client`` is shared across every ``RaycastAgent``
(bearer + device-id are process-wide), so we open it lazily - only
when at least one ``RaycastAgent`` is present - and close it via
the caller's exit stack.
Each :class:`ClaudeAgent` gets its own :class:`ClaudeSdkBackend`
(own live-session pool, own prompt and MCP set); all of them share
the session store and the usage sink.
"""
backends: dict[str, Backend] = {}
raycast_agents = [a for a in agents if isinstance(a, RaycastAgent)]
if raycast_agents:
client = await _try_open_raycast_client(settings, stack)
if client is not None:
raycast_backend = RaycastBackend(
client, mcp_servers=mcp_servers, mcp_tools=mcp_tools
)
for a in raycast_agents:
backends[a.name] = raycast_backend
runner = RunnerConfig(user=settings.claude_runner_user, home=settings.claude_home)
mcp_tool_names = {
name: [t.name for t in tools] for name, tools in mcp_tools.items()
}
async def record_usage(event: UsageEvent) -> None:
row = Usage(
agent_name=event.agent_name,
conversation_id=event.conversation_id,
session_id=event.session_id,
model=event.model,
effort=event.effort,
input_tokens=event.usage.input_tokens,
output_tokens=event.usage.output_tokens,
cache_read_tokens=event.usage.cache_read_tokens,
cache_creation_tokens=event.usage.cache_creation_tokens,
context_tokens=event.usage.context_tokens,
cost_usd=event.usage.cost_usd,
duration_ms=event.usage.duration_ms,
num_turns=event.usage.num_turns,
model_usage=event.usage.model_usage,
)
try:
async with db.session() as session:
await append_usage(session, row)
except Exception: # noqa: BLE001
_log.exception("usage write failed for %s", event.agent_name)
async def record_tool(event: ToolAudit) -> None:
detail = {
"conversation": event.conversation,
"kind": event.kind,
"tool": event.tool,
"decision": event.decision,
"reason": event.reason,
"brief": event.brief,
}
try:
async with db.session() as session:
await append_audit(
session,
actor=f"agent:{event.agent}",
kind="tool_call",
agent_name=event.agent,
detail=detail,
)
except Exception: # noqa: BLE001
_log.exception("tool audit write failed for %s", event.agent)
for a in agents:
if isinstance(a, ClaudeAgent):
adapter = ClaudeSdkBackend(
agent=a,
mcp_internal_urls=mcp_internal_urls,
session_store=session_store,
mcp_tool_names=mcp_tool_names,
runner=runner,
usage_sink=record_usage,
pool=pool,
tool_server=functools.partial(late.server, names=a.gateway_tools),
asker=late.ask,
audit_sink=record_tool,
)
await stack.enter_async_context(adapter)
backends[a.name] = adapter
return backends
async def _try_open_raycast_client(
settings: Settings, stack: AsyncExitStack
) -> RaycastClient | None:
"""Open a ``raycast_api.Client`` if creds are available, else warn + skip.
Skipping is gentler than failing startup: the user can bring up the
gateway, list their agents, and still hit a ``ClaudeAgent``; affected
``RaycastAgent`` instances 503 with a specific message at request time.
Missing ``raycast.json`` falls into the same bucket — first-run
users won't have it yet.
"""
if not settings.raycast_bearer:
_log.warning(
"RaycastAgent present but RAYCAST_BEARER is unset — those agents will 503"
)
return None
if not settings.raycast_device_id:
_log.warning(
"RaycastAgent present but RAYCAST_DEVICE_ID is unset — those agents "
"will 503 (generate once with `python -c 'import secrets; "
"print(secrets.token_hex(32))'`)"
)
return None
if not settings.raycast_config_path.exists():
_log.warning(
"RaycastAgent present but %s is missing — those agents will 503",
settings.raycast_config_path,
)
return None
config = RaycastConfig.load(settings.raycast_config_path)
client = RaycastClient(
config=config,
bearer_token=settings.raycast_bearer,
device_id=settings.raycast_device_id,
locale=settings.raycast_locale,
)
return await stack.enter_async_context(client)
@@ -1,20 +1,9 @@
"""Load the user's ``/config/config.py``. """``Gateway`` - the one object a setup's ``config.py`` assembles - and its loader."""
The config file is regular Python. We ``exec`` it in a namespace seeded
with the public surface the user is expected to use (``ClaudeAgent``,
``RaycastAgent``, ``McpServer``, ``ExposedMcp``, ``Gateway``). The file
must assign a top-level ``gateway = Gateway(...)``.
Element-level validation already happens at construction time agents
and ``McpServer`` factories are pydantic models that reject garbage.
What we *can't* validate at construction is the ``Gateway`` container
itself (deliberately a plain dataclass per PRD), so we type-check its
contents here before handing it back.
"""
from __future__ import annotations from __future__ import annotations
import sys import sys
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from croniter import croniter from croniter import croniter
@@ -22,18 +11,58 @@ from croniter import croniter
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.conversations.texts import ConversationTexts
from beaver_gateway.core.registry import Gateway
from beaver_gateway.core.scheduler import Job
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.jobs.scheduler import Job
from beaver_gateway.mcp.types import HttpMcp, McpServer, PythonToolMcp, StdioMcp from beaver_gateway.mcp.types import HttpMcp, McpServer, PythonToolMcp, StdioMcp
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from pathlib import Path from pathlib import Path
from beaver_gateway.conversations.distill import Distiller
from beaver_gateway.conversations.envelope import RecallContext
from beaver_gateway.conversations.rotation import RotationPolicy
from beaver_gateway.conversations.texts import UserSaid
from beaver_gateway.jobs.scheduler import Budget
from beaver_gateway.mcp.types import McpServerT
from beaver_gateway.vault.watch import VaultWatch
__all__ = ["ConfigError", "Gateway", "load"]
@dataclass(slots=True)
class Gateway:
agents: list[BaseAgent] = field(default_factory=list)
mcps: list[McpServerT] = field(default_factory=list)
frontends: list[Frontend] = field(default_factory=list)
texts: ConversationTexts | None = None
"""Every string the gateway says to a model; ``None`` = the English defaults."""
jobs: list[Job] = field(default_factory=list)
"""Cron, webhook and event jobs run by ``jobs.scheduler``."""
rotation: RotationPolicy | None = None
"""When a master is replaced by a fresh one; ``None`` keeps the defaults."""
watch: VaultWatch | None = None
"""Directory watcher feeding the envelope; ``None`` = no change block."""
recall: Callable[[RecallContext], str | None] | None = None
"""Envelope lookup on the user's text: pointers the setup derives from its files."""
user_sink: Callable[[UserSaid], Awaitable[None] | None] | None = None
"""Sees every user message entering a master or branch turn."""
budget: Budget | None = None
"""Subscription utilisation past which non-critical jobs wait."""
distiller: Distiller | None = None
"""Who closes deep chats and where digests and the index live."""
tz: str = "UTC"
"""Zone for the envelope clock, cron expressions and the rotation hour."""
host: str = "0.0.0.0" # noqa: S104
port: int = 8000
"""The one listener; every HTTP frontend is mounted under its ``path``."""
public_url: str | None = None
"""Origin the reverse proxy shows the world; ``None`` derives it per request."""
class ConfigError(Exception): class ConfigError(Exception):
"""User config file is missing, unreadable, or structurally wrong.""" pass
_PUBLIC_NAMES: dict[str, Any] = { _PUBLIC_NAMES: dict[str, Any] = {
@@ -49,7 +78,7 @@ _McpInstance = StdioMcp | HttpMcp | PythonToolMcp
def load(path: Path) -> Gateway: def load(path: Path) -> Gateway:
"""Execute ``path`` and return its top-level ``gateway`` object.""" """Execute a setup's ``config.py`` and return its top-level ``gateway``."""
try: try:
source = path.read_text(encoding="utf-8") source = path.read_text(encoding="utf-8")
except FileNotFoundError as exc: except FileNotFoundError as exc:
@@ -60,12 +89,11 @@ def load(path: Path) -> Gateway:
raise ConfigError(msg) from exc raise ConfigError(msg) from exc
code = compile(source, str(path), "exec") code = compile(source, str(path), "exec")
# Siblings of the config (``policy.py``, ``mcps/``) import by name.
parent = str(path.resolve().parent) parent = str(path.resolve().parent)
if parent not in sys.path: if parent not in sys.path:
sys.path.insert(0, parent) sys.path.insert(0, parent)
namespace: dict[str, Any] = {"__file__": str(path), **_PUBLIC_NAMES} namespace: dict[str, Any] = {"__file__": str(path), **_PUBLIC_NAMES}
exec(code, namespace) # noqa: S102 - exec'ing user config is the feature exec(code, namespace) # noqa: S102
try: try:
gw = namespace["gateway"] gw = namespace["gateway"]
@@ -0,0 +1 @@
"""Conversations: rows, queue, seeds, turns, questions, closing, rotation, envelope."""
+322
View File
@@ -0,0 +1,322 @@
"""Ending conversations: the distiller, the line cap, the master handover."""
from __future__ import annotations
import asyncio
import contextlib
import inspect
import logging
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from typing import TYPE_CHECKING, Any, cast
from beaver_gateway.conversations.distill import (
Digest,
DistillContext,
LineCap,
append_index,
check_digest,
find_digest,
index_line,
trim_summary,
written_paths,
)
from beaver_gateway.conversations.questions import Questions
from beaver_gateway.conversations.state import aware
from beaver_gateway.conversations.texts import NewDayContext
if TYPE_CHECKING:
from beaver_gateway.conversations.rotation import HandoutContext
from beaver_gateway.storage.models import Conversation, InjectQueueItem
__all__ = ["Closing", "DistillResult"]
_log = logging.getLogger(__name__)
CLOSE_WAIT = 0.25
CLOSE_TRIES = 40
CAP_TRIES = 3
@dataclass(frozen=True, slots=True)
class DistillResult:
conversation: Conversation
fork: Conversation
text: str
digest: Digest | None
error: str | None
trimmed: bool
class Closing(Questions):
async def close(self, conv: Conversation) -> Conversation:
row = await self.set_status(conv, "closed")
with contextlib.suppress(LookupError):
await self._backend(conv.agent_name).close(conv.external_id)
return row
async def request_close(self, conv: Conversation) -> Conversation:
"""``close_chat`` from inside a turn: the chat closes once the turn ends."""
if conv.kind != "deep":
msg = f"only deep chats close this way, {conv.external_id} is {conv.kind}"
raise ValueError(msg)
return await self.set_flags(conv, {"close_requested": True})
async def idle(
self,
*,
kind: str,
days: int,
since: datetime | None = None,
limit: int | None = None,
) -> list[Conversation]:
"""Open conversations of ``kind`` with a session, quiet for ``days``."""
now = datetime.now(UTC)
cutoff = now - timedelta(days=days)
out: list[tuple[datetime, Conversation]] = []
for conv in await self.find(status="open", kind=kind, limit=10_000):
if conv.session_id is None:
continue
last = aware(conv.last_activity_at or conv.created_at)
if last > cutoff or (since is not None and last < since):
continue
out.append((last, conv))
out.sort(key=lambda pair: pair[0])
rows = [conv for _, conv in out]
return rows[:limit] if limit is not None else rows
async def distill(
self, conv: Conversation, *, reason: str = "api"
) -> DistillResult:
"""Fork under the distiller: digest checked and indexed, merge to the master."""
if self._distiller is None:
msg = "no distiller configured (Gateway(distiller=...))"
raise RuntimeError(msg)
if conv.kind != "deep":
msg = f"only deep chats are distilled, {conv.external_id} is {conv.kind}"
raise ValueError(msg)
row = await self.get_row(cast("int", conv.id)) or conv
if row.status != "open":
msg = f"conversation {row.external_id} is {row.status}"
raise ValueError(msg)
if await self.busy(row):
msg = f"conversation {row.external_id} is busy"
raise RuntimeError(msg)
memory = bool(row.flags.get("memory", True))
chat_name = await self.chat_name(row)
ctx = DistillContext(
conversation=row,
title=await self.implied_title(row),
source=await self.window_of(row),
chat_name=chat_name,
memory=memory,
reason=reason,
day=datetime.now(UTC).astimezone().date(),
)
prompt = await self._distill_prompt(ctx)
started = datetime.now(UTC)
self._bus.publish(
"distill.start",
conversation_id=row.external_id,
reason=reason,
memory=memory,
)
result = await self.fork(
row,
prompt,
strip_tools=True,
agent=self._distiller.agent,
title=f"digest: {chat_name}",
)
text, trimmed = trim_summary(result.text)
digest: Digest | None = None
error: str | None = None
if memory:
written = written_paths(result.capture.synthesized_messages)
path = find_digest(self._distiller, since=started, written=written)
if path is None:
error = self._texts.digest_missing
else:
checked = check_digest(path, self._distiller)
if isinstance(checked, str):
error = f"{path.name}: {checked}"
else:
digest = checked
append_index(self._distiller, index_line(digest, chat_name))
if error is not None:
_log.warning("distill of %s: %s", row.external_id, error)
master = await self.open_master()
if master is not None and text:
note = self._texts.closed.format(
chat=chat_name,
digest=(
self._texts.closed_digest.format(digest=digest.path.stem)
if digest
else ""
),
text=text,
)
await self.inject(master, note, urgency="normal", origin="digest")
await self.close(row)
row = await self.set_flags(
row,
{
"close_requested": None,
"closed_reason": reason,
"digest": str(digest.path) if digest else None,
"digest_error": error,
},
)
self._bus.publish(
"conversation.distilled",
conversation_id=row.external_id,
fork=result.conversation.external_id,
reason=reason,
memory=memory,
digest=str(digest.path) if digest else None,
error=error,
text=text,
trimmed=trimmed,
master=master.external_id if master is not None else None,
)
return DistillResult(
conversation=row,
fork=result.conversation,
text=text,
digest=digest,
error=error,
trimmed=trimmed,
)
async def _distill_prompt(self, ctx: DistillContext) -> str:
source = self._texts.distill
if source is None:
template = (
self._texts.distill_prompt
if ctx.memory
else self._texts.distill_prompt_no_memory
)
return template.format(
chat=ctx.chat_name, reason=ctx.reason, day=ctx.day.isoformat()
)
produced: Any = source(ctx)
return await produced if inspect.isawaitable(produced) else produced
async def _close_after_turn(self, conv: Conversation) -> None:
for _ in range(CLOSE_TRIES):
if await self.busy(conv):
await asyncio.sleep(CLOSE_WAIT)
continue
try:
await self.distill(conv, reason="close_chat")
except RuntimeError as exc:
_log.info("closing %s: %s, retrying", conv.external_id, exc)
await asyncio.sleep(CLOSE_WAIT)
continue
except Exception: # noqa: BLE001
_log.exception("closing %s after its turn failed", conv.external_id)
return
_log.warning("closing %s: still busy, giving up", conv.external_id)
async def before_turn(self, conv: Conversation) -> str | None:
cap = LineCap.from_flags(conv.flags.get("line_cap"))
if cap is None:
return None
try:
return cap.path.read_text(encoding="utf-8") if cap.path.exists() else ""
except OSError:
_log.exception("line cap: cannot read %s", cap.path)
return None
async def after_turn(self, conv: Conversation, before: str | None) -> None:
row = await self.get_row(cast("int", conv.id))
if row is None:
return
if row.kind == "deep" and row.flags.get("close_requested"):
self._track(asyncio.create_task(self._close_after_turn(row)))
cap = LineCap.from_flags(row.flags.get("line_cap"))
if cap is not None and before is not None:
await self._enforce_cap(row, cap, before)
async def _enforce_cap(self, conv: Conversation, cap: LineCap, before: str) -> None:
if not cap.path.exists():
return
after = cap.path.read_text(encoding="utf-8")
lines = sum(1 for line in after.splitlines() if line.strip())
if lines <= cap.max_lines:
return
if before:
cap.path.write_text(before, encoding="utf-8")
else:
cap.path.unlink()
attempts = int(conv.flags.get("line_cap_attempts", 0) or 0) + 1
await self.set_flags(conv, {"line_cap_attempts": attempts})
self._bus.publish(
"line_cap.bounced",
conversation_id=conv.external_id,
path=str(cap.path),
lines=lines,
max_lines=cap.max_lines,
attempt=attempts,
)
_log.warning(
"line cap: %s came back with %d lines (cap %d), restored; attempt %d",
cap.path,
lines,
cap.max_lines,
attempts,
)
if attempts > CAP_TRIES:
return
await self.inject(
conv,
self._texts.too_long.format(
name=cap.path.name, lines=lines, max_lines=cap.max_lines
),
urgency="urgent",
origin="cap",
interrupt=False,
)
async def handout(self, conv: Conversation, ctx: HandoutContext) -> str:
"""The closing master's last turn."""
source = self._texts.handout
if isinstance(source, str):
prompt = source.format(day=ctx.day.isoformat(), reason=ctx.reason)
else:
produced: Any = source(ctx)
prompt = await produced if inspect.isawaitable(produced) else produced
self._bus.publish(
"handout.start", conversation_id=conv.external_id, day=ctx.day.isoformat()
)
try:
text, _ = await self.run_text_turn(conv, prompt, origin="handout")
except Exception: # noqa: BLE001
_log.exception("handout turn on %s failed", conv.external_id)
text = ""
self._bus.publish(
"handout.end",
conversation_id=conv.external_id,
day=ctx.day.isoformat(),
text=text[:2000],
)
return text
async def new_day(
self, conv: Conversation, *, reason: str = "night", moved: int = 0
) -> InjectQueueItem:
"""The new master's first inject."""
ctx = NewDayContext(
day=datetime.now(UTC).astimezone().date(), reason=reason, moved=moved
)
source = self._texts.new_day
if isinstance(source, str):
text = source.format(day=ctx.day.isoformat(), reason=ctx.reason)
else:
produced: Any = source(ctx)
text = await produced if inspect.isawaitable(produced) else produced
if moved:
text += self._texts.moved_injects.format(moved=moved)
return await self.inject(
conv, text, urgency="urgent", origin="rotation", interrupt=False
)
@@ -49,10 +49,10 @@ class Distiller:
agent: str agent: str
dir: Path dir: Path
index: Path index: Path
type: str = "выжимка" type: str = "digest"
"""Value the ``type`` frontmatter key must carry.""" """Value the ``type`` frontmatter key must carry."""
index_header: str = "# индекс\n\nстрока на выжимку: чат → его выжимка.\n" # noqa: RUF001 index_header: str = "# index\n\none line per digest: chat → its digest.\n"
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -150,13 +150,13 @@ def check_digest(path: Path, digests: Distiller) -> Digest | str:
try: try:
post = frontmatter.load(str(path)) post = frontmatter.load(str(path))
except (OSError, ValueError) as exc: except (OSError, ValueError) as exc:
return f"не читается: {exc}" return f"unreadable: {exc}"
meta = post.metadata meta = post.metadata
if meta.get("type") != digests.type: if meta.get("type") != digests.type:
return f"`type` должен быть `{digests.type}`, не {meta.get('type')!r}" return f"`type` must be `{digests.type}`, not {meta.get('type')!r}"
source = meta.get("source") source = meta.get("source")
if not isinstance(source, str) or not source.strip(): if not isinstance(source, str) or not source.strip():
return "`source` пустой" return "`source` is empty"
when = meta.get("date") when = meta.get("date")
if isinstance(when, datetime): if isinstance(when, datetime):
when = when.date() when = when.date()
@@ -164,11 +164,11 @@ def check_digest(path: Path, digests: Distiller) -> Digest | str:
try: try:
when = date.fromisoformat(when.strip()) when = date.fromisoformat(when.strip())
except ValueError: except ValueError:
return f"`date` не дата: {when!r}" return f"`date` is not a date: {when!r}"
if not isinstance(when, date): if not isinstance(when, date):
return "`date` отсутствует" return "`date` is missing"
if not post.content.strip(): if not post.content.strip():
return "тело пустое" return "empty body"
return Digest(path=path, source=source.strip(), date=when) return Digest(path=path, source=source.strip(), date=when)
@@ -1,39 +1,27 @@
"""The envelope (§3.3): a background block after the user's text. """The envelope: a background block under the user's text - clock, changes, recall."""
Assembled when the turn starts, never when the message is queued: the
time and what changed in the vault since the last envelope (added lines
for the ``full`` files, names and counts for the rest). Ceilings keep it
a signal, not a document; the injects that ride along are bundled below
it by the queue, with their own header.
"""
from __future__ import annotations from __future__ import annotations
import logging import logging
from dataclasses import dataclass from dataclasses import dataclass, field
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from beaver_gateway.conversations.texts import EnvelopeTexts
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable, Sequence from collections.abc import Callable, Sequence
from beaver_gateway.core.watch import Change, VaultWatch from beaver_gateway.vault.watch import Change, VaultWatch
__all__ = ["Envelope", "RecallContext", "render"] __all__ = ["Envelope", "RecallContext", "render"]
_log = logging.getLogger(__name__) _log = logging.getLogger(__name__)
HEADER = (
"[конверт - фоновый сигнал, не обращение; "
"реагируй, только если относится к вопросу]"
)
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class RecallContext: class RecallContext:
"""What the setup's ``recall`` hook sees: the user's text and where it landed."""
text: str text: str
kind: str kind: str
now: datetime now: datetime
@@ -48,10 +36,8 @@ class Envelope:
names_only_within: float = 600.0 names_only_within: float = 600.0
last_at: datetime | None = None last_at: datetime | None = None
recall: Callable[[RecallContext], str | None] | None = None recall: Callable[[RecallContext], str | None] | None = None
"""Setup-side lookup run on the user's text at turn start: pointers into """Setup-side lookup on the user's text; its lines go under the change block."""
the vault (a person's card, the agent's own notes, due dates) that the texts: EnvelopeTexts = field(default_factory=EnvelopeTexts)
gateway cannot know the paths of. Its lines go under the vault block;
a failure is logged and the envelope goes out without them."""
def build( def build(
self, *, now: datetime | None = None, text: str = "", kind: str = "master" self, *, now: datetime | None = None, text: str = "", kind: str = "master"
@@ -70,6 +56,7 @@ class Envelope:
names_only=names_only, names_only=names_only,
max_lines=self.max_lines, max_lines=self.max_lines,
per_file=self.per_file, per_file=self.per_file,
texts=self.texts,
) )
self.last_at = now self.last_at = now
block = self.recall_block(text=text, kind=kind, now=now) block = self.recall_block(text=text, kind=kind, now=now)
@@ -78,9 +65,8 @@ class Envelope:
def recall_only( def recall_only(
self, *, text: str, kind: str, now: datetime | None = None self, *, text: str, kind: str, now: datetime | None = None
) -> str | None: ) -> str | None:
"""The recall lines under the header, without the vault diff (branches)."""
block = self.recall_block(text=text, kind=kind, now=now or datetime.now(UTC)) block = self.recall_block(text=text, kind=kind, now=now or datetime.now(UTC))
return f"{HEADER}\n{block}" if block else None return f"{self.texts.header}\n{block}" if block else None
def recall_block(self, *, text: str, kind: str, now: datetime) -> str | None: def recall_block(self, *, text: str, kind: str, now: datetime) -> str | None:
if self.recall is None or not text.strip(): if self.recall is None or not text.strip():
@@ -102,38 +88,52 @@ def render(
names_only: bool, names_only: bool,
max_lines: int = 120, max_lines: int = 120,
per_file: int = 30, per_file: int = 30,
texts: EnvelopeTexts | None = None,
) -> str: ) -> str:
texts = texts or EnvelopeTexts()
zone = ZoneInfo(tz) zone = ZoneInfo(tz)
stamp = now.astimezone(zone) stamp = now.astimezone(zone)
lines = [HEADER, f"время: {stamp:%Y-%m-%d %H:%M} ({_zone_label(tz)})"] lines = [
texts.header,
texts.time.format(stamp=f"{stamp:%Y-%m-%d %H:%M}", zone=_zone_label(tz)),
]
ordered = sorted(changes, key=lambda c: (not c.full, c.path)) ordered = sorted(changes, key=lambda c: (not c.full, c.path))
since_label = ( since_label = (
f"с {since.astimezone(zone):%H:%M}" if since is not None else "со старта" # noqa: RUF001 texts.since.format(time=f"{since.astimezone(zone):%H:%M}")
if since is not None
else texts.since_start
) )
if ordered: if ordered:
names = ", ".join(f"{c.path} (+{c.added_count})" for c in ordered) names = ", ".join(f"{c.path} (+{c.added_count})" for c in ordered)
lines.append(f"vault, изменено {since_label}: {names}") lines.append(texts.changed.format(since=since_label, names=names))
if not names_only: if not names_only:
_append_diffs(lines, ordered, max_lines=max_lines, per_file=per_file) _append_diffs(
lines, ordered, max_lines=max_lines, per_file=per_file, texts=texts
)
return "\n".join(lines[:max_lines]) return "\n".join(lines[:max_lines])
def _append_diffs( def _append_diffs(
lines: list[str], changes: Sequence[Change], *, max_lines: int, per_file: int lines: list[str],
changes: Sequence[Change],
*,
max_lines: int,
per_file: int,
texts: EnvelopeTexts,
) -> None: ) -> None:
budget = max_lines - len(lines) - 1 budget = max_lines - len(lines) - 1
for change in changes: for change in changes:
if not change.full or not change.added: if not change.full or not change.added:
continue continue
if budget < 3: if budget < 3:
lines.append("… (потолок конверта)") lines.append(texts.truncated)
return return
shown = change.added[: min(per_file, budget - 2)] shown = change.added[: min(per_file, budget - 2)]
lines.append(f"--- {change.path}, только добавленное ---") lines.append(texts.file_header.format(path=change.path))
lines.extend(f"+ {line}" for line in shown) lines.extend(f"+ {line}" for line in shown)
budget -= 1 + len(shown) budget -= 1 + len(shown)
if len(change.added) > len(shown): if len(change.added) > len(shown):
lines.append(f"+ … ещё {len(change.added) - len(shown)}") lines.append(texts.more_lines.format(count=len(change.added) - len(shown)))
budget -= 1 budget -= 1
@@ -42,8 +42,8 @@ URGENCY: tuple[Priority, ...] = ("normal", "wake", "urgent")
"""What the API and the tools accept for ``urgency``: every priority but ``user``.""" """What the API and the tools accept for ``urgency``: every priority but ``user``."""
INTERRUPTED_TURN = ( INTERRUPTED_TURN = (
"[этот инжект прервал предыдущий тёрн: «Request interrupted» выше - " "[this inject cut the previous turn: the «Request interrupted» above is "
"прерывание, не отказ от тулзы]" "an interruption, not a refused tool call]"
) )
@@ -58,9 +58,7 @@ class InjectContext:
def inject_header(ctx: InjectContext) -> str: def inject_header(ctx: InjectContext) -> str:
"""Default framing; a setup overrides it via ``ConversationTexts.inject_header``.""" """Default framing; a setup overrides it via ``ConversationTexts.inject_header``."""
head = ( head = f"[inject: {ctx.origin} - not the user, no reply needed]"
f"[инжект: {ctx.origin} - это не Бобёр, отвечать не нужно, голос не обязателен]"
)
return f"{head}\n{INTERRUPTED_TURN}" if ctx.interrupted_turn else head return f"{head}\n{INTERRUPTED_TURN}" if ctx.interrupted_turn else head
@@ -0,0 +1,116 @@
"""Putting words into a conversation: a message, an inject, ``say``, ``schedule``."""
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Any, cast
from beaver_gateway.conversations.turns import Turns
if TYPE_CHECKING:
from datetime import datetime
from beaver_gateway.conversations.injects import Priority
from beaver_gateway.storage.models import Conversation, InjectQueueItem
__all__ = ["Messaging"]
_log = logging.getLogger(__name__)
class Messaging(Turns):
async def post(
self, conv: Conversation, text: str, *, origin: str = "user"
) -> InjectQueueItem:
item = await self._queue.push(
conversation_id=cast("int", conv.id),
priority="user",
origin=origin,
text=text,
)
await self.touch_user(conv)
self._bus.publish(
"message.queued",
conversation_id=conv.external_id,
item=item.id,
origin=origin,
)
self._ensure_worker(cast("int", conv.id))
return item
async def inject(
self,
conv: Conversation,
text: str,
*,
urgency: Priority = "normal",
origin: str = "system",
interrupt: bool = True,
) -> InjectQueueItem:
if conv.kind == "master" and conv.status != "open":
live = await self.open_master()
if live is None:
_log.error(
"inject (%s) for closed master %s: no open master, it stays there",
origin,
conv.external_id,
)
else:
_log.info(
"inject (%s) for closed master %s goes to %s",
origin,
conv.external_id,
live.external_id,
)
conv = live
item = await self._queue.push(
conversation_id=cast("int", conv.id),
priority=urgency,
origin=origin,
text=text,
)
self._bus.publish(
"inject.queued",
conversation_id=conv.external_id,
item=item.id,
priority=urgency,
origin=origin,
)
if urgency == "urgent" and interrupt:
backend = self._backend(conv.agent_name)
if await backend.interrupt(conv.external_id):
_log.info(
"conversation %s: interrupted for urgent inject", conv.external_id
)
await self._queue.mark_interrupting(item)
self._ensure_worker(cast("int", conv.id))
return item
async def say(self, conv: Conversation, text: str) -> dict[str, Any]:
runner = self._runners.get(cast("int", conv.id))
_log.info("say[%s]: %s", conv.external_id, text[:200])
return self._bus.publish(
"say",
conversation_id=conv.external_id,
text=text,
turn_id=runner.turn_id if runner is not None else None,
)
async def schedule(
self,
conv: Conversation,
at: str,
text: str,
*,
urgency: Priority = "wake",
dedupe_key: str | None = None,
) -> tuple[int | None, datetime]:
if self.scheduler is None:
msg = "no scheduler; `schedule` is unavailable"
raise RuntimeError(msg)
return await self.scheduler.schedule(
conv, at, text, urgency=urgency, dedupe_key=dedupe_key
)
async def schedules(self, conv: Conversation | None = None) -> list[dict[str, Any]]:
return await self.scheduler.scheduled(conv) if self.scheduler else []
@@ -0,0 +1,85 @@
"""``AskUserQuestion`` from inside a turn: answered or timed out."""
from __future__ import annotations
import asyncio
import contextlib
from typing import TYPE_CHECKING, Any, cast
from uuid import uuid4
from beaver_gateway.conversations.spawning import Spawning
from beaver_gateway.conversations.state import Question
if TYPE_CHECKING:
from beaver_gateway.storage.models import Conversation
__all__ = ["Questions"]
class Questions(Spawning):
async def ask(self, key: str, payload: dict[str, Any]) -> str | None:
"""Show the question and wait; ``None`` when nobody answered in time."""
conv = await self.get(key)
if conv is None:
return None
runner = self._runners.get(cast("int", conv.id))
question_id = f"q_{uuid4().hex[:10]}"
pending = Question(
conversation_id=key,
turn_id=runner.turn_id if runner is not None else None,
questions=list(payload.get("questions") or []),
answer=asyncio.get_running_loop().create_future(),
)
self._questions[question_id] = pending
await self._set_pending_question(conv, value=True)
self._bus.publish(
"question",
conversation_id=key,
turn_id=pending.turn_id,
question_id=question_id,
questions=pending.questions,
timeout=self._question_timeout,
)
try:
async with asyncio.timeout(self._question_timeout):
answer = await pending.answer
except TimeoutError:
self._bus.publish(
"question.timeout",
conversation_id=key,
turn_id=pending.turn_id,
question_id=question_id,
)
return None
finally:
self._questions.pop(question_id, None)
await self._set_pending_question(conv, value=False)
self._bus.publish(
"question.answered",
conversation_id=key,
turn_id=pending.turn_id,
question_id=question_id,
answer=answer,
)
return answer
def answer(self, question_id: str, answer: str) -> bool:
pending = self._questions.get(question_id)
if pending is None or pending.answer.done():
return False
pending.answer.set_result(answer)
return True
def answer_text(self, answer: str | None) -> str:
if answer is None:
return self._texts.unanswered.format(
minutes=round(self._question_timeout / 60)
)
return self._texts.answered.format(answer=answer)
async def _set_pending_question(self, conv: Conversation, *, value: bool) -> None:
async def apply(row: Conversation) -> None:
row.pending_question = value
with contextlib.suppress(LookupError):
await self._update(conv, apply)
@@ -17,12 +17,12 @@ from typing import TYPE_CHECKING
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
if TYPE_CHECKING: if TYPE_CHECKING:
from beaver_gateway.core.conversations import Conversations from beaver_gateway.conversations.service import Conversations
from beaver_gateway.storage.models import Conversation from beaver_gateway.storage.models import Conversation
__all__ = ["HandoutContext", "Rotation", "RotationPolicy"] __all__ = ["HandoutContext", "Rotation", "RotationPolicy"]
_log = logging.getLogger("beaver_gateway.core.rotation") _log = logging.getLogger("beaver_gateway.conversations.rotation")
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -46,11 +46,11 @@ class RotationPolicy:
if now.astimezone(zone) < boundary: if now.astimezone(zone) < boundary:
boundary -= timedelta(days=1) boundary -= timedelta(days=1)
if created < boundary and silence > self.night_silence: if created < boundary and silence > self.night_silence:
return "ночь" return "night"
if now - created > self.max_age and silence > self.short_silence: if now - created > self.max_age and silence > self.short_silence:
return "возраст" return "age"
if context_tokens > self.max_context_tokens and silence > self.short_silence: if context_tokens > self.max_context_tokens and silence > self.short_silence:
return "транскрипт" return "context"
return None return None
def day_of(self, master: Conversation) -> date: def day_of(self, master: Conversation) -> date:
+514
View File
@@ -0,0 +1,514 @@
"""The conversation rows: create, find, bind to windows, flags, status, history."""
from __future__ import annotations
import json
import logging
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any, cast
from uuid import uuid4
from sqlmodel import col, select
from beaver_gateway.backends.transcript import (
messages_from_entries,
render_messages,
text_of,
)
from beaver_gateway.conversations.kinds import KINDS, Kind
from beaver_gateway.conversations.state import State, iso
from beaver_gateway.frontends.markdown.history import load_messages
from beaver_gateway.storage.models import (
Conversation,
ConversationBinding,
ConversationMessage,
RateLimit,
Usage,
)
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable, Iterable
from beaver_gateway.frontends.base import Frontend
__all__ = [
"MASTER_ALIAS",
"PARENT_ALIAS",
"STATUSES",
"TITLE_MAX",
"Rows",
"context_of",
"implied_title",
]
_log = logging.getLogger(__name__)
MASTER_ALIAS = "master"
PARENT_ALIAS = "parent"
STATUSES = ("open", "merged", "closed", "archived")
TITLE_MAX = 80
class Rows(State):
async def create(
self,
*,
kind: Kind,
agent: str,
parent: Conversation | None = None,
title: str | None = None,
origin: str = "api",
session_id: str | None = None,
flags: dict[str, Any] | None = None,
) -> Conversation:
if kind not in KINDS:
msg = f"unknown conversation kind {kind!r}"
raise ValueError(msg)
if not self._claude_agent(agent).serves(kind):
msg = f"agent {agent!r} does not serve kind {kind!r}"
raise ValueError(msg)
now = datetime.now(UTC)
row = Conversation(
frontend=origin,
external_id=str(uuid4()),
agent_name=agent,
kind=kind,
parent_id=parent.id if parent is not None else None,
title=title,
session_id=session_id,
flags=dict(flags or {}),
last_activity_at=now,
)
async with self._db.session() as session:
session.add(row)
await session.commit()
await session.refresh(row)
self._bus.publish("conversation.created", **self.public(row))
return row
async def get(self, public_id: str) -> Conversation | None:
async with self._db.session() as session:
result = await session.exec(
select(Conversation).where(Conversation.external_id == public_id)
)
return result.first()
async def get_row(self, row_id: int) -> Conversation | None:
async with self._db.session() as session:
return await session.get(Conversation, row_id)
async def resolve(
self, key: str, *, origin: Conversation | None = None
) -> Conversation | None:
"""By public id, or ``master`` / ``parent`` relative to ``origin``."""
key = key.strip()
if key == MASTER_ALIAS:
return await self.open_master()
if key == PARENT_ALIAS:
if origin is None or origin.parent_id is None:
return None
return await self.get_row(origin.parent_id)
return await self.get(key)
async def open_master(self) -> Conversation | None:
masters = await self.find(kind="master", status="open", limit=1)
return masters[0] if masters else None
async def find(
self,
*,
status: str | None = None,
kind: str | None = None,
parent: Conversation | None = None,
limit: int = 200,
) -> list[Conversation]:
stmt = select(Conversation).order_by(col(Conversation.id).desc()).limit(limit)
if status is not None:
stmt = stmt.where(Conversation.status == status)
if kind is not None:
stmt = stmt.where(Conversation.kind == kind)
if parent is not None:
stmt = stmt.where(Conversation.parent_id == parent.id)
async with self._db.session() as session:
return list((await session.exec(stmt)).all())
async def bindings(self, conv: Conversation) -> list[ConversationBinding]:
async with self._db.session() as session:
result = await session.exec(
select(ConversationBinding)
.where(ConversationBinding.conversation_id == conv.id)
.order_by(col(ConversationBinding.id))
)
return list(result.all())
async def bind(
self,
conv: Conversation,
*,
frontend: str,
external_id: str,
visible: bool = True,
) -> ConversationBinding:
if conv.kind not in self.frontend(frontend).kinds:
msg = f"frontend {frontend!r} does not show kind {conv.kind!r}"
raise ValueError(msg)
async with self._db.session() as session:
existing = list(
(
await session.exec(
select(ConversationBinding).where(
ConversationBinding.conversation_id == conv.id,
ConversationBinding.frontend == frontend,
)
)
).all()
)
row = next((b for b in existing if b.external_id == external_id), None)
if visible:
for other in existing:
if other is not row and other.visible:
other.visible = False
session.add(other)
same_window = await session.exec(
select(ConversationBinding).where(
ConversationBinding.frontend == frontend,
ConversationBinding.external_id == external_id,
ConversationBinding.conversation_id != conv.id,
col(ConversationBinding.visible).is_(True),
)
)
for other in same_window.all():
other.visible = False
session.add(other)
if row is None:
row = ConversationBinding(
conversation_id=cast("int", conv.id),
frontend=frontend,
external_id=external_id,
visible=visible,
)
else:
row.visible = visible
session.add(row)
await session.commit()
await session.refresh(row)
self._bus.publish(
"conversation.bound",
conversation_id=conv.external_id,
frontend=frontend,
external_id=external_id,
visible=visible,
)
return row
async def find_bound(
self, *, frontend: str, external_id: str
) -> Conversation | None:
async with self._db.session() as session:
result = await session.exec(
select(Conversation)
.join(
ConversationBinding,
col(ConversationBinding.conversation_id) == col(Conversation.id),
)
.where(
ConversationBinding.frontend == frontend,
ConversationBinding.external_id == external_id,
col(ConversationBinding.visible).is_(True),
)
.order_by(col(Conversation.id).desc())
)
return result.first()
async def last_binding(
self, *, frontend: str, kind: str
) -> ConversationBinding | None:
"""The window ``frontend`` last used for ``kind``; outlives a rotation."""
async with self._db.session() as session:
result = await session.exec(
select(ConversationBinding)
.join(
Conversation,
col(Conversation.id) == col(ConversationBinding.conversation_id),
)
.where(
ConversationBinding.frontend == frontend, Conversation.kind == kind
)
.order_by(col(ConversationBinding.id).desc())
)
return result.first()
async def window_of(self, conv: Conversation) -> str | None:
for binding in await self.bindings(conv):
if binding.visible:
return binding.external_id
return None
async def set_flags(
self, conv: Conversation, flags: dict[str, Any]
) -> Conversation:
async def apply(row: Conversation) -> None:
row.flags = {**row.flags, **flags}
return await self._update(conv, apply)
async def set_status(self, conv: Conversation, status: str) -> Conversation:
if status not in STATUSES:
msg = f"unknown status {status!r}"
raise ValueError(msg)
async def apply(row: Conversation) -> None:
row.status = status
return await self._update(conv, apply)
async def set_title(self, conv: Conversation, title: str) -> Conversation:
async def apply(row: Conversation) -> None:
row.title = title
return await self._update(conv, apply)
async def touch_user(self, conv: Conversation) -> Conversation:
async def apply(row: Conversation) -> None:
row.last_user_activity_at = datetime.now(UTC)
return await self._update(conv, apply)
async def reparent(self, conv: Conversation, parent: Conversation) -> Conversation:
async def apply(row: Conversation) -> None:
row.parent_id = parent.id
return await self._update(conv, apply)
async def _update(
self, conv: Conversation, apply: Callable[[Conversation], Awaitable[None]]
) -> Conversation:
async with self._db.session() as session:
row = await session.get(Conversation, conv.id)
if row is None:
msg = f"conversation {conv.external_id} vanished"
raise LookupError(msg)
await apply(row)
row.updated_at = datetime.now(UTC)
session.add(row)
await session.commit()
await session.refresh(row)
self._bus.publish("conversation.updated", **self.public(row))
return row
def public(self, conv: Conversation) -> dict[str, Any]:
return {
"id": conv.external_id,
"kind": conv.kind,
"agent": conv.agent_name,
"title": conv.title,
"status": conv.status,
"parent_row": conv.parent_id,
"session_id": conv.session_id,
"running_turn": conv.running_turn,
"pending_question": conv.pending_question,
"flags": conv.flags,
"origin": conv.frontend,
"created_at": iso(conv.created_at),
"last_user_activity_at": iso(conv.last_user_activity_at),
"last_activity_at": iso(conv.last_activity_at),
}
async def describe(self, conv: Conversation) -> dict[str, Any]:
out = self.public(conv)
out["title"] = await self.implied_title(conv)
parent = await self.get_row(conv.parent_id) if conv.parent_id else None
out["parent"] = parent.external_id if parent is not None else None
out["bindings"] = [
{"frontend": b.frontend, "external_id": b.external_id, "visible": b.visible}
for b in await self.bindings(conv)
]
live = self._pool.get(conv.external_id)
out["live"] = live is not None
out["busy"] = live.busy if live is not None else False
runner = self._runners.get(cast("int", conv.id))
out["turn"] = runner.snapshot() if runner is not None else None
pending = self.pending_question(conv.external_id)
out["question"] = (
{"id": pending[0], "questions": pending[1]} if pending else None
)
return out
def pending_question(self, key: str) -> tuple[str, list[dict[str, Any]]] | None:
for question_id, pending in self._questions.items():
if pending.conversation_id == key and not pending.answer.done():
return question_id, pending.questions
return None
async def rate_limits(self, *, limit: int = 100) -> list[RateLimit]:
async with self._db.session() as session:
result = await session.exec(
select(RateLimit).order_by(col(RateLimit.id).desc()).limit(limit)
)
return list(result.all())
async def context_tokens(self, conv: Conversation) -> int:
async with self._db.session() as session:
row = (
await session.exec(
select(Usage)
.where(Usage.conversation_id == conv.external_id)
.order_by(col(Usage.id).desc())
.limit(1)
)
).first()
return context_of(row)
async def usage_tokens(self, since: datetime) -> int:
async with self._db.session() as session:
rows = (
await session.exec(
select(Usage).where(
col(Usage.ts) >= since.astimezone(UTC).replace(tzinfo=None)
)
)
).all()
return sum(
r.input_tokens + r.output_tokens + r.cache_creation_tokens for r in rows
)
async def busy(self, conv: Conversation) -> bool:
row = await self.get_row(cast("int", conv.id)) or conv
if row.running_turn or row.pending_question:
return True
live = self._pool.get(row.external_id)
if live is not None and live.busy:
return True
pending = await self._queue.pending(cast("int", row.id))
return any(i.priority in ("user", "urgent", "wake") for i in pending)
@property
def frontends(self) -> list[Frontend]:
return list(self._frontends)
def frontend(self, name: str) -> Frontend:
for fe in self._frontends:
if fe.name == name:
return fe
msg = f"unknown frontend {name!r}"
raise ValueError(msg)
def default_agent(self, kind: Kind) -> str | None:
for fe in self._frontends:
if kind in fe.kinds and (agent := fe.agent_for(kind)):
return agent
return None
async def materialize(self, conv: Conversation) -> ConversationBinding | None:
for fe in self._frontends:
if conv.kind not in fe.kinds:
continue
binding = await fe.materialize(conv)
if binding is not None:
return binding
return None
async def mark_closed(self, conv: Conversation) -> bool:
marked = False
for fe in self._frontends:
if conv.kind in fe.kinds:
try:
marked = await fe.mark_closed(conv) or marked
except Exception: # noqa: BLE001
_log.exception("%s could not mark %s", fe.name, conv.external_id)
return marked
async def read(self, conv: Conversation, *, window: int | None = None) -> str:
return render_messages(await self.history(conv), window=window)
async def history(self, conv: Conversation) -> list[dict[str, Any]]:
if conv.session_id is None:
async with self._db.session() as session:
return await load_messages(
session, conversation_id=cast("int", conv.id)
)
return messages_from_entries(cast("Any", await self.entries(conv)))
async def entries(self, conv: Conversation, *, subpath: str = "") -> list[Any]:
if conv.session_id is None:
return []
key = {**self._store_key(conv), "subpath": subpath}
return list(await self._store.load(cast("Any", key)) or [])
async def subpaths(self, conv: Conversation) -> list[str]:
if conv.session_id is None:
return []
return list(await self._store.list_subkeys(cast("Any", self._store_key(conv))))
async def first_user_texts(self, ids: Iterable[int]) -> dict[int, str]:
wanted = list(ids)
if not wanted:
return {}
async with self._db.session() as session:
rows = (
await session.exec(
select(ConversationMessage).where(
col(ConversationMessage.conversation_id).in_(wanted),
ConversationMessage.seq == 0,
ConversationMessage.role == "user",
)
)
).all()
return {
r.conversation_id: text_of(json.loads(r.content_json)).strip() for r in rows
}
async def implied_title(self, conv: Conversation) -> str | None:
if conv.title:
return conv.title
text = (await self.first_user_texts([cast("int", conv.id)])).get(
cast("int", conv.id)
)
return implied_title(text)
async def chat_name(self, conv: Conversation) -> str:
"""What a ``[[wikilink]]`` to the chat says: the file's stem when it has one."""
window = await self.window_of(conv)
if window and window.endswith(".md"):
return window.rsplit("/", 1)[-1][: -len(".md")]
return await self.implied_title(conv) or conv.external_id
async def adopt(self, *, kind: Kind, first_user_text: str) -> Conversation | None:
"""The one unbound, session-less conversation whose history starts here."""
text = first_user_text.strip()
if not text:
return None
bound = select(ConversationBinding.conversation_id).where(
col(ConversationBinding.visible).is_(True)
)
async with self._db.session() as session:
rows = (
await session.exec(
select(Conversation).where(
Conversation.kind == kind,
Conversation.status == "open",
col(Conversation.session_id).is_(None),
col(Conversation.id).not_in(bound),
)
)
).all()
firsts = await self.first_user_texts(cast("int", r.id) for r in rows)
hits = [r for r in rows if firsts.get(cast("int", r.id)) == text]
return hits[0] if len(hits) == 1 else None
def implied_title(text: str | None) -> str | None:
if not text:
return None
line = text.strip().splitlines()[0].strip()
return line if len(line) <= TITLE_MAX else line[: TITLE_MAX - 1] + ""
def context_of(row: Usage | None) -> int:
"""The last API call's input, or the per-call average for older rows."""
if row is None:
return 0
if row.context_tokens:
return row.context_tokens
total = row.input_tokens + row.cache_read_tokens + row.cache_creation_tokens
return round(total / max(row.num_turns or 1, 1))
+71
View File
@@ -0,0 +1,71 @@
"""How a new conversation starts: the seed rendered into its first prompt."""
from __future__ import annotations
import inspect
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any, cast
from beaver_gateway.conversations.kinds import as_kind
from beaver_gateway.conversations.rows import Rows
from beaver_gateway.conversations.texts import SeedContext
if TYPE_CHECKING:
from beaver_gateway.storage.models import Conversation
__all__ = ["SEEDS", "Seeds"]
SEEDS = ("clean", "morning", "copy", "brief")
class Seeds(Rows):
async def pending_seed(self, conv: Conversation) -> str | None:
"""A seed nobody has spoken after yet: rendered now, spent once."""
seed = conv.flags.get("seed")
if not seed:
return None
parent = await self.get_row(conv.parent_id) if conv.parent_id else None
ctx = SeedContext(
kind=as_kind(conv.kind),
seed=str(seed),
agent=conv.agent_name,
parent=parent,
text=None,
title=conv.title,
)
window = conv.flags.get("seed_window")
text = await self.seed_text(
ctx, window=window if isinstance(window, int) else None
)
await self.set_flags(conv, {"seed": None, "seed_window": None})
return text
async def seed_text(self, ctx: SeedContext, *, window: int | None) -> str:
texts = self._texts
stamp = datetime.now(UTC).astimezone().strftime("%Y-%m-%d %H:%M")
head = texts.seed_head.format(
seed=ctx.seed,
kind=ctx.kind,
title=f" «{ctx.title}»" if ctx.title else "",
stamp=stamp,
)
body: str | None = None
if texts.seed is not None:
produced: Any = texts.seed(ctx)
if inspect.isawaitable(produced):
produced = await produced
body = cast("str | None", produced)
if body is None:
if ctx.seed == "brief":
body = ctx.text
elif ctx.seed == "copy":
scope = (
texts.seed_copy_window.format(window=window)
if window
else texts.seed_copy_all
)
body = texts.seed_copy.format(scope=scope)
elif ctx.seed == "morning":
body = texts.seed_morning_missing
parts = [head, body, ctx.text if ctx.seed != "brief" else None]
return "\n\n".join(p for p in parts if p)
+158
View File
@@ -0,0 +1,158 @@
"""``Conversations`` - the service every frontend, job and gateway tool talks to.
Built as layers, one file each: rows → seeds → turns → messaging →
spawning → questions → closing; this file adds start, stop and restart
recovery. A turn started by a user message streams back to whoever asked;
a turn started by an inject streams nowhere and can only speak via ``say``.
"""
from __future__ import annotations
import asyncio
import contextlib
import logging
import re
from datetime import UTC, datetime, timedelta, tzinfo
from sqlmodel import col, select
from beaver_gateway.conversations.closing import Closing, DistillResult
from beaver_gateway.conversations.kinds import KINDS
from beaver_gateway.conversations.rows import context_of, implied_title
from beaver_gateway.conversations.seeds import SEEDS
from beaver_gateway.conversations.spawning import ForkResult
from beaver_gateway.conversations.state import aware
from beaver_gateway.conversations.texts import (
ConversationTexts,
NewDayContext,
SeedContext,
UserSaid,
)
from beaver_gateway.storage.models import Conversation
__all__ = [
"KINDS",
"SEEDS",
"ConversationTexts",
"Conversations",
"DistillResult",
"ForkResult",
"NewDayContext",
"SeedContext",
"UserSaid",
"context_of",
"implied_title",
"parse_at",
]
_log = logging.getLogger(__name__)
_RELATIVE = re.compile(r"^\+(\d+)\s*([smhd])$")
_UNITS = {"s": 1, "m": 60, "h": 3600, "d": 86400}
class Conversations(Closing):
async def start(self) -> None:
await self.recover()
for row_id in await self._queue.conversations_with_pending():
self._ensure_worker(row_id)
self._idle_task = asyncio.create_task(self._idle_loop())
async def stop(self) -> None:
tasks = list(self._tasks)
if self._idle_task is not None:
tasks.append(self._idle_task)
for task in tasks:
task.cancel()
for task in tasks:
with contextlib.suppress(BaseException):
await task
self._tasks.clear()
self._idle_task = None
async def recover(self) -> list[Conversation]:
"""Repair the transcripts of turns a restart cut and tell each conversation."""
async with self._db.session() as session:
result = await session.exec(
select(Conversation).where(col(Conversation.running_turn).is_not(None))
)
cut = list(result.all())
for conv in cut:
fixed = 0
if conv.session_id is not None:
backend = self._backend(conv.agent_name)
try:
fixed = await backend.repair_session(
conv.session_id, text=self._texts.interrupted
)
except Exception: # noqa: BLE001
_log.exception("repair of %s failed", conv.session_id)
turn_id = conv.running_turn
async def clear(row: Conversation) -> None:
row.running_turn = None
row.pending_question = False
await self._update(conv, clear)
note = self._texts.cut_by_restart.format(turn_id=turn_id)
if fixed:
note += self._texts.repaired_tools.format(
fixed=fixed, interrupted=self._texts.interrupted
)
await self.inject(conv, note, urgency="normal", origin="system")
_log.warning("conversation %s: %s", conv.external_id, note)
for item in await self._queue.interrupted():
_log.warning(
"queue item #%s (%s) was running at restart; marked interrupted",
item.id,
item.priority,
)
return cut
async def _idle_loop(self) -> None:
while True:
try:
await self._emit_idle()
except Exception: # noqa: BLE001
_log.exception("idle watcher failed")
await asyncio.sleep(self._idle_interval)
async def _emit_idle(self) -> None:
if not self._idle_days:
return
now = datetime.now(UTC)
for conv in await self.find(status="open", limit=10_000):
last = aware(conv.last_activity_at or conv.created_at)
days = int((now - last).total_seconds() // 86400)
due = [d for d in self._idle_days if days >= d]
if not due:
continue
notified = int(conv.flags.get("idle_notified", 0) or 0)
if due[-1] <= notified:
continue
await self.set_flags(conv, {"idle_notified": due[-1]})
bindings = await self.bindings(conv)
self._bus.publish(
"conversation.idle",
conversation_id=conv.external_id,
kind=conv.kind,
agent=conv.agent_name,
days=due[-1],
bindings=[
{"frontend": b.frontend, "external_id": b.external_id}
for b in bindings
if b.visible
],
)
def parse_at(at: str, tz: tzinfo = UTC) -> datetime:
raw = at.strip()
match = _RELATIVE.match(raw.replace(" ", ""))
if match:
amount, unit = match.groups()
return datetime.now(UTC) + timedelta(seconds=int(amount) * _UNITS[unit])
parsed = datetime.fromisoformat(raw)
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=tz)
return parsed.astimezone(UTC)
@@ -0,0 +1,201 @@
"""New conversations from old ones: spawn with a seed, fork a copy, merge back."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, cast
from claude_agent_sdk import fork_session_via_store, project_key_for_directory
from beaver_gateway.backends.transcript import strip_tool_entries, window_entries
from beaver_gateway.conversations.messaging import Messaging
from beaver_gateway.conversations.seeds import SEEDS
from beaver_gateway.conversations.texts import SeedContext
if TYPE_CHECKING:
from beaver_gateway.backends.capture import TurnCapture
from beaver_gateway.conversations.kinds import Kind
from beaver_gateway.storage.models import Conversation
__all__ = ["ForkResult", "Spawning"]
_log = logging.getLogger(__name__)
@dataclass(frozen=True, slots=True)
class ForkResult:
conversation: Conversation
text: str
capture: TurnCapture
class Spawning(Messaging):
async def spawn(
self,
*,
kind: Kind,
agent: str | None = None,
seed: str = "clean",
parent: Conversation | None = None,
text: str | None = None,
title: str | None = None,
window: int | None = None,
origin: str = "api",
binding: tuple[str, str] | None = None,
flags: dict[str, Any] | None = None,
) -> Conversation:
"""Create a conversation in a window and queue its seed.
``binding`` reuses a window that already exists instead of asking the
home frontend for one. Without ``text`` the seed waits in ``flags``
and opens the first turn, so a fresh window costs nothing until
someone speaks.
"""
if seed not in SEEDS:
msg = f"unknown seed {seed!r}"
raise ValueError(msg)
if seed == "brief" and not text:
msg = "seed=brief needs text"
raise ValueError(msg)
if agent is None and kind == "branch" and parent is not None:
agent = parent.agent_name
agent = agent or self.default_agent(kind)
if agent is None:
msg = f"no default agent for kind {kind!r}; pass `agent`"
raise ValueError(msg)
session_id: str | None = None
if seed == "copy":
if parent is None or parent.session_id is None:
msg = "seed=copy needs a parent with a session"
raise ValueError(msg)
session_id = await self._copy_session(
parent, window=window, strip_tools=False
)
conv = await self.create(
kind=kind,
agent=agent,
parent=parent,
title=title,
origin=origin,
session_id=session_id,
flags=flags,
)
if binding is not None:
await self.bind(conv, frontend=binding[0], external_id=binding[1])
else:
await self.materialize(conv)
ctx = SeedContext(
kind=kind, seed=seed, agent=agent, parent=parent, text=text, title=title
)
if text is None:
return await self.set_flags(conv, {"seed": seed, "seed_window": window})
await self._queue.push(
conversation_id=cast("int", conv.id),
priority="user",
origin=f"seed:{seed}" if seed == "brief" else origin,
text=await self.seed_text(ctx, window=window),
)
self._ensure_worker(cast("int", conv.id))
return conv
async def fork(
self,
conv: Conversation,
prompt: str,
*,
window: int | None = None,
strip_tools: bool = False,
title: str | None = None,
agent: str | None = None,
) -> ForkResult:
"""Copy the history into a one-off session, run ``prompt`` on it, close it."""
agent = agent or conv.agent_name
session_id = await self._copy_session(
conv, window=window, strip_tools=strip_tools, agent=agent
)
child = await self.create(
kind="fork",
agent=agent,
parent=conv,
title=title or f"fork: {conv.title or conv.external_id}",
origin="system",
session_id=session_id,
)
try:
text, capture = await self.run_text_turn(
child, prompt, origin="fork", tools=False
)
finally:
await self._backend(agent).close(child.external_id)
child = await self.set_status(child, "closed")
return ForkResult(conversation=child, text=text, capture=capture)
async def merge(self, conv: Conversation) -> ForkResult:
if conv.parent_id is None:
msg = "merge needs a parent conversation"
raise ValueError(msg)
parent = await self.get_row(conv.parent_id)
if parent is None:
msg = "parent conversation vanished"
raise LookupError(msg)
result = await self.fork(
conv,
self._texts.merge_prompt,
title=f"merge: {conv.title or conv.external_id}",
)
if result.text.strip():
await self.inject(parent, result.text, urgency="normal", origin="merge")
await self.set_status(conv, "merged")
await self.mark_closed(conv)
self._bus.publish(
"conversation.merged",
conversation_id=conv.external_id,
parent=parent.external_id,
fork=result.conversation.external_id,
)
return result
async def _copy_session(
self,
conv: Conversation,
*,
window: int | None,
strip_tools: bool,
agent: str | None = None,
) -> str:
if conv.session_id is None:
msg = f"conversation {conv.external_id} has no session to copy"
raise ValueError(msg)
live = self._pool.get(conv.external_id)
if live is not None and live.dirty:
msg = f"conversation {conv.external_id} has a mirror gap; not forking"
raise RuntimeError(msg)
source = self._claude_agent(conv.agent_name)
target = self._claude_agent(agent) if agent else source
forked = await fork_session_via_store(
self._store, conv.session_id, directory=str(source.cwd)
)
source_key = {
"project_key": project_key_for_directory(str(source.cwd)),
"session_id": forked.session_id,
}
target_key = {
"project_key": project_key_for_directory(str(target.cwd)),
"session_id": forked.session_id,
}
if window is not None or strip_tools or target_key != source_key:
entries = await self._store.load(cast("Any", source_key)) or []
trimmed = window_entries(cast("Any", entries), window=window)
if strip_tools:
trimmed = strip_tool_entries(trimmed)
await self._store.delete(cast("Any", source_key))
await self._store.append(cast("Any", target_key), cast("Any", trimmed))
_log.info(
"forked session %s -> %s (window=%s, strip_tools=%s)",
conv.session_id,
forked.session_id,
window,
strip_tools,
)
return forked.session_id
+185
View File
@@ -0,0 +1,185 @@
"""What every layer of the conversations service shares: wiring and lookups."""
from __future__ import annotations
import asyncio
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any, cast
from claude_agent_sdk import project_key_for_directory
from beaver_gateway.conversations.injects import InjectQueue
from beaver_gateway.conversations.texts import ConversationTexts
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable, Sequence
from claude_agent_sdk import SessionStore
from beaver_gateway.agents.claude import ClaudeAgent
from beaver_gateway.app import AgentRegistry
from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend
from beaver_gateway.backends.sessions import SessionPool
from beaver_gateway.conversations.distill import Distiller
from beaver_gateway.conversations.envelope import Envelope
from beaver_gateway.conversations.texts import UserSaid
from beaver_gateway.events.bus import EventBus
from beaver_gateway.frontends.base import Frontend
from beaver_gateway.jobs.scheduler import Scheduler
from beaver_gateway.storage.db import Database
from beaver_gateway.storage.models import Conversation
__all__ = ["Question", "Runner", "State", "aware", "iso"]
@dataclass
class Runner:
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
wake: asyncio.Event = field(default_factory=asyncio.Event)
task: asyncio.Task[None] | None = None
turn_id: str | None = None
origin: str | None = None
text: str | None = None
started_at: datetime | None = None
tools: dict[str, dict[str, Any]] = field(default_factory=dict)
def snapshot(self) -> dict[str, Any] | None:
if self.turn_id is None:
return None
return {
"id": self.turn_id,
"origin": self.origin,
"text": self.text,
"started_at": iso(self.started_at),
"tools": list(self.tools.values()),
}
@dataclass(frozen=True, slots=True)
class Question:
conversation_id: str
turn_id: str | None
questions: list[dict[str, Any]]
answer: asyncio.Future[str]
class State:
_db: Database
_agents: AgentRegistry
_backends: dict[str, Any]
_bus: EventBus
_pool: SessionPool
_store: SessionStore
_texts: ConversationTexts
_frontends: list[Frontend]
_normal_window: float
_idle_days: tuple[int, ...]
_idle_interval: float
_question_timeout: float
_envelope: Envelope | None
_distiller: Distiller | None
_user_sink: Callable[[UserSaid], Awaitable[None] | None] | None
_queue: InjectQueue
_runners: dict[int, Runner]
_questions: dict[str, Question]
_tasks: set[asyncio.Task[None]]
_idle_task: asyncio.Task[None] | None
scheduler: Scheduler | None
def __init__(
self,
*,
db: Database,
agents: AgentRegistry,
backends: dict[str, Any],
bus: EventBus,
pool: SessionPool,
store: SessionStore,
texts: ConversationTexts | None = None,
frontends: Sequence[Frontend] = (),
normal_window: float = 3600.0,
idle_days: Sequence[int] = (2,),
idle_interval: float = 3600.0,
question_timeout: float = 600.0,
envelope: Envelope | None = None,
distiller: Distiller | None = None,
user_sink: Callable[[UserSaid], Awaitable[None] | None] | None = None,
) -> None:
self._db = db
self._agents = agents
self._backends = backends
self._bus = bus
self._pool = pool
self._store = store
self._texts = texts or ConversationTexts()
self._frontends = [f for f in frontends if f.name]
self._normal_window = normal_window
self._idle_days = tuple(sorted(idle_days))
self._idle_interval = idle_interval
self._question_timeout = question_timeout
self._envelope = envelope
self._distiller = distiller
self._user_sink = user_sink
self._queue = InjectQueue(db)
self._runners = {}
self._questions = {}
self._tasks = set()
self._idle_task = None
self.scheduler = None
@property
def db(self) -> Database:
return self._db
@property
def queue(self) -> InjectQueue:
return self._queue
@property
def bus(self) -> EventBus:
return self._bus
@property
def pool(self) -> SessionPool:
return self._pool
def _backend(self, agent: str) -> ClaudeSdkBackend:
backend = self._backends.get(agent)
if backend is None or not hasattr(backend, "repair_session"):
msg = f"agent {agent!r} has no Claude SDK backend"
raise LookupError(msg)
return cast("ClaudeSdkBackend", backend)
def _claude_agent(self, name: str) -> ClaudeAgent:
agent = self._agents.get(name)
if agent is None or not hasattr(agent, "cwd"):
msg = f"unknown Claude agent {name!r}"
raise LookupError(msg)
return cast("ClaudeAgent", agent)
def _store_key(self, conv: Conversation) -> dict[str, str]:
agent = self._claude_agent(conv.agent_name)
return {
"project_key": project_key_for_directory(str(agent.cwd)),
"session_id": cast("str", conv.session_id),
}
def _runner(self, row_id: int) -> Runner:
runner = self._runners.get(row_id)
if runner is None:
runner = Runner()
self._runners[row_id] = runner
return runner
def _track(self, task: asyncio.Task[None]) -> None:
self._tasks.add(task)
task.add_done_callback(self._tasks.discard)
def aware(value: datetime) -> datetime:
return value if value.tzinfo is not None else value.replace(tzinfo=UTC)
def iso(value: datetime | None) -> str | None:
return aware(value).isoformat(timespec="seconds") if value is not None else None
+125
View File
@@ -0,0 +1,125 @@
"""Every string the gateway puts in front of a model; English defaults, overridable."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
from beaver_gateway.conversations import injects
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from datetime import date, datetime
from beaver_gateway.conversations.distill import DistillContext
from beaver_gateway.conversations.kinds import Kind
from beaver_gateway.conversations.rotation import HandoutContext
from beaver_gateway.storage.models import Conversation
__all__ = [
"ConversationTexts",
"EnvelopeTexts",
"NewDayContext",
"SeedContext",
"UserSaid",
]
@dataclass(frozen=True, slots=True)
class SeedContext:
kind: Kind
seed: str
agent: str
parent: Conversation | None
text: str | None
title: str | None
@dataclass(frozen=True, slots=True)
class NewDayContext:
"""``reason`` is ``night``, ``age`` or ``context``; only ``night`` is a new day."""
day: date
reason: str
moved: int
@dataclass(frozen=True, slots=True)
class UserSaid:
conversation_id: str
kind: str
title: str | None
text: str
at: datetime
@dataclass(frozen=True, slots=True)
class EnvelopeTexts:
header: str = (
"[envelope - background signal, not a message; react only if it bears "
"on the question]"
)
time: str = "time: {stamp} ({zone})"
changed: str = "vault, changed {since}: {names}"
since: str = "since {time}"
since_start: str = "since start"
truncated: str = "… (envelope cap)"
file_header: str = "--- {path}, added lines only ---"
more_lines: str = "+ … {count} more"
@dataclass(frozen=True, slots=True)
class ConversationTexts:
merge_prompt: str = (
"This branch is closing. Write a merge note for the master: what was "
"decided, what was done, what was not and why, open questions. "
"Identifiers and links verbatim. Brief, past tense."
)
inject_header: Callable[[injects.InjectContext], str] = injects.inject_header
bundle_header: str = "[injects accumulated since {since}; not the user]"
interrupted: str = "interrupted"
answered: str = "The user answered: {answer}"
unanswered: str = (
"The user did not answer within {minutes} min. The question was shown "
"to them as text; finish the turn now, the answer comes as the next "
"message."
)
seed: Callable[[SeedContext], Awaitable[str | None] | str | None] | None = None
"""Body of a seed by mode; ``None`` from it falls back to the defaults below."""
seed_head: str = "[seed: {seed}] {kind}{title}, {stamp}."
seed_copy: str = "The parent's history is copied ({scope}); continue in it."
seed_copy_window: str = "last {window} turns"
seed_copy_all: str = "whole history"
seed_morning_missing: str = "No handout arrived."
handout: Callable[[HandoutContext], Awaitable[str] | str] | str = (
"This master is closing ({reason}). Write the handout for {day}: a "
"briefing for the morning, not a task list - past tense, no imperatives."
)
new_day: Callable[[NewDayContext], Awaitable[str] | str] | str = (
"The master was replaced ({reason}); the handout for {day} is written."
)
moved_injects: str = " {moved} queued injects moved over from the old master."
distill: Callable[[DistillContext], Awaitable[str] | str] | None = None
"""The distiller fork's first message; ``None`` uses the two templates below."""
distill_prompt: str = (
"Deep chat «{chat}» is closed ({reason}), today is {day}. Write the "
"digest as a file and the merge as your reply: up to 5 lines, third "
"person."
)
distill_prompt_no_memory: str = (
"Deep chat «{chat}» is closed ({reason}), today is {day}. Memory is off "
"for it: write no file, only the merge as your reply - up to 5 lines, "
"third person."
)
closed: str = "Deep chat [[{chat}]] closed{digest}.\n{text}"
closed_digest: str = ", digest [[{digest}]]"
digest_missing: str = "the digest file did not appear"
too_long: str = (
"`{name}`: {lines} lines against a cap of {max_lines}. The write was "
"rejected and the file restored. Shorten and rewrite."
)
cut_by_restart: str = "turn {turn_id} was cut by a gateway restart"
repaired_tools: str = (
"; {fixed} open tool calls received tool_result «{interrupted}»"
)
envelope: EnvelopeTexts = field(default_factory=EnvelopeTexts)
@@ -13,9 +13,9 @@ from typing import TYPE_CHECKING, Any, cast
from claude_agent_sdk import create_sdk_mcp_server, tool from claude_agent_sdk import create_sdk_mcp_server, tool
from beaver_gateway.core.injects import URGENCY from beaver_gateway.conversations.injects import URGENCY
from beaver_gateway.core.kinds import as_kind from beaver_gateway.conversations.kinds import as_kind
from beaver_gateway.core.redact import redact_data from beaver_gateway.security.redact import redact_data
URGENCY_HELP = ( URGENCY_HELP = (
"normal waits for the hourly batch or rides with the next turn, wake " "normal waits for the hourly batch or rides with the next turn, wake "
@@ -31,11 +31,11 @@ if TYPE_CHECKING:
from claude_agent_sdk import McpSdkServerConfig, SdkMcpTool from claude_agent_sdk import McpSdkServerConfig, SdkMcpTool
from beaver_gateway.core.conversations import Conversations from beaver_gateway.conversations.service import Conversations
__all__ = ["SERVER_NAME", "TOOL_NAMES", "build_tool_server"] __all__ = ["SERVER_NAME", "TOOL_NAMES", "build_tool_server"]
_log = logging.getLogger("beaver_gateway.core.gateway_tools") _log = logging.getLogger("beaver_gateway.conversations.tools")
SERVER_NAME = "gateway" SERVER_NAME = "gateway"
SAY_IN_USER_TURN = ( SAY_IN_USER_TURN = (
@@ -230,7 +230,7 @@ def _tools(conversations: Conversations, key: str) -> list[SdkMcpTool[Any]]:
conv, conv,
str(args["text"]), str(args["text"]),
urgency=cast("Any", args.get("urgency") or "normal"), urgency=cast("Any", args.get("urgency") or "normal"),
origin="агент", origin="agent",
) )
return _text(f"queued #{item.id}") return _text(f"queued #{item.id}")
+444
View File
@@ -0,0 +1,444 @@
"""Running turns: one worker per conversation over its queue, the backend call."""
from __future__ import annotations
import asyncio
import contextlib
import inspect
import logging
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any, cast
from uuid import uuid4
from claude_agent_sdk import (
AssistantMessage,
RateLimitEvent,
ResultMessage,
StreamEvent,
ToolResultBlock,
ToolUseBlock,
UserMessage,
)
from beaver_gateway.backends.capture import TurnCapture
from beaver_gateway.backends.transcript import text_of
from beaver_gateway.conversations import injects
from beaver_gateway.conversations.seeds import Seeds
from beaver_gateway.conversations.state import Runner, aware, iso
from beaver_gateway.conversations.texts import UserSaid
from beaver_gateway.frontends.accumulate import StreamAccumulator
from beaver_gateway.storage.models import Conversation, InjectQueueItem, RateLimit
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Callable, Sequence
from beaver_gateway.events.stream import MessageStreamEvent
__all__ = ["Turns"]
_log = logging.getLogger(__name__)
class Turns(Seeds):
async def turn(
self,
conv: Conversation,
*,
messages: Sequence[Any],
origin: str,
capture: TurnCapture | None = None,
session_id: str | None = None,
use_session: bool = True,
tools: bool = True,
turn_id: str | None = None,
item_origin: str | None = None,
) -> AsyncIterator[MessageStreamEvent]:
"""Run one turn under the conversation's lock; the only path to the backend."""
row_id = cast("int", conv.id)
runner = self._runner(row_id)
backend = self._backend(conv.agent_name)
turn_id = turn_id or f"turn_{uuid4().hex[:12]}"
capture = capture or TurnCapture()
resume = session_id if session_id is not None else conv.session_id
async with runner.lock:
runner.turn_id = turn_id
runner.origin = origin
runner.text = _prompt_preview(messages)
runner.started_at = datetime.now(UTC)
runner.tools = {}
await self._mark_running(conv, turn_id)
before = await self.before_turn(conv)
self._bus.publish(
"turn.start",
conversation_id=conv.external_id,
turn_id=turn_id,
origin=origin,
item_origin=item_origin,
text=runner.text,
)
stop = "error"
cut = False
try:
events = backend.complete(
agent=self._claude_agent(conv.agent_name),
messages=messages,
conversation_id=conv.external_id,
session_id=resume if use_session else None,
reseed=not use_session,
capture=capture,
kind=conv.kind,
pinned=conv.kind == "master",
tools=tools,
observer=self._observer(conv, runner, turn_id, origin),
turn_id=turn_id,
)
async for event in events:
yield event
stop = "interrupted" if capture.interrupted else "end_turn"
except asyncio.CancelledError:
cut = True
raise
finally:
runner.turn_id = None
await self._mark_done(conv, capture, cut=cut)
if stop != "error":
try:
await self.after_turn(conv, before)
except Exception: # noqa: BLE001
_log.exception("after-turn hook on %s failed", conv.external_id)
self._bus.publish(
"turn.end",
conversation_id=conv.external_id,
turn_id=turn_id,
origin=origin,
item_origin=item_origin,
stop=stop,
usage=_usage_dict(capture),
)
async def run_text_turn(
self,
conv: Conversation,
text: str,
*,
origin: str,
tools: bool = True,
turn_id: str | None = None,
item_origin: str | None = None,
) -> tuple[str, TurnCapture]:
capture = TurnCapture()
acc = StreamAccumulator()
agent = self._claude_agent(conv.agent_name)
async for event in self.turn(
conv,
messages=[{"role": "user", "content": text}],
origin=origin,
capture=capture,
tools=tools,
turn_id=turn_id,
item_origin=item_origin,
):
acc.feed(event)
message = acc.finalize(model=agent.model)
reply = "\n\n".join(
getattr(b, "text", "")
for b in message.content
if getattr(b, "type", "") == "text"
).strip()
return reply, capture
async def before_turn(self, conv: Conversation) -> str | None: # noqa: ARG002
return None
async def after_turn(self, conv: Conversation, before: str | None) -> None: # noqa: ARG002
return
def turn_origin(self, conv: Conversation) -> str | None:
runner = self._runners.get(cast("int", conv.id))
return runner.origin if runner is not None and runner.turn_id else None
def _ensure_worker(self, row_id: int) -> None:
runner = self._runner(row_id)
runner.wake.set()
if runner.task is None or runner.task.done():
runner.task = asyncio.create_task(self._worker(row_id))
self._track(runner.task)
async def _worker(self, row_id: int) -> None:
runner = self._runner(row_id)
while True:
items = await self._queue.pending(row_id)
batch, wait = self._pick(items)
if batch is None:
runner.wake.clear()
if wait is None:
return
with contextlib.suppress(TimeoutError):
await asyncio.wait_for(runner.wake.wait(), timeout=wait)
continue
conv = await self.get_row(row_id)
if conv is None:
await self._queue.finish(batch, status="failed")
return
await self._run_batch(conv, batch)
def _pick(
self, items: list[InjectQueueItem]
) -> tuple[list[InjectQueueItem] | None, float | None]:
if not items:
return None, None
head = items[0]
if head.priority == "urgent":
return [head], None
tail = [i for i in items if i is not head and i.priority in ("wake", "normal")]
if head.priority in ("user", "wake"):
return [head, *tail], None
age = (datetime.now(UTC) - aware(head.created_at)).total_seconds()
if age >= self._normal_window:
return [head, *tail], None
return None, max(self._normal_window - age, 1.0)
async def _run_batch(
self, conv: Conversation, batch: list[InjectQueueItem]
) -> None:
head = batch[0]
turn_id = f"turn_{uuid4().hex[:12]}"
await self._queue.start(batch, turn_id)
if head.priority == "user":
origin = "user"
prompt = head.text
await self._note_user(conv, head.text)
envelope = self._envelope_for(conv, head.text)
if envelope:
prompt += "\n\n" + envelope
if len(batch) > 1:
prompt += "\n\n" + self._bundle(batch[1:])
else:
origin = "inject"
prompt = "\n\n".join(
f"{self._texts.inject_header(injects.context_of(i))}\n{i.text}"
for i in batch
)
seed = await self.pending_seed(conv)
if seed:
prompt = f"{seed}\n\n{prompt}"
try:
text, capture = await self.run_text_turn(
conv, prompt, origin=origin, turn_id=turn_id, item_origin=head.origin
)
except Exception: # noqa: BLE001
_log.exception("turn %s on %s failed", turn_id, conv.external_id)
await self._queue.finish(batch, status="failed")
return
await self._queue.finish(
batch, status="interrupted" if capture.interrupted else "done"
)
if origin == "user":
self._bus.publish(
"reply",
conversation_id=conv.external_id,
turn_id=turn_id,
item=head.id,
item_origin=head.origin,
source="queue",
prompt=prompt,
user_text=head.text,
text=text,
)
def _bundle(self, items: Sequence[InjectQueueItem]) -> str:
lines = [self._texts.bundle_header.format(since=iso(items[0].created_at))]
lines.extend(f"- [{i.origin}] {i.text}" for i in items)
return "\n".join(lines)
def _envelope_for(self, conv: Conversation, text: str = "") -> str | None:
if self._envelope is None:
return None
if conv.kind == "master":
return self._envelope.build(text=text, kind="master")
if conv.kind == "branch":
return self._envelope.recall_only(text=text, kind="branch")
return None
async def _note_user(self, conv: Conversation, text: str) -> None:
if self._user_sink is None:
return
message = UserSaid(
conversation_id=conv.external_id,
kind=conv.kind,
title=conv.title,
text=text,
at=datetime.now(UTC),
)
try:
result = self._user_sink(message)
if inspect.isawaitable(result):
await result
except Exception: # noqa: BLE001
_log.exception("user sink failed for %s", conv.external_id)
def _observer(
self, conv: Conversation, runner: Runner, turn_id: str, origin: str
) -> Callable[[Any], None]:
conversation_id = conv.external_id
def observe(message: Any) -> None:
parent = getattr(message, "parent_tool_use_id", None)
if isinstance(message, RateLimitEvent):
self._observe_rate_limit(conv, message)
elif isinstance(message, StreamEvent):
self._bus.publish(
"stream",
conversation_id=conversation_id,
turn_id=turn_id,
origin=origin,
parent_tool_use_id=parent,
event=message.event,
)
elif isinstance(message, AssistantMessage):
for block in message.content:
if isinstance(block, ToolUseBlock):
event = self._bus.publish(
"tool",
conversation_id=conversation_id,
turn_id=turn_id,
origin=origin,
parent_tool_use_id=parent,
tool_use_id=block.id,
name=block.name,
input=block.input,
)
runner.tools[block.id] = {
"tool_use_id": block.id,
"name": block.name,
"input": block.input,
"parent_tool_use_id": parent,
"started_at": event["ts"],
"ended_at": None,
"is_error": None,
"content": None,
}
elif isinstance(message, UserMessage):
blocks = message.content if isinstance(message.content, list) else ()
for block in blocks:
if isinstance(block, ToolResultBlock):
event = self._bus.publish(
"tool.result",
conversation_id=conversation_id,
turn_id=turn_id,
origin=origin,
parent_tool_use_id=parent,
tool_use_id=block.tool_use_id,
is_error=bool(block.is_error),
content=_result_preview(block.content),
)
node = runner.tools.get(block.tool_use_id)
if node is not None:
node["ended_at"] = event["ts"]
node["is_error"] = event["is_error"]
node["content"] = event["content"]
elif isinstance(message, ResultMessage) and parent is None:
self._bus.publish(
"result",
conversation_id=conversation_id,
turn_id=turn_id,
origin=origin,
subtype=message.subtype,
is_error=message.is_error,
num_turns=message.num_turns,
)
return observe
def _observe_rate_limit(self, conv: Conversation, message: RateLimitEvent) -> None:
info = message.rate_limit_info
row = RateLimit(
window=info.rate_limit_type or "unknown",
status=info.status,
utilization=info.utilization,
resets_at=_from_unix(info.resets_at),
overage_status=info.overage_status,
overage_resets_at=_from_unix(info.overage_resets_at),
agent_name=conv.agent_name,
session_id=message.session_id,
raw=dict(info.raw),
)
self._bus.publish(
"rate_limit",
conversation_id=conv.external_id,
window=row.window,
status=row.status,
utilization=row.utilization,
resets_at=iso(row.resets_at),
overage_status=row.overage_status,
)
self._track(asyncio.create_task(self._record_rate_limit(row)))
async def _record_rate_limit(self, row: RateLimit) -> None:
try:
async with self._db.session() as session:
session.add(row)
await session.commit()
except Exception: # noqa: BLE001
_log.exception("rate limit write failed")
async def _mark_running(self, conv: Conversation, turn_id: str) -> None:
async def apply(row: Conversation) -> None:
row.running_turn = turn_id
row.last_activity_at = datetime.now(UTC)
await self._update(conv, apply)
async def _mark_done(
self, conv: Conversation, capture: TurnCapture, *, cut: bool = False
) -> None:
async def apply(row: Conversation) -> None:
if not cut:
row.running_turn = None
row.last_activity_at = datetime.now(UTC)
if capture.session_id is not None:
row.session_id = capture.session_id
await self._update(conv, apply)
def _prompt_preview(messages: Sequence[Any], limit: int = 400) -> str | None:
if not messages:
return None
text = text_of(messages[-1].get("content"))
return text[:limit] if text else None
def _from_unix(value: int | None) -> datetime | None:
return datetime.fromtimestamp(value, tz=UTC) if value is not None else None
def _result_preview(
content: str | list[dict[str, Any]] | None, limit: int = 400
) -> str:
if content is None:
return ""
text = (
content
if isinstance(content, str)
else "\n".join(
str(part.get("text", ""))
for part in content
if isinstance(part, dict) and part.get("type") == "text"
)
)
return text if len(text) <= limit else text[:limit] + ""
def _usage_dict(capture: TurnCapture) -> dict[str, Any] | None:
usage = capture.usage
if usage is None:
return None
return {
"input": usage.input_tokens,
"output": usage.output_tokens,
"cache_read": usage.cache_read_tokens,
"cache_creation": usage.cache_creation_tokens,
"cost_usd": usage.cost_usd,
"duration_ms": usage.duration_ms,
}
-7
View File
@@ -1,7 +0,0 @@
"""Cross-cutting machinery: registries, event protocol, auth, sessions."""
from __future__ import annotations
from beaver_gateway.core.registry import AgentRegistry, Gateway, McpRegistry
__all__ = ["AgentRegistry", "Gateway", "McpRegistry"]
File diff suppressed because it is too large Load Diff
-117
View File
@@ -1,117 +0,0 @@
"""Agent / MCP registries + the user-facing ``Gateway`` collector.
The user's ``/config/config.py`` ends with::
gateway = Gateway(agents=[...], mcps=[...], frontends=[...])
``cli.main`` picks that object up and builds the registries.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable, Iterable, Iterator
from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.core.conversations import ConversationTexts, UserSaid
from beaver_gateway.core.distill import Distiller
from beaver_gateway.core.envelope import RecallContext
from beaver_gateway.core.rotation import RotationPolicy
from beaver_gateway.core.scheduler import Budget, Job
from beaver_gateway.core.watch import VaultWatch
from beaver_gateway.frontends.base import Frontend
from beaver_gateway.mcp.types import McpServerT
class AgentRegistry:
"""Name → agent lookup with duplicate detection."""
def __init__(self, agents: Iterable[BaseAgent]) -> None:
self._agents: dict[str, BaseAgent] = {}
for a in agents:
if a.name in self._agents:
msg = f"duplicate agent name: {a.name!r}"
raise ValueError(msg)
self._agents[a.name] = a
def __getitem__(self, name: str) -> BaseAgent:
return self._agents[name]
def get(self, name: str) -> BaseAgent | None:
return self._agents.get(name)
def __iter__(self) -> Iterator[BaseAgent]:
return iter(self._agents.values())
def __len__(self) -> int:
return len(self._agents)
def __contains__(self, name: object) -> bool:
return name in self._agents
class McpRegistry:
"""Name → MCP server lookup with duplicate detection."""
def __init__(self, mcps: Iterable[McpServerT]) -> None:
self._mcps: dict[str, McpServerT] = {}
for m in mcps:
if m.name in self._mcps:
msg = f"duplicate mcp name: {m.name!r}"
raise ValueError(msg)
self._mcps[m.name] = m
def __getitem__(self, name: str) -> McpServerT:
return self._mcps[name]
def get(self, name: str) -> McpServerT | None:
return self._mcps.get(name)
def __iter__(self) -> Iterator[McpServerT]:
return iter(self._mcps.values())
def __len__(self) -> int:
return len(self._mcps)
def __contains__(self, name: object) -> bool:
return name in self._mcps
@dataclass(slots=True)
class Gateway:
"""Top-level object the user assembles in ``/config/config.py``."""
agents: list[BaseAgent] = field(default_factory=list)
mcps: list[McpServerT] = 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)."""
jobs: list[Job] = field(default_factory=list)
"""Cron / webhook / event jobs for ``core/scheduler`` (§3.6, §4.5)."""
rotation: RotationPolicy | None = None
"""When a master is rotated (§4.5); ``None`` keeps the defaults."""
watch: VaultWatch | None = None
"""Vault watcher feeding the envelope (§3.5, §4.6); ``None`` = no vault block."""
recall: Callable[[RecallContext], str | None] | None = None
"""Envelope lookup on the user's text: pointers into the vault (cards,
the agent's notes, due dates) the gateway knows no paths for (§3.3)."""
user_sink: Callable[[UserSaid], Awaitable[None] | None] | None = None
"""Sees every user message as it enters a master or branch turn - the
setup's own grep-able log of what the user said, outside the transcript."""
budget: Budget | None = None
"""Subscription window past which non-critical jobs wait (§4.5)."""
distiller: Distiller | None = None
"""Who closes deep chats and where the digests and the index live (§8.4)."""
tz: str = "UTC"
"""Local zone for the envelope clock and the rotation hour."""
host: str = "0.0.0.0" # noqa: S104
port: int = 8000
"""The one listener; every HTTP frontend is mounted under its ``path``."""
public_url: str | None = None
"""Origin the reverse proxy shows the world (``https://b.example.com``).
Advertised endpoints and MCP discovery are built on it; ``None``
derives the origin from each request."""
@@ -33,7 +33,7 @@ from anthropic.types import (
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from beaver_gateway.core.events import MessageStreamEvent, StopReason from beaver_gateway.events.stream import MessageStreamEvent, StopReason
__all__ = ["StreamAccumulator", "accumulate"] __all__ = ["StreamAccumulator", "accumulate"]
@@ -27,8 +27,8 @@ import itsdangerous
from fastapi import FastAPI, HTTPException, Request, status from fastapi import FastAPI, HTTPException, Request, status
from fastapi.responses import FileResponse, JSONResponse, Response from fastapi.responses import FileResponse, JSONResponse, Response
from beaver_gateway.core import audit
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.security import audit
if TYPE_CHECKING: if TYPE_CHECKING:
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
+8 -8
View File
@@ -25,21 +25,21 @@ from fastapi import FastAPI, HTTPException, Request, status
from fastapi.responses import JSONResponse, StreamingResponse from fastapi.responses import JSONResponse, StreamingResponse
from beaver_gateway.agents.claude import ClaudeAgent from beaver_gateway.agents.claude import ClaudeAgent
from beaver_gateway.core import audit from beaver_gateway.backends.capture import TurnCapture
from beaver_gateway.core.transcript import fingerprint, text_of from beaver_gateway.backends.transcript import fingerprint, text_of
from beaver_gateway.core.turn_capture import TurnCapture from beaver_gateway.frontends.accumulate import StreamAccumulator
from beaver_gateway.core.turn_record import TurnRecord
from beaver_gateway.frontends._accumulate import StreamAccumulator
from beaver_gateway.frontends._auth import require_token
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.frontends.bearer import require_token
from beaver_gateway.frontends.turn_record import TurnRecord
from beaver_gateway.security import audit
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator, Awaitable, Callable from collections.abc import AsyncIterator, Awaitable, Callable
from anthropic.types import Message, MessageParam from anthropic.types import Message, MessageParam
from beaver_gateway.core.conversations import Conversations from beaver_gateway.conversations.service import Conversations
from beaver_gateway.core.events import MessageStreamEvent from beaver_gateway.events.stream import MessageStreamEvent
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.storage.models import Conversation from beaver_gateway.storage.models import Conversation
+11 -11
View File
@@ -29,20 +29,20 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, StreamingResponse from fastapi.responses import JSONResponse, StreamingResponse
from sqlmodel import col, select from sqlmodel import col, select
from beaver_gateway.core import audit from beaver_gateway.conversations.injects import URGENCY
from beaver_gateway.core.auth import VALID_SCOPES, hash_token from beaver_gateway.conversations.kinds import Kind, as_kind
from beaver_gateway.core.conversations import SEEDS, implied_title from beaver_gateway.conversations.service import SEEDS, implied_title
from beaver_gateway.core.injects import URGENCY from beaver_gateway.frontends.base import Frontend
from beaver_gateway.core.kinds import Kind, as_kind from beaver_gateway.frontends.bearer import require_token
from beaver_gateway.frontends._auth import require_token from beaver_gateway.frontends.sse import (
from beaver_gateway.frontends._sse import (
KEEPALIVE, KEEPALIVE,
SSE_HEADERS, SSE_HEADERS,
events_with_heartbeat, events_with_heartbeat,
sse_pack, sse_pack,
) )
from beaver_gateway.frontends._urls import frontend_url from beaver_gateway.frontends.urls import frontend_url
from beaver_gateway.frontends.base import Frontend from beaver_gateway.security import audit
from beaver_gateway.security.auth import VALID_SCOPES, hash_token
from beaver_gateway.storage import ( from beaver_gateway.storage import (
create_token, create_token,
list_audit_records, list_audit_records,
@@ -61,9 +61,9 @@ if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterable, Sequence from collections.abc import AsyncIterator, Iterable, Sequence
from pathlib import Path from pathlib import Path
from beaver_gateway.core.conversations import Conversations from beaver_gateway.conversations.service import Conversations
from beaver_gateway.core.scheduler import Job, Scheduler
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.jobs.scheduler import Job, Scheduler
_log = logging.getLogger("beaver_gateway.frontends.api") _log = logging.getLogger("beaver_gateway.frontends.api")
+4 -4
View File
@@ -21,11 +21,11 @@ if TYPE_CHECKING:
from starlette.types import ASGIApp from starlette.types import ASGIApp
from beaver_gateway.app import AgentRegistry, McpRegistry
from beaver_gateway.backends.base import Backend from beaver_gateway.backends.base import Backend
from beaver_gateway.core.auth import TokenStore from beaver_gateway.conversations.kinds import Kind
from beaver_gateway.core.kinds import Kind from beaver_gateway.frontends.turn_record import TurnRecord
from beaver_gateway.core.registry import AgentRegistry, McpRegistry from beaver_gateway.security.auth import TokenStore
from beaver_gateway.core.turn_record import TurnRecord
from beaver_gateway.storage import Database from beaver_gateway.storage import Database
from beaver_gateway.storage.models import Conversation, ConversationBinding from beaver_gateway.storage.models import Conversation, ConversationBinding
@@ -32,7 +32,7 @@ if TYPE_CHECKING:
from anthropic.types import MessageParam from anthropic.types import MessageParam
from beaver_gateway.core.turn_record import TurnRecord from beaver_gateway.frontends.turn_record import TurnRecord
_log = logging.getLogger("beaver_gateway.frontends.markdown.crossfront") _log = logging.getLogger("beaver_gateway.frontends.markdown.crossfront")
@@ -45,23 +45,10 @@ from fastapi import FastAPI, HTTPException, Request, status
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, StreamingResponse from fastapi.responses import JSONResponse, StreamingResponse
from beaver_gateway.core import audit from beaver_gateway.backends.capture import TurnCapture
from beaver_gateway.core.conversation_store import ( from beaver_gateway.frontends.accumulate import StreamAccumulator
diff_and_fork,
load_messages,
rewrite_messages,
)
from beaver_gateway.core.turn_capture import TurnCapture
from beaver_gateway.core.turn_record import TurnRecord
from beaver_gateway.frontends._accumulate import StreamAccumulator
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.bearer import require_token
from beaver_gateway.frontends.markdown import parser, renderer from beaver_gateway.frontends.markdown import parser, renderer
from beaver_gateway.frontends.markdown.crossfront import CrossFrontendLogger from beaver_gateway.frontends.markdown.crossfront import CrossFrontendLogger
from beaver_gateway.frontends.markdown.files import ( from beaver_gateway.frontends.markdown.files import (
@@ -69,12 +56,25 @@ from beaver_gateway.frontends.markdown.files import (
reattach_frontmatter, reattach_frontmatter,
write_atomic, write_atomic,
) )
from beaver_gateway.frontends.markdown.history import (
diff_and_fork,
load_messages,
rewrite_messages,
)
from beaver_gateway.frontends.markdown.mirror import FRONTEND, ChatMirror from beaver_gateway.frontends.markdown.mirror import FRONTEND, ChatMirror
from beaver_gateway.frontends.sse import (
KEEPALIVE,
SSE_HEADERS,
events_with_heartbeat,
sse_pack,
)
from beaver_gateway.frontends.turn_record import TurnRecord
from beaver_gateway.security import audit
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator, Callable from collections.abc import AsyncIterator, Callable
from beaver_gateway.core.kinds import Kind from beaver_gateway.conversations.kinds import Kind
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.storage.models import Conversation, ConversationBinding from beaver_gateway.storage.models import Conversation, ConversationBinding
@@ -19,8 +19,6 @@ from typing import TYPE_CHECKING, Any
import frontmatter import frontmatter
from beaver_gateway.core.conversation_store import load_messages, rewrite_messages
from beaver_gateway.core.turn_record import slugify
from beaver_gateway.frontends.markdown import renderer from beaver_gateway.frontends.markdown import renderer
from beaver_gateway.frontends.markdown.crossfront import strip_trailing_user_scaffold from beaver_gateway.frontends.markdown.crossfront import strip_trailing_user_scaffold
from beaver_gateway.frontends.markdown.files import ( from beaver_gateway.frontends.markdown.files import (
@@ -28,12 +26,14 @@ from beaver_gateway.frontends.markdown.files import (
reattach_frontmatter, reattach_frontmatter,
write_atomic, write_atomic,
) )
from beaver_gateway.frontends.markdown.history import load_messages, rewrite_messages
from beaver_gateway.frontends.turn_record import slugify
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable, Sequence from collections.abc import Callable, Sequence
from pathlib import Path from pathlib import Path
from beaver_gateway.core.bus import Event from beaver_gateway.events.bus import Event
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.storage.models import Conversation, ConversationBinding from beaver_gateway.storage.models import Conversation, ConversationBinding
+2 -2
View File
@@ -45,10 +45,10 @@ from starlette.applications import Starlette
from starlette.responses import HTMLResponse, JSONResponse, StreamingResponse from starlette.responses import HTMLResponse, JSONResponse, StreamingResponse
from starlette.routing import Route from starlette.routing import Route
from beaver_gateway.core import audit
from beaver_gateway.frontends._urls import external_base
from beaver_gateway.frontends.base import Frontend from beaver_gateway.frontends.base import Frontend
from beaver_gateway.frontends.urls import external_base
from beaver_gateway.mcp.internal_app import ALL_NAMESPACE from beaver_gateway.mcp.internal_app import ALL_NAMESPACE
from beaver_gateway.security import audit
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator, Mapping from collections.abc import AsyncIterator, Mapping
@@ -44,9 +44,9 @@ from beaver_gateway.frontends.telegram.outbox import Outbox
from beaver_gateway.frontends.telegram.render import chunks, status_label from beaver_gateway.frontends.telegram.render import chunks, status_label
if TYPE_CHECKING: if TYPE_CHECKING:
from beaver_gateway.core.bus import Event, EventBus from beaver_gateway.conversations.kinds import Kind
from beaver_gateway.core.conversations import Conversations from beaver_gateway.conversations.service import Conversations
from beaver_gateway.core.kinds import Kind from beaver_gateway.events.bus import Event, EventBus
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.storage.models import Conversation, ConversationBinding from beaver_gateway.storage.models import Conversation, ConversationBinding
@@ -773,7 +773,7 @@ class TelegramFrontend(Frontend):
draft = self._drafts.pop(conv.external_id, None) draft = self._drafts.pop(conv.external_id, None)
if draft is not None: if draft is not None:
await draft.finish(chunks(text)[-1] if text.strip() else "") await draft.finish(chunks(text)[-1] if text.strip() else "")
if origin != FRONTEND and not origin.startswith("сид"): if origin != FRONTEND and not origin.startswith("seed"):
user_text = str(event.get("user_text") or "") user_text = str(event.get("user_text") or "")
if user_text: if user_text:
await self._deliver( await self._deliver(
@@ -32,7 +32,7 @@ from beaver_gateway.storage.models import Delivery
if TYPE_CHECKING: if TYPE_CHECKING:
from aiogram import Bot from aiogram import Bot
from beaver_gateway.core.bus import EventBus from beaver_gateway.events.bus import EventBus
from beaver_gateway.storage.db import Database from beaver_gateway.storage.db import Database
__all__ = ["Outbox"] __all__ = ["Outbox"]
View File
@@ -33,7 +33,7 @@ from starlette.applications import Starlette
from starlette.responses import JSONResponse from starlette.responses import JSONResponse
from starlette.routing import Route from starlette.routing import Route
from beaver_gateway.core.conversations import parse_at from beaver_gateway.conversations.service import parse_at
from beaver_gateway.storage.models import JobRunRecord from beaver_gateway.storage.models import JobRunRecord
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -44,15 +44,15 @@ if TYPE_CHECKING:
from pgqueuer.ports.driver import Driver from pgqueuer.ports.driver import Driver
from starlette.requests import Request from starlette.requests import Request
from beaver_gateway.core.conversations import Conversations, DistillResult from beaver_gateway.conversations.distill import LineCap
from beaver_gateway.core.distill import LineCap from beaver_gateway.conversations.injects import Priority
from beaver_gateway.core.injects import Priority from beaver_gateway.conversations.rotation import Rotation
from beaver_gateway.core.rotation import Rotation from beaver_gateway.conversations.service import Conversations, DistillResult
from beaver_gateway.storage.models import Conversation from beaver_gateway.storage.models import Conversation
__all__ = ["INJECT", "Budget", "Job", "JobRun", "LocalCron", "Scheduler", "next_run"] __all__ = ["INJECT", "Budget", "Job", "JobRun", "LocalCron", "Scheduler", "next_run"]
_log = logging.getLogger("beaver_gateway.core.scheduler") _log = logging.getLogger("beaver_gateway.jobs.scheduler")
INJECT = "inject" INJECT = "inject"
RETRY = timedelta(minutes=15) RETRY = timedelta(minutes=15)
+2 -2
View File
@@ -20,7 +20,7 @@ remembering to list it here.
What this does not reach: tools that never touch a FastMCP server — the What this does not reach: tools that never touch a FastMCP server — the
gateway's own ``gateway`` tools, and everything claude-code runs inside gateway's own ``gateway`` tools, and everything claude-code runs inside
its own process (``Bash``, ``Read``). Those are guarded by its own process (``Bash``, ``Read``). Those are guarded by
:mod:`beaver_gateway.core.policy` and the vault mounts instead. :mod:`beaver_gateway.agents.policy` and the vault mounts instead.
""" """
from __future__ import annotations from __future__ import annotations
@@ -31,7 +31,7 @@ import mcp.types as mt
from fastmcp.server.middleware import Middleware from fastmcp.server.middleware import Middleware
from fastmcp.tools.base import ToolResult from fastmcp.tools.base import ToolResult
from beaver_gateway.core.redact import redact, redact_data from beaver_gateway.security.redact import redact, redact_data
if TYPE_CHECKING: if TYPE_CHECKING:
from fastmcp.server.middleware import CallNext, MiddlewareContext from fastmcp.server.middleware import CallNext, MiddlewareContext
@@ -29,7 +29,7 @@ if TYPE_CHECKING:
__all__ = ["Change", "VaultWatch", "WatchRules"] __all__ = ["Change", "VaultWatch", "WatchRules"]
_log = logging.getLogger("beaver_gateway.core.watch") _log = logging.getLogger("beaver_gateway.vault.watch")
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
+6 -6
View File
@@ -18,11 +18,11 @@ from claude_agent_sdk import (
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from test_conversations import ScriptedClient, World from test_conversations import ScriptedClient, World
from beaver_gateway.core.conversation_store import rewrite_messages from beaver_gateway.frontends.markdown.history import rewrite_messages
from beaver_gateway.core.auth import TokenStore from beaver_gateway.security.auth import TokenStore
from beaver_gateway.core.registry import McpRegistry from beaver_gateway.app import McpRegistry
from beaver_gateway.core.scheduler import Job, JobRun, Scheduler from beaver_gateway.jobs.scheduler import Job, JobRun, Scheduler
from beaver_gateway.core.transcript import build_entries from beaver_gateway.backends.transcript import build_entries
from beaver_gateway.frontends.admin import AdminFrontend from beaver_gateway.frontends.admin import AdminFrontend
from beaver_gateway.frontends.admin.frontend import build_app as build_admin from beaver_gateway.frontends.admin.frontend import build_app as build_admin
from beaver_gateway.frontends.api import ApiFrontend from beaver_gateway.frontends.api import ApiFrontend
@@ -473,7 +473,7 @@ async def test_close_distills_a_deep_chat(world: World) -> None:
assert body["digest"] == str(config.dir / "2026-08-29 - тема.md") assert body["digest"] == str(config.dir / "2026-08-29 - тема.md")
assert body["text"].count("\n") == 2 assert body["text"].count("\n") == 2
assert (await world.conversations.get(chat.external_id)).status == "closed" assert (await world.conversations.get(chat.external_id)).status == "closed"
assert (await world.conversations.queue.recent(master.id))[0].origin == "выжимка" assert (await world.conversations.queue.recent(master.id))[0].origin == "digest"
again = await api.http.post( again = await api.http.post(
f"/conversations/{chat.external_id}/close", headers=HEADERS f"/conversations/{chat.external_id}/close", headers=HEADERS
) )
+3 -3
View File
@@ -7,8 +7,8 @@ import pytest
from fastapi import HTTPException from fastapi import HTTPException
from starlette.requests import Request from starlette.requests import Request
from beaver_gateway.core.auth import TokenStore from beaver_gateway.security.auth import TokenStore
from beaver_gateway.frontends._auth import require_token from beaver_gateway.frontends.bearer import require_token
def _request(query: str = "", headers: dict[str, str] | None = None) -> Request: def _request(query: str = "", headers: dict[str, str] | None = None) -> Request:
@@ -61,7 +61,7 @@ async def test_bootstrap_entry_can_carry_a_scope() -> None:
def test_access_log_filter_masks_query_tokens() -> None: def test_access_log_filter_masks_query_tokens() -> None:
from beaver_gateway.core.redact import RedactFilter from beaver_gateway.security.redact import RedactFilter
record = logging.LogRecord( record = logging.LogRecord(
"uvicorn.access", "uvicorn.access",
+3 -3
View File
@@ -32,8 +32,8 @@ from beaver_gateway.backends.claude_sdk import (
UsageEvent, UsageEvent,
fingerprint, fingerprint,
) )
from beaver_gateway.core.transcript import messages_from_entries from beaver_gateway.backends.transcript import messages_from_entries
from beaver_gateway.core.turn_capture import TurnCapture from beaver_gateway.backends.capture import TurnCapture
def _stream(index: int, text: str) -> list[StreamEvent]: def _stream(index: int, text: str) -> list[StreamEvent]:
@@ -506,7 +506,7 @@ async def test_deltas_reach_the_caller_before_the_turn_ends(cwd: Path) -> None:
async def test_policy_hook_denies_and_audits(cwd: Path) -> None: async def test_policy_hook_denies_and_audits(cwd: Path) -> None:
from beaver_gateway.core.policy import Deny, ToolCall from beaver_gateway.agents.policy import Deny, ToolCall
def no_days(c: ToolCall): def no_days(c: ToolCall):
p = c.path() p = c.path()
@@ -1,7 +1,7 @@
import tempfile import tempfile
from pathlib import Path from pathlib import Path
from beaver_gateway import config_loader from beaver_gateway import config
def test_config_imports_sibling_modules(): def test_config_imports_sibling_modules():
@@ -12,5 +12,5 @@ def test_config_imports_sibling_modules():
"assert RULES == ('x',)\n" "assert RULES == ('x',)\n"
"gateway = Gateway()\n" "gateway = Gateway()\n"
) )
gw = config_loader.load(root / "config.py") gw = config.load(root / "config.py")
assert gw.agents == [] assert gw.agents == []
+23 -19
View File
@@ -21,12 +21,16 @@ from claude_agent_sdk import (
from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions
from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend
from beaver_gateway.core.bus import EventBus from beaver_gateway.events.bus import EventBus
from beaver_gateway.core.conversations import Conversations, ConversationTexts, parse_at from beaver_gateway.conversations.service import (
from beaver_gateway.core.gateway_tools import SAY_IN_USER_TURN, _tools Conversations,
from beaver_gateway.core.registry import AgentRegistry ConversationTexts,
from beaver_gateway.core.sessions import SessionPool parse_at,
from beaver_gateway.core.transcript import ( )
from beaver_gateway.conversations.tools import SAY_IN_USER_TURN, _tools
from beaver_gateway.app import AgentRegistry
from beaver_gateway.backends.sessions import SessionPool
from beaver_gateway.backends.transcript import (
build_entries, build_entries,
close_open_tool_uses, close_open_tool_uses,
open_tool_uses, open_tool_uses,
@@ -312,14 +316,14 @@ async def test_urgent_interrupts_and_goes_first(world: World) -> None:
("urgent", "done"), ("urgent", "done"),
] ]
assert client.prompts[0] == "first" assert client.prompts[0] == "first"
assert client.prompts[1].startswith("[инжект: крон") assert client.prompts[1].startswith("[inject: крон")
assert "прервал предыдущий тёрн" in client.prompts[1] assert "cut the previous turn" in client.prompts[1]
assert client.prompts[1].endswith("ALERT") assert client.prompts[1].endswith("ALERT")
assert client.prompts[2] == "second" assert client.prompts[2] == "second"
async def test_inject_header_is_configurable(world: World) -> None: async def test_inject_header_is_configurable(world: World) -> None:
from beaver_gateway.core.conversations import ConversationTexts from beaver_gateway.conversations.service import ConversationTexts
world.conversations._texts = ConversationTexts( # noqa: SLF001 world.conversations._texts = ConversationTexts( # noqa: SLF001
inject_header=lambda ctx: ( inject_header=lambda ctx: (
@@ -343,7 +347,7 @@ async def test_normal_injects_ride_with_the_next_user_message(world: World) -> N
await world.settle(conv, 2) await world.settle(conv, 2)
prompts = ScriptedClient.instances[0].prompts prompts = ScriptedClient.instances[0].prompts
assert len(prompts) == 1 assert len(prompts) == 1
assert prompts[0].startswith("hello\n\n[инжекты") assert prompts[0].startswith("hello\n\n[injects")
assert "- [watch] vault changed" in prompts[0] assert "- [watch] vault changed" in prompts[0]
@@ -373,7 +377,7 @@ async def test_spawn_seeds_first_user_message(world: World) -> None:
) )
await world.settle(conv, 1) await world.settle(conv, 1)
prompt = ScriptedClient.instances[0].prompts[0] prompt = ScriptedClient.instances[0].prompts[0]
assert prompt.startswith("[сид: brief] branch «t», ") assert prompt.startswith("[seed: brief] branch «t», ")
assert prompt.endswith("\n\ndo X") assert prompt.endswith("\n\ndo X")
assert ScriptedClient.instances[0].options.system_prompt == "hi" assert ScriptedClient.instances[0].options.system_prompt == "hi"
@@ -459,8 +463,8 @@ async def test_copy_seed_forks_parent_with_window(world: World) -> None:
await world.settle(child, 1) await world.settle(child, 1)
assert ScriptedClient.instances[0].options.resume == child.session_id assert ScriptedClient.instances[0].options.resume == child.session_id
prompt = ScriptedClient.instances[0].prompts[0] prompt = ScriptedClient.instances[0].prompts[0]
assert prompt.startswith("[сид: copy] branch, ") assert prompt.startswith("[seed: copy] branch, ")
assert "последние 1 тёрнов" in prompt and prompt.endswith("\n\ngo") assert "last 1 turns" in prompt and prompt.endswith("\n\ngo")
assert (await world.conversations.get(child.external_id)).flags["seed"] is None assert (await world.conversations.get(child.external_id)).flags["seed"] is None
@@ -484,7 +488,7 @@ async def test_merge_injects_summary_into_parent(world: World) -> None:
assert (await world.conversations.get(branch.external_id)).status == "merged" assert (await world.conversations.get(branch.external_id)).status == "merged"
assert await world.statuses(master) == [("normal", "queued")] assert await world.statuses(master) == [("normal", "queued")]
item = (await world.conversations.queue.recent(master.id))[0] item = (await world.conversations.queue.recent(master.id))[0]
assert item.origin == "слив" and item.text == result.text assert item.origin == "merge" and item.text == result.text
async def test_recover_closes_open_tool_use_and_injects_interrupted( async def test_recover_closes_open_tool_use_and_injects_interrupted(
@@ -538,7 +542,7 @@ async def test_recover_closes_open_tool_use_and_injects_interrupted(
assert tail["message"]["content"][0] == { assert tail["message"]["content"][0] == {
"type": "tool_result", "type": "tool_result",
"tool_use_id": "t9", "tool_use_id": "t9",
"content": "прервано", "content": "interrupted",
"is_error": True, "is_error": True,
} }
assert tail["parentUuid"] == entries[-2]["uuid"] assert tail["parentUuid"] == entries[-2]["uuid"]
@@ -547,8 +551,8 @@ async def test_recover_closes_open_tool_use_and_injects_interrupted(
note = (await world.conversations.queue.recent(conv.id))[0] note = (await world.conversations.queue.recent(conv.id))[0]
assert ( assert (
"turn_dead" in note.text "turn_dead" in note.text
and "оборван" in note.text and "cut by a gateway restart" in note.text
and "1 незакрытых" in note.text and "1 open tool calls" in note.text
) )
await asyncio.sleep(0.2) await asyncio.sleep(0.2)
assert ScriptedClient.instances == [] assert ScriptedClient.instances == []
@@ -798,9 +802,9 @@ async def test_wake_injects_start_a_turn_and_take_normals_along(world: World) ->
await world.settle(conv, 2) await world.settle(conv, 2)
prompts = ScriptedClient.instances[0].prompts prompts = ScriptedClient.instances[0].prompts
assert len(prompts) == 1 assert len(prompts) == 1
assert prompts[0].startswith("[инжект: schedule") assert prompts[0].startswith("[inject: schedule")
assert "reminder" in prompts[0] assert "reminder" in prompts[0]
assert "[инжект: watch" in prompts[0] assert "[inject: watch" in prompts[0]
assert "digest" in prompts[0] assert "digest" in prompts[0]
assert await world.statuses(conv) == [("normal", "done"), ("wake", "done")] assert await world.statuses(conv) == [("normal", "done"), ("wake", "done")]
+22 -22
View File
@@ -17,24 +17,24 @@ from test_conversations import ScriptedClient, World, world
from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions
from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend
from beaver_gateway.core.conversations import ConversationTexts from beaver_gateway.conversations.service import ConversationTexts
from beaver_gateway.core.distill import ( from beaver_gateway.conversations.distill import (
Distiller, Distiller,
DistillContext, DistillContext,
LineCap, LineCap,
check_digest, check_digest,
trim_summary, trim_summary,
) )
from beaver_gateway.core.gateway_tools import _tools, build_tool_server from beaver_gateway.conversations.tools import _tools, build_tool_server
from beaver_gateway.core.registry import AgentRegistry from beaver_gateway.app import AgentRegistry
from beaver_gateway.core.scheduler import Job, JobRun, Scheduler from beaver_gateway.jobs.scheduler import Job, JobRun, Scheduler
from beaver_gateway.core.transcript import build_entries from beaver_gateway.backends.transcript import build_entries
from beaver_gateway.storage.models import Conversation from beaver_gateway.storage.models import Conversation
__all__ = ["world"] __all__ = ["world"]
DIGEST = """--- DIGEST = """---
type: выжимка type: digest
source: "[[{chat}]]" source: "[[{chat}]]"
date: 2026-08-29 date: 2026-08-29
--- ---
@@ -59,7 +59,7 @@ class DistillerClient(ScriptedClient):
async def receive_response(self): async def receive_response(self):
prompt = self.prompts[-1] prompt = self.prompts[-1]
written: list[ToolUseBlock] = [] written: list[ToolUseBlock] = []
if self.write and self.digest_dir is not None and "файл не пиши" not in prompt: if self.write and self.digest_dir is not None and "write no file" not in prompt:
chat = prompt.split("«", 1)[1].split("»", 1)[0] if "«" in prompt else "чат" chat = prompt.split("«", 1)[1].split("»", 1)[0] if "«" in prompt else "чат"
path = self.digest_dir / "2026-08-29 - тема.md" path = self.digest_dir / "2026-08-29 - тема.md"
path.write_text(self.frontmatter.format(chat=chat), encoding="utf-8") path.write_text(self.frontmatter.format(chat=chat), encoding="utf-8")
@@ -226,7 +226,7 @@ async def test_distill_writes_the_digest_indexes_it_and_merges_short(
assert result.digest.path == config.dir / "2026-08-29 - тема.md" assert result.digest.path == config.dir / "2026-08-29 - тема.md"
assert result.digest.source == "[[2026-08-20 - тема чата]]" assert result.digest.source == "[[2026-08-20 - тема чата]]"
index = config.index.read_text(encoding="utf-8") index = config.index.read_text(encoding="utf-8")
assert index.startswith("# индекс") assert index.startswith("# index")
assert "- 2026-08-29 [[2026-08-20 - тема чата]] → [[2026-08-29 - тема]]" in index assert "- 2026-08-29 [[2026-08-20 - тема чата]] → [[2026-08-29 - тема]]" in index
assert result.text.count("\n") == 2 and not result.trimmed assert result.text.count("\n") == 2 and not result.trimmed
closed = await world.conversations.get(chat.external_id) closed = await world.conversations.get(chat.external_id)
@@ -236,9 +236,9 @@ async def test_distill_writes_the_digest_indexes_it_and_merges_short(
fork = await world.conversations.get(result.fork.external_id) fork = await world.conversations.get(result.fork.external_id)
assert fork.kind == "fork" and fork.agent_name == "x" and fork.status == "closed" assert fork.kind == "fork" and fork.agent_name == "x" and fork.status == "closed"
items = await world.conversations.queue.recent(master.id) items = await world.conversations.queue.recent(master.id)
assert items[0].origin == "выжимка" assert items[0].origin == "digest"
assert items[0].text.startswith( assert items[0].text.startswith(
"Закрыт глубокий чат [[2026-08-20 - тема чата]], выжимка [[2026-08-29 - тема]]." "Deep chat [[2026-08-20 - тема чата]] closed, digest [[2026-08-29 - тема]]."
) )
assert items[0].text.endswith(result.text) assert items[0].text.endswith(result.text)
forked = ScriptedClient.instances[-1] forked = ScriptedClient.instances[-1]
@@ -270,8 +270,8 @@ async def test_memory_off_merges_without_a_file(world: World) -> None:
assert (await world.conversations.get(chat.external_id)).status == "closed" assert (await world.conversations.get(chat.external_id)).status == "closed"
items = await world.conversations.queue.recent(master.id) items = await world.conversations.queue.recent(master.id)
assert len(items) == 1 and items[0].text.endswith(result.text) assert len(items) == 1 and items[0].text.endswith(result.text)
assert ", выжимка" not in items[0].text assert ", digest" not in items[0].text
assert "файл не пиши" in ScriptedClient.instances[-1].prompts[0] assert "write no file" in ScriptedClient.instances[-1].prompts[0]
async def test_bad_frontmatter_and_long_merge_are_reported(world: World) -> None: async def test_bad_frontmatter_and_long_merge_are_reported(world: World) -> None:
@@ -294,13 +294,13 @@ async def test_bad_frontmatter_and_long_merge_are_reported(world: World) -> None
def test_check_digest_rejects_what_is_not_a_digest(tmp_path: Path) -> None: def test_check_digest_rejects_what_is_not_a_digest(tmp_path: Path) -> None:
config = Distiller(agent="x", dir=tmp_path, index=tmp_path / "i.md") config = Distiller(agent="x", dir=tmp_path, index=tmp_path / "i.md")
path = tmp_path / "d.md" path = tmp_path / "d.md"
path.write_text("---\ntype: выжимка\nsource: ''\ndate: 2026-08-29\n---\nx\n") path.write_text("---\ntype: digest\nsource: ''\ndate: 2026-08-29\n---\nx\n")
assert check_digest(path, config) == "`source` пустой" assert check_digest(path, config) == "`source` is empty"
path.write_text("---\ntype: выжимка\nsource: '[[a]]'\ndate: вчера\n---\nx\n") path.write_text("---\ntype: digest\nsource: '[[a]]'\ndate: вчера\n---\nx\n")
assert "`date`" in check_digest(path, config) assert "`date`" in check_digest(path, config)
path.write_text("---\ntype: выжимка\nsource: '[[a]]'\ndate: 2026-08-29\n---\n\n") path.write_text("---\ntype: digest\nsource: '[[a]]'\ndate: 2026-08-29\n---\n\n")
assert check_digest(path, config) == "тело пустое" assert check_digest(path, config) == "empty body"
path.write_text("---\ntype: выжимка\nsource: '[[a]]'\ndate: 2026-08-29\n---\nx\n") path.write_text("---\ntype: digest\nsource: '[[a]]'\ndate: 2026-08-29\n---\nx\n")
digest = check_digest(path, config) digest = check_digest(path, config)
assert not isinstance(digest, str) and digest.date.isoformat() == "2026-08-29" assert not isinstance(digest, str) and digest.date.isoformat() == "2026-08-29"
assert trim_summary("a\n\nb\nc\nd\ne\nf") == ("a\nb\nc\nd\ne", True) assert trim_summary("a\n\nb\nc\nd\ne\nf") == ("a\nb\nc\nd\ne", True)
@@ -347,7 +347,7 @@ async def test_close_chat_tool_closes_after_the_reply(world: World) -> None:
assert row.flags["closed_reason"] == "close_chat" assert row.flags["closed_reason"] == "close_chat"
assert row.flags["digest"] is not None assert row.flags["digest"] is not None
items = await world.conversations.queue.recent(master.id) items = await world.conversations.queue.recent(master.id)
assert items[0].origin == "выжимка" assert items[0].origin == "digest"
async def test_idle_picks_quiet_chats_after_launch_at_most_limit(world: World) -> None: async def test_idle_picks_quiet_chats_after_launch_at_most_limit(world: World) -> None:
@@ -412,8 +412,8 @@ async def test_line_cap_bounces_a_long_rewrite_and_asks_to_shorten(
client = ScriptedClient.instances[-1] client = ScriptedClient.instances[-1]
assert len(client.prompts) == 2 assert len(client.prompts) == 2
assert "70 строк при потолке 60" in client.prompts[1] assert "70 lines against a cap of 60" in client.prompts[1]
assert "[инжект: потолок" in client.prompts[1] assert "[inject: cap" in client.prompts[1]
assert state.read_text(encoding="utf-8").count("\n") == 10 - 1 assert state.read_text(encoding="utf-8").count("\n") == 10 - 1
row = await world.conversations.get(job.external_id) row = await world.conversations.get(job.external_id)
assert row.flags["line_cap_attempts"] == 1 assert row.flags["line_cap_attempts"] == 1
+13 -10
View File
@@ -5,12 +5,15 @@ from pathlib import Path
from test_conversations import ScriptedClient, World, world from test_conversations import ScriptedClient, World, world
from beaver_gateway.core.conversations import UserSaid from beaver_gateway.conversations.service import UserSaid
from beaver_gateway.core.envelope import HEADER, Envelope, RecallContext, render from beaver_gateway.conversations.envelope import Envelope, RecallContext, render
from beaver_gateway.core.watch import Change, VaultWatch, WatchRules from beaver_gateway.conversations.texts import EnvelopeTexts
from beaver_gateway.vault.watch import Change, VaultWatch, WatchRules
__all__ = ["world"] __all__ = ["world"]
HEADER = EnvelopeTexts().header
RULES = WatchRules( RULES = WatchRules(
full=("дни/{today}.md",), full=("дни/{today}.md",),
names=("дни/*", "люди/*", "мета/бобер/*"), names=("дни/*", "люди/*", "мета/бобер/*"),
@@ -88,11 +91,11 @@ def test_envelope_respects_ceilings_and_names_only_window() -> None:
lines = text.splitlines() lines = text.splitlines()
assert lines[0] == HEADER assert lines[0] == HEADER
assert "(Warsaw)" in lines[1] assert "(Warsaw)" in lines[1]
assert lines[2].startswith("vault, изменено со старта: ") assert lines[2].startswith("vault, changed since start: ")
assert f"дни/{today}.md (+200)" in lines[2] assert f"дни/{today}.md (+200)" in lines[2]
assert "люди/Петя.md (+50)" in lines[2] assert "люди/Петя.md (+50)" in lines[2]
assert sum(1 for line in lines if line.startswith("+ ")) == 31 assert sum(1 for line in lines if line.startswith("+ ")) == 31
assert "+ … ещё 170" in lines assert "+ … 170 more" in lines
assert len(lines) <= 120 assert len(lines) <= 120
append(diary, "- ещё одна\n") append(diary, "- ещё одна\n")
watch.note(diary) watch.note(diary)
@@ -122,8 +125,8 @@ def test_render_hits_total_ceiling() -> None:
) )
lines = text.splitlines() lines = text.splitlines()
assert len(lines) <= 120 assert len(lines) <= 120
assert "… (потолок конверта)" in lines assert "… (envelope cap)" in lines
assert lines[1] == "время: 2026-08-26 13:04 (Warsaw)" assert lines[1] == "time: 2026-08-26 13:04 (Warsaw)"
async def test_master_turn_gets_envelope_after_text_and_before_injects( async def test_master_turn_gets_envelope_after_text_and_before_injects(
@@ -142,12 +145,12 @@ async def test_master_turn_gets_envelope_after_text_and_before_injects(
assert head == "hello" assert head == "hello"
assert rest.startswith(HEADER) assert rest.startswith(HEADER)
assert "люди/Прохор.md (+1)" in rest assert "люди/Прохор.md (+1)" in rest
assert rest.index("[инжекты") > rest.index(HEADER) assert rest.index("[injects") > rest.index(HEADER)
branch = await world.conversations.spawn( branch = await world.conversations.spawn(
kind="branch", parent=master, seed="brief", text="do X" kind="branch", parent=master, seed="brief", text="do X"
) )
await world.settle(branch, 1) await world.settle(branch, 1)
assert "[конверт" not in ScriptedClient.instances[-1].prompts[0] assert "[envelope" not in ScriptedClient.instances[-1].prompts[0]
await asyncio.sleep(0) await asyncio.sleep(0)
@@ -187,7 +190,7 @@ async def test_recall_lines_follow_the_vault_block_and_reach_branches(
branch_prompt = ScriptedClient.instances[-1].prompts[0] branch_prompt = ScriptedClient.instances[-1].prompts[0]
assert "про Прохор подробнее\n\n" + HEADER in branch_prompt assert "про Прохор подробнее\n\n" + HEADER in branch_prompt
assert branch_prompt.endswith(f"{HEADER}\n👤 Прохор - карточка `люди/Прохор.md`") assert branch_prompt.endswith(f"{HEADER}\n👤 Прохор - карточка `люди/Прохор.md`")
assert "vault, изменено" not in branch_prompt assert "vault, changed" not in branch_prompt
# the first branch turn carries the seed line above the text # the first branch turn carries the seed line above the text
assert seen[-1][0] == "branch" and seen[-1][1].endswith("про Прохор подробнее") assert seen[-1][0] == "branch" and seen[-1][1].endswith("про Прохор подробнее")
assert noted[-1].kind == "branch" assert noted[-1].kind == "branch"
+3 -3
View File
@@ -5,8 +5,8 @@ import pytest
from fastmcp import FastMCP from fastmcp import FastMCP
from fastmcp.tools.base import ToolResult from fastmcp.tools.base import ToolResult
from beaver_gateway.core import redact as redact_mod from beaver_gateway.security import redact as redact_mod
from beaver_gateway.core.redact import redact from beaver_gateway.security.redact import redact
from beaver_gateway.mcp.internal_app import build_internal_app from beaver_gateway.mcp.internal_app import build_internal_app
from beaver_gateway.mcp.redacting import RedactingMiddleware from beaver_gateway.mcp.redacting import RedactingMiddleware
from beaver_gateway.mcp.types import McpServer from beaver_gateway.mcp.types import McpServer
@@ -179,7 +179,7 @@ async def test_gateway_own_tools_are_filtered_too() -> None:
# in-process — so they carry their own wrapper. # in-process — so they carry their own wrapper.
from claude_agent_sdk import SdkMcpTool from claude_agent_sdk import SdkMcpTool
from beaver_gateway.core.gateway_tools import _redacting from beaver_gateway.conversations.tools import _redacting
async def handler(_args: dict[str, object]) -> dict[str, object]: async def handler(_args: dict[str, object]) -> dict[str, object]:
return {"content": [{"type": "text", "text": KOMODO_DEPLOY}]} return {"content": [{"type": "text", "text": KOMODO_DEPLOY}]}
+1 -1
View File
@@ -2,7 +2,7 @@ from pathlib import Path
import pytest import pytest
from beaver_gateway.core.policy import Deny, ToolCall, brief, evaluate, hook_output from beaver_gateway.agents.policy import Deny, ToolCall, brief, evaluate, hook_output
def call(tool: str, **tool_input) -> ToolCall: def call(tool: str, **tool_input) -> ToolCall:
+2 -2
View File
@@ -5,8 +5,8 @@ import logging
import pytest import pytest
from beaver_gateway.core import redact as redact_mod from beaver_gateway.security import redact as redact_mod
from beaver_gateway.core.redact import ( from beaver_gateway.security.redact import (
RedactFilter, RedactFilter,
RedactingFormatter, RedactingFormatter,
env_secrets, env_secrets,
+17 -13
View File
@@ -3,8 +3,12 @@ from datetime import UTC, date, datetime, timedelta
from test_conversations import ScriptedClient, StubFrontend, World, world from test_conversations import ScriptedClient, StubFrontend, World, world
from beaver_gateway.core.conversations import ConversationTexts from beaver_gateway.conversations.service import ConversationTexts
from beaver_gateway.core.rotation import HandoutContext, Rotation, RotationPolicy from beaver_gateway.conversations.rotation import (
HandoutContext,
Rotation,
RotationPolicy,
)
from beaver_gateway.storage.models import Conversation, Usage from beaver_gateway.storage.models import Conversation, Usage
__all__ = ["world"] __all__ = ["world"]
@@ -28,7 +32,7 @@ def master(
def test_night_rule_needs_silence_and_a_master_from_before_four() -> None: def test_night_rule_needs_silence_and_a_master_from_before_four() -> None:
now = datetime(2026, 8, 29, 2, 30, tzinfo=UTC) now = datetime(2026, 8, 29, 2, 30, tzinfo=UTC)
quiet = master(created_ago=timedelta(hours=20), silence=timedelta(hours=4), now=now) quiet = master(created_ago=timedelta(hours=20), silence=timedelta(hours=4), now=now)
assert POLICY.reason(quiet, now=now, context_tokens=0) == "ночь" assert POLICY.reason(quiet, now=now, context_tokens=0) == "night"
active = master( active = master(
created_ago=timedelta(hours=20), silence=timedelta(minutes=5), now=now created_ago=timedelta(hours=20), silence=timedelta(minutes=5), now=now
) )
@@ -49,13 +53,13 @@ def test_age_and_context_rules_need_thirty_minutes_of_silence() -> None:
old = master( old = master(
created_ago=timedelta(hours=37), silence=timedelta(minutes=31), now=now created_ago=timedelta(hours=37), silence=timedelta(minutes=31), now=now
) )
assert POLICY.reason(old, now=now, context_tokens=0) == "возраст" assert POLICY.reason(old, now=now, context_tokens=0) == "age"
busy = master( busy = master(
created_ago=timedelta(hours=37), silence=timedelta(minutes=5), now=now created_ago=timedelta(hours=37), silence=timedelta(minutes=5), now=now
) )
assert POLICY.reason(busy, now=now, context_tokens=0) is None assert POLICY.reason(busy, now=now, context_tokens=0) is None
big = master(created_ago=timedelta(hours=2), silence=timedelta(minutes=31), now=now) big = master(created_ago=timedelta(hours=2), silence=timedelta(minutes=31), now=now)
assert POLICY.reason(big, now=now, context_tokens=90_000) == "транскрипт" assert POLICY.reason(big, now=now, context_tokens=90_000) == "context"
assert POLICY.reason(big, now=now, context_tokens=70_000) is None assert POLICY.reason(big, now=now, context_tokens=70_000) is None
@@ -73,7 +77,7 @@ async def test_rotation_does_not_touch_a_master_mid_turn(world: World) -> None:
await world.conversations.post(conv, "working") await world.conversations.post(conv, "working")
await asyncio.sleep(0.2) await asyncio.sleep(0.2)
rotation = Rotation(world.conversations, POLICY) rotation = Rotation(world.conversations, POLICY)
assert await rotation.rotate(conv, "возраст") is None assert await rotation.rotate(conv, "age") is None
assert (await world.conversations.get(conv.external_id)).status == "open" assert (await world.conversations.get(conv.external_id)).status == "open"
assert len(await world.conversations.find(kind="master")) == 1 assert len(await world.conversations.find(kind="master")) == 1
ScriptedClient.hold.set() ScriptedClient.hold.set()
@@ -119,7 +123,7 @@ async def test_rotation_order_handout_close_marks_moves_and_new_day(
await world.conversations.inject(old, "later", urgency="normal", origin="крон") await world.conversations.inject(old, "later", urgency="normal", origin="крон")
old_client = ScriptedClient.instances[0] old_client = ScriptedClient.instances[0]
new = await Rotation(world.conversations, POLICY).rotate(old, "ночь") new = await Rotation(world.conversations, POLICY).rotate(old, "night")
assert new is not None and new.kind == "master" assert new is not None and new.kind == "master"
assert handouts[0].day == date(2026, 8, 27) assert handouts[0].day == date(2026, 8, 27)
assert old_client.prompts[-1] == "напиши хендаут за 2026-08-27" assert old_client.prompts[-1] == "напиши хендаут за 2026-08-27"
@@ -138,10 +142,10 @@ async def test_rotation_order_handout_close_marks_moves_and_new_day(
await world.settle(new, 1) await world.settle(new, 1)
new_client = ScriptedClient.instances[-1] new_client = ScriptedClient.instances[-1]
prompt = new_client.prompts[0] prompt = new_client.prompts[0]
assert prompt.startswith("[сид: morning] master") assert prompt.startswith("[seed: morning] master")
assert "[инжект: ротация" in prompt assert "[inject: rotation" in prompt
assert "Новый день 20" in prompt and "(ночь)" in prompt assert "Новый день 20" in prompt and "(night)" in prompt
assert "переехало из старого мастера: 1" in prompt assert "1 queued injects moved over" in prompt
moved = await world.conversations.queue.pending(new.id) moved = await world.conversations.queue.pending(new.id)
assert [(i.priority, i.text) for i in moved] == [("normal", "later")] assert [(i.priority, i.text) for i in moved] == [("normal", "later")]
assert (await world.conversations.find(kind="master", status="open")) == [ assert (await world.conversations.find(kind="master", status="open")) == [
@@ -171,11 +175,11 @@ async def test_due_uses_last_usage_row_for_context_size(world: World) -> None:
assert await rotation.due() == [] assert await rotation.due() == []
rotation = Rotation(world.conversations, RotationPolicy(max_context_tokens=5)) rotation = Rotation(world.conversations, RotationPolicy(max_context_tokens=5))
(pair,) = await rotation.due() (pair,) = await rotation.due()
assert pair[0].id == conv.id and pair[1] == "транскрипт" assert pair[0].id == conv.id and pair[1] == "context"
def test_context_of_prefers_last_call_and_averages_old_rows() -> None: def test_context_of_prefers_last_call_and_averages_old_rows() -> None:
from beaver_gateway.core.conversations import context_of from beaver_gateway.conversations.service import context_of
from beaver_gateway.storage.models import Usage from beaver_gateway.storage.models import Usage
fresh = Usage( fresh = Usage(
+6 -6
View File
@@ -7,10 +7,10 @@ import frontmatter
import httpx import httpx
import pytest import pytest
from beaver_gateway.core.auth import TokenStore from beaver_gateway.security.auth import TokenStore
from beaver_gateway.core.conversation_store import load_messages, rewrite_messages from beaver_gateway.frontends.markdown.history import load_messages, rewrite_messages
from beaver_gateway.core.gateway_tools import _tools from beaver_gateway.conversations.tools import _tools
from beaver_gateway.core.registry import McpRegistry from beaver_gateway.app import McpRegistry
from beaver_gateway.frontends.anthropic import AnthropicMessagesFrontend from beaver_gateway.frontends.anthropic import AnthropicMessagesFrontend
from beaver_gateway.frontends.api import ApiFrontend from beaver_gateway.frontends.api import ApiFrontend
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
@@ -127,10 +127,10 @@ async def test_api_spawn_deep_lands_in_vault(stack: Stack) -> None:
) )
assert r.status_code == 400 assert r.status_code == 400
path = stack.vault / rel path = stack.vault / rel
text = await stack.wait_file(path, "ok:[сид: brief]") text = await stack.wait_file(path, "ok:[seed: brief]")
post = frontmatter.loads(text) post = frontmatter.loads(text)
assert post.metadata == {"agent": "d", "conversation_id": body["id"]} assert post.metadata == {"agent": "d", "conversation_id": body["id"]}
assert post.content.startswith("### User:\n\n[сид: brief] deep «Тема», ") assert post.content.startswith("### User:\n\n[seed: brief] deep «Тема», ")
assert post.content.rstrip().endswith("### User:") assert post.content.rstrip().endswith("### User:")
+3 -3
View File
@@ -13,8 +13,8 @@ from httpx import ASGITransport, AsyncClient
from pgqueuer import PsycopgDriver, Queries from pgqueuer import PsycopgDriver, Queries
from test_conversations import World from test_conversations import World
from beaver_gateway.core.conversations import parse_at from beaver_gateway.conversations.service import parse_at
from beaver_gateway.core.scheduler import Budget, Job, JobRun, Scheduler, next_run from beaver_gateway.jobs.scheduler import Budget, Job, JobRun, Scheduler, next_run
from beaver_gateway.storage.models import RateLimit from beaver_gateway.storage.models import RateLimit
DATABASE_URL = os.environ.get("TEST_DATABASE_URL") DATABASE_URL = os.environ.get("TEST_DATABASE_URL")
@@ -165,7 +165,7 @@ async def test_schedule_survives_a_restart(world: World, pg: Pg) -> None:
world.conversations._normal_window = 0.05 # noqa: SLF001 world.conversations._normal_window = 0.05 # noqa: SLF001
await until(lambda: len(ScriptedClient_prompts(world)) == 1) await until(lambda: len(ScriptedClient_prompts(world)) == 1)
prompt = ScriptedClient_prompts(world)[0] prompt = ScriptedClient_prompts(world)[0]
assert prompt.startswith("[инжект: schedule") assert prompt.startswith("[inject: schedule")
assert prompt.endswith("push X") assert prompt.endswith("push X")
assert await world.statuses(conv) == [("wake", "done")] assert await world.statuses(conv) == [("wake", "done")]
assert await world.conversations.schedules(conv) == [] assert await world.conversations.schedules(conv) == []
+9 -9
View File
@@ -10,8 +10,8 @@ from aiogram.methods import SendMessage
from aiogram.types import Update from aiogram.types import Update
from claude_agent_sdk import PermissionResultAllow, PermissionResultDeny from claude_agent_sdk import PermissionResultAllow, PermissionResultDeny
from beaver_gateway.core.registry import McpRegistry from beaver_gateway.app import McpRegistry
from beaver_gateway.core.transcript import build_entries from beaver_gateway.backends.transcript import build_entries
from beaver_gateway.frontends.base import GatewayRuntime from beaver_gateway.frontends.base import GatewayRuntime
from beaver_gateway.frontends.telegram import Attachments, TelegramFrontend from beaver_gateway.frontends.telegram import Attachments, TelegramFrontend
from beaver_gateway.frontends.telegram.render import ( from beaver_gateway.frontends.telegram.render import (
@@ -279,7 +279,7 @@ async def stack() -> Stack:
async def test_general_is_master_and_reply_has_no_thread(stack: Stack) -> None: async def test_general_is_master_and_reply_has_no_thread(stack: Stack) -> None:
stack.bot.message("hi") stack.bot.message("hi")
reply = await stack.until(lambda: stack.sent_with("hi"), what="reply") reply = await stack.until(lambda: stack.sent_with("hi"), what="reply")
assert reply["text"].startswith("ok:[сид: clean] master") assert reply["text"].startswith("ok:[seed: clean] master")
assert reply["thread"] == 901 assert reply["thread"] == 901
assert reply["parse_mode"] == "HTML" assert reply["parse_mode"] == "HTML"
assert stack.bot.topics == ["🦫 General"] assert stack.bot.topics == ["🦫 General"]
@@ -316,8 +316,8 @@ async def test_new_topic_becomes_morning_branch_and_replies_in_thread(
) )
assert branch.parent_id == master.id assert branch.parent_id == master.id
prompts = [p for c in ScriptedClient.instances for p in c.prompts] prompts = [p for c in ScriptedClient.instances for p in c.prompts]
seed = next(p for p in prompts if p.startswith("[сид: morning] branch «план»")) seed = next(p for p in prompts if p.startswith("[seed: morning] branch «план»"))
assert "Хендаут не приехал." in seed assert "No handout arrived." in seed
assert seed.endswith("hello topic") assert seed.endswith("hello topic")
assert stack.bot.drafts[-1]["message_thread_id"] == 7 assert stack.bot.drafts[-1]["message_thread_id"] == 7
@@ -334,7 +334,7 @@ async def test_first_message_without_service_message_seeds_with_text(
client = next( client = next(
c for c in ScriptedClient.instances if c.prompts and "сразу" in c.prompts[0] c for c in ScriptedClient.instances if c.prompts and "сразу" in c.prompts[0]
) )
assert client.prompts[0].startswith("[сид: morning]") assert client.prompts[0].startswith("[seed: morning]")
assert client.prompts[0].endswith("сразу текстом") assert client.prompts[0].endswith("сразу текстом")
assert len(client.prompts) == 1 assert len(client.prompts) == 1
@@ -432,7 +432,7 @@ async def test_question_becomes_buttons_and_callback_answers(stack: Stack) -> No
stack.bot.callback(f"q:{pending[0]}:0:1", question["message_id"]) stack.bot.callback(f"q:{pending[0]}:0:1", question["message_id"])
assert await asyncio.wait_for(asking, 5) == "Красный" assert await asyncio.wait_for(asking, 5) == "Красный"
assert stack.world.conversations.answer_text("Красный") == ( assert stack.world.conversations.answer_text("Красный") == (
"Пользователь ответил: Красный" "The user answered: Красный"
) )
await stack.until( await stack.until(
lambda: any("✅ Красный" in e.get("text", "") for e in stack.bot.edits), lambda: any("✅ Красный" in e.get("text", "") for e in stack.bot.edits),
@@ -454,7 +454,7 @@ async def test_question_timeout_renders_text_and_free_text_answers(
payload = {"questions": [{"header": "Q", "question": "Сколько?", "options": []}]} payload = {"questions": [{"header": "Q", "question": "Сколько?", "options": []}]}
result = await stack.world.conversations.ask(master.external_id, payload) result = await stack.world.conversations.ask(master.external_id, payload)
assert result is None assert result is None
assert "не ответил" in stack.world.conversations.answer_text(None) assert "did not answer" in stack.world.conversations.answer_text(None)
await stack.until( await stack.until(
lambda: any("время вышло" in e.get("text", "") for e in stack.bot.edits), lambda: any("время вышло" in e.get("text", "") for e in stack.bot.edits),
what="timeout edit", what="timeout edit",
@@ -510,7 +510,7 @@ async def test_commands_status_merge_and_new(stack: Stack) -> None:
stack.bot.message("первое в новый топик", thread=902) stack.bot.message("первое в новый топик", thread=902)
reply = await stack.until(lambda: stack.sent_with("первое в новый топик"), what="r") reply = await stack.until(lambda: stack.sent_with("первое в новый топик"), what="r")
assert reply["thread"] == 902 assert reply["thread"] == 902
assert reply["text"].startswith("ok:[сид: morning] branch «отчёт»") assert reply["text"].startswith("ok:[seed: morning] branch «отчёт»")
assert await stack.tg.mark_topic(child) assert await stack.tg.mark_topic(child)
assert stack.bot.topics[-1] == "edit:902:✅ отчёт" assert stack.bot.topics[-1] == "edit:902:✅ отчёт"
+1 -1
View File
@@ -1,4 +1,4 @@
from beaver_gateway.core.transcript import ( from beaver_gateway.backends.transcript import (
CLI_VERSION, CLI_VERSION,
build_entries, build_entries,
messages_from_entries, messages_from_entries,