refactor: split flat core into capability packages, layer the conversations service, English defaults for every model-facing text
This commit is contained in:
@@ -17,9 +17,9 @@ from pathlib import Path # noqa: TC003 - pydantic runtime
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
from beaver_gateway.agents.base import BaseAgent
|
||||
from beaver_gateway.core.kinds import KINDS, Kind
|
||||
from beaver_gateway.core.policy import PolicyRule # noqa: TC001 - pydantic runtime
|
||||
from beaver_gateway.core.prompt import PromptSource # noqa: TC001 - pydantic runtime
|
||||
from beaver_gateway.agents.policy import PolicyRule # noqa: TC001 - pydantic runtime
|
||||
from beaver_gateway.agents.prompts import PromptSource # noqa: TC001 - pydantic runtime
|
||||
from beaver_gateway.conversations.kinds import KINDS, Kind
|
||||
|
||||
__all__ = ["ClaudeAgent", "ClaudeOptions", "Prompts", "SkillSets"]
|
||||
|
||||
|
||||
@@ -21,7 +21,7 @@ if TYPE_CHECKING:
|
||||
|
||||
__all__ = ["PromptSource", "assemble"]
|
||||
|
||||
_log = logging.getLogger("beaver_gateway.core.prompt")
|
||||
_log = logging.getLogger("beaver_gateway.agents.prompts")
|
||||
|
||||
PromptSource = str | Path | tuple[str, str | Path]
|
||||
|
||||
@@ -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,6 +1,6 @@
|
||||
"""Backend adapters.
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Backend protocol.
|
||||
|
||||
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
|
||||
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
|
||||
the backend may pin a live session to), ``session_id`` (backend session to
|
||||
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).
|
||||
"""
|
||||
|
||||
@@ -26,7 +26,7 @@ if TYPE_CHECKING:
|
||||
from anthropic.types import MessageParam
|
||||
|
||||
from beaver_gateway.agents.base import BaseAgent
|
||||
from beaver_gateway.core.events import MessageStreamEvent
|
||||
from beaver_gateway.events.stream import MessageStreamEvent
|
||||
|
||||
|
||||
class Backend(Protocol):
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
One :class:`ClaudeSdkBackend` per :class:`ClaudeAgent`. A live session is
|
||||
one ``ClaudeSDKClient`` (one claude subprocess) and runs one turn at a
|
||||
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
|
||||
the caller passes one or by a text-only fingerprint of ``messages[:-1]``
|
||||
for stateless callers (``/v1/messages``). Without a live session the
|
||||
@@ -67,9 +67,18 @@ from claude_agent_sdk import (
|
||||
project_key_for_directory,
|
||||
)
|
||||
|
||||
from beaver_gateway.core import policy as policy_mod
|
||||
from beaver_gateway.core import prompt as prompt_assembly
|
||||
from beaver_gateway.core.events import (
|
||||
from beaver_gateway.agents import policy as policy_mod
|
||||
from beaver_gateway.agents import prompts as prompt_assembly
|
||||
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,
|
||||
build_content_block_stop,
|
||||
build_input_json_delta,
|
||||
@@ -83,15 +92,6 @@ from beaver_gateway.core.events import (
|
||||
build_thinking_delta,
|
||||
build_tool_use_block_start,
|
||||
)
|
||||
from beaver_gateway.core.kinds import as_kind
|
||||
from beaver_gateway.core.sessions import Session, SessionClient, SessionPool
|
||||
from beaver_gateway.core.transcript import (
|
||||
build_entries,
|
||||
close_open_tool_uses,
|
||||
fingerprint,
|
||||
text_of,
|
||||
)
|
||||
from beaver_gateway.core.turn_capture import TurnCapture, TurnUsage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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.claude import ClaudeAgent
|
||||
from beaver_gateway.core.events import MessageStreamEvent
|
||||
from beaver_gateway.core.kinds import Kind
|
||||
from beaver_gateway.conversations.kinds import Kind
|
||||
from beaver_gateway.events.stream import MessageStreamEvent
|
||||
|
||||
|
||||
_log = logging.getLogger("beaver_gateway.backends.claude_sdk")
|
||||
@@ -281,13 +281,17 @@ class ClaudeSdkBackend:
|
||||
def live(self, key: str) -> Session | None:
|
||||
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."""
|
||||
key = self._store_key(session_id)
|
||||
entries = await self._store.load(cast("Any", key))
|
||||
if not entries:
|
||||
return 0
|
||||
fixes = close_open_tool_uses(cast("list[Mapping[str, Any]]", entries))
|
||||
fixes = close_open_tool_uses(
|
||||
cast("list[Mapping[str, Any]]", entries), text=text
|
||||
)
|
||||
if fixes:
|
||||
await self._store.append(cast("Any", key), cast("Any", fixes))
|
||||
_log.warning(
|
||||
|
||||
@@ -42,7 +42,7 @@ from raycast_api import Message as RaycastMessage
|
||||
from raycast_api import RemoteTool, Tool, ToolCall
|
||||
|
||||
from beaver_gateway.agents.raycast import RaycastAgent
|
||||
from beaver_gateway.core.events import (
|
||||
from beaver_gateway.events.stream import (
|
||||
StopReason,
|
||||
build_content_block_stop,
|
||||
build_input_json_delta,
|
||||
@@ -65,7 +65,7 @@ if TYPE_CHECKING:
|
||||
from raycast_api import ChatStreamChunk, Client
|
||||
|
||||
from beaver_gateway.agents.base import BaseAgent
|
||||
from beaver_gateway.core.events import MessageStreamEvent
|
||||
from beaver_gateway.events.stream import MessageStreamEvent
|
||||
else:
|
||||
from collections.abc import Mapping
|
||||
|
||||
|
||||
@@ -28,7 +28,7 @@ if TYPE_CHECKING:
|
||||
|
||||
__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] = {
|
||||
"master": None,
|
||||
@@ -296,7 +296,7 @@ def _zero_usage() -> dict[str, Any]:
|
||||
# ---- repair, windows, projections ---------------------------------------
|
||||
|
||||
_PROMPT_TYPES = ("user", "assistant")
|
||||
_INTERRUPTED = "прервано"
|
||||
_INTERRUPTED = "interrupted"
|
||||
|
||||
|
||||
def open_tool_uses(
|
||||
+13
-505
@@ -1,533 +1,41 @@
|
||||
"""Process entrypoint.
|
||||
|
||||
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.
|
||||
"""
|
||||
"""Process entrypoint: logging, signals, ``.env``, the config, then ``app.run``."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import functools
|
||||
import logging
|
||||
import signal
|
||||
from contextlib import AsyncExitStack
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import psycopg
|
||||
import uvicorn
|
||||
import uvloop
|
||||
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.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.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 import app, config
|
||||
from beaver_gateway.security.redact import install as install_redaction
|
||||
from beaver_gateway.security.redact import load_secrets as load_secrets_to_mask
|
||||
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:
|
||||
"""Sync wrapper: uvloop loop factory + asyncio.run."""
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s"
|
||||
)
|
||||
install_redaction()
|
||||
_install_sigterm_handler()
|
||||
asyncio.run(_async_main(), loop_factory=uvloop.new_event_loop)
|
||||
_sigterm_as_interrupt()
|
||||
asyncio.run(_run(), loop_factory=uvloop.new_event_loop)
|
||||
|
||||
|
||||
def _install_sigterm_handler() -> None:
|
||||
"""Turn SIGTERM into a normal interpreter exit.
|
||||
async def _run() -> None:
|
||||
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:
|
||||
raise KeyboardInterrupt
|
||||
|
||||
with contextlib.suppress(ValueError, OSError):
|
||||
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``.
|
||||
|
||||
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.
|
||||
"""
|
||||
"""``Gateway`` - the one object a setup's ``config.py`` assembles - and its loader."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from croniter import croniter
|
||||
@@ -22,18 +11,58 @@ from croniter import croniter
|
||||
from beaver_gateway.agents.base import BaseAgent, ExposedMcp
|
||||
from beaver_gateway.agents.claude import ClaudeAgent
|
||||
from beaver_gateway.agents.raycast import RaycastAgent
|
||||
from beaver_gateway.core.conversations import ConversationTexts
|
||||
from beaver_gateway.core.registry import Gateway
|
||||
from beaver_gateway.core.scheduler import Job
|
||||
from beaver_gateway.conversations.texts import ConversationTexts
|
||||
from beaver_gateway.frontends.base import Frontend
|
||||
from beaver_gateway.jobs.scheduler import Job
|
||||
from beaver_gateway.mcp.types import HttpMcp, McpServer, PythonToolMcp, StdioMcp
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
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):
|
||||
"""User config file is missing, unreadable, or structurally wrong."""
|
||||
pass
|
||||
|
||||
|
||||
_PUBLIC_NAMES: dict[str, Any] = {
|
||||
@@ -49,7 +78,7 @@ _McpInstance = StdioMcp | HttpMcp | PythonToolMcp
|
||||
|
||||
|
||||
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:
|
||||
source = path.read_text(encoding="utf-8")
|
||||
except FileNotFoundError as exc:
|
||||
@@ -60,12 +89,11 @@ def load(path: Path) -> Gateway:
|
||||
raise ConfigError(msg) from exc
|
||||
|
||||
code = compile(source, str(path), "exec")
|
||||
# Siblings of the config (``policy.py``, ``mcps/``) import by name.
|
||||
parent = str(path.resolve().parent)
|
||||
if parent not in sys.path:
|
||||
sys.path.insert(0, parent)
|
||||
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:
|
||||
gw = namespace["gateway"]
|
||||
@@ -0,0 +1 @@
|
||||
"""Conversations: rows, queue, seeds, turns, questions, closing, rotation, envelope."""
|
||||
@@ -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
|
||||
dir: Path
|
||||
index: Path
|
||||
type: str = "выжимка"
|
||||
type: str = "digest"
|
||||
"""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)
|
||||
@@ -150,13 +150,13 @@ def check_digest(path: Path, digests: Distiller) -> Digest | str:
|
||||
try:
|
||||
post = frontmatter.load(str(path))
|
||||
except (OSError, ValueError) as exc:
|
||||
return f"не читается: {exc}"
|
||||
return f"unreadable: {exc}"
|
||||
meta = post.metadata
|
||||
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")
|
||||
if not isinstance(source, str) or not source.strip():
|
||||
return "`source` пустой"
|
||||
return "`source` is empty"
|
||||
when = meta.get("date")
|
||||
if isinstance(when, datetime):
|
||||
when = when.date()
|
||||
@@ -164,11 +164,11 @@ def check_digest(path: Path, digests: Distiller) -> Digest | str:
|
||||
try:
|
||||
when = date.fromisoformat(when.strip())
|
||||
except ValueError:
|
||||
return f"`date` не дата: {when!r}"
|
||||
return f"`date` is not a date: {when!r}"
|
||||
if not isinstance(when, date):
|
||||
return "`date` отсутствует"
|
||||
return "`date` is missing"
|
||||
if not post.content.strip():
|
||||
return "тело пустое"
|
||||
return "empty body"
|
||||
return Digest(path=path, source=source.strip(), date=when)
|
||||
|
||||
|
||||
+31
-31
@@ -1,39 +1,27 @@
|
||||
"""The envelope (§3.3): a background block after the user's text.
|
||||
|
||||
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.
|
||||
"""
|
||||
"""The envelope: a background block under the user's text - clock, changes, recall."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from beaver_gateway.conversations.texts import EnvelopeTexts
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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"]
|
||||
|
||||
_log = logging.getLogger(__name__)
|
||||
|
||||
HEADER = (
|
||||
"[конверт - фоновый сигнал, не обращение; "
|
||||
"реагируй, только если относится к вопросу]"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RecallContext:
|
||||
"""What the setup's ``recall`` hook sees: the user's text and where it landed."""
|
||||
|
||||
text: str
|
||||
kind: str
|
||||
now: datetime
|
||||
@@ -48,10 +36,8 @@ class Envelope:
|
||||
names_only_within: float = 600.0
|
||||
last_at: datetime | None = None
|
||||
recall: Callable[[RecallContext], str | None] | None = None
|
||||
"""Setup-side lookup run on the user's text at turn start: pointers into
|
||||
the vault (a person's card, the agent's own notes, due dates) that the
|
||||
gateway cannot know the paths of. Its lines go under the vault block;
|
||||
a failure is logged and the envelope goes out without them."""
|
||||
"""Setup-side lookup on the user's text; its lines go under the change block."""
|
||||
texts: EnvelopeTexts = field(default_factory=EnvelopeTexts)
|
||||
|
||||
def build(
|
||||
self, *, now: datetime | None = None, text: str = "", kind: str = "master"
|
||||
@@ -70,6 +56,7 @@ class Envelope:
|
||||
names_only=names_only,
|
||||
max_lines=self.max_lines,
|
||||
per_file=self.per_file,
|
||||
texts=self.texts,
|
||||
)
|
||||
self.last_at = now
|
||||
block = self.recall_block(text=text, kind=kind, now=now)
|
||||
@@ -78,9 +65,8 @@ class Envelope:
|
||||
def recall_only(
|
||||
self, *, text: str, kind: str, now: datetime | None = 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))
|
||||
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:
|
||||
if self.recall is None or not text.strip():
|
||||
@@ -102,38 +88,52 @@ def render(
|
||||
names_only: bool,
|
||||
max_lines: int = 120,
|
||||
per_file: int = 30,
|
||||
texts: EnvelopeTexts | None = None,
|
||||
) -> str:
|
||||
texts = texts or EnvelopeTexts()
|
||||
zone = ZoneInfo(tz)
|
||||
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))
|
||||
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:
|
||||
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:
|
||||
_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])
|
||||
|
||||
|
||||
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:
|
||||
budget = max_lines - len(lines) - 1
|
||||
for change in changes:
|
||||
if not change.full or not change.added:
|
||||
continue
|
||||
if budget < 3:
|
||||
lines.append("… (потолок конверта)")
|
||||
lines.append(texts.truncated)
|
||||
return
|
||||
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)
|
||||
budget -= 1 + 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
|
||||
|
||||
|
||||
@@ -42,8 +42,8 @@ URGENCY: tuple[Priority, ...] = ("normal", "wake", "urgent")
|
||||
"""What the API and the tools accept for ``urgency``: every priority but ``user``."""
|
||||
|
||||
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:
|
||||
"""Default framing; a setup overrides it via ``ConversationTexts.inject_header``."""
|
||||
head = (
|
||||
f"[инжект: {ctx.origin} - это не Бобёр, отвечать не нужно, голос не обязателен]"
|
||||
)
|
||||
head = f"[inject: {ctx.origin} - not the user, no reply needed]"
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from beaver_gateway.core.conversations import Conversations
|
||||
from beaver_gateway.conversations.service import Conversations
|
||||
from beaver_gateway.storage.models import Conversation
|
||||
|
||||
__all__ = ["HandoutContext", "Rotation", "RotationPolicy"]
|
||||
|
||||
_log = logging.getLogger("beaver_gateway.core.rotation")
|
||||
_log = logging.getLogger("beaver_gateway.conversations.rotation")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -46,11 +46,11 @@ class RotationPolicy:
|
||||
if now.astimezone(zone) < boundary:
|
||||
boundary -= timedelta(days=1)
|
||||
if created < boundary and silence > self.night_silence:
|
||||
return "ночь"
|
||||
return "night"
|
||||
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:
|
||||
return "транскрипт"
|
||||
return "context"
|
||||
return None
|
||||
|
||||
def day_of(self, master: Conversation) -> date:
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
+6
-6
@@ -13,9 +13,9 @@ from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from claude_agent_sdk import create_sdk_mcp_server, tool
|
||||
|
||||
from beaver_gateway.core.injects import URGENCY
|
||||
from beaver_gateway.core.kinds import as_kind
|
||||
from beaver_gateway.core.redact import redact_data
|
||||
from beaver_gateway.conversations.injects import URGENCY
|
||||
from beaver_gateway.conversations.kinds import as_kind
|
||||
from beaver_gateway.security.redact import redact_data
|
||||
|
||||
URGENCY_HELP = (
|
||||
"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 beaver_gateway.core.conversations import Conversations
|
||||
from beaver_gateway.conversations.service import Conversations
|
||||
|
||||
__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"
|
||||
SAY_IN_USER_TURN = (
|
||||
@@ -230,7 +230,7 @@ def _tools(conversations: Conversations, key: str) -> list[SdkMcpTool[Any]]:
|
||||
conv,
|
||||
str(args["text"]),
|
||||
urgency=cast("Any", args.get("urgency") or "normal"),
|
||||
origin="агент",
|
||||
origin="agent",
|
||||
)
|
||||
return _text(f"queued #{item.id}")
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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
@@ -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."""
|
||||
+1
-1
@@ -33,7 +33,7 @@ from anthropic.types import (
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
from beaver_gateway.core.events import MessageStreamEvent, StopReason
|
||||
from beaver_gateway.events.stream import MessageStreamEvent, StopReason
|
||||
|
||||
|
||||
__all__ = ["StreamAccumulator", "accumulate"]
|
||||
@@ -27,8 +27,8 @@ import itsdangerous
|
||||
from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi.responses import FileResponse, JSONResponse, Response
|
||||
|
||||
from beaver_gateway.core import audit
|
||||
from beaver_gateway.frontends.base import Frontend
|
||||
from beaver_gateway.security import audit
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from beaver_gateway.frontends.base import GatewayRuntime
|
||||
|
||||
@@ -25,21 +25,21 @@ from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
|
||||
from beaver_gateway.agents.claude import ClaudeAgent
|
||||
from beaver_gateway.core import audit
|
||||
from beaver_gateway.core.transcript import fingerprint, text_of
|
||||
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.backends.capture import TurnCapture
|
||||
from beaver_gateway.backends.transcript import fingerprint, text_of
|
||||
from beaver_gateway.frontends.accumulate import StreamAccumulator
|
||||
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:
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
|
||||
from anthropic.types import Message, MessageParam
|
||||
|
||||
from beaver_gateway.core.conversations import Conversations
|
||||
from beaver_gateway.core.events import MessageStreamEvent
|
||||
from beaver_gateway.conversations.service import Conversations
|
||||
from beaver_gateway.events.stream import MessageStreamEvent
|
||||
from beaver_gateway.frontends.base import GatewayRuntime
|
||||
from beaver_gateway.storage.models import Conversation
|
||||
|
||||
|
||||
@@ -29,20 +29,20 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from sqlmodel import col, select
|
||||
|
||||
from beaver_gateway.core import audit
|
||||
from beaver_gateway.core.auth import VALID_SCOPES, hash_token
|
||||
from beaver_gateway.core.conversations import SEEDS, implied_title
|
||||
from beaver_gateway.core.injects import URGENCY
|
||||
from beaver_gateway.core.kinds import Kind, as_kind
|
||||
from beaver_gateway.frontends._auth import require_token
|
||||
from beaver_gateway.frontends._sse import (
|
||||
from beaver_gateway.conversations.injects import URGENCY
|
||||
from beaver_gateway.conversations.kinds import Kind, as_kind
|
||||
from beaver_gateway.conversations.service import SEEDS, implied_title
|
||||
from beaver_gateway.frontends.base import Frontend
|
||||
from beaver_gateway.frontends.bearer import require_token
|
||||
from beaver_gateway.frontends.sse import (
|
||||
KEEPALIVE,
|
||||
SSE_HEADERS,
|
||||
events_with_heartbeat,
|
||||
sse_pack,
|
||||
)
|
||||
from beaver_gateway.frontends._urls import frontend_url
|
||||
from beaver_gateway.frontends.base import Frontend
|
||||
from beaver_gateway.frontends.urls import frontend_url
|
||||
from beaver_gateway.security import audit
|
||||
from beaver_gateway.security.auth import VALID_SCOPES, hash_token
|
||||
from beaver_gateway.storage import (
|
||||
create_token,
|
||||
list_audit_records,
|
||||
@@ -61,9 +61,9 @@ if TYPE_CHECKING:
|
||||
from collections.abc import AsyncIterator, Iterable, Sequence
|
||||
from pathlib import Path
|
||||
|
||||
from beaver_gateway.core.conversations import Conversations
|
||||
from beaver_gateway.core.scheduler import Job, Scheduler
|
||||
from beaver_gateway.conversations.service import Conversations
|
||||
from beaver_gateway.frontends.base import GatewayRuntime
|
||||
from beaver_gateway.jobs.scheduler import Job, Scheduler
|
||||
|
||||
_log = logging.getLogger("beaver_gateway.frontends.api")
|
||||
|
||||
|
||||
@@ -21,11 +21,11 @@ if TYPE_CHECKING:
|
||||
|
||||
from starlette.types import ASGIApp
|
||||
|
||||
from beaver_gateway.app import AgentRegistry, McpRegistry
|
||||
from beaver_gateway.backends.base import Backend
|
||||
from beaver_gateway.core.auth import TokenStore
|
||||
from beaver_gateway.core.kinds import Kind
|
||||
from beaver_gateway.core.registry import AgentRegistry, McpRegistry
|
||||
from beaver_gateway.core.turn_record import TurnRecord
|
||||
from beaver_gateway.conversations.kinds import Kind
|
||||
from beaver_gateway.frontends.turn_record import TurnRecord
|
||||
from beaver_gateway.security.auth import TokenStore
|
||||
from beaver_gateway.storage import Database
|
||||
from beaver_gateway.storage.models import Conversation, ConversationBinding
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ if TYPE_CHECKING:
|
||||
|
||||
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")
|
||||
|
||||
@@ -45,23 +45,10 @@ from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
|
||||
from beaver_gateway.core import audit
|
||||
from beaver_gateway.core.conversation_store import (
|
||||
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.backends.capture import TurnCapture
|
||||
from beaver_gateway.frontends.accumulate import StreamAccumulator
|
||||
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.crossfront import CrossFrontendLogger
|
||||
from beaver_gateway.frontends.markdown.files import (
|
||||
@@ -69,12 +56,25 @@ from beaver_gateway.frontends.markdown.files import (
|
||||
reattach_frontmatter,
|
||||
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.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:
|
||||
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.storage.models import Conversation, ConversationBinding
|
||||
|
||||
|
||||
@@ -19,8 +19,6 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
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.crossfront import strip_trailing_user_scaffold
|
||||
from beaver_gateway.frontends.markdown.files import (
|
||||
@@ -28,12 +26,14 @@ from beaver_gateway.frontends.markdown.files import (
|
||||
reattach_frontmatter,
|
||||
write_atomic,
|
||||
)
|
||||
from beaver_gateway.frontends.markdown.history import load_messages, rewrite_messages
|
||||
from beaver_gateway.frontends.turn_record import slugify
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable, Sequence
|
||||
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.storage.models import Conversation, ConversationBinding
|
||||
|
||||
|
||||
@@ -45,10 +45,10 @@ from starlette.applications import Starlette
|
||||
from starlette.responses import HTMLResponse, JSONResponse, StreamingResponse
|
||||
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.urls import external_base
|
||||
from beaver_gateway.mcp.internal_app import ALL_NAMESPACE
|
||||
from beaver_gateway.security import audit
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from beaver_gateway.core.bus import Event, EventBus
|
||||
from beaver_gateway.core.conversations import Conversations
|
||||
from beaver_gateway.core.kinds import Kind
|
||||
from beaver_gateway.conversations.kinds import Kind
|
||||
from beaver_gateway.conversations.service import Conversations
|
||||
from beaver_gateway.events.bus import Event, EventBus
|
||||
from beaver_gateway.frontends.base import GatewayRuntime
|
||||
from beaver_gateway.storage.models import Conversation, ConversationBinding
|
||||
|
||||
@@ -773,7 +773,7 @@ class TelegramFrontend(Frontend):
|
||||
draft = self._drafts.pop(conv.external_id, None)
|
||||
if draft is not None:
|
||||
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 "")
|
||||
if user_text:
|
||||
await self._deliver(
|
||||
|
||||
@@ -32,7 +32,7 @@ from beaver_gateway.storage.models import Delivery
|
||||
if TYPE_CHECKING:
|
||||
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
|
||||
|
||||
__all__ = ["Outbox"]
|
||||
|
||||
@@ -33,7 +33,7 @@ from starlette.applications import Starlette
|
||||
from starlette.responses import JSONResponse
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -44,15 +44,15 @@ if TYPE_CHECKING:
|
||||
from pgqueuer.ports.driver import Driver
|
||||
from starlette.requests import Request
|
||||
|
||||
from beaver_gateway.core.conversations import Conversations, DistillResult
|
||||
from beaver_gateway.core.distill import LineCap
|
||||
from beaver_gateway.core.injects import Priority
|
||||
from beaver_gateway.core.rotation import Rotation
|
||||
from beaver_gateway.conversations.distill import LineCap
|
||||
from beaver_gateway.conversations.injects import Priority
|
||||
from beaver_gateway.conversations.rotation import Rotation
|
||||
from beaver_gateway.conversations.service import Conversations, DistillResult
|
||||
from beaver_gateway.storage.models import Conversation
|
||||
|
||||
__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"
|
||||
RETRY = timedelta(minutes=15)
|
||||
@@ -20,7 +20,7 @@ remembering to list it here.
|
||||
What this does not reach: tools that never touch a FastMCP server — the
|
||||
gateway's own ``gateway`` tools, and everything claude-code runs inside
|
||||
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
|
||||
@@ -31,7 +31,7 @@ import mcp.types as mt
|
||||
from fastmcp.server.middleware import Middleware
|
||||
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:
|
||||
from fastmcp.server.middleware import CallNext, MiddlewareContext
|
||||
|
||||
@@ -29,7 +29,7 @@ if TYPE_CHECKING:
|
||||
|
||||
__all__ = ["Change", "VaultWatch", "WatchRules"]
|
||||
|
||||
_log = logging.getLogger("beaver_gateway.core.watch")
|
||||
_log = logging.getLogger("beaver_gateway.vault.watch")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
Reference in New Issue
Block a user