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
7 changes: 6 additions & 1 deletion backend/app/api/agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
8 changes: 7 additions & 1 deletion backend/app/api/kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
10 changes: 10 additions & 0 deletions backend/app/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
9 changes: 6 additions & 3 deletions backend/app/services/channel_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@
context.
"""

import hashlib
import hmac
import logging
import re
Expand All @@ -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__)

Expand Down Expand Up @@ -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.
Expand Down
4 changes: 4 additions & 0 deletions backend/app/services/invocation_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
49 changes: 33 additions & 16 deletions backend/app/services/llm_credentials_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
108 changes: 108 additions & 0 deletions backend/app/services/session_binding.py
Original file line number Diff line number Diff line change
@@ -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
139 changes: 139 additions & 0 deletions scripts/check_session_binding_authz.py
Original file line number Diff line number Diff line change
@@ -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")
Loading