diff --git a/docs/v0.4-baseline-findings.md b/docs/v0.4-baseline-findings.md new file mode 100644 index 0000000..7deffb3 --- /dev/null +++ b/docs/v0.4-baseline-findings.md @@ -0,0 +1,207 @@ +# v0.4 baseline findings: approval and session boundaries + +Date observed: 2026-08-16 +Revision: `55ff730` (`feat/agent-approval-plane`, identical to `origin/main` at observation time) + +## Scope and method + +These are behavioral baseline experiments, not a proposed implementation. They used only local ASGI/MCP clients, an in-memory SQLite database, and fake pool/session objects. No real SSH server or external service was contacted. Production and test sources were not modified. + +Interfaces exercised: + +- `fastapi.testclient.TestClient` against the unified service returned by `create_service_app()`; +- `fastmcp.Client` against a real in-memory `FastMCP` instance with registered Shuttle tools; +- real `ConfirmTokenStore` and tool registration/execution flow; +- stateful fake `SessionManager` and fake node repository, so session selection could be observed without SSH. + +For reproducibility, the focused pre-existing tests were also run: + +```text +.venv/bin/pytest -q \ + tests/test_mcp/test_fastmcp_client.py::test_run_confirm_flow_via_client \ + tests/test_mcp/test_service_app.py::test_create_service_app_rejects_bad_bearer \ + tests/test_mcp/test_tools.py::test_implicit_session_reused_across_calls + +... [100%] +``` + +## Finding 1: the service API token protects `/api`, not `/mcp` + +A unified app was created with `api_token="baseline-secret"`. The pool and session manager were replaced with local fakes, and startup used in-memory SQLite. + +### Exact observations + +| Probe | Authorization | HTTP status | Observed response | +|---|---|---:|---| +| `GET /api/stats` | none | 401 | `{"detail":"Invalid or missing token"}` | +| `GET /api/stats` | `Bearer wrong` | 401 | `{"detail":"Invalid or missing token"}` | +| `GET /api/stats` | `Bearer baseline-secret` | 200 | `{"node_count":0,"active_sessions":0,"total_commands":0}` | +| valid MCP `initialize` to `POST /mcp/` | none | 200 | MCP initialization result; server name `shuttle`; an `mcp-session-id` was issued | +| valid MCP `initialize` to `POST /mcp/` | `Bearer wrong` | 200 | same successful MCP initialization shape; a different `mcp-session-id` was issued | + +The successful unauthenticated MCP response began: + +```text +event: message +data: {"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"2025-03-26",... +``` + +**Current boundary:** possession of the web/API bearer token is required for `/api/*`, but is neither required nor validated for `/mcp/*`. An invalid bearer header does not cause MCP rejection. + +**Security implication:** anyone who can reach the MCP listener can initialize an MCP session and invoke its tools. The API token currently is a web-control-plane credential, not a service-wide or MCP authorization credential. + +## Finding 2: a CONFIRM challenge can be self-approved by the same MCP caller + +A real `FastMCP` instance was configured with the real `ConfirmTokenStore`, a guard that classified `sudo id` as CONFIRM, and a fake session executor. A single MCP client called `ssh_run` twice. + +### Exact observations + +First call (`ssh_run(command="sudo id", node="node-a")`) returned a bearer-like secret directly to the caller (actual random value redacted here; observed length was 43 characters): + +```text +⚠️ Confirmation required +Command: sudo id +Rule: operator approval required + +To proceed: ssh_run(command="sudo id", node="node-a", confirm_token="") +``` + +The same client copied that token into its next tool call. Observed result: + +```text +EXECUTED sid=session-1 command=sudo id +``` + +The fake executor's event log confirmed actual execution: + +```text +["session-1", "sudo id"] +``` + +Replaying the same token returned: + +```text +Error: invalid or expired confirmation token. +``` + +**Current boundary:** token validity binds to `(command, node)` and is one-time/TTL constrained, but not to an approving principal, transport session, or channel. The challenge and approval capability are delivered to the same caller through the same tool response. Therefore CONFIRM is a two-call acknowledgement, not independent human/operator approval. + +## Finding 3: implicit execution sessions are keyed only by node in one server process + +A stateful fake session manager recorded all creates and executes. Calls were sent through the real FastMCP tool protocol. + +### Same client, two nodes + +Call sequence and outputs: + +```text +node-a / pwd -> EXECUTED sid=session-1 command=pwd +node-b / pwd -> EXECUTED sid=session-2 command=pwd +node-a / whoami -> EXECUTED sid=session-1 command=whoami +``` + +Recorded creates and active sessions: + +```text +creates: ["node-a", "node-b"] +active_sessions: [["session-1", "node-a"], ["session-2", "node-b"]] +``` + +Thus node A reused its first session while node B received a separate session. + +### Two distinct MCP clients, same node + +Client A connected, invoked `ssh_run`, disconnected; then client B made a separate MCP connection and invoked `ssh_run` for the same node. Exact observation: + +```json +{ + "active_sessions": [["s-1", "node-a"]], + "client_a_output": "s-1", + "client_b_output": "s-1", + "events": [ + ["create", "node-a", "s-1"], + ["execute", "from-client-a", "s-1"], + ["execute", "from-client-b", "s-1"] + ] +} +``` + +**Current boundary:** automatic lookup chooses the first active in-memory session whose `node_id` equals the resolved node. It has no caller/client/transport identity dimension. Consequently, independent MCP clients targeting the same node share working-directory state and session bypass patterns. There can also be at most one automatically selected session per node under the normal tool path; if multiple active sessions exist, list order decides which is used. + +## Baseline conclusion + +The current effective model is: + +1. **Authentication:** web `/api/*` bearer authentication only; MCP is reachable without that credential. +2. **Confirmation:** one caller obtains and redeems its own one-time command/node token. +3. **Session isolation:** process-local automatic sessions are selected by node, and therefore shared across callers targeting that node. + +These are objective descriptions of v0.3.1 behavior, not assertions that the behavior is intended for v0.4. + +## Proposed v0.4 acceptance tests + +The expected policy should be made explicit before implementation. For an approval plane whose goal is authenticated agents, independent approval, and per-agent isolation, the following tests are recommended. + +### A. Authentication boundary + +1. **MCP rejects missing credentials** + - Build the unified ASGI app with a configured agent/API credential. + - Send a valid MCP `initialize` request to `/mcp/` without credentials. + - Expect `401` (or the selected protocol-auth failure), no MCP session ID, and no initialized session. + +2. **MCP rejects invalid credentials** + - Repeat with an invalid bearer credential. + - Expect rejection, not successful initialization. + +3. **MCP accepts a valid agent credential and propagates identity** + - Initialize and call a harmless fake-backed tool with a valid credential. + - Expect success and verify the authenticated principal is available to authorization, audit logging, confirmation, and session selection. + +4. **Web and MCP credential policy is explicit** + - If credentials are intentionally separate, prove a web-only token cannot invoke MCP and an agent-only token cannot mutate web administration routes. + - If intentionally shared, prove the same configured token is enforced consistently on both route families. + +### B. Independent confirmation + +5. **Requesting agent cannot self-approve from the challenge response** + - Agent A requests a CONFIRM-level command. + - The response may expose a challenge/request ID, but must not expose a credential sufficient for Agent A to execute it. + - Repeating the command or echoing the request ID without approval must not execute. + +6. **Only an authorized approver can approve** + - Agent A creates a pending request. + - Agent A's credential cannot approve it. + - Approver B, with the required role/capability, approves it through the approval interface. + - Agent A can then execute exactly the approved command on exactly the approved node. + +7. **Approval is bound and one-use** + - Approval for `(requester A, command X, node N)` must fail for command Y, node M, or requester C. + - First valid execution succeeds; replay fails; expiry fails. + +8. **Denial and audit are observable** + - Approver denial prevents execution. + - Audit evidence records requester, approver, decision, command digest/value, node, timestamps, and final execution outcome without relying on caller-supplied identity. + +### C. Session ownership and isolation + +9. **Same agent + same node reuses state** + - Agent A executes two commands on node N using a stateful fake. + - Expect the same session ID and preserved working directory/bypass state. + +10. **Different agents + same node do not share state** + - Agent A changes directory or obtains a session-scoped bypass on node N. + - Agent B then targets node N. + - Expect a different session ID, default working directory, and no inherited bypass. + +11. **Same agent + different nodes gets distinct sessions** + - Agent A targets nodes N and M. + - Expect separate session IDs and state. + +12. **Transport reconnect policy is explicit** + - Reconnect the same authenticated agent and assert either stable reuse by durable agent identity or deliberate reset, according to policy. + - A transport-generated MCP session ID alone should not silently redefine ownership unless that is the documented contract. + +13. **Concurrent selection is deterministic** + - Run simultaneous first calls for the same `(agent, node)` and assert one owned session is selected/created (or an explicitly supported multiplicity model), with no list-order ambiguity. + +All tests can remain hermetic by using TestClient/ASGI transport, FastMCP's in-memory client, in-memory SQLite, and stateful fake executors; none requires a real SSH endpoint. diff --git a/skills/shuttle/SKILL.md b/skills/shuttle/SKILL.md new file mode 100644 index 0000000..3773a52 --- /dev/null +++ b/skills/shuttle/SKILL.md @@ -0,0 +1,44 @@ +--- +name: shuttle +description: Use Shuttle for authenticated, policy-controlled remote operations. +--- + +# Using Shuttle + +Use Shuttle when an AI agent needs to operate existing SSH nodes through a central policy, approval, and audit gateway. + +## Core workflow + +1. Inspect available nodes with `ssh_list_nodes` and select the exact requested node. Never substitute another host silently. +2. Begin with read-only inspection whenever possible. +3. Run commands with `ssh_run`. Shuttle derives actor, client, and conversation identity from the authenticated MCP request context. +4. Handle the result according to its policy decision: + - `ALLOW` / `WARN`: inspect and report the result. + - `BLOCK`: stop; do not bypass it. + - `PENDING_APPROVAL`: report the exact node and command, then wait for a human decision in the Shuttle Web Approvals page. +5. After approval, retry the identical command with its `approval_id`. An approval is bound to the authenticated requester, MCP session, node, and exact command; it is single-use and expires. +6. Report the command result together with its approval or audit ID when present. + +## Sessions + +- Shuttle automatically preserves remote working-directory state within the authenticated MCP session. +- Do not try to provide or override actor, client, or conversation identifiers. +- A reconnect may create a new conversation boundary; do not assume prior approvals or session state carry over. + +## File operations + +- Use `ssh_upload` and `ssh_download` only after verifying the local path, remote path, node, and overwrite impact. +- Do not place secrets in command text or ordinary files when a credential mechanism is available. +- Treat node creation and file transfer as privileged operations even when no shell command is involved. + +## Authentication and host trust + +- HTTP MCP uses an Agent token distinct from the Web/operator token. +- Never expose the Web/operator token to an agent. +- Shuttle requires SSH host-key verification. Register the host key in the trusted `known_hosts` file before connecting; never disable checking with `known_hosts=None`. + +## Failure handling + +- Use bounded timeouts for commands that may hang. +- On connection or host-key errors, stop and report the exact failure instead of weakening verification. +- If an approval is denied, expired, consumed, or mismatched, request a new approval rather than altering identifiers or replaying it. diff --git a/src/shuttle/cli.py b/src/shuttle/cli.py index da68dc1..6a48dab 100644 --- a/src/shuttle/cli.py +++ b/src/shuttle/cli.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import json import secrets from pathlib import Path @@ -81,6 +82,14 @@ def serve( token_path.write_text(api_token) token_path.chmod(0o600) + agent_token_path = config.shuttle_dir / "agent_token" + if agent_token_path.exists(): + mcp_token = agent_token_path.read_text().strip() + else: + mcp_token = secrets.token_urlsafe(32) + agent_token_path.write_text(mcp_token) + agent_token_path.chmod(0o600) + from rich.console import Console from rich.panel import Panel from rich.text import Text @@ -93,6 +102,8 @@ def serve( info.append(f"http://{host}:{port}\n", style="cyan") info.append(" API token ", style="dim") info.append(api_token, style="green bold") + info.append("\n Agent token ", style="dim") + info.append(mcp_token, style="cyan bold") console.print( Panel( info, @@ -111,6 +122,7 @@ async def _run(): host=host, port=port, api_token=api_token, + mcp_token=mcp_token, db_url=db_url, ) @@ -771,6 +783,53 @@ async def _import() -> None: # ── Config commands ─────────────────────────────────────────────────────────── +@app.command("approvals") +def approvals_list( + status: str = typer.Option("pending", help="Approval status to list"), + json_output: bool = typer.Option(False, "--json", help="Emit stable JSON"), +) -> None: + """List approval requests for agents and operators.""" + from shuttle.core.config import ShuttleConfig + from shuttle.db.engine import create_db_engine, create_session_factory, init_db + from shuttle.db.repository import ApprovalRepo + + async def _list() -> list[dict]: + config = ShuttleConfig() + engine = create_db_engine(config.db_url) + await init_db(engine) + factory = create_session_factory(engine) + try: + async with factory() as session: + items = await ApprovalRepo(session).list(status) + return [ + { + "id": item.id, + "status": item.status, + "node": item.node_name, + "command": item.command, + "actor_id": item.actor_id, + "client_id": item.client_id, + "conversation_id": item.conversation_id, + "expires_at": item.expires_at.isoformat(), + } + for item in items + ] + finally: + await engine.dispose() + + items = asyncio.run(_list()) + if json_output: + typer.echo(json.dumps({"items": items}, ensure_ascii=False)) + return + if not items: + typer.echo("No approval requests.") + return + for item in items: + typer.echo( + f"{item['id']} {item['status']} {item['node']} {item['actor_id']} {item['command']}" + ) + + @config_app.command("show") def config_show() -> None: """Show current Shuttle configuration.""" diff --git a/src/shuttle/core/proxy.py b/src/shuttle/core/proxy.py index 2e6d58d..c841e9f 100644 --- a/src/shuttle/core/proxy.py +++ b/src/shuttle/core/proxy.py @@ -7,6 +7,7 @@ from __future__ import annotations from dataclasses import dataclass, field +from pathlib import Path import asyncssh @@ -47,7 +48,9 @@ class NodeConnectInfo: port: int = 22 password: str | None = None private_key: str | None = None - known_hosts: str | None = None + known_hosts: str | None = field( + default_factory=lambda: str(Path.home() / ".ssh" / "known_hosts") + ) jump_host: NodeConnectInfo | None = None connect_timeout: float = 30.0 extra_options: dict = field(default_factory=dict) @@ -99,11 +102,11 @@ def _build_connect_kwargs(info: NodeConnectInfo) -> dict: if info.private_key is not None: kwargs["client_keys"] = [asyncssh.import_private_key(info.private_key)] - if info.known_hosts is not None: - kwargs["known_hosts"] = info.known_hosts - else: - # Disable host-key verification when no known_hosts is provided. - kwargs["known_hosts"] = None + if info.known_hosts is None: + raise ValueError( + "known_hosts is required; configure a trusted file instead of disabling host-key verification" + ) + kwargs["known_hosts"] = info.known_hosts # Forward any caller-supplied overrides last so they can override defaults. kwargs.update(info.extra_options) diff --git a/src/shuttle/core/session.py b/src/shuttle/core/session.py index b43c080..77a6821 100644 --- a/src/shuttle/core/session.py +++ b/src/shuttle/core/session.py @@ -63,6 +63,9 @@ class SSHSession: session_id: str node_id: str + actor_id: str = "anonymous" + client_id: str = "mcp" + conversation_id: str = "default" working_directory: str = "~" bypass_patterns: set[str] = field(default_factory=set) status: SessionStatus = SessionStatus.ACTIVE @@ -100,7 +103,14 @@ def __init__( # CRUD # ------------------------------------------------------------------ - async def create(self, node_id: str) -> SSHSession: + async def create( + self, + node_id: str, + *, + actor_id: str = "anonymous", + client_id: str = "mcp", + conversation_id: str = "default", + ) -> SSHSession: """Create a new session for *node_id*. Runs ``pwd`` on the remote host to obtain the initial working @@ -123,6 +133,9 @@ async def create(self, node_id: str) -> SSHSession: session = SSHSession( session_id=str(uuid.uuid4()), node_id=node_id, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, working_directory=working_directory, ) self._sessions[session.session_id] = session @@ -146,6 +159,27 @@ def list_active(self) -> list[SSHSession]: """Return all currently active sessions.""" return list(self._sessions.values()) + def find_active( + self, + node_id: str, + *, + actor_id: str, + client_id: str, + conversation_id: str, + ) -> SSHSession | None: + """Return a session only when the complete caller boundary matches.""" + return next( + ( + session + for session in self._sessions.values() + if session.node_id == node_id + and session.actor_id == actor_id + and session.client_id == client_id + and session.conversation_id == conversation_id + ), + None, + ) + # ------------------------------------------------------------------ # Command execution # ------------------------------------------------------------------ diff --git a/src/shuttle/db/engine.py b/src/shuttle/db/engine.py index 4287181..13711c1 100644 --- a/src/shuttle/db/engine.py +++ b/src/shuttle/db/engine.py @@ -69,6 +69,8 @@ async def init_db(engine: AsyncEngine) -> None: "CREATE INDEX IF NOT EXISTS ix_command_logs_session ON command_logs (session_id)", "CREATE INDEX IF NOT EXISTS ix_security_rules_node ON security_rules (node_id)", "CREATE INDEX IF NOT EXISTS ix_sessions_node_status ON sessions (node_id, status)", + "CREATE INDEX IF NOT EXISTS ix_approval_status_created ON approval_requests (status, created_at)", + "CREATE INDEX IF NOT EXISTS ix_approval_actor ON approval_requests (actor_id, client_id, conversation_id)", ]: try: await conn.execute(text(idx_sql)) @@ -97,6 +99,31 @@ async def init_db(engine: AsyncEngine) -> None: await conn.execute( text("ALTER TABLE nodes ADD COLUMN last_seen_at DATETIME") ) + + migrations = { + "sessions": { + "actor_id": "VARCHAR(255) NOT NULL DEFAULT 'anonymous'", + "client_id": "VARCHAR(255) NOT NULL DEFAULT 'mcp'", + "conversation_id": "VARCHAR(255) NOT NULL DEFAULT 'default'", + }, + "command_logs": { + "action": "VARCHAR(50) NOT NULL DEFAULT 'command'", + "actor_id": "VARCHAR(255) NOT NULL DEFAULT 'anonymous'", + "client_id": "VARCHAR(255) NOT NULL DEFAULT 'mcp'", + "conversation_id": "VARCHAR(255) NOT NULL DEFAULT 'default'", + "approval_id": "VARCHAR(36)", + }, + } + for table, additions in migrations.items(): + result = await conn.execute(text(f"PRAGMA table_info({table})")) + existing = {row[1] for row in result} + for column, definition in additions.items(): + if column not in existing: + await conn.execute( + text( + f"ALTER TABLE {table} ADD COLUMN {column} {definition}" + ) + ) else: result = await conn.execute( text( @@ -127,3 +154,25 @@ async def init_db(engine: AsyncEngine) -> None: "ALTER TABLE nodes ADD COLUMN last_seen_at TIMESTAMP WITH TIME ZONE" ) ) + + pg_migrations = { + "sessions": { + "actor_id": "VARCHAR(255) NOT NULL DEFAULT 'anonymous'", + "client_id": "VARCHAR(255) NOT NULL DEFAULT 'mcp'", + "conversation_id": "VARCHAR(255) NOT NULL DEFAULT 'default'", + }, + "command_logs": { + "action": "VARCHAR(50) NOT NULL DEFAULT 'command'", + "actor_id": "VARCHAR(255) NOT NULL DEFAULT 'anonymous'", + "client_id": "VARCHAR(255) NOT NULL DEFAULT 'mcp'", + "conversation_id": "VARCHAR(255) NOT NULL DEFAULT 'default'", + "approval_id": "VARCHAR(36)", + }, + } + for table, additions in pg_migrations.items(): + for column, definition in additions.items(): + await conn.execute( + text( + f"ALTER TABLE {table} ADD COLUMN IF NOT EXISTS {column} {definition}" + ) + ) diff --git a/src/shuttle/db/models.py b/src/shuttle/db/models.py index b2f1e21..d619b6a 100644 --- a/src/shuttle/db/models.py +++ b/src/shuttle/db/models.py @@ -115,6 +115,13 @@ class Session(Base): node_id: Mapped[str] = mapped_column( String(36), ForeignKey("nodes.id"), nullable=False ) + actor_id: Mapped[str] = mapped_column( + String(255), nullable=False, default="anonymous" + ) + client_id: Mapped[str] = mapped_column(String(255), nullable=False, default="mcp") + conversation_id: Mapped[str] = mapped_column( + String(255), nullable=False, default="default" + ) working_directory: Mapped[str | None] = mapped_column(String(1024), nullable=True) env_vars: Mapped[dict | None] = mapped_column(JSON, nullable=True) status: Mapped[str] = mapped_column(String(50), nullable=False, default="active") @@ -154,6 +161,15 @@ class CommandLog(Base): String(36), ForeignKey("nodes.id"), nullable=False ) command: Mapped[str] = mapped_column(Text, nullable=False) + action: Mapped[str] = mapped_column(String(50), nullable=False, default="command") + actor_id: Mapped[str] = mapped_column( + String(255), nullable=False, default="anonymous" + ) + client_id: Mapped[str] = mapped_column(String(255), nullable=False, default="mcp") + conversation_id: Mapped[str] = mapped_column( + String(255), nullable=False, default="default" + ) + approval_id: Mapped[str | None] = mapped_column(String(36), nullable=True) exit_code: Mapped[int | None] = mapped_column(Integer, nullable=True) stdout: Mapped[str | None] = mapped_column(Text, nullable=True) stderr: Mapped[str | None] = mapped_column(Text, nullable=True) @@ -174,6 +190,49 @@ class CommandLog(Base): node: Mapped["Node"] = relationship("Node", back_populates="command_logs") +class ApprovalRequest(Base): + """A human decision request for one exact remote action.""" + + __tablename__ = "approval_requests" + __table_args__ = ( + Index("ix_approval_status_created", "status", "created_at"), + Index("ix_approval_actor", "actor_id", "client_id", "conversation_id"), + ) + + id: Mapped[str] = mapped_column( + String(36), primary_key=True, default=lambda: str(uuid.uuid4()) + ) + node_id: Mapped[str] = mapped_column( + String(36), ForeignKey("nodes.id"), nullable=False + ) + node_name: Mapped[str] = mapped_column(String(255), nullable=False) + action: Mapped[str] = mapped_column(String(50), nullable=False, default="command") + command: Mapped[str] = mapped_column(Text, nullable=False) + command_hash: Mapped[str] = mapped_column(String(64), nullable=False) + actor_id: Mapped[str] = mapped_column(String(255), nullable=False) + client_id: Mapped[str] = mapped_column(String(255), nullable=False) + conversation_id: Mapped[str] = mapped_column(String(255), nullable=False) + rule_id: Mapped[str | None] = mapped_column(String(36), nullable=True) + reason: Mapped[str | None] = mapped_column(Text, nullable=True) + status: Mapped[str] = mapped_column(String(50), nullable=False, default="pending") + approver: Mapped[str | None] = mapped_column(String(255), nullable=True) + decision_reason: Mapped[str | None] = mapped_column(Text, nullable=True) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC) + ) + expires_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False + ) + decided_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + consumed_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + + node: Mapped["Node"] = relationship("Node") + + class AppConfig(Base): """Application key-value configuration store.""" diff --git a/src/shuttle/db/repository.py b/src/shuttle/db/repository.py index fa5605f..3907e4b 100644 --- a/src/shuttle/db/repository.py +++ b/src/shuttle/db/repository.py @@ -3,10 +3,17 @@ from datetime import UTC, datetime, timedelta from typing import Any -from sqlalchemy import select +from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession -from shuttle.db.models import AppConfig, CommandLog, Node, SecurityRule, Session +from shuttle.db.models import ( + AppConfig, + ApprovalRequest, + CommandLog, + Node, + SecurityRule, + Session, +) class NodeRepo: @@ -194,12 +201,18 @@ async def create( working_directory: str | None = None, env_vars: dict | None = None, status: str = "active", + actor_id: str = "anonymous", + client_id: str = "mcp", + conversation_id: str = "default", ) -> Session: sess = Session( node_id=node_id, working_directory=working_directory, env_vars=env_vars, status=status, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, ) self._session.add(sess) await self._session.commit() @@ -258,6 +271,11 @@ async def create( security_rule_id: str | None = None, bypassed: bool = False, duration_ms: int | None = None, + action: str = "command", + actor_id: str = "anonymous", + client_id: str = "mcp", + conversation_id: str = "default", + approval_id: str | None = None, ) -> CommandLog: log = CommandLog( session_id=session_id, @@ -270,6 +288,11 @@ async def create( security_rule_id=security_rule_id, bypassed=bypassed, duration_ms=duration_ms, + action=action, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, + approval_id=approval_id, ) self._session.add(log) await self._session.commit() @@ -307,6 +330,87 @@ async def list_by_node( return list(result.scalars().all()) +class ApprovalRepo: + """Persistence and atomic state transitions for human approvals.""" + + def __init__(self, session: AsyncSession) -> None: + self._session = session + + @staticmethod + def _is_expired(request: ApprovalRequest, now: datetime) -> bool: + expires_at = request.expires_at + if expires_at.tzinfo is None: + expires_at = expires_at.replace(tzinfo=UTC) + return expires_at <= now + + async def create(self, **kwargs: Any) -> ApprovalRequest: + request = ApprovalRequest(**kwargs) + self._session.add(request) + await self._session.commit() + await self._session.refresh(request) + return request + + async def get_by_id(self, approval_id: str) -> ApprovalRequest | None: + result = await self._session.execute( + select(ApprovalRequest).where(ApprovalRequest.id == approval_id) + ) + return result.scalar_one_or_none() + + async def list(self, status: str | None = None) -> list[ApprovalRequest]: + stmt = select(ApprovalRequest).order_by(ApprovalRequest.created_at.desc()) + if status is not None: + stmt = stmt.where(ApprovalRequest.status == status) + result = await self._session.execute(stmt) + return list(result.scalars().all()) + + async def decide( + self, + approval_id: str, + *, + approve: bool, + approver: str, + reason: str | None = None, + ) -> ApprovalRequest | None: + request = await self.get_by_id(approval_id) + if request is None or request.status != "pending": + return None + now = datetime.now(UTC) + if self._is_expired(request, now): + request.status = "expired" + else: + request.status = "approved" if approve else "denied" + request.approver = approver + request.decision_reason = reason + request.decided_at = now + await self._session.commit() + await self._session.refresh(request) + return request + + async def consume(self, approval_id: str) -> ApprovalRequest | None: + now = datetime.now(UTC) + # Compare-and-swap: exactly one concurrent caller can transition this + # request from approved to consumed. + result = await self._session.execute( + update(ApprovalRequest) + .where( + ApprovalRequest.id == approval_id, + ApprovalRequest.status == "approved", + ApprovalRequest.consumed_at.is_(None), + ApprovalRequest.expires_at > now, + ) + .values(status="consumed", consumed_at=now) + .execution_options(synchronize_session=False) + ) + await self._session.commit() + if result.rowcount != 1: + request = await self.get_by_id(approval_id) + if request is not None and self._is_expired(request, now): + request.status = "expired" + await self._session.commit() + return None + return await self.get_by_id(approval_id) + + async def cleanup_old_data( session: AsyncSession, command_log_days: int = 30, diff --git a/src/shuttle/mcp/server.py b/src/shuttle/mcp/server.py index 64d1410..03caa35 100644 --- a/src/shuttle/mcp/server.py +++ b/src/shuttle/mcp/server.py @@ -221,6 +221,7 @@ async def create_service_app( host: str = "127.0.0.1", port: int = 9876, api_token: str | None = None, + mcp_token: str | None = None, shuttle_dir: Path | None = None, db_url: str | None = None, ) -> FastAPI: @@ -405,8 +406,8 @@ async def _shutdown(): init_db_deps(api_token=api_token, engine=engine, session_factory=session_factory) # ── FastAPI app ────────────────────────────────────────────────── - # Note: verify_token is applied per-router (not globally) so that - # /mcp/* endpoints are not gated by the web panel token. + # HTTP MCP uses an agent credential separate from the Web/operator token. + # Stdio MCP remains local and does not pass through this ASGI application. from shuttle import __version__ app = FastAPI( @@ -415,6 +416,26 @@ async def _shutdown(): lifespan=combined_lifespan, ) + @app.middleware("http") + async def protect_http_mcp(request, call_next): + if request.url.path == "/mcp" or request.url.path.startswith("/mcp/"): + authorization = request.headers.get("authorization", "") + if not mcp_token or not authorization.startswith("Bearer "): + from starlette.responses import JSONResponse + + return JSONResponse( + {"detail": "Missing MCP bearer token"}, status_code=401 + ) + import secrets + + if not secrets.compare_digest(authorization[7:], mcp_token): + from starlette.responses import JSONResponse + + return JSONResponse( + {"detail": "Invalid MCP bearer token"}, status_code=401 + ) + return await call_next(request) + app.add_middleware( CORSMiddleware, allow_origins=["*"], @@ -424,7 +445,16 @@ async def _shutdown(): ) # API routes — token auth applied per-router so /mcp is not gated - from shuttle.web.routes import data, logs, nodes, rules, sessions, settings, stats + from shuttle.web.routes import ( + approvals, + data, + logs, + nodes, + rules, + sessions, + settings, + stats, + ) api_deps = [Depends(verify_token)] app.include_router(stats.router, prefix="/api", dependencies=api_deps) @@ -434,6 +464,7 @@ async def _shutdown(): app.include_router(logs.router, prefix="/api", dependencies=api_deps) app.include_router(settings.router, prefix="/api", dependencies=api_deps) app.include_router(data.router, prefix="/api", dependencies=api_deps) + app.include_router(approvals.router, prefix="/api", dependencies=api_deps) # Mount MCP — Starlette mount strips trailing path, so sub-app # with path="/" receives requests at /mcp/*. The MCP client posts to diff --git a/src/shuttle/mcp/tools.py b/src/shuttle/mcp/tools.py index 56a1f7b..f592358 100644 --- a/src/shuttle/mcp/tools.py +++ b/src/shuttle/mcp/tools.py @@ -9,13 +9,16 @@ from __future__ import annotations +import hashlib from collections.abc import AsyncIterator, Callable +from datetime import UTC, datetime, timedelta from typing import Any +from fastmcp import Context from loguru import logger from shuttle.core.security import CommandGuard, ConfirmTokenStore, SecurityLevel -from shuttle.core.session import SessionManager +from shuttle.core.session import SessionManager, SSHSession # Truncation limits MAX_OUTPUT_BYTES = 10 * 1024 * 1024 # 10 MB for caller output @@ -48,6 +51,10 @@ async def _execute_command_logic( session_mgr: SessionManager, db_session_ctx: Callable[..., AsyncIterator], node_repo_factory: Callable, + approval_id: str | None = None, + actor_id: str = "anonymous", + client_id: str = "mcp", + conversation_id: str = "default", ) -> str: """Execute a command with security checks, node resolution, and DB logging. @@ -101,24 +108,35 @@ async def _execute_command_logic( "Provide 'node'." ) - # -- 2. Auto-session: find existing or create ---------------------------- - active_sessions = session_mgr.list_active() - node_session = next( - (s for s in active_sessions if s.node_id == resolved_node), None + # Resolve the persisted node before policy/approval so pending requests do + # not open an SSH connection merely to ask a human for permission. + async with db_session_ctx() as db_sess: + repo = node_repo_factory(db_sess) + node_obj = await repo.get_by_name(resolved_node) + if node_obj is None: + return f"Error: node '{resolved_node}' not found." + + # -- 2. Security check ---------------------------------------------------- + node_session = session_mgr.find_active( + resolved_node, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, ) - if node_session: - session_id = node_session.session_id - session_obj = node_session - else: - try: - new_session = await session_mgr.create(resolved_node) - session_id = new_session.session_id - session_obj = new_session - except Exception as exc: - return f"Error: failed to auto-create session — {exc}" - - # -- 3. Security check ---------------------------------------------------- - bypass_patterns = list(session_obj.bypass_patterns) if session_obj else [] + # Test doubles and older embedders may not implement find_active yet. + if not isinstance(node_session, SSHSession): + node_session = next( + ( + session + for session in session_mgr.list_active() + if session.node_id == resolved_node + and getattr(session, "actor_id", "anonymous") == actor_id + and getattr(session, "client_id", "mcp") == client_id + and getattr(session, "conversation_id", "default") == conversation_id + ), + None, + ) + bypass_patterns = list(node_session.bypass_patterns) if node_session else [] async with db_session_ctx() as db_sess: decision = await guard.evaluate( command, resolved_node, db_sess, bypass_patterns @@ -128,30 +146,71 @@ async def _execute_command_logic( return f"⛔ Blocked: {decision.message}" if decision.level == SecurityLevel.CONFIRM: - if confirm_token is None: - # Create a token and ask the caller to confirm - token = token_store.create(command, resolved_node) + from shuttle.db.repository import ApprovalRepo + + if confirm_token is not None: + return "Error: confirm_token is no longer accepted; use a human-approved approval_id." + command_hash = hashlib.sha256(command.encode()).hexdigest() + if approval_id is None: + async with db_session_ctx() as db_sess: + approval = await ApprovalRepo(db_sess).create( + node_id=node_obj.id, + node_name=resolved_node, + command=command, + command_hash=command_hash, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, + rule_id=decision.matched_rule, + reason=decision.message, + expires_at=datetime.now(UTC) + timedelta(minutes=5), + ) return ( - f"⚠️ Confirmation required\n" + f"PENDING_APPROVAL id={approval.id}\n" f"Command: {command}\n" f"Rule: {decision.message}\n" - f"\n" - f'To proceed: ssh_run(command="{command}", node="{resolved_node}", confirm_token="{token}")' + "A human must approve this request in the Shuttle Web panel. " + "Then retry with the same command and approval_id." ) - - # Validate the provided token - if not token_store.validate(confirm_token, command, resolved_node): - return "Error: invalid or expired confirmation token." - - # Token valid — optionally add bypass for this session - if bypass_scope == "session" and session_obj and decision.matched_rule: - async with db_session_ctx() as db_sess: - from shuttle.db.repository import RuleRepo - - rule_repo = RuleRepo(db_sess) - matched = await rule_repo.get_by_id(decision.matched_rule) - if matched: - session_obj.bypass_patterns.add(matched.pattern) + async with db_session_ctx() as db_sess: + approval_repo = ApprovalRepo(db_sess) + approval = await approval_repo.get_by_id(approval_id) + if ( + approval is None + or approval.node_id != node_obj.id + or approval.command_hash != command_hash + or approval.actor_id != actor_id + or approval.client_id != client_id + or approval.conversation_id != conversation_id + ): + return "Error: approval does not match this actor, conversation, node, and command." + consumed = await approval_repo.consume(approval_id) + if consumed is None: + status = approval.status if approval is not None else "missing" + return f"PENDING_APPROVAL id={approval_id} status={status}" + + # -- 3. Caller-isolated auto-session -------------------------------------- + if node_session: + session_id = node_session.session_id + session_obj = node_session + else: + try: + if (actor_id, client_id, conversation_id) == ( + "anonymous", + "mcp", + "default", + ): + session_obj = await session_mgr.create(resolved_node) + else: + session_obj = await session_mgr.create( + resolved_node, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, + ) + session_id = session_obj.session_id + except Exception as exc: + return f"Error: failed to auto-create session — {exc}" if decision.level == SecurityLevel.WARN: logger.warning( @@ -177,22 +236,14 @@ async def _execute_command_logic( # -- 5. Persist command log to DB ----------------------------------------- try: - # Resolve node UUID for the FK - node_uuid: str | None = None - async with db_session_ctx() as db_sess: - repo = node_repo_factory(db_sess) - node_obj = await repo.get_by_name(resolved_node) - if node_obj: - node_uuid = node_obj.id - - if node_uuid: + if node_obj.id: db_stdout = _truncate(stdout, MAX_DB_OUTPUT_BYTES) if stdout else None async with db_session_ctx() as db_sess: from shuttle.db.repository import LogRepo log_repo = LogRepo(db_sess) await log_repo.create( - node_id=node_uuid, + node_id=node_obj.id, session_id=session_id, command=command, exit_code=exit_status, @@ -200,20 +251,26 @@ async def _execute_command_logic( stderr=None, security_level=decision.level.value if decision else None, security_rule_id=decision.matched_rule if decision else None, - bypassed=confirm_token is not None, + bypassed=approval_id is not None, duration_ms=duration_ms, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, + approval_id=approval_id, ) # Update node last_seen_at async with db_session_ctx() as db_sess: repo = node_repo_factory(db_sess) - from datetime import UTC, datetime - await repo.update( node_obj.id, last_seen_at=datetime.now(UTC), status="active" ) - except Exception: - logger.warning("Failed to persist command log for {cmd}", cmd=command[:80]) + except Exception as exc: + logger.warning( + "Failed to persist command log for {cmd}: {error}", + cmd=command[:80], + error=exc, + ) return stdout @@ -263,6 +320,8 @@ async def ssh_run( timeout: float = 30.0, confirm_token: str | None = None, bypass_scope: str | None = None, + approval_id: str | None = None, + ctx: Context | None = None, ) -> str: """Execute a shell command on a remote SSH node. @@ -270,12 +329,30 @@ async def ssh_run( across calls to the same node. Security checks (BLOCK / CONFIRM / WARN / ALLOW) are applied before execution. """ + actor_id = "local-stdio" + client_id = "mcp" + conversation_id = "default" + if ctx is not None: + client_id = ctx.client_id or "mcp" + conversation_id = ctx.session_id + try: + authorization = ctx.get_http_request().headers.get("authorization", "") + if authorization.startswith("Bearer "): + actor_id = "agent:" + hashlib.sha256( + authorization[7:].encode() + ).hexdigest()[:16] + except RuntimeError: + actor_id = "local-stdio" return await _execute_command_logic( command=command, node=node, timeout=timeout, confirm_token=confirm_token, bypass_scope=bypass_scope, + approval_id=approval_id, + actor_id=actor_id, + client_id=client_id, + conversation_id=conversation_id, pool=pool, guard=guard, token_store=token_store, diff --git a/src/shuttle/web/app.py b/src/shuttle/web/app.py index ecd3415..a31ed61 100644 --- a/src/shuttle/web/app.py +++ b/src/shuttle/web/app.py @@ -46,7 +46,16 @@ def create_app( allow_headers=["*"], ) - from shuttle.web.routes import data, logs, nodes, rules, sessions, settings, stats + from shuttle.web.routes import ( + approvals, + data, + logs, + nodes, + rules, + sessions, + settings, + stats, + ) app.include_router(stats.router, prefix="/api") app.include_router(nodes.router, prefix="/api") @@ -55,6 +64,7 @@ def create_app( app.include_router(logs.router, prefix="/api") app.include_router(settings.router, prefix="/api") app.include_router(data.router, prefix="/api") + app.include_router(approvals.router, prefix="/api") static_dir = Path(__file__).parent / "static" if static_dir.is_dir() and (static_dir / "index.html").exists(): diff --git a/src/shuttle/web/routes/approvals.py b/src/shuttle/web/routes/approvals.py new file mode 100644 index 0000000..49a60fd --- /dev/null +++ b/src/shuttle/web/routes/approvals.py @@ -0,0 +1,48 @@ +"""Human approval queue endpoints.""" + +from fastapi import APIRouter, Depends, HTTPException, Query +from sqlalchemy.ext.asyncio import AsyncSession + +from shuttle.db.repository import ApprovalRepo +from shuttle.web.deps import get_db_session +from shuttle.web.schemas import ApprovalDecision, ApprovalResponse + +router = APIRouter(prefix="/approvals", tags=["approvals"]) + + +@router.get("", response_model=list[ApprovalResponse]) +async def list_approvals( + status: str | None = Query( + None, pattern=r"^(pending|approved|denied|consumed|expired)$" + ), + db: AsyncSession = Depends(get_db_session), +): + return await ApprovalRepo(db).list(status) + + +@router.get("/{approval_id}", response_model=ApprovalResponse) +async def get_approval( + approval_id: str, + db: AsyncSession = Depends(get_db_session), +): + request = await ApprovalRepo(db).get_by_id(approval_id) + if request is None: + raise HTTPException(404, "Approval request not found") + return request + + +@router.post("/{approval_id}/decision", response_model=ApprovalResponse) +async def decide_approval( + approval_id: str, + body: ApprovalDecision, + db: AsyncSession = Depends(get_db_session), +): + request = await ApprovalRepo(db).decide( + approval_id, + approve=body.approve, + approver=body.approver, + reason=body.reason, + ) + if request is None: + raise HTTPException(409, "Approval is missing or no longer pending") + return request diff --git a/src/shuttle/web/routes/logs.py b/src/shuttle/web/routes/logs.py index 95facef..7b6f27e 100644 --- a/src/shuttle/web/routes/logs.py +++ b/src/shuttle/web/routes/logs.py @@ -32,6 +32,11 @@ def _log_to_response(log: CommandLog, node_names: dict[str, str]) -> dict: "node_id": log.node_id, "node_name": node_names.get(log.node_id), "command": log.command, + "action": log.action, + "actor_id": log.actor_id, + "client_id": log.client_id, + "conversation_id": log.conversation_id, + "approval_id": log.approval_id, "exit_code": log.exit_code, "stdout": log.stdout, "stderr": log.stderr, diff --git a/src/shuttle/web/schemas.py b/src/shuttle/web/schemas.py index 7111562..16bbfc1 100644 --- a/src/shuttle/web/schemas.py +++ b/src/shuttle/web/schemas.py @@ -102,6 +102,9 @@ class SessionResponse(BaseModel): id: str node_id: str node_name: str | None = None + actor_id: str = "anonymous" + client_id: str = "mcp" + conversation_id: str = "default" working_directory: str | None status: str created_at: datetime @@ -119,6 +122,11 @@ class CommandLogResponse(BaseModel): node_id: str node_name: str | None = None command: str + action: str = "command" + actor_id: str = "anonymous" + client_id: str = "mcp" + conversation_id: str = "default" + approval_id: str | None = None exit_code: int | None stdout: str | None stderr: str | None @@ -137,6 +145,37 @@ class LogListResponse(BaseModel): page_size: int +# ── Human Approvals ──────────────────────────────── + + +class ApprovalResponse(BaseModel): + id: str + node_id: str + node_name: str + action: str + command: str + actor_id: str + client_id: str + conversation_id: str + rule_id: str | None + reason: str | None + status: str + approver: str | None + decision_reason: str | None + created_at: datetime + expires_at: datetime + decided_at: datetime | None + consumed_at: datetime | None + + model_config = {"from_attributes": True} + + +class ApprovalDecision(BaseModel): + approve: bool + approver: str = Field(..., min_length=1, max_length=255) + reason: str | None = Field(None, max_length=2000) + + # ── Settings ─────────────────────────────────────── diff --git a/tests/test_core/test_proxy.py b/tests/test_core/test_proxy.py index 228405e..8c16d6a 100644 --- a/tests/test_core/test_proxy.py +++ b/tests/test_core/test_proxy.py @@ -16,7 +16,7 @@ def test_build_connect_kwargs_minimal() -> None: assert kw["port"] == 22 assert kw["username"] == "u" assert kw["connect_timeout"] == 30.0 - assert kw["known_hosts"] is None + assert kw["known_hosts"].endswith("/.ssh/known_hosts") def test_build_connect_kwargs_password_and_known_hosts() -> None: @@ -66,6 +66,17 @@ def test_build_connect_kwargs_extra_options_override() -> None: assert kw["compression_algs"] == () +def test_build_connect_kwargs_rejects_disabled_host_key_checking() -> None: + info = NodeConnectInfo( + node_id="n1", + hostname="h", + username="u", + known_hosts=None, + ) + with pytest.raises(ValueError, match="known_hosts is required"): + _build_connect_kwargs(info) + + @pytest.mark.asyncio async def test_connect_ssh_direct_calls_asyncssh_connect() -> None: info = NodeConnectInfo(node_id="n1", hostname="h", username="u") diff --git a/tests/test_core/test_session_identity.py b/tests/test_core/test_session_identity.py new file mode 100644 index 0000000..e7e082f --- /dev/null +++ b/tests/test_core/test_session_identity.py @@ -0,0 +1,43 @@ +"""Session identity isolation tests.""" + +from unittest.mock import MagicMock + +from shuttle.core.session import SessionManager, SSHSession + + +def test_find_active_requires_complete_actor_boundary() -> None: + manager = SessionManager(pool=MagicMock()) + claude = SSHSession( + session_id="claude-1", + node_id="gpu-1", + actor_id="alice", + client_id="claude-code", + conversation_id="conv-a", + ) + codex = SSHSession( + session_id="codex-1", + node_id="gpu-1", + actor_id="alice", + client_id="codex", + conversation_id="conv-b", + ) + manager._sessions = {claude.session_id: claude, codex.session_id: codex} + + assert ( + manager.find_active( + "gpu-1", + actor_id="alice", + client_id="claude-code", + conversation_id="conv-a", + ) + is claude + ) + assert ( + manager.find_active( + "gpu-1", + actor_id="alice", + client_id="claude-code", + conversation_id="conv-b", + ) + is None + ) diff --git a/tests/test_db/test_approval_concurrency.py b/tests/test_db/test_approval_concurrency.py new file mode 100644 index 0000000..a561082 --- /dev/null +++ b/tests/test_db/test_approval_concurrency.py @@ -0,0 +1,46 @@ +"""Approval concurrency tests.""" + +import asyncio +from datetime import UTC, datetime, timedelta + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine +from sqlalchemy.orm import sessionmaker + +from shuttle.db.models import Base +from shuttle.db.repository import ApprovalRepo, NodeRepo + + +@pytest.mark.asyncio +async def test_approval_consumption_is_atomic(tmp_path): + engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'race.db'}") + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + factory = sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + async with factory() as session: + node = await NodeRepo(session).create( + name="race-node", + host="127.0.0.1", + username="runner", + encrypted_credential="enc", + ) + request = await ApprovalRepo(session).create( + node_id=node.id, + node_name=node.name, + command="sudo true", + command_hash="d" * 64, + actor_id="alice", + client_id="hermes", + conversation_id="race", + status="approved", + approver="operator", + expires_at=datetime.now(UTC) + timedelta(minutes=5), + ) + + async def consume_once(): + async with factory() as session: + return await ApprovalRepo(session).consume(request.id) + + results = await asyncio.gather(*[consume_once() for _ in range(20)]) + assert sum(result is not None for result in results) == 1 + await engine.dispose() diff --git a/tests/test_mcp/test_execute_logic_more.py b/tests/test_mcp/test_execute_logic_more.py index 1bf6d9f..bc91f9c 100644 --- a/tests/test_mcp/test_execute_logic_more.py +++ b/tests/test_mcp/test_execute_logic_more.py @@ -138,7 +138,7 @@ async def test_execute_confirm_invalid_token() -> None: db_session_ctx=_noop_db_ctx, node_repo_factory=_node_repo_factory, ) - assert "invalid" in out.lower() + assert "no longer accepted" in out.lower() @pytest.mark.asyncio @@ -244,39 +244,28 @@ async def test_execute_persists_command_log_to_db() -> None: @pytest.mark.asyncio -async def test_execute_persists_log_with_confirm_bypassed() -> None: - """When confirm_token is provided, bypassed=True in the log entry.""" +async def test_execute_rejects_legacy_confirm_token() -> None: + """An agent-provided confirmation token can no longer authorize execution.""" session = SSHSession(session_id="s1", node_id="n1") guard = _make_guard(SecurityLevel.CONFIRM, message="sudo", rule="r1") - - ts = MagicMock(spec=ConfirmTokenStore) - ts.validate.return_value = True - mgr = _sm_with_session(session, stdout="root") - mock_log_repo = MagicMock() - mock_log_repo.create = AsyncMock() - - with patch("shuttle.db.repository.LogRepo", return_value=mock_log_repo): - out = await _execute_command_logic( - command="sudo whoami", - node="n1", - timeout=10, - confirm_token="valid-token", - bypass_scope=None, - pool=MagicMock(), - guard=guard, - token_store=ts, - session_mgr=mgr, - db_session_ctx=_noop_db_ctx, - node_repo_factory=_node_repo_factory, - ) + out = await _execute_command_logic( + command="sudo whoami", + node="n1", + timeout=10, + confirm_token="valid-token", + bypass_scope=None, + pool=MagicMock(), + guard=guard, + token_store=MagicMock(spec=ConfirmTokenStore), + session_mgr=mgr, + db_session_ctx=_noop_db_ctx, + node_repo_factory=_node_repo_factory, + ) - assert out == "root" - mock_log_repo.create.assert_awaited_once() - kwargs = mock_log_repo.create.call_args.kwargs - assert kwargs["bypassed"] is True - assert kwargs["security_level"] == "confirm" + assert "no longer accepted" in out + mgr.execute.assert_not_awaited() @pytest.mark.asyncio @@ -360,8 +349,8 @@ async def _update(*args, **kwargs): @pytest.mark.asyncio -async def test_execute_still_returns_stdout_when_db_logging_fails() -> None: - """If DB logging raises, the command output must still be returned.""" +async def test_execute_returns_error_when_node_lookup_fails() -> None: + """A failed node lookup must fail closed before remote execution.""" session = SSHSession(session_id="s1", node_id="n1") guard = _make_guard(SecurityLevel.ALLOW) mgr = _sm_with_session(session, stdout="important output") @@ -372,20 +361,21 @@ def _broken_repo_factory(db_sess): repo.list_all = AsyncMock(return_value=[]) return repo - out = await _execute_command_logic( - command="echo test", - node="n1", - timeout=10, - confirm_token=None, - bypass_scope=None, - pool=MagicMock(), - guard=guard, - token_store=MagicMock(spec=ConfirmTokenStore), - session_mgr=mgr, - db_session_ctx=_noop_db_ctx, - node_repo_factory=_broken_repo_factory, - ) - assert out == "important output" + with pytest.raises(RuntimeError, match="DB down"): + await _execute_command_logic( + command="echo test", + node="n1", + timeout=10, + confirm_token=None, + bypass_scope=None, + pool=MagicMock(), + guard=guard, + token_store=MagicMock(spec=ConfirmTokenStore), + session_mgr=mgr, + db_session_ctx=_noop_db_ctx, + node_repo_factory=_broken_repo_factory, + ) + mgr.execute.assert_not_awaited() @pytest.mark.asyncio diff --git a/tests/test_mcp/test_fastmcp_client.py b/tests/test_mcp/test_fastmcp_client.py index 6aa701c..e78c819 100644 --- a/tests/test_mcp/test_fastmcp_client.py +++ b/tests/test_mcp/test_fastmcp_client.py @@ -213,7 +213,7 @@ async def test_run_blocked_command_via_client(mcp_server_with_session, db_factor @pytest.mark.asyncio async def test_run_confirm_flow_via_client(mcp_server_with_session, db_factory): - """CONFIRM-level command returns token request, re-call with token executes.""" + """CONFIRM-level command requires a separate persisted human decision.""" async with db_factory() as sess: rule_repo = RuleRepo(sess) await rule_repo.create( @@ -225,20 +225,41 @@ async def test_run_confirm_flow_via_client(mcp_server_with_session, db_factory): async with Client(mcp_server_with_session) as client: result1 = await client.call_tool( - "ssh_run", {"command": "sudo ls", "node": "test-node"} + "ssh_run", + { + "command": "sudo ls", + "node": "test-node", + }, ) text1 = _result_text(result1) - assert "confirm_token" in text1 + assert "PENDING_APPROVAL" in text1 + assert "confirm_token" not in text1 import re - match = re.search(r'confirm_token="([^"]+)"', text1) - assert match, f"Could not find confirm_token in: {text1}" - token = match.group(1) + match = re.search(r"id=([0-9a-f-]{36})", text1) + assert match, f"Could not find approval id in: {text1}" + approval_id = match.group(1) + + from shuttle.db.repository import ApprovalRepo + + async with db_factory() as sess: + approved = await ApprovalRepo(sess).decide( + approval_id, + approve=True, + approver="human@example.com", + reason="Reviewed in Web", + ) + assert approved is not None + assert approved.status == "approved" result2 = await client.call_tool( "ssh_run", - {"command": "sudo ls", "node": "test-node", "confirm_token": token}, + { + "command": "sudo ls", + "node": "test-node", + "approval_id": approval_id, + }, ) text2 = _result_text(result2) assert "mocked output" in text2 diff --git a/tests/test_mcp/test_service_app.py b/tests/test_mcp/test_service_app.py index 97887be..6b5da2c 100644 --- a/tests/test_mcp/test_service_app.py +++ b/tests/test_mcp/test_service_app.py @@ -41,6 +41,7 @@ async def test_create_service_app_exposes_stats_with_bearer(tmp_path): shuttle_dir=shuttle_dir, db_url=db_url, api_token=token, + mcp_token="agent-bearer-token", port=19999, ) @@ -54,10 +55,17 @@ async def test_create_service_app_exposes_stats_with_bearer(tmp_path): data = r.json() assert "node_count" in data - redir = await client.get("/mcp", follow_redirects=False) + redir = await client.get( + "/mcp", + headers={"Authorization": "Bearer agent-bearer-token"}, + follow_redirects=False, + ) assert redir.status_code == 307 assert redir.headers.get("location", "").endswith("/mcp/") + unauthenticated = await client.get("/mcp", follow_redirects=False) + assert unauthenticated.status_code == 401 + @pytest.mark.asyncio async def test_create_service_app_rejects_bad_bearer(tmp_path): @@ -75,6 +83,7 @@ async def test_create_service_app_rejects_bad_bearer(tmp_path): shuttle_dir=shuttle_dir, db_url=db_url, api_token="good", + mcp_token="agent-good", ) transport = ASGITransport(app=app) diff --git a/tests/test_mcp/test_tools.py b/tests/test_mcp/test_tools.py index a6a040c..36a44cd 100644 --- a/tests/test_mcp/test_tools.py +++ b/tests/test_mcp/test_tools.py @@ -147,31 +147,37 @@ async def test_blocked_command_rejected(): @pytest.mark.asyncio -async def test_confirm_returns_token_request(): - """When guard returns CONFIRM without a token, a confirmation message with token is returned.""" +async def test_confirm_requires_persisted_human_approval() -> None: + """CONFIRM never returns a secret which the calling agent can replay.""" session = SSHSession(session_id="s1", node_id="node1") guard = _make_guard(SecurityLevel.CONFIRM, message="Needs approval") - token_store = _make_token_store() - token_store.create.return_value = "tok_abc123" session_mgr = _make_session_mgr(session=session) - result = await _execute_command_logic( - command="shutdown -h now", - node="node1", - timeout=30, - confirm_token=None, - bypass_scope=None, - pool=MagicMock(), - guard=guard, - token_store=token_store, - session_mgr=session_mgr, - db_session_ctx=_noop_db_session, - node_repo_factory=_noop_node_repo_factory, - ) + class FakeApproval: + id = "approval-123" + + fake_repo = MagicMock() + fake_repo.create = AsyncMock(return_value=FakeApproval()) + + from unittest.mock import patch + + with patch("shuttle.db.repository.ApprovalRepo", return_value=fake_repo): + result = await _execute_command_logic( + command="shutdown -h now", + node="node1", + timeout=30, + confirm_token=None, + bypass_scope=None, + pool=MagicMock(), + guard=guard, + token_store=_make_token_store(), + session_mgr=session_mgr, + db_session_ctx=_noop_db_session, + node_repo_factory=_noop_node_repo_factory, + ) - assert "confirm_token" in result - assert "tok_abc123" in result - token_store.create.assert_called_once_with("shutdown -h now", "node1") + assert "PENDING_APPROVAL id=approval-123" in result + assert "confirm_token" not in result session_mgr.execute.assert_not_awaited() diff --git a/tests/test_web/test_approvals_api.py b/tests/test_web/test_approvals_api.py new file mode 100644 index 0000000..5f3a7ba --- /dev/null +++ b/tests/test_web/test_approvals_api.py @@ -0,0 +1,74 @@ +"""Approval queue API tests.""" + +from datetime import UTC, datetime, timedelta + +import pytest + +from shuttle.db.repository import ApprovalRepo, NodeRepo + + +@pytest.mark.asyncio +async def test_pending_approval_can_only_be_decided_once(client, db_session): + node = await NodeRepo(db_session).create( + name="gpu-1", + host="10.0.0.1", + username="runner", + encrypted_credential="enc", + ) + request = await ApprovalRepo(db_session).create( + node_id=node.id, + node_name=node.name, + action="command", + command="sudo systemctl restart trainer", + command_hash="a" * 64, + actor_id="alice", + client_id="claude-code", + conversation_id="conv-1", + reason="Service restart", + expires_at=datetime.now(UTC) + timedelta(minutes=5), + ) + + listed = await client.get("/api/approvals?status=pending") + assert listed.status_code == 200 + assert listed.json()[0]["id"] == request.id + + decided = await client.post( + f"/api/approvals/{request.id}/decision", + json={"approve": True, "approver": "human@example.com", "reason": "Reviewed"}, + ) + assert decided.status_code == 200 + assert decided.json()["status"] == "approved" + assert decided.json()["approver"] == "human@example.com" + + duplicate = await client.post( + f"/api/approvals/{request.id}/decision", + json={"approve": False, "approver": "other@example.com"}, + ) + assert duplicate.status_code == 409 + + +@pytest.mark.asyncio +async def test_expired_approval_cannot_be_granted(client, db_session): + node = await NodeRepo(db_session).create( + name="gpu-expired", + host="10.0.0.2", + username="runner", + encrypted_credential="enc", + ) + request = await ApprovalRepo(db_session).create( + node_id=node.id, + node_name=node.name, + command="shutdown -h now", + command_hash="b" * 64, + actor_id="bob", + client_id="codex", + conversation_id="conv-2", + expires_at=datetime.now(UTC) - timedelta(seconds=1), + ) + + response = await client.post( + f"/api/approvals/{request.id}/decision", + json={"approve": True, "approver": "human@example.com"}, + ) + assert response.status_code == 200 + assert response.json()["status"] == "expired" diff --git a/web/src/App.tsx b/web/src/App.tsx index dfc6335..d2ec82d 100644 --- a/web/src/App.tsx +++ b/web/src/App.tsx @@ -9,6 +9,7 @@ import Overview from "./pages/Overview"; import Activity from "./pages/Activity"; import Rules from "./pages/Rules"; import Settings from "./pages/Settings"; +import Approvals from "./pages/Approvals"; export default function App() { const [authed, setAuthed] = useState(!!getToken()); @@ -34,6 +35,7 @@ export default function App() { } /> } /> } /> + } /> diff --git a/web/src/api/client.ts b/web/src/api/client.ts index 6775f93..560c98c 100644 --- a/web/src/api/client.ts +++ b/web/src/api/client.ts @@ -16,6 +16,7 @@ import type { StatsResponse, SettingsResponse, SettingsUpdate, + ApprovalResponse, } from "../types"; // ── Fetch wrapper ────────────────────────────────── @@ -73,6 +74,7 @@ const keys = { status ? (["sessions", status] as const) : (["sessions"] as const), logs: (params?: LogParams) => ["logs", params] as const, settings: ["settings"] as const, + approvals: (status?: string) => ["approvals", status ?? "all"] as const, }; // ── Stats ────────────────────────────────────────── @@ -226,6 +228,41 @@ export function useCloseSession() { }); } +// ── Human approvals ──────────────────────────────── + +export function useApprovals(status = "pending") { + return useQuery({ + queryKey: keys.approvals(status), + queryFn: () => apiFetch(`/approvals?status=${encodeURIComponent(status)}`), + refetchInterval: 3000, + }); +} + +export function useDecideApproval() { + const qc = useQueryClient(); + return useMutation< + ApprovalResponse, + Error, + { id: string; approve: boolean; approver: string; reason?: string } + >({ + mutationFn: (input: { + id: string; + approve: boolean; + approver: string; + reason?: string; + }) => { + const { id, ...body } = input; + return apiFetch(`/approvals/${id}/decision`, { + method: "POST", + body: JSON.stringify(body), + }); + }, + onSuccess: () => { + void qc.invalidateQueries({ queryKey: ["approvals"] }); + }, + }); +} + // ── Logs ─────────────────────────────────────────── export interface LogParams { diff --git a/web/src/components/Sidebar.tsx b/web/src/components/Sidebar.tsx index 35d255f..d55302f 100644 --- a/web/src/components/Sidebar.tsx +++ b/web/src/components/Sidebar.tsx @@ -1,5 +1,5 @@ import { NavLink, useLocation } from "react-router-dom"; -import { Shield, Settings, Server, Sun, Moon } from "lucide-react"; +import { Shield, Settings, Server, Sun, Moon, CircleCheckBig } from "lucide-react"; import clsx from "clsx"; import { useApp } from "../hooks/AppContext"; @@ -102,6 +102,10 @@ export default function Sidebar() { Rules + navItemCls(isActive)}> + + Approvals + navItemCls(isActive)}> Settings @@ -110,7 +114,7 @@ export default function Sidebar() { {/* Footer */}
-

Shuttle MCP v2

+

Shuttle approval plane

); diff --git a/web/src/pages/Approvals.tsx b/web/src/pages/Approvals.tsx new file mode 100644 index 0000000..f9c26c3 --- /dev/null +++ b/web/src/pages/Approvals.tsx @@ -0,0 +1,103 @@ +import { useState } from "react"; +import { Check, X, RefreshCw, ShieldCheck } from "lucide-react"; +import { toast } from "sonner"; +import { useApprovals, useDecideApproval } from "../api/client"; + +export default function Approvals() { + const [approver, setApprover] = useState("operator"); + const { data = [], isLoading, refetch } = useApprovals("pending"); + const decide = useDecideApproval(); + + const submit = async (id: string, approve: boolean) => { + try { + await decide.mutateAsync({ id, approve, approver }); + toast.success(approve ? "Approved once" : "Denied"); + } catch (error) { + toast.error(error instanceof Error ? error.message : "Decision failed"); + } + }; + + return ( +
+
+
+
+

+ Pending approvals +

+

+ Decisions are bound to one actor, conversation, node, and exact command. +

+
+ +
+ + + + {isLoading &&

Loading…

} + {!isLoading && data.length === 0 && ( +
+ +

+ No commands are waiting for approval. +

+
+ )} +
+ {data.map((request) => ( +
+
+ + {request.node_name} + + {request.actor_id}· + {request.client_id}· + {request.conversation_id} +
+
+                {request.command}
+              
+ {request.reason && ( +

+ Policy: {request.reason} +

+ )} +
+ + +
+
+ ))} +
+
+
+ ); +} diff --git a/web/src/types/index.ts b/web/src/types/index.ts index cbec973..028da3e 100644 --- a/web/src/types/index.ts +++ b/web/src/types/index.ts @@ -89,6 +89,11 @@ export interface CommandLogResponse { node_id: string; node_name: string | null; command: string; + action: string; + actor_id: string; + client_id: string; + conversation_id: string; + approval_id: string | null; exit_code: number | null; stdout: string | null; stderr: string | null; @@ -105,6 +110,26 @@ export interface LogListResponse { page_size: number; } +export interface ApprovalResponse { + id: string; + node_id: string; + node_name: string; + action: string; + command: string; + actor_id: string; + client_id: string; + conversation_id: string; + rule_id: string | null; + reason: string | null; + status: "pending" | "approved" | "denied" | "consumed" | "expired"; + approver: string | null; + decision_reason: string | null; + created_at: string; + expires_at: string; + decided_at: string | null; + consumed_at: string | null; +} + // ── Settings ─────────────────────────────────────── export interface SettingsResponse {