""" A2A server for the coding agent. """ from __future__ import annotations import json import os import secrets import uuid from contextlib import asynccontextmanager from datetime import datetime from typing import Any, AsyncGenerator, Optional from fastapi import FastAPI, HTTPException, Request from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse, StreamingResponse from pydantic import BaseModel, Field from coding_a2a_agent.common.agent_callback_utils import AgentCallbackHandler, CallbackContextManager from coding_a2a_agent.agent import CodingA2ARuntime from coding_a2a_agent.config import AgentMetadata, CodingRequestConfig, LiteLLMConfig SERVICE_HOST = os.getenv("SERVICE_HOST", "0.0.0.0") SERVICE_PORT = int(os.getenv("SERVICE_PORT", "8000")) POD_NAME = os.getenv("POD_NAME", "coding-a2a-agent") TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "coding_a2a_agent") USER_ID = os.getenv("USER_ID", "") HEICODE_AGENT_ID = os.getenv("HEICODE_AGENT_ID", "") AGENT_ACCESS_TOKEN = os.getenv("AGENT_ACCESS_TOKEN", "") AGENT_ACCESS_HEADER = "X-Agent-Access-Token" class A2APart(BaseModel): kind: str = "text" text: Optional[str] = None data: Optional[dict[str, Any]] = None mime_type: Optional[str] = None class A2AMessage(BaseModel): role: str parts: list[A2APart] messageId: str = Field(default_factory=lambda: uuid.uuid4().hex) class A2ARequest(BaseModel): jsonrpc: str = "2.0" id: str method: str params: Optional[dict[str, Any]] = None class A2AArtifact(BaseModel): artifactId: str = Field(default_factory=lambda: uuid.uuid4().hex) name: str = "coding-result" parts: list[A2APart] class A2ATaskStatus(BaseModel): state: str timestamp: str = Field(default_factory=lambda: datetime.utcnow().isoformat() + "Z") message: Optional[str] = None class A2ATask(BaseModel): kind: str = "task" id: str = Field(default_factory=lambda: uuid.uuid4().hex) contextId: str = Field(default_factory=lambda: uuid.uuid4().hex) status: A2ATaskStatus artifacts: Optional[list[A2AArtifact]] = None metadata: Optional[dict[str, Any]] = None class AgentSkill(BaseModel): id: str name: str description: str class AgentCapabilities(BaseModel): text: bool = True streaming: bool = True push_notifications: bool = False forms: bool = False files: bool = True class AgentCard(BaseModel): name: str description: str version: str url: str capabilities: AgentCapabilities skills: list[AgentSkill] authentication: Optional[dict[str, Any]] = None class CodingA2AServer: def __init__(self, api_key: Optional[str] = None, model: Optional[str] = None): self.llm_config = LiteLLMConfig(api_key=api_key, model=model or LiteLLMConfig().model) self.metadata = AgentMetadata() self.runtime = CodingA2ARuntime(self.llm_config, self.metadata) self.callback_handler = AgentCallbackHandler(agent_name=POD_NAME, user_id=USER_ID) self.tasks: dict[str, A2ATask] = {} self.app = self._create_app() def _agent_access_required(self) -> bool: return bool(AGENT_ACCESS_TOKEN) def _agent_authentication_card(self) -> Optional[dict[str, Any]]: if not self._agent_access_required(): return None return { "type": "header", "header": AGENT_ACCESS_HEADER, "required": True, "agent_id": HEICODE_AGENT_ID or POD_NAME, } def _authorize_agent_request(self, request: Request) -> Optional[JSONResponse]: expected_token = AGENT_ACCESS_TOKEN if not expected_token: return None provided_token = request.headers.get(AGENT_ACCESS_HEADER, "") if not provided_token: return JSONResponse( status_code=401, content={"detail": f"missing {AGENT_ACCESS_HEADER}"}, ) if not secrets.compare_digest(expected_token, provided_token): return JSONResponse( status_code=403, content={"detail": "agent access denied"}, ) return None def _create_app(self) -> FastAPI: @asynccontextmanager async def lifespan(app: FastAPI): yield app = FastAPI( title=f"{self.metadata.name} - A2A", version=self.metadata.version, lifespan=lifespan, ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) self._register_routes(app) return app def _register_routes(self, app: FastAPI): @app.get("/") async def root(): runtime_defaults = CodingRequestConfig() return { "name": self.metadata.name, "version": self.metadata.version, "protocol": "A2A", "status": "running", "pod_name": POD_NAME, "template_type": TEMPLATE_TYPE, "role_name": self.metadata.role_name, "instruction_source": self.metadata.instruction_source, "enabled_resources": runtime_defaults.resources.enabled_resource_names, "auth_required": self._agent_access_required(), "agent_id": HEICODE_AGENT_ID or POD_NAME, } @app.get("/health") async def health(): runtime_defaults = CodingRequestConfig() return { "status": "healthy", "template_type": TEMPLATE_TYPE, "role_name": self.metadata.role_name, "instruction_source": self.metadata.instruction_source, "enabled_resources": runtime_defaults.resources.enabled_resource_names, "auth_required": self._agent_access_required(), "agent_id": HEICODE_AGENT_ID or POD_NAME, "timestamp": datetime.utcnow().isoformat() + "Z", } @app.get("/.well-known/agent.json") async def agent_card(request: Request): base_url = str(request.base_url).rstrip("/") card = AgentCard( name=self.metadata.name, description=self.metadata.description, version=self.metadata.version, url=base_url, capabilities=AgentCapabilities(streaming=self.metadata.enable_streaming), authentication=self._agent_authentication_card(), skills=[ AgentSkill(id="coding", name="Coding", description="Inspect, edit, and verify repositories like a Claude Code style coding agent."), AgentSkill(id="git", name="Git", description="Prepare workspaces, inspect git state, branch, commit, and push for Gitea, GitHub, and GitLab."), AgentSkill(id="data", name="Data Resources", description="Inspect MySQL/PostgreSQL schemas and Azure Blob artifacts when configured."), ], ) return card.model_dump() @app.post("/message/send") async def message_send(request: Request): auth_error = self._authorize_agent_request(request) if auth_error: return auth_error rpc_request = A2ARequest(**(await request.json())) if rpc_request.method != "message/send": return JSONResponse( {"jsonrpc": "2.0", "id": rpc_request.id, "error": {"code": -32601, "message": f"Method not found: {rpc_request.method}"}} ) return await self._handle_message_send(rpc_request) @app.post("/message/stream") async def message_stream(request: Request): auth_error = self._authorize_agent_request(request) if auth_error: return auth_error rpc_request = A2ARequest(**(await request.json())) return await self._handle_message_stream(rpc_request) @app.get("/tasks/{task_id}") async def get_task(task_id: str, request: Request): auth_error = self._authorize_agent_request(request) if auth_error: return auth_error if task_id not in self.tasks: raise HTTPException(status_code=404, detail="Task not found") return self.tasks[task_id].model_dump() async def _handle_message_send(self, request: A2ARequest) -> JSONResponse: params = request.params or {} message_text = self._extract_message_text(params.get("message", {})) if not message_text: return JSONResponse( {"jsonrpc": "2.0", "id": request.id, "error": {"code": -32602, "message": "Invalid params: no text content found"}} ) task_id = uuid.uuid4().hex context_id = params.get("contextId", uuid.uuid4().hex) task = A2ATask(id=task_id, contextId=context_id, status=A2ATaskStatus(state="working")) self.tasks[task_id] = task runtime_config = CodingRequestConfig.model_validate(params.get("configuration") or {}) api_key = params.get("api_key") or runtime_config.api_key or self.llm_config.api_key model = params.get("model") or runtime_config.model or self.llm_config.model try: with CallbackContextManager( handler=self.callback_handler, user_id=params.get("user_id") or USER_ID, request_id=task_id, ) as callback: callback.add_tool("a2a_message_send") result = await self.runtime.run_task( message_text, runtime_config, api_key=api_key, model=model, ) task.status = A2ATaskStatus(state="completed") task.artifacts = [ A2AArtifact( name="coding-result", parts=[A2APart(kind="text", text=result.response_text)], ) ] task.metadata = { "summary": result.summary, "workspace_root": result.workspace_root, "files_changed": result.files_changed, "tool_log": [entry.model_dump() for entry in result.tool_log], "resources_used": result.resources_used, "role_name": self.metadata.role_name, "instruction_source": self.metadata.instruction_source, } self.tasks[task_id] = task return JSONResponse({"jsonrpc": "2.0", "id": request.id, "result": task.model_dump()}) except Exception as exc: task.status = A2ATaskStatus(state="failed", message=str(exc)) self.tasks[task_id] = task return JSONResponse( { "jsonrpc": "2.0", "id": request.id, "error": {"code": -32000, "message": f"Agent error: {exc}"}, } ) async def _handle_message_stream(self, request: A2ARequest) -> StreamingResponse: params = request.params or {} message_text = self._extract_message_text(params.get("message", {})) runtime_config = CodingRequestConfig.model_validate(params.get("configuration") or {}) api_key = params.get("api_key") or runtime_config.api_key or self.llm_config.api_key model = params.get("model") or runtime_config.model or self.llm_config.model task_id = uuid.uuid4().hex context_id = params.get("contextId", uuid.uuid4().hex) async def event_stream() -> AsyncGenerator[str, None]: start_event = {"kind": "task-start", "taskId": task_id, "contextId": context_id} yield f"data: {json.dumps(start_event, ensure_ascii=False)}\n\n" try: result = await self.runtime.run_task( message_text, runtime_config, api_key=api_key, model=model, ) artifact_event = { "kind": "artifact", "taskId": task_id, "contextId": context_id, "data": { "text": result.response_text, "summary": result.summary, "files_changed": result.files_changed, }, } yield f"data: {json.dumps(artifact_event, ensure_ascii=False)}\n\n" finish_event = {"kind": "task-complete", "taskId": task_id, "contextId": context_id} yield f"data: {json.dumps(finish_event, ensure_ascii=False)}\n\n" except Exception as exc: error_event = { "kind": "task-failed", "taskId": task_id, "contextId": context_id, "data": {"message": str(exc)}, } yield f"data: {json.dumps(error_event, ensure_ascii=False)}\n\n" return StreamingResponse(event_stream(), media_type="text/event-stream") def _extract_message_text(self, message: dict[str, Any]) -> str: parts = message.get("parts", []) return "".join(part.get("text", "") for part in parts if part.get("kind") == "text") def create_app(api_key: Optional[str] = None, model: Optional[str] = None) -> FastAPI: return CodingA2AServer(api_key=api_key, model=model).app