Skip to content

Commit 751ae20

Browse files
author
Krzysztof Dziedzic
committed
test: setup itk resubscribe tests
1 parent c0c6c08 commit 751ae20

2 files changed

Lines changed: 218 additions & 53 deletions

File tree

itk/main.py

Lines changed: 200 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -3,17 +3,20 @@
33
import base64
44
import logging
55
import os
6+
import signal
67
import uuid
78

89
import grpc
910
import httpx
1011
import uvicorn
1112

1213
from fastapi import FastAPI
14+
from typing import Any
1315

1416
from pyproto import instruction_pb2
1517

16-
from a2a.client import ClientConfig, create_client
18+
from a2a.client import Client, ClientConfig, create_client
19+
from a2a.client.errors import A2AClientError
1720
from a2a.compat.v0_3 import a2a_v0_3_pb2_grpc
1821
from a2a.compat.v0_3.grpc_handler import CompatGrpcHandler
1922
from a2a.server.agent_execution import AgentExecutor, RequestContext
@@ -31,14 +34,17 @@
3134
InMemoryPushNotificationConfigStore,
3235
)
3336
from a2a.server.tasks.inmemory_task_store import InMemoryTaskStore
37+
from a2a.server.context import ServerCallContext
3438
from a2a.types import a2a_pb2_grpc
3539
from a2a.types.a2a_pb2 import (
3640
AgentCapabilities,
3741
AgentCard,
3842
AgentInterface,
43+
CancelTaskRequest,
3944
Message,
4045
Part,
4146
SendMessageRequest,
47+
SubscribeToTaskRequest,
4248
Task,
4349
TaskState,
4450
TaskStatus,
@@ -98,6 +104,95 @@ def extract_instruction(
98104
return None
99105

100106

107+
def _extract_text_from_event(event: Any) -> list[str]:
108+
"""Extracts text parts from an event's message."""
109+
if isinstance(event, tuple):
110+
results = []
111+
for item in event:
112+
results.extend(_extract_text_from_event(item))
113+
return results
114+
115+
message = None
116+
if hasattr(event, 'HasField'):
117+
if event.HasField('message'):
118+
message = event.message
119+
elif event.HasField('task') and event.task.status.HasField('message'):
120+
message = event.task.status.message
121+
elif event.HasField(
122+
'status_update'
123+
) and event.status_update.status.HasField('message'):
124+
message = event.status_update.status.message
125+
126+
results = []
127+
if message:
128+
results.extend(part.text for part in message.parts if part.text)
129+
return results
130+
131+
132+
async def _handle_call_agent_with_resubscribe(
133+
client: Client, request: SendMessageRequest
134+
) -> list[str]:
135+
"""Handles the send-disconnect-resubscribe flow."""
136+
results = []
137+
logger.info('Executing re-subscribe behavior')
138+
agen = client.send_message(request)
139+
task_id = None
140+
141+
async for event in agen:
142+
logger.info('Event before disconnect: %s', event)
143+
if event.HasField('task'):
144+
task_id = event.task.id
145+
elif event.HasField('status_update'):
146+
task_id = event.status_update.task_id
147+
break
148+
149+
await agen.aclose()
150+
logger.info('Disconnected from task %s. Now re-subscribing.', task_id)
151+
152+
resub_agen = client.subscribe(SubscribeToTaskRequest(id=task_id))
153+
154+
task_obj = None
155+
finished = False
156+
async for event in resub_agen:
157+
logger.info('Event after re-subscribe: %s', event)
158+
if hasattr(event, 'task') or (
159+
hasattr(event, 'HasField') and event.HasField('task')
160+
):
161+
task_obj = event.task
162+
163+
extracted_text = _extract_text_from_event(event)
164+
for text in extracted_text:
165+
processed_text = text.replace('task-finished', '')
166+
results.append(processed_text)
167+
if any('task-finished' in text for text in extracted_text):
168+
logger.info(
169+
'Received task-finished after re-subscribe, breaking loop.'
170+
)
171+
finished = True
172+
break
173+
174+
if not results and task_obj and hasattr(task_obj, 'history'):
175+
logger.info('Results empty after loop, reading from history.')
176+
for msg in task_obj.history:
177+
if msg.role in {'ROLE_AGENT', 'agent'}:
178+
results.extend(
179+
part.text.replace('task-finished', '')
180+
for part in msg.parts
181+
if part.text
182+
)
183+
184+
if not finished:
185+
logger.info('Canceling task %s after retrieval.', task_id)
186+
try:
187+
await client.cancel_task(CancelTaskRequest(id=task_id))
188+
logger.info('Task cancelled successfully: %s', task_id)
189+
except A2AClientError:
190+
logger.exception('Failed to cancel task %s', task_id)
191+
raise
192+
193+
return results
194+
195+
101196
def wrap_instruction_to_request(inst: instruction_pb2.Instruction) -> Message:
102197
"""Wraps an Instruction proto into an A2A Message."""
103198
inst_bytes = inst.SerializeToString()
@@ -129,18 +224,22 @@ async def handle_call_agent(
129224
'GRPC': TransportProtocol.GRPC,
130225
}
131226

132-
selected_transport = transport_map.get(call.transport.upper())
227+
selected_transport = transport_map.get(
228+
call.transport.upper(), TransportProtocol.JSONRPC
229+
)
133230
if selected_transport is None:
134231
raise ValueError(f'Unsupported transport: {call.transport}')
135232

136233
config = ClientConfig()
137-
config.httpx_client = httpx.AsyncClient(timeout=30.0)
138234
config.grpc_channel_factory = grpc.aio.insecure_channel
139235
config.supported_protocol_bindings = [selected_transport]
140236
config.streaming = call.streaming or (
141237
selected_transport == TransportProtocol.GRPC
142238
)
143239

240+
if call.HasField('resubscribe') and not config.streaming:
241+
raise ValueError('Re-subscription requires streaming to be enabled')
242+
144243
if call.HasField('push_notification'):
145244
url = call.push_notification.url
146245
if not url:
@@ -152,44 +251,45 @@ async def handle_call_agent(
152251
token='itk-token', # noqa: S106
153252
)
154253

155-
try:
156-
client = await create_client(
157-
call.agent_card_uri,
158-
client_config=config,
159-
)
254+
async with httpx.AsyncClient(timeout=30.0) as httpx_client:
255+
config.httpx_client = httpx_client
256+
try:
257+
client = await create_client(
258+
call.agent_card_uri,
259+
client_config=config,
260+
)
160261

161-
# Wrap nested instruction
162-
nested_msg = wrap_instruction_to_request(call.instruction)
163-
request = SendMessageRequest(message=nested_msg)
262+
# Wrap nested instruction
263+
nested_msg = wrap_instruction_to_request(call.instruction)
264+
request = SendMessageRequest(message=nested_msg)
164265

165-
results = []
166-
async for event in client.send_message(request):
167-
# Event is streaming response and task
168-
logger.info('Event: %s', event)
169-
stream_resp = event
170-
171-
message = None
172-
if stream_resp.HasField('message'):
173-
message = stream_resp.message
174-
elif stream_resp.HasField(
175-
'task'
176-
) and stream_resp.task.status.HasField('message'):
177-
message = stream_resp.task.status.message
178-
elif stream_resp.HasField(
179-
'status_update'
180-
) and stream_resp.status_update.status.HasField('message'):
181-
message = stream_resp.status_update.status.message
182-
183-
if message:
184-
results.extend(part.text for part in message.parts if part.text)
185-
186-
except Exception as e:
187-
logger.exception('Failed to call outbound agent')
188-
raise RuntimeError(
189-
f'Outbound call to {call.agent_card_uri} failed: {e!s}'
190-
) from e
191-
else:
192-
return results
266+
results = []
267+
268+
if call.HasField('resubscribe'):
269+
results.extend(
270+
await _handle_call_agent_with_resubscribe(client, request)
271+
)
272+
else:
273+
async for event in client.send_message(request):
274+
logger.info('Event: %s', event)
275+
results.extend(_extract_text_from_event(event))
276+
277+
except Exception as e:
278+
logger.exception('Failed to call outbound agent')
279+
raise RuntimeError(
280+
f'Outbound call to {call.agent_card_uri} failed: {e!s}'
281+
) from e
282+
else:
283+
return results
284+
285+
286+
def _should_hold(inst: instruction_pb2.Instruction) -> bool:
287+
"""Recursively checks if any part of the instruction requests holding the task."""
288+
if inst.HasField('return_response') and inst.return_response.hold_task:
289+
return True
290+
if inst.HasField('steps'):
291+
return any(_should_hold(step) for step in inst.steps.instructions)
292+
return False
193293

194294

195295
async def handle_instruction(
@@ -245,23 +345,58 @@ async def execute(
245345
)
246346
return
247347

348+
should_hold_task = _should_hold(instruction)
349+
248350
try:
249351
logger.info('Instruction: %s', instruction)
250352
results = await handle_instruction(instruction)
353+
251354
response_text = '\n'.join(results)
252355
logger.info('Response: %s', response_text)
253-
await task_updater.update_status(
254-
TaskState.TASK_STATE_COMPLETED,
255-
message=task_updater.new_agent_message(
256-
[Part(text=response_text)]
257-
),
258-
)
259-
logger.info('Task %s completed', context.task_id)
260-
except Exception as e:
356+
357+
if should_hold_task:
358+
logger.info('Holding task %s as requested', context.task_id)
359+
# Emitted event: response + task-finished
360+
logger.info(
361+
'Emitting response and task-finished for held task %s',
362+
context.task_id,
363+
)
364+
await task_updater.update_status(
365+
TaskState.TASK_STATE_WORKING,
366+
message=task_updater.new_agent_message(
367+
[Part(text=response_text + '\n' + 'task-finished')]
368+
),
369+
)
370+
await asyncio.sleep(2)
371+
372+
# Continue emitting "task-finished" every 2 seconds
373+
try:
374+
while True:
375+
logger.info(
376+
'Emitting periodic status update for held task %s',
377+
context.task_id,
378+
)
379+
await task_updater.update_status(
380+
TaskState.TASK_STATE_WORKING,
381+
message=None,
382+
)
383+
await asyncio.sleep(2)
384+
except asyncio.CancelledError:
385+
logger.info('Task %s cancelled', context.task_id)
386+
return
387+
else:
388+
await task_updater.update_status(
389+
TaskState.TASK_STATE_COMPLETED,
390+
message=task_updater.new_agent_message(
391+
[Part(text=response_text)]
392+
),
393+
)
394+
logger.info('Task %s completed', context.task_id)
395+
except Exception:
261396
logger.exception('Error during instruction handling')
262397
await task_updater.update_status(
263398
TaskState.TASK_STATE_FAILED,
264-
message=task_updater.new_agent_message([Part(text=str(e))]),
399+
message=None,
265400
)
266401

267402
async def cancel(
@@ -325,19 +460,19 @@ async def main_async(http_port: int, grpc_port: int) -> None:
325460
name='ITK v10 Agent',
326461
description='Python agent using SDK 1.0.',
327462
version='1.0.0',
328-
capabilities=AgentCapabilities(
329-
streaming=True, push_notifications=True, extended_agent_card=True
330-
),
463+
capabilities=AgentCapabilities(streaming=True),
331464
default_input_modes=['text/plain'],
332465
default_output_modes=['text/plain'],
333466
supported_interfaces=interfaces,
334467
)
335468

336469
task_store = InMemoryTaskStore()
337470
push_config_store = InMemoryPushNotificationConfigStore()
471+
httpx_client = httpx.AsyncClient()
338472
push_sender = BasePushNotificationSender(
339-
httpx_client=httpx.AsyncClient(),
473+
httpx_client=httpx_client,
340474
config_store=push_config_store,
475+
context=ServerCallContext(),
341476
)
342477

343478
handler = DefaultRequestHandler(
@@ -396,10 +531,22 @@ async def main_async(http_port: int, grpc_port: int) -> None:
396531
)
397532

398533
config = uvicorn.Config(
399-
app, host='127.0.0.1', port=http_port, log_level=log_level_str.lower()
534+
app, host='127.0.0.1', port=http_port, log_level='info'
400535
)
401536
uvicorn_server = uvicorn.Server(config)
402537

538+
# Signal handling
539+
loop = asyncio.get_running_loop()
540+
541+
async def shutdown() -> None:
542+
logger.info('Shutting down...')
543+
uvicorn_server.should_exit = True
544+
await server.stop(5)
545+
await httpx_client.aclose()
546+
547+
for sig in (signal.SIGINT, signal.SIGTERM):
548+
loop.add_signal_handler(sig, lambda: asyncio.create_task(shutdown()))
549+
403550
await uvicorn_server.serve()
404551

405552

itk/run_itk.sh

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -163,6 +163,24 @@ RESPONSE=$(curl -s -X POST http://127.0.0.1:8000/run \
163163
"edges": ["0->1", "0->2", "1->0", "2->0"],
164164
"protocols": ["http_json"],
165165
"behavior": "push_notification"
166+
},
167+
{
168+
"name": "Resubscribe Test - JSONRPC",
169+
"sdks": ["current", "python_v10", "python_v03", "go_v10", "go_v03"],
170+
"traversal": "euler",
171+
"edges": ["0->1", "0->2", "0->3", "0->4", "1->0", "2->0", "3->0", "4->0"],
172+
"protocols": ["jsonrpc"],
173+
"streaming": true,
174+
"behavior": "resubscribe"
175+
},
176+
{
177+
"name": "Resubscribe Test - Python & Go Non-JSONRPC Protocols",
178+
"sdks": ["current", "python_v10", "python_v03", "go_v10"],
179+
"traversal": "euler",
180+
"edges": ["0->1", "0->2", "0->3", "1->0", "2->0", "3->0"],
181+
"protocols": ["grpc", "http_json"],
182+
"streaming": true,
183+
"behavior": "resubscribe"
166184
}
167185
]
168186
}')

0 commit comments

Comments
 (0)