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