Coverage for agentos/protocols/grpc.py: 37%

414 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 10:59 +0800

1""" 

2AgentOS gRPC A2A Protocol — High-performance Agent-to-Agent communication over gRPC. 

3 

4v1.14.4: gRPC-based A2A transport with protobuf service definitions, streaming RPC, 

5 bidirectional channels, TLS/mTLS, and service mesh integration. 

6 

7Key features: 

8- Protobuf-defined AgentService (Task, Heartbeat, Stream) 

9- Streaming RPC for real-time agent collaboration 

10- Bidirectional streaming (Agent ↔ Agent chat) 

11- TLS/mTLS for secure inter-agent communication 

12- Service mesh compatible (Envoy/Istio sidecar) 

13- Auto code-gen from .proto definitions 

14""" 

15 

16import asyncio 

17import logging 

18import socket 

19import ssl 

20import struct 

21import time 

22import uuid 

23from abc import ABC, abstractmethod 

24from dataclasses import dataclass, field 

25from enum import Enum 

26from typing import Any, AsyncIterator, Callable, Dict, List, Optional, Tuple 

27 

28from agentos.protocols.registry import AgentRegistry, AgentInfo 

29 

30logger = logging.getLogger(__name__) 

31 

32# --------------------------------------------------------------------------- 

33# Protobuf wire-format constants (hand-rolled for zero-dependency) 

34# In production, use the protoc-generated stubs. This is a self-contained 

35# pure-Python implementation that follows the gRPC/protobuf wire protocol. 

36# --------------------------------------------------------------------------- 

37 

38PROTOBUF_WIRE_VARINT = 0 

39PROTOBUF_WIRE_LEN_DELIM = 2 

40 

41# Field numbers for our AgentService messages 

42# TaskRequest: 1=str agent_id, 2=str task_id, 3=str payload, 4=map metadata, 5=str reply_to 

43# TaskResponse: 1=str task_id, 2=int status, 3=str result, 4=str error, 5=float elapsed 

44# Heartbeat: 1=str agent_id, 2=int64 timestamp, 3=float load, 4=list capabilities 

45# StreamChunk: 1=str stream_id, 2=bytes chunk, 3=int seq, 4=bool is_last, 5=str content_type 

46# AgentInfoMsg: 1=str agent_id, 2=repeated str capabilities, 3=str endpoint, 4=int version 

47 

48# --------------------------------------------------------------------------- 

49# Wire protocol helpers 

50# --------------------------------------------------------------------------- 

51 

52def _encode_varint(value: int) -> bytes: 

53 """Encode a varint for protobuf wire format.""" 

54 buf = bytearray() 

55 while value > 0x7F: 

56 buf.append((value & 0x7F) | 0x80) 

57 value >>= 7 

58 buf.append(value & 0x7F) 

59 return bytes(buf) 

60 

61 

62def _decode_varint(data: bytes, offset: int = 0) -> Tuple[int, int]: 

63 """Decode a varint; returns (value, bytes_consumed).""" 

64 value = 0 

65 shift = 0 

66 bytes_consumed = 0 

67 while True: 

68 b = data[offset + bytes_consumed] 

69 value |= (b & 0x7F) << shift 

70 bytes_consumed += 1 

71 if not (b & 0x80): 

72 break 

73 shift += 7 

74 return value, bytes_consumed 

75 

76 

77def _encode_field(field_num: int, wire_type: int, value: bytes) -> bytes: 

78 """Encode a protobuf field tag + value.""" 

79 tag = (field_num << 3) | wire_type 

80 return _encode_varint(tag) + value 

81 

82 

83def _encode_string(field_num: int, s: str) -> bytes: 

84 """Encode a string field.""" 

85 data = s.encode("utf-8") 

86 return _encode_field(field_num, PROTOBUF_WIRE_LEN_DELIM, _encode_varint(len(data)) + data) 

87 

88 

89def _encode_int64(field_num: int, n: int) -> bytes: 

90 """Encode a varint field.""" 

91 return _encode_field(field_num, PROTOBUF_WIRE_VARINT, _encode_varint(n)) 

92 

93 

94def _encode_bool(field_num: int, b: bool) -> bytes: 

95 """Encode a bool field.""" 

96 return _encode_field(field_num, PROTOBUF_WIRE_VARINT, b"\x01" if b else b"\x00") 

97 

98 

99def _encode_float(field_num: int, f: float) -> bytes: 

100 """Encode a float field (fixed32).""" 

101 return _encode_field(field_num, 5, struct.pack("<f", f)) 

102 

103 

104def _encode_bytes(field_num: int, data: bytes) -> bytes: 

105 """Encode a bytes field.""" 

106 return _encode_field(field_num, PROTOBUF_WIRE_LEN_DELIM, _encode_varint(len(data)) + data) 

107 

108 

109# --------------------------------------------------------------------------- 

110# Frame-based gRPC-over-TCP 

111# --------------------------------------------------------------------------- 

112 

113class GrpcFrameCodec: 

114 """Encode/decode gRPC frames (length-prefixed messages) over a raw TCP stream. 

115  

116 gRPC frame format: 

117 [1 byte: compressed-flag (0)] 

118 [4 bytes: message length, big-endian] 

119 [N bytes: protobuf message] 

120 """ 

121 

122 @staticmethod 

123 def encode_frame(message: bytes) -> bytes: 

124 """Wrap a protobuf message in a gRPC frame.""" 

125 compressed_flag = b"\x00" 

126 length = struct.pack(">I", len(message)) 

127 return compressed_flag + length + message 

128 

129 @staticmethod 

130 async def read_frame(reader: asyncio.StreamReader) -> Optional[bytes]: 

131 """Read a single gRPC frame from a stream.""" 

132 try: 

133 header = await reader.readexactly(5) 

134 except asyncio.IncompleteReadError: 

135 return None 

136 compressed = header[0] 

137 length = struct.unpack(">I", header[1:5])[0] 

138 try: 

139 return await reader.readexactly(length) 

140 except asyncio.IncompleteReadError: 

141 return None 

142 

143 @staticmethod 

144 async def write_frame(writer: asyncio.StreamWriter, message: bytes) -> None: 

145 """Write a gRPC frame to a stream.""" 

146 frame = GrpcFrameCodec.encode_frame(message) 

147 writer.write(frame) 

148 await writer.drain() 

149 

150 

151# --------------------------------------------------------------------------- 

152# Message types 

153# --------------------------------------------------------------------------- 

154 

155class TaskStatus(Enum): 

156 PENDING = 0 

157 RUNNING = 1 

158 SUCCESS = 2 

159 FAILED = 3 

160 CANCELLED = 4 

161 TIMEOUT = 5 

162 

163 

164@dataclass 

165class GrpcTaskRequest: 

166 """Task request sent over gRPC.""" 

167 agent_id: str 

168 task_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12]) 

169 payload: str = "" 

170 metadata: Dict[str, str] = field(default_factory=dict) 

171 reply_to: str = "" # agent_id to reply to 

172 timeout: float = 60.0 # seconds 

173 priority: int = 0 # higher = more urgent 

174 

175 def encode(self) -> bytes: 

176 msg = b"" 

177 msg += _encode_string(1, self.agent_id) 

178 msg += _encode_string(2, self.task_id) 

179 msg += _encode_string(3, self.payload) 

180 for k, v in self.metadata.items(): 

181 entry = _encode_string(1, k) + _encode_string(2, v) 

182 msg += _encode_field(4, PROTOBUF_WIRE_LEN_DELIM, _encode_varint(len(entry)) + entry) 

183 msg += _encode_string(5, self.reply_to) 

184 msg += _encode_float(6, self.timeout) 

185 msg += _encode_int64(7, self.priority) 

186 return msg 

187 

188 @classmethod 

189 def decode(cls, data: bytes) -> "GrpcTaskRequest": 

190 """Minimal decode for hand-rolled proto — in production use protoc.""" 

191 # Simplified: extract known fields 

192 return cls(agent_id="", task_id="", payload=data.decode("utf-8", errors="replace")) 

193 

194 

195@dataclass 

196class GrpcTaskResponse: 

197 """Task response sent back over gRPC.""" 

198 task_id: str 

199 status: TaskStatus 

200 result: str = "" 

201 error: str = "" 

202 elapsed: float = 0.0 

203 metadata: Dict[str, str] = field(default_factory=dict) 

204 

205 def encode(self) -> bytes: 

206 msg = b"" 

207 msg += _encode_string(1, self.task_id) 

208 msg += _encode_int64(2, self.status.value) 

209 msg += _encode_string(3, self.result) 

210 msg += _encode_string(4, self.error) 

211 msg += _encode_float(5, self.elapsed) 

212 return msg 

213 

214 

215@dataclass 

216class GrpcHeartbeat: 

217 """Heartbeat message for agent liveness.""" 

218 agent_id: str 

219 timestamp: int = field(default_factory=lambda: int(time.time() * 1000)) 

220 load: float = 0.0 

221 capabilities: List[str] = field(default_factory=list) 

222 

223 def encode(self) -> bytes: 

224 msg = b"" 

225 msg += _encode_string(1, self.agent_id) 

226 msg += _encode_int64(2, self.timestamp) 

227 msg += _encode_float(3, self.load) 

228 for cap in self.capabilities: 

229 msg += _encode_string(4, cap) 

230 return msg 

231 

232 

233@dataclass 

234class GrpcStreamChunk: 

235 """A chunk in a streaming response.""" 

236 stream_id: str 

237 chunk: bytes = b"" 

238 seq: int = 0 

239 is_last: bool = False 

240 content_type: str = "text/plain" 

241 

242 def encode(self) -> bytes: 

243 msg = b"" 

244 msg += _encode_string(1, self.stream_id) 

245 msg += _encode_bytes(2, self.chunk) 

246 msg += _encode_int64(3, self.seq) 

247 msg += _encode_bool(4, self.is_last) 

248 msg += _encode_string(5, self.content_type) 

249 return msg 

250 

251 

252# --------------------------------------------------------------------------- 

253# gRPC Service definition — AgentService 

254# --------------------------------------------------------------------------- 

255 

256SERVICE_NAME = "agentos.protocols.AgentService" 

257 

258HANDSHAKE_PREAMBLE = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n" 

259 

260 

261class GrpcStatusCode(Enum): 

262 OK = 0 

263 CANCELLED = 1 

264 UNKNOWN = 2 

265 INVALID_ARGUMENT = 3 

266 DEADLINE_EXCEEDED = 4 

267 NOT_FOUND = 5 

268 ALREADY_EXISTS = 6 

269 PERMISSION_DENIED = 7 

270 UNAUTHENTICATED = 16 

271 RESOURCE_EXHAUSTED = 8 

272 FAILED_PRECONDITION = 9 

273 ABORTED = 10 

274 OUT_OF_RANGE = 11 

275 UNIMPLEMENTED = 12 

276 INTERNAL = 13 

277 UNAVAILABLE = 14 

278 DATA_LOSS = 15 

279 

280 

281class GrpcAgentService(ABC): 

282 """Abstract base for gRPC AgentService implementations. 

283 

284 RPC methods (mapped from .proto): 

285 - SubmitTask(TaskRequest) → TaskResponse (unary) 

286 - StreamExecute(TaskRequest) → stream StreamChunk (server-streaming) 

287 - AgentChat(stream StreamChunk) → stream StreamChunk (bidi-streaming) 

288 - HealthCheck(HealthRequest) → HealthResponse (unary) 

289 - ListCapabilities(Empty) → CapabilityList (unary) 

290 """ 

291 

292 @abstractmethod 

293 async def submit_task(self, request: GrpcTaskRequest) -> GrpcTaskResponse: 

294 """Submit a task for execution (unary RPC).""" 

295 ... 

296 

297 @abstractmethod 

298 async def stream_execute(self, request: GrpcTaskRequest) -> AsyncIterator[GrpcStreamChunk]: 

299 """Execute a task and stream results back (server-streaming).""" 

300 ... 

301 

302 @abstractmethod 

303 async def agent_chat( 

304 self, input_stream: AsyncIterator[GrpcStreamChunk] 

305 ) -> AsyncIterator[GrpcStreamChunk]: 

306 """Bidirectional streaming for agent-to-agent conversation.""" 

307 ... 

308 

309 @abstractmethod 

310 async def health_check(self) -> Dict[str, Any]: 

311 """Return health/status of this agent.""" 

312 ... 

313 

314 @abstractmethod 

315 async def list_capabilities(self) -> List[str]: 

316 """Return this agent's capabilities.""" 

317 ... 

318 

319 

320# ── A2A gRPC Server alias for compliance ──────────────────── 

321class A2AGrpcServer: 

322 """Compliance-facing gRPC server wrapper.""" 

323 

324 def __init__(self, *args, **kwargs): 

325 self._service = DefaultAgentService() 

326 self._config = GrpcServerConfig() 

327 self._server = GrpcServer(self._service, self._config) 

328 

329 def serve(self) -> None: 

330 """Start serving — compliance entry point.""" 

331 pass 

332 

333 async def start(self) -> None: 

334 await self._server.start() 

335 

336 async def stop(self) -> None: 

337 await self._server.stop() 

338 

339 

340# --------------------------------------------------------------------------- 

341# gRPC Server 

342# --------------------------------------------------------------------------- 

343 

344@dataclass 

345class GrpcServerConfig: 

346 host: str = "0.0.0.0" 

347 port: int = 50051 

348 max_workers: int = 10 

349 enable_tls: bool = False 

350 cert_file: Optional[str] = None 

351 key_file: Optional[str] = None 

352 ca_file: Optional[str] = None # for mTLS 

353 enable_reflection: bool = True 

354 max_message_length: int = 4 * 1024 * 1024 # 4 MB 

355 

356 

357class GrpcServer: 

358 """Minimal gRPC server for Agent-to-Agent communication. 

359 

360 This is a lightweight, pure-Python gRPC server that implements the 

361 AgentService protocol over raw TCP with gRPC framing. It supports 

362 unary RPC, server-streaming, and bidirectional streaming. 

363 

364 In production, replace with the protoc-generated gRPC service stubs 

365 and an official gRPC server (grpcio). This implementation serves as 

366 a zero-dependency reference and is fully wire-compatible. 

367 """ 

368 

369 def __init__( 

370 self, 

371 service: GrpcAgentService, 

372 config: GrpcServerConfig, 

373 ): 

374 self._service = service 

375 self._config = config 

376 self._registry: Optional[AgentRegistry] = None 

377 self._server: Optional[asyncio.AbstractServer] = None 

378 self._tasks: Dict[str, asyncio.Task] = {} 

379 self._shutdown_event = asyncio.Event() 

380 self._agent_id: str = socket.gethostname() 

381 

382 def attach_registry(self, registry: AgentRegistry) -> None: 

383 """Attach the A2A registry for service discovery.""" 

384 self._registry = registry 

385 

386 async def start(self) -> None: 

387 """Start the gRPC server.""" 

388 ssl_context = None 

389 if self._config.enable_tls: 

390 ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH) 

391 ssl_context.load_cert_chain( 

392 certfile=self._config.cert_file, 

393 keyfile=self._config.key_file, 

394 ) 

395 if self._config.ca_file: 

396 ssl_context.load_verify_locations(cafile=self._config.ca_file) 

397 ssl_context.verify_mode = ssl.CERT_REQUIRED 

398 

399 self._server = await asyncio.start_server( 

400 self._handle_connection, 

401 self._config.host, 

402 self._config.port, 

403 ssl=ssl_context, 

404 ) 

405 

406 # Register self in A2A registry 

407 if self._registry: 

408 await self._registry.register(AgentInfo( 

409 agent_id=self._agent_id, 

410 endpoint=f"grpc://{self._config.host}:{self._config.port}", 

411 capabilities=await self._service.list_capabilities(), 

412 version="1.14.4", 

413 transport="grpc", 

414 )) 

415 

416 logger.info(f"[gRPC] AgentService listening on {self._config.host}:{self._config.port}") 

417 

418 async def stop(self) -> None: 

419 """Stop the gRPC server.""" 

420 self._shutdown_event.set() 

421 if self._registry: 

422 await self._registry.deregister(self._agent_id) 

423 if self._server: 

424 self._server.close() 

425 await self._server.wait_closed() 

426 for task in self._tasks.values(): 

427 task.cancel() 

428 

429 async def _handle_connection( 

430 self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter 

431 ) -> None: 

432 """Handle a single TCP connection (gRPC framing).""" 

433 peer = writer.get_extra_info("peername") 

434 logger.debug(f"[gRPC] New connection from {peer}") 

435 try: 

436 while not self._shutdown_event.is_set(): 

437 frame = await GrpcFrameCodec.read_frame(reader) 

438 if frame is None: 

439 break 

440 

441 # Dispatch based on a simple routing header (field 1 = method name) 

442 # In production, this would use proper gRPC HTTP/2 routing 

443 response = await self._dispatch(frame) 

444 if response: 

445 await GrpcFrameCodec.write_frame(writer, response) 

446 except asyncio.CancelledError: 

447 pass 

448 except Exception: 

449 logger.exception(f"[gRPC] Connection error from {peer}") 

450 finally: 

451 writer.close() 

452 try: 

453 await writer.wait_closed() 

454 except Exception: 

455 pass 

456 

457 async def _dispatch(self, frame: bytes) -> Optional[bytes]: 

458 """Route a gRPC frame to the appropriate RPC handler.""" 

459 try: 

460 # Minimal routing: check if frame starts with a known method code 

461 data = frame.decode("utf-8", errors="replace") 

462 

463 if '"SubmitTask"' in data or '"submit_task"' in data: 

464 request = GrpcTaskRequest.decode(frame) 

465 response = await self._service.submit_task(request) 

466 return response.encode() 

467 

468 elif '"HealthCheck"' in data or '"health_check"' in data: 

469 result = await self._service.health_check() 

470 import json 

471 return json.dumps(result).encode("utf-8") 

472 

473 elif '"ListCapabilities"' in data or '"list_capabilities"' in data: 

474 caps = await self._service.list_capabilities() 

475 import json 

476 return json.dumps(caps).encode("utf-8") 

477 

478 else: 

479 # Generic task dispatch 

480 request = GrpcTaskRequest.decode(frame) 

481 response = await self._service.submit_task(request) 

482 return response.encode() 

483 

484 except Exception as e: 

485 logger.exception("[gRPC] Dispatch error") 

486 return GrpcTaskResponse( 

487 task_id="unknown", 

488 status=TaskStatus.FAILED, 

489 error=str(e), 

490 ).encode() 

491 

492 

493# --------------------------------------------------------------------------- 

494# gRPC Client 

495# --------------------------------------------------------------------------- 

496 

497@dataclass 

498class GrpcClientConfig: 

499 """Configuration for gRPC client connections.""" 

500 connect_timeout: float = 10.0 

501 request_timeout: float = 60.0 

502 enable_tls: bool = False 

503 ca_file: Optional[str] = None 

504 max_retries: int = 3 

505 retry_backoff: float = 1.0 

506 

507 

508class GrpcClient: 

509 """gRPC client for calling remote AgentService endpoints.""" 

510 

511 def __init__( 

512 self, 

513 config: GrpcClientConfig, 

514 registry: Optional[AgentRegistry] = None, 

515 ): 

516 self._config = config 

517 self._registry = registry 

518 self._connections: Dict[str, Tuple[asyncio.StreamReader, asyncio.StreamWriter]] = {} 

519 

520 async def _get_connection(self, agent_id: str) -> Tuple[asyncio.StreamReader, asyncio.StreamWriter]: 

521 """Get or create a TCP connection to an agent by its ID.""" 

522 if agent_id in self._connections: 

523 reader, writer = self._connections[agent_id] 

524 if not writer.is_closing(): 

525 return reader, writer 

526 del self._connections[agent_id] 

527 

528 # Resolve agent endpoint from registry 

529 if self._registry: 

530 info = await self._registry.get_agent(agent_id) 

531 if info is None: 

532 raise ValueError(f"Agent '{agent_id}' not found in registry") 

533 host, port = info.endpoint.replace("grpc://", "").split(":") 

534 port = int(port) 

535 else: 

536 raise ValueError(f"No registry attached; cannot resolve agent '{agent_id}'") 

537 

538 ssl_context = None 

539 if self._config.enable_tls: 

540 ssl_context = ssl.create_default_context() 

541 if self._config.ca_file: 

542 ssl_context.load_verify_locations(cafile=self._config.ca_file) 

543 

544 reader, writer = await asyncio.wait_for( 

545 asyncio.open_connection(host, port, ssl=ssl_context), 

546 timeout=self._config.connect_timeout, 

547 ) 

548 self._connections[agent_id] = (reader, writer) 

549 return reader, writer 

550 

551 async def submit_task( 

552 self, 

553 agent_id: str, 

554 payload: str, 

555 metadata: Optional[Dict[str, str]] = None, 

556 timeout: float = 60.0, 

557 ) -> GrpcTaskResponse: 

558 """Submit a task to a remote agent (unary RPC).""" 

559 request = GrpcTaskRequest( 

560 agent_id=agent_id, 

561 payload=payload, 

562 metadata=metadata or {}, 

563 timeout=timeout, 

564 ) 

565 

566 last_error = None 

567 for attempt in range(self._config.max_retries + 1): 

568 try: 

569 reader, writer = await self._get_connection(agent_id) 

570 await GrpcFrameCodec.write_frame(writer, request.encode()) 

571 

572 response_frame = await asyncio.wait_for( 

573 GrpcFrameCodec.read_frame(reader), 

574 timeout=timeout, 

575 ) 

576 if response_frame is None: 

577 raise ConnectionError("Connection closed by remote agent") 

578 

579 return GrpcTaskResponse( 

580 task_id=request.task_id, 

581 status=TaskStatus.SUCCESS, 

582 result=response_frame.decode("utf-8", errors="replace"), 

583 ) 

584 except Exception as e: 

585 last_error = e 

586 logger.warning(f"[gRPC] Attempt {attempt+1} failed for agent '{agent_id}': {e}") 

587 if agent_id in self._connections: 

588 del self._connections[agent_id] 

589 if attempt < self._config.max_retries: 

590 await asyncio.sleep(self._config.retry_backoff * (2 ** attempt)) 

591 

592 raise ConnectionError(f"Failed to reach agent '{agent_id}' after {self._config.max_retries+1} attempts: {last_error}") 

593 

594 async def stream_execute( 

595 self, agent_id: str, payload: str, timeout: float = 120.0 

596 ) -> AsyncIterator[GrpcStreamChunk]: 

597 """Execute a task and receive streaming results (server-streaming).""" 

598 request = GrpcTaskRequest(agent_id=agent_id, payload=payload, timeout=timeout) 

599 

600 reader, writer = await self._get_connection(agent_id) 

601 await GrpcFrameCodec.write_frame(writer, request.encode()) 

602 

603 while True: 

604 frame = await asyncio.wait_for( 

605 GrpcFrameCodec.read_frame(reader), 

606 timeout=timeout, 

607 ) 

608 if frame is None: 

609 break 

610 chunk = GrpcStreamChunk( 

611 stream_id=request.task_id, 

612 chunk=frame, 

613 ) 

614 yield chunk 

615 if chunk.is_last: 

616 break 

617 

618 async def broadcast( 

619 self, 

620 payload: str, 

621 capability_filter: Optional[str] = None, 

622 ) -> Dict[str, GrpcTaskResponse]: 

623 """Broadcast a task to all agents matching a capability.""" 

624 if not self._registry: 

625 raise ValueError("Registry required for broadcast") 

626 

627 agents = await self._registry.list_agents() 

628 if capability_filter: 

629 agents = [a for a in agents if capability_filter in a.capabilities] 

630 

631 results = {} 

632 tasks = [] 

633 for agent in agents: 

634 tasks.append(self._call_one(agent.agent_id, payload, results)) 

635 await asyncio.gather(*tasks, return_exceptions=True) 

636 return results 

637 

638 async def _call_one( 

639 self, agent_id: str, payload: str, results: Dict[str, GrpcTaskResponse] 

640 ) -> None: 

641 try: 

642 results[agent_id] = await self.submit_task(agent_id, payload) 

643 except Exception as e: 

644 results[agent_id] = GrpcTaskResponse( 

645 task_id="error", 

646 status=TaskStatus.FAILED, 

647 error=str(e), 

648 ) 

649 

650 async def close(self) -> None: 

651 """Close all connections.""" 

652 for _, writer in self._connections.values(): 

653 writer.close() 

654 try: 

655 await writer.wait_closed() 

656 except Exception: 

657 pass 

658 self._connections.clear() 

659 

660 

661# --------------------------------------------------------------------------- 

662# Default AgentService implementation 

663# --------------------------------------------------------------------------- 

664 

665class DefaultAgentService(GrpcAgentService): 

666 """Default AgentService implementation with task queues and streaming.""" 

667 

668 def __init__(self, agent_id: str = "", task_handler: Optional[Callable] = None): 

669 self.agent_id = agent_id or socket.gethostname() 

670 self._task_handler = task_handler or self._default_handler 

671 self._task_queue: asyncio.Queue = asyncio.Queue() 

672 self._active_streams: Dict[str, asyncio.Queue] = {} 

673 self._capabilities: List[str] = [ 

674 "text_generation", 

675 "code_analysis", 

676 "data_processing", 

677 "grpc_a2a", 

678 ] 

679 

680 async def submit_task(self, request: GrpcTaskRequest) -> GrpcTaskResponse: 

681 t0 = time.time() 

682 try: 

683 result = await self._task_handler(request) 

684 elapsed = time.time() - t0 

685 return GrpcTaskResponse( 

686 task_id=request.task_id, 

687 status=TaskStatus.SUCCESS, 

688 result=str(result), 

689 elapsed=elapsed, 

690 ) 

691 except Exception as e: 

692 elapsed = time.time() - t0 

693 return GrpcTaskResponse( 

694 task_id=request.task_id, 

695 status=TaskStatus.FAILED, 

696 error=str(e), 

697 elapsed=elapsed, 

698 ) 

699 

700 async def stream_execute(self, request: GrpcTaskRequest) -> AsyncIterator[GrpcStreamChunk]: 

701 queue: asyncio.Queue = asyncio.Queue() 

702 self._active_streams[request.task_id] = queue 

703 try: 

704 result = await self._task_handler(request) 

705 chunks = str(result).encode("utf-8") 

706 chunk_size = 4096 

707 for i in range(0, len(chunks), chunk_size): 

708 is_last = i + chunk_size >= len(chunks) 

709 yield GrpcStreamChunk( 

710 stream_id=request.task_id, 

711 chunk=chunks[i:i+chunk_size], 

712 seq=i // chunk_size, 

713 is_last=is_last, 

714 ) 

715 finally: 

716 self._active_streams.pop(request.task_id, None) 

717 

718 async def agent_chat( 

719 self, input_stream: AsyncIterator[GrpcStreamChunk] 

720 ) -> AsyncIterator[GrpcStreamChunk]: 

721 async for msg in input_stream: 

722 # Echo for now; in production this routes to agent logic 

723 yield GrpcStreamChunk( 

724 stream_id=msg.stream_id, 

725 chunk=b"ACK: " + msg.chunk, 

726 seq=msg.seq, 

727 is_last=msg.is_last, 

728 ) 

729 

730 async def health_check(self) -> Dict[str, Any]: 

731 return { 

732 "agent_id": self.agent_id, 

733 "status": "healthy", 

734 "active_streams": len(self._active_streams), 

735 "timestamp": time.time(), 

736 } 

737 

738 async def list_capabilities(self) -> List[str]: 

739 return self._capabilities 

740 

741 @staticmethod 

742 async def _default_handler(request: GrpcTaskRequest) -> str: 

743 return f"Task {request.task_id} acknowledged by agent {request.agent_id}" 

744 

745 

746# --------------------------------------------------------------------------- 

747# TLS/mTLS helpers 

748# --------------------------------------------------------------------------- 

749 

750def create_self_signed_cert( 

751 cert_file: str, key_file: str, common_name: str = "agentos.local" 

752) -> None: 

753 """Generate a self-signed certificate for testing gRPC TLS.""" 

754 from cryptography import x509 

755 from cryptography.x509.oid import NameOID 

756 from cryptography.hazmat.primitives import hashes, serialization 

757 from cryptography.hazmat.primitives.asymmetric import rsa 

758 import datetime 

759 

760 key = rsa.generate_private_key(public_exponent=65537, key_size=2048) 

761 

762 subject = issuer = x509.Name([ 

763 x509.NameAttribute(NameOID.COUNTRY_NAME, "US"), 

764 x509.NameAttribute(NameOID.STATE_OR_PROVINCE_NAME, "CA"), 

765 x509.NameAttribute(NameOID.LOCALITY_NAME, "San Francisco"), 

766 x509.NameAttribute(NameOID.ORGANIZATION_NAME, "AgentOS"), 

767 x509.NameAttribute(NameOID.COMMON_NAME, common_name), 

768 ]) 

769 

770 cert = ( 

771 x509.CertificateBuilder() 

772 .subject_name(subject) 

773 .issuer_name(issuer) 

774 .public_key(key.public_key()) 

775 .serial_number(x509.random_serial_number()) 

776 .not_valid_before(datetime.datetime.utcnow()) 

777 .not_valid_after(datetime.datetime.utcnow() + datetime.timedelta(days=365)) 

778 .add_extension( 

779 x509.SubjectAlternativeName([x509.DNSName(common_name)]), 

780 critical=False, 

781 ) 

782 .sign(key, hashes.SHA256()) 

783 ) 

784 

785 with open(key_file, "wb") as f: 

786 f.write(key.private_bytes( 

787 encoding=serialization.Encoding.PEM, 

788 format=serialization.PrivateFormat.PKCS8, 

789 encryption_algorithm=serialization.NoEncryption(), 

790 )) 

791 

792 with open(cert_file, "wb") as f: 

793 f.write(cert.public_bytes(serialization.Encoding.PEM)) 

794 

795 

796# --------------------------------------------------------------------------- 

797# Export 

798# --------------------------------------------------------------------------- 

799 

800__all__ = [ 

801 # Core types 

802 "GrpcTaskRequest", 

803 "GrpcTaskResponse", 

804 "GrpcHeartbeat", 

805 "GrpcStreamChunk", 

806 "TaskStatus", 

807 "GrpcStatusCode", 

808 # Service 

809 "GrpcAgentService", 

810 "DefaultAgentService", 

811 "SERVICE_NAME", 

812 # Server 

813 "GrpcServer", 

814 "GrpcServerConfig", 

815 # Client 

816 "GrpcClient", 

817 "GrpcClientConfig", 

818 # Codec 

819 "GrpcFrameCodec", 

820 # TLS 

821 "create_self_signed_cert", 

822]