mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
feat(a2a): reach Microsoft Foundry agents with Entra auth and versioned card discovery
Foundry serves its agent card only at agentCard/v1.0, accepts only an Entra ID bearer, and defaults to a non-blocking send, so the A2A relay and the chat completions route could not use it. The relay gains an agent_card_path litellm_param plus agentCard/v1.0 as a third discovery probe, mints a bearer from flat Entra fields on the agent (tenant_id, client_id, client_secret, azure_ad_token, azure_username, azure_password, azure_scope) for https://ai.azure.com/.default, and sends it on the card fetch, message/send, message/stream, tasks/* and the chat bridge. Chat completions look the registered agent up by its provider-stripped name so its api_key and headers reach the request, tag every message with its kind, ask for a blocking send, fall back to a blocking send when the registered card says streaming: false, and fail the call on a JSON-RPC error inside a stream instead of yielding an empty one. Entra fields stay out of the chat bridge's logged parameters. Resolves LIT-5122
This commit is contained in:
parent
16bbff6643
commit
e243237a7c
17 changed files with 983 additions and 261 deletions
|
|
@ -9,6 +9,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
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
|
||||
|
|
@ -145,9 +148,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 (Microsoft Foundry agents, which serve no well-known card)
|
||||
"""
|
||||
|
||||
async def get_agent_card(
|
||||
|
|
@ -158,18 +162,18 @@ class LiteLLMA2ACardResolver(_A2ACardResolver):
|
|||
"""
|
||||
Fetch the agent card, trying multiple well-known paths.
|
||||
|
||||
First tries the standard path, then falls back to the previous path.
|
||||
First tries the standard path, then the previous path, then Foundry's documented path.
|
||||
|
||||
Args:
|
||||
relative_card_path: Optional path to the agent card endpoint.
|
||||
If None, tries both well-known paths.
|
||||
If None, tries every known path in order.
|
||||
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
|
||||
A2AAgentCardDiscoveryError naming every probed path and its error when no path answers
|
||||
"""
|
||||
# If a specific path is provided, use the parent implementation
|
||||
if relative_card_path is not None:
|
||||
|
|
@ -178,28 +182,26 @@ class LiteLLMA2ACardResolver(_A2ACardResolver):
|
|||
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: dict[str, Any] | None,
|
||||
failures: tuple[tuple[str, Exception], ...],
|
||||
) -> "AgentCard":
|
||||
if not paths:
|
||||
raise A2AAgentCardDiscoveryError(base_url=self.base_url, failures=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))
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
@ -112,6 +114,15 @@ class A2AAgentCardError(A2AError):
|
|||
)
|
||||
|
||||
|
||||
class A2AAgentCardDiscoveryError(A2AAgentCardError):
|
||||
"""Raised when no known agent card path answered; names every path probed and why each failed."""
|
||||
|
||||
def __init__(self, base_url: str, failures: tuple[tuple[str, Exception], ...]) -> 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)
|
||||
|
||||
|
||||
class A2ALocalhostURLError(A2AConnectionError):
|
||||
"""
|
||||
Raised when an agent card contains a localhost/internal URL.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -3,11 +3,16 @@ 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,
|
||||
has_azure_entra_params,
|
||||
)
|
||||
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
|
||||
|
|
@ -26,6 +31,25 @@ 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 capabilities.get("streaming") is False
|
||||
|
||||
|
||||
def _registry_api_key(agent_litellm_params: dict[str, object]) -> str | None:
|
||||
configured_api_key: Final = agent_litellm_params.get("api_key")
|
||||
if isinstance(configured_api_key, str):
|
||||
return configured_api_key
|
||||
if has_azure_entra_params(agent_litellm_params):
|
||||
return get_azure_ai_agent_entra_token(agent_litellm_params)
|
||||
return None
|
||||
|
||||
|
||||
class A2AConfig(BaseConfig):
|
||||
"""
|
||||
Configuration for A2A (Agent-to-Agent) Protocol.
|
||||
|
|
@ -35,20 +59,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 +80,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 +98,25 @@ 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:
|
||||
if not headers:
|
||||
agent_headers: Final = agent.litellm_params.get("headers")
|
||||
if agent_headers:
|
||||
headers = agent_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)
|
||||
|
||||
|
|
@ -226,6 +253,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 +265,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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
@ -41,6 +43,76 @@ 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`, 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:
|
||||
"""
|
||||
Mint the Entra ID bearer for a Microsoft Foundry agent endpoint from the agent's own `litellm_params`.
|
||||
|
||||
Unlike the `azure` provider's `get_azure_ad_token`, this never falls back to the process-wide
|
||||
`AZURE_*` environment variables: only the credentials registered on the agent (literal values or
|
||||
`os.environ/` references) may authenticate a call to that agent's URL. Foundry agents accept only
|
||||
the `https://ai.azure.com/.default` scope, so that scope applies unless `azure_scope` is set.
|
||||
"""
|
||||
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
|
||||
)()
|
||||
if azure_ad_token and azure_ad_token.startswith("oidc/"):
|
||||
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:
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -2190,7 +2190,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,
|
||||
|
|
|
|||
|
|
@ -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.azure_ai.common_utils import has_azure_entra_params, resolve_azure_ai_agent_auth_header
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.a2a.version_convert import (
|
||||
A2AVersion,
|
||||
|
|
@ -157,10 +158,30 @@ 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:
|
||||
"""
|
||||
Mint the bearer the agent's backend requires, when the agent is configured for one.
|
||||
|
||||
Databricks Apps take a short-lived OAuth M2M token from a ``databricks_oauth`` block. Microsoft
|
||||
Foundry agents take an Entra ID token from the agent's own Entra credentials, but only when the
|
||||
proxy speaks A2A to that URL itself: for completion-bridge agents (``custom_llm_provider`` set)
|
||||
those same fields belong to the model provider and travel with the completion call instead.
|
||||
"""
|
||||
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
|
||||
return await resolve_databricks_app_auth_header(litellm_params)
|
||||
if not custom_llm_provider and has_azure_entra_params(litellm_params):
|
||||
return await resolve_azure_ai_agent_auth_header(litellm_params)
|
||||
return None
|
||||
|
||||
|
||||
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:
|
||||
passthrough: Final = tuple(
|
||||
(name, value)
|
||||
|
|
@ -169,7 +190,8 @@ def _forwarding_headers(
|
|||
)
|
||||
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))
|
||||
backend_auth: Final = backend_auth_header.items() if backend_auth_header else ()
|
||||
merged: Final = dict((*passthrough, *caller_identity.items(), *trace, *backend_auth))
|
||||
return merged or None
|
||||
|
||||
|
||||
|
|
@ -795,26 +817,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
|
||||
|
|
|
|||
|
|
@ -138,3 +138,64 @@ 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"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_falls_through_to_the_foundry_card_path():
|
||||
"""Microsoft Foundry agents serve their card only at /agentCard/v1.0 and 404 both well-known
|
||||
paths, so discovery must reach that path after the two well-known probes fail."""
|
||||
mock_agent_card = MagicMock()
|
||||
paths_called = []
|
||||
|
||||
async def mock_parent_get_agent_card(self, relative_card_path=None, http_kwargs=None):
|
||||
paths_called.append(relative_card_path)
|
||||
if relative_card_path == "/agentCard/v1.0":
|
||||
return mock_agent_card
|
||||
raise Exception("404 Not Found")
|
||||
|
||||
with patch.object(LiteLLMA2ACardResolver.__bases__[0], "get_agent_card", mock_parent_get_agent_card):
|
||||
resolver = LiteLLMA2ACardResolver(httpx_client=MagicMock(), base_url="https://foundry.example.com/a2a")
|
||||
result = await resolver.get_agent_card()
|
||||
|
||||
assert paths_called == ["/.well-known/agent-card.json", "/.well-known/agent.json", "/agentCard/v1.0"]
|
||||
assert result is mock_agent_card
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_explicit_path_skips_the_probes():
|
||||
mock_agent_card = MagicMock()
|
||||
paths_called = []
|
||||
|
||||
async def mock_parent_get_agent_card(self, relative_card_path=None, http_kwargs=None):
|
||||
paths_called.append(relative_card_path)
|
||||
return mock_agent_card
|
||||
|
||||
with patch.object(LiteLLMA2ACardResolver.__bases__[0], "get_agent_card", mock_parent_get_agent_card):
|
||||
resolver = LiteLLMA2ACardResolver(httpx_client=MagicMock(), base_url="https://foundry.example.com/a2a")
|
||||
result = await resolver.get_agent_card(relative_card_path="agentCard/v1.0")
|
||||
|
||||
assert paths_called == ["agentCard/v1.0"]
|
||||
assert result is mock_agent_card
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_names_every_probed_path_when_discovery_fails():
|
||||
"""A Foundry agent 401s its well-known paths and 404s the rest; surfacing only the last probe's
|
||||
error would hide the auth failure that actually explains the outage."""
|
||||
from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError
|
||||
|
||||
async def mock_parent_get_agent_card(self, relative_card_path=None, http_kwargs=None):
|
||||
if relative_card_path == "/.well-known/agent.json":
|
||||
raise Exception("HTTP 401 Unauthorized")
|
||||
raise Exception("HTTP 404 Not Found")
|
||||
|
||||
with patch.object(LiteLLMA2ACardResolver.__bases__[0], "get_agent_card", mock_parent_get_agent_card):
|
||||
resolver = LiteLLMA2ACardResolver(httpx_client=MagicMock(), base_url="https://foundry.example.com/a2a")
|
||||
with pytest.raises(A2AAgentCardDiscoveryError) as raised:
|
||||
await resolver.get_agent_card()
|
||||
|
||||
message = str(raised.value)
|
||||
assert "https://foundry.example.com/a2a" in message
|
||||
assert "/.well-known/agent-card.json (HTTP 404 Not Found)" in message
|
||||
assert "/.well-known/agent.json (HTTP 401 Unauthorized)" in message
|
||||
assert "/agentCard/v1.0 (HTTP 404 Not Found)" in message
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,114 @@ 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"})
|
||||
|
||||
|
||||
@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"}
|
||||
|
|
|
|||
|
|
@ -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,11 @@ 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"
|
||||
|
|
|
|||
|
|
@ -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,166 @@ 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_streaming_chat_to_an_agent_whose_card_declines_streaming_uses_a_blocking_send():
|
||||
"""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, and an agent whose card is silent about streaming keeps message/stream."""
|
||||
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={"url": "https://foundry.example.com/a2a", "capabilities": {"streaming": False}},
|
||||
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"
|
||||
|
||||
|
||||
def test_registry_lookup_leaves_streaming_alone_when_the_card_does_not_decline_it():
|
||||
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={"url": "https://agent.example.com/a2a"},
|
||||
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}
|
||||
|
||||
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue