Merge pull request #41511 from BerriAI/litellm_foundry_a2a_entra_agents

feat(a2a): reach Microsoft Foundry agents with Entra auth and versioned card discovery
This commit is contained in:
Mateo Wang 2026-09-18 11:44:14 -07:00 • committed by GitHub
commit 88c9dd1294
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
20 changed files with 1309 additions and 293 deletions

View file

@ -6,9 +6,10 @@ Extends the A2A SDK's card resolver to support multiple well-known paths.
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol, runtime_checkable
from litellm._logging import verbose_logger
from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError
from litellm.constants import LOCALHOST_URL_PATTERNS
if TYPE_CHECKING:
@ -18,6 +19,8 @@ if TYPE_CHECKING:
_A2ACardResolver: Any = None
AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent-card.json"
PREV_AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent.json"
FOUNDRY_AGENT_CARD_PATH: Final = "/agentCard/v1.0"
AGENT_CARD_PATH_PARAM: Final = "agent_card_path"
try:
from a2a.client import A2ACardResolver as _A2ACardResolver
@ -29,6 +32,20 @@ except ImportError:
pass
@runtime_checkable
class _HasStatusCode(Protocol):
status_code: int | None
def _discovery_status_code(failures: tuple[tuple[str, Exception], ...]) -> int:
statuses: Final = tuple(
error.status_code
for _, error in failures
if isinstance(error, _HasStatusCode) and error.status_code is not None and error.status_code != 404
)
return statuses[0] if statuses else 404
def is_localhost_or_internal_url(url: str | None) -> bool:
"""
Check if a URL is a localhost or internal URL.
@ -145,9 +162,10 @@ class LiteLLMA2ACardResolver(_A2ACardResolver):
"""
Custom A2A card resolver that supports multiple well-known paths.
Extends the base A2ACardResolver to try both:
Extends the base A2ACardResolver to try, in order:
- /.well-known/agent-card.json (standard)
- /.well-known/agent.json (previous/alternative)
- /agentCard/v1.0
"""
async def get_agent_card(
@ -155,51 +173,37 @@ class LiteLLMA2ACardResolver(_A2ACardResolver):
relative_card_path: str | None = None,
http_kwargs: Mapping[str, object] | None = None,
) -> "AgentCard":
"""
Fetch the agent card, trying multiple well-known paths.
First tries the standard path, then falls back to the previous path.
Args:
relative_card_path: Optional path to the agent card endpoint.
If None, tries both well-known paths.
http_kwargs: Optional dictionary of keyword arguments to pass to httpx.get
Returns:
AgentCard from the A2A agent
Raises:
A2AClientHTTPError or A2AClientJSONError if both paths fail
"""
# If a specific path is provided, use the parent implementation
"""Fetch the agent card, probing every known path when none is given."""
if relative_card_path is not None:
return await super().get_agent_card(
relative_card_path=relative_card_path,
http_kwargs=http_kwargs,
)
# Try both well-known paths
paths: Final = [
AGENT_CARD_WELL_KNOWN_PATH,
PREV_AGENT_CARD_WELL_KNOWN_PATH,
]
return await self._get_agent_card_from_first_reachable_path(
paths=(AGENT_CARD_WELL_KNOWN_PATH, PREV_AGENT_CARD_WELL_KNOWN_PATH, FOUNDRY_AGENT_CARD_PATH),
http_kwargs=http_kwargs,
failures=(),
)
last_error = None
for path in paths:
try:
verbose_logger.debug("Attempting to fetch agent card from %s%s", self.base_url, path)
return await super().get_agent_card(
relative_card_path=path,
http_kwargs=http_kwargs,
)
except Exception as e:
verbose_logger.debug("Failed to fetch agent card from %s%s: %s", self.base_url, path, e)
last_error = e
continue
# If we get here, all paths failed - re-raise the last error
if last_error is not None:
raise last_error
# This shouldn't happen, but just in case
raise Exception(f"Failed to fetch agent card from {self.base_url}. Tried paths: {', '.join(paths)}")
async def _get_agent_card_from_first_reachable_path(
self,
paths: tuple[str, ...],
http_kwargs: Mapping[str, object] | None,
failures: tuple[tuple[str, Exception], ...],
) -> "AgentCard":
if not paths:
raise A2AAgentCardDiscoveryError(
base_url=self.base_url,
failures=failures,
status_code=_discovery_status_code(failures),
)
path: Final = paths[0]
try:
verbose_logger.debug("Attempting to fetch agent card from %s%s", self.base_url, path)
return await super().get_agent_card(relative_card_path=path, http_kwargs=http_kwargs)
except Exception as e:
verbose_logger.debug("Failed to fetch agent card from %s%s: %s", self.base_url, path, e)
return await self._get_agent_card_from_first_reachable_path(
paths=paths[1:], http_kwargs=http_kwargs, failures=(*failures, (path, e))
)

View file

@ -4,6 +4,8 @@ A2A Protocol Exceptions.
Custom exception types for A2A protocol operations, following LiteLLM's exception pattern.
"""
from typing import Final
import httpx
@ -100,11 +102,12 @@ class A2AAgentCardError(A2AError):
model: str | None = None,
response: httpx.Response | None = None,
litellm_debug_info: str | None = None,
status_code: int = 404,
):
self.url = url
super().__init__(
message=message,
status_code=404,
status_code=status_code,
llm_provider="a2a_agent",
model=model,
response=response,
@ -112,6 +115,17 @@ class A2AAgentCardError(A2AError):
)
class A2AAgentCardDiscoveryError(A2AAgentCardError):
def __init__(self, base_url: str, failures: tuple[tuple[str, Exception], ...], status_code: int) -> None:
self.failures = failures
attempts: Final = ", ".join(f"{path} ({error})" for path, error in failures)
super().__init__(
message=f"Failed to fetch agent card from {base_url}. Tried {attempts}",
url=base_url,
status_code=status_code,
)
class A2ALocalhostURLError(A2AConnectionError):
"""
Raised when an agent card contains a localhost/internal URL.

View file

@ -15,6 +15,7 @@ from typing import Any, Final
import litellm
from litellm._logging import verbose_logger
from litellm.a2a_protocol.card_resolver import AGENT_CARD_PATH_PARAM
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
A2ACompletionBridgeTransformation,
A2AStreamingContext,
@ -36,6 +37,7 @@ _AGENT_ONLY_PARAMS: Final = frozenset(
"agent_name",
"agent_id",
"agent_card_params",
AGENT_CARD_PATH_PARAM,
A2A_USER_API_KEY_HASH_PARAM,
}
)

View file

@ -13,7 +13,7 @@ import asyncio
import datetime
import uuid
from collections.abc import AsyncIterator, Coroutine, Mapping
from types import ModuleType
from types import MappingProxyType, ModuleType
from typing import TYPE_CHECKING, Any, Final, Optional, cast
import litellm
@ -72,6 +72,7 @@ except ImportError:
# Import our custom card resolver that supports multiple well-known paths
from litellm.a2a_protocol.card_resolver import (
AGENT_CARD_PATH_PARAM,
LiteLLMA2ACardResolver,
get_agent_card_url,
normalize_agent_card_interfaces,
@ -132,6 +133,26 @@ def _set_agent_id_on_logging_obj(
_A2A_COST_PARAM_KEYS: Final = ("cost_per_query", "input_cost_per_token", "output_cost_per_token")
def _a2a_cost_params(litellm_params: Mapping[str, object] | None) -> Mapping[str, object]:
"""Only the agent's pricing keys reach the logging object; its credentials never do."""
return MappingProxyType(
{
key: litellm_params[key]
for key in _A2A_COST_PARAM_KEYS
if litellm_params is not None and litellm_params.get(key) is not None
}
)
def _card_http_kwargs(extra_headers: dict[str, str] | None) -> dict[str, object] | None:
return {"headers": extra_headers} if extra_headers else None # mutable-ok: a2a-sdk's get_agent_card takes a dict
def _agent_card_path(litellm_params: Mapping[str, object]) -> str | None:
configured_path: Final = litellm_params.get(AGENT_CARD_PATH_PARAM)
return configured_path if isinstance(configured_path, str) and configured_path else None
def _set_litellm_params_on_logging_obj(
kwargs: Mapping[str, object],
litellm_params: Mapping[str, object],
@ -148,9 +169,7 @@ def _set_litellm_params_on_logging_obj(
if not isinstance(logging_obj, Logging):
return
cost_params: Final = {
key: litellm_params[key] for key in _A2A_COST_PARAM_KEYS if litellm_params.get(key) is not None
}
cost_params: Final = _a2a_cost_params(litellm_params)
if not cost_params:
return
@ -475,7 +494,11 @@ async def asend_message(
# Overlay agent-level headers (agent headers take precedence over LiteLLM internal ones)
if agent_extra_headers:
extra_headers.update(agent_extra_headers)
a2a_client = await create_a2a_client(base_url=api_base, extra_headers=extra_headers)
a2a_client = await create_a2a_client(
base_url=api_base,
extra_headers=extra_headers,
relative_card_path=_agent_card_path(litellm_params),
)
# Type assertion: a2a_client is guaranteed to be non-None here
assert a2a_client is not None
@ -588,11 +611,10 @@ def _build_streaming_logging_obj(
if agent_id:
logging_obj.model_call_details["agent_id"] = agent_id
_litellm_params: Final = litellm_params.copy() if litellm_params else {}
if metadata:
_litellm_params["metadata"] = metadata
if proxy_server_request:
_litellm_params["proxy_server_request"] = proxy_server_request
_request_context: Final = (("metadata", metadata), ("proxy_server_request", proxy_server_request))
_litellm_params: Final = dict( # mutable-ok: Logging.litellm_params is declared as a dict
(*_a2a_cost_params(litellm_params).items(), *((key, value) for key, value in _request_context if value))
)
logging_obj.litellm_params = _litellm_params
logging_obj.optional_params = _litellm_params
@ -700,6 +722,7 @@ async def asend_message_streaming(
base_url=api_base,
extra_headers=extra_headers,
streaming=True,
relative_card_path=_agent_card_path(litellm_params),
)
assert a2a_client is not None
@ -746,6 +769,7 @@ async def create_a2a_client(
timeout: float = DEFAULT_A2A_AGENT_TIMEOUT,
extra_headers: dict[str, str] | None = None,
streaming: bool = False,
relative_card_path: str | None = None,
) -> "A2AClientType":
"""
Create an A2A client for the given agent URL.
@ -757,6 +781,8 @@ async def create_a2a_client(
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
timeout: Request timeout in seconds (default: ``DEFAULT_A2A_AGENT_TIMEOUT`` / env ``DEFAULT_A2A_AGENT_TIMEOUT``)
extra_headers: Optional additional headers to include in requests
relative_card_path: Optional card path relative to ``base_url`` (e.g. ``agentCard/v1.0`` for a
Microsoft Foundry agent); when None the well-known paths are probed in order
Returns:
An initialized a2a.client.A2AClient instance
@ -790,7 +816,10 @@ async def create_a2a_client(
resolver: Final = A2ACardResolver(httpx_client=httpx_client, base_url=base_url)
agent_card: Final = normalize_agent_card_interfaces(
await resolver.get_agent_card(http_kwargs={"headers": extra_headers} if extra_headers else None)
await resolver.get_agent_card(
relative_card_path=relative_card_path,
http_kwargs=_card_http_kwargs(extra_headers),
)
)
a2a_client: Final = await create_client( # pyright: ignore[reportOptionalCall]
@ -820,6 +849,7 @@ async def aget_agent_card(
base_url: str,
timeout: float = DEFAULT_A2A_AGENT_TIMEOUT,
extra_headers: dict[str, str] | None = None,
relative_card_path: str | None = None,
) -> "AgentCard":
"""
Fetch the agent card from an A2A agent.
@ -828,6 +858,7 @@ async def aget_agent_card(
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
timeout: Request timeout in seconds (default: ``DEFAULT_A2A_AGENT_TIMEOUT`` / env ``DEFAULT_A2A_AGENT_TIMEOUT``)
extra_headers: Optional additional headers to include in requests
relative_card_path: Optional card path relative to ``base_url``; when None the well-known paths are probed
Returns:
AgentCard from the A2A agent
@ -850,7 +881,10 @@ async def aget_agent_card(
httpx_client=httpx_client,
base_url=base_url,
)
agent_card: Final = await resolver.get_agent_card()
agent_card: Final = await resolver.get_agent_card(
relative_card_path=relative_card_path,
http_kwargs=_card_http_kwargs(extra_headers),
)
verbose_logger.info("Fetched agent card: %s", agent_card.name if hasattr(agent_card, "name") else "unknown")
return agent_card

View file

@ -7,7 +7,7 @@ from typing import Final
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
from ..common_utils import extract_text_from_a2a_response
from ..common_utils import A2AError, extract_text_from_a2a_response
class A2AModelResponseIterator(BaseModelResponseIterator):
@ -56,6 +56,10 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
}
}
"""
error: Final = chunk.get("error")
if isinstance(error, dict):
raise A2AError(status_code=500, message=f"A2A error: {error.get('message', 'Unknown error')}")
try:
# Extract text from A2A response
text: Final = extract_text_from_a2a_response(chunk)

View file

@ -3,11 +3,12 @@ A2A Protocol Transformation for LiteLLM
"""
import uuid
from collections.abc import Iterator
from collections.abc import Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm.llms.azure_ai.common_utils import AZURE_ENTRA_LITELLM_PARAM_KEYS, get_azure_ai_agent_entra_token
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.openai import AllMessageValues
@ -15,6 +16,7 @@ from litellm.types.utils import Choices, Message, ModelResponse, Usage
from ..common_utils import (
A2AError,
a2a_hop_uses_entra,
convert_messages_to_prompt,
extract_text_from_a2a_response,
)
@ -26,6 +28,39 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
_REGISTRY_PARAMS_KEPT_OUT_OF_OPTIONAL_PARAMS: Final = (
frozenset({"api_key", "api_base", "headers", "model"}) | AZURE_ENTRA_LITELLM_PARAM_KEYS
)
def _card_declares_no_streaming(agent_card_params: Mapping[str, object]) -> bool:
capabilities: Final = agent_card_params.get("capabilities")
return isinstance(capabilities, Mapping) and not capabilities.get("streaming")
def _agent_authenticates_with_entra(agent_litellm_params: Mapping[str, object]) -> bool:
return a2a_hop_uses_entra(agent_litellm_params, agent_litellm_params.get("custom_llm_provider"))
def _registry_api_key(agent_litellm_params: Mapping[str, object]) -> str | None:
if _agent_authenticates_with_entra(agent_litellm_params):
return get_azure_ai_agent_entra_token(agent_litellm_params)
configured_api_key: Final = agent_litellm_params.get("api_key")
return configured_api_key if isinstance(configured_api_key, str) else None
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, Any] | None:
stored_headers: Final = agent_litellm_params.get("headers")
if not isinstance(stored_headers, Mapping):
return None
entra_owns_authorization: Final = _agent_authenticates_with_entra(agent_litellm_params)
return { # mutable-ok: completion() and httpx take the request headers as a dict
name: value
for name, value in stored_headers.items()
if not (entra_owns_authorization and str(name).lower() == "authorization")
}
class A2AConfig(BaseConfig):
"""
Configuration for A2A (Agent-to-Agent) Protocol.
@ -35,20 +70,19 @@ class A2AConfig(BaseConfig):
@staticmethod
def resolve_agent_config_from_registry(
model: str,
agent_name: str,
api_base: str | None,
api_key: str | None,
headers: dict[str, Any] | None,
optional_params: dict[str, Any],
) -> tuple[str | None, str | None, dict[str, Any] | None]:
"""
Resolve agent configuration from registry if model format is "a2a/<agent-name>".
Extracts agent name from model string and looks up configuration in the
agent registry (if available in proxy context).
Resolve agent configuration from the registry for a registered agent.
Args:
model: Model string (e.g., "a2a/my-agent")
agent_name: The model string with the provider prefix already stripped by
get_llm_provider ("a2a/my-agent" -> "my-agent"), the name the agent was
registered under
api_base: Explicit api_base (takes precedence over registry)
api_key: Explicit api_key (takes precedence over registry)
headers: Explicit headers (takes precedence over registry)
@ -57,11 +91,7 @@ class A2AConfig(BaseConfig):
Returns:
Tuple of (api_base, api_key, headers) with registry values filled in
"""
# Extract agent name from model (e.g., "a2a/my-agent" -> "my-agent")
agent_name: Final = model.split("/", 1)[1] if "/" in model else None
# Only lookup if agent name exists and some config is missing
if not agent_name or (api_base is not None and api_key is not None and headers is not None):
if not agent_name or (api_base is not None and api_key is not None and headers):
return api_base, api_key, headers
# Try registry lookup (only available in proxy context)
@ -79,17 +109,23 @@ class A2AConfig(BaseConfig):
# Get api_key, headers, and other params from litellm_params
if agent.litellm_params:
if api_key is None:
api_key = agent.litellm_params.get("api_key")
api_key = _registry_api_key(agent.litellm_params)
if headers is None:
agent_headers: Final = agent.litellm_params.get("headers")
if agent_headers:
headers = agent_headers
if not headers:
headers = _registry_headers(agent.litellm_params) or headers
# Merge other litellm_params (timeout, max_retries, etc.)
for key, value in agent.litellm_params.items():
if key not in ["api_key", "api_base", "headers", "model"] and key not in optional_params:
optional_params[key] = value
# Merge other litellm_params (timeout, max_retries, etc.)
registry_params: Final = tuple(
(key, value)
for key, value in (agent.litellm_params.items() if agent.litellm_params else ())
if key not in _REGISTRY_PARAMS_KEPT_OUT_OF_OPTIONAL_PARAMS and key not in optional_params
)
streaming_fallback: Final = (
(("stream", False), ("fake_stream", True))
if optional_params.get("stream") and _card_declares_no_streaming(agent.agent_card_params)
else ()
)
optional_params.update((*registry_params, *streaming_fallback))
except ImportError:
pass # Registry not available (not running in proxy context)
@ -147,17 +183,13 @@ class A2AConfig(BaseConfig):
api_base: API base URL
Returns:
Updated headers dict
A new headers dict; the caller's dict is left untouched
"""
# Ensure Content-Type is set to application/json for JSON-RPC 2.0
if "content-type" not in headers and "Content-Type" not in headers:
headers["Content-Type"] = "application/json"
# Add Authorization header if API key is provided
if api_key is not None:
headers["Authorization"] = f"Bearer {api_key}"
return headers
content_type_default: Final = (
() if "content-type" in headers or "Content-Type" in headers else (("Content-Type", "application/json"),)
)
bearer: Final = () if api_key is None else (("Authorization", f"Bearer {api_key}"),)
return dict((*headers.items(), *content_type_default, *bearer))
def get_complete_url(
self,
@ -226,6 +258,7 @@ class A2AConfig(BaseConfig):
# Create single A2A message with full conversation context
a2a_message: Final = {
"kind": "message",
"role": "user",
"parts": [{"kind": "text", "text": full_context}],
"messageId": str(uuid.uuid4()),
@ -237,11 +270,14 @@ class A2AConfig(BaseConfig):
stream: Final = optional_params.get("stream", False)
method: Final = "message/stream" if stream else "message/send"
params: Final = (
{"message": a2a_message} if stream else {"message": a2a_message, "configuration": {"blocking": True}}
)
request_data: Final = {
"jsonrpc": "2.0",
"id": request_id,
"method": method,
"params": {"message": a2a_message},
"params": params,
}
return request_data

View file

@ -2,7 +2,7 @@
Common utilities for A2A (Agent-to-Agent) Protocol
"""
from collections.abc import Mapping
from collections.abc import Awaitable, Callable, Mapping
from typing import Any, Final
from pydantic import BaseModel
@ -10,6 +10,7 @@ from pydantic import BaseModel
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
)
from litellm.llms.azure_ai.common_utils import has_azure_entra_params, resolve_azure_ai_agent_auth_header
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.openai import AllMessageValues
@ -142,3 +143,21 @@ def extract_text_from_a2a_response(response_dict: Mapping[str, object], max_dept
return extract_text_from_a2a_message(first_artifact, depth=0, max_depth=max_depth)
return ""
AgentAuthHeaderResolver = Callable[[Mapping[str, object]], Awaitable[Mapping[str, str]]]
def a2a_hop_uses_entra(litellm_params: Mapping[str, object], custom_llm_provider: object) -> bool:
return not custom_llm_provider and has_azure_entra_params(litellm_params)
async def resolve_a2a_hop_auth_header(
litellm_params: Mapping[str, object],
custom_llm_provider: object,
resolve_entra_header: AgentAuthHeaderResolver = resolve_azure_ai_agent_auth_header,
) -> Mapping[str, str] | None:
"""Entra credentials authenticate the A2A hop only; a completion-bridge agent hands them to the model provider it bridges to."""
if not a2a_hop_uses_entra(litellm_params, custom_llm_provider):
return None
return await resolve_entra_header(litellm_params)

View file

@ -1,4 +1,6 @@
import asyncio
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final, Literal
from urllib.parse import urlparse
@ -44,6 +46,70 @@ def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None)
return get_azure_ad_token(params)
AZURE_AI_AGENTS_SCOPE: Final = "https://ai.azure.com/.default"
AZURE_ENTRA_CREDENTIAL_PARAM_KEYS: Final = frozenset({"azure_ad_token", "client_secret", "azure_password"})
AZURE_ENTRA_LITELLM_PARAM_KEYS: Final = AZURE_ENTRA_CREDENTIAL_PARAM_KEYS | frozenset(
{"tenant_id", "client_id", "azure_username", "azure_scope"}
)
AZURE_ENTRA_CREDENTIAL_HELP: Final = (
"Set `tenant_id` + `client_id` + `client_secret`, `azure_ad_token` (an `oidc/` token also needs "
"`tenant_id` + `client_id`), or `client_id` + `azure_username` + `azure_password` in the agent's `litellm_params`"
)
def has_azure_entra_params(litellm_params: Mapping[str, object] | None) -> bool:
if not litellm_params:
return False
return any(litellm_params.get(key) for key in AZURE_ENTRA_CREDENTIAL_PARAM_KEYS)
def _resolve_config_secret(value: object) -> str | None:
if not isinstance(value, str) or not value:
return None
return get_secret_str(value) if value.startswith("os.environ/") else value
def get_azure_ai_agent_entra_token(litellm_params: Mapping[str, object]) -> str:
"""Mints the Entra bearer from the agent's own litellm_params, never from process-wide AZURE_* env vars."""
from litellm.llms.azure.common_utils import (
get_azure_ad_token_from_entra_id,
get_azure_ad_token_from_oidc,
get_azure_ad_token_from_username_password,
)
resolved: Final = MappingProxyType(
{key: _resolve_config_secret(litellm_params.get(key)) for key in AZURE_ENTRA_LITELLM_PARAM_KEYS}
)
scope: Final = resolved["azure_scope"] or AZURE_AI_AGENTS_SCOPE
tenant_id: Final = resolved["tenant_id"]
client_id: Final = resolved["client_id"]
client_secret: Final = resolved["client_secret"]
azure_username: Final = resolved["azure_username"]
azure_password: Final = resolved["azure_password"]
azure_ad_token: Final = resolved["azure_ad_token"]
if tenant_id and client_id and client_secret:
return get_azure_ad_token_from_entra_id(
tenant_id=tenant_id, client_id=client_id, client_secret=client_secret, scope=scope
)()
if client_id and azure_username and azure_password:
return get_azure_ad_token_from_username_password(
client_id=client_id, azure_username=azure_username, azure_password=azure_password, scope=scope
)()
federated: Final = azure_ad_token is not None and azure_ad_token.startswith("oidc/")
if azure_ad_token and federated and tenant_id and client_id:
return get_azure_ad_token_from_oidc(
azure_ad_token=azure_ad_token, azure_client_id=client_id, azure_tenant_id=tenant_id, scope=scope
)
if azure_ad_token and not federated:
return azure_ad_token
raise ValueError(f"Azure AI agent Entra ID credentials did not resolve to a token. {AZURE_ENTRA_CREDENTIAL_HELP}")
async def resolve_azure_ai_agent_auth_header(litellm_params: Mapping[str, object]) -> Mapping[str, str]:
token: Final = await asyncio.to_thread(get_azure_ai_agent_entra_token, litellm_params)
return MappingProxyType({"Authorization": f"Bearer {token}"})
def get_azure_ai_auth_headers(
api_key: str | None,
litellm_params: Mapping[str, object] | None = None,

View file

@ -2193,7 +2193,7 @@ def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
api_key,
headers,
) = litellm.A2AConfig.resolve_agent_config_from_registry(
model=model,
agent_name=model,
api_base=api_base,
api_key=api_key,
headers=headers,

View file

@ -24,6 +24,7 @@ from pydantic import ValidationError
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.llms.a2a.common_utils import resolve_a2a_hop_auth_header
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.a2a.version_convert import (
A2AVersion,
@ -157,19 +158,31 @@ def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str,
)
async def _resolve_backend_auth_header(
litellm_params: dict[str, object],
custom_llm_provider: object,
) -> Mapping[str, str] | None:
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
return await resolve_databricks_app_auth_header(litellm_params)
return await resolve_a2a_hop_auth_header(litellm_params, custom_llm_provider)
def _forwarding_headers(
caller_identity: Mapping[str, str],
request_data: Mapping[str, object],
agent_extra_headers: Mapping[str, str] | None,
backend_auth_header: Mapping[str, str] | None,
) -> dict[str, str] | None:
backend_auth: Final = tuple(backend_auth_header.items()) if backend_auth_header else ()
minted_names: Final = frozenset(name.lower() for name, _ in backend_auth)
passthrough: Final = tuple(
(name, value)
for name, value in (agent_extra_headers.items() if agent_extra_headers else ())
if not name.lower().startswith("x-litellm-")
if not name.lower().startswith("x-litellm-") and name.lower() not in minted_names
)
trace_id: Final = request_data.get("litellm_trace_id")
trace: Final = (("X-LiteLLM-Trace-Id", str(trace_id)),) if trace_id else ()
merged: Final = dict((*passthrough, *caller_identity.items(), *trace))
merged: Final = dict((*passthrough, *caller_identity.items(), *trace, *backend_auth))
return merged or None
@ -795,26 +808,16 @@ async def invoke_agent_a2a(
if header_name:
dynamic_headers[header_name] = val
agent_extra_headers = _forwarding_headers(
agent_extra_headers: Final = _forwarding_headers(
caller_identity=caller_identity,
request_data=data,
agent_extra_headers=merge_agent_headers(
dynamic_headers=dynamic_headers or None,
static_headers=static_headers or None,
),
backend_auth_header=await _resolve_backend_auth_header(litellm_params, custom_llm_provider),
)
# Databricks App endpoints require a short-lived OAuth M2M token rather
# than a static bearer. Only agents explicitly configured with a
# ``databricks_oauth`` block get one; every other agent is left untouched.
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
databricks_auth: Final = await resolve_databricks_app_auth_header(litellm_params)
if databricks_auth:
agent_extra_headers = {
**(agent_extra_headers or {}),
**databricks_auth,
}
# Merge agent-level guardrails into data so post_call_success_hook and
# _handle_stream_message both pick them up. A2A agents use model
# a2a_agent/*, which is not an llm_router deployment, so

View file

@ -57,7 +57,7 @@ def mock_a2a_client(monkeypatch):
import litellm.a2a_protocol.main as a2a_main
async def _fake_create_a2a_client(
base_url, timeout=60.0, extra_headers=None, streaming=False
base_url, timeout=60.0, extra_headers=None, streaming=False, relative_card_path=None
):
return MockA2AClient()

View file

@ -5,8 +5,10 @@ Tests that the card resolver tries both old and new well-known paths.
"""
from types import SimpleNamespace
from typing import Any, Final
from unittest.mock import MagicMock, patch
import httpx
import pytest
from litellm.a2a_protocol.card_resolver import (
@ -16,6 +18,7 @@ from litellm.a2a_protocol.card_resolver import (
normalize_agent_card_interfaces,
set_agent_card_url,
)
from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError
@pytest.mark.asyncio
@ -138,3 +141,109 @@ def test_normalize_agent_card_interfaces_downgrades_miscased_interfaces_to_the_0
]
assert card.supported_interfaces[0].protocol_binding == "jsonrpc"
assert card.supported_interfaces[0].protocol_version == "1.0"
_FOUNDRY_BASE_URL: Final = "https://foundry.example.com/a2a"
_FOUNDRY_CARD_JSON: Final = {
"name": "Foundry Agent",
"description": "A test agent",
"url": "https://foundry.example.com/a2a",
"version": "1.0",
"capabilities": {"streaming": True},
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"skills": [{"id": "chat", "name": "chat", "description": "Chat", "tags": ["chat"]}],
"protocolVersion": "1.0",
}
class _FakeHttpxClient:
"""Answers GETs from a path -> (status, body) map and records the path of each call."""
def __init__(self, base_url: str, responses: dict[str, tuple[int, dict[str, Any]]]) -> None:
self._base_url = base_url.rstrip("/")
self._responses = responses
self.calls: list[str] = []
async def get(self, url: str, **kwargs: Any) -> httpx.Response:
path: Final = url.removeprefix(self._base_url)
self.calls.append(path)
status_code, body = self._responses[path]
return httpx.Response(status_code, json=body, request=httpx.Request("GET", url))
@pytest.mark.asyncio
async def test_card_resolver_falls_through_to_the_foundry_card_path():
httpx_client = _FakeHttpxClient(
base_url=_FOUNDRY_BASE_URL,
responses={
"/.well-known/agent-card.json": (404, {"error": "not found"}),
"/.well-known/agent.json": (404, {"error": "not found"}),
"/agentCard/v1.0": (200, dict(_FOUNDRY_CARD_JSON)),
},
)
resolver = LiteLLMA2ACardResolver(httpx_client=httpx_client, base_url=_FOUNDRY_BASE_URL)
result = await resolver.get_agent_card()
assert httpx_client.calls == ["/.well-known/agent-card.json", "/.well-known/agent.json", "/agentCard/v1.0"]
assert result.name == "Foundry Agent"
assert result.supported_interfaces[0].url == "https://foundry.example.com/a2a"
@pytest.mark.asyncio
async def test_card_resolver_explicit_path_skips_the_probes():
httpx_client = _FakeHttpxClient(
base_url=_FOUNDRY_BASE_URL,
responses={"/agentCard/v1.0": (200, dict(_FOUNDRY_CARD_JSON))},
)
resolver = LiteLLMA2ACardResolver(httpx_client=httpx_client, base_url=_FOUNDRY_BASE_URL)
result = await resolver.get_agent_card(relative_card_path="agentCard/v1.0")
assert httpx_client.calls == ["/agentCard/v1.0"]
assert result.name == "Foundry Agent"
@pytest.mark.asyncio
async def test_card_resolver_names_every_probed_path_when_discovery_fails():
httpx_client = _FakeHttpxClient(
base_url=_FOUNDRY_BASE_URL,
responses={
"/.well-known/agent-card.json": (404, {"error": "not found"}),
"/.well-known/agent.json": (401, {"error": "unauthorized"}),
"/agentCard/v1.0": (404, {"error": "not found"}),
},
)
resolver = LiteLLMA2ACardResolver(httpx_client=httpx_client, base_url=_FOUNDRY_BASE_URL)
with pytest.raises(A2AAgentCardDiscoveryError) as raised:
await resolver.get_agent_card()
assert raised.value.status_code == 401
message = str(raised.value)
assert _FOUNDRY_BASE_URL in message
assert "/.well-known/agent-card.json (" in message and "HTTP 404" in message
assert "/.well-known/agent.json (" in message and "HTTP 401" in message
assert "/agentCard/v1.0 (" in message
@pytest.mark.asyncio
async def test_card_resolver_discovery_error_is_404_when_every_probe_is_404():
resolver = LiteLLMA2ACardResolver(
httpx_client=_FakeHttpxClient(
base_url=_FOUNDRY_BASE_URL,
responses={
"/.well-known/agent-card.json": (404, {"error": "not found"}),
"/.well-known/agent.json": (404, {"error": "not found"}),
"/agentCard/v1.0": (404, {"error": "not found"}),
},
),
base_url=_FOUNDRY_BASE_URL,
)
with pytest.raises(A2AAgentCardDiscoveryError) as raised:
await resolver.get_agent_card()
assert raised.value.status_code == 404

View file

@ -26,9 +26,7 @@ class TestA2AStreamingTransformation:
"parts": [{"text": "Reply to ticket #4823"}],
"metadata": {"skillId": "draft_reply"},
}
openai_messages = (
A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
)
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
# Metadata is forwarded on the run payload only, not duplicated on messages.
assert "metadata" not in openai_messages[0]
@ -174,10 +172,7 @@ class TestA2AStreamingTransformation:
assert "artifactId" in event["result"]["artifact"]
assert event["result"]["artifact"]["name"] == "response"
assert event["result"]["artifact"]["parts"][0]["kind"] == "text"
assert (
event["result"]["artifact"]["parts"][0]["text"]
== "Hello, I am an AI assistant."
)
assert event["result"]["artifact"]["parts"][0]["text"] == "Hello, I am an AI assistant."
@pytest.mark.asyncio
@ -332,3 +327,43 @@ async def test_handle_non_streaming_forwards_api_key():
assert call_kwargs["api_key"] == "my-secret-api-key"
assert call_kwargs["api_base"] == "https://my-azure.com/"
assert call_kwargs["model"] == "azure_ai/agents/asst_456"
@pytest.mark.asyncio
async def test_handle_streaming_keeps_agent_card_path_out_of_the_completion_call():
"""agent_card_path describes where an A2A agent serves its card; a completion-bridge agent carrying
it must not pass it to litellm.acompletion, where an unknown kwarg breaks the provider call."""
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
async def mock_streaming_response():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta = MagicMock()
chunk.choices[0].delta.content = "Hello"
yield chunk
with (
patch( # test-quality-ok: the bridge calls litellm.acompletion directly; the sibling tests capture its kwargs through the same seam
"litellm.acompletion", new_callable=AsyncMock
) as mock_acompletion
):
mock_acompletion.return_value = mock_streaming_response()
events = [
event
async for event in A2ACompletionBridgeHandler.handle_streaming(
request_id="req-card-path",
params={"message": {"role": "user", "parts": [{"kind": "text", "text": "Hi"}], "messageId": "m1"}},
litellm_params={
"custom_llm_provider": "langgraph",
"model": "agent",
"agent_card_path": "agentCard/v1.0",
},
api_base="http://localhost:2024",
)
]
assert len(events) == 4
assert "agent_card_path" not in mock_acompletion.call_args.kwargs

View file

@ -16,7 +16,13 @@ from a2a.compat.v0_3.types import (
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.a2a_protocol.main import _send_message, _stream_messages, asend_message, create_a2a_client
from litellm.a2a_protocol.main import (
_send_message,
_stream_messages,
aget_agent_card,
asend_message,
create_a2a_client,
)
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT
from litellm.llms.custom_httpx.http_handler import (
@ -236,6 +242,7 @@ class _RequestRecorder:
self.card = card
self.rpc_reply = rpc_reply
self.card_requests = []
self.card_urls = []
self.rpc_requests = []
self.client = None
@ -243,16 +250,19 @@ class _RequestRecorder:
headers = {k.lower(): v for k, v in request.headers.items()}
if request.method == "GET":
self.card_requests.append(headers)
self.card_urls.append(str(request.url))
return httpx.Response(200, json=self.card)
self.rpc_requests.append(headers)
return httpx.Response(200, json=self.rpc_reply)
def _a2a_client_cache_key(timeout: float) -> str:
return "async_httpx_client" + f"timeout_{timeout}" + httpxSpecialProvider.A2AProvider
def _a2a_client_cache_key(timeout: float, provider: str = httpxSpecialProvider.A2AProvider) -> str:
return "async_httpx_client" + f"timeout_{timeout}" + provider
async def _seed_shared_a2a_client(card=_AGENT_CARD, rpc_reply=_RPC_REPLY) -> _RequestRecorder:
async def _seed_shared_a2a_client(
card=_AGENT_CARD, rpc_reply=_RPC_REPLY, provider: str = httpxSpecialProvider.A2AProvider
) -> _RequestRecorder:
"""Put the one A2A client the cache will hand out behind a mock transport.
Seeding has to happen on the test's own event loop, because the client cache keys on
@ -265,9 +275,11 @@ async def _seed_shared_a2a_client(card=_AGENT_CARD, rpc_reply=_RPC_REPLY) -> _Re
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(recorder))
await owned_client.aclose()
litellm.in_memory_llm_clients_cache.set_cache(key=_a2a_client_cache_key(DEFAULT_A2A_AGENT_TIMEOUT), value=handler)
litellm.in_memory_llm_clients_cache.set_cache(
key=_a2a_client_cache_key(DEFAULT_A2A_AGENT_TIMEOUT, provider), value=handler
)
seeded = get_async_httpx_client(
llm_provider=httpxSpecialProvider.A2AProvider,
llm_provider=provider,
params={"timeout": DEFAULT_A2A_AGENT_TIMEOUT},
)
assert seeded is handler, "cache key drifted from get_async_httpx_client; these tests would test nothing"
@ -397,6 +409,36 @@ async def test_agent_card_fetch_carries_the_callers_headers(isolated_client_cach
assert recorder.card_requests[-1]["x-agent-token"] == "token-for-a"
@pytest.mark.asyncio
async def test_agent_card_path_param_fetches_that_path_with_the_agents_headers(isolated_client_cache):
"""A Microsoft Foundry agent serves its card only at agentCard/v1.0 behind the same Entra bearer
as the agent, so an agent registered with agent_card_path fetches exactly that path, authenticated,
instead of probing the well-known paths."""
recorder = await _seed_shared_a2a_client()
await asend_message(
request=_send_request("req-foundry"),
api_base="http://127.0.0.1:9",
litellm_params={"agent_card_path": "agentCard/v1.0"},
agent_extra_headers=_AGENT_A_HEADERS,
)
assert recorder.card_urls == ["http://127.0.0.1:9/agentCard/v1.0"]
assert recorder.card_requests[-1]["x-agent-token"] == "token-for-a"
@pytest.mark.asyncio
async def test_aget_agent_card_carries_the_callers_headers_and_path(isolated_client_cache):
recorder = await _seed_shared_a2a_client(provider=httpxSpecialProvider.A2A)
await aget_agent_card(
base_url="http://127.0.0.1:9", extra_headers=_AGENT_A_HEADERS, relative_card_path="agentCard/v1.0"
)
assert recorder.card_urls == ["http://127.0.0.1:9/agentCard/v1.0"]
assert recorder.card_requests[-1]["x-agent-token"] == "token-for-a"
@pytest.mark.asyncio
async def test_the_pooled_a2a_client_arrives_with_cookie_persistence_disabled(isolated_client_cache):
"""create_a2a_client takes its client from the shared builder rather than building one,
@ -464,3 +506,41 @@ async def test_asend_message_counts_usage_off_the_event_loop(monkeypatch):
assert recorder.payload["prompt_tokens"] > 100_000
assert recorder.payload["completion_tokens"] > 100_000
assert_loop_stayed_free(took, lags)
def test_streaming_logging_obj_keeps_agent_credentials_out_of_logging_params():
"""Callbacks receive the streaming logging object's litellm_params as raw kwargs, so an agent's
Entra, Databricks, or static credentials must never be copied into it; only pricing keys are."""
from litellm.a2a_protocol.main import _build_streaming_logging_obj
request = SendStreamingMessageRequest(
id="rpc-secrets",
params=MessageSendParams(
message={"messageId": "m1", "role": "user", "parts": [{"kind": "text", "text": "hi"}]}
),
)
logging_obj = _build_streaming_logging_obj(
request=request,
agent_name="foundry-agent",
agent_id="agent-1",
litellm_params={
"client_secret": "sp-secret",
"azure_ad_token": "entra-token",
"tenant_id": "tenant",
"databricks_oauth": {"client_secret": "dbx-secret"},
"api_key": "static-key",
"cost_per_query": 0.25,
},
metadata={"user_api_key": "hashed"},
proxy_server_request={"url": "http://localhost:4000"},
)
expected = {
"cost_per_query": 0.25,
"metadata": {"user_api_key": "hashed"},
"proxy_server_request": {"url": "http://localhost:4000"},
}
assert logging_obj.litellm_params == expected
assert logging_obj.optional_params == expected
assert logging_obj.model_call_details["litellm_params"] == expected

View file

@ -0,0 +1,36 @@
"""Tests for litellm/llms/a2a/chat/streaming_iterator.py."""
import pytest
from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator
from litellm.llms.a2a.common_utils import A2AError
def _iterator(lines: list[str]) -> A2AModelResponseIterator:
return A2AModelResponseIterator(streaming_response=iter(lines), sync_stream=True)
def test_a_jsonrpc_error_in_the_stream_fails_the_call():
"""An agent that answers message/stream with a JSON-RPC error (Microsoft Foundry replies -32004
"operation not supported") must fail the call with that message instead of ending an empty stream."""
iterator = _iterator(
['{"jsonrpc":"2.0","id":"1","error":{"code":-32004,"message":"This operation is not supported"}}']
)
with pytest.raises(A2AError, match="This operation is not supported"):
next(iterator)
def test_a_completed_task_chunk_yields_its_text_and_stops():
iterator = _iterator(
[
'{"jsonrpc":"2.0","id":"1","result":{"kind":"task","status":{"state":"completed"},'
'"artifacts":[{"parts":[{"kind":"text","text":"7"}]}]}}'
]
)
chunk = next(iterator)
assert chunk["text"] == "7"
assert chunk["is_finished"] is True
assert chunk["finish_reason"] == "stop"

View file

@ -2,6 +2,8 @@
from unittest.mock import MagicMock
import pytest
from litellm.llms.a2a.chat.transformation import A2AConfig
from litellm.types.utils import ModelResponse
@ -40,3 +42,46 @@ def test_transform_response_sets_usage():
assert result.usage.prompt_tokens > 0
assert result.usage.completion_tokens > 0
assert result.usage.total_tokens == (result.usage.prompt_tokens + result.usage.completion_tokens)
def test_transform_request_asks_the_agent_for_a_blocking_send():
"""Chat completions need the final answer in one response. Microsoft Foundry agents default to a
non-blocking send that returns a submitted task, so the request must opt into blocking."""
request = A2AConfig().transform_request(
model="a2a/test-agent",
messages=[{"role": "user", "content": "hi there agent"}],
optional_params={},
litellm_params={},
headers={},
)
assert request["method"] == "message/send"
assert request["params"]["configuration"] == {"blocking": True}
def test_transform_request_streams_without_a_send_configuration():
request = A2AConfig().transform_request(
model="a2a/test-agent",
messages=[{"role": "user", "content": "hi there agent"}],
optional_params={"stream": True},
litellm_params={},
headers={},
)
assert request["method"] == "message/stream"
assert "configuration" not in request["params"]
@pytest.mark.parametrize("optional_params", [{}, {"stream": True}])
def test_transform_request_tags_the_message_with_its_kind(optional_params: dict):
"""A2A 0.3 messages carry a `kind` discriminator; Microsoft Foundry rejects a message without it as
missing a required property, so both send methods must tag the message."""
request = A2AConfig().transform_request(
model="a2a/test-agent",
messages=[{"role": "user", "content": "hi there agent"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
assert request["params"]["message"]["kind"] == "message"

View file

@ -0,0 +1,52 @@
"""Tests for litellm/llms/a2a/common_utils.py."""
from collections.abc import Mapping
from types import MappingProxyType
import pytest
from litellm.llms.a2a.common_utils import resolve_a2a_hop_auth_header
class _RecordingEntraResolver:
def __init__(self) -> None:
self.calls: list[Mapping[str, object]] = []
async def __call__(self, litellm_params: Mapping[str, object]) -> Mapping[str, str]:
self.calls.append(litellm_params)
return MappingProxyType({"Authorization": "Bearer minted-entra-token"})
_SERVICE_PRINCIPAL = MappingProxyType({"tenant_id": "tenant", "client_id": "client", "client_secret": "sp-secret"})
@pytest.mark.asyncio
async def test_entra_agent_gets_a_minted_bearer_for_the_a2a_hop():
resolver = _RecordingEntraResolver()
header = await resolve_a2a_hop_auth_header(_SERVICE_PRINCIPAL, None, resolver)
assert header == {"Authorization": "Bearer minted-entra-token"}
assert resolver.calls == [_SERVICE_PRINCIPAL]
@pytest.mark.asyncio
async def test_completion_bridge_agent_keeps_its_entra_credentials_for_the_model_provider():
"""A bridged agent's tenant_id/client_id/client_secret authenticate the model it bridges to, so the A2A hop
must not spend them on a bearer of its own."""
resolver = _RecordingEntraResolver()
header = await resolve_a2a_hop_auth_header(_SERVICE_PRINCIPAL, "azure_ai", resolver)
assert header is None
assert resolver.calls == []
@pytest.mark.asyncio
async def test_agent_without_entra_credentials_gets_no_bearer():
resolver = _RecordingEntraResolver()
header = await resolve_a2a_hop_auth_header({"api_base": "https://agent.example.com"}, None, resolver)
assert header is None
assert resolver.calls == []

View file

@ -10,7 +10,12 @@ from unittest.mock import patch
import pytest
import litellm
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
from litellm.llms.azure_ai.common_utils import (
get_azure_ai_agent_entra_token,
get_azure_ai_auth_headers,
has_azure_entra_params,
resolve_azure_ai_agent_auth_header,
)
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
ENTRA_PARAMS = {"azure_ad_token": "entra-token"}
@ -152,3 +157,148 @@ def test_image_generation_still_uses_api_key_header():
headers = mock_image_generation.call_args.kwargs["headers"]
assert headers["api-key"] == "my-key"
assert "Authorization" not in headers
def test_agents_without_entra_credentials_are_not_treated_as_entra_agents():
"""Only a credential-bearing field opts an agent into Entra auth: scope or identity fields alone
must never make the proxy mint a bearer for that agent's URL."""
assert has_azure_entra_params({"api_key": "static", "headers": {"x": "y"}}) is False
assert has_azure_entra_params(None) is False
assert has_azure_entra_params({"azure_scope": "https://ai.azure.com/.default"}) is False
assert has_azure_entra_params({"tenant_id": "t", "client_id": "c"}) is False
assert has_azure_entra_params({"azure_ad_token": "entra-token"}) is True
assert has_azure_entra_params({"tenant_id": "t", "client_id": "c", "client_secret": "s"}) is True
assert has_azure_entra_params({"client_id": "c", "azure_username": "u", "azure_password": "p"}) is True
def test_agent_entra_token_ignores_the_process_wide_azure_credentials(monkeypatch):
"""The azure provider's token helper falls back to AZURE_* env vars. An agent's bearer must come
from that agent's own litellm_params only, or the host's service principal would authenticate to
whatever URL an agent registers."""
monkeypatch.setenv("AZURE_TENANT_ID", "host-tenant")
monkeypatch.setenv("AZURE_CLIENT_ID", "host-client")
monkeypatch.setenv("AZURE_CLIENT_SECRET", "host-secret")
monkeypatch.setenv("AZURE_AD_TOKEN", "host-token")
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch so a host-credential leak would show up as a call instead of a network round trip
mock_entra_id.return_value = lambda: "host-sp-token"
with pytest.raises(ValueError, match="client_secret"):
get_azure_ai_agent_entra_token({"azure_scope": "https://ai.azure.com/.default"})
assert get_azure_ai_agent_entra_token({"azure_ad_token": "agent-token"}) == "agent-token"
mock_entra_id.assert_not_called()
def test_agent_service_principal_fields_resolve_os_environ_references(monkeypatch):
monkeypatch.setenv("FOUNDRY_AGENT_TENANT_ID", "tenant-from-env")
monkeypatch.setenv("FOUNDRY_AGENT_CLIENT_ID", "client-from-env")
monkeypatch.setenv("FOUNDRY_AGENT_CLIENT_SECRET", "secret-from-env")
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to assert the resolved secret values reach the credential; live SP path proven by the PR's Azure Foundry e2e QA
mock_entra_id.return_value = lambda: "sp-token"
token = get_azure_ai_agent_entra_token(
{
"tenant_id": "os.environ/FOUNDRY_AGENT_TENANT_ID",
"client_id": "os.environ/FOUNDRY_AGENT_CLIENT_ID",
"client_secret": "os.environ/FOUNDRY_AGENT_CLIENT_SECRET",
}
)
mock_entra_id.assert_called_once_with(
tenant_id="tenant-from-env",
client_id="client-from-env",
client_secret="secret-from-env",
scope="https://ai.azure.com/.default",
)
assert token == "sp-token"
def test_agent_service_principal_wins_over_a_static_token_on_the_same_agent():
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to pin the precedence between a refreshing credential and a static token
mock_entra_id.return_value = lambda: "sp-token"
token = get_azure_ai_agent_entra_token(
{"tenant_id": "tenant", "client_id": "client", "client_secret": "secret", "azure_ad_token": "stale-token"}
)
assert token == "sp-token"
def test_agent_service_principal_token_defaults_to_the_foundry_agents_scope():
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to assert the scope Foundry agents require reaches the credential; live SP path proven by the PR's Azure Foundry e2e QA
mock_entra_id.return_value = lambda: "sp-token"
token = get_azure_ai_agent_entra_token({"tenant_id": "tenant", "client_id": "client", "client_secret": "secret"})
mock_entra_id.assert_called_once_with(
tenant_id="tenant",
client_id="client",
client_secret="secret",
scope="https://ai.azure.com/.default",
)
assert token == "sp-token"
def test_agent_azure_scope_overrides_the_foundry_agents_default():
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to assert an explicit azure_scope wins over the agents default; live SP path proven by the PR's Azure Foundry e2e QA
mock_entra_id.return_value = lambda: "sp-token"
get_azure_ai_agent_entra_token(
{"tenant_id": "tenant", "client_id": "client", "client_secret": "secret", "azure_scope": "custom/.default"}
)
assert mock_entra_id.call_args.kwargs["scope"] == "custom/.default"
def test_agent_entra_values_resolve_os_environ_references(monkeypatch):
monkeypatch.setenv("FOUNDRY_AGENT_AD_TOKEN", "token-from-env")
assert get_azure_ai_agent_entra_token({"azure_ad_token": "os.environ/FOUNDRY_AGENT_AD_TOKEN"}) == "token-from-env"
def test_agent_entra_token_failure_names_the_credential_fields():
with pytest.raises(ValueError, match="client_secret"):
get_azure_ai_agent_entra_token({"azure_scope": "https://ai.azure.com/.default"})
def test_agent_oidc_token_without_agent_ids_never_borrows_the_host_identity(monkeypatch):
"""The shared OIDC helper fills a missing client and tenant id from AZURE_CLIENT_ID and AZURE_TENANT_ID,
which would exchange the host's federated token for the host's identity at that agent's URL."""
monkeypatch.setenv("AZURE_TENANT_ID", "host-tenant")
monkeypatch.setenv("AZURE_CLIENT_ID", "host-client")
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_oidc") as mock_oidc: # test-quality-ok: stubs the OIDC exchange so a host-identity leak would show up as a call instead of a network round trip
mock_oidc.return_value = "host-minted-token"
with pytest.raises(ValueError, match="oidc/"):
get_azure_ai_agent_entra_token({"azure_ad_token": "oidc/github"})
with pytest.raises(ValueError, match="oidc/"):
get_azure_ai_agent_entra_token({"azure_ad_token": "oidc/github", "tenant_id": "agent-tenant"})
mock_oidc.assert_not_called()
def test_agent_oidc_token_exchanges_with_the_agent_ids_and_scope():
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_oidc") as mock_oidc: # test-quality-ok: stubs the OIDC exchange to assert the agent's own ids and the Foundry scope reach it
mock_oidc.return_value = "agent-minted-token"
token = get_azure_ai_agent_entra_token(
{"azure_ad_token": "oidc/github", "tenant_id": "agent-tenant", "client_id": "agent-client"}
)
assert token == "agent-minted-token"
mock_oidc.assert_called_once_with(
azure_ad_token="oidc/github",
azure_client_id="agent-client",
azure_tenant_id="agent-tenant",
scope="https://ai.azure.com/.default",
)
@pytest.mark.asyncio
async def test_agent_auth_header_is_the_entra_bearer():
headers = await resolve_azure_ai_agent_auth_header({"azure_ad_token": "entra-token"})
assert headers == {"Authorization": "Bearer entra-token"}

View file

@ -124,9 +124,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
MessageSendParams = make_mock_pydantic_class("MessageSendParams")
SendMessageRequest = make_mock_pydantic_class("SendMessageRequest")
SendStreamingMessageRequest = make_mock_pydantic_class(
"SendStreamingMessageRequest"
)
SendStreamingMessageRequest = make_mock_pydantic_class("SendStreamingMessageRequest")
# Create a mock module for a2a.types
mock_a2a_types = MagicMock()
@ -359,10 +357,9 @@ async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge():
user_api_key_dict=mock_user_api_key_dict,
)
assert (
captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM)
== mock_user_api_key_dict.api_key
), "authenticated key hash was not forwarded to the completion bridge"
assert captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM) == mock_user_api_key_dict.api_key, (
"authenticated key hash was not forwarded to the completion bridge"
)
def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
@ -376,9 +373,7 @@ def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
return agent
def _make_request_mock(
method: str, params: Mapping[str, object], request_id: object = "req-1"
) -> MagicMock:
def _make_request_mock(method: str, params: Mapping[str, object], request_id: object = "req-1") -> MagicMock:
req = MagicMock()
req.headers = {}
req.json = AsyncMock(
@ -436,6 +431,7 @@ async def _invoke_message_method(
mock_request: MagicMock,
user_api_key_dict: UserAPIKeyAuth,
add_litellm_data: AddLiteLLMData | None = None,
agent: MagicMock | None = None,
) -> CapturedAgentCall:
from fastapi.responses import JSONResponse
@ -466,7 +462,7 @@ async def _invoke_message_method(
downstream: Final = AsyncMock(side_effect=fake_asend_message if is_send else fake_stream_message)
with ExitStack() as stack:
for p in _base_patches(_make_agent_mock(), add_litellm_data):
for p in _base_patches(agent or _make_agent_mock(), add_litellm_data):
stack.enter_context(p)
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
if is_send:
@ -515,6 +511,98 @@ async def test_message_methods_forward_caller_identity_headers(method: str):
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "team-xyz"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_send_the_entra_bearer_for_azure_agents(method: str):
"""A Microsoft Foundry agent accepts only an Entra ID bearer, so an agent registered with
Entra credentials in litellm_params must reach the backend with that bearer on every call."""
agent = _make_agent_mock()
agent.litellm_params = {"azure_ad_token": "entra-token"}
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
assert (captured.agent_extra_headers or {}).get("Authorization") == "Bearer entra-token"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_leave_agents_without_entra_params_unauthenticated(method: str):
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
assert "Authorization" not in (captured.agent_extra_headers or {})
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_leave_entra_fields_to_the_model_provider_for_bridge_agents(method: str):
"""A completion-bridge agent's tenant_id/client_id/client_secret belong to the model provider it
calls through litellm, so the proxy must not mint a Foundry bearer for them."""
agent = _make_agent_mock()
agent.litellm_params = {
"custom_llm_provider": "azure_ai",
"model": "azure_ai/foundry-model",
"tenant_id": "tenant",
"client_id": "client",
"client_secret": "sp-secret",
}
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
assert "Authorization" not in (captured.agent_extra_headers or {})
@pytest.mark.asyncio
async def test_message_send_reports_an_unresolvable_entra_credential_as_internal_error(monkeypatch):
"""An agent whose Entra credential points at an unset environment variable must fail the call
with the JSON-RPC internal error naming the credential fields, never reach the backend unauthenticated."""
monkeypatch.delenv("LITELLM_TEST_UNSET_FOUNDRY_TOKEN", raising=False)
agent = _make_agent_mock()
agent.litellm_params = {"azure_ad_token": "os.environ/LITELLM_TEST_UNSET_FOUNDRY_TOKEN"}
mock_request = _make_request_mock("message/send", _HELLO_MESSAGE_PARAMS)
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)
downstream = AsyncMock()
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 tests use; the request must fail before any backend call is made
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
)
)
stack.enter_context(
patch( # test-quality-ok: the observation point proving the backend is never called; the sibling send tests use the same seam
"litellm.a2a_protocol.asend_message", new=downstream
)
)
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 == 500
assert body["error"]["code"] == -32603
assert "client_secret" in body["error"]["message"]
downstream.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_caller_identity_headers_cannot_be_spoofed(method: str):
@ -528,12 +616,12 @@ async def test_message_methods_caller_identity_headers_cannot_be_spoofed(method:
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
forwarded_headers = captured.agent_extra_headers or {}
assert (
forwarded_headers.get("X-LiteLLM-User-Id") == "real-user"
), "authenticated user id must not be overridden by forwarded client headers"
assert (
forwarded_headers.get("X-LiteLLM-Team-Id") == "real-team"
), "authenticated team id must not be overridden by forwarded client headers"
assert forwarded_headers.get("X-LiteLLM-User-Id") == "real-user", (
"authenticated user id must not be overridden by forwarded client headers"
)
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "real-team", (
"authenticated team id must not be overridden by forwarded client headers"
)
@pytest.mark.asyncio
@ -637,6 +725,47 @@ async def test_task_methods_forward_jsonrpc(method: str, params: dict):
assert forwarded_body["method"] == method
@pytest.mark.asyncio
async def test_task_methods_forward_the_entra_bearer_for_azure_agents():
"""tasks/get on a Foundry agent polls the task the agent created, so the forwarded call needs
the same Entra bearer as message/send."""
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
agent.litellm_params = {"azure_ad_token": "entra-token"}
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
mock_http_response = MagicMock()
mock_http_response.json.return_value = {"jsonrpc": "2.0", "id": "req-1", "result": {"id": "task-1"}}
mock_http_response.is_success = True
mock_http_response.raise_for_status = MagicMock()
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = MagicMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch( # test-quality-ok: the task route builds its own httpx client; the sibling task tests capture the post through the same seam
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=mock_handler
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1"),
)
posted_headers = mock_handler.post.call_args.kwargs["headers"]
assert posted_headers["Authorization"] == "Bearer entra-token"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["tasks/get", "tasks/resubscribe"])
async def test_task_methods_extract_litellm_params_before_forwarding(method: str):
@ -808,9 +937,7 @@ async def test_subscribe_to_task_calls_pre_call_hook():
yield chunk
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.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
mock_proxy_logging.async_post_call_streaming_iterator_hook = _passthrough_iterator
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
@ -866,9 +993,7 @@ async def test_subscribe_to_task_runs_post_call_streaming_guardrail():
inspected.append(response)
return response
guardrail = _RecordingGuardrail(
guardrail_name="record-a2a", default_on=True, event_hook="post_call"
)
guardrail = _RecordingGuardrail(guardrail_name="record-a2a", default_on=True, event_hook="post_call")
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"})
@ -918,8 +1043,7 @@ async def test_subscribe_to_task_runs_post_call_streaming_guardrail():
pass
assert any("resubscribe-secret" in str(r) for r in inspected), (
"tasks/resubscribe streamed content was not passed to the post-call "
"streaming guardrail hook"
"tasks/resubscribe streamed content was not passed to the post-call streaming guardrail hook"
)
@ -946,9 +1070,7 @@ async def test_task_method_failure_hook_uses_enriched_request_data():
mock_handler.post = AsyncMock(side_effect=RuntimeError("upstream failed"))
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.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:
@ -984,9 +1106,7 @@ async def test_task_method_failure_hook_uses_enriched_request_data():
body = json.loads(response.body.decode())
assert body["error"]["code"] == -32603
failure_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs[
"request_data"
]
failure_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs["request_data"]
assert failure_data.get("litellm_call_id")
assert failure_data.get("agent_id") == "test-agent"
@ -1015,9 +1135,7 @@ async def test_agentcore_invalid_context_id_returns_jsonrpc_invalid_params_400()
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.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:
@ -1129,10 +1247,7 @@ async def test_get_agent_card_uses_proxy_base_url_when_set(monkeypatch):
body = json.loads(response.body.decode())
assert body["url"] == "https://litellm.example.com/a2a/test-agent"
assert (
body["supportedInterfaces"][0]["url"]
== "https://litellm.example.com/a2a/test-agent"
)
assert body["supportedInterfaces"][0]["url"] == "https://litellm.example.com/a2a/test-agent"
@pytest.mark.asyncio
@ -1182,9 +1297,7 @@ async def test_get_agent_card_0_3_card_with_a2a_version_1_0_header():
"url": "http://backend-agent:10001",
"version": "1.0.0",
"capabilities": {"streaming": True},
"skills": [
{"id": "s1", "name": "skill one", "description": "d", "tags": ["t"]}
],
"skills": [{"id": "s1", "name": "skill one", "description": "d", "tags": ["t"]}],
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
}
@ -1207,9 +1320,7 @@ async def test_get_agent_card_0_3_card_with_a2a_version_1_0_header():
body = json.loads(response.body.decode())
assert "url" not in body
assert body["supportedInterfaces"][0]["url"] == (
"http://localhost:4000/a2a/test-agent"
)
assert body["supportedInterfaces"][0]["url"] == ("http://localhost:4000/a2a/test-agent")
@pytest.mark.asyncio
@ -1278,9 +1389,7 @@ def test_build_merged_agent_card_uses_proxy_base_url_for_supported_interfaces(
http_request=mock_request,
)
assert merged["supportedInterfaces"][0]["url"] == (
"https://litellm.example.com/a2a/jenkins_agent"
)
assert merged["supportedInterfaces"][0]["url"] == ("https://litellm.example.com/a2a/jenkins_agent")
@pytest.mark.asyncio
@ -1324,9 +1433,7 @@ async def test_unknown_method_returns_jsonrpc_error():
("GetExtendedAgentCard", "agent/getAuthenticatedExtendedCard"),
],
)
async def test_pascal_method_names_normalize_to_wire_format(
pascal_method: str, expected_wire_method: str
):
async def test_pascal_method_names_normalize_to_wire_format(pascal_method: str, expected_wire_method: str):
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
@ -1448,9 +1555,7 @@ async def test_handle_stream_message_rejects_invalid_params_with_32602():
)
assert response.media_type == "text/event-stream"
chunks = [chunk async for chunk in response.body_iterator]
body = "".join(
chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks
)
body = "".join(chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks)
assert body.startswith("data: ")
assert body.endswith("\n\n")
payload = json.loads(body.removeprefix("data: ").strip())
@ -1504,10 +1609,7 @@ async def test_handle_stream_message_frames_events_as_sse():
)
assert response.media_type == "text/event-stream"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == len(events)
for chunk, event in zip(chunks, events):
@ -1530,10 +1632,7 @@ async def test_handle_stream_message_sdk_unavailable_frames_error_as_sse():
)
assert response.media_type == "text/event-stream"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 1
assert chunks[0].startswith("data: ")
assert chunks[0].endswith("\n\n")
@ -1569,9 +1668,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse():
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
@ -1589,10 +1686,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse():
)
assert response.media_type == "text/event-stream"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == len(events)
for chunk, event in zip(chunks, events):
@ -1620,9 +1714,7 @@ async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once():
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
@ -1636,10 +1728,7 @@ async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once():
},
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 1
payload = json.loads(chunks[0].removeprefix("data: ").strip())
@ -1661,9 +1750,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse():
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
@ -1680,10 +1767,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse():
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 2
assert chunks[-1].startswith("data: ")
@ -1707,9 +1791,7 @@ async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error()
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
@ -1726,10 +1808,7 @@ async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error()
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 1
error_payload = json.loads(chunks[0].removeprefix("data: ").strip())
@ -1749,9 +1828,7 @@ async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event():
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
@ -1765,10 +1842,7 @@ async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event():
},
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert chunks == ['data: "not json at all"\n\n']
@ -1785,9 +1859,7 @@ async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error():
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
@ -1801,10 +1873,7 @@ async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error():
},
)
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 2
error_payload = json.loads(chunks[-1].removeprefix("data: ").strip())
@ -1911,10 +1980,7 @@ def test_normalize_response_keeps_wire_format_for_0_3():
"role": "agent",
},
}
assert (
normalize_jsonrpc_response(wire_response, "0.3", method="message/send")
is wire_response
)
assert normalize_jsonrpc_response(wire_response, "0.3", method="message/send") is wire_response
@pytest.mark.asyncio
@ -1936,9 +2002,7 @@ async def test_task_method_upstream_jsonrpc_error_on_http_4xx_is_relayed():
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_error
mock_http_response.is_success = False
mock_http_response.raise_for_status = MagicMock(
side_effect=Exception("404 Not Found")
)
mock_http_response.raise_for_status = MagicMock(side_effect=Exception("404 Not Found"))
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
@ -1982,9 +2046,7 @@ async def test_subscribe_to_task_upstream_error_yields_jsonrpc_error_event():
mock_resp.is_success = False
mock_resp.status_code = 404
mock_resp.reason_phrase = "Not Found"
mock_resp.aread = AsyncMock(
return_value=b'{"jsonrpc":"2.0","error":{"code":-32001,"message":"Task not found"}}'
)
mock_resp.aread = AsyncMock(return_value=b'{"jsonrpc":"2.0","error":{"code":-32001,"message":"Task not found"}}')
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
@ -2076,9 +2138,7 @@ async def test_task_methods_forward_caller_identity_headers():
}
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test", user_id="user-abc", team_id="team-xyz"
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="user-abc", team_id="team-xyz")
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
@ -2364,9 +2424,7 @@ async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers()
"x-a2a-test-agent-x-litellm-user-id": "attacker-user",
"x-a2a-test-agent-x-litellm-team-id": "attacker-team",
}
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test", user_id="real-user", team_id="real-team"
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="real-user", team_id="real-team")
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
@ -2395,19 +2453,17 @@ async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers()
)
posted_headers = mock_handler.post.call_args.kwargs.get("headers") or {}
assert (
posted_headers.get("X-LiteLLM-User-Id") == "real-user"
), "authenticated user id must not be overridden by forwarded client headers"
assert (
posted_headers.get("X-LiteLLM-Team-Id") == "real-team"
), "authenticated team id must not be overridden by forwarded client headers"
assert posted_headers.get("X-LiteLLM-User-Id") == "real-user", (
"authenticated user id must not be overridden by forwarded client headers"
)
assert posted_headers.get("X-LiteLLM-Team-Id") == "real-team", (
"authenticated team id must not be overridden by forwarded client headers"
)
def _agent(protocol_version):
agent = MagicMock()
agent.agent_card_params = (
{"protocolVersion": protocol_version} if protocol_version is not None else {}
)
agent.agent_card_params = {"protocolVersion": protocol_version} if protocol_version is not None else {}
return agent
@ -2553,16 +2609,11 @@ async def test_handle_stream_message_pings_while_the_upstream_agent_is_still_sil
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _stream_message_response()
assert response.headers["x-accel-buffering"] == "no"
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert chunks[0] == ": ping\n\n"
assert chunks.count(": ping\n\n") >= 3
@ -2583,16 +2634,26 @@ async def test_handle_stream_message_is_untouched_while_keepalives_are_unconfigu
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
)
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _stream_message_response()
assert "x-accel-buffering" not in response.headers
chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in response.body_iterator
]
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert not any(chunk.startswith(":") for chunk in chunks)
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
def test_forwarding_headers_minted_bearer_replaces_a_forwarded_authorization_of_any_case():
"""A client header the admin chose to forward keeps the casing the config named it with, so a forwarded
`authorization` must not travel next to the minted `Authorization` as a second header line."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _forwarding_headers
merged = _forwarding_headers(
caller_identity={},
request_data={},
agent_extra_headers={"authorization": "Bearer client-token", "X-Custom": "kept"},
backend_auth_header={"Authorization": "Bearer minted-token"},
)
assert merged == {"X-Custom": "kept", "Authorization": "Bearer minted-token"}

View file

@ -4,8 +4,10 @@ Test A2A provider registry lookup functionality.
Maps to: litellm/llms/a2a/chat/transformation.py
"""
import json
from unittest.mock import patch
import httpx
import pytest
import litellm
@ -15,19 +17,20 @@ from litellm.llms.a2a.chat.transformation import A2AConfig
def test_resolve_agent_config_from_registry_static_method():
"""Test the static helper method for registry resolution"""
# Test 1: No agent name in model
# Test 1: Unregistered agent name keeps the explicit config
api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry(
model="a2a",
agent_name="not-registered",
api_base="http://test.com",
api_key=None,
headers=None,
optional_params={},
)
assert api_base == "http://test.com"
assert api_key is None
# Test 2: All params provided - should not lookup registry
api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry(
model="a2a/test-agent",
agent_name="test-agent",
api_base="http://explicit.com",
api_key="explicit-key",
headers={"X-Test": "value"},
@ -38,34 +41,297 @@ def test_resolve_agent_config_from_registry_static_method():
def test_a2a_registry_integration():
"""Test registry lookup in proxy context"""
"""A chat call for a registered agent must post to the registered url with the registered key as the
bearer even though completion() strips the a2a/ prefix before the lookup runs."""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
test_agent = AgentResponse(
agent_id="test-id",
agent_name="test-agent",
agent_card_params={"url": "http://registry-url.example.com:9999"},
litellm_params={"api_key": "registry-key", "headers": {"X-Agent": "static"}},
)
client = HTTPHandler()
agent_reply = httpx.Response(
200,
json={"jsonrpc": "2.0", "id": "1", "result": {"kind": "message", "parts": [{"kind": "text", "text": "4"}]}},
)
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(test_agent)
try:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
# Create test agent
test_agent = AgentResponse(
agent_id="test-id",
agent_name="test-agent",
agent_card_params={"url": "http://registry-url.example.com:9999"},
litellm_params={"api_key": "registry-key"},
)
# Register and test
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(test_agent)
try:
litellm.completion(
model="a2a/test-agent", messages=[{"role": "user", "content": "Hello"}]
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
response = litellm.completion(
model="a2a/test-agent", messages=[{"role": "user", "content": "What is 2+2?"}], client=client
)
except Exception as e:
# Should use registry URL (connection error expected)
if "registry-url.example.com" not in str(e) and "APIConnectionError" not in type(e).__name__:
raise
finally:
global_agent_registry.agent_list = original_agents
finally:
global_agent_registry.agent_list = original_agents
except ImportError:
pytest.skip("Registry not available (not in proxy context)")
assert response.choices[0].message.content == "4"
assert post.call_args.kwargs["url"] == "http://registry-url.example.com:9999"
assert post.call_args.kwargs["headers"]["Authorization"] == "Bearer registry-key"
assert post.call_args.kwargs["headers"]["X-Agent"] == "static"
def test_one_callers_bearer_never_reaches_another_caller_of_the_same_registered_agent():
"""The registered headers dict is shared by every request to the agent, so the bearer one caller
supplies must be written to that request alone and never persisted onto the agent for the next
caller, who has no key of their own."""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
shared_agent = AgentResponse(
agent_id="shared-id",
agent_name="shared-agent",
agent_card_params={"url": "http://registry-url.example.com:9999"},
litellm_params={"headers": {"X-Agent": "static"}},
)
client = HTTPHandler()
agent_reply = httpx.Response(
200,
json={"jsonrpc": "2.0", "id": "1", "result": {"kind": "message", "parts": [{"kind": "text", "text": "ok"}]}},
)
messages = [{"role": "user", "content": "hi"}]
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(shared_agent)
try:
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
litellm.completion(model="a2a/shared-agent", messages=messages, api_key="caller-one-key", client=client)
litellm.completion(model="a2a/shared-agent", messages=messages, client=client)
finally:
global_agent_registry.agent_list = original_agents
first_call_headers, second_call_headers = (call.kwargs["headers"] for call in post.call_args_list)
assert first_call_headers["Authorization"] == "Bearer caller-one-key"
assert "Authorization" not in second_call_headers
assert second_call_headers["X-Agent"] == "static"
assert shared_agent.litellm_params == {"headers": {"X-Agent": "static"}}
def _foundry_card_stored_through_the_agents_api() -> dict:
from litellm.proxy.a2a.agent_card import merge_agent_card
return merge_agent_card(
{"name": "Foundry", "url": "https://foundry.example.com/a2a", "capabilities": {"streaming": False}},
proxy_url="http://localhost:4000/a2a/foundry-agent",
proxy_base_url="http://localhost:4000",
)
@pytest.mark.parametrize(
"agent_card_params",
[
{"url": "https://foundry.example.com/a2a", "capabilities": {"streaming": False}},
_foundry_card_stored_through_the_agents_api(),
],
ids=["card registered verbatim from config.yaml", "card stored through POST /v1/agents"],
)
def test_streaming_chat_to_an_agent_whose_card_declines_streaming_uses_a_blocking_send(agent_card_params: dict):
"""Microsoft Foundry agents publish `capabilities.streaming: false` and answer message/stream with a
JSON-RPC error. A streaming chat call to such an agent must post a blocking message/send and hand the
caller the answer as a stream, whether the card was registered verbatim from config.yaml or stored
through POST /v1/agents, which keeps only truthy capabilities and so drops the `false` itself."""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
foundry_agent = AgentResponse(
agent_id="foundry-id",
agent_name="foundry-agent",
agent_card_params=agent_card_params,
litellm_params={"api_key": "registry-key"},
)
client = HTTPHandler()
agent_reply = httpx.Response(
200,
json={
"jsonrpc": "2.0",
"id": "1",
"result": {
"kind": "task",
"status": {"state": "completed"},
"artifacts": [{"parts": [{"kind": "text", "text": "4"}]}],
},
},
)
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(foundry_agent)
try:
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
chunks = list(
litellm.completion(
model="a2a/foundry-agent",
messages=[{"role": "user", "content": "What is 2+2?"}],
stream=True,
client=client,
)
)
finally:
global_agent_registry.agent_list = original_agents
posted = json.loads(post.call_args.kwargs["data"])
assert posted["method"] == "message/send"
assert posted["params"]["configuration"] == {"blocking": True}
assert post.call_args.kwargs.get("stream", False) is False
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "4"
assert chunks[-1].choices[0].finish_reason == "stop"
@pytest.mark.parametrize(
"agent_card_params",
[
{"url": "https://agent.example.com/a2a"},
{"url": "https://agent.example.com/a2a", "capabilities": {"streaming": True}},
],
ids=["card without a capabilities block", "card says streaming true"],
)
def test_registry_lookup_leaves_streaming_alone_when_the_card_does_not_decline_it(agent_card_params: dict):
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
silent_agent = AgentResponse(
agent_id="silent-id",
agent_name="silent-agent",
agent_card_params=agent_card_params,
litellm_params={"api_key": "registry-key"},
)
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(silent_agent)
optional_params: dict = {"stream": True}
try:
A2AConfig.resolve_agent_config_from_registry(
agent_name="silent-agent", api_base=None, api_key=None, headers=None, optional_params=optional_params
)
finally:
global_agent_registry.agent_list = original_agents
assert optional_params == {"stream": True}
def test_registry_entra_agent_authenticates_with_the_entra_token_and_keeps_its_secrets_private():
"""An agent registered with Entra credentials has no api_key, so the chat route must resolve the
bearer from those credentials, and the credential fields must not ride along into optional_params
where they would reach spend logs and callbacks."""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
entra_agent = AgentResponse(
agent_id="entra-id",
agent_name="entra-agent",
agent_card_params={"url": "https://foundry.example.com/a2a"},
litellm_params={"azure_ad_token": "entra-token", "tenant_id": "tenant", "timeout": 30},
)
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(entra_agent)
optional_params: dict = {}
try:
api_base, api_key, _headers = A2AConfig.resolve_agent_config_from_registry(
agent_name="entra-agent",
api_base=None,
api_key=None,
headers=None,
optional_params=optional_params,
)
finally:
global_agent_registry.agent_list = original_agents
assert api_base == "https://foundry.example.com/a2a"
assert api_key == "entra-token"
assert optional_params == {"timeout": 30}
_STORED_STATIC_CREDENTIALS: dict = {
"api_key": "stored-key",
"headers": {"authorization": "Bearer stored-header", "X-Agent": "static"},
}
@pytest.mark.parametrize(
("litellm_params", "expected_authorization_lines"),
[
(
{**_STORED_STATIC_CREDENTIALS, "azure_ad_token": "entra-token"},
{"Authorization": "Bearer entra-token"},
),
(
_STORED_STATIC_CREDENTIALS,
{"authorization": "Bearer stored-header", "Authorization": "Bearer stored-key"},
),
(
{**_STORED_STATIC_CREDENTIALS, "azure_ad_token": "model-provider-token", "custom_llm_provider": "azure_ai"},
{"authorization": "Bearer stored-header", "Authorization": "Bearer stored-key"},
),
],
ids=[
"entra agent: the minted bearer is the only authorization line",
"agent without entra credentials: static credentials sent as before",
"bridge agent: its entra credentials belong to the model provider, never to the a2a hop",
],
)
def test_entra_credentials_beat_the_static_credentials_stored_next_to_them_on_the_chat_route(
litellm_params: dict, expected_authorization_lines: dict
):
"""The relay sends the minted Entra bearer over any static Authorization stored on the agent; the chat
route must agree, or an api_key or authorization header left next to the Entra fields makes the same
agent answer on /a2a and fail with the backend's 401 on /v1/chat/completions."""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
agent = AgentResponse(
agent_id="mixed-credentials-id",
agent_name="mixed-credentials-agent",
agent_card_params={"url": "https://foundry.example.com/a2a"},
litellm_params=litellm_params,
)
client = HTTPHandler()
agent_reply = httpx.Response(
200,
json={"jsonrpc": "2.0", "id": "1", "result": {"kind": "message", "parts": [{"kind": "text", "text": "ok"}]}},
)
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(agent)
try:
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
litellm.completion(
model="a2a/mixed-credentials-agent", messages=[{"role": "user", "content": "hi"}], client=client
)
finally:
global_agent_registry.agent_list = original_agents
sent_headers = post.call_args.kwargs["headers"]
assert {
name: value for name, value in sent_headers.items() if name.lower() == "authorization"
} == expected_authorization_lines
assert sent_headers["X-Agent"] == "static"
def test_registry_entra_agent_with_an_unresolvable_credential_fails_the_chat_call(monkeypatch):
"""The chat route mints the Foundry bearer from the registered credentials; when they resolve to
nothing the caller must get the credential error instead of an unauthenticated backend call."""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
monkeypatch.delenv("LITELLM_TEST_UNSET_FOUNDRY_TOKEN", raising=False)
entra_agent = AgentResponse(
agent_id="entra-unset-id",
agent_name="entra-unset-agent",
agent_card_params={"url": "https://foundry.example.com/a2a"},
litellm_params={"azure_ad_token": "os.environ/LITELLM_TEST_UNSET_FOUNDRY_TOKEN"},
)
original_agents = global_agent_registry.agent_list.copy()
global_agent_registry.register_agent(entra_agent)
try:
with pytest.raises(litellm.APIConnectionError, match="client_secret"):
litellm.completion(model="a2a/entra-unset-agent", messages=[{"role": "user", "content": "hi"}])
finally:
global_agent_registry.agent_list = original_agents