fix(backends,storage,core): setuid in exec wrapper, missing-column migration, runner config dir for mirror, tagged prompt sources

This commit is contained in:
hh
2026-08-28 02:10:19 +02:00
parent e1f242a87a
commit 50b7057fa4
6 changed files with 127 additions and 23 deletions
+5 -3
View File
@@ -1,8 +1,9 @@
"""Claude agent definition, backed by the Claude Agent SDK. """Claude agent definition, backed by the Claude Agent SDK.
The system prompt is either ``system_prompt`` verbatim or, when The system prompt is either ``system_prompt`` verbatim or, when
``prompt_sources`` is set, the concatenation of those files assembled at ``prompt_sources`` is set, the concatenation of those files (or
every session spawn (see ``core/prompt.py``). ``skill_sets`` are ``(tag, file)`` pairs) assembled at every session spawn (see
``core/prompt.py``). ``skill_sets`` are
directories of ``<skill>/SKILL.md`` folders; each becomes a local SDK directories of ``<skill>/SKILL.md`` folders; each becomes a local SDK
plugin. Nothing from disk is loaded otherwise: the adapter runs with plugin. Nothing from disk is loaded otherwise: the adapter runs with
``setting_sources=[]``. ``setting_sources=[]``.
@@ -16,6 +17,7 @@ from pathlib import Path # noqa: TC003 - pydantic runtime
from pydantic import BaseModel, ConfigDict, Field from pydantic import BaseModel, ConfigDict, Field
from beaver_gateway.agents.base import BaseAgent from beaver_gateway.agents.base import BaseAgent
from beaver_gateway.core.prompt import PromptSource # noqa: TC001 - pydantic runtime
__all__ = ["ClaudeAgent", "ClaudeOptions"] __all__ = ["ClaudeAgent", "ClaudeOptions"]
@@ -49,6 +51,6 @@ class ClaudeOptions(BaseModel):
class ClaudeAgent(BaseAgent): class ClaudeAgent(BaseAgent):
cwd: Path cwd: Path
system_prompt: str = "" system_prompt: str = ""
prompt_sources: tuple[Path, ...] = () prompt_sources: tuple[PromptSource, ...] = ()
skill_sets: tuple[Path, ...] = () skill_sets: tuple[Path, ...] = ()
options: ClaudeOptions = Field(default_factory=ClaudeOptions) options: ClaudeOptions = Field(default_factory=ClaudeOptions)
+17 -9
View File
@@ -13,8 +13,9 @@ Events on the wire are the Anthropic ``MessageStreamEvent`` family: one
rebased across the API calls claude makes inside the turn. rebased across the API calls claude makes inside the turn.
Process isolation: claude is spawned through a small exec wrapper that Process isolation: claude is spawned through a small exec wrapper that
drops every inherited environment variable outside a whitelist, and, drops every inherited environment variable outside a whitelist and, when
when ``RunnerConfig.user`` is set, under that uid. ``RunnerConfig.user`` is set, switches to that uid before exec (done in the
wrapper rather than via ``subprocess(user=...)``, which uvloop rejects).
""" """
from __future__ import annotations from __future__ import annotations
@@ -204,7 +205,7 @@ class ClaudeSdkBackend:
self._mcp_disallowed = _mcp_disallowed(agent, mcp_tool_names or {}) self._mcp_disallowed = _mcp_disallowed(agent, mcp_tool_names or {})
self._sessions: dict[str, _Live] = {} self._sessions: dict[str, _Live] = {}
self._reaper: asyncio.Task[None] | None = None self._reaper: asyncio.Task[None] | None = None
self._uid = _resolve_uid(self._runner.user) self._uid, self._gid = _resolve_ids(self._runner.user)
self._wrapper: Path | None = None self._wrapper: Path | None = None
@property @property
@@ -411,6 +412,7 @@ class ClaudeSdkBackend:
env = dict(opt.env) env = dict(opt.env)
if self._runner.home is not None: if self._runner.home is not None:
env["HOME"] = str(self._runner.home) env["HOME"] = str(self._runner.home)
env.setdefault("CLAUDE_CONFIG_DIR", str(self._runner.home / ".claude"))
plugins = self._plugins() plugins = self._plugins()
system_prompt = ( system_prompt = (
prompt_assembly.assemble(agent.prompt_sources) prompt_assembly.assemble(agent.prompt_sources)
@@ -430,7 +432,6 @@ class ClaudeSdkBackend:
cwd=str(agent.cwd), cwd=str(agent.cwd),
add_dirs=list(opt.add_dirs), add_dirs=list(opt.add_dirs),
env=env, env=env,
user=self._runner.user,
cli_path=str(self._exec_wrapper(extra_keep=tuple(env))), cli_path=str(self._exec_wrapper(extra_keep=tuple(env))),
include_partial_messages=opt.include_partial_messages, include_partial_messages=opt.include_partial_messages,
session_store=self._store, session_store=self._store,
@@ -472,6 +473,8 @@ class ClaudeSdkBackend:
target=json.dumps(target), target=json.dumps(target),
keep=json.dumps(keep), keep=json.dumps(keep),
prefixes=json.dumps(list(ENV_KEEP_PREFIXES)), prefixes=json.dumps(list(ENV_KEEP_PREFIXES)),
uid=json.dumps(self._uid),
gid=json.dumps(self._gid),
) )
digest = hashlib.sha256(script.encode("utf-8")).hexdigest()[:12] digest = hashlib.sha256(script.encode("utf-8")).hexdigest()[:12]
path = self._work_dir / f"claude-exec-{digest}.py" path = self._work_dir / f"claude-exec-{digest}.py"
@@ -525,9 +528,15 @@ import sys
TARGET = {target} TARGET = {target}
KEEP = set({keep}) KEEP = set({keep})
PREFIXES = tuple({prefixes}) PREFIXES = tuple({prefixes})
UID = {uid}
GID = {gid}
env = {{ env = {{
k: v for k, v in os.environ.items() if k in KEEP or k.startswith(PREFIXES) k: v for k, v in os.environ.items() if k in KEEP or k.startswith(PREFIXES)
}} }}
if UID is not None and os.getuid() != UID:
os.setgroups([])
os.setgid(GID)
os.setuid(UID)
os.execve(TARGET, [TARGET, *sys.argv[1:]], env) os.execve(TARGET, [TARGET, *sys.argv[1:]], env)
""" """
@@ -550,12 +559,11 @@ def _claude_binary() -> str:
return found return found
def _resolve_uid(user: str | None) -> int | None: def _resolve_ids(user: str | None) -> tuple[int | None, int | None]:
if user is None: if user is None:
return None return None, None
if user.isdigit(): record = pwd.getpwuid(int(user)) if user.isdigit() else pwd.getpwnam(user)
return int(user) return record.pw_uid, record.pw_gid
return pwd.getpwnam(user).pw_uid
def _chown_tree(root: Path, uid: int) -> None: def _chown_tree(root: Path, uid: int) -> None:
+15 -7
View File
@@ -3,7 +3,9 @@
The gateway holds no prompt text: an agent names its granules (paths from The gateway holds no prompt text: an agent names its granules (paths from
``config.py``) and :func:`assemble` concatenates them in that order, so the ``config.py``) and :func:`assemble` concatenates them in that order, so the
result is byte-for-byte identical for every session of the same agent as result is byte-for-byte identical for every session of the same agent as
long as the files are. Each granule's hash is logged at assembly so a long as the files are. A source is a path, or a ``(tag, path)`` pair whose
content is wrapped in ``<tag>...</tag>`` - the markup lives here, the vault
keeps plain markdown. Each granule's hash is logged at assembly so a
drifted prompt can be traced to the file that changed. drifted prompt can be traced to the file that changed.
""" """
@@ -17,18 +19,24 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Iterable from collections.abc import Iterable
__all__ = ["assemble"] __all__ = ["PromptSource", "assemble"]
_log = logging.getLogger("beaver_gateway.core.prompt") _log = logging.getLogger("beaver_gateway.core.prompt")
PromptSource = str | Path | tuple[str, str | Path]
def assemble(sources: Iterable[str | Path]) -> str:
def assemble(sources: Iterable[PromptSource]) -> str:
parts: list[str] = [] parts: list[str] = []
for source in sources: for source in sources:
path = Path(source) tag, raw = source if isinstance(source, tuple) else (None, source)
path = Path(raw)
text = path.read_text(encoding="utf-8").strip() text = path.read_text(encoding="utf-8").strip()
digest = hashlib.sha256(text.encode("utf-8")).hexdigest()[:12] digest = hashlib.sha256(text.encode("utf-8")).hexdigest()[:12]
_log.info("prompt granule %s sha=%s bytes=%d", path, digest, len(text)) _log.info(
if text: "prompt granule %s tag=%s sha=%s bytes=%d", path, tag, digest, len(text)
parts.append(text) )
if not text:
continue
parts.append(f"<{tag}>\n{text}\n</{tag}>" if tag else text)
return "\n\n".join(parts) + "\n" return "\n\n".join(parts) + "\n"
+21 -1
View File
@@ -18,6 +18,7 @@ import json
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from sqlalchemy import inspect, text
from sqlalchemy.ext.asyncio import create_async_engine from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import SQLModel, select from sqlmodel import SQLModel, select
from sqlmodel.ext.asyncio.session import AsyncSession from sqlmodel.ext.asyncio.session import AsyncSession
@@ -27,6 +28,7 @@ from beaver_gateway.storage.models import AuditLog, Token, Usage
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Sequence from collections.abc import Sequence
from sqlalchemy import Connection
from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio import AsyncEngine
@@ -72,9 +74,14 @@ class Database:
self._engine: AsyncEngine = create_async_engine(_to_async_url(url), echo=False) self._engine: AsyncEngine = create_async_engine(_to_async_url(url), echo=False)
async def create_all(self) -> None: async def create_all(self) -> None:
"""Idempotent ``CREATE TABLE IF NOT EXISTS`` for every model.""" """Create missing tables, then add columns models gained since.
Nullable columns only - that is the whole migration story until
Alembic is worth it.
"""
async with self._engine.begin() as conn: async with self._engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all) await conn.run_sync(SQLModel.metadata.create_all)
await conn.run_sync(_add_missing_columns)
def session(self) -> AsyncSession: def session(self) -> AsyncSession:
"""Open a fresh :class:`AsyncSession` (use as ``async with``). """Open a fresh :class:`AsyncSession` (use as ``async with``).
@@ -90,6 +97,19 @@ class Database:
await self._engine.dispose() await self._engine.dispose()
def _add_missing_columns(conn: Connection) -> None:
inspector = inspect(conn)
for table in SQLModel.metadata.sorted_tables:
existing = {c["name"] for c in inspector.get_columns(table.name)}
for column in table.columns:
if column.name in existing:
continue
kind = column.type.compile(conn.dialect)
conn.execute(
text(f"ALTER TABLE {table.name} ADD COLUMN {column.name} {kind}")
)
# ---- Token CRUD --------------------------------------------------------- # ---- Token CRUD ---------------------------------------------------------
+36 -3
View File
@@ -25,7 +25,12 @@ from claude_agent_sdk import (
from beaver_gateway.agents.base import ExposedMcp from beaver_gateway.agents.base import ExposedMcp
from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions from beaver_gateway.agents.claude import ClaudeAgent, ClaudeOptions
from beaver_gateway.backends.claude_sdk import ClaudeSdkBackend, UsageEvent, fingerprint from beaver_gateway.backends.claude_sdk import (
ClaudeSdkBackend,
RunnerConfig,
UsageEvent,
fingerprint,
)
from beaver_gateway.core.transcript import messages_from_entries from beaver_gateway.core.transcript import messages_from_entries
from beaver_gateway.core.turn_capture import TurnCapture from beaver_gateway.core.turn_capture import TurnCapture
@@ -359,7 +364,9 @@ async def test_prompt_sources_are_assembled(cwd: Path) -> None:
(cwd / "a.md").write_text("alpha\n") (cwd / "a.md").write_text("alpha\n")
(cwd / "b.md").write_text("\nbeta\n\n") (cwd / "b.md").write_text("\nbeta\n\n")
backend = _backend( backend = _backend(
cwd, InMemorySessionStore(), prompt_sources=(cwd / "a.md", cwd / "b.md") cwd,
InMemorySessionStore(),
prompt_sources=(("role", cwd / "a.md"), cwd / "b.md"),
) )
await _drain( await _drain(
backend.complete( backend.complete(
@@ -368,7 +375,33 @@ async def test_prompt_sources_are_assembled(cwd: Path) -> None:
conversation_id="c", conversation_id="c",
) )
) )
assert FakeClient.instances[0].options.system_prompt == "alpha\n\nbeta\n" assert (
FakeClient.instances[0].options.system_prompt
== "<role>\nalpha\n</role>\n\nbeta\n"
)
async def test_runner_user_lands_in_wrapper(cwd: Path) -> None:
import os
import pwd
me = pwd.getpwuid(os.getuid())
backend = _backend(cwd, InMemorySessionStore())
backend._runner = RunnerConfig(user=me.pw_name, home=cwd)
backend._uid, backend._gid = me.pw_uid, me.pw_gid
await _drain(
backend.complete(
agent=backend.agent,
messages=[{"role": "user", "content": "x"}],
conversation_id="c",
)
)
opts = FakeClient.instances[0].options
assert opts.env["HOME"] == str(cwd)
assert opts.env["CLAUDE_CONFIG_DIR"] == str(cwd / ".claude")
wrapper = Path(opts.cli_path).read_text()
assert f"UID = {me.pw_uid}" in wrapper
assert "os.setuid(UID)" in wrapper
async def test_close_disconnects(cwd: Path) -> None: async def test_close_disconnects(cwd: Path) -> None:
+33
View File
@@ -0,0 +1,33 @@
import sqlite3
import tempfile
from pathlib import Path
from sqlmodel import select
from beaver_gateway.storage import Database
from beaver_gateway.storage.models import Conversation
async def test_create_all_adds_missing_columns() -> None:
path = Path(tempfile.mkdtemp(prefix="beaver-migrate-")) / "old.db"
raw = sqlite3.connect(path)
raw.execute(
"CREATE TABLE conversations (id INTEGER PRIMARY KEY, frontend VARCHAR NOT NULL, "
"external_id VARCHAR NOT NULL, agent_name VARCHAR NOT NULL, "
"created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL)"
)
raw.execute(
"INSERT INTO conversations VALUES (1, 'markdown', 'x', 'a', '2026-01-01', '2026-01-01')"
)
raw.commit()
raw.close()
db = Database(f"sqlite:///{path}")
await db.create_all()
async with db.session() as session:
conv = (await session.exec(select(Conversation))).one()
assert conv.session_id is None
conv.session_id = "sid"
session.add(conv)
await session.commit()
await db.dispose()