|
12 | 12 | # See the License for the specific language governing permissions and |
13 | 13 | # limitations under the License. |
14 | 14 |
|
| 15 | +import inspect |
15 | 16 | import json |
16 | 17 | import logging |
| 18 | +from collections.abc import AsyncIterator |
17 | 19 | from contextlib import asynccontextmanager |
18 | 20 | from typing import Any |
19 | 21 |
|
|
62 | 64 | logger = logging.getLogger(__name__) |
63 | 65 |
|
64 | 66 |
|
| 67 | +async def _call_lifecycle_handler(handler: Any) -> None: |
| 68 | + result = handler() |
| 69 | + if inspect.isawaitable(result): |
| 70 | + await result |
| 71 | + |
| 72 | + |
| 73 | +@asynccontextmanager |
| 74 | +async def _run_a2a_app_lifespan(a2a_app: Any) -> AsyncIterator[None]: |
| 75 | + router = getattr(a2a_app, "router", None) |
| 76 | + if router is None: |
| 77 | + raise RuntimeError("A2A server app has no router; cannot initialize lifecycle.") |
| 78 | + |
| 79 | + lifespan_context = getattr(router, "lifespan_context", None) |
| 80 | + if lifespan_context is not None: |
| 81 | + async with lifespan_context(a2a_app): |
| 82 | + yield |
| 83 | + return |
| 84 | + |
| 85 | + startup_handlers = tuple(getattr(router, "on_startup", ()) or ()) |
| 86 | + shutdown_handlers = tuple(getattr(router, "on_shutdown", ()) or ()) |
| 87 | + for handler in startup_handlers: |
| 88 | + await _call_lifecycle_handler(handler) |
| 89 | + try: |
| 90 | + yield |
| 91 | + finally: |
| 92 | + for handler in shutdown_handlers: |
| 93 | + await _call_lifecycle_handler(handler) |
| 94 | + |
| 95 | + |
65 | 96 | class AgentKitAgentLoader(BaseAgentLoader): |
66 | 97 | def __init__(self, agent_or_app: BaseAgent | App) -> None: |
67 | 98 | super().__init__() |
@@ -157,11 +188,10 @@ def __init__( |
157 | 188 | async def lifespan(app: FastAPI): |
158 | 189 | # trigger A2A server app startup |
159 | 190 | logger.info( |
160 | | - "Triggering A2A server app startup within API server..." |
| 191 | + "Triggering A2A server app lifespan within API server..." |
161 | 192 | ) |
162 | | - for handler in _a2a_server_app.router.on_startup: |
163 | | - await handler() |
164 | | - yield |
| 193 | + async with _run_a2a_app_lifespan(_a2a_server_app): |
| 194 | + yield |
165 | 195 |
|
166 | 196 | resolved_allow_origins = resolve_agentkit_allow_origins( |
167 | 197 | allow_origins=allow_origins, |
@@ -280,78 +310,6 @@ async def event_generator(): |
280 | 310 | routes.insert(0, routes.pop(i)) |
281 | 311 | break |
282 | 312 |
|
283 | | - @self.app.post("/run_sse") |
284 | | - async def run_agent_sse(req: RunAgentRequest) -> StreamingResponse: |
285 | | - print("my run sse !!!") |
286 | | - # SSE endpoint |
287 | | - session = await self.server.session_service.get_session( |
288 | | - app_name=req.app_name, |
289 | | - user_id=req.user_id, |
290 | | - session_id=req.session_id, |
291 | | - ) |
292 | | - if not session: |
293 | | - raise HTTPException(status_code=404, detail="Session not found") |
294 | | - |
295 | | - # Convert the events to properly formatted SSE |
296 | | - async def event_generator(): |
297 | | - try: |
298 | | - stream_mode = ( |
299 | | - StreamingMode.SSE |
300 | | - if req.streaming |
301 | | - else StreamingMode.NONE |
302 | | - ) |
303 | | - runner = await self.server.get_runner_async(req.app_name) |
304 | | - async with Aclosing( |
305 | | - runner.run_async( |
306 | | - user_id=req.user_id, |
307 | | - session_id=req.session_id, |
308 | | - new_message=req.new_message, |
309 | | - state_delta=req.state_delta, |
310 | | - run_config=RunConfig(streaming_mode=stream_mode), |
311 | | - invocation_id=req.invocation_id, |
312 | | - ) |
313 | | - ) as agen: |
314 | | - async for event in agen: |
315 | | - # ADK Web renders artifacts from `actions.artifactDelta` |
316 | | - # during part processing *and* during action processing |
317 | | - # 1) the original event with `artifactDelta` cleared (content) |
318 | | - # 2) a content-less "action-only" event carrying `artifactDelta` |
319 | | - events_to_stream = [event] |
320 | | - if ( |
321 | | - event.actions.artifact_delta |
322 | | - and event.content |
323 | | - and event.content.parts |
324 | | - ): |
325 | | - content_event = event.model_copy(deep=True) |
326 | | - content_event.actions.artifact_delta = {} |
327 | | - artifact_event = event.model_copy(deep=True) |
328 | | - artifact_event.content = None |
329 | | - events_to_stream = [ |
330 | | - content_event, |
331 | | - artifact_event, |
332 | | - ] |
333 | | - |
334 | | - for event_to_stream in events_to_stream: |
335 | | - sse_event = event_to_stream.model_dump_json( |
336 | | - exclude_none=True, |
337 | | - by_alias=True, |
338 | | - ) |
339 | | - logger.debug( |
340 | | - "Generated event in agent run streaming: %s", |
341 | | - sse_event, |
342 | | - ) |
343 | | - yield f"data: {sse_event}\n\n" |
344 | | - except Exception as e: |
345 | | - logger.exception("Error in event_generator: %s", e) |
346 | | - yield f"data: {json.dumps({'error': str(e)})}\n\n" |
347 | | - |
348 | | - # Returns a streaming response with the proper media type for SSE |
349 | | - |
350 | | - return StreamingResponse( |
351 | | - event_generator(), |
352 | | - media_type="text/event-stream", |
353 | | - ) |
354 | | - |
355 | 313 | # Attach ASGI middleware for unified telemetry across all routes |
356 | 314 | self.app.add_middleware(AgentkitTelemetryHTTPMiddleware) |
357 | 315 |
|
|
0 commit comments