mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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
This commit is contained in:
parent
f035984dd7
commit
4370f6fb74
11 changed files with 1148 additions and 9 deletions
337
enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
Normal file
337
enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
Normal file
|
|
@ -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))
|
||||
59
litellm/a2a/__init__.py
Normal file
59
litellm/a2a/__init__.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
107
litellm/a2a/client.py
Normal file
107
litellm/a2a/client.py
Normal file
|
|
@ -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
|
||||
261
litellm/a2a/main.py
Normal file
261
litellm/a2a/main.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
189
litellm/proxy/agent_endpoints/a2a_endpoints.py
Normal file
189
litellm/proxy/agent_endpoints/a2a_endpoints.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ class httpxSpecialProvider(str, Enum):
|
|||
Search = "search"
|
||||
MCP = "mcp"
|
||||
RAG = "rag"
|
||||
A2A = "a2a"
|
||||
|
||||
|
||||
VerifyTypes = Union[str, bool, ssl.SSLContext]
|
||||
|
|
|
|||
138
tests/agent_tests/test_a2a.py
Normal file
138
tests/agent_tests/test_a2a.py
Normal file
|
|
@ -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
|
||||
|
||||
|
||||
Loading…
Add table
Reference in a new issue