Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions src/cmcp_runtime/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,9 @@ def build_server(ctx: RuntimeContext, *, trace_gate: TraceGate | None = None) ->
session=session,
bearer_token=ctx.config.bearer_token,
operator_token=ctx.config.operator_token,
audit_store=ctx.audit_store,
kill_switch_store=ctx.kill_switch_store,
session_state_store=ctx.session_state_store,
)


Expand Down
4 changes: 4 additions & 0 deletions src/cmcp_runtime/kill_switch.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,10 @@ def __init__(self, db_path: Path) -> None:
self._conn.commit()
logger.info("Kill switch block store opened: path=%s", db_path)

def close(self) -> None:
with self._lock:
self._conn.close()

def block(self, agent_id: str, *, reason: str) -> None:
"""Record a block. Blocking an identity that is already blocked keeps the first record."""
with self._lock:
Expand Down
8 changes: 8 additions & 0 deletions src/cmcp_runtime/mcp/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -345,6 +345,7 @@ def __init__(
# persist. Keep this separate from drain state: shutdown can still reap.
self._failed_terminal_call: str | None = None
self._shutting_down = False
self._shutdown_drained = False
# Set when the kill switch has stopped this gateway: the session it
# served is closed and no successor exists, so admission is refused
# rather than held open for a rotation that is never coming.
Expand Down Expand Up @@ -499,6 +500,11 @@ async def _drain_calls(self, drain_timeout: float) -> None:
) from None
self._drain_incomplete = False

@property
def shutdown_drained(self) -> bool:
"""Admission is permanently sealed and all store writers have drained."""
return self._shutdown_drained

async def shutdown(self, *, drain_timeout: float = SESSION_CLOSE_DRAIN_SECONDS) -> None:
"""Permanently reject admission, drain calls, then close owned resources.

Expand All @@ -513,6 +519,8 @@ async def shutdown(self, *, drain_timeout: float = SESSION_CLOSE_DRAIN_SECONDS)
async with self._lifecycle_condition:
self._session_rotation_in_progress = True
await self._drain_calls(drain_timeout)
# Cleanup may fail after draining; durable stores are still safe to close.
self._shutdown_drained = True
async with self._stdio_spawn_lock:
await self.aclose()

Expand Down
41 changes: 40 additions & 1 deletion src/cmcp_runtime/mcp/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,11 @@

if TYPE_CHECKING:
from cmcp_runtime.audit.chain import AuditChain
from cmcp_runtime.audit.store import SqliteAuditStore
from cmcp_runtime.kill_switch import KillSwitchBlockStore
from cmcp_runtime.session.manager import SessionManager
from cmcp_runtime.session.state import ClosedSessionRecord, SessionState
from cmcp_runtime.session.store import SqliteSessionStateStore

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -374,8 +377,18 @@ def __init__(
operator_token: str | None = None,
session: SessionState | None = None,
max_request_bytes: int = _DEFAULT_MAX_REQUEST_BYTES,
audit_store: SqliteAuditStore | None = None,
kill_switch_store: KillSwitchBlockStore | None = None,
session_state_store: SqliteSessionStateStore | None = None,
) -> None:
self._proxy = proxy
# Explicit transfer of process-lifetime ownership. Stores merely referenced
# by an embedded proxy/chain are not implicitly owned by this server.
self._durable_stores = [
store for store in (audit_store, kill_switch_store, session_state_store)
if store is not None
]
self._shutdown_lock = asyncio.Lock()
# Read once. With no gate the /trace routes are not registered and
# tools/call ignores _cmcp.trace, so the server is unchanged.
self._trace_gate: TraceGate | None = (
Expand Down Expand Up @@ -477,7 +490,33 @@ async def _lifespan(self, app: Starlette) -> AsyncIterator[None]:
try:
yield
finally:
await self._proxy.shutdown(drain_timeout=self._session_close_drain_s)
await self.shutdown()

async def shutdown(self) -> None:
"""Release runtime-owned stores after writers stop; failed closes can be retried."""
async with self._shutdown_lock:
shutdown_error: BaseException | None = None
try:
await self._proxy.shutdown(drain_timeout=self._session_close_drain_s)
except BaseException as exc:
shutdown_error = exc
raise
finally:
# A timeout or cancellation before drain completion leaves writers
# alive. A later retry must retain their audit/persistence handles.
if self._proxy.shutdown_drained is True:
close_error: Exception | None = None
for store in tuple(self._durable_stores):
try:
store.close()
except Exception as exc:
logger.exception("Failed to close a durable store during shutdown")
if close_error is None:
close_error = exc
else:
self._durable_stores.remove(store)
if shutdown_error is None and close_error is not None:
raise close_error

async def _parse_mcp_envelope(self, request: Request) -> dict[str, Any] | Response:
"""Read, size-check, and parse the request body.
Expand Down
137 changes: 137 additions & 0 deletions tests/unit/test_cli_wiring.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,3 +119,140 @@ def test_proxy_receives_attestation_timestamps(ctx):
proxy = server._proxy
assert proxy._attestation_generated_at is not None
assert proxy._attestation_validity_seconds == 86400


@pytest.fixture
def durable_ctx(ctx, tmp_path):
from cmcp_runtime.audit.keys import SigningKey
from cmcp_runtime.kill_switch import KillSwitchBlockStore
from cmcp_runtime.session.store import SqliteSessionStateStore

ctx.signing_key = SigningKey()
ctx.kill_switch_store = KillSwitchBlockStore(tmp_path / "blocks.db")
ctx.session_state_store = SqliteSessionStateStore(tmp_path / "sessions.db")
yield ctx
for store in (ctx.audit_store, ctx.kill_switch_store, ctx.session_state_store):
store.close()


def test_shutdown_closes_production_stores(durable_ctx, monkeypatch):
stores = (durable_ctx.audit_store, durable_ctx.kill_switch_store, durable_ctx.session_state_store)
closes = [MagicMock(wraps=store.close) for store in stores]
for store, close in zip(stores, closes, strict=True):
monkeypatch.setattr(store, "close", close)
server = build_server(durable_ctx)
with TestClient(server.app):
assert not any(close.called for close in closes)
for store, close in zip(stores, closes, strict=True):
close.assert_called_once_with()
with pytest.raises(sqlite3.ProgrammingError, match="closed"):
store._conn.execute("SELECT 1")


@pytest.mark.asyncio
async def test_shutdown_retries_only_failed_store_closes(durable_ctx, monkeypatch):
server = build_server(durable_ctx)
stores = (durable_ctx.audit_store, durable_ctx.kill_switch_store, durable_ctx.session_state_store)
closes = [MagicMock(wraps=store.close) for store in stores]
failure = RuntimeError("audit close failed")
closes[0].side_effect = failure
for store, close in zip(stores, closes, strict=True):
monkeypatch.setattr(store, "close", close)
with pytest.raises(RuntimeError) as caught:
await server.shutdown()
assert caught.value is failure
assert [close.call_count for close in closes] == [1, 1, 1]
closes[0].side_effect = None
await server.shutdown()
await server.shutdown()
assert [close.call_count for close in closes] == [2, 1, 1]


@pytest.mark.asyncio
async def test_shutdown_preserves_cleanup_error_and_attempts_all_stores(durable_ctx, monkeypatch):
from unittest.mock import AsyncMock

server = build_server(durable_ctx)
failure = RuntimeError("upstream cleanup failed")
monkeypatch.setattr(server._proxy, "aclose", AsyncMock(side_effect=failure))
close = MagicMock(side_effect=RuntimeError("audit close failed"))
original_close = durable_ctx.audit_store.close
monkeypatch.setattr(durable_ctx.audit_store, "close", close)
try:
with pytest.raises(RuntimeError) as caught:
await server.shutdown()
assert caught.value is failure
close.assert_called_once_with()
for store in (durable_ctx.kill_switch_store, durable_ctx.session_state_store):
with pytest.raises(sqlite3.ProgrammingError, match="closed"):
store._conn.execute("SELECT 1")
finally:
monkeypatch.setattr(durable_ctx.audit_store, "close", original_close)


@pytest.mark.asyncio
async def test_incomplete_shutdown_retains_stores_for_outcome_and_retry(durable_ctx, monkeypatch, tmp_path):
import asyncio
from types import SimpleNamespace

from cmcp_runtime.errors import SessionDrainIncomplete, UpstreamUnavailable
from cmcp_runtime.mcp import proxy as proxy_module
from cmcp_runtime.session.store import StoredSensitivity

server = build_server(durable_ctx)
server._session_close_drain_s = 0.01
monkeypatch.setattr(proxy_module, "SESSION_CANCELLATION_GRACE_SECONDS", 0.01)
entered, release = asyncio.Event(), asyncio.Event()
stores = (durable_ctx.audit_store, durable_ctx.kill_switch_store, durable_ctx.session_state_store)
closes = [MagicMock(wraps=store.close) for store in stores]
for store, close in zip(stores, closes, strict=True):
monkeypatch.setattr(store, "close", close)

async def call(*args, **kwargs):
entered.set()
try:
await release.wait()
except asyncio.CancelledError:
await release.wait()
server._audit_chain.append("fault", call_id="late-outcome", detail={"outcome": "finished"})
durable_ctx.kill_switch_store.block("late-agent", reason="test outcome")
durable_ctx.session_state_store.save(
"late-session", StoredSensitivity("PUBLIC", None, "late-outcome", 0)
)
return SimpleNamespace()

monkeypatch.setattr(server._proxy, "_call_tool_impl", call)
task = asyncio.create_task(server._proxy.call_tool("late-outcome", "test.tool", {}))
await entered.wait()
try:
async with server._lifespan(server.app):
pass
except SessionDrainIncomplete:
assert not any(close.called for close in closes)
with pytest.raises(UpstreamUnavailable):
await server._proxy._enter_call()
else:
pytest.fail("shutdown claimed success while a writer remained")
finally:
release.set()
await asyncio.wait_for(task, 1)

assert durable_ctx.kill_switch_store.is_blocked("late-agent")
assert durable_ctx.session_state_store.load("late-session").sensitivity_raised_by_call == "late-outcome"
with sqlite3.connect(tmp_path / "audit.db") as conn:
assert conn.execute("SELECT COUNT(*) FROM audit_entries WHERE payload LIKE '%late-outcome%'").fetchone()[0] >= 1
await server.shutdown()
await server.shutdown()
assert [close.call_count for close in closes] == [1, 1, 1]


@pytest.mark.asyncio
async def test_embedded_server_does_not_close_borrowed_stores(durable_ctx):
from cmcp_runtime.mcp.server import MCPServer

owner = build_server(durable_ctx)
embedded = MCPServer(owner._proxy, audit_chain=owner._audit_chain)
await embedded.shutdown()
for store in (durable_ctx.audit_store, durable_ctx.kill_switch_store, durable_ctx.session_state_store):
assert store._conn.execute("SELECT 1").fetchone() == (1,)
await owner.shutdown()
Loading