mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
a677242d6f
commit
cde9d94c36
6 changed files with 325 additions and 27 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue