feat(agentcore-a2a): derive runtime session id from A2A message.contextId (#39371)

Native AgentCore A2A always sent either a fresh generated runtime session id
or the single configured runtimeSessionId, so related turns lost context and
unrelated callers shared one AgentCore microVM. The runtime session id is now
params.message.contextId scoped to the calling key hash, then runtimeSessionId,
then generated, and is length-validated (33-256) before the header is signed.
Invalid ids surface as JSON-RPC -32602 / HTTP 400 instead of a 500.

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-02 12:40:24 -07:00 • committed by GitHub
parent a677242d6f
commit cde9d94c36
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 325 additions and 27 deletions

View file

@ -10,8 +10,19 @@ from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
from litellm._logging import verbose_logger
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
from litellm.a2a_protocol.utils import (
get_session_id_from_a2a_params,
scope_session_to_principal,
)
from litellm.exceptions import BadRequestError
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
RUNTIME_SESSION_ID_MIN_LENGTH: Final = 33
RUNTIME_SESSION_ID_MAX_LENGTH: Final = 256
# Reserved outbound header names that must never be sourced from per-request
# ``agent_extra_headers`` for AgentCore requests. ``agent_extra_headers`` carries
# values rewritten from the client-controlled ``x-a2a-{agent}-*`` convention, so
@ -19,8 +30,9 @@ from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreCo
# request identity / SigV4 metadata by overwriting headers the proxy sets from
# trusted server-side config.
#
# The runtime headers (session / user id) are derived server-side from
# ``runtimeSessionId`` / ``runtimeUserId`` in the agent's ``litellm_params``;
# The runtime headers (session / user id) are derived server-side from the A2A
# ``message.contextId`` and ``runtimeSessionId`` / ``runtimeUserId`` in the
# agent's ``litellm_params``;
# ``authorization`` is set by the AgentCore signer (JWT or SigV4); ``host`` and
# the ``x-amz-*`` family are owned by SigV4 itself.
_RESERVED_EXACT_HEADERS: Final = frozenset(
@ -66,6 +78,31 @@ def _filter_reserved_headers(
return filtered or None
def _request_scoped_runtime_session_id(
params: Mapping[str, Any],
litellm_params: Mapping[str, Any],
) -> str | None:
context_id: Final = get_session_id_from_a2a_params(params)
if not isinstance(context_id, str) or not context_id:
return None
return scope_session_to_principal(context_id, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM))
def _validate_runtime_session_id(session_id: str, model: str) -> str:
if RUNTIME_SESSION_ID_MIN_LENGTH <= len(session_id) <= RUNTIME_SESSION_ID_MAX_LENGTH:
return session_id
raise BadRequestError(
message=(
f"Invalid AgentCore runtime session id {session_id!r}: AWS requires "
f"{RUNTIME_SESSION_ID_MIN_LENGTH}-{RUNTIME_SESSION_ID_MAX_LENGTH} characters. It is built from the A2A "
"message.contextId (prefixed with a 16-hex-char hash of the calling key and '-') when set, "
"otherwise from the agent's configured runtimeSessionId."
),
model=model,
llm_provider="bedrock",
)
class BedrockAgentCoreA2ATransformation:
"""
Request/response transformation for Bedrock AgentCore A2A agents.
@ -100,7 +137,9 @@ class BedrockAgentCoreA2ATransformation:
here to prevent a caller-controlled ``x-a2a-{agent}-*`` header from
spoofing the AgentCore runtime user id or other SigV4 metadata. Use
``api_key`` / ``runtimeUserId`` / ``runtimeSessionId`` in litellm_params
(not ``agent_extra_headers``) to override those values.
(not ``agent_extra_headers``) to override those values. The runtime
session id is taken from ``params["message"]["contextId"]`` (scoped to
the calling key) when present, then ``runtimeSessionId``, else generated.
Returns:
Tuple of (url, signed_headers, signed_body_bytes)
@ -139,7 +178,11 @@ class BedrockAgentCoreA2ATransformation:
# Set required AgentCore session headers (normally set by transform_request,
# which we skip because it also builds {"prompt": "..."})
headers: Final[dict] = {}
session_id: Final = agentcore_config._get_runtime_session_id(optional_params)
session_id: Final = _validate_runtime_session_id(
_request_scoped_runtime_session_id(params, litellm_params)
or agentcore_config._get_runtime_session_id(optional_params),
model=model,
)
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id
runtime_user_id: Final = agentcore_config._get_runtime_user_id(optional_params)
if runtime_user_id:

View file

@ -2,6 +2,8 @@
Utility functions for A2A protocol.
"""
import hashlib
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import litellm
@ -140,6 +142,29 @@ class A2ARequestUtils:
return prompt_tokens, completion_tokens, total_tokens
def get_session_id_from_a2a_params(params: Mapping[str, Any]) -> str | None:
message: Final = params.get("message", {})
if isinstance(message, dict):
return message.get("contextId")
return getattr(message, "contextId", None)
def scope_session_to_principal(session_id: str, principal: str | None) -> str:
"""
Bind a client-supplied A2A contextId to the authenticated principal.
Without this, two distinct keys authorized for the same agent could set the
same contextId and read/append to each other's backend memory. The
principal is hashed (it is already a hashed token) so the raw value is never
sent to the agent backend, while the original contextId is kept as a suffix
for operator-side correlation.
"""
if not principal:
return session_id
principal_prefix: Final = hashlib.sha256(principal.encode("utf-8")).hexdigest()[:16]
return f"{principal_prefix}-{session_id}"
# Backwards compatibility aliases
def extract_text_from_a2a_message(message: Any) -> str:
return A2ARequestUtils.extract_text_from_message(message)

View file

@ -1,28 +1,9 @@
import hashlib
from typing import Any, Final
def get_session_id_from_a2a_params(params: dict[str, Any]) -> str | None:
message: Final = params.get("message", {})
if isinstance(message, dict):
return message.get("contextId")
return getattr(message, "contextId", None)
def scope_session_to_principal(session_id: str, principal: str | None) -> str:
"""
Bind a client-supplied A2A contextId to the authenticated principal.
Without this, two distinct keys authorized for the same LangFlow agent could
set the same contextId and read/append to each other's LangFlow memory. The
principal is hashed (it is already a hashed token) so the raw value is never
sent to the LangFlow backend, while the original contextId is kept as a
suffix for operator-side correlation.
"""
if not principal:
return session_id
principal_prefix: Final = hashlib.sha256(principal.encode("utf-8")).hexdigest()[:16]
return f"{principal_prefix}-{session_id}"
from litellm.a2a_protocol.utils import (
get_session_id_from_a2a_params,
scope_session_to_principal,
)
def merge_a2a_session_into_litellm_params(

View file

@ -1019,4 +1019,6 @@ async def invoke_agent_a2a(
)
except Exception:
pass
if isinstance(e, litellm.BadRequestError):
return _jsonrpc_error(body.get("id"), -32602, e.message, 400)
return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {e}", 500)

View file

@ -11,7 +11,9 @@ Verifies that:
import json
import httpx
import pytest
import respx
from unittest.mock import AsyncMock, MagicMock, patch
@ -295,6 +297,195 @@ class TestTransformation:
assert headers["Authorization"].startswith("AWS4-HMAC-SHA256")
SESSION_HEADER = "X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"
CONTEXT_ID = "conversation-alpha-0001-0000000000000000"
KEY_HASH = "hashed-key-of-caller-one"
def _params_with_context(context_id: object) -> dict:
return {"message": {**SAMPLE_PARAMS["message"], "contextId": context_id}}
def _scoped(context_id: str, key_hash: str) -> str:
import hashlib
return f"{hashlib.sha256(key_hash.encode()).hexdigest()[:16]}-{context_id}"
def _session_header(params: dict, litellm_params: dict) -> str:
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=params,
litellm_params=litellm_params,
)
return headers[SESSION_HEADER]
@pytest.fixture
def httpx_transport(monkeypatch):
import litellm
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
class TestRequestScopedRuntimeSession:
"""message.contextId selects the AgentCore runtime session, scoped to the calling key."""
def test_context_id_scoped_to_calling_key(self):
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
litellm_params = {**SAMPLE_LITELLM_PARAMS, A2A_USER_API_KEY_HASH_PARAM: KEY_HASH}
assert _session_header(_params_with_context(CONTEXT_ID), litellm_params) == _scoped(CONTEXT_ID, KEY_HASH)
def test_context_id_used_verbatim_without_principal(self):
assert _session_header(_params_with_context(CONTEXT_ID), SAMPLE_LITELLM_PARAMS) == CONTEXT_ID
def test_same_context_id_reuses_session_and_other_context_isolated(self):
first = _session_header(_params_with_context(CONTEXT_ID), SAMPLE_LITELLM_PARAMS)
second = _session_header(_params_with_context(CONTEXT_ID), SAMPLE_LITELLM_PARAMS)
other = _session_header(
_params_with_context("conversation-beta-00002-0000000000000000"),
SAMPLE_LITELLM_PARAMS,
)
assert first == second
assert other != first
def test_same_context_id_from_different_keys_is_isolated(self):
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
params = _params_with_context(CONTEXT_ID)
caller_one = _session_header(params, {**SAMPLE_LITELLM_PARAMS, A2A_USER_API_KEY_HASH_PARAM: KEY_HASH})
caller_two = _session_header(
params, {**SAMPLE_LITELLM_PARAMS, A2A_USER_API_KEY_HASH_PARAM: "hashed-key-of-caller-two"}
)
assert caller_one != caller_two
assert caller_one.endswith(f"-{CONTEXT_ID}")
assert caller_two.endswith(f"-{CONTEXT_ID}")
def test_context_id_takes_precedence_over_configured_session(self):
litellm_params = {**SAMPLE_LITELLM_PARAMS, "runtimeSessionId": "a" * 40}
assert _session_header(_params_with_context(CONTEXT_ID), litellm_params) == CONTEXT_ID
def test_configured_session_is_fallback_without_context_id(self):
litellm_params = {**SAMPLE_LITELLM_PARAMS, "runtimeSessionId": "a" * 40}
assert _session_header(SAMPLE_PARAMS, litellm_params) == "a" * 40
assert _session_header(_params_with_context(""), litellm_params) == "a" * 40
def test_no_context_id_and_no_config_generates_new_session_per_request(self):
first = _session_header(SAMPLE_PARAMS, SAMPLE_LITELLM_PARAMS)
second = _session_header(SAMPLE_PARAMS, SAMPLE_LITELLM_PARAMS)
assert first != second
assert 33 <= len(first) <= 256
@pytest.mark.parametrize(
"context_id",
[
"short-context-id",
"x" * 257,
],
)
def test_invalid_context_id_rejected_with_clear_error(self, context_id):
import litellm
with pytest.raises(litellm.BadRequestError, match="Invalid AgentCore runtime session id") as exc_info:
_session_header(_params_with_context(context_id), SAMPLE_LITELLM_PARAMS)
assert exc_info.value.status_code == 400
assert "33-256" in str(exc_info.value)
def test_scoped_context_id_shorter_than_33_rejected(self):
import litellm
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
litellm_params = {**SAMPLE_LITELLM_PARAMS, A2A_USER_API_KEY_HASH_PARAM: KEY_HASH}
with pytest.raises(litellm.BadRequestError, match=_scoped("c" * 15, KEY_HASH)):
_session_header(_params_with_context("c" * 15), litellm_params)
assert _session_header(_params_with_context("c" * 16), litellm_params) == _scoped("c" * 16, KEY_HASH)
def test_invalid_configured_session_rejected(self):
import litellm
litellm_params = {**SAMPLE_LITELLM_PARAMS, "runtimeSessionId": "too-short"}
with pytest.raises(litellm.BadRequestError, match="Invalid AgentCore runtime session id"):
_session_header(SAMPLE_PARAMS, litellm_params)
def test_non_string_context_id_falls_back(self):
litellm_params = {**SAMPLE_LITELLM_PARAMS, "runtimeSessionId": "a" * 40}
assert _session_header(_params_with_context(12345), litellm_params) == "a" * 40
def test_spoofed_session_header_does_not_override_context_id(self):
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=_params_with_context(CONTEXT_ID),
litellm_params=SAMPLE_LITELLM_PARAMS,
agent_extra_headers={SESSION_HEADER: "s" * 40},
)
assert headers[SESSION_HEADER] == CONTEXT_ID
@pytest.mark.asyncio
async def test_context_id_session_header_on_outbound_non_streaming_post(self, httpx_transport):
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
with respx.mock(assert_all_called=True) as router:
route = router.post(url__regex=r".*/invocations.*").mock(
return_value=httpx.Response(200, json={"jsonrpc": "2.0", "id": "req-001", "result": {}})
)
await BedrockAgentCoreA2AConfig().handle_non_streaming(
request_id="req-001",
params=_params_with_context(CONTEXT_ID),
litellm_params={**SAMPLE_LITELLM_PARAMS, A2A_USER_API_KEY_HASH_PARAM: KEY_HASH},
)
assert route.calls.last.request.headers[SESSION_HEADER] == _scoped(CONTEXT_ID, KEY_HASH)
@pytest.mark.asyncio
async def test_context_id_session_header_on_outbound_streaming_post(self, httpx_transport):
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
with respx.mock(assert_all_called=True) as router:
route = router.post(url__regex=r".*/invocations.*").mock(
return_value=httpx.Response(200, json={"jsonrpc": "2.0", "id": "req-001", "result": {}})
)
events = [
event
async for event in BedrockAgentCoreA2AConfig().handle_streaming(
request_id="req-001",
params=_params_with_context(CONTEXT_ID),
litellm_params={**SAMPLE_LITELLM_PARAMS, A2A_USER_API_KEY_HASH_PARAM: KEY_HASH},
)
]
assert events == [{"jsonrpc": "2.0", "id": "req-001", "result": {}}]
assert route.calls.last.request.headers[SESSION_HEADER] == _scoped(CONTEXT_ID, KEY_HASH)
class TestNonStreaming:
"""Test end-to-end non-streaming flow."""

View file

@ -918,6 +918,62 @@ async def test_task_method_failure_hook_uses_enriched_request_data():
assert failure_data.get("agent_id") == "test-agent"
@pytest.mark.asyncio
async def test_agentcore_invalid_context_id_returns_jsonrpc_invalid_params_400():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
agent.litellm_params = {
"custom_llm_provider": "bedrock",
"model": "bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/demo",
"api_key": "test-jwt-token",
}
mock_request = _make_request_mock(
"message/send",
{
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "Hello"}],
"messageId": "msg-1",
"contextId": "too-short",
}
},
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch( # test-quality-ok: same proxy_logging_obj injection the sibling failure-hook test uses; no HTTP call is made because the request is rejected before signing
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert response.status_code == 400
assert body["id"] == "req-1"
assert body["error"]["code"] == -32602
assert "Invalid AgentCore runtime session id" in body["error"]["message"]
assert "Internal error" not in body["error"]["message"]
mock_proxy_logging.post_call_failure_hook.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_extended_agent_card_rewrites_url():
from litellm.proxy._types import UserAPIKeyAuth