finish a2a
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
from .server import A2AServer
|
||||
from .task_manager import TaskManager, InMemoryTaskManager
|
||||
|
||||
__all__ = ["A2AServer", "TaskManager", "InMemoryTaskManager"]
|
||||
@@ -0,0 +1,168 @@
|
||||
from starlette.applications import Starlette
|
||||
from starlette.responses import JSONResponse
|
||||
from sse_starlette.sse import EventSourceResponse
|
||||
from starlette.requests import Request
|
||||
from common.types import (
|
||||
A2ARequest,
|
||||
JSONRPCResponse,
|
||||
InvalidRequestError,
|
||||
JSONParseError,
|
||||
GetTaskRequest,
|
||||
CancelTaskRequest,
|
||||
SendTaskRequest,
|
||||
SetTaskPushNotificationRequest,
|
||||
GetTaskPushNotificationRequest,
|
||||
InternalError,
|
||||
AgentCard,
|
||||
TaskResubscriptionRequest,
|
||||
SendTaskStreamingRequest,
|
||||
Message,
|
||||
)
|
||||
from pydantic import ValidationError
|
||||
import json
|
||||
from typing import AsyncIterable, Any
|
||||
from common.server.task_manager import TaskManager
|
||||
|
||||
import logging
|
||||
|
||||
# Configure a logger specific to the server
|
||||
logger = logging.getLogger("A2AServer")
|
||||
|
||||
|
||||
class A2AServer:
|
||||
def __init__(
|
||||
self,
|
||||
host="0.0.0.0",
|
||||
port=5000,
|
||||
endpoint="/",
|
||||
agent_card: AgentCard = None,
|
||||
task_manager: TaskManager = None,
|
||||
):
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.endpoint = endpoint
|
||||
self.task_manager = task_manager
|
||||
self.agent_card = agent_card
|
||||
self.app = Starlette()
|
||||
self.app.add_route(self.endpoint, self._process_request, methods=["POST"])
|
||||
self.app.add_route(
|
||||
"/.well-known/agent.json", self._get_agent_card, methods=["GET"]
|
||||
)
|
||||
|
||||
def start(self):
|
||||
if self.agent_card is None:
|
||||
raise ValueError("agent_card is not defined")
|
||||
|
||||
if self.task_manager is None:
|
||||
raise ValueError("request_handler is not defined")
|
||||
|
||||
import uvicorn
|
||||
|
||||
# Basic logging config moved to __main__.py for application-level control
|
||||
uvicorn.run(self.app, host=self.host, port=self.port)
|
||||
|
||||
def _get_agent_card(self, request: Request) -> JSONResponse:
|
||||
logger.info("Serving Agent Card request")
|
||||
return JSONResponse(self.agent_card.model_dump(exclude_none=True))
|
||||
|
||||
async def _process_request(self, request: Request):
|
||||
request_id_for_log = "N/A" # Default if parsing fails early
|
||||
raw_body = b""
|
||||
try:
|
||||
# Log raw body first
|
||||
raw_body = await request.body()
|
||||
body = json.loads(raw_body) # Attempt parsing
|
||||
request_id_for_log = body.get("id", "N/A") # Get ID if possible
|
||||
logger.info(f"<- Received Request (ID: {request_id_for_log}):\n{json.dumps(body, indent=2)}")
|
||||
|
||||
json_rpc_request = A2ARequest.validate_python(body)
|
||||
|
||||
# Route based on method (same as before)
|
||||
if isinstance(json_rpc_request, GetTaskRequest):
|
||||
result = await self.task_manager.on_get_task(json_rpc_request)
|
||||
elif isinstance(json_rpc_request, SendTaskRequest):
|
||||
result = await self.task_manager.on_send_task(json_rpc_request)
|
||||
elif isinstance(json_rpc_request, SendTaskStreamingRequest):
|
||||
result = await self.task_manager.on_send_task_subscribe(
|
||||
json_rpc_request
|
||||
)
|
||||
elif isinstance(json_rpc_request, CancelTaskRequest):
|
||||
result = await self.task_manager.on_cancel_task(json_rpc_request)
|
||||
elif isinstance(json_rpc_request, SetTaskPushNotificationRequest):
|
||||
result = await self.task_manager.on_set_task_push_notification(json_rpc_request)
|
||||
elif isinstance(json_rpc_request, GetTaskPushNotificationRequest):
|
||||
result = await self.task_manager.on_get_task_push_notification(json_rpc_request)
|
||||
elif isinstance(json_rpc_request, TaskResubscriptionRequest):
|
||||
result = await self.task_manager.on_resubscribe_to_task(
|
||||
json_rpc_request
|
||||
)
|
||||
else:
|
||||
logger.warning(f"Unexpected request type: {type(json_rpc_request)}")
|
||||
raise ValueError(f"Unexpected request type: {type(request)}")
|
||||
|
||||
return self._create_response(result) # Pass result to response creation
|
||||
|
||||
except json.decoder.JSONDecodeError as e:
|
||||
logger.error(f"JSON Parse Error for Request body: <<<{raw_body.decode('utf-8', errors='replace')}>>>\nError: {e}")
|
||||
return self._handle_exception(e, request_id_for_log) # Pass ID if known
|
||||
except ValidationError as e:
|
||||
logger.error(f"Request Validation Error (ID: {request_id_for_log}): {e.json()}")
|
||||
return self._handle_exception(e, request_id_for_log)
|
||||
except Exception as e:
|
||||
logger.error(f"Unhandled Exception processing request (ID: {request_id_for_log}): {e}", exc_info=True)
|
||||
return self._handle_exception(e, request_id_for_log) # Pass ID if known
|
||||
|
||||
def _handle_exception(self, e: Exception, req_id=None) -> JSONResponse: # Accept req_id
|
||||
if isinstance(e, json.decoder.JSONDecodeError):
|
||||
json_rpc_error = JSONParseError()
|
||||
elif isinstance(e, ValidationError):
|
||||
json_rpc_error = InvalidRequestError(data=json.loads(e.json()))
|
||||
else:
|
||||
# Log the full exception details
|
||||
logger.error(f"Internal Server Error (ReqID: {req_id}): {e}", exc_info=True)
|
||||
json_rpc_error = InternalError(message=f"Internal Server Error: {type(e).__name__}")
|
||||
|
||||
response = JSONRPCResponse(id=req_id, error=json_rpc_error)
|
||||
response_dump = response.model_dump(exclude_none=True)
|
||||
logger.info(f"-> Sending Error Response (ReqID: {req_id}):\n{json.dumps(response_dump, indent=2)}")
|
||||
# A2A errors are still sent with HTTP 200
|
||||
return JSONResponse(response_dump, status_code=200)
|
||||
|
||||
def _create_response(self, result: Any) -> JSONResponse | EventSourceResponse:
|
||||
if isinstance(result, AsyncIterable):
|
||||
# Streaming response
|
||||
async def event_generator(result_stream) -> AsyncIterable[dict[str, str]]:
|
||||
stream_request_id = None # Capture ID from the first event if possible
|
||||
try:
|
||||
async for item in result_stream:
|
||||
# Log each streamed item
|
||||
response_json = item.model_dump_json(exclude_none=True)
|
||||
stream_request_id = item.id # Update ID
|
||||
logger.info(f"-> Sending SSE Event (ID: {stream_request_id}):\n{json.dumps(json.loads(response_json), indent=2)}")
|
||||
yield {"data": response_json}
|
||||
logger.info(f"SSE Stream ended for request ID: {stream_request_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error during SSE generation (ReqID: {stream_request_id}): {e}", exc_info=True)
|
||||
# Optionally yield an error event if the protocol allows/requires it
|
||||
# error_payload = JSONRPCResponse(id=stream_request_id, error=InternalError(message=f"SSE Error: {e}"))
|
||||
# yield {"data": error_payload.model_dump_json(exclude_none=True)}
|
||||
|
||||
logger.info("Starting SSE stream...") # Log stream start
|
||||
return EventSourceResponse(event_generator(result))
|
||||
elif isinstance(result, JSONRPCResponse):
|
||||
# Standard JSON response
|
||||
response_dump = result.model_dump(exclude_none=True)
|
||||
log_id = result.id if result.id is not None else "N/A (Notification?)"
|
||||
log_prefix = "->"
|
||||
log_type = "Response"
|
||||
if result.error:
|
||||
log_prefix = "-> Sending Error"
|
||||
log_type = "Error Response"
|
||||
|
||||
logger.info(f"{log_prefix} {log_type} (ID: {log_id}):\n{json.dumps(response_dump, indent=2)}")
|
||||
return JSONResponse(response_dump)
|
||||
else:
|
||||
# This should ideally not happen if task manager returns correctly
|
||||
logger.error(f"Task manager returned unexpected type: {type(result)}")
|
||||
err_resp = JSONRPCResponse(id=None, error=InternalError(message="Invalid internal response type"))
|
||||
return JSONResponse(err_resp.model_dump(exclude_none=True), status_code=500)
|
||||
@@ -0,0 +1,277 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Union, AsyncIterable, List
|
||||
from common.types import Task
|
||||
from common.types import (
|
||||
JSONRPCResponse,
|
||||
TaskIdParams,
|
||||
TaskQueryParams,
|
||||
GetTaskRequest,
|
||||
TaskNotFoundError,
|
||||
SendTaskRequest,
|
||||
CancelTaskRequest,
|
||||
TaskNotCancelableError,
|
||||
SetTaskPushNotificationRequest,
|
||||
GetTaskPushNotificationRequest,
|
||||
GetTaskResponse,
|
||||
CancelTaskResponse,
|
||||
SendTaskResponse,
|
||||
SetTaskPushNotificationResponse,
|
||||
GetTaskPushNotificationResponse,
|
||||
PushNotificationNotSupportedError,
|
||||
TaskSendParams,
|
||||
TaskStatus,
|
||||
TaskState,
|
||||
TaskResubscriptionRequest,
|
||||
SendTaskStreamingRequest,
|
||||
SendTaskStreamingResponse,
|
||||
Artifact,
|
||||
PushNotificationConfig,
|
||||
TaskStatusUpdateEvent,
|
||||
JSONRPCError,
|
||||
TaskPushNotificationConfig,
|
||||
InternalError,
|
||||
)
|
||||
from common.server.utils import new_not_implemented_error
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class TaskManager(ABC):
|
||||
@abstractmethod
|
||||
async def on_get_task(self, request: GetTaskRequest) -> GetTaskResponse:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_cancel_task(self, request: CancelTaskRequest) -> CancelTaskResponse:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_send_task(self, request: SendTaskRequest) -> SendTaskResponse:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_send_task_subscribe(
|
||||
self, request: SendTaskStreamingRequest
|
||||
) -> Union[AsyncIterable[SendTaskStreamingResponse], JSONRPCResponse]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_set_task_push_notification(
|
||||
self, request: SetTaskPushNotificationRequest
|
||||
) -> SetTaskPushNotificationResponse:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_get_task_push_notification(
|
||||
self, request: GetTaskPushNotificationRequest
|
||||
) -> GetTaskPushNotificationResponse:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_resubscribe_to_task(
|
||||
self, request: TaskResubscriptionRequest
|
||||
) -> Union[AsyncIterable[SendTaskResponse], JSONRPCResponse]:
|
||||
pass
|
||||
|
||||
|
||||
class InMemoryTaskManager(TaskManager):
|
||||
def __init__(self):
|
||||
self.tasks: dict[str, Task] = {}
|
||||
self.push_notification_infos: dict[str, PushNotificationConfig] = {}
|
||||
self.lock = asyncio.Lock()
|
||||
self.task_sse_subscribers: dict[str, List[asyncio.Queue]] = {}
|
||||
self.subscriber_lock = asyncio.Lock()
|
||||
|
||||
async def on_get_task(self, request: GetTaskRequest) -> GetTaskResponse:
|
||||
logger.info(f"Getting task {request.params.id}")
|
||||
task_query_params: TaskQueryParams = request.params
|
||||
|
||||
async with self.lock:
|
||||
task = self.tasks.get(task_query_params.id)
|
||||
if task is None:
|
||||
return GetTaskResponse(id=request.id, error=TaskNotFoundError())
|
||||
|
||||
task_result = self.append_task_history(
|
||||
task, task_query_params.historyLength
|
||||
)
|
||||
|
||||
return GetTaskResponse(id=request.id, result=task_result)
|
||||
|
||||
async def on_cancel_task(self, request: CancelTaskRequest) -> CancelTaskResponse:
|
||||
logger.info(f"Cancelling task {request.params.id}")
|
||||
task_id_params: TaskIdParams = request.params
|
||||
|
||||
async with self.lock:
|
||||
task = self.tasks.get(task_id_params.id)
|
||||
if task is None:
|
||||
return CancelTaskResponse(id=request.id, error=TaskNotFoundError())
|
||||
|
||||
return CancelTaskResponse(id=request.id, error=TaskNotCancelableError())
|
||||
|
||||
@abstractmethod
|
||||
async def on_send_task(self, request: SendTaskRequest) -> SendTaskResponse:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def on_send_task_subscribe(
|
||||
self, request: SendTaskStreamingRequest
|
||||
) -> Union[AsyncIterable[SendTaskStreamingResponse], JSONRPCResponse]:
|
||||
pass
|
||||
|
||||
async def set_push_notification_info(self, task_id: str, notification_config: PushNotificationConfig):
|
||||
async with self.lock:
|
||||
task = self.tasks.get(task_id)
|
||||
if task is None:
|
||||
raise ValueError(f"Task not found for {task_id}")
|
||||
|
||||
self.push_notification_infos[task_id] = notification_config
|
||||
|
||||
return
|
||||
|
||||
async def get_push_notification_info(self, task_id: str) -> PushNotificationConfig:
|
||||
async with self.lock:
|
||||
task = self.tasks.get(task_id)
|
||||
if task is None:
|
||||
raise ValueError(f"Task not found for {task_id}")
|
||||
|
||||
return self.push_notification_infos[task_id]
|
||||
|
||||
return
|
||||
|
||||
async def has_push_notification_info(self, task_id: str) -> bool:
|
||||
async with self.lock:
|
||||
return task_id in self.push_notification_infos
|
||||
|
||||
|
||||
async def on_set_task_push_notification(
|
||||
self, request: SetTaskPushNotificationRequest
|
||||
) -> SetTaskPushNotificationResponse:
|
||||
logger.info(f"Setting task push notification {request.params.id}")
|
||||
task_notification_params: TaskPushNotificationConfig = request.params
|
||||
|
||||
try:
|
||||
await self.set_push_notification_info(task_notification_params.id, task_notification_params.pushNotificationConfig)
|
||||
except Exception as e:
|
||||
logger.error(f"Error while setting push notification info: {e}")
|
||||
return JSONRPCResponse(
|
||||
id=request.id,
|
||||
error=InternalError(
|
||||
message="An error occurred while setting push notification info"
|
||||
),
|
||||
)
|
||||
|
||||
return SetTaskPushNotificationResponse(id=request.id, result=task_notification_params)
|
||||
|
||||
async def on_get_task_push_notification(
|
||||
self, request: GetTaskPushNotificationRequest
|
||||
) -> GetTaskPushNotificationResponse:
|
||||
logger.info(f"Getting task push notification {request.params.id}")
|
||||
task_params: TaskIdParams = request.params
|
||||
|
||||
try:
|
||||
notification_info = await self.get_push_notification_info(task_params.id)
|
||||
except Exception as e:
|
||||
logger.error(f"Error while getting push notification info: {e}")
|
||||
return GetTaskPushNotificationResponse(
|
||||
id=request.id,
|
||||
error=InternalError(
|
||||
message="An error occurred while getting push notification info"
|
||||
),
|
||||
)
|
||||
|
||||
return GetTaskPushNotificationResponse(id=request.id, result=TaskPushNotificationConfig(id=task_params.id, pushNotificationConfig=notification_info))
|
||||
|
||||
async def upsert_task(self, task_send_params: TaskSendParams) -> Task:
|
||||
logger.info(f"Upserting task {task_send_params.id}")
|
||||
async with self.lock:
|
||||
task = self.tasks.get(task_send_params.id)
|
||||
if task is None:
|
||||
task = Task(
|
||||
id=task_send_params.id,
|
||||
sessionId = task_send_params.sessionId,
|
||||
messages=[task_send_params.message],
|
||||
status=TaskStatus(state=TaskState.SUBMITTED),
|
||||
history=[task_send_params.message],
|
||||
)
|
||||
self.tasks[task_send_params.id] = task
|
||||
else:
|
||||
task.history.append(task_send_params.message)
|
||||
|
||||
return task
|
||||
|
||||
async def on_resubscribe_to_task(
|
||||
self, request: TaskResubscriptionRequest
|
||||
) -> Union[AsyncIterable[SendTaskStreamingResponse], JSONRPCResponse]:
|
||||
return new_not_implemented_error(request.id)
|
||||
|
||||
async def update_store(
|
||||
self, task_id: str, status: TaskStatus, artifacts: list[Artifact]
|
||||
) -> Task:
|
||||
async with self.lock:
|
||||
try:
|
||||
task = self.tasks[task_id]
|
||||
except KeyError:
|
||||
logger.error(f"Task {task_id} not found for updating the task")
|
||||
raise ValueError(f"Task {task_id} not found")
|
||||
|
||||
task.status = status
|
||||
|
||||
if status.message is not None:
|
||||
task.history.append(status.message)
|
||||
|
||||
if artifacts is not None:
|
||||
if task.artifacts is None:
|
||||
task.artifacts = []
|
||||
task.artifacts.extend(artifacts)
|
||||
|
||||
return task
|
||||
|
||||
def append_task_history(self, task: Task, historyLength: int | None):
|
||||
new_task = task.model_copy()
|
||||
if historyLength is not None and historyLength > 0:
|
||||
new_task.history = new_task.history[-historyLength:]
|
||||
else:
|
||||
new_task.history = []
|
||||
|
||||
return new_task
|
||||
|
||||
async def setup_sse_consumer(self, task_id: str, is_resubscribe: bool = False):
|
||||
async with self.subscriber_lock:
|
||||
if task_id not in self.task_sse_subscribers:
|
||||
if is_resubscribe:
|
||||
raise ValueError("Task not found for resubscription")
|
||||
else:
|
||||
self.task_sse_subscribers[task_id] = []
|
||||
|
||||
sse_event_queue = asyncio.Queue(maxsize=0) # <=0 is unlimited
|
||||
self.task_sse_subscribers[task_id].append(sse_event_queue)
|
||||
return sse_event_queue
|
||||
|
||||
async def enqueue_events_for_sse(self, task_id, task_update_event):
|
||||
async with self.subscriber_lock:
|
||||
if task_id not in self.task_sse_subscribers:
|
||||
return
|
||||
|
||||
current_subscribers = self.task_sse_subscribers[task_id]
|
||||
for subscriber in current_subscribers:
|
||||
await subscriber.put(task_update_event)
|
||||
|
||||
async def dequeue_events_for_sse(
|
||||
self, request_id, task_id, sse_event_queue: asyncio.Queue
|
||||
) -> AsyncIterable[SendTaskStreamingResponse] | JSONRPCResponse:
|
||||
try:
|
||||
while True:
|
||||
event = await sse_event_queue.get()
|
||||
if isinstance(event, JSONRPCError):
|
||||
yield SendTaskStreamingResponse(id=request_id, error=event)
|
||||
break
|
||||
|
||||
yield SendTaskStreamingResponse(id=request_id, result=event)
|
||||
if isinstance(event, TaskStatusUpdateEvent) and event.final:
|
||||
break
|
||||
finally:
|
||||
async with self.subscriber_lock:
|
||||
if task_id in self.task_sse_subscribers:
|
||||
self.task_sse_subscribers[task_id].remove(sse_event_queue)
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
from common.types import (
|
||||
JSONRPCResponse,
|
||||
ContentTypeNotSupportedError,
|
||||
UnsupportedOperationError,
|
||||
)
|
||||
from typing import List
|
||||
|
||||
|
||||
def are_modalities_compatible(
|
||||
server_output_modes: List[str], client_output_modes: List[str]
|
||||
):
|
||||
"""Modalities are compatible if they are both non-empty
|
||||
and there is at least one common element."""
|
||||
if client_output_modes is None or len(client_output_modes) == 0:
|
||||
return True
|
||||
|
||||
if server_output_modes is None or len(server_output_modes) == 0:
|
||||
return True
|
||||
|
||||
return any(x in server_output_modes for x in client_output_modes)
|
||||
|
||||
|
||||
def new_incompatible_types_error(request_id):
|
||||
return JSONRPCResponse(id=request_id, error=ContentTypeNotSupportedError())
|
||||
|
||||
|
||||
def new_not_implemented_error(request_id):
|
||||
return JSONRPCResponse(id=request_id, error=UnsupportedOperationError())
|
||||
Reference in New Issue
Block a user