Merge remote-tracking branch 'origin/main' into litellm_ryan-march-31

This commit is contained in:
Ryan Crabbe 2026-04-04 10:07:59 -07:00
commit ad81892f05
No known key found for this signature in database
24 changed files with 910 additions and 28 deletions

View file

@ -42,6 +42,6 @@ jobs:
retention-days: 5
- name: Upload to code scanning
uses: github/codeql-action/upload-sarif@c10b806170c8ee63ea24152429041b5624f0baf5 # v4.35.1
uses: github/codeql-action/upload-sarif@c10b8064de6f491fea524254123dbe5e09572f13 # v4.35.1
with:
sarif_file: results.sarif

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-proxy-extras"
version = "0.4.63"
version = "0.4.64"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.4.63"
version = "0.4.64"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-proxy-extras==",

View file

@ -48,20 +48,19 @@ class A2ACompletionBridgeHandler:
# Get provider config for custom_llm_provider
custom_llm_provider = litellm_params.get("custom_llm_provider")
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider=custom_llm_provider
custom_llm_provider=custom_llm_provider,
model=litellm_params.get("model"),
)
# If provider config exists, use it
if a2a_provider_config is not None:
if api_base is None:
raise ValueError(f"api_base is required for {custom_llm_provider}")
verbose_logger.info(f"A2A: Using provider config for {custom_llm_provider}")
response_data = await a2a_provider_config.handle_non_streaming(
request_id=request_id,
params=params,
api_base=api_base,
litellm_params=litellm_params,
)
return response_data
@ -147,14 +146,12 @@ class A2ACompletionBridgeHandler:
# Get provider config for custom_llm_provider
custom_llm_provider = litellm_params.get("custom_llm_provider")
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider=custom_llm_provider
custom_llm_provider=custom_llm_provider,
model=litellm_params.get("model"),
)
# If provider config exists, use it
if a2a_provider_config is not None:
if api_base is None:
raise ValueError(f"api_base is required for {custom_llm_provider}")
verbose_logger.info(
f"A2A: Using provider config for {custom_llm_provider} (streaming)"
)
@ -163,6 +160,7 @@ class A2ACompletionBridgeHandler:
request_id=request_id,
params=params,
api_base=api_base,
litellm_params=litellm_params,
):
yield chunk

View file

@ -3,7 +3,7 @@ Base configuration for A2A protocol providers.
"""
from abc import ABC, abstractmethod
from typing import Any, AsyncIterator, Dict
from typing import Any, AsyncIterator, Dict, Optional
class BaseA2AProviderConfig(ABC):
@ -19,7 +19,7 @@ class BaseA2AProviderConfig(ABC):
self,
request_id: str,
params: Dict[str, Any],
api_base: str,
api_base: Optional[str] = None,
**kwargs,
) -> Dict[str, Any]:
"""
@ -41,7 +41,7 @@ class BaseA2AProviderConfig(ABC):
self,
request_id: str,
params: Dict[str, Any],
api_base: str,
api_base: Optional[str] = None,
**kwargs,
) -> AsyncIterator[Dict[str, Any]]:
"""

View file

@ -0,0 +1,22 @@
"""
Bedrock AgentCore A2A provider.
Preserves JSON-RPC envelopes for AgentCore agents that speak A2A natively,
bypassing the completion bridge that would otherwise strip the envelope.
"""
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
from litellm.a2a_protocol.providers.bedrock_agentcore.handler import (
BedrockAgentCoreA2AHandler,
)
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
__all__ = [
"BedrockAgentCoreA2AConfig",
"BedrockAgentCoreA2AHandler",
"BedrockAgentCoreA2ATransformation",
]

View file

@ -0,0 +1,61 @@
"""
Bedrock AgentCore A2A provider configuration.
"""
from typing import Any, AsyncIterator, Dict, Optional
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
from litellm.a2a_protocol.providers.bedrock_agentcore.handler import (
BedrockAgentCoreA2AHandler,
)
class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
"""
Provider configuration for Bedrock AgentCore A2A-native agents.
AgentCore agents that speak A2A natively expect the full JSON-RPC envelope.
This config bypasses the completion bridge and forwards requests directly,
deriving the endpoint URL from the model ARN and signing with SigV4/JWT.
"""
async def handle_non_streaming(
self,
request_id: str,
params: Dict[str, Any],
api_base: Optional[str] = None,
**kwargs,
) -> Dict[str, Any]:
"""Handle non-streaming request to AgentCore A2A agent."""
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for BedrockAgentCoreA2AConfig "
"(must contain model with AgentCore ARN)"
)
return await BedrockAgentCoreA2AHandler.handle_non_streaming(
request_id=request_id,
params=params,
litellm_params=litellm_params,
)
async def handle_streaming(
self,
request_id: str,
params: Dict[str, Any],
api_base: Optional[str] = None,
**kwargs,
) -> AsyncIterator[Dict[str, Any]]:
"""Handle streaming request to AgentCore A2A agent."""
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for BedrockAgentCoreA2AConfig "
"(must contain model with AgentCore ARN)"
)
async for chunk in BedrockAgentCoreA2AHandler.handle_streaming(
request_id=request_id,
params=params,
litellm_params=litellm_params,
):
yield chunk

View file

@ -0,0 +1,134 @@
"""
Handler for Bedrock AgentCore A2A-native agents.
Sends JSON-RPC envelopes directly to AgentCore endpoints, bypassing the
completion bridge that would otherwise strip the envelope.
"""
import json
from typing import Any, AsyncIterator, Dict, cast
from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
class BedrockAgentCoreA2AHandler:
"""
Handler for Bedrock AgentCore A2A requests.
Constructs JSON-RPC envelopes, signs them via AmazonAgentCoreConfig,
and POSTs directly to the AgentCore endpoint.
"""
@staticmethod
async def handle_non_streaming(
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
) -> Dict[str, Any]:
"""
Handle non-streaming A2A request to AgentCore.
Args:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (model, api_key, etc.)
Returns:
A2A JSON-RPC response dict from the AgentCore agent
"""
url, headers, body = (
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id=request_id,
params=params,
litellm_params=litellm_params,
method="message/send",
)
)
verbose_logger.info(
f"BedrockAgentCore A2A: Sending non-streaming request to {url}"
)
client = get_async_httpx_client(
llm_provider=cast(Any, httpxSpecialProvider.A2AProvider),
)
response = await client.post(
url,
headers=headers,
data=body,
)
response.raise_for_status()
response_data = response.json()
if "error" in response_data:
verbose_logger.warning(
f"BedrockAgentCore A2A: Agent returned error: {response_data['error']}"
)
return response_data
@staticmethod
async def handle_streaming(
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
) -> AsyncIterator[Dict[str, Any]]:
"""
Handle streaming A2A request to AgentCore.
Args:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (model, api_key, etc.)
Yields:
A2A streaming response events from the AgentCore agent
"""
url, headers, body = (
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id=request_id,
params=params,
litellm_params=litellm_params,
method="message/send",
stream=True,
)
)
verbose_logger.info(
f"BedrockAgentCore A2A: Sending streaming request to {url}"
)
client = get_async_httpx_client(
llm_provider=cast(Any, httpxSpecialProvider.A2AProvider),
)
response = await client.post(
url,
headers=headers,
data=body,
stream=True,
)
response.raise_for_status()
# Check content type — AgentCore may return JSON instead of SSE
content_type = response.headers.get("content-type", "").lower()
if "application/json" in content_type:
# Single JSON response fallback (not SSE)
verbose_logger.debug(
"BedrockAgentCore A2A streaming: received JSON instead of SSE, "
"yielding as single event"
)
response_body = await response.aread()
response_data = json.loads(response_body)
yield response_data
else:
# SSE stream — parse data: lines
async for event in BedrockAgentCoreA2ATransformation.parse_sse_events(
response
):
yield event

View file

@ -0,0 +1,134 @@
"""
Transformation layer for Bedrock AgentCore A2A provider.
Constructs JSON-RPC envelopes, derives AgentCore URLs from model ARNs,
and signs requests via AmazonAgentCoreConfig (SigV4 or JWT).
"""
import json
from typing import Any, AsyncIterator, Dict, Tuple
from litellm._logging import verbose_logger
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
class BedrockAgentCoreA2ATransformation:
"""
Request/response transformation for Bedrock AgentCore A2A agents.
Reuses AmazonAgentCoreConfig for URL construction, ARN parsing,
and request signing. No logic is duplicated.
"""
@staticmethod
def get_url_and_signed_request(
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
method: str = "message/send",
stream: bool = False,
) -> Tuple[str, dict, bytes]:
"""
Build the AgentCore URL, construct a JSON-RPC envelope, and sign the request.
Args:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams
litellm_params: Agent's litellm_params (model, api_key, etc.)
method: JSON-RPC method name (default: "message/send")
stream: Whether this is a streaming request
Returns:
Tuple of (url, signed_headers, signed_body_bytes)
"""
# Extract model and strip the "bedrock/" prefix
# "bedrock/agentcore/arn:aws:..." → "agentcore/arn:aws:..."
model = litellm_params.get("model", "")
if model.startswith("bedrock/"):
agentcore_model = model[len("bedrock/") :]
else:
agentcore_model = model
# Build optional_params from litellm_params (everything except model and custom_llm_provider)
optional_params = {
k: v
for k, v in litellm_params.items()
if k not in ("model", "custom_llm_provider")
}
agentcore_config = AmazonAgentCoreConfig()
# Derive URL from ARN
url = agentcore_config.get_complete_url(
api_base=optional_params.get("api_base"),
api_key=optional_params.get("api_key"),
model=agentcore_model,
optional_params=optional_params,
litellm_params=litellm_params,
stream=stream,
)
# Construct JSON-RPC 2.0 envelope
json_rpc_body = {
"jsonrpc": "2.0",
"method": method,
"id": request_id,
"params": params,
}
# Set required AgentCore session headers (normally set by transform_request,
# which we skip because it also builds {"prompt": "..."})
headers: dict = {}
session_id = agentcore_config._get_runtime_session_id(optional_params)
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id
runtime_user_id = agentcore_config._get_runtime_user_id(optional_params)
if runtime_user_id:
headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id
# Sign the request (SigV4 or JWT depending on api_key presence)
signed_headers, signed_body = agentcore_config.sign_request(
headers=headers,
optional_params=optional_params,
request_data=json_rpc_body,
api_base=url,
api_key=optional_params.get("api_key"),
model=agentcore_model,
stream=stream,
)
# sign_request returns Optional[bytes] — ensure we have bytes
if signed_body is None:
signed_body = json.dumps(json_rpc_body).encode()
return url, signed_headers, signed_body
@staticmethod
async def parse_sse_events(response: Any) -> AsyncIterator[Dict[str, Any]]:
"""
Parse SSE events from an httpx streaming response.
Reads line-by-line, parses `data:` lines as JSON, and yields each parsed dict.
Args:
response: httpx streaming response
Yields:
Parsed JSON dicts from SSE data lines
"""
async for line in response.aiter_lines():
line = line.strip()
if not line:
continue
if line.startswith("data:"):
data_str = line[len("data:") :].strip()
if not data_str:
continue
try:
event = json.loads(data_str)
yield event
except json.JSONDecodeError:
verbose_logger.debug(
f"BedrockAgentCore A2A: Skipping non-JSON SSE line: {data_str[:100]}"
)
continue

View file

@ -19,12 +19,14 @@ class A2AProviderConfigManager:
@staticmethod
def get_provider_config(
custom_llm_provider: Optional[str],
model: Optional[str] = None,
) -> Optional[BaseA2AProviderConfig]:
"""
Get the provider configuration for a given custom_llm_provider.
Args:
custom_llm_provider: The provider identifier (e.g., "pydantic_ai_agents")
model: The model string (used to distinguish sub-providers, e.g. agentcore vs other bedrock)
Returns:
Provider configuration instance or None if not found
@ -39,9 +41,11 @@ class A2AProviderConfigManager:
return PydanticAIProviderConfig()
# Add more providers here as needed
# elif custom_llm_provider == "another_provider":
# from litellm.a2a_protocol.providers.another_provider.config import AnotherProviderConfig
# return AnotherProviderConfig()
if custom_llm_provider == "bedrock" and model and "agentcore" in model:
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
return BedrockAgentCoreA2AConfig()
return None

View file

@ -2753,6 +2753,8 @@ class NewProjectRequest(LiteLLM_BudgetTable):
budget_id: Optional[str] = None
metadata: Optional[dict] = None
tags: Optional[List[str]] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
models: List[str] = []
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None
@ -2785,6 +2787,8 @@ class UpdateProjectRequest(LiteLLM_BudgetTable):
team_id: Optional[str] = None
metadata: Optional[dict] = None
tags: Optional[List[str]] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
models: Optional[List[str]] = None
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None

View file

@ -1472,17 +1472,19 @@ def _add_guardrails_from_key_or_team_metadata(
team_metadata: Optional[dict],
data: dict,
metadata_variable_name: str,
project_metadata: Optional[dict] = None,
) -> None:
"""
Helper add guardrails from key or team metadata to request data
Helper add guardrails from key, team, or project metadata to request data
Key guardrails are set first, then team guardrails are appended (without duplicates).
Key guardrails are set first, then team and project guardrails are appended (without duplicates).
Args:
key_metadata: The key metadata dictionary to check for guardrails
team_metadata: The team metadata dictionary to check for guardrails
data: The request data to update
metadata_variable_name: The name of the metadata field in data
project_metadata: The project metadata dictionary to check for guardrails
"""
from litellm.proxy.utils import _premium_user_check
@ -1508,6 +1510,15 @@ def _add_guardrails_from_key_or_team_metadata(
_premium_user_check()
combined_guardrails.update(team_metadata["guardrails"])
# Add project-level guardrails (set automatically handles duplicates)
if project_metadata and "guardrails" in project_metadata:
if (
isinstance(project_metadata["guardrails"], list)
and len(project_metadata["guardrails"]) > 0
):
_premium_user_check()
combined_guardrails.update(project_metadata["guardrails"])
# Set combined guardrails in metadata as list
if combined_guardrails:
data[metadata_variable_name]["guardrails"] = list(combined_guardrails)
@ -1518,12 +1529,13 @@ def _add_guardrails_from_policies_in_metadata(
team_metadata: Optional[dict],
data: dict,
metadata_variable_name: str,
project_metadata: Optional[dict] = None,
) -> None:
"""
Helper to resolve guardrails from policies attached to key/team metadata.
Helper to resolve guardrails from policies attached to key/team/project metadata.
This function:
1. Gets policy names from key and team metadata
1. Gets policy names from key, team, and project metadata
2. Resolves guardrails from those policies (including inheritance)
3. Adds resolved guardrails to request metadata
@ -1532,6 +1544,7 @@ def _add_guardrails_from_policies_in_metadata(
team_metadata: The team metadata dictionary to check for policies
data: The request data to update
metadata_variable_name: The name of the metadata field in data
project_metadata: The project metadata dictionary to check for policies
"""
from litellm._logging import verbose_proxy_logger
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
@ -1560,6 +1573,15 @@ def _add_guardrails_from_policies_in_metadata(
_premium_user_check()
policy_names.update(team_metadata["policies"])
# Add project-level policies
if project_metadata and "policies" in project_metadata:
if (
isinstance(project_metadata["policies"], list)
and len(project_metadata["policies"]) > 0
):
_premium_user_check()
policy_names.update(project_metadata["policies"])
if not policy_names:
return
@ -1641,6 +1663,7 @@ async def move_guardrails_to_metadata(
# Early-out: skip all guardrails processing when nothing is configured
key_metadata = user_api_key_dict.metadata
team_metadata = user_api_key_dict.team_metadata
project_metadata = user_api_key_dict.project_metadata or {}
has_key_config = key_metadata and (
"guardrails" in key_metadata or "policies" in key_metadata
@ -1648,12 +1671,15 @@ async def move_guardrails_to_metadata(
has_team_config = team_metadata and (
"guardrails" in team_metadata or "policies" in team_metadata
)
has_project_config = project_metadata and (
"guardrails" in project_metadata or "policies" in project_metadata
)
has_request_config = (
"guardrails" in data or "guardrail_config" in data or "policies" in data
)
# Only check policy engine if no local config (avoid import + registry lookup)
if not (has_key_config or has_team_config or has_request_config):
if not (has_key_config or has_team_config or has_project_config or has_request_config):
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
if not get_policy_registry().is_initialized():
@ -1661,20 +1687,22 @@ async def move_guardrails_to_metadata(
data.pop("policies", None)
return
# Check key-level guardrails
# Check key/team/project-level guardrails
_add_guardrails_from_key_or_team_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=project_metadata,
data=data,
metadata_variable_name=_metadata_variable_name,
)
#########################################################################################
# Add guardrails from policies attached to key/team metadata
# Add guardrails from policies attached to key/team/project metadata
#########################################################################################
_add_guardrails_from_policies_in_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=project_metadata,
data=data,
metadata_variable_name=_metadata_variable_name,
)

View file

@ -63,7 +63,7 @@ boto3 = { version = "1.42.80", optional = true }
redisvl = {version = "0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
mcp = {version = "1.26.0", optional = true, python = ">=3.10"}
a2a-sdk = {version = "0.3.25", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "0.4.63", optional = true}
litellm-proxy-extras = {version = "0.4.64", optional = true}
rich = {version = "13.9.4", optional = true}
litellm-enterprise = {version = "0.1.35", optional = true}
diskcache = {version = "5.6.3", optional = true}

View file

@ -56,7 +56,7 @@ grpcio==1.80.0
sentry_sdk==2.21.0 # for sentry error handling
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
tzdata==2025.1 # IANA time zone database
litellm-proxy-extras==0.4.63 # for proxy extras - e.g. prisma migrations
litellm-proxy-extras==0.4.64 # for proxy extras - e.g. prisma migrations
llm-sandbox==0.3.31 # for skill execution in sandbox
### LITELLM PACKAGE DEPENDENCIES
python-dotenv==1.0.1 # for env

View file

@ -0,0 +1,327 @@
"""
Tests for Bedrock AgentCore A2A provider.
Verifies that:
- JSON-RPC envelopes are preserved (not stripped by the completion bridge)
- URLs are derived from the model ARN
- Auth uses JWT Bearer or SigV4
- Config manager routes "bedrock" correctly
- Handler passes litellm_params and allows api_base=None
"""
import json
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
SAMPLE_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789:runtime/my_agent"
SAMPLE_MODEL = f"bedrock/agentcore/{SAMPLE_ARN}"
SAMPLE_PARAMS = {
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "what is 1+1?"}],
"messageId": "msg-001",
}
}
SAMPLE_LITELLM_PARAMS = {
"model": SAMPLE_MODEL,
"custom_llm_provider": "bedrock",
"api_key": "test-jwt-token",
}
class TestTransformation:
"""Test URL construction and JSON-RPC envelope building."""
def test_json_rpc_envelope_structure(self):
"""Verify JSON-RPC body has jsonrpc, method, id, and params."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
url, headers, body = (
BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
method="message/send",
)
)
body_dict = json.loads(body)
assert body_dict["jsonrpc"] == "2.0"
assert body_dict["method"] == "message/send"
assert body_dict["id"] == "req-001"
assert body_dict["params"] == SAMPLE_PARAMS
def test_url_derived_from_arn(self):
"""Verify URL is constructed from the ARN, not from api_base."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
url, _, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
)
assert "bedrock-agentcore.us-west-2.amazonaws.com" in url
assert "/runtimes/" in url
assert "/invocations" in url
def test_jwt_auth_uses_bearer_header(self):
"""When api_key is set, Authorization header uses Bearer token."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
)
assert headers["Authorization"] == "Bearer test-jwt-token"
def test_session_id_header_set(self):
"""Verify X-Amzn-Bedrock-AgentCore-Runtime-Session-Id is set."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
)
session_id = headers.get("X-Amzn-Bedrock-AgentCore-Runtime-Session-Id", "")
assert len(session_id) >= 33
def test_custom_session_id_header(self):
"""Verify custom runtimeSessionId is used when provided."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
params_with_session = {**SAMPLE_LITELLM_PARAMS, "runtimeSessionId": "a" * 40}
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=params_with_session,
)
assert headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] == "a" * 40
def test_sigv4_auth_when_no_api_key(self):
"""When no api_key, falls through to SigV4 signing."""
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
litellm_params_no_key = {
"model": SAMPLE_MODEL,
"custom_llm_provider": "bedrock",
"aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
"aws_region_name": "us-west-2",
}
# Mock _sign_request to avoid hitting real botocore credential resolution
fake_sigv4_headers = {
"Authorization": "AWS4-HMAC-SHA256 Credential=AKIA.../bedrock-agentcore/aws4_request",
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
}
fake_body = b'{"jsonrpc":"2.0"}'
with patch(
"litellm.llms.bedrock.chat.agentcore.transformation.AmazonAgentCoreConfig._sign_request",
return_value=(fake_sigv4_headers, fake_body),
):
_, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=litellm_params_no_key,
)
# SigV4 produces an Authorization header starting with "AWS4-HMAC-SHA256"
assert "Authorization" in headers
assert headers["Authorization"].startswith("AWS4-HMAC-SHA256")
class TestNonStreaming:
"""Test end-to-end non-streaming flow."""
@pytest.mark.asyncio
async def test_json_rpc_body_sent_to_agentcore(self):
"""Verify the full JSON-RPC envelope is POSTed, not {"prompt": "..."}."""
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
mock_response = MagicMock()
mock_response.json.return_value = {
"jsonrpc": "2.0",
"id": "req-001",
"result": {
"message": {
"role": "agent",
"parts": [{"kind": "text", "text": "2"}],
"messageId": "resp-001",
}
},
}
mock_response.raise_for_status = MagicMock()
with patch(
"litellm.a2a_protocol.providers.bedrock_agentcore.handler.get_async_httpx_client"
) as mock_get_client:
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
config = BedrockAgentCoreA2AConfig()
result = await config.handle_non_streaming(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
)
# Verify the POST was called
mock_client.post.assert_called_once()
call_kwargs = mock_client.post.call_args
# Verify sent body is JSON-RPC, not {"prompt": "..."}
sent_body = json.loads(call_kwargs.kwargs["data"])
assert "jsonrpc" in sent_body
assert "method" in sent_body
assert sent_body["method"] == "message/send"
assert sent_body["params"]["message"]["parts"][0]["text"] == "what is 1+1?"
# Verify response is passed through
assert result["result"]["message"]["parts"][0]["text"] == "2"
@pytest.mark.asyncio
async def test_a2a_error_response_passthrough(self):
"""JSON-RPC error responses from the agent are returned as-is."""
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
error_response = {
"jsonrpc": "2.0",
"id": "req-001",
"error": {"code": -32600, "message": "Bad request"},
}
mock_response = MagicMock()
mock_response.json.return_value = error_response
mock_response.raise_for_status = MagicMock()
with patch(
"litellm.a2a_protocol.providers.bedrock_agentcore.handler.get_async_httpx_client"
) as mock_get_client:
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
config = BedrockAgentCoreA2AConfig()
result = await config.handle_non_streaming(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
)
assert result["error"]["code"] == -32600
assert result["error"]["message"] == "Bad request"
class TestConfigManager:
"""Test that config manager routes 'bedrock' correctly."""
def test_bedrock_returns_config(self):
from litellm.a2a_protocol.providers.bedrock_agentcore.config import (
BedrockAgentCoreA2AConfig,
)
from litellm.a2a_protocol.providers.config_manager import (
A2AProviderConfigManager,
)
config = A2AProviderConfigManager.get_provider_config(
"bedrock", model=SAMPLE_MODEL
)
assert config is not None
assert isinstance(config, BedrockAgentCoreA2AConfig)
def test_bedrock_non_agentcore_returns_none(self):
"""Non-agentcore bedrock models should fall through to completion bridge."""
from litellm.a2a_protocol.providers.config_manager import (
A2AProviderConfigManager,
)
config = A2AProviderConfigManager.get_provider_config(
"bedrock", model="bedrock/anthropic.claude-3-sonnet"
)
assert config is None
def test_unknown_provider_returns_none(self):
from litellm.a2a_protocol.providers.config_manager import (
A2AProviderConfigManager,
)
assert A2AProviderConfigManager.get_provider_config("unknown") is None
class TestHandlerIntegration:
"""Test handler.py changes — litellm_params passed through, api_base not required."""
@pytest.mark.asyncio
async def test_provider_config_receives_litellm_params(self):
"""Verify handler passes litellm_params to provider config via kwargs."""
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
mock_config = AsyncMock()
mock_config.handle_non_streaming = AsyncMock(
return_value={"jsonrpc": "2.0", "id": "req-001", "result": {}}
)
with patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config",
return_value=mock_config,
):
await A2ACompletionBridgeHandler.handle_non_streaming(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
api_base=None,
)
mock_config.handle_non_streaming.assert_called_once_with(
request_id="req-001",
params=SAMPLE_PARAMS,
api_base=None,
litellm_params=SAMPLE_LITELLM_PARAMS,
)
@pytest.mark.asyncio
async def test_api_base_none_allowed_with_provider_config(self):
"""api_base=None no longer raises when a provider config is registered."""
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
mock_config = AsyncMock()
mock_config.handle_non_streaming = AsyncMock(
return_value={"jsonrpc": "2.0", "id": "req-001", "result": {}}
)
with patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config",
return_value=mock_config,
):
# Should NOT raise ValueError
result = await A2ACompletionBridgeHandler.handle_non_streaming(
request_id="req-001",
params=SAMPLE_PARAMS,
litellm_params=SAMPLE_LITELLM_PARAMS,
api_base=None,
)
assert result is not None

View file

@ -1363,6 +1363,101 @@ async def test_request_guardrails_do_not_override_key_guardrails():
assert len(requested_guardrails) == 1
@pytest.mark.asyncio
async def test_project_guardrails_merge_with_key_and_team():
"""
Test that project guardrails are merged with key and team guardrails (union semantics).
All three levels should contribute to the final guardrails list without duplicates.
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={"guardrails": ["key-guardrail-1"]},
team_metadata={"guardrails": ["team-guardrail-1", "key-guardrail-1"]},
project_metadata={"guardrails": ["project-guardrail-1", "team-guardrail-1"]},
)
with patch("litellm.proxy.utils._premium_user_check"):
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
guardrails = metadata.get("guardrails", [])
# All three sources contribute
assert "key-guardrail-1" in guardrails
assert "team-guardrail-1" in guardrails
assert "project-guardrail-1" in guardrails
# No duplicates
assert guardrails.count("key-guardrail-1") == 1
assert guardrails.count("team-guardrail-1") == 1
@pytest.mark.asyncio
async def test_project_guardrails_only():
"""
Test that project guardrails work when key and team have no guardrails configured.
"""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={},
team_metadata={},
project_metadata={"guardrails": ["project-guardrail-1", "project-guardrail-2"]},
)
with patch("litellm.proxy.utils._premium_user_check"):
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
guardrails = metadata.get("guardrails", [])
assert "project-guardrail-1" in guardrails
assert "project-guardrail-2" in guardrails
assert len(guardrails) == 2
def test_update_model_if_key_alias_exists():
"""
Test that _update_model_if_key_alias_exists properly updates the model when a key alias exists.

View file

@ -17,6 +17,7 @@ export interface ProjectCreateParams {
models?: string[];
max_budget?: number;
blocked?: boolean;
guardrails?: string[];
metadata?: Record<string, unknown>;
model_rpm_limit?: Record<string, number>;
model_tpm_limit?: Record<string, number>;

View file

@ -17,6 +17,7 @@ export interface ProjectUpdateParams {
models?: string[];
max_budget?: number;
blocked?: boolean;
guardrails?: string[];
metadata?: Record<string, unknown>;
model_rpm_limit?: Record<string, number>;
model_tpm_limit?: Record<string, number>;

View file

@ -33,6 +33,9 @@ export function EditProjectModal({
const metadataObj = (project.metadata ?? {}) as Record<string, unknown>;
const rpmLimits = (metadataObj.model_rpm_limit ?? {}) as Record<string, number>;
const tpmLimits = (metadataObj.model_tpm_limit ?? {}) as Record<string, number>;
const guardrails = (Array.isArray(metadataObj.guardrails)
? metadataObj.guardrails
: []) as string[];
const modelLimits: ProjectFormValues["modelLimits"] = [];
const allLimitModels = new Set([
@ -48,7 +51,7 @@ export function EditProjectModal({
}
// Filter out internal keys from user-facing metadata
const internalKeys = new Set(["model_rpm_limit", "model_tpm_limit"]);
const internalKeys = new Set(["model_rpm_limit", "model_tpm_limit", "guardrails"]);
const metadata: ProjectFormValues["metadata"] = [];
for (const [key, value] of Object.entries(metadataObj)) {
if (!internalKeys.has(key)) {
@ -63,6 +66,7 @@ export function EditProjectModal({
models: project.models ?? [],
max_budget: project.litellm_budget_table?.max_budget ?? undefined,
isBlocked: project.blocked,
guardrails: guardrails.length > 0 ? guardrails : undefined,
modelLimits: modelLimits.length > 0 ? modelLimits : undefined,
metadata: metadata.length > 0 ? metadata : undefined,
});

View file

@ -14,6 +14,10 @@ vi.mock("@/components/organisms/create_key_button", () => ({
fetchTeamModels: vi.fn().mockResolvedValue([]),
}));
vi.mock("@/components/networking", () => ({
getGuardrailsList: vi.fn().mockResolvedValue({ guardrails: [] }),
}));
vi.mock("@/components/key_team_helpers/fetch_available_models_team_key", () => ({
getModelDisplayName: (model: string) => model,
}));
@ -86,4 +90,13 @@ describe("ProjectBaseForm", () => {
renderWithProviders(<FormWrapper />);
expect(screen.getByText("Advanced Settings")).toBeInTheDocument();
});
it("should show a Guardrails field in the Advanced Settings section", async () => {
const user = userEvent.setup();
renderWithProviders(<FormWrapper />);
await user.click(screen.getByText("Advanced Settings"));
await waitFor(() => {
expect(screen.getByText("Guardrails")).toBeInTheDocument();
});
});
});

View file

@ -22,6 +22,7 @@ import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
import { Team } from "../../key_team_helpers/key_list";
import { fetchTeamModels } from "../../organisms/create_key_button";
import { getModelDisplayName } from "../../key_team_helpers/fetch_available_models_team_key";
import { getGuardrailsList } from "@/components/networking";
export interface ProjectFormValues {
project_alias: string;
@ -30,6 +31,7 @@ export interface ProjectFormValues {
models: string[];
max_budget?: number;
isBlocked: boolean;
guardrails?: string[];
modelLimits?: { model: string; tpm?: number; rpm?: number }[];
metadata?: { key: string; value: string }[];
}
@ -46,6 +48,23 @@ export function ProjectBaseForm({
const [selectedTeam, setSelectedTeam] = useState<Team | null>(null);
const [modelsToPick, setModelsToPick] = useState<string[]>([]);
const [guardrailsList, setGuardrailsList] = useState<string[]>([]);
useEffect(() => {
const fetchGuardrails = async () => {
if (!accessToken) return;
try {
const response = await getGuardrailsList(accessToken);
const names = response.guardrails.map(
(g: { guardrail_name: string }) => g.guardrail_name
);
setGuardrailsList(names);
} catch (error) {
console.error("Failed to fetch guardrails:", error);
}
};
fetchGuardrails();
}, [accessToken]);
// Sync selectedTeam from form value (needed for edit mode pre-fill)
const teamIdValue = Form.useWatch("team_id", form);
@ -259,6 +278,24 @@ export function ProjectBaseForm({
<Divider />
<Form.Item
label="Guardrails"
name="guardrails"
help="Select existing guardrails or enter new ones"
>
<Select
mode="tags"
style={{ width: "100%" }}
placeholder="Select or enter guardrails"
options={guardrailsList.map((name) => ({
value: name,
label: name,
}))}
/>
</Form.Item>
<Divider />
<Typography.Text
strong
style={{ display: "block", marginBottom: 12 }}

View file

@ -98,4 +98,20 @@ describe("buildProjectApiParams", () => {
});
expect(result).not.toHaveProperty("metadata");
});
it("should include guardrails as a top-level field when provided", () => {
const result = buildProjectApiParams({
...baseValues,
guardrails: ["pii-check", "content-filter"],
});
expect(result.guardrails).toEqual(["pii-check", "content-filter"]);
});
it("should omit guardrails when the array is empty", () => {
const result = buildProjectApiParams({
...baseValues,
guardrails: [],
});
expect(result).not.toHaveProperty("guardrails");
});
});

View file

@ -25,6 +25,9 @@ export function buildProjectApiParams(values: ProjectFormValues) {
models: values.models ?? [],
max_budget: values.max_budget,
blocked: values.isBlocked ?? false,
...(values.guardrails && values.guardrails.length > 0 && {
guardrails: values.guardrails,
}),
...(Object.keys(modelRpmLimit).length > 0 && {
model_rpm_limit: modelRpmLimit,
}),