diff --git a/backend/app/api/agents.py b/backend/app/api/agents.py index ef71b33..a8e0a52 100644 --- a/backend/app/api/agents.py +++ b/backend/app/api/agents.py @@ -102,7 +102,12 @@ def invoke_agent(agent_id: str, req: AgentInvokeRequest, user: Principal = Depen source="api", target=f"agent:{agent_id}", prompt=req.prompt, - runtime_session_id=req.session_id, + # namespaced under the caller so a client-supplied session_id + # cannot land on another tenant's warm microVM (see + # resolve_session_id / resolve_memory_actor) + runtime_session_id=invocation_service.resolve_session_id( + user, req.session_id + ), # a published agent carries its own memory_id, so the actor ID is # all that separates callers' memory lines (see resolve_memory_actor) memory_actor_id=invocation_service.resolve_memory_actor( diff --git a/backend/app/api/kernels.py b/backend/app/api/kernels.py index 984b931..c582780 100644 --- a/backend/app/api/kernels.py +++ b/backend/app/api/kernels.py @@ -26,7 +26,13 @@ def invoke_sdk_kernel(req: InvokeRequest, user: Principal = Depends(get_current_ prompt=req.prompt, system=req.system, max_turns=req.max_turns, - runtime_session_id=req.session_id, + # a caller-submitted session_id selects which warm microVM/process + # (and its /tmp, secret cache and gateway grant) the call lands on, + # so it is namespaced under the caller — never passed through + # verbatim (see resolve_session_id / resolve_memory_actor) + runtime_session_id=invocation_service.resolve_session_id( + user, req.session_id + ), mcp_server_ids=req.mcp_server_ids, skill_ids=req.skill_ids, memory_id=req.memory_id, diff --git a/backend/app/config.py b/backend/app/config.py index b6c8ed7..56cf623 100644 --- a/backend/app/config.py +++ b/backend/app/config.py @@ -113,6 +113,16 @@ class Settings(BaseSettings): # Dockerfile creates it; empty = keep the backend's own identity. workflow_runner_user: str = "workflow" + # Keys the caller-binding of client-controllable AgentCore session ids + # (see session_binding). A caller-submitted session_id and a channel + # conversation id are both folded through an HMAC under this secret so + # they cannot resolve onto another tenant's warm microVM, and a channel + # session id cannot be predicted offline from the public webhook URL. + # Empty = a fallback keyed on deployment identifiers (runtime ARNs, table + # name); set explicitly in production. Only needs to be stable across + # replicas and restarts. + session_binding_secret: str = "" + # CORS origins for the portal frontend cors_origins: str = "http://localhost:5173" diff --git a/backend/app/services/channel_service.py b/backend/app/services/channel_service.py index 7799578..95bdebd 100644 --- a/backend/app/services/channel_service.py +++ b/backend/app/services/channel_service.py @@ -11,7 +11,6 @@ context. """ -import hashlib import hmac import logging import re @@ -22,6 +21,7 @@ import boto3 from app.config import settings +from app.services.session_binding import derive_channel_session_id logger = logging.getLogger(__name__) @@ -166,8 +166,11 @@ def _route(self, item: dict, channel_id: str, *, user: str, message: str, runtime_session_id = None memory_actor_id = "" if conversation_id: - digest = hashlib.sha256(f"{channel_id}:{conversation_id}".encode()).hexdigest() - runtime_session_id = f"chn-{digest[:44]}" # ≥33 chars for AgentCore + # Keyed with the platform binding secret: the public channel_id is + # in the webhook URL, so an unkeyed digest of channel_id + + # conversation_id would let anyone compute this session id offline + # and land on the conversation's warm microVM. + runtime_session_id = derive_channel_session_id(channel_id, conversation_id) # One memory line per conversation (not per channel): the actor is # channel-scoped so equal conversation_ids on different channels # don't share memory. diff --git a/backend/app/services/invocation_service.py b/backend/app/services/invocation_service.py index 69681a3..4e0fffe 100644 --- a/backend/app/services/invocation_service.py +++ b/backend/app/services/invocation_service.py @@ -23,9 +23,13 @@ from app.services.kernel_service import kernel_service from app.services.model_config_service import model_config_service from app.services.observability_service import observability_service +from app.services.session_binding import resolve_session_id # re-exported for the API routes logger = logging.getLogger(__name__) +__all__ = ["invoke", "invoke_async_and_wait", "resolve_memory_actor", + "resolve_session_id", "forward_identity", "IdentityRequired"] + class IdentityRequired(Exception): """An attached MCP server forwards the caller's identity, but this call diff --git a/backend/app/services/llm_credentials_service.py b/backend/app/services/llm_credentials_service.py index 77d5b79..8bceaf4 100644 --- a/backend/app/services/llm_credentials_service.py +++ b/backend/app/services/llm_credentials_service.py @@ -113,22 +113,39 @@ def mint( token = secrets.token_urlsafe(32) expires_at = int(time.time()) + max(60, int(ttl_s)) - self.table.put_item( - Item={ - "PK": LLM_TOKEN_PK, - "SK": f"RSID#{runtime_session_id}", - # Only the digest is stored: a reader of this table cannot - # replay the grant. - "token_sha256": _sha256(token), - "expires_at": expires_at, - "runtime_session_id": runtime_session_id, - "user": user, - "team": team, - "upstream_base_url": base_url, - "gateway_secret_name": secret_name, - "allowed_models": models, - } - ) + try: + self.table.put_item( + Item={ + "PK": LLM_TOKEN_PK, + "SK": f"RSID#{runtime_session_id}", + # Only the digest is stored: a reader of this table cannot + # replay the grant. + "token_sha256": _sha256(token), + "expires_at": expires_at, + "runtime_session_id": runtime_session_id, + "user": user, + "team": team, + "upstream_base_url": base_url, + "gateway_secret_name": secret_name, + "allowed_models": models, + }, + # A grant belongs to the session's owner. Only mint a new one + # or re-mint the owner's own — never overwrite a live grant held + # by a different principal. Session ids are already caller-bound + # upstream (see session_binding), so a collision here should be + # unreachable; this makes overwriting another user's grant (and + # the targeted DoS it enables) impossible even if that breaks. + ConditionExpression="attribute_not_exists(PK) OR #u = :user", + ExpressionAttributeNames={"#u": "user"}, + ExpressionAttributeValues={":user": user}, + ) + except self.table.meta.client.exceptions.ConditionalCheckFailedException: + logger.error( + "refusing to mint gateway credentials for %s: session grant is " + "held by a different principal", + runtime_session_id, + ) + return None return { "endpoint": settings.llm_edge_url, # Echoed back by the kernel as the x-platform-session-id header so diff --git a/backend/app/services/session_binding.py b/backend/app/services/session_binding.py new file mode 100644 index 0000000..dfbd939 --- /dev/null +++ b/backend/app/services/session_binding.py @@ -0,0 +1,108 @@ +"""Caller-binding for client-controllable AgentCore session identifiers. + +An AgentCore ``runtimeSessionId`` is not a label: reusing one routes to the +*same warm microVM and the same kernel process*, so two callers that share a +runtime session id share ``/tmp`` working files, the process secret cache and +any in-process model-gateway grant. That makes the id an authorization +boundary in exactly the way the memory actor id is (see +``invocation_service.resolve_memory_actor``) — and it must be treated the same +way: a value the caller can influence must never resolve onto another tenant's +session. + +Two entry points, mirroring the two ways a session id enters the pipeline: + +``resolve_session_id`` — for the API boundary (Debug console, published-agent +invoke), where the caller submits ``session_id`` verbatim. The submitted value +is namespaced under the *authenticated* caller (an HMAC that folds in the +caller identity), so caller A's ``"foo"`` and caller B's ``"foo"`` can never +collide onto one microVM, and a raw id lifted from someone else's response or +ledger entry re-namespaces under whoever replays it instead of landing on the +original session. The mapping is idempotent: the resolved id is echoed back to +the client (``InvokeResponse.runtime_session_id``) and the Debug console +resends it as ``session_id`` for conversation continuity, so feeding a +resolved id back in must return it unchanged. + +``derive_channel_session_id`` — for the channel webhook path, where the id is +built server-side from a ``conversation_id``. The digest is keyed with the +platform binding secret so it is not offline-computable from the public +``channel_id`` (which is embedded in the webhook URL) plus a guessed +``conversation_id``. + +The binding secret is ``PLATFORM_SESSION_BINDING_SECRET`` when set. The +fallback folds in deployment identifiers that do not appear in any public +webhook URL (the runtime ARNs, the table name, the API token), so an external +caller who knows only the channel URL still cannot predict a session id. +Setting the secret explicitly is strongly recommended for production; the +value only has to be stable across replicas and restarts, not rotated. +""" + +import hashlib +import hmac +import secrets + +from app.config import settings + +# runtimeSessionId charset is [A-Za-z0-9_-] and AgentCore requires >= 33 chars. +# Both forms below are 47/48 chars: prefix + 12-char caller tag + 32-char body. +_SID_PREFIX = "s" + + +def _secret() -> bytes: + explicit = getattr(settings, "session_binding_secret", "") or "" + if explicit: + return explicit.encode("utf-8") + seed = "|".join( + str(x) + for x in ( + settings.sdk_runtime_arn, + settings.interactive_runtime_arn, + settings.dynamo_table, + settings.aws_region, + settings.api_token, + ) + ) + return ("session-binding:" + seed).encode("utf-8") + + +def _hmac(label: bytes, message: bytes) -> str: + return hmac.new(_secret(), label + b"\x00" + message, hashlib.sha256).hexdigest() + + +def _caller_tag(user: str) -> str: + """A stable, per-caller prefix segment. Identifies whose namespace a + resolved id lives in without revealing the caller name.""" + return _hmac(b"utag", str(user).encode("utf-8"))[:12] + + +def resolve_session_id(user, requested: str | None) -> str: + """Map a caller-submitted ``session_id`` onto a caller-bound runtime + session id. Never returns another caller's id. + + - empty request -> a fresh id already inside this caller's namespace + (so the very first turn's echoed id round-trips cleanly); + - a value this caller already owns -> returned unchanged (continuity); + - anything else (a raw client id, or an id minted for another caller) + -> namespaced under this caller, which can never collide onto another + tenant's microVM. + """ + tag = _caller_tag(str(user)) + if not requested: + return f"{_SID_PREFIX}-{tag}-{secrets.token_hex(16)}" + parts = requested.split("-", 2) + if ( + len(parts) == 3 + and parts[0] == _SID_PREFIX + and hmac.compare_digest(parts[1], tag) + ): + # one of this caller's own ids, echoed back for continuity + return requested + return f"{_SID_PREFIX}-{tag}-{_hmac(b'sid', str(user).encode('utf-8') + b'|' + requested.encode('utf-8'))[:32]}" + + +def derive_channel_session_id(channel_id: str, conversation_id: str) -> str: + """Stable, non-guessable runtime session id for a channel conversation. + + Keyed with the binding secret so it cannot be reproduced offline from the + public ``channel_id`` plus a guessed ``conversation_id``.""" + digest = _hmac(b"chn", f"{channel_id}\x00{conversation_id}".encode("utf-8")) + return f"chn-{digest[:44]}" # >= 33 chars for AgentCore diff --git a/scripts/check_session_binding_authz.py b/scripts/check_session_binding_authz.py new file mode 100755 index 0000000..828ce50 --- /dev/null +++ b/scripts/check_session_binding_authz.py @@ -0,0 +1,139 @@ +#!/usr/bin/env python3 +"""Authorization checks for the session-binding boundary. + +An AgentCore ``runtimeSessionId`` selects which warm microVM/process (and thus +which ``/tmp`` files, secret cache and model-gateway grant) a call lands on, so +a caller-controllable session id is an authorization boundary, not a label. +This asserts that: + + * a caller-submitted session_id is namespaced under the *authenticated* + caller, so it can never resolve onto another tenant's microVM; + * the mapping is idempotent, so the id echoed back to the client (and resent + for conversation continuity) round-trips unchanged; + * channel session ids are keyed with the binding secret, so they cannot be + predicted offline from the public channel_id; + * the routes actually apply the mapping — a correct function is worthless if + a call site passes the request value through verbatim (the bug this guards + against regressing). + +No third-party dependencies and no AWS calls: ``app.config`` is stubbed so the +module loads without pydantic/boto3. + +Run: python3 scripts/check_session_binding_authz.py +""" + +import importlib.util +import os +import re +import sys +import types + +ROOT = os.path.join(os.path.dirname(__file__), "..", "backend", "app") + + +def _load_session_binding(secret: str = "test-secret"): + """Load session_binding.py with a stubbed app.config.settings.""" + cfg = types.ModuleType("app.config") + cfg.settings = types.SimpleNamespace( + session_binding_secret=secret, + sdk_runtime_arn="arn:sdk", + interactive_runtime_arn="arn:int", + dynamo_table="agent-platform", + aws_region="us-east-1", + api_token="", + ) + app_pkg = types.ModuleType("app") + sys.modules["app"] = app_pkg + sys.modules["app.config"] = cfg + path = os.path.join(ROOT, "services", "session_binding.py") + spec = importlib.util.spec_from_file_location("app.services.session_binding", path) + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + return mod + + +FAILURES: list[str] = [] + + +def check(label: str, ok: bool, detail: str = "") -> None: + if not ok: + FAILURES.append(f"{label}{': ' + detail if detail else ''}") + print(f" {'ok ' if ok else 'FAIL'} {label}") + + +sb = _load_session_binding() +rs = sb.resolve_session_id +VALID = re.compile(r"^[A-Za-z0-9_-]+$") + +print("resolve_session_id (API boundary):") +alice_conv = rs("alice", "conv-1") +bob_conv = rs("bob", "conv-1") +# 1. the case that matters: same submitted id, two callers, never the same VM +check("two callers submitting the same id never collide", alice_conv != bob_conv) +# 2. an id lifted from alice's response re-namespaces under whoever replays it +check( + "replaying another caller's id does not land on it", + rs("bob", alice_conv) != alice_conv, +) +check( + "...and alice's own id is unchanged by bob replaying it", + rs("alice", alice_conv) == alice_conv, +) +# 3. idempotent: the resolved id is echoed back and resent for continuity +check("idempotent for a raw id", rs("alice", rs("alice", "conv-1")) == alice_conv) +fresh = rs("alice", "") +fresh2 = rs("alice", "") +check("empty request mints a fresh id", fresh != fresh2) +check("...that round-trips unchanged", rs("alice", fresh) == fresh) +# 4. AgentCore shape: charset + >= 33 chars +for label, sid in [("raw", alice_conv), ("fresh", fresh)]: + check(f"{label} id >= 33 chars", len(sid) >= 33, f"len={len(sid)}") + check(f"{label} id charset", bool(VALID.match(sid)), sid) + +print("derive_channel_session_id (channel path):") +dc = sb.derive_channel_session_id +sid_a = dc("chan_marketing_bot", "thread-8842") +check("deterministic", dc("chan_marketing_bot", "thread-8842") == sid_a) +check("distinct conversations differ", dc("chan_marketing_bot", "thread-1") != sid_a) +check("distinct channels differ", dc("other", "thread-8842") != sid_a) +check("charset + length", bool(VALID.match(sid_a)) and len(sid_a) >= 33, sid_a) +# the whole point: not the old, offline-computable unkeyed sha256 +import hashlib +old = "chn-" + hashlib.sha256(b"chan_marketing_bot:thread-8842").hexdigest()[:44] +check("not the unkeyed sha256 an attacker can precompute", sid_a != old) +# secret actually keys it +sb2 = _load_session_binding(secret="different-secret") +check( + "changing the binding secret changes the id", + sb2.derive_channel_session_id("chan_marketing_bot", "thread-8842") != sid_a, +) + +print("route wiring (regression guard):") +def _src(rel: str) -> str: + with open(os.path.join(ROOT, rel), encoding="utf-8") as fh: + return fh.read() + +kernels = _src(os.path.join("api", "kernels.py")) +agents = _src(os.path.join("api", "agents.py")) +channel = _src(os.path.join("services", "channel_service.py")) +check( + "kernels route resolves the session id", + "resolve_session_id(" in kernels and "runtime_session_id=req.session_id" not in kernels, +) +check( + "agents route resolves the session id", + "resolve_session_id(" in agents and "runtime_session_id=req.session_id" not in agents, +) +check( + "channel path derives a keyed session id", + "derive_channel_session_id(" in channel + and "sha256(f\"{channel_id}:{conversation_id}\"" not in channel, +) + +print() +if FAILURES: + print(f"FAILED ({len(FAILURES)}):") + for f in FAILURES: + print(" -", f) + sys.exit(1) +print("all session-binding authorization checks passed") diff --git a/scripts/e2e_session_isolation.py b/scripts/e2e_session_isolation.py new file mode 100755 index 0000000..a9b9d6a --- /dev/null +++ b/scripts/e2e_session_isolation.py @@ -0,0 +1,179 @@ +#!/usr/bin/env python3 +"""End-to-end test for the cross-tenant session-hijack / grant-DoS fix (H6). + +Runs the *real* backend routes and services (invocation → kernel → llm +credentials) with AWS mocked (moto) and the AgentCore invoke replaced by a +recording fake kernel, so the only stubbed things are authentication (we need +two distinct tenants; open mode has one identity) and the runtime that is not +deployed here. + +Asserts, over the public invoke API and the credential service: + + [attack A] two tenants submitting the *same* session_id land on different + runtime sessions (microVMs); replaying a tenant's echoed id from + another tenant does not land on it; a tenant resending its own id + keeps continuity (same runtime session). + [attack B] a tenant cannot overwrite another tenant's live gateway grant + (the DoS): mint refuses, the victim's stored digest is untouched, + and the owner can still re-mint its own. +""" +import io +import json +import os +import sys + +# ---- env must be set before app.config instantiates Settings() ---- +os.environ.update( + PLATFORM_AWS_REGION="us-east-1", + AWS_DEFAULT_REGION="us-east-1", + AWS_ACCESS_KEY_ID="testing", + AWS_SECRET_ACCESS_KEY="testing", + PLATFORM_DYNAMO_TABLE="agent-platform", + PLATFORM_SDK_RUNTIME_ARN="local_sdk_runtime", + PLATFORM_LLM_EDGE_URL="http://llm-edge.local", + PLATFORM_SESSION_BINDING_SECRET="e2e-binding-secret", +) + +from moto import mock_aws + +FAIL: list[str] = [] + + +def check(label: str, ok: bool, detail: str = "") -> None: + print(f" {'ok ' if ok else 'FAIL'} {label}" + (f" ({detail})" if detail else "")) + if not ok: + FAIL.append(label) + + +class _Body: + def __init__(self, data: bytes): + self._d = data + + def read(self) -> bytes: + return self._d + + +class FakeKernelRuntime: + """Stand-in for the bedrock-agentcore client: records the runtimeSessionId + each invocation is routed to and returns a minimal valid kernel reply.""" + + def __init__(self): + self.seen: list[str] = [] + + def invoke_agent_runtime(self, **kwargs): + self.seen.append(kwargs["runtimeSessionId"]) + body = json.dumps({"ok": True, "result": "ok", "usage": {}}).encode() + return {"response": _Body(body)} + + +def main() -> int: + with mock_aws(): + import boto3 + + ddb = boto3.client("dynamodb", region_name="us-east-1") + ddb.create_table( + TableName="agent-platform", + KeySchema=[ + {"AttributeName": "PK", "KeyType": "HASH"}, + {"AttributeName": "SK", "KeyType": "RANGE"}, + ], + AttributeDefinitions=[ + {"AttributeName": "PK", "AttributeType": "S"}, + {"AttributeName": "SK", "AttributeType": "S"}, + ], + BillingMode="PAY_PER_REQUEST", + ) + + from fastapi import Header + from fastapi.testclient import TestClient + + from app.dependencies import Principal, get_current_user + from app.main import app + from app.services import kernel_service as ks_mod + + fake = FakeKernelRuntime() + ks_mod.kernel_service.agentcore = fake + + def fake_user(x_test_user: str = Header(default="demo")) -> Principal: + p = Principal(x_test_user) + p.is_admin = False + return p + + app.dependency_overrides[get_current_user] = fake_user + client = TestClient(app) + + def invoke(user: str, session_id): + body = {"prompt": "hi"} + if session_id is not None: + body["session_id"] = session_id + r = client.post( + "/api/v1/kernels/agent-sdk/invoke", + json=body, + headers={"X-Test-User": user}, + ) + assert r.status_code == 200, (r.status_code, r.text) + return r.json()["runtime_session_id"] + + print("[attack A] cross-tenant session isolation + continuity") + # both tenants submit the *same* client-chosen session_id + a1 = invoke("alice", "shared-convo") + sid_alice = fake.seen[-1] + b1 = invoke("bob", "shared-convo") + sid_bob = fake.seen[-1] + check("echoed id equals the routed runtime session", a1 == sid_alice and b1 == sid_bob) + check("same submitted id, two tenants -> different microVMs", sid_alice != sid_bob, + f"{sid_alice} vs {sid_bob}") + + # alice replays bob's echoed id — must NOT land on bob's session + invoke("alice", b1) + check("replaying another tenant's echoed id does not land on it", + fake.seen[-1] != sid_bob, fake.seen[-1]) + + # each tenant resending its own echoed id keeps continuity + invoke("alice", a1) + check("alice resending her own id keeps the same microVM", fake.seen[-1] == sid_alice) + invoke("bob", b1) + check("bob resending his own id keeps the same microVM", fake.seen[-1] == sid_bob) + + # first-turn (no session_id) still round-trips for the same caller + f1 = invoke("carol", None) + invoke("carol", f1) + check("first-turn generated id round-trips", fake.seen[-1] == f1) + check("...and is caller-bound (bob cannot reuse carol's)", + invoke("bob", f1) != f1) + + print("[attack B] gateway grant cannot be overwritten across tenants") + from app.services.llm_credentials_service import llm_credentials_service as llm + spec = { + "backend": "gateway", + "base_url": "https://gw.local/v1", + "secret_name": "agent-platform/llm-gateway-key", + "model": "claude-fable-5-1", + } + grant_sid = "s-victimtag-deadbeef" # a concrete live session + c_alice = llm.mint(grant_sid, "alice", spec) + check("victim (alice) mints a grant", bool(c_alice)) + item0 = llm.table.get_item(Key={"PK": "LLMTOKEN", "SK": f"RSID#{grant_sid}"})["Item"] + dig0, owner0 = item0["token_sha256"], item0["user"] + + c_bob = llm.mint(grant_sid, "bob", spec) # the overwrite attempt + check("attacker (bob) mint on victim's session is refused", c_bob is None) + item1 = llm.table.get_item(Key={"PK": "LLMTOKEN", "SK": f"RSID#{grant_sid}"})["Item"] + check("victim's grant digest is untouched (no DoS)", item1["token_sha256"] == dig0) + check("victim's grant still owned by victim", item1["user"] == owner0 == "alice") + + c_alice2 = llm.mint(grant_sid, "alice", spec) # owner re-mint (refresh) + item2 = llm.table.get_item(Key={"PK": "LLMTOKEN", "SK": f"RSID#{grant_sid}"})["Item"] + check("owner can still re-mint its own grant", bool(c_alice2)) + check("...and that rotates the digest", item2["token_sha256"] != dig0) + + print() + if FAIL: + print(f"FAILED ({len(FAIL)}): " + ", ".join(FAIL)) + return 1 + print("all session-isolation e2e checks passed") + return 0 + + +if __name__ == "__main__": + sys.exit(main())