mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge pull request #41511 from BerriAI/litellm_foundry_a2a_entra_agents
feat(a2a): reach Microsoft Foundry agents with Entra auth and versioned card discovery
This commit is contained in:
commit
88c9dd1294
20 changed files with 1309 additions and 293 deletions
|
|
@ -6,9 +6,10 @@ Extends the A2A SDK's card resolver to support multiple well-known paths.
|
|||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, runtime_checkable
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError
|
||||
from litellm.constants import LOCALHOST_URL_PATTERNS
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -18,6 +19,8 @@ if TYPE_CHECKING:
|
|||
_A2ACardResolver: Any = None
|
||||
AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent-card.json"
|
||||
PREV_AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent.json"
|
||||
FOUNDRY_AGENT_CARD_PATH: Final = "/agentCard/v1.0"
|
||||
AGENT_CARD_PATH_PARAM: Final = "agent_card_path"
|
||||
|
||||
try:
|
||||
from a2a.client import A2ACardResolver as _A2ACardResolver
|
||||
|
|
@ -29,6 +32,20 @@ except ImportError:
|
|||
pass
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _HasStatusCode(Protocol):
|
||||
status_code: int | None
|
||||
|
||||
|
||||
def _discovery_status_code(failures: tuple[tuple[str, Exception], ...]) -> int:
|
||||
statuses: Final = tuple(
|
||||
error.status_code
|
||||
for _, error in failures
|
||||
if isinstance(error, _HasStatusCode) and error.status_code is not None and error.status_code != 404
|
||||
)
|
||||
return statuses[0] if statuses else 404
|
||||
|
||||
|
||||
def is_localhost_or_internal_url(url: str | None) -> bool:
|
||||
"""
|
||||
Check if a URL is a localhost or internal URL.
|
||||
|
|
@ -145,9 +162,10 @@ class LiteLLMA2ACardResolver(_A2ACardResolver):
|
|||
"""
|
||||
Custom A2A card resolver that supports multiple well-known paths.
|
||||
|
||||
Extends the base A2ACardResolver to try both:
|
||||
Extends the base A2ACardResolver to try, in order:
|
||||
- /.well-known/agent-card.json (standard)
|
||||
- /.well-known/agent.json (previous/alternative)
|
||||
- /agentCard/v1.0
|
||||
"""
|
||||
|
||||
async def get_agent_card(
|
||||
|
|
@ -155,51 +173,37 @@ class LiteLLMA2ACardResolver(_A2ACardResolver):
|
|||
relative_card_path: str | None = None,
|
||||
http_kwargs: Mapping[str, object] | None = None,
|
||||
) -> "AgentCard":
|
||||
"""
|
||||
Fetch the agent card, trying multiple well-known paths.
|
||||
|
||||
First tries the standard path, then falls back to the previous path.
|
||||
|
||||
Args:
|
||||
relative_card_path: Optional path to the agent card endpoint.
|
||||
If None, tries both well-known paths.
|
||||
http_kwargs: Optional dictionary of keyword arguments to pass to httpx.get
|
||||
|
||||
Returns:
|
||||
AgentCard from the A2A agent
|
||||
|
||||
Raises:
|
||||
A2AClientHTTPError or A2AClientJSONError if both paths fail
|
||||
"""
|
||||
# If a specific path is provided, use the parent implementation
|
||||
"""Fetch the agent card, probing every known path when none is given."""
|
||||
if relative_card_path is not None:
|
||||
return await super().get_agent_card(
|
||||
relative_card_path=relative_card_path,
|
||||
http_kwargs=http_kwargs,
|
||||
)
|
||||
|
||||
# Try both well-known paths
|
||||
paths: Final = [
|
||||
AGENT_CARD_WELL_KNOWN_PATH,
|
||||
PREV_AGENT_CARD_WELL_KNOWN_PATH,
|
||||
]
|
||||
return await self._get_agent_card_from_first_reachable_path(
|
||||
paths=(AGENT_CARD_WELL_KNOWN_PATH, PREV_AGENT_CARD_WELL_KNOWN_PATH, FOUNDRY_AGENT_CARD_PATH),
|
||||
http_kwargs=http_kwargs,
|
||||
failures=(),
|
||||
)
|
||||
|
||||
last_error = None
|
||||
for path in paths:
|
||||
try:
|
||||
verbose_logger.debug("Attempting to fetch agent card from %s%s", self.base_url, path)
|
||||
return await super().get_agent_card(
|
||||
relative_card_path=path,
|
||||
http_kwargs=http_kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Failed to fetch agent card from %s%s: %s", self.base_url, path, e)
|
||||
last_error = e
|
||||
continue
|
||||
|
||||
# If we get here, all paths failed - re-raise the last error
|
||||
if last_error is not None:
|
||||
raise last_error
|
||||
|
||||
# This shouldn't happen, but just in case
|
||||
raise Exception(f"Failed to fetch agent card from {self.base_url}. Tried paths: {', '.join(paths)}")
|
||||
async def _get_agent_card_from_first_reachable_path(
|
||||
self,
|
||||
paths: tuple[str, ...],
|
||||
http_kwargs: Mapping[str, object] | None,
|
||||
failures: tuple[tuple[str, Exception], ...],
|
||||
) -> "AgentCard":
|
||||
if not paths:
|
||||
raise A2AAgentCardDiscoveryError(
|
||||
base_url=self.base_url,
|
||||
failures=failures,
|
||||
status_code=_discovery_status_code(failures),
|
||||
)
|
||||
path: Final = paths[0]
|
||||
try:
|
||||
verbose_logger.debug("Attempting to fetch agent card from %s%s", self.base_url, path)
|
||||
return await super().get_agent_card(relative_card_path=path, http_kwargs=http_kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Failed to fetch agent card from %s%s: %s", self.base_url, path, e)
|
||||
return await self._get_agent_card_from_first_reachable_path(
|
||||
paths=paths[1:], http_kwargs=http_kwargs, failures=(*failures, (path, e))
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ A2A Protocol Exceptions.
|
|||
Custom exception types for A2A protocol operations, following LiteLLM's exception pattern.
|
||||
"""
|
||||
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
|
|
@ -100,11 +102,12 @@ class A2AAgentCardError(A2AError):
|
|||
model: str | None = None,
|
||||
response: httpx.Response | None = None,
|
||||
litellm_debug_info: str | None = None,
|
||||
status_code: int = 404,
|
||||
):
|
||||
self.url = url
|
||||
super().__init__(
|
||||
message=message,
|
||||
status_code=404,
|
||||
status_code=status_code,
|
||||
llm_provider="a2a_agent",
|
||||
model=model,
|
||||
response=response,
|
||||
|
|
@ -112,6 +115,17 @@ class A2AAgentCardError(A2AError):
|
|||
)
|
||||
|
||||
|
||||
class A2AAgentCardDiscoveryError(A2AAgentCardError):
|
||||
def __init__(self, base_url: str, failures: tuple[tuple[str, Exception], ...], status_code: int) -> None:
|
||||
self.failures = failures
|
||||
attempts: Final = ", ".join(f"{path} ({error})" for path, error in failures)
|
||||
super().__init__(
|
||||
message=f"Failed to fetch agent card from {base_url}. Tried {attempts}",
|
||||
url=base_url,
|
||||
status_code=status_code,
|
||||
)
|
||||
|
||||
|
||||
class A2ALocalhostURLError(A2AConnectionError):
|
||||
"""
|
||||
Raised when an agent card contains a localhost/internal URL.
|
||||
|
|
|
|||
|
|
@ -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,12 @@ A2A Protocol Transformation for LiteLLM
|
|||
"""
|
||||
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from collections.abc import Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.azure_ai.common_utils import AZURE_ENTRA_LITELLM_PARAM_KEYS, get_azure_ai_agent_entra_token
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -15,6 +16,7 @@ from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
|||
|
||||
from ..common_utils import (
|
||||
A2AError,
|
||||
a2a_hop_uses_entra,
|
||||
convert_messages_to_prompt,
|
||||
extract_text_from_a2a_response,
|
||||
)
|
||||
|
|
@ -26,6 +28,39 @@ if TYPE_CHECKING:
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
_REGISTRY_PARAMS_KEPT_OUT_OF_OPTIONAL_PARAMS: Final = (
|
||||
frozenset({"api_key", "api_base", "headers", "model"}) | AZURE_ENTRA_LITELLM_PARAM_KEYS
|
||||
)
|
||||
|
||||
|
||||
def _card_declares_no_streaming(agent_card_params: Mapping[str, object]) -> bool:
|
||||
capabilities: Final = agent_card_params.get("capabilities")
|
||||
return isinstance(capabilities, Mapping) and not capabilities.get("streaming")
|
||||
|
||||
|
||||
def _agent_authenticates_with_entra(agent_litellm_params: Mapping[str, object]) -> bool:
|
||||
return a2a_hop_uses_entra(agent_litellm_params, agent_litellm_params.get("custom_llm_provider"))
|
||||
|
||||
|
||||
def _registry_api_key(agent_litellm_params: Mapping[str, object]) -> str | None:
|
||||
if _agent_authenticates_with_entra(agent_litellm_params):
|
||||
return get_azure_ai_agent_entra_token(agent_litellm_params)
|
||||
configured_api_key: Final = agent_litellm_params.get("api_key")
|
||||
return configured_api_key if isinstance(configured_api_key, str) else None
|
||||
|
||||
|
||||
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, Any] | None:
|
||||
stored_headers: Final = agent_litellm_params.get("headers")
|
||||
if not isinstance(stored_headers, Mapping):
|
||||
return None
|
||||
entra_owns_authorization: Final = _agent_authenticates_with_entra(agent_litellm_params)
|
||||
return { # mutable-ok: completion() and httpx take the request headers as a dict
|
||||
name: value
|
||||
for name, value in stored_headers.items()
|
||||
if not (entra_owns_authorization and str(name).lower() == "authorization")
|
||||
}
|
||||
|
||||
|
||||
class A2AConfig(BaseConfig):
|
||||
"""
|
||||
Configuration for A2A (Agent-to-Agent) Protocol.
|
||||
|
|
@ -35,20 +70,19 @@ class A2AConfig(BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def resolve_agent_config_from_registry(
|
||||
model: str,
|
||||
agent_name: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
headers: dict[str, Any] | None,
|
||||
optional_params: dict[str, Any],
|
||||
) -> tuple[str | None, str | None, dict[str, Any] | None]:
|
||||
"""
|
||||
Resolve agent configuration from registry if model format is "a2a/<agent-name>".
|
||||
|
||||
Extracts agent name from model string and looks up configuration in the
|
||||
agent registry (if available in proxy context).
|
||||
Resolve agent configuration from the registry for a registered agent.
|
||||
|
||||
Args:
|
||||
model: Model string (e.g., "a2a/my-agent")
|
||||
agent_name: The model string with the provider prefix already stripped by
|
||||
get_llm_provider ("a2a/my-agent" -> "my-agent"), the name the agent was
|
||||
registered under
|
||||
api_base: Explicit api_base (takes precedence over registry)
|
||||
api_key: Explicit api_key (takes precedence over registry)
|
||||
headers: Explicit headers (takes precedence over registry)
|
||||
|
|
@ -57,11 +91,7 @@ class A2AConfig(BaseConfig):
|
|||
Returns:
|
||||
Tuple of (api_base, api_key, headers) with registry values filled in
|
||||
"""
|
||||
# Extract agent name from model (e.g., "a2a/my-agent" -> "my-agent")
|
||||
agent_name: Final = model.split("/", 1)[1] if "/" in model else None
|
||||
|
||||
# Only lookup if agent name exists and some config is missing
|
||||
if not agent_name or (api_base is not None and api_key is not None and headers is not None):
|
||||
if not agent_name or (api_base is not None and api_key is not None and headers):
|
||||
return api_base, api_key, headers
|
||||
|
||||
# Try registry lookup (only available in proxy context)
|
||||
|
|
@ -79,17 +109,23 @@ class A2AConfig(BaseConfig):
|
|||
# Get api_key, headers, and other params from litellm_params
|
||||
if agent.litellm_params:
|
||||
if api_key is None:
|
||||
api_key = agent.litellm_params.get("api_key")
|
||||
api_key = _registry_api_key(agent.litellm_params)
|
||||
|
||||
if headers is None:
|
||||
agent_headers: Final = agent.litellm_params.get("headers")
|
||||
if agent_headers:
|
||||
headers = agent_headers
|
||||
if not headers:
|
||||
headers = _registry_headers(agent.litellm_params) or headers
|
||||
|
||||
# Merge other litellm_params (timeout, max_retries, etc.)
|
||||
for key, value in agent.litellm_params.items():
|
||||
if key not in ["api_key", "api_base", "headers", "model"] and key not in optional_params:
|
||||
optional_params[key] = value
|
||||
# Merge other litellm_params (timeout, max_retries, etc.)
|
||||
registry_params: Final = tuple(
|
||||
(key, value)
|
||||
for key, value in (agent.litellm_params.items() if agent.litellm_params else ())
|
||||
if key not in _REGISTRY_PARAMS_KEPT_OUT_OF_OPTIONAL_PARAMS and key not in optional_params
|
||||
)
|
||||
streaming_fallback: Final = (
|
||||
(("stream", False), ("fake_stream", True))
|
||||
if optional_params.get("stream") and _card_declares_no_streaming(agent.agent_card_params)
|
||||
else ()
|
||||
)
|
||||
optional_params.update((*registry_params, *streaming_fallback))
|
||||
except ImportError:
|
||||
pass # Registry not available (not running in proxy context)
|
||||
|
||||
|
|
@ -147,17 +183,13 @@ class A2AConfig(BaseConfig):
|
|||
api_base: API base URL
|
||||
|
||||
Returns:
|
||||
Updated headers dict
|
||||
A new headers dict; the caller's dict is left untouched
|
||||
"""
|
||||
# Ensure Content-Type is set to application/json for JSON-RPC 2.0
|
||||
if "content-type" not in headers and "Content-Type" not in headers:
|
||||
headers["Content-Type"] = "application/json"
|
||||
|
||||
# Add Authorization header if API key is provided
|
||||
if api_key is not None:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
return headers
|
||||
content_type_default: Final = (
|
||||
() if "content-type" in headers or "Content-Type" in headers else (("Content-Type", "application/json"),)
|
||||
)
|
||||
bearer: Final = () if api_key is None else (("Authorization", f"Bearer {api_key}"),)
|
||||
return dict((*headers.items(), *content_type_default, *bearer))
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
|
|
@ -226,6 +258,7 @@ class A2AConfig(BaseConfig):
|
|||
|
||||
# Create single A2A message with full conversation context
|
||||
a2a_message: Final = {
|
||||
"kind": "message",
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": full_context}],
|
||||
"messageId": str(uuid.uuid4()),
|
||||
|
|
@ -237,11 +270,14 @@ class A2AConfig(BaseConfig):
|
|||
stream: Final = optional_params.get("stream", False)
|
||||
method: Final = "message/stream" if stream else "message/send"
|
||||
|
||||
params: Final = (
|
||||
{"message": a2a_message} if stream else {"message": a2a_message, "configuration": {"blocking": True}}
|
||||
)
|
||||
request_data: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"method": method,
|
||||
"params": {"message": a2a_message},
|
||||
"params": params,
|
||||
}
|
||||
|
||||
return request_data
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Common utilities for A2A (Agent-to-Agent) Protocol
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Any, Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -10,6 +10,7 @@ from pydantic import BaseModel
|
|||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
from litellm.llms.azure_ai.common_utils import has_azure_entra_params, resolve_azure_ai_agent_auth_header
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
|
@ -142,3 +143,21 @@ def extract_text_from_a2a_response(response_dict: Mapping[str, object], max_dept
|
|||
return extract_text_from_a2a_message(first_artifact, depth=0, max_depth=max_depth)
|
||||
|
||||
return ""
|
||||
|
||||
|
||||
AgentAuthHeaderResolver = Callable[[Mapping[str, object]], Awaitable[Mapping[str, str]]]
|
||||
|
||||
|
||||
def a2a_hop_uses_entra(litellm_params: Mapping[str, object], custom_llm_provider: object) -> bool:
|
||||
return not custom_llm_provider and has_azure_entra_params(litellm_params)
|
||||
|
||||
|
||||
async def resolve_a2a_hop_auth_header(
|
||||
litellm_params: Mapping[str, object],
|
||||
custom_llm_provider: object,
|
||||
resolve_entra_header: AgentAuthHeaderResolver = resolve_azure_ai_agent_auth_header,
|
||||
) -> Mapping[str, str] | None:
|
||||
"""Entra credentials authenticate the A2A hop only; a completion-bridge agent hands them to the model provider it bridges to."""
|
||||
if not a2a_hop_uses_entra(litellm_params, custom_llm_provider):
|
||||
return None
|
||||
return await resolve_entra_header(litellm_params)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
|
@ -44,6 +46,70 @@ def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None)
|
|||
return get_azure_ad_token(params)
|
||||
|
||||
|
||||
AZURE_AI_AGENTS_SCOPE: Final = "https://ai.azure.com/.default"
|
||||
AZURE_ENTRA_CREDENTIAL_PARAM_KEYS: Final = frozenset({"azure_ad_token", "client_secret", "azure_password"})
|
||||
AZURE_ENTRA_LITELLM_PARAM_KEYS: Final = AZURE_ENTRA_CREDENTIAL_PARAM_KEYS | frozenset(
|
||||
{"tenant_id", "client_id", "azure_username", "azure_scope"}
|
||||
)
|
||||
AZURE_ENTRA_CREDENTIAL_HELP: Final = (
|
||||
"Set `tenant_id` + `client_id` + `client_secret`, `azure_ad_token` (an `oidc/` token also needs "
|
||||
"`tenant_id` + `client_id`), or `client_id` + `azure_username` + `azure_password` in the agent's `litellm_params`"
|
||||
)
|
||||
|
||||
|
||||
def has_azure_entra_params(litellm_params: Mapping[str, object] | None) -> bool:
|
||||
if not litellm_params:
|
||||
return False
|
||||
return any(litellm_params.get(key) for key in AZURE_ENTRA_CREDENTIAL_PARAM_KEYS)
|
||||
|
||||
|
||||
def _resolve_config_secret(value: object) -> str | None:
|
||||
if not isinstance(value, str) or not value:
|
||||
return None
|
||||
return get_secret_str(value) if value.startswith("os.environ/") else value
|
||||
|
||||
|
||||
def get_azure_ai_agent_entra_token(litellm_params: Mapping[str, object]) -> str:
|
||||
"""Mints the Entra bearer from the agent's own litellm_params, never from process-wide AZURE_* env vars."""
|
||||
from litellm.llms.azure.common_utils import (
|
||||
get_azure_ad_token_from_entra_id,
|
||||
get_azure_ad_token_from_oidc,
|
||||
get_azure_ad_token_from_username_password,
|
||||
)
|
||||
|
||||
resolved: Final = MappingProxyType(
|
||||
{key: _resolve_config_secret(litellm_params.get(key)) for key in AZURE_ENTRA_LITELLM_PARAM_KEYS}
|
||||
)
|
||||
scope: Final = resolved["azure_scope"] or AZURE_AI_AGENTS_SCOPE
|
||||
tenant_id: Final = resolved["tenant_id"]
|
||||
client_id: Final = resolved["client_id"]
|
||||
client_secret: Final = resolved["client_secret"]
|
||||
azure_username: Final = resolved["azure_username"]
|
||||
azure_password: Final = resolved["azure_password"]
|
||||
azure_ad_token: Final = resolved["azure_ad_token"]
|
||||
if tenant_id and client_id and client_secret:
|
||||
return get_azure_ad_token_from_entra_id(
|
||||
tenant_id=tenant_id, client_id=client_id, client_secret=client_secret, scope=scope
|
||||
)()
|
||||
if client_id and azure_username and azure_password:
|
||||
return get_azure_ad_token_from_username_password(
|
||||
client_id=client_id, azure_username=azure_username, azure_password=azure_password, scope=scope
|
||||
)()
|
||||
federated: Final = azure_ad_token is not None and azure_ad_token.startswith("oidc/")
|
||||
if azure_ad_token and federated and tenant_id and client_id:
|
||||
return get_azure_ad_token_from_oidc(
|
||||
azure_ad_token=azure_ad_token, azure_client_id=client_id, azure_tenant_id=tenant_id, scope=scope
|
||||
)
|
||||
if azure_ad_token and not federated:
|
||||
return azure_ad_token
|
||||
raise ValueError(f"Azure AI agent Entra ID credentials did not resolve to a token. {AZURE_ENTRA_CREDENTIAL_HELP}")
|
||||
|
||||
|
||||
async def resolve_azure_ai_agent_auth_header(litellm_params: Mapping[str, object]) -> Mapping[str, str]:
|
||||
token: Final = await asyncio.to_thread(get_azure_ai_agent_entra_token, litellm_params)
|
||||
return MappingProxyType({"Authorization": f"Bearer {token}"})
|
||||
|
||||
|
||||
def get_azure_ai_auth_headers(
|
||||
api_key: str | None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
|
|
|
|||
|
|
@ -2193,7 +2193,7 @@ def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
|
|||
api_key,
|
||||
headers,
|
||||
) = litellm.A2AConfig.resolve_agent_config_from_registry(
|
||||
model=model,
|
||||
agent_name=model,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
headers=headers,
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from pydantic import ValidationError
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
from litellm.llms.a2a.common_utils import resolve_a2a_hop_auth_header
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.a2a.version_convert import (
|
||||
A2AVersion,
|
||||
|
|
@ -157,19 +158,31 @@ def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str,
|
|||
)
|
||||
|
||||
|
||||
async def _resolve_backend_auth_header(
|
||||
litellm_params: dict[str, object],
|
||||
custom_llm_provider: object,
|
||||
) -> Mapping[str, str] | None:
|
||||
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
|
||||
return await resolve_databricks_app_auth_header(litellm_params)
|
||||
return await resolve_a2a_hop_auth_header(litellm_params, custom_llm_provider)
|
||||
|
||||
|
||||
def _forwarding_headers(
|
||||
caller_identity: Mapping[str, str],
|
||||
request_data: Mapping[str, object],
|
||||
agent_extra_headers: Mapping[str, str] | None,
|
||||
backend_auth_header: Mapping[str, str] | None,
|
||||
) -> dict[str, str] | None:
|
||||
backend_auth: Final = tuple(backend_auth_header.items()) if backend_auth_header else ()
|
||||
minted_names: Final = frozenset(name.lower() for name, _ in backend_auth)
|
||||
passthrough: Final = tuple(
|
||||
(name, value)
|
||||
for name, value in (agent_extra_headers.items() if agent_extra_headers else ())
|
||||
if not name.lower().startswith("x-litellm-")
|
||||
if not name.lower().startswith("x-litellm-") and name.lower() not in minted_names
|
||||
)
|
||||
trace_id: Final = request_data.get("litellm_trace_id")
|
||||
trace: Final = (("X-LiteLLM-Trace-Id", str(trace_id)),) if trace_id else ()
|
||||
merged: Final = dict((*passthrough, *caller_identity.items(), *trace))
|
||||
merged: Final = dict((*passthrough, *caller_identity.items(), *trace, *backend_auth))
|
||||
return merged or None
|
||||
|
||||
|
||||
|
|
@ -795,26 +808,16 @@ async def invoke_agent_a2a(
|
|||
if header_name:
|
||||
dynamic_headers[header_name] = val
|
||||
|
||||
agent_extra_headers = _forwarding_headers(
|
||||
agent_extra_headers: Final = _forwarding_headers(
|
||||
caller_identity=caller_identity,
|
||||
request_data=data,
|
||||
agent_extra_headers=merge_agent_headers(
|
||||
dynamic_headers=dynamic_headers or None,
|
||||
static_headers=static_headers or None,
|
||||
),
|
||||
backend_auth_header=await _resolve_backend_auth_header(litellm_params, custom_llm_provider),
|
||||
)
|
||||
|
||||
# Databricks App endpoints require a short-lived OAuth M2M token rather
|
||||
# than a static bearer. Only agents explicitly configured with a
|
||||
# ``databricks_oauth`` block get one; every other agent is left untouched.
|
||||
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
|
||||
databricks_auth: Final = await resolve_databricks_app_auth_header(litellm_params)
|
||||
if databricks_auth:
|
||||
agent_extra_headers = {
|
||||
**(agent_extra_headers or {}),
|
||||
**databricks_auth,
|
||||
}
|
||||
|
||||
# Merge agent-level guardrails into data so post_call_success_hook and
|
||||
# _handle_stream_message both pick them up. A2A agents use model
|
||||
# a2a_agent/*, which is not an llm_router deployment, so
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ def mock_a2a_client(monkeypatch):
|
|||
import litellm.a2a_protocol.main as a2a_main
|
||||
|
||||
async def _fake_create_a2a_client(
|
||||
base_url, timeout=60.0, extra_headers=None, streaming=False
|
||||
base_url, timeout=60.0, extra_headers=None, streaming=False, relative_card_path=None
|
||||
):
|
||||
return MockA2AClient()
|
||||
|
||||
|
|
|
|||
|
|
@ -5,8 +5,10 @@ Tests that the card resolver tries both old and new well-known paths.
|
|||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Final
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.a2a_protocol.card_resolver import (
|
||||
|
|
@ -16,6 +18,7 @@ from litellm.a2a_protocol.card_resolver import (
|
|||
normalize_agent_card_interfaces,
|
||||
set_agent_card_url,
|
||||
)
|
||||
from litellm.a2a_protocol.exceptions import A2AAgentCardDiscoveryError
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -138,3 +141,109 @@ def test_normalize_agent_card_interfaces_downgrades_miscased_interfaces_to_the_0
|
|||
]
|
||||
assert card.supported_interfaces[0].protocol_binding == "jsonrpc"
|
||||
assert card.supported_interfaces[0].protocol_version == "1.0"
|
||||
|
||||
|
||||
_FOUNDRY_BASE_URL: Final = "https://foundry.example.com/a2a"
|
||||
|
||||
_FOUNDRY_CARD_JSON: Final = {
|
||||
"name": "Foundry Agent",
|
||||
"description": "A test agent",
|
||||
"url": "https://foundry.example.com/a2a",
|
||||
"version": "1.0",
|
||||
"capabilities": {"streaming": True},
|
||||
"defaultInputModes": ["text"],
|
||||
"defaultOutputModes": ["text"],
|
||||
"skills": [{"id": "chat", "name": "chat", "description": "Chat", "tags": ["chat"]}],
|
||||
"protocolVersion": "1.0",
|
||||
}
|
||||
|
||||
|
||||
class _FakeHttpxClient:
|
||||
"""Answers GETs from a path -> (status, body) map and records the path of each call."""
|
||||
|
||||
def __init__(self, base_url: str, responses: dict[str, tuple[int, dict[str, Any]]]) -> None:
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._responses = responses
|
||||
self.calls: list[str] = []
|
||||
|
||||
async def get(self, url: str, **kwargs: Any) -> httpx.Response:
|
||||
path: Final = url.removeprefix(self._base_url)
|
||||
self.calls.append(path)
|
||||
status_code, body = self._responses[path]
|
||||
return httpx.Response(status_code, json=body, request=httpx.Request("GET", url))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_falls_through_to_the_foundry_card_path():
|
||||
httpx_client = _FakeHttpxClient(
|
||||
base_url=_FOUNDRY_BASE_URL,
|
||||
responses={
|
||||
"/.well-known/agent-card.json": (404, {"error": "not found"}),
|
||||
"/.well-known/agent.json": (404, {"error": "not found"}),
|
||||
"/agentCard/v1.0": (200, dict(_FOUNDRY_CARD_JSON)),
|
||||
},
|
||||
)
|
||||
|
||||
resolver = LiteLLMA2ACardResolver(httpx_client=httpx_client, base_url=_FOUNDRY_BASE_URL)
|
||||
result = await resolver.get_agent_card()
|
||||
|
||||
assert httpx_client.calls == ["/.well-known/agent-card.json", "/.well-known/agent.json", "/agentCard/v1.0"]
|
||||
assert result.name == "Foundry Agent"
|
||||
assert result.supported_interfaces[0].url == "https://foundry.example.com/a2a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_explicit_path_skips_the_probes():
|
||||
httpx_client = _FakeHttpxClient(
|
||||
base_url=_FOUNDRY_BASE_URL,
|
||||
responses={"/agentCard/v1.0": (200, dict(_FOUNDRY_CARD_JSON))},
|
||||
)
|
||||
|
||||
resolver = LiteLLMA2ACardResolver(httpx_client=httpx_client, base_url=_FOUNDRY_BASE_URL)
|
||||
result = await resolver.get_agent_card(relative_card_path="agentCard/v1.0")
|
||||
|
||||
assert httpx_client.calls == ["/agentCard/v1.0"]
|
||||
assert result.name == "Foundry Agent"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_names_every_probed_path_when_discovery_fails():
|
||||
httpx_client = _FakeHttpxClient(
|
||||
base_url=_FOUNDRY_BASE_URL,
|
||||
responses={
|
||||
"/.well-known/agent-card.json": (404, {"error": "not found"}),
|
||||
"/.well-known/agent.json": (401, {"error": "unauthorized"}),
|
||||
"/agentCard/v1.0": (404, {"error": "not found"}),
|
||||
},
|
||||
)
|
||||
|
||||
resolver = LiteLLMA2ACardResolver(httpx_client=httpx_client, base_url=_FOUNDRY_BASE_URL)
|
||||
with pytest.raises(A2AAgentCardDiscoveryError) as raised:
|
||||
await resolver.get_agent_card()
|
||||
|
||||
assert raised.value.status_code == 401
|
||||
message = str(raised.value)
|
||||
assert _FOUNDRY_BASE_URL in message
|
||||
assert "/.well-known/agent-card.json (" in message and "HTTP 404" in message
|
||||
assert "/.well-known/agent.json (" in message and "HTTP 401" in message
|
||||
assert "/agentCard/v1.0 (" in message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_card_resolver_discovery_error_is_404_when_every_probe_is_404():
|
||||
resolver = LiteLLMA2ACardResolver(
|
||||
httpx_client=_FakeHttpxClient(
|
||||
base_url=_FOUNDRY_BASE_URL,
|
||||
responses={
|
||||
"/.well-known/agent-card.json": (404, {"error": "not found"}),
|
||||
"/.well-known/agent.json": (404, {"error": "not found"}),
|
||||
"/agentCard/v1.0": (404, {"error": "not found"}),
|
||||
},
|
||||
),
|
||||
base_url=_FOUNDRY_BASE_URL,
|
||||
)
|
||||
|
||||
with pytest.raises(A2AAgentCardDiscoveryError) as raised:
|
||||
await resolver.get_agent_card()
|
||||
|
||||
assert raised.value.status_code == 404
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
52
tests/test_litellm/llms/a2a/test_common_utils.py
Normal file
52
tests/test_litellm/llms/a2a/test_common_utils.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
"""Tests for litellm/llms/a2a/common_utils.py."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.a2a.common_utils import resolve_a2a_hop_auth_header
|
||||
|
||||
|
||||
class _RecordingEntraResolver:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[Mapping[str, object]] = []
|
||||
|
||||
async def __call__(self, litellm_params: Mapping[str, object]) -> Mapping[str, str]:
|
||||
self.calls.append(litellm_params)
|
||||
return MappingProxyType({"Authorization": "Bearer minted-entra-token"})
|
||||
|
||||
|
||||
_SERVICE_PRINCIPAL = MappingProxyType({"tenant_id": "tenant", "client_id": "client", "client_secret": "sp-secret"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_entra_agent_gets_a_minted_bearer_for_the_a2a_hop():
|
||||
resolver = _RecordingEntraResolver()
|
||||
|
||||
header = await resolve_a2a_hop_auth_header(_SERVICE_PRINCIPAL, None, resolver)
|
||||
|
||||
assert header == {"Authorization": "Bearer minted-entra-token"}
|
||||
assert resolver.calls == [_SERVICE_PRINCIPAL]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_completion_bridge_agent_keeps_its_entra_credentials_for_the_model_provider():
|
||||
"""A bridged agent's tenant_id/client_id/client_secret authenticate the model it bridges to, so the A2A hop
|
||||
must not spend them on a bearer of its own."""
|
||||
resolver = _RecordingEntraResolver()
|
||||
|
||||
header = await resolve_a2a_hop_auth_header(_SERVICE_PRINCIPAL, "azure_ai", resolver)
|
||||
|
||||
assert header is None
|
||||
assert resolver.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_without_entra_credentials_gets_no_bearer():
|
||||
resolver = _RecordingEntraResolver()
|
||||
|
||||
header = await resolve_a2a_hop_auth_header({"api_base": "https://agent.example.com"}, None, resolver)
|
||||
|
||||
assert header is None
|
||||
assert resolver.calls == []
|
||||
|
|
@ -10,7 +10,12 @@ from unittest.mock import patch
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
|
||||
from litellm.llms.azure_ai.common_utils import (
|
||||
get_azure_ai_agent_entra_token,
|
||||
get_azure_ai_auth_headers,
|
||||
has_azure_entra_params,
|
||||
resolve_azure_ai_agent_auth_header,
|
||||
)
|
||||
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
||||
|
||||
ENTRA_PARAMS = {"azure_ad_token": "entra-token"}
|
||||
|
|
@ -152,3 +157,148 @@ def test_image_generation_still_uses_api_key_header():
|
|||
headers = mock_image_generation.call_args.kwargs["headers"]
|
||||
assert headers["api-key"] == "my-key"
|
||||
assert "Authorization" not in headers
|
||||
|
||||
|
||||
def test_agents_without_entra_credentials_are_not_treated_as_entra_agents():
|
||||
"""Only a credential-bearing field opts an agent into Entra auth: scope or identity fields alone
|
||||
must never make the proxy mint a bearer for that agent's URL."""
|
||||
assert has_azure_entra_params({"api_key": "static", "headers": {"x": "y"}}) is False
|
||||
assert has_azure_entra_params(None) is False
|
||||
assert has_azure_entra_params({"azure_scope": "https://ai.azure.com/.default"}) is False
|
||||
assert has_azure_entra_params({"tenant_id": "t", "client_id": "c"}) is False
|
||||
assert has_azure_entra_params({"azure_ad_token": "entra-token"}) is True
|
||||
assert has_azure_entra_params({"tenant_id": "t", "client_id": "c", "client_secret": "s"}) is True
|
||||
assert has_azure_entra_params({"client_id": "c", "azure_username": "u", "azure_password": "p"}) is True
|
||||
|
||||
|
||||
def test_agent_entra_token_ignores_the_process_wide_azure_credentials(monkeypatch):
|
||||
"""The azure provider's token helper falls back to AZURE_* env vars. An agent's bearer must come
|
||||
from that agent's own litellm_params only, or the host's service principal would authenticate to
|
||||
whatever URL an agent registers."""
|
||||
monkeypatch.setenv("AZURE_TENANT_ID", "host-tenant")
|
||||
monkeypatch.setenv("AZURE_CLIENT_ID", "host-client")
|
||||
monkeypatch.setenv("AZURE_CLIENT_SECRET", "host-secret")
|
||||
monkeypatch.setenv("AZURE_AD_TOKEN", "host-token")
|
||||
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch so a host-credential leak would show up as a call instead of a network round trip
|
||||
mock_entra_id.return_value = lambda: "host-sp-token"
|
||||
|
||||
with pytest.raises(ValueError, match="client_secret"):
|
||||
get_azure_ai_agent_entra_token({"azure_scope": "https://ai.azure.com/.default"})
|
||||
assert get_azure_ai_agent_entra_token({"azure_ad_token": "agent-token"}) == "agent-token"
|
||||
|
||||
mock_entra_id.assert_not_called()
|
||||
|
||||
|
||||
def test_agent_service_principal_fields_resolve_os_environ_references(monkeypatch):
|
||||
monkeypatch.setenv("FOUNDRY_AGENT_TENANT_ID", "tenant-from-env")
|
||||
monkeypatch.setenv("FOUNDRY_AGENT_CLIENT_ID", "client-from-env")
|
||||
monkeypatch.setenv("FOUNDRY_AGENT_CLIENT_SECRET", "secret-from-env")
|
||||
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to assert the resolved secret values reach the credential; live SP path proven by the PR's Azure Foundry e2e QA
|
||||
mock_entra_id.return_value = lambda: "sp-token"
|
||||
|
||||
token = get_azure_ai_agent_entra_token(
|
||||
{
|
||||
"tenant_id": "os.environ/FOUNDRY_AGENT_TENANT_ID",
|
||||
"client_id": "os.environ/FOUNDRY_AGENT_CLIENT_ID",
|
||||
"client_secret": "os.environ/FOUNDRY_AGENT_CLIENT_SECRET",
|
||||
}
|
||||
)
|
||||
|
||||
mock_entra_id.assert_called_once_with(
|
||||
tenant_id="tenant-from-env",
|
||||
client_id="client-from-env",
|
||||
client_secret="secret-from-env",
|
||||
scope="https://ai.azure.com/.default",
|
||||
)
|
||||
assert token == "sp-token"
|
||||
|
||||
|
||||
def test_agent_service_principal_wins_over_a_static_token_on_the_same_agent():
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to pin the precedence between a refreshing credential and a static token
|
||||
mock_entra_id.return_value = lambda: "sp-token"
|
||||
|
||||
token = get_azure_ai_agent_entra_token(
|
||||
{"tenant_id": "tenant", "client_id": "client", "client_secret": "secret", "azure_ad_token": "stale-token"}
|
||||
)
|
||||
|
||||
assert token == "sp-token"
|
||||
|
||||
|
||||
def test_agent_service_principal_token_defaults_to_the_foundry_agents_scope():
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to assert the scope Foundry agents require reaches the credential; live SP path proven by the PR's Azure Foundry e2e QA
|
||||
mock_entra_id.return_value = lambda: "sp-token"
|
||||
|
||||
token = get_azure_ai_agent_entra_token({"tenant_id": "tenant", "client_id": "client", "client_secret": "secret"})
|
||||
|
||||
mock_entra_id.assert_called_once_with(
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
scope="https://ai.azure.com/.default",
|
||||
)
|
||||
assert token == "sp-token"
|
||||
|
||||
|
||||
def test_agent_azure_scope_overrides_the_foundry_agents_default():
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id") as mock_entra_id: # test-quality-ok: stubs the Entra token fetch to assert an explicit azure_scope wins over the agents default; live SP path proven by the PR's Azure Foundry e2e QA
|
||||
mock_entra_id.return_value = lambda: "sp-token"
|
||||
|
||||
get_azure_ai_agent_entra_token(
|
||||
{"tenant_id": "tenant", "client_id": "client", "client_secret": "secret", "azure_scope": "custom/.default"}
|
||||
)
|
||||
|
||||
assert mock_entra_id.call_args.kwargs["scope"] == "custom/.default"
|
||||
|
||||
|
||||
def test_agent_entra_values_resolve_os_environ_references(monkeypatch):
|
||||
monkeypatch.setenv("FOUNDRY_AGENT_AD_TOKEN", "token-from-env")
|
||||
|
||||
assert get_azure_ai_agent_entra_token({"azure_ad_token": "os.environ/FOUNDRY_AGENT_AD_TOKEN"}) == "token-from-env"
|
||||
|
||||
|
||||
def test_agent_entra_token_failure_names_the_credential_fields():
|
||||
with pytest.raises(ValueError, match="client_secret"):
|
||||
get_azure_ai_agent_entra_token({"azure_scope": "https://ai.azure.com/.default"})
|
||||
|
||||
|
||||
def test_agent_oidc_token_without_agent_ids_never_borrows_the_host_identity(monkeypatch):
|
||||
"""The shared OIDC helper fills a missing client and tenant id from AZURE_CLIENT_ID and AZURE_TENANT_ID,
|
||||
which would exchange the host's federated token for the host's identity at that agent's URL."""
|
||||
monkeypatch.setenv("AZURE_TENANT_ID", "host-tenant")
|
||||
monkeypatch.setenv("AZURE_CLIENT_ID", "host-client")
|
||||
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_oidc") as mock_oidc: # test-quality-ok: stubs the OIDC exchange so a host-identity leak would show up as a call instead of a network round trip
|
||||
mock_oidc.return_value = "host-minted-token"
|
||||
|
||||
with pytest.raises(ValueError, match="oidc/"):
|
||||
get_azure_ai_agent_entra_token({"azure_ad_token": "oidc/github"})
|
||||
with pytest.raises(ValueError, match="oidc/"):
|
||||
get_azure_ai_agent_entra_token({"azure_ad_token": "oidc/github", "tenant_id": "agent-tenant"})
|
||||
|
||||
mock_oidc.assert_not_called()
|
||||
|
||||
|
||||
def test_agent_oidc_token_exchanges_with_the_agent_ids_and_scope():
|
||||
with patch("litellm.llms.azure.common_utils.get_azure_ad_token_from_oidc") as mock_oidc: # test-quality-ok: stubs the OIDC exchange to assert the agent's own ids and the Foundry scope reach it
|
||||
mock_oidc.return_value = "agent-minted-token"
|
||||
|
||||
token = get_azure_ai_agent_entra_token(
|
||||
{"azure_ad_token": "oidc/github", "tenant_id": "agent-tenant", "client_id": "agent-client"}
|
||||
)
|
||||
|
||||
assert token == "agent-minted-token"
|
||||
mock_oidc.assert_called_once_with(
|
||||
azure_ad_token="oidc/github",
|
||||
azure_client_id="agent-client",
|
||||
azure_tenant_id="agent-tenant",
|
||||
scope="https://ai.azure.com/.default",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_auth_header_is_the_entra_bearer():
|
||||
headers = await resolve_azure_ai_agent_auth_header({"azure_ad_token": "entra-token"})
|
||||
|
||||
assert headers == {"Authorization": "Bearer entra-token"}
|
||||
|
|
|
|||
|
|
@ -124,9 +124,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
|
||||
MessageSendParams = make_mock_pydantic_class("MessageSendParams")
|
||||
SendMessageRequest = make_mock_pydantic_class("SendMessageRequest")
|
||||
SendStreamingMessageRequest = make_mock_pydantic_class(
|
||||
"SendStreamingMessageRequest"
|
||||
)
|
||||
SendStreamingMessageRequest = make_mock_pydantic_class("SendStreamingMessageRequest")
|
||||
|
||||
# Create a mock module for a2a.types
|
||||
mock_a2a_types = MagicMock()
|
||||
|
|
@ -359,10 +357,9 @@ async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge():
|
|||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
assert (
|
||||
captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM)
|
||||
== mock_user_api_key_dict.api_key
|
||||
), "authenticated key hash was not forwarded to the completion bridge"
|
||||
assert captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM) == mock_user_api_key_dict.api_key, (
|
||||
"authenticated key hash was not forwarded to the completion bridge"
|
||||
)
|
||||
|
||||
|
||||
def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
|
||||
|
|
@ -376,9 +373,7 @@ def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
|
|||
return agent
|
||||
|
||||
|
||||
def _make_request_mock(
|
||||
method: str, params: Mapping[str, object], request_id: object = "req-1"
|
||||
) -> MagicMock:
|
||||
def _make_request_mock(method: str, params: Mapping[str, object], request_id: object = "req-1") -> MagicMock:
|
||||
req = MagicMock()
|
||||
req.headers = {}
|
||||
req.json = AsyncMock(
|
||||
|
|
@ -436,6 +431,7 @@ async def _invoke_message_method(
|
|||
mock_request: MagicMock,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
add_litellm_data: AddLiteLLMData | None = None,
|
||||
agent: MagicMock | None = None,
|
||||
) -> CapturedAgentCall:
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
|
|
@ -466,7 +462,7 @@ async def _invoke_message_method(
|
|||
downstream: Final = AsyncMock(side_effect=fake_asend_message if is_send else fake_stream_message)
|
||||
|
||||
with ExitStack() as stack:
|
||||
for p in _base_patches(_make_agent_mock(), add_litellm_data):
|
||||
for p in _base_patches(agent or _make_agent_mock(), add_litellm_data):
|
||||
stack.enter_context(p)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
if is_send:
|
||||
|
|
@ -515,6 +511,98 @@ async def test_message_methods_forward_caller_identity_headers(method: str):
|
|||
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "team-xyz"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
|
||||
async def test_message_methods_send_the_entra_bearer_for_azure_agents(method: str):
|
||||
"""A Microsoft Foundry agent accepts only an Entra ID bearer, so an agent registered with
|
||||
Entra credentials in litellm_params must reach the backend with that bearer on every call."""
|
||||
agent = _make_agent_mock()
|
||||
agent.litellm_params = {"azure_ad_token": "entra-token"}
|
||||
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
|
||||
|
||||
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
|
||||
|
||||
assert (captured.agent_extra_headers or {}).get("Authorization") == "Bearer entra-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
|
||||
async def test_message_methods_leave_agents_without_entra_params_unauthenticated(method: str):
|
||||
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
|
||||
|
||||
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
|
||||
|
||||
assert "Authorization" not in (captured.agent_extra_headers or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
|
||||
async def test_message_methods_leave_entra_fields_to_the_model_provider_for_bridge_agents(method: str):
|
||||
"""A completion-bridge agent's tenant_id/client_id/client_secret belong to the model provider it
|
||||
calls through litellm, so the proxy must not mint a Foundry bearer for them."""
|
||||
agent = _make_agent_mock()
|
||||
agent.litellm_params = {
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"model": "azure_ai/foundry-model",
|
||||
"tenant_id": "tenant",
|
||||
"client_id": "client",
|
||||
"client_secret": "sp-secret",
|
||||
}
|
||||
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
|
||||
|
||||
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
|
||||
|
||||
assert "Authorization" not in (captured.agent_extra_headers or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_send_reports_an_unresolvable_entra_credential_as_internal_error(monkeypatch):
|
||||
"""An agent whose Entra credential points at an unset environment variable must fail the call
|
||||
with the JSON-RPC internal error naming the credential fields, never reach the backend unauthenticated."""
|
||||
monkeypatch.delenv("LITELLM_TEST_UNSET_FOUNDRY_TOKEN", raising=False)
|
||||
agent = _make_agent_mock()
|
||||
agent.litellm_params = {"azure_ad_token": "os.environ/LITELLM_TEST_UNSET_FOUNDRY_TOKEN"}
|
||||
mock_request = _make_request_mock("message/send", _HELLO_MESSAGE_PARAMS)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
downstream = AsyncMock()
|
||||
|
||||
with ExitStack() as stack:
|
||||
for p in _base_patches(agent):
|
||||
stack.enter_context(p)
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: same proxy_logging_obj injection the sibling failure-hook tests use; the request must fail before any backend call is made
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: the observation point proving the backend is never called; the sibling send tests use the same seam
|
||||
"litellm.a2a_protocol.asend_message", new=downstream
|
||||
)
|
||||
)
|
||||
|
||||
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
|
||||
|
||||
response = await invoke_agent_a2a(
|
||||
agent_id="test-agent",
|
||||
request=mock_request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
body = json.loads(response.body.decode())
|
||||
assert response.status_code == 500
|
||||
assert body["error"]["code"] == -32603
|
||||
assert "client_secret" in body["error"]["message"]
|
||||
downstream.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
|
||||
async def test_message_methods_caller_identity_headers_cannot_be_spoofed(method: str):
|
||||
|
|
@ -528,12 +616,12 @@ async def test_message_methods_caller_identity_headers_cannot_be_spoofed(method:
|
|||
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
|
||||
|
||||
forwarded_headers = captured.agent_extra_headers or {}
|
||||
assert (
|
||||
forwarded_headers.get("X-LiteLLM-User-Id") == "real-user"
|
||||
), "authenticated user id must not be overridden by forwarded client headers"
|
||||
assert (
|
||||
forwarded_headers.get("X-LiteLLM-Team-Id") == "real-team"
|
||||
), "authenticated team id must not be overridden by forwarded client headers"
|
||||
assert forwarded_headers.get("X-LiteLLM-User-Id") == "real-user", (
|
||||
"authenticated user id must not be overridden by forwarded client headers"
|
||||
)
|
||||
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "real-team", (
|
||||
"authenticated team id must not be overridden by forwarded client headers"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -637,6 +725,47 @@ async def test_task_methods_forward_jsonrpc(method: str, params: dict):
|
|||
assert forwarded_body["method"] == method
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_task_methods_forward_the_entra_bearer_for_azure_agents():
|
||||
"""tasks/get on a Foundry agent polls the task the agent created, so the forwarded call needs
|
||||
the same Entra bearer as message/send."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
agent = _make_agent_mock()
|
||||
agent.litellm_params = {"azure_ad_token": "entra-token"}
|
||||
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
|
||||
|
||||
mock_http_response = MagicMock()
|
||||
mock_http_response.json.return_value = {"jsonrpc": "2.0", "id": "req-1", "result": {"id": "task-1"}}
|
||||
mock_http_response.is_success = True
|
||||
mock_http_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_handler = MagicMock()
|
||||
mock_handler.post = AsyncMock(return_value=mock_http_response)
|
||||
mock_handler.client = MagicMock()
|
||||
|
||||
with ExitStack() as stack:
|
||||
for p in _base_patches(agent):
|
||||
stack.enter_context(p)
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: the task route builds its own httpx client; the sibling task tests capture the post through the same seam
|
||||
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=mock_handler
|
||||
)
|
||||
)
|
||||
|
||||
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
|
||||
|
||||
await invoke_agent_a2a(
|
||||
agent_id="test-agent",
|
||||
request=mock_request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1"),
|
||||
)
|
||||
|
||||
posted_headers = mock_handler.post.call_args.kwargs["headers"]
|
||||
assert posted_headers["Authorization"] == "Bearer entra-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ["tasks/get", "tasks/resubscribe"])
|
||||
async def test_task_methods_extract_litellm_params_before_forwarding(method: str):
|
||||
|
|
@ -808,9 +937,7 @@ async def test_subscribe_to_task_calls_pre_call_hook():
|
|||
yield chunk
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(
|
||||
side_effect=lambda user_api_key_dict, data, call_type: data
|
||||
)
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.async_post_call_streaming_iterator_hook = _passthrough_iterator
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -866,9 +993,7 @@ async def test_subscribe_to_task_runs_post_call_streaming_guardrail():
|
|||
inspected.append(response)
|
||||
return response
|
||||
|
||||
guardrail = _RecordingGuardrail(
|
||||
guardrail_name="record-a2a", default_on=True, event_hook="post_call"
|
||||
)
|
||||
guardrail = _RecordingGuardrail(guardrail_name="record-a2a", default_on=True, event_hook="post_call")
|
||||
|
||||
agent = _make_agent_mock()
|
||||
mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"})
|
||||
|
|
@ -918,8 +1043,7 @@ async def test_subscribe_to_task_runs_post_call_streaming_guardrail():
|
|||
pass
|
||||
|
||||
assert any("resubscribe-secret" in str(r) for r in inspected), (
|
||||
"tasks/resubscribe streamed content was not passed to the post-call "
|
||||
"streaming guardrail hook"
|
||||
"tasks/resubscribe streamed content was not passed to the post-call streaming guardrail hook"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -946,9 +1070,7 @@ async def test_task_method_failure_hook_uses_enriched_request_data():
|
|||
mock_handler.post = AsyncMock(side_effect=RuntimeError("upstream failed"))
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(
|
||||
side_effect=lambda user_api_key_dict, data, call_type: data
|
||||
)
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
with ExitStack() as stack:
|
||||
|
|
@ -984,9 +1106,7 @@ async def test_task_method_failure_hook_uses_enriched_request_data():
|
|||
|
||||
body = json.loads(response.body.decode())
|
||||
assert body["error"]["code"] == -32603
|
||||
failure_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs[
|
||||
"request_data"
|
||||
]
|
||||
failure_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs["request_data"]
|
||||
assert failure_data.get("litellm_call_id")
|
||||
assert failure_data.get("agent_id") == "test-agent"
|
||||
|
||||
|
|
@ -1015,9 +1135,7 @@ async def test_agentcore_invalid_context_id_returns_jsonrpc_invalid_params_400()
|
|||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(
|
||||
side_effect=lambda user_api_key_dict, data, call_type: data
|
||||
)
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
with ExitStack() as stack:
|
||||
|
|
@ -1129,10 +1247,7 @@ async def test_get_agent_card_uses_proxy_base_url_when_set(monkeypatch):
|
|||
|
||||
body = json.loads(response.body.decode())
|
||||
assert body["url"] == "https://litellm.example.com/a2a/test-agent"
|
||||
assert (
|
||||
body["supportedInterfaces"][0]["url"]
|
||||
== "https://litellm.example.com/a2a/test-agent"
|
||||
)
|
||||
assert body["supportedInterfaces"][0]["url"] == "https://litellm.example.com/a2a/test-agent"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1182,9 +1297,7 @@ async def test_get_agent_card_0_3_card_with_a2a_version_1_0_header():
|
|||
"url": "http://backend-agent:10001",
|
||||
"version": "1.0.0",
|
||||
"capabilities": {"streaming": True},
|
||||
"skills": [
|
||||
{"id": "s1", "name": "skill one", "description": "d", "tags": ["t"]}
|
||||
],
|
||||
"skills": [{"id": "s1", "name": "skill one", "description": "d", "tags": ["t"]}],
|
||||
"defaultInputModes": ["text"],
|
||||
"defaultOutputModes": ["text"],
|
||||
}
|
||||
|
|
@ -1207,9 +1320,7 @@ async def test_get_agent_card_0_3_card_with_a2a_version_1_0_header():
|
|||
|
||||
body = json.loads(response.body.decode())
|
||||
assert "url" not in body
|
||||
assert body["supportedInterfaces"][0]["url"] == (
|
||||
"http://localhost:4000/a2a/test-agent"
|
||||
)
|
||||
assert body["supportedInterfaces"][0]["url"] == ("http://localhost:4000/a2a/test-agent")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1278,9 +1389,7 @@ def test_build_merged_agent_card_uses_proxy_base_url_for_supported_interfaces(
|
|||
http_request=mock_request,
|
||||
)
|
||||
|
||||
assert merged["supportedInterfaces"][0]["url"] == (
|
||||
"https://litellm.example.com/a2a/jenkins_agent"
|
||||
)
|
||||
assert merged["supportedInterfaces"][0]["url"] == ("https://litellm.example.com/a2a/jenkins_agent")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1324,9 +1433,7 @@ async def test_unknown_method_returns_jsonrpc_error():
|
|||
("GetExtendedAgentCard", "agent/getAuthenticatedExtendedCard"),
|
||||
],
|
||||
)
|
||||
async def test_pascal_method_names_normalize_to_wire_format(
|
||||
pascal_method: str, expected_wire_method: str
|
||||
):
|
||||
async def test_pascal_method_names_normalize_to_wire_format(pascal_method: str, expected_wire_method: str):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
agent = _make_agent_mock()
|
||||
|
|
@ -1448,9 +1555,7 @@ async def test_handle_stream_message_rejects_invalid_params_with_32602():
|
|||
)
|
||||
assert response.media_type == "text/event-stream"
|
||||
chunks = [chunk async for chunk in response.body_iterator]
|
||||
body = "".join(
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks
|
||||
)
|
||||
body = "".join(chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks)
|
||||
assert body.startswith("data: ")
|
||||
assert body.endswith("\n\n")
|
||||
payload = json.loads(body.removeprefix("data: ").strip())
|
||||
|
|
@ -1504,10 +1609,7 @@ async def test_handle_stream_message_frames_events_as_sse():
|
|||
)
|
||||
|
||||
assert response.media_type == "text/event-stream"
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert len(chunks) == len(events)
|
||||
for chunk, event in zip(chunks, events):
|
||||
|
|
@ -1530,10 +1632,7 @@ async def test_handle_stream_message_sdk_unavailable_frames_error_as_sse():
|
|||
)
|
||||
|
||||
assert response.media_type == "text/event-stream"
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
assert len(chunks) == 1
|
||||
assert chunks[0].startswith("data: ")
|
||||
assert chunks[0].endswith("\n\n")
|
||||
|
|
@ -1569,9 +1668,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse():
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _handle_stream_message(
|
||||
api_base="http://upstream.local",
|
||||
|
|
@ -1589,10 +1686,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse():
|
|||
)
|
||||
|
||||
assert response.media_type == "text/event-stream"
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert len(chunks) == len(events)
|
||||
for chunk, event in zip(chunks, events):
|
||||
|
|
@ -1620,9 +1714,7 @@ async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once():
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _handle_stream_message(
|
||||
api_base="http://upstream.local",
|
||||
|
|
@ -1636,10 +1728,7 @@ async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once():
|
|||
},
|
||||
)
|
||||
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert len(chunks) == 1
|
||||
payload = json.loads(chunks[0].removeprefix("data: ").strip())
|
||||
|
|
@ -1661,9 +1750,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse():
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _handle_stream_message(
|
||||
api_base="http://upstream.local",
|
||||
|
|
@ -1680,10 +1767,7 @@ async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse():
|
|||
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
|
||||
)
|
||||
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert len(chunks) == 2
|
||||
assert chunks[-1].startswith("data: ")
|
||||
|
|
@ -1707,9 +1791,7 @@ async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error()
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _handle_stream_message(
|
||||
api_base="http://upstream.local",
|
||||
|
|
@ -1726,10 +1808,7 @@ async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error()
|
|||
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
|
||||
)
|
||||
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert len(chunks) == 1
|
||||
error_payload = json.loads(chunks[0].removeprefix("data: ").strip())
|
||||
|
|
@ -1749,9 +1828,7 @@ async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event():
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _handle_stream_message(
|
||||
api_base="http://upstream.local",
|
||||
|
|
@ -1765,10 +1842,7 @@ async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event():
|
|||
},
|
||||
)
|
||||
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert chunks == ['data: "not json at all"\n\n']
|
||||
|
||||
|
|
@ -1785,9 +1859,7 @@ async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error():
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _handle_stream_message(
|
||||
api_base="http://upstream.local",
|
||||
|
|
@ -1801,10 +1873,7 @@ async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error():
|
|||
},
|
||||
)
|
||||
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert len(chunks) == 2
|
||||
error_payload = json.loads(chunks[-1].removeprefix("data: ").strip())
|
||||
|
|
@ -1911,10 +1980,7 @@ def test_normalize_response_keeps_wire_format_for_0_3():
|
|||
"role": "agent",
|
||||
},
|
||||
}
|
||||
assert (
|
||||
normalize_jsonrpc_response(wire_response, "0.3", method="message/send")
|
||||
is wire_response
|
||||
)
|
||||
assert normalize_jsonrpc_response(wire_response, "0.3", method="message/send") is wire_response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1936,9 +2002,7 @@ async def test_task_method_upstream_jsonrpc_error_on_http_4xx_is_relayed():
|
|||
mock_http_response = MagicMock()
|
||||
mock_http_response.json.return_value = upstream_error
|
||||
mock_http_response.is_success = False
|
||||
mock_http_response.raise_for_status = MagicMock(
|
||||
side_effect=Exception("404 Not Found")
|
||||
)
|
||||
mock_http_response.raise_for_status = MagicMock(side_effect=Exception("404 Not Found"))
|
||||
|
||||
mock_handler = MagicMock()
|
||||
mock_handler.post = AsyncMock(return_value=mock_http_response)
|
||||
|
|
@ -1982,9 +2046,7 @@ async def test_subscribe_to_task_upstream_error_yields_jsonrpc_error_event():
|
|||
mock_resp.is_success = False
|
||||
mock_resp.status_code = 404
|
||||
mock_resp.reason_phrase = "Not Found"
|
||||
mock_resp.aread = AsyncMock(
|
||||
return_value=b'{"jsonrpc":"2.0","error":{"code":-32001,"message":"Task not found"}}'
|
||||
)
|
||||
mock_resp.aread = AsyncMock(return_value=b'{"jsonrpc":"2.0","error":{"code":-32001,"message":"Task not found"}}')
|
||||
mock_resp.aclose = AsyncMock()
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
|
|
@ -2076,9 +2138,7 @@ async def test_task_methods_forward_caller_identity_headers():
|
|||
}
|
||||
agent = _make_agent_mock()
|
||||
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-test", user_id="user-abc", team_id="team-xyz"
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="user-abc", team_id="team-xyz")
|
||||
|
||||
mock_http_response = MagicMock()
|
||||
mock_http_response.json.return_value = upstream_response
|
||||
|
|
@ -2364,9 +2424,7 @@ async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers()
|
|||
"x-a2a-test-agent-x-litellm-user-id": "attacker-user",
|
||||
"x-a2a-test-agent-x-litellm-team-id": "attacker-team",
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-test", user_id="real-user", team_id="real-team"
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="real-user", team_id="real-team")
|
||||
|
||||
mock_http_response = MagicMock()
|
||||
mock_http_response.json.return_value = upstream_response
|
||||
|
|
@ -2395,19 +2453,17 @@ async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers()
|
|||
)
|
||||
|
||||
posted_headers = mock_handler.post.call_args.kwargs.get("headers") or {}
|
||||
assert (
|
||||
posted_headers.get("X-LiteLLM-User-Id") == "real-user"
|
||||
), "authenticated user id must not be overridden by forwarded client headers"
|
||||
assert (
|
||||
posted_headers.get("X-LiteLLM-Team-Id") == "real-team"
|
||||
), "authenticated team id must not be overridden by forwarded client headers"
|
||||
assert posted_headers.get("X-LiteLLM-User-Id") == "real-user", (
|
||||
"authenticated user id must not be overridden by forwarded client headers"
|
||||
)
|
||||
assert posted_headers.get("X-LiteLLM-Team-Id") == "real-team", (
|
||||
"authenticated team id must not be overridden by forwarded client headers"
|
||||
)
|
||||
|
||||
|
||||
def _agent(protocol_version):
|
||||
agent = MagicMock()
|
||||
agent.agent_card_params = (
|
||||
{"protocolVersion": protocol_version} if protocol_version is not None else {}
|
||||
)
|
||||
agent.agent_card_params = {"protocolVersion": protocol_version} if protocol_version is not None else {}
|
||||
return agent
|
||||
|
||||
|
||||
|
|
@ -2553,16 +2609,11 @@ async def test_handle_stream_message_pings_while_the_upstream_agent_is_still_sil
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _stream_message_response()
|
||||
assert response.headers["x-accel-buffering"] == "no"
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert chunks[0] == ": ping\n\n"
|
||||
assert chunks.count(": ping\n\n") >= 3
|
||||
|
|
@ -2583,16 +2634,26 @@ async def test_handle_stream_message_is_untouched_while_keepalives_are_unconfigu
|
|||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
|
||||
stack.enter_context(
|
||||
patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream)
|
||||
)
|
||||
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
|
||||
|
||||
response = await _stream_message_response()
|
||||
assert "x-accel-buffering" not in response.headers
|
||||
chunks = [
|
||||
chunk.decode() if isinstance(chunk, bytes) else chunk
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
|
||||
|
||||
assert not any(chunk.startswith(":") for chunk in chunks)
|
||||
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
|
||||
|
||||
|
||||
def test_forwarding_headers_minted_bearer_replaces_a_forwarded_authorization_of_any_case():
|
||||
"""A client header the admin chose to forward keeps the casing the config named it with, so a forwarded
|
||||
`authorization` must not travel next to the minted `Authorization` as a second header line."""
|
||||
from litellm.proxy.agent_endpoints.a2a_endpoints import _forwarding_headers
|
||||
|
||||
merged = _forwarding_headers(
|
||||
caller_identity={},
|
||||
request_data={},
|
||||
agent_extra_headers={"authorization": "Bearer client-token", "X-Custom": "kept"},
|
||||
backend_auth_header={"Authorization": "Bearer minted-token"},
|
||||
)
|
||||
|
||||
assert merged == {"X-Custom": "kept", "Authorization": "Bearer minted-token"}
|
||||
|
|
|
|||
|
|
@ -4,8 +4,10 @@ Test A2A provider registry lookup functionality.
|
|||
Maps to: litellm/llms/a2a/chat/transformation.py
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
|
@ -15,19 +17,20 @@ from litellm.llms.a2a.chat.transformation import A2AConfig
|
|||
def test_resolve_agent_config_from_registry_static_method():
|
||||
"""Test the static helper method for registry resolution"""
|
||||
|
||||
# Test 1: No agent name in model
|
||||
# Test 1: Unregistered agent name keeps the explicit config
|
||||
api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry(
|
||||
model="a2a",
|
||||
agent_name="not-registered",
|
||||
api_base="http://test.com",
|
||||
api_key=None,
|
||||
headers=None,
|
||||
optional_params={},
|
||||
)
|
||||
assert api_base == "http://test.com"
|
||||
assert api_key is None
|
||||
|
||||
# Test 2: All params provided - should not lookup registry
|
||||
api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry(
|
||||
model="a2a/test-agent",
|
||||
agent_name="test-agent",
|
||||
api_base="http://explicit.com",
|
||||
api_key="explicit-key",
|
||||
headers={"X-Test": "value"},
|
||||
|
|
@ -38,34 +41,297 @@ def test_resolve_agent_config_from_registry_static_method():
|
|||
|
||||
|
||||
def test_a2a_registry_integration():
|
||||
"""Test registry lookup in proxy context"""
|
||||
"""A chat call for a registered agent must post to the registered url with the registered key as the
|
||||
bearer even though completion() strips the a2a/ prefix before the lookup runs."""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
test_agent = AgentResponse(
|
||||
agent_id="test-id",
|
||||
agent_name="test-agent",
|
||||
agent_card_params={"url": "http://registry-url.example.com:9999"},
|
||||
litellm_params={"api_key": "registry-key", "headers": {"X-Agent": "static"}},
|
||||
)
|
||||
client = HTTPHandler()
|
||||
agent_reply = httpx.Response(
|
||||
200,
|
||||
json={"jsonrpc": "2.0", "id": "1", "result": {"kind": "message", "parts": [{"kind": "text", "text": "4"}]}},
|
||||
)
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(test_agent)
|
||||
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
# Create test agent
|
||||
test_agent = AgentResponse(
|
||||
agent_id="test-id",
|
||||
agent_name="test-agent",
|
||||
agent_card_params={"url": "http://registry-url.example.com:9999"},
|
||||
litellm_params={"api_key": "registry-key"},
|
||||
)
|
||||
|
||||
# Register and test
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(test_agent)
|
||||
|
||||
try:
|
||||
litellm.completion(
|
||||
model="a2a/test-agent", messages=[{"role": "user", "content": "Hello"}]
|
||||
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
|
||||
response = litellm.completion(
|
||||
model="a2a/test-agent", messages=[{"role": "user", "content": "What is 2+2?"}], client=client
|
||||
)
|
||||
except Exception as e:
|
||||
# Should use registry URL (connection error expected)
|
||||
if "registry-url.example.com" not in str(e) and "APIConnectionError" not in type(e).__name__:
|
||||
raise
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
except ImportError:
|
||||
pytest.skip("Registry not available (not in proxy context)")
|
||||
assert response.choices[0].message.content == "4"
|
||||
assert post.call_args.kwargs["url"] == "http://registry-url.example.com:9999"
|
||||
assert post.call_args.kwargs["headers"]["Authorization"] == "Bearer registry-key"
|
||||
assert post.call_args.kwargs["headers"]["X-Agent"] == "static"
|
||||
|
||||
|
||||
def test_one_callers_bearer_never_reaches_another_caller_of_the_same_registered_agent():
|
||||
"""The registered headers dict is shared by every request to the agent, so the bearer one caller
|
||||
supplies must be written to that request alone and never persisted onto the agent for the next
|
||||
caller, who has no key of their own."""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
shared_agent = AgentResponse(
|
||||
agent_id="shared-id",
|
||||
agent_name="shared-agent",
|
||||
agent_card_params={"url": "http://registry-url.example.com:9999"},
|
||||
litellm_params={"headers": {"X-Agent": "static"}},
|
||||
)
|
||||
client = HTTPHandler()
|
||||
agent_reply = httpx.Response(
|
||||
200,
|
||||
json={"jsonrpc": "2.0", "id": "1", "result": {"kind": "message", "parts": [{"kind": "text", "text": "ok"}]}},
|
||||
)
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(shared_agent)
|
||||
|
||||
try:
|
||||
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
|
||||
litellm.completion(model="a2a/shared-agent", messages=messages, api_key="caller-one-key", client=client)
|
||||
litellm.completion(model="a2a/shared-agent", messages=messages, client=client)
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
first_call_headers, second_call_headers = (call.kwargs["headers"] for call in post.call_args_list)
|
||||
assert first_call_headers["Authorization"] == "Bearer caller-one-key"
|
||||
assert "Authorization" not in second_call_headers
|
||||
assert second_call_headers["X-Agent"] == "static"
|
||||
assert shared_agent.litellm_params == {"headers": {"X-Agent": "static"}}
|
||||
|
||||
|
||||
def _foundry_card_stored_through_the_agents_api() -> dict:
|
||||
from litellm.proxy.a2a.agent_card import merge_agent_card
|
||||
|
||||
return merge_agent_card(
|
||||
{"name": "Foundry", "url": "https://foundry.example.com/a2a", "capabilities": {"streaming": False}},
|
||||
proxy_url="http://localhost:4000/a2a/foundry-agent",
|
||||
proxy_base_url="http://localhost:4000",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"agent_card_params",
|
||||
[
|
||||
{"url": "https://foundry.example.com/a2a", "capabilities": {"streaming": False}},
|
||||
_foundry_card_stored_through_the_agents_api(),
|
||||
],
|
||||
ids=["card registered verbatim from config.yaml", "card stored through POST /v1/agents"],
|
||||
)
|
||||
def test_streaming_chat_to_an_agent_whose_card_declines_streaming_uses_a_blocking_send(agent_card_params: dict):
|
||||
"""Microsoft Foundry agents publish `capabilities.streaming: false` and answer message/stream with a
|
||||
JSON-RPC error. A streaming chat call to such an agent must post a blocking message/send and hand the
|
||||
caller the answer as a stream, whether the card was registered verbatim from config.yaml or stored
|
||||
through POST /v1/agents, which keeps only truthy capabilities and so drops the `false` itself."""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
foundry_agent = AgentResponse(
|
||||
agent_id="foundry-id",
|
||||
agent_name="foundry-agent",
|
||||
agent_card_params=agent_card_params,
|
||||
litellm_params={"api_key": "registry-key"},
|
||||
)
|
||||
client = HTTPHandler()
|
||||
agent_reply = httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": "1",
|
||||
"result": {
|
||||
"kind": "task",
|
||||
"status": {"state": "completed"},
|
||||
"artifacts": [{"parts": [{"kind": "text", "text": "4"}]}],
|
||||
},
|
||||
},
|
||||
)
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(foundry_agent)
|
||||
|
||||
try:
|
||||
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
|
||||
chunks = list(
|
||||
litellm.completion(
|
||||
model="a2a/foundry-agent",
|
||||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
posted = json.loads(post.call_args.kwargs["data"])
|
||||
assert posted["method"] == "message/send"
|
||||
assert posted["params"]["configuration"] == {"blocking": True}
|
||||
assert post.call_args.kwargs.get("stream", False) is False
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "4"
|
||||
assert chunks[-1].choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"agent_card_params",
|
||||
[
|
||||
{"url": "https://agent.example.com/a2a"},
|
||||
{"url": "https://agent.example.com/a2a", "capabilities": {"streaming": True}},
|
||||
],
|
||||
ids=["card without a capabilities block", "card says streaming true"],
|
||||
)
|
||||
def test_registry_lookup_leaves_streaming_alone_when_the_card_does_not_decline_it(agent_card_params: dict):
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
silent_agent = AgentResponse(
|
||||
agent_id="silent-id",
|
||||
agent_name="silent-agent",
|
||||
agent_card_params=agent_card_params,
|
||||
litellm_params={"api_key": "registry-key"},
|
||||
)
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(silent_agent)
|
||||
optional_params: dict = {"stream": True}
|
||||
|
||||
try:
|
||||
A2AConfig.resolve_agent_config_from_registry(
|
||||
agent_name="silent-agent", api_base=None, api_key=None, headers=None, optional_params=optional_params
|
||||
)
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
assert optional_params == {"stream": True}
|
||||
|
||||
|
||||
def test_registry_entra_agent_authenticates_with_the_entra_token_and_keeps_its_secrets_private():
|
||||
"""An agent registered with Entra credentials has no api_key, so the chat route must resolve the
|
||||
bearer from those credentials, and the credential fields must not ride along into optional_params
|
||||
where they would reach spend logs and callbacks."""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
entra_agent = AgentResponse(
|
||||
agent_id="entra-id",
|
||||
agent_name="entra-agent",
|
||||
agent_card_params={"url": "https://foundry.example.com/a2a"},
|
||||
litellm_params={"azure_ad_token": "entra-token", "tenant_id": "tenant", "timeout": 30},
|
||||
)
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(entra_agent)
|
||||
optional_params: dict = {}
|
||||
|
||||
try:
|
||||
api_base, api_key, _headers = A2AConfig.resolve_agent_config_from_registry(
|
||||
agent_name="entra-agent",
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
headers=None,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
assert api_base == "https://foundry.example.com/a2a"
|
||||
assert api_key == "entra-token"
|
||||
assert optional_params == {"timeout": 30}
|
||||
|
||||
|
||||
_STORED_STATIC_CREDENTIALS: dict = {
|
||||
"api_key": "stored-key",
|
||||
"headers": {"authorization": "Bearer stored-header", "X-Agent": "static"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("litellm_params", "expected_authorization_lines"),
|
||||
[
|
||||
(
|
||||
{**_STORED_STATIC_CREDENTIALS, "azure_ad_token": "entra-token"},
|
||||
{"Authorization": "Bearer entra-token"},
|
||||
),
|
||||
(
|
||||
_STORED_STATIC_CREDENTIALS,
|
||||
{"authorization": "Bearer stored-header", "Authorization": "Bearer stored-key"},
|
||||
),
|
||||
(
|
||||
{**_STORED_STATIC_CREDENTIALS, "azure_ad_token": "model-provider-token", "custom_llm_provider": "azure_ai"},
|
||||
{"authorization": "Bearer stored-header", "Authorization": "Bearer stored-key"},
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"entra agent: the minted bearer is the only authorization line",
|
||||
"agent without entra credentials: static credentials sent as before",
|
||||
"bridge agent: its entra credentials belong to the model provider, never to the a2a hop",
|
||||
],
|
||||
)
|
||||
def test_entra_credentials_beat_the_static_credentials_stored_next_to_them_on_the_chat_route(
|
||||
litellm_params: dict, expected_authorization_lines: dict
|
||||
):
|
||||
"""The relay sends the minted Entra bearer over any static Authorization stored on the agent; the chat
|
||||
route must agree, or an api_key or authorization header left next to the Entra fields makes the same
|
||||
agent answer on /a2a and fail with the backend's 401 on /v1/chat/completions."""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
agent = AgentResponse(
|
||||
agent_id="mixed-credentials-id",
|
||||
agent_name="mixed-credentials-agent",
|
||||
agent_card_params={"url": "https://foundry.example.com/a2a"},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
client = HTTPHandler()
|
||||
agent_reply = httpx.Response(
|
||||
200,
|
||||
json={"jsonrpc": "2.0", "id": "1", "result": {"kind": "message", "parts": [{"kind": "text", "text": "ok"}]}},
|
||||
)
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(agent)
|
||||
|
||||
try:
|
||||
with patch.object(client, "post", return_value=agent_reply) as post: # test-quality-ok: injected client
|
||||
litellm.completion(
|
||||
model="a2a/mixed-credentials-agent", messages=[{"role": "user", "content": "hi"}], client=client
|
||||
)
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
sent_headers = post.call_args.kwargs["headers"]
|
||||
assert {
|
||||
name: value for name, value in sent_headers.items() if name.lower() == "authorization"
|
||||
} == expected_authorization_lines
|
||||
assert sent_headers["X-Agent"] == "static"
|
||||
|
||||
|
||||
def test_registry_entra_agent_with_an_unresolvable_credential_fails_the_chat_call(monkeypatch):
|
||||
"""The chat route mints the Foundry bearer from the registered credentials; when they resolve to
|
||||
nothing the caller must get the credential error instead of an unauthenticated backend call."""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
monkeypatch.delenv("LITELLM_TEST_UNSET_FOUNDRY_TOKEN", raising=False)
|
||||
entra_agent = AgentResponse(
|
||||
agent_id="entra-unset-id",
|
||||
agent_name="entra-unset-agent",
|
||||
agent_card_params={"url": "https://foundry.example.com/a2a"},
|
||||
litellm_params={"azure_ad_token": "os.environ/LITELLM_TEST_UNSET_FOUNDRY_TOKEN"},
|
||||
)
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(entra_agent)
|
||||
|
||||
try:
|
||||
with pytest.raises(litellm.APIConnectionError, match="client_secret"):
|
||||
litellm.completion(model="a2a/entra-unset-agent", messages=[{"role": "user", "content": "hi"}])
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue