From 4370f6fb74aa2d51d32c542e0314f078b65ec75e Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 3 Dec 2025 18:54:55 -0800 Subject: [PATCH] [Feat] Agent Gateway - Allow invoking agents through AI Gateway (#17440) * init litellm A2a client * simpler a2a client interface * test a2a * move a2a invoking tests * test fix * ensure a2a send message is tracked n logs * rename tags * add streaming handlng * add a2a invocation --- .../proxy/vector_stores/endpoints.py | 337 ++++++++++++++++++ litellm/a2a/__init__.py | 59 +++ litellm/a2a/client.py | 107 ++++++ litellm/a2a/main.py | 261 ++++++++++++++ litellm/litellm_core_utils/litellm_logging.py | 2 + .../proxy/agent_endpoints/a2a_endpoints.py | 189 ++++++++++ litellm/proxy/agent_endpoints/endpoints.py | 14 +- litellm/proxy/proxy_server.py | 2 + litellm/types/agents.py | 47 ++- litellm/types/llms/custom_http.py | 1 + tests/agent_tests/test_a2a.py | 138 +++++++ 11 files changed, 1148 insertions(+), 9 deletions(-) create mode 100644 enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py create mode 100644 litellm/a2a/__init__.py create mode 100644 litellm/a2a/client.py create mode 100644 litellm/a2a/main.py create mode 100644 litellm/proxy/agent_endpoints/a2a_endpoints.py create mode 100644 tests/agent_tests/test_a2a.py diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py new file mode 100644 index 00000000000..fdb1dba372f --- /dev/null +++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py @@ -0,0 +1,337 @@ +""" +VECTOR STORE MANAGEMENT + +All /vector_store management endpoints + +/vector_store/new +/vector_store/delete +/vector_store/list +""" + +import copy +import json +from typing import List, Optional + +from fastapi import APIRouter, Depends, HTTPException + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy._types import ( + LiteLLM_ManagedVectorStoresTable, + ResponseLiteLLM_ManagedVectorStore, + UserAPIKeyAuth, +) +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.types.vector_stores import ( + LiteLLM_ManagedVectorStore, + LiteLLM_ManagedVectorStoreListResponse, + VectorStoreDeleteRequest, + VectorStoreInfoRequest, + VectorStoreUpdateRequest, +) +from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + +router = APIRouter() + + +######################################################## +# Management Endpoints +######################################################## +@router.post( + "/vector_store/new", + tags=["vector store management"], + dependencies=[Depends(user_api_key_auth)], +) +async def new_vector_store( + vector_store: LiteLLM_ManagedVectorStore, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new vector store. + + Parameters: + - vector_store_id: str - Unique identifier for the vector store + - custom_llm_provider: str - Provider of the vector store + - vector_store_name: Optional[str] - Name of the vector store + - vector_store_description: Optional[str] - Description of the vector store + - vector_store_metadata: Optional[Dict] - Additional metadata for the vector store + """ + from litellm.proxy.proxy_server import prisma_client + from litellm.types.router import GenericLiteLLMParams + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if vector store already exists + existing_vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.find_unique( + where={"vector_store_id": vector_store.get("vector_store_id")} + ) + ) + if existing_vector_store is not None: + raise HTTPException( + status_code=400, + detail=f"Vector store with ID {vector_store.get('vector_store_id')} already exists", + ) + + if vector_store.get("vector_store_metadata") is not None: + vector_store["vector_store_metadata"] = safe_dumps( + vector_store.get("vector_store_metadata") + ) + + # Safely handle JSON serialization of litellm_params + litellm_params_json: Optional[str] = None + _input_litellm_params: dict = vector_store.get("litellm_params", {}) or {} + if _input_litellm_params is not None: + litellm_params_dict = GenericLiteLLMParams( + **_input_litellm_params + ).model_dump(exclude_none=True) + litellm_params_json = safe_dumps(litellm_params_dict) + del vector_store["litellm_params"] + + _new_vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.create( + data={ + **vector_store, + "litellm_params": litellm_params_json, + } + ) + ) + + new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore( + **_new_vector_store.model_dump() + ) + + # Add vector store to registry + if litellm.vector_store_registry is not None: + litellm.vector_store_registry.add_vector_store_to_registry( + vector_store=new_vector_store + ) + + return { + "status": "success", + "message": f"Vector store {vector_store.get('vector_store_id')} created successfully", + "vector_store": new_vector_store, + } + except Exception as e: + verbose_proxy_logger.exception(f"Error creating vector store: {str(e)}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/vector_store/list", + tags=["vector store management"], + dependencies=[Depends(user_api_key_auth)], + response_model=LiteLLM_ManagedVectorStoreListResponse, +) +async def list_vector_stores( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + page: int = 1, + page_size: int = 100, +): + """ + List all available vector stores with optional filtering and pagination. + Combines both in-memory vector stores and those stored in the database. + + Parameters: + - page: int - Page number for pagination (default: 1) + - page_size: int - Number of items per page (default: 100) + """ + from litellm.proxy.proxy_server import prisma_client + + seen_vector_store_ids = set() + + try: + # Get in-memory vector stores + in_memory_vector_stores: List[LiteLLM_ManagedVectorStore] = [] + if litellm.vector_store_registry is not None: + in_memory_vector_stores = copy.deepcopy( + litellm.vector_store_registry.vector_stores + ) + + # Get vector stores from database + vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( + prisma_client=prisma_client + ) + + # Combine in-memory and database vector stores + combined_vector_stores: List[LiteLLM_ManagedVectorStore] = [] + for vector_store in in_memory_vector_stores + vector_stores_from_db: + vector_store_id = vector_store.get("vector_store_id", None) + if vector_store_id not in seen_vector_store_ids: + combined_vector_stores.append(vector_store) + seen_vector_store_ids.add(vector_store_id) + + total_count = len(combined_vector_stores) + total_pages = (total_count + page_size - 1) // page_size + + # Format response using LiteLLM_ManagedVectorStoreListResponse + response = LiteLLM_ManagedVectorStoreListResponse( + object="list", + data=combined_vector_stores, + total_count=total_count, + current_page=page, + total_pages=total_pages, + ) + + return response + except Exception as e: + verbose_proxy_logger.exception(f"Error listing vector stores: {str(e)}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/vector_store/delete", + tags=["vector store management"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_vector_store( + data: VectorStoreDeleteRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Delete a vector store. + + Parameters: + - vector_store_id: str - ID of the vector store to delete + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if vector store exists + existing_vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.find_unique( + where={"vector_store_id": data.vector_store_id} + ) + ) + if existing_vector_store is None: + raise HTTPException( + status_code=404, + detail=f"Vector store with ID {data.vector_store_id} not found", + ) + + # Delete vector store + await prisma_client.db.litellm_managedvectorstorestable.delete( + where={"vector_store_id": data.vector_store_id} + ) + + # Delete vector store from registry + if litellm.vector_store_registry is not None: + litellm.vector_store_registry.delete_vector_store_from_registry( + vector_store_id=data.vector_store_id + ) + + return {"message": f"Vector store {data.vector_store_id} deleted successfully"} + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/vector_store/info", + tags=["vector store management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ResponseLiteLLM_ManagedVectorStore, +) +async def get_vector_store_info( + data: VectorStoreInfoRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """Return a single vector store's details""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + if litellm.vector_store_registry is not None: + vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( + vector_store_id=data.vector_store_id + ) + if vector_store is not None: + vector_store_metadata = vector_store.get("vector_store_metadata") + # Parse metadata if it's a JSON string + parsed_metadata: Optional[dict] = None + if isinstance(vector_store_metadata, str): + parsed_metadata = json.loads(vector_store_metadata) + elif isinstance(vector_store_metadata, dict): + parsed_metadata = vector_store_metadata + + vector_store_pydantic_obj = LiteLLM_ManagedVectorStoresTable( + vector_store_id=vector_store.get("vector_store_id") or "", + custom_llm_provider=vector_store.get("custom_llm_provider") or "", + vector_store_name=vector_store.get("vector_store_name") or None, + vector_store_description=vector_store.get( + "vector_store_description" + ) + or None, + vector_store_metadata=parsed_metadata, + created_at=vector_store.get("created_at") or None, + updated_at=vector_store.get("updated_at") or None, + litellm_credential_name=vector_store.get("litellm_credential_name"), + litellm_params=vector_store.get("litellm_params") or None, + ) + return {"vector_store": vector_store_pydantic_obj} + + vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.find_unique( + where={"vector_store_id": data.vector_store_id} + ) + ) + if vector_store is None: + raise HTTPException( + status_code=404, + detail=f"Vector store with ID {data.vector_store_id} not found", + ) + + vector_store_dict = vector_store.model_dump() # type: ignore[attr-defined] + return {"vector_store": vector_store_dict} + except Exception as e: + verbose_proxy_logger.exception(f"Error getting vector store info: {str(e)}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/vector_store/update", + tags=["vector store management"], + dependencies=[Depends(user_api_key_auth)], +) +async def update_vector_store( + data: VectorStoreUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """Update vector store details""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + update_data = data.model_dump(exclude_unset=True) + vector_store_id = update_data.pop("vector_store_id") + if update_data.get("vector_store_metadata") is not None: + update_data["vector_store_metadata"] = safe_dumps( + update_data["vector_store_metadata"] + ) + + updated = await prisma_client.db.litellm_managedvectorstorestable.update( + where={"vector_store_id": vector_store_id}, + data=update_data, + ) + + updated_vs = LiteLLM_ManagedVectorStore(**updated.model_dump()) + + if litellm.vector_store_registry is not None: + litellm.vector_store_registry.update_vector_store_in_registry( + vector_store_id=vector_store_id, + updated_data=updated_vs, + ) + + return {"vector_store": updated_vs} + except Exception as e: + verbose_proxy_logger.exception(f"Error updating vector store: {str(e)}") + raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/a2a/__init__.py b/litellm/a2a/__init__.py new file mode 100644 index 00000000000..a2e9770cd20 --- /dev/null +++ b/litellm/a2a/__init__.py @@ -0,0 +1,59 @@ +""" +LiteLLM A2A - Wrapper for invoking A2A protocol agents. + +This module provides a thin wrapper around the official `a2a` SDK that: +- Handles httpx client creation and agent card resolution +- Adds LiteLLM logging via @client decorator +- Matches the A2A SDK interface (SendMessageRequest, SendMessageResponse, etc.) + +Example usage (standalone functions with @client decorator): + ```python + from litellm.a2a import asend_message + from a2a.types import SendMessageRequest, MessageSendParams + from uuid import uuid4 + + request = SendMessageRequest( + id=str(uuid4()), + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": "Hello!"}], + "messageId": uuid4().hex, + } + ) + ) + response = await asend_message( + base_url="http://localhost:10001", + request=request, + ) + print(response.model_dump(mode='json', exclude_none=True)) + ``` + +Example usage (class-based): + ```python + from litellm.a2a import A2AClient + + client = A2AClient(base_url="http://localhost:10001") + response = await client.send_message(request) + ``` +""" + +from litellm.a2a.client import A2AClient +from litellm.a2a.main import ( + aget_agent_card, + asend_message, + asend_message_streaming, + create_a2a_client, + send_message, +) +from litellm.types.agents import LiteLLMSendMessageResponse + +__all__ = [ + "A2AClient", + "asend_message", + "send_message", + "asend_message_streaming", + "aget_agent_card", + "create_a2a_client", + "LiteLLMSendMessageResponse", +] diff --git a/litellm/a2a/client.py b/litellm/a2a/client.py new file mode 100644 index 00000000000..0c3b8fff2ba --- /dev/null +++ b/litellm/a2a/client.py @@ -0,0 +1,107 @@ +""" +LiteLLM A2A Client class. + +Provides a class-based interface for A2A agent invocation. +""" + +from typing import TYPE_CHECKING, AsyncIterator, Dict, Optional + +from litellm.types.agents import LiteLLMSendMessageResponse + +if TYPE_CHECKING: + from a2a.client import A2AClient as A2AClientType + from a2a.types import ( + AgentCard, + SendMessageRequest, + SendStreamingMessageRequest, + SendStreamingMessageResponse, + ) + + +class A2AClient: + """ + LiteLLM wrapper for A2A agent invocation. + + Creates the underlying A2A client once on first use and reuses it. + + Example: + ```python + from litellm.a2a import A2AClient + from a2a.types import SendMessageRequest, MessageSendParams + from uuid import uuid4 + + client = A2AClient(base_url="http://localhost:10001") + + request = SendMessageRequest( + id=str(uuid4()), + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": "Hello!"}], + "messageId": uuid4().hex, + } + ) + ) + response = await client.send_message(request) + ``` + """ + + def __init__( + self, + base_url: str, + timeout: float = 60.0, + extra_headers: Optional[Dict[str, str]] = None, + ): + """ + Initialize the A2A client wrapper. + + Args: + base_url: The base URL of the A2A agent (e.g., "http://localhost:10001") + timeout: Request timeout in seconds (default: 60.0) + extra_headers: Optional additional headers to include in requests + """ + self.base_url = base_url + self.timeout = timeout + self.extra_headers = extra_headers + self._a2a_client: Optional["A2AClientType"] = None + + async def _get_client(self) -> "A2AClientType": + """Get or create the underlying A2A client.""" + if self._a2a_client is None: + from litellm.a2a.main import create_a2a_client + + self._a2a_client = await create_a2a_client( + base_url=self.base_url, + timeout=self.timeout, + extra_headers=self.extra_headers, + ) + return self._a2a_client + + async def get_agent_card(self) -> "AgentCard": + """Fetch the agent card from the server.""" + from litellm.a2a.main import aget_agent_card + + return await aget_agent_card( + base_url=self.base_url, + timeout=self.timeout, + extra_headers=self.extra_headers, + ) + + async def send_message( + self, request: "SendMessageRequest" + ) -> LiteLLMSendMessageResponse: + """Send a message to the A2A agent.""" + from litellm.a2a.main import asend_message + + a2a_client = await self._get_client() + return await asend_message(a2a_client=a2a_client, request=request) + + async def send_message_streaming( + self, request: "SendStreamingMessageRequest" + ) -> AsyncIterator["SendStreamingMessageResponse"]: + """Send a streaming message to the A2A agent.""" + from litellm.a2a.main import asend_message_streaming + + a2a_client = await self._get_client() + async for chunk in asend_message_streaming(a2a_client=a2a_client, request=request): + yield chunk diff --git a/litellm/a2a/main.py b/litellm/a2a/main.py new file mode 100644 index 00000000000..7e293dec81e --- /dev/null +++ b/litellm/a2a/main.py @@ -0,0 +1,261 @@ +""" +LiteLLM A2A SDK functions. + +Provides standalone functions with @client decorator for LiteLLM logging integration. +""" + +import asyncio +from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union + +from litellm._logging import verbose_logger +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.agents import LiteLLMSendMessageResponse +from litellm.utils import client + +if TYPE_CHECKING: + from a2a.client import A2AClient as A2AClientType + from a2a.types import ( + AgentCard, + SendMessageRequest, + SendStreamingMessageRequest, + SendStreamingMessageResponse, + ) + +# Runtime imports with availability check +A2A_SDK_AVAILABLE = False +A2ACardResolver: Any = None +_A2AClient: Any = None + +try: + from a2a.client import A2ACardResolver # type: ignore[no-redef] + from a2a.client import A2AClient as _A2AClient # type: ignore[no-redef] + + A2A_SDK_AVAILABLE = True +except ImportError: + pass + + +@client +async def asend_message( + a2a_client: "A2AClientType", + request: "SendMessageRequest", + **kwargs: Any, +) -> LiteLLMSendMessageResponse: + """ + Async: Send a message to an A2A agent. + + Uses the @client decorator for LiteLLM logging and tracking. + + Args: + a2a_client: An initialized a2a.client.A2AClient instance + request: SendMessageRequest from a2a.types + **kwargs: Additional arguments passed to the client decorator + + Returns: + LiteLLMSendMessageResponse (wraps a2a SendMessageResponse with _hidden_params) + + Example: + ```python + from litellm.a2a import asend_message, create_a2a_client + from a2a.types import SendMessageRequest, MessageSendParams + from uuid import uuid4 + + # Create client once + a2a_client = await create_a2a_client(base_url="http://localhost:10001") + + # Use it for multiple requests + request = SendMessageRequest( + id=str(uuid4()), + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": "Hello!"}], + "messageId": uuid4().hex, + } + ) + ) + response = await asend_message(a2a_client=a2a_client, request=request) + ``` + """ + verbose_logger.info(f"A2A send_message request_id={request.id}") + + a2a_response = await a2a_client.send_message(request) + + verbose_logger.info(f"A2A send_message completed, request_id={request.id}") + + # Wrap in LiteLLM response type for _hidden_params support + response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response) + + return response + + +@client +def send_message( + a2a_client: "A2AClientType", + request: "SendMessageRequest", + **kwargs: Any, +) -> Union[LiteLLMSendMessageResponse, Coroutine[Any, Any, LiteLLMSendMessageResponse]]: + """ + Sync: Send a message to an A2A agent. + + Uses the @client decorator for LiteLLM logging and tracking. + + Args: + a2a_client: An initialized a2a.client.A2AClient instance + request: SendMessageRequest from a2a.types + **kwargs: Additional arguments passed to the client decorator + + Returns: + LiteLLMSendMessageResponse (wraps a2a SendMessageResponse with _hidden_params) + """ + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + + if loop is not None: + return asend_message(a2a_client=a2a_client, request=request, **kwargs) + else: + return asyncio.run(asend_message(a2a_client=a2a_client, request=request, **kwargs)) + + +async def asend_message_streaming( + a2a_client: "A2AClientType", + request: "SendStreamingMessageRequest", +) -> AsyncIterator["SendStreamingMessageResponse"]: + """ + Async: Send a streaming message to an A2A agent. + + Args: + a2a_client: An initialized a2a.client.A2AClient instance + request: SendStreamingMessageRequest from a2a.types + + Yields: + SendStreamingMessageResponse chunks from the agent + """ + verbose_logger.info(f"A2A send_message_streaming request_id={request.id}") + + stream = a2a_client.send_message_streaming(request) + + chunk_count = 0 + async for chunk in stream: + chunk_count += 1 + yield chunk + + verbose_logger.info( + f"A2A send_message_streaming completed, request_id={request.id}, chunks={chunk_count}" + ) + + +async def create_a2a_client( + base_url: str, + timeout: float = 60.0, + extra_headers: Optional[Dict[str, str]] = None, +) -> "A2AClientType": + """ + Create an A2A client for the given agent URL. + + This resolves the agent card and returns a ready-to-use A2A client. + The client can be reused for multiple requests. + + Args: + base_url: The base URL of the A2A agent (e.g., "http://localhost:10001") + timeout: Request timeout in seconds (default: 60.0) + extra_headers: Optional additional headers to include in requests + + Returns: + An initialized a2a.client.A2AClient instance + + Example: + ```python + from litellm.a2a import create_a2a_client, asend_message + + # Create client once + client = await create_a2a_client(base_url="http://localhost:10001") + + # Reuse for multiple requests + response1 = await asend_message(a2a_client=client, request=request1) + response2 = await asend_message(a2a_client=client, request=request2) + ``` + """ + if not A2A_SDK_AVAILABLE: + raise ImportError( + "The 'a2a' package is required for A2A agent invocation. " + "Install it with: pip install a2a" + ) + + verbose_logger.info(f"Creating A2A client for {base_url}") + + # Use LiteLLM's cached httpx client + http_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.A2A, + params={"timeout": timeout}, + ) + httpx_client = http_handler.client + + # Resolve agent card + resolver = A2ACardResolver( + httpx_client=httpx_client, + base_url=base_url, + ) + agent_card = await resolver.get_agent_card() + + verbose_logger.debug( + f"Resolved agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}" + ) + + # Create and return A2A client + a2a_client = _A2AClient( + httpx_client=httpx_client, + agent_card=agent_card, + ) + + verbose_logger.info(f"A2A client created for {base_url}") + + return a2a_client + + +async def aget_agent_card( + base_url: str, + timeout: float = 60.0, + extra_headers: Optional[Dict[str, str]] = None, +) -> "AgentCard": + """ + Fetch the agent card from an A2A agent. + + Args: + base_url: The base URL of the A2A agent (e.g., "http://localhost:10001") + timeout: Request timeout in seconds (default: 60.0) + extra_headers: Optional additional headers to include in requests + + Returns: + AgentCard from the A2A agent + """ + if not A2A_SDK_AVAILABLE: + raise ImportError( + "The 'a2a' package is required for A2A agent invocation. " + "Install it with: pip install a2a" + ) + + verbose_logger.info(f"Fetching agent card from {base_url}") + + # Use LiteLLM's cached httpx client + http_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.A2A, + params={"timeout": timeout}, + ) + httpx_client = http_handler.client + + resolver = A2ACardResolver( + httpx_client=httpx_client, + base_url=base_url, + ) + agent_card = await resolver.get_agent_card() + + verbose_logger.info( + f"Fetched agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}" + ) + return agent_card diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 0b0f483ff74..4fd2c988797 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -71,6 +71,7 @@ from litellm.litellm_core_utils.redact_messages import ( from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.llms.base_llm.search.transformation import SearchResponse from litellm.responses.utils import ResponseAPILoggingUtils +from litellm.types.agents import LiteLLMSendMessageResponse from litellm.types.containers.main import ContainerObject from litellm.types.llms.openai import ( AllMessageValues, @@ -1738,6 +1739,7 @@ class Logging(LiteLLMLoggingBaseClass): and logging_result.get("object") == "search" # Search API (dict format) or isinstance(logging_result, VideoObject) or isinstance(logging_result, ContainerObject) + or isinstance(logging_result, LiteLLMSendMessageResponse) # A2A or (self.call_type == CallTypes.call_mcp_tool.value) ): return True diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py new file mode 100644 index 00000000000..4534e313c29 --- /dev/null +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -0,0 +1,189 @@ +""" +A2A Protocol endpoints for LiteLLM Proxy. + +Allows clients to invoke agents through LiteLLM using the A2A protocol. +The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM. +""" + +import json +from typing import Any, Optional + +from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi.responses import JSONResponse, StreamingResponse + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +def _jsonrpc_error( + request_id: Optional[str], + code: int, + message: str, + status_code: int = 400, +) -> JSONResponse: + """Create a JSON-RPC 2.0 error response.""" + return JSONResponse( + content={ + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": code, "message": message}, + }, + status_code=status_code, + ) + + +def _get_agent(agent_id: str): + """Look up an agent by ID or name. Returns None if not found.""" + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + + agent = global_agent_registry.get_agent_by_id(agent_id=agent_id) + if agent is None: + agent = global_agent_registry.get_agent_by_name(agent_name=agent_id) + return agent + + +async def _handle_send_message( + a2a_client: Any, + request_id: str, + params: dict, +) -> JSONResponse: + """Handle message/send method.""" + from a2a.types import MessageSendParams, SendMessageRequest + + from litellm.a2a import asend_message + + a2a_request = SendMessageRequest( + id=request_id, + params=MessageSendParams(**params), + ) + response = await asend_message(a2a_client=a2a_client, request=a2a_request) + return JSONResponse(content=response.model_dump(mode="json", exclude_none=True)) + + +async def _handle_stream_message( + a2a_client: Any, + request_id: str, + params: dict, +) -> StreamingResponse: + """Handle message/stream method.""" + from a2a.types import MessageSendParams, SendStreamingMessageRequest + + a2a_request = SendStreamingMessageRequest( + id=request_id, + params=MessageSendParams(**params), + ) + + async def stream_response(): + try: + async for chunk in a2a_client.send_message_streaming(a2a_request): + yield json.dumps(chunk.model_dump(mode="json", exclude_none=True)) + "\n" + except Exception as e: + verbose_proxy_logger.exception(f"Error streaming A2A response: {e}") + yield json.dumps({ + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": -32603, "message": f"Streaming error: {str(e)}"}, + }) + "\n" + + return StreamingResponse(stream_response(), media_type="application/x-ndjson") + + +@router.get( + "/a2a/{agent_id}/.well-known/agent-card.json", + tags=["[beta] A2A Agents"], + dependencies=[Depends(user_api_key_auth)], +) +async def get_agent_card( + agent_id: str, + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get the agent card for an agent (A2A discovery endpoint). + + The URL in the agent card is rewritten to point to the LiteLLM proxy, + so all subsequent A2A calls go through LiteLLM for logging and cost tracking. + """ + try: + agent = _get_agent(agent_id) + if agent is None: + raise HTTPException(status_code=404, detail=f"Agent '{agent_id}' not found") + + # Copy and rewrite URL to point to LiteLLM proxy + agent_card = dict(agent.agent_card_params) + agent_card["url"] = f"{str(request.base_url).rstrip('/')}/a2a/{agent_id}" + + verbose_proxy_logger.debug( + f"Returning agent card for '{agent_id}' with proxy URL: {agent_card['url']}" + ) + return JSONResponse(content=agent_card) + + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error getting agent card: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/a2a/{agent_id}", + tags=["[beta] A2A Agents"], + dependencies=[Depends(user_api_key_auth)], +) +async def invoke_agent_a2a( + agent_id: str, + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Invoke an agent using the A2A protocol (JSON-RPC 2.0). + + Supported methods: + - message/send: Send a message and get a response + - message/stream: Send a message and stream the response + """ + from litellm.a2a import create_a2a_client + + body = {} + try: + body = await request.json() + verbose_proxy_logger.debug(f"A2A request for agent '{agent_id}': {body}") + + # Validate JSON-RPC format + if body.get("jsonrpc") != "2.0": + return _jsonrpc_error(body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'") + + request_id = body.get("id") + method = body.get("method") + params = body.get("params", {}) + + # Find the agent + agent = _get_agent(agent_id) + if agent is None: + return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404) + + # Get backend URL + agent_url = agent.agent_card_params.get("url") + if not agent_url: + return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) + + verbose_proxy_logger.info(f"Proxying A2A request to agent '{agent_id}' at {agent_url}") + + # Create A2A client and dispatch to handler + a2a_client = await create_a2a_client(base_url=agent_url) + + if method == "message/send": + return await _handle_send_message(a2a_client, request_id, params) + elif method == "message/stream": + return await _handle_stream_message(a2a_client, request_id, params) + else: + return _jsonrpc_error(request_id, -32601, f"Method '{method}' not found") + + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error invoking agent: {e}") + return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {str(e)}", 500) diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 489c8e82302..90688036d84 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -29,7 +29,7 @@ router = APIRouter() @router.get( "/v1/agents", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=List[AgentResponse], ) @@ -101,7 +101,7 @@ from litellm.proxy.agent_endpoints.agent_registry import ( @router.post( "/v1/agents", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=AgentResponse, ) @@ -196,7 +196,7 @@ async def create_agent( @router.get( "/v1/agents/{agent_id}", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=AgentResponse, ) @@ -239,7 +239,7 @@ async def get_agent_by_id(agent_id: str): @router.put( "/v1/agents/{agent_id}", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=AgentResponse, ) @@ -328,7 +328,7 @@ async def update_agent( @router.patch( "/v1/agents/{agent_id}", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=AgentResponse, ) @@ -471,7 +471,7 @@ async def delete_agent(agent_id: str): @router.post( "/v1/agents/{agent_id}/make_public", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=AgentMakePublicResponse, ) @@ -585,7 +585,7 @@ async def make_agent_public( @router.post( "/v1/agents/make_public", - tags=["[beta] Agents"], + tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], response_model=AgentMakePublicResponse, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4d971e8ce42..8e2c229c914 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -188,6 +188,7 @@ from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) from litellm.proxy._types import * +from litellm.proxy.agent_endpoints.a2a_endpoints import router as a2a_router from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry from litellm.proxy.agent_endpoints.endpoints import router as agent_endpoints_router from litellm.proxy.analytics_endpoints.analytics_endpoints import ( @@ -10068,6 +10069,7 @@ app.include_router(user_agent_analytics_router) app.include_router(enterprise_router) app.include_router(ui_discovery_endpoints_router) app.include_router(agent_endpoints_router) +app.include_router(a2a_router) ######################################################## # MCP Server ######################################################## diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 850dab0ea7f..2eb26dc6227 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -1,9 +1,14 @@ from datetime import datetime -from typing import Any, Dict, List, Literal, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union -from pydantic import BaseModel +from pydantic import BaseModel, PrivateAttr from typing_extensions import Required, TypedDict +from litellm.types.llms.base import LiteLLMPydanticObjectBase + +if TYPE_CHECKING: + from a2a.types import SendMessageResponse + # AgentProvider class AgentProvider(TypedDict, total=False): @@ -200,3 +205,41 @@ class AgentMakePublicResponse(BaseModel): class MakeAgentsPublicRequest(BaseModel): agent_ids: List[str] + + +class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): + """ + LiteLLM wrapper for A2A SendMessageResponse. + + Wraps the a2a SDK's SendMessageResponse with LiteLLM's _hidden_params + for cost tracking and logging integration. + """ + + # A2A response fields + id: str + jsonrpc: str = "2.0" + result: Optional[Dict[str, Any]] = None + error: Optional[Dict[str, Any]] = None + + model_config = {"extra": "allow"} + + # LiteLLM private attributes for logging/cost tracking + _hidden_params: dict = PrivateAttr(default_factory=dict) + + @classmethod + def from_a2a_response( + cls, response: "SendMessageResponse" + ) -> "LiteLLMSendMessageResponse": + """ + Create a LiteLLMSendMessageResponse from an a2a SDK SendMessageResponse. + + Args: + response: The a2a SDK SendMessageResponse + + Returns: + LiteLLMSendMessageResponse with _hidden_params support + """ + # Convert the a2a response to a dict + response_dict = response.model_dump(mode="json", exclude_none=True) + + return cls(**response_dict) diff --git a/litellm/types/llms/custom_http.py b/litellm/types/llms/custom_http.py index 4d5ca02f032..9ed25005c05 100644 --- a/litellm/types/llms/custom_http.py +++ b/litellm/types/llms/custom_http.py @@ -24,6 +24,7 @@ class httpxSpecialProvider(str, Enum): Search = "search" MCP = "mcp" RAG = "rag" + A2A = "a2a" VerifyTypes = Union[str, bool, ssl.SSLContext] diff --git a/tests/agent_tests/test_a2a.py b/tests/agent_tests/test_a2a.py new file mode 100644 index 00000000000..fb5620e44a1 --- /dev/null +++ b/tests/agent_tests/test_a2a.py @@ -0,0 +1,138 @@ +""" +Test for LiteLLM A2A module. + +Run with: + pytest tests/agent_tests/test_a2a.py -v -s +""" + +import asyncio +import os +import sys +import json +from typing import Optional +from uuid import uuid4 + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.utils import StandardLoggingPayload + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + +from a2a.types import MessageSendParams, SendMessageRequest + + +@pytest.mark.asyncio +async def test_asend_message_with_client_decorator(): + """ + Test asend_message standalone function with @client decorator. + This tests the LiteLLM logging integration. + """ + litellm._turn_on_debug() + from litellm.a2a import asend_message, create_a2a_client + + # Create the A2A client first + a2a_client = await create_a2a_client(base_url="http://localhost:10001") + + # Build the request matching A2A SDK spec + send_message_payload = { + "message": { + "role": "user", + "parts": [ + { + "kind": "text", + "text": "Hello from @client decorated function!", + } + ], + "messageId": uuid4().hex, + }, + } + + request = SendMessageRequest( + id=str(uuid4()), + params=MessageSendParams(**send_message_payload), + ) + + # Send message using standalone function with @client decorator + response = await asend_message(a2a_client=a2a_client, request=request) + + # Print response for debugging + print("\n=== A2A Response (standalone with @client) ===") + print(response.model_dump(mode="json", exclude_none=True)) + + # Basic assertions + assert response is not None + + +class TestA2ALogger(CustomLogger): + """Custom logger to capture A2A logging payloads for testing.""" + + def __init__(self): + self.standard_logging_payload: Optional[StandardLoggingPayload] = None + self.logged_kwargs: Optional[dict] = None + self.log_success_called = False + super().__init__() + + async def async_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): + print("TestA2ALogger: async_log_success_event called") + self.log_success_called = True + self.logged_kwargs = kwargs + self.standard_logging_payload = kwargs.get("standard_logging_object", None) + print(f"Captured standard_logging_payload: {self.standard_logging_payload}") + + +@pytest.mark.asyncio +async def test_a2a_logging_payload(): + """ + Test that A2A calls create a standard logging payload. + Validates the @client decorator integration with LiteLLM logging. + """ + # Reset callbacks and set up custom logger + litellm.logging_callback_manager._reset_all_callbacks() + test_logger = TestA2ALogger() + litellm.callbacks = [test_logger] + + from litellm.a2a import asend_message, create_a2a_client + + # Create the A2A client first + a2a_client = await create_a2a_client(base_url="http://localhost:10001") + + # Build the request + send_message_payload = { + "message": { + "role": "user", + "parts": [ + { + "kind": "text", + "text": "Hello! Testing logging payload.", + } + ], + "messageId": uuid4().hex, + }, + } + + request = SendMessageRequest( + id=str(uuid4()), + params=MessageSendParams(**send_message_payload), + ) + + # Send message + response = await asend_message(a2a_client=a2a_client, request=request) + + # Give async logging time to complete + await asyncio.sleep(1) + + # Print debug info + print("\n=== Logging Validation ===") + print(f"log_success_called: {test_logger.log_success_called}") + print(f"standard_logging_payload: {test_logger.standard_logging_payload}") + print(f"logged kwargs: {json.dumps(test_logger.logged_kwargs, indent=4, default=str)}") + + assert test_logger.standard_logging_payload is not None + + \ No newline at end of file