|
| 1 | +""" |
| 2 | +Module for A2A Agent. |
| 3 | +""" |
| 4 | + |
| 5 | +import logging |
| 6 | +import sys |
| 7 | +import traceback |
| 8 | +from typing import Callable |
| 9 | + |
| 10 | +import uvicorn |
| 11 | +from crewai_tools import MCPServerAdapter |
| 12 | +from crewai_tools.adapters.tool_collection import ToolCollection |
| 13 | + |
| 14 | +from mcp import ClientSession |
| 15 | +from mcp.client.streamable_http import streamablehttp_client |
| 16 | + |
| 17 | +from a2a.server.agent_execution import AgentExecutor, RequestContext |
| 18 | +from a2a.server.apps import A2AStarletteApplication |
| 19 | +from a2a.server.events.event_queue import EventQueue |
| 20 | +from a2a.server.request_handlers import DefaultRequestHandler |
| 21 | +from a2a.server.tasks import InMemoryTaskStore, TaskUpdater |
| 22 | +from a2a.types import AgentCapabilities, AgentCard, AgentSkill, TaskState, TextPart, SecurityScheme, HTTPAuthSecurityScheme |
| 23 | +from a2a.utils import new_agent_text_message, new_task |
| 24 | + |
| 25 | +from starlette.authentication import AuthCredentials, SimpleUser, AuthenticationBackend |
| 26 | +from starlette.middleware.authentication import AuthenticationMiddleware |
| 27 | + |
| 28 | +from git_issue_agent.auth import on_auth_error, BearerAuthBackend, auth_headers |
| 29 | +from git_issue_agent.config import settings, Settings |
| 30 | +from git_issue_agent.event import Event |
| 31 | +from git_issue_agent.main import GitIssueAgent |
| 32 | + |
| 33 | +logger = logging.getLogger(__name__) |
| 34 | +logging.basicConfig(level=settings.LOG_LEVEL, stream=sys.stdout, format='%(levelname)s: %(message)s') |
| 35 | + |
| 36 | +class BearerAuthBackend(AuthenticationBackend): |
| 37 | + """ Very temporary demo to grab auth token and print it""" |
| 38 | + async def authenticate(self, conn): |
| 39 | + try: |
| 40 | + auth = conn.headers.get("authorization") |
| 41 | + if not auth or not auth.lower().startswith("bearer "): |
| 42 | + print("No bearer token provided") |
| 43 | + return |
| 44 | + token = auth.split(" ", 1)[1] |
| 45 | + print(f"TOKEN: {token}") |
| 46 | + |
| 47 | + # Storing the token as the username - not a real life scenario - just demo-ing the passing of creds |
| 48 | + user = SimpleUser(token) |
| 49 | + return AuthCredentials(["authenticated"]), user |
| 50 | + except Exception as e: |
| 51 | + logger.error("Exception when attempting to obtain user token") |
| 52 | + logger.error(e) |
| 53 | + |
| 54 | + |
| 55 | +def get_agent_card(host: str, port: int): |
| 56 | + """Returns the Agent Card for the AG2 Agent.""" |
| 57 | + capabilities = AgentCapabilities(streaming=True) |
| 58 | + skill = AgentSkill( |
| 59 | + id="github_issue_agent", |
| 60 | + name="Github issue agent", |
| 61 | + description="Answer queries by searching through a given slack server", |
| 62 | + tags=["git", "github", "issues"], |
| 63 | + examples=[ |
| 64 | + "Find me the issues with the most comments in kubernetes/kubernetes", |
| 65 | + "Show all issues assigned to me across any repository", |
| 66 | + ], |
| 67 | + ) |
| 68 | + return AgentCard( |
| 69 | + name="Github issue agent", |
| 70 | + description="Answer queries about Github issues", |
| 71 | + url=f"http://{host}:{port}/", |
| 72 | + version="1.0.0", |
| 73 | + default_input_modes=["text"], |
| 74 | + default_output_modes=["text"], |
| 75 | + capabilities=capabilities, |
| 76 | + skills=[skill], |
| 77 | + securitySchemes={ |
| 78 | + "Bearer": SecurityScheme( |
| 79 | + root=HTTPAuthSecurityScheme( |
| 80 | + type="http", |
| 81 | + scheme="bearer", |
| 82 | + bearerFormat="JWT", |
| 83 | + description="OAuth 2.0 JWT token" |
| 84 | + ) |
| 85 | + ) |
| 86 | + }, |
| 87 | + ) |
| 88 | + |
| 89 | + |
| 90 | +class A2AEvent(Event): |
| 91 | + """ |
| 92 | + A class to handle events for A2A Agent. |
| 93 | +
|
| 94 | + Attributes: |
| 95 | + task_updater (TaskUpdater): The task updater instance. |
| 96 | + """ |
| 97 | + |
| 98 | + def __init__(self, task_updater: TaskUpdater): |
| 99 | + """ |
| 100 | + Initializes the A2AEvent instance. |
| 101 | +
|
| 102 | + Args: |
| 103 | + task_updater (TaskUpdater): The task updater instance. |
| 104 | + """ |
| 105 | + self.task_updater = task_updater |
| 106 | + |
| 107 | + async def emit_event(self, message: str, final: bool = False) -> None: |
| 108 | + """ |
| 109 | + Emits an event with the given message. |
| 110 | +
|
| 111 | + Args: |
| 112 | + message (str): The event message. |
| 113 | + final (bool): Whether the event is final. Defaults to False. |
| 114 | + """ |
| 115 | + logger.info("Emitting event %s", message) |
| 116 | + |
| 117 | + if final: |
| 118 | + parts = [TextPart(text=message)] |
| 119 | + await self.task_updater.add_artifact(parts) |
| 120 | + await self.task_updater.complete() |
| 121 | + else: |
| 122 | + await self.task_updater.update_status( |
| 123 | + TaskState.working, |
| 124 | + new_agent_text_message( |
| 125 | + message, |
| 126 | + self.task_updater.context_id, |
| 127 | + self.task_updater.task_id, |
| 128 | + ), |
| 129 | + ) |
| 130 | + |
| 131 | + |
| 132 | +class GithubExecutor(AgentExecutor): |
| 133 | + """ |
| 134 | + A class to handle research execution for A2A Agent. |
| 135 | + """ |
| 136 | + async def _run_agent(self, |
| 137 | + messages: dict, |
| 138 | + settings: Settings, |
| 139 | + event_emitter: Event, |
| 140 | + toolkit: ToolCollection): |
| 141 | + |
| 142 | + git_issue_agent = GitIssueAgent( |
| 143 | + config=settings, |
| 144 | + eventer=event_emitter, |
| 145 | + mcp_toolkit=toolkit, |
| 146 | + ) |
| 147 | + result = await git_issue_agent.execute(messages) |
| 148 | + await event_emitter.emit_event(result, True) |
| 149 | + |
| 150 | + async def execute(self, context: RequestContext, event_queue: EventQueue): |
| 151 | + """ |
| 152 | + Executes the task. |
| 153 | +
|
| 154 | + Args: |
| 155 | + context (RequestContext): The request context. |
| 156 | + event_queue (EventQueue): The event queue instance. |
| 157 | +
|
| 158 | + Returns: |
| 159 | + None |
| 160 | + """ |
| 161 | + ### |
| 162 | + # commenting this out for now since we have external github MCP. |
| 163 | + # in the future we need to figure out the token exchange story for this scenario |
| 164 | + # |
| 165 | + #if settings.JWKS_URI: |
| 166 | + # user_token = context.call_context.user._user.access_token |
| 167 | + #else: |
| 168 | + user_token = settings.GITHUB_TOKEN |
| 169 | + user_input = [context.get_user_input()] |
| 170 | + task = context.current_task |
| 171 | + if not task: |
| 172 | + task = new_task(context.message) |
| 173 | + await event_queue.enqueue_event(task) |
| 174 | + task_updater = TaskUpdater(event_queue, task.id, task.context_id) |
| 175 | + event_emitter = A2AEvent(task_updater) |
| 176 | + messages = [] |
| 177 | + for message in user_input: |
| 178 | + messages.append( |
| 179 | + { |
| 180 | + "role": "User", |
| 181 | + "content": message, |
| 182 | + } |
| 183 | + ) |
| 184 | + |
| 185 | + # Hook up MCP tools |
| 186 | + try: |
| 187 | + if settings.MCP_URL: |
| 188 | + logging.info("Connecting to MCP server at %s", settings.MCP_URL) |
| 189 | + |
| 190 | + headers = await auth_headers( |
| 191 | + user_token, |
| 192 | + target_audience=settings.TARGET_AUDIENCE, |
| 193 | + target_scopes=settings.TARGET_SCOPES |
| 194 | + ) |
| 195 | + |
| 196 | + server_params = { |
| 197 | + "url": settings.MCP_URL, |
| 198 | + "transport": "streamable-http", |
| 199 | + "headers": headers, |
| 200 | + } |
| 201 | + with MCPServerAdapter(server_params, connect_timeout=60) as mcp_tools: |
| 202 | + # Keep only search and list issue-related tools. |
| 203 | + issue_tools = [ |
| 204 | + tool |
| 205 | + for tool in mcp_tools |
| 206 | + if ("issue" in tool.name.lower() or "label" in tool.name.lower()) and |
| 207 | + ("search" in tool.name.lower() or "list" in tool.name.lower()) |
| 208 | + ] |
| 209 | + |
| 210 | + if not issue_tools: |
| 211 | + raise RuntimeError( |
| 212 | + "No issue-related tools found from the GitHub MCP server. " |
| 213 | + "Ensure your PAT scopes allow issue access and the server is reachable." |
| 214 | + ) |
| 215 | + await self._run_agent(messages, settings, event_emitter, issue_tools) |
| 216 | + else: |
| 217 | + await self._run_agent(messages, settings, |
| 218 | + event_emitter, |
| 219 | + None) |
| 220 | + |
| 221 | + except Exception as e: |
| 222 | + traceback.print_exc() |
| 223 | + await event_emitter.emit_event(f"I'm sorry I was unable to fulfill your request. I encountered the following exception: {str(e)}", True) |
| 224 | + |
| 225 | + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: |
| 226 | + """ |
| 227 | + Not implemented |
| 228 | + """ |
| 229 | + raise Exception("cancel not supported") |
| 230 | + |
| 231 | + |
| 232 | +def run(): |
| 233 | + """ |
| 234 | + Runs the A2A Agent application. |
| 235 | + """ |
| 236 | + agent_card = get_agent_card(host="0.0.0.0", port=settings.SERVICE_PORT) |
| 237 | + |
| 238 | + request_handler = DefaultRequestHandler( |
| 239 | + agent_executor=GithubExecutor(), |
| 240 | + task_store=InMemoryTaskStore(), |
| 241 | + ) |
| 242 | + |
| 243 | + server = A2AStarletteApplication( |
| 244 | + agent_card=agent_card, |
| 245 | + http_handler=request_handler, |
| 246 | + ) |
| 247 | + |
| 248 | + app = server.build() # this returns a Starlette app |
| 249 | + if not settings.JWKS_URI is None: |
| 250 | + logging.info("JWKS_URI is set - using JWT Validation middleware") |
| 251 | + app.add_middleware(AuthenticationMiddleware, backend=BearerAuthBackend(), on_error=on_auth_error) |
| 252 | + |
| 253 | + uvicorn.run(app, host="0.0.0.0", port=settings.SERVICE_PORT) |
0 commit comments