Coverage for src/quickmcp/autodiscovery.py: 60%

281 statements  

« prev     ^ index     » next       coverage.py v7.10.4, created at 2025-08-20 20:56 +0200

1""" 

2QuickMCP Autodiscovery - Allow MCP servers to be discovered on the network 

3""" 

4 

5import asyncio 

6import json 

7import socket 

8import time 

9import threading 

10import logging 

11from typing import Optional, Dict, Any, Callable, List 

12from dataclasses import dataclass, asdict 

13import platform 

14import uuid 

15 

16logger = logging.getLogger(__name__) 

17 

18# Default multicast group and port for MCP discovery 

19MCP_MULTICAST_GROUP = "239.255.41.42" # Private multicast address 

20MCP_DISCOVERY_PORT = 42424 

21MCP_BROADCAST_INTERVAL = 5.0 # seconds 

22MCP_DISCOVERY_MAGIC = b"MCP_DISCOVER_V1" 

23 

24 

25@dataclass 

26class ServerInfo: 

27 """Information about an MCP server for discovery.""" 

28 id: str 

29 name: str 

30 version: str 

31 description: str 

32 host: str 

33 port: int 

34 transport: str 

35 capabilities: Dict[str, Any] 

36 metadata: Dict[str, Any] 

37 timestamp: float 

38 

39 def to_json(self) -> str: 

40 """Convert to JSON string.""" 

41 return json.dumps(asdict(self)) 

42 

43 @classmethod 

44 def from_json(cls, data: str) -> "ServerInfo": 

45 """Create from JSON string.""" 

46 return cls(**json.loads(data)) 

47 

48 

49class DiscoveryBroadcaster: 

50 """Broadcasts MCP server information for autodiscovery.""" 

51 

52 def __init__( 

53 self, 

54 server_info: ServerInfo, 

55 multicast_group: str = MCP_MULTICAST_GROUP, 

56 port: int = MCP_DISCOVERY_PORT, 

57 interval: float = MCP_BROADCAST_INTERVAL, 

58 enable_broadcast: bool = True, 

59 enable_multicast: bool = True, 

60 ): 

61 """ 

62 Initialize the discovery broadcaster. 

63  

64 Args: 

65 server_info: Information about the server to broadcast 

66 multicast_group: Multicast group address 

67 port: Port for discovery 

68 interval: Broadcast interval in seconds 

69 enable_broadcast: Enable UDP broadcast 

70 enable_multicast: Enable multicast 

71 """ 

72 self.server_info = server_info 

73 self.multicast_group = multicast_group 

74 self.port = port 

75 self.interval = interval 

76 self.enable_broadcast = enable_broadcast 

77 self.enable_multicast = enable_multicast 

78 

79 self._running = False 

80 self._thread: Optional[threading.Thread] = None 

81 self._sockets: List[socket.socket] = [] 

82 

83 logger.info(f"Discovery broadcaster initialized for {server_info.name}") 

84 

85 def start(self) -> None: 

86 """Start broadcasting discovery information.""" 

87 if self._running: 

88 logger.warning("Discovery broadcaster already running") 

89 return 

90 

91 self._running = True 

92 self._setup_sockets() 

93 self._thread = threading.Thread(target=self._broadcast_loop, daemon=True) 

94 self._thread.start() 

95 logger.info(f"Discovery broadcasting started for {self.server_info.name}") 

96 

97 def stop(self) -> None: 

98 """Stop broadcasting discovery information.""" 

99 self._running = False 

100 if self._thread: 

101 self._thread.join(timeout=1.0) 

102 self._cleanup_sockets() 

103 logger.info(f"Discovery broadcasting stopped for {self.server_info.name}") 

104 

105 def _setup_sockets(self) -> None: 

106 """Set up broadcast and multicast sockets.""" 

107 # Broadcast socket 

108 if self.enable_broadcast: 

109 try: 

110 sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) 

111 sock.setsockopt(socket.SOL_SOCKET, socket.SO_BROADCAST, 1) 

112 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) 

113 # Platform-specific reuse port 

114 if hasattr(socket, 'SO_REUSEPORT'): 

115 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1) 

116 self._sockets.append(sock) 

117 logger.debug("Broadcast socket created") 

118 except Exception as e: 

119 logger.error(f"Failed to create broadcast socket: {e}") 

120 

121 # Multicast socket 

122 if self.enable_multicast: 

123 try: 

124 sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP) 

125 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) 

126 if hasattr(socket, 'SO_REUSEPORT'): 

127 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1) 

128 # Set multicast TTL 

129 sock.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_TTL, 2) 

130 # Enable multicast loop (receive own messages) 

131 sock.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_LOOP, 1) 

132 self._sockets.append(sock) 

133 logger.debug("Multicast socket created") 

134 except Exception as e: 

135 logger.error(f"Failed to create multicast socket: {e}") 

136 

137 def _cleanup_sockets(self) -> None: 

138 """Clean up sockets.""" 

139 for sock in self._sockets: 

140 try: 

141 sock.close() 

142 except: 

143 pass 

144 self._sockets.clear() 

145 

146 def _broadcast_loop(self) -> None: 

147 """Main broadcast loop.""" 

148 while self._running: 

149 try: 

150 # Update timestamp 

151 self.server_info.timestamp = time.time() 

152 

153 # Create discovery packet 

154 packet = self._create_packet() 

155 

156 # Broadcast 

157 if self.enable_broadcast: 

158 self._send_broadcast(packet) 

159 

160 # Multicast 

161 if self.enable_multicast: 

162 self._send_multicast(packet) 

163 

164 except Exception as e: 

165 logger.error(f"Error in broadcast loop: {e}") 

166 

167 # Wait for next broadcast 

168 time.sleep(self.interval) 

169 

170 def _create_packet(self) -> bytes: 

171 """Create a discovery packet.""" 

172 # Packet format: MAGIC + JSON data 

173 data = self.server_info.to_json().encode('utf-8') 

174 return MCP_DISCOVERY_MAGIC + b'\n' + data 

175 

176 def _send_broadcast(self, packet: bytes) -> None: 

177 """Send broadcast packet.""" 

178 if not self._sockets: 

179 return 

180 

181 try: 

182 sock = self._sockets[0] # Use first socket for broadcast 

183 sock.sendto(packet, ('<broadcast>', self.port)) 

184 logger.debug(f"Broadcast sent for {self.server_info.name}") 

185 except Exception as e: 

186 logger.error(f"Failed to send broadcast: {e}") 

187 

188 def _send_multicast(self, packet: bytes) -> None: 

189 """Send multicast packet.""" 

190 if len(self._sockets) < 2 and not self.enable_broadcast: 

191 sock = self._sockets[0] 

192 elif len(self._sockets) >= 2: 

193 sock = self._sockets[1] 

194 else: 

195 return 

196 

197 try: 

198 sock.sendto(packet, (self.multicast_group, self.port)) 

199 logger.debug(f"Multicast sent for {self.server_info.name}") 

200 except Exception as e: 

201 logger.error(f"Failed to send multicast: {e}") 

202 

203 

204class DiscoveryListener: 

205 """Listens for MCP server discovery broadcasts.""" 

206 

207 def __init__( 

208 self, 

209 callback: Optional[Callable[[ServerInfo], None]] = None, 

210 multicast_group: str = MCP_MULTICAST_GROUP, 

211 port: int = MCP_DISCOVERY_PORT, 

212 timeout: float = 30.0, 

213 ): 

214 """ 

215 Initialize the discovery listener. 

216  

217 Args: 

218 callback: Callback function when server is discovered 

219 multicast_group: Multicast group address 

220 port: Port for discovery 

221 timeout: Timeout for removing stale servers 

222 """ 

223 self.callback = callback 

224 self.multicast_group = multicast_group 

225 self.port = port 

226 self.timeout = timeout 

227 

228 self._running = False 

229 self._servers: Dict[str, ServerInfo] = {} 

230 self._thread: Optional[threading.Thread] = None 

231 self._socket: Optional[socket.socket] = None 

232 

233 logger.info("Discovery listener initialized") 

234 

235 def start(self) -> None: 

236 """Start listening for discovery broadcasts.""" 

237 if self._running: 

238 logger.warning("Discovery listener already running") 

239 return 

240 

241 self._running = True 

242 self._setup_socket() 

243 self._thread = threading.Thread(target=self._listen_loop, daemon=True) 

244 self._thread.start() 

245 logger.info("Discovery listener started") 

246 

247 def stop(self) -> None: 

248 """Stop listening for discovery broadcasts.""" 

249 self._running = False 

250 if self._thread: 

251 self._thread.join(timeout=1.0) 

252 self._cleanup_socket() 

253 logger.info("Discovery listener stopped") 

254 

255 def get_servers(self) -> List[ServerInfo]: 

256 """Get list of discovered servers.""" 

257 current_time = time.time() 

258 # Remove stale servers 

259 stale_ids = [ 

260 sid for sid, info in self._servers.items() 

261 if current_time - info.timestamp > self.timeout 

262 ] 

263 for sid in stale_ids: 

264 del self._servers[sid] 

265 logger.debug(f"Removed stale server: {sid}") 

266 

267 return list(self._servers.values()) 

268 

269 def _setup_socket(self) -> None: 

270 """Set up listening socket.""" 

271 try: 

272 # Create socket 

273 self._socket = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) 

274 self._socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) 

275 if hasattr(socket, 'SO_REUSEPORT'): 

276 self._socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1) 

277 

278 # Bind to port 

279 self._socket.bind(('', self.port)) 

280 

281 # Join multicast group 

282 group = socket.inet_aton(self.multicast_group) 

283 mreq = group + socket.inet_aton('0.0.0.0') 

284 self._socket.setsockopt(socket.IPPROTO_IP, socket.IP_ADD_MEMBERSHIP, mreq) 

285 

286 # Set timeout for socket operations 

287 self._socket.settimeout(1.0) 

288 

289 logger.debug(f"Listening socket created on port {self.port}") 

290 except Exception as e: 

291 logger.error(f"Failed to create listening socket: {e}") 

292 raise 

293 

294 def _cleanup_socket(self) -> None: 

295 """Clean up listening socket.""" 

296 if self._socket: 

297 try: 

298 # Leave multicast group 

299 group = socket.inet_aton(self.multicast_group) 

300 mreq = group + socket.inet_aton('0.0.0.0') 

301 self._socket.setsockopt(socket.IPPROTO_IP, socket.IP_DROP_MEMBERSHIP, mreq) 

302 except: 

303 pass 

304 

305 try: 

306 self._socket.close() 

307 except: 

308 pass 

309 

310 self._socket = None 

311 

312 def _listen_loop(self) -> None: 

313 """Main listening loop.""" 

314 while self._running: 

315 try: 

316 # Receive packet 

317 data, addr = self._socket.recvfrom(65535) 

318 

319 # Process packet 

320 self._process_packet(data, addr) 

321 

322 except socket.timeout: 

323 # Timeout is expected, continue 

324 continue 

325 except Exception as e: 

326 if self._running: 

327 logger.error(f"Error in listen loop: {e}") 

328 

329 def _process_packet(self, data: bytes, addr: tuple) -> None: 

330 """Process a received discovery packet.""" 

331 try: 

332 # Check magic bytes 

333 if not data.startswith(MCP_DISCOVERY_MAGIC): 

334 return 

335 

336 # Extract JSON data 

337 json_data = data[len(MCP_DISCOVERY_MAGIC) + 1:] 

338 

339 # Parse server info 

340 server_info = ServerInfo.from_json(json_data.decode('utf-8')) 

341 

342 # Update or add server 

343 is_new = server_info.id not in self._servers 

344 self._servers[server_info.id] = server_info 

345 

346 # Call callback for new servers 

347 if is_new and self.callback: 

348 self.callback(server_info) 

349 

350 logger.debug(f"Discovered server: {server_info.name} from {addr[0]}") 

351 

352 except Exception as e: 

353 logger.error(f"Failed to process discovery packet: {e}") 

354 

355 

356class AutoDiscovery: 

357 """High-level autodiscovery interface for QuickMCP servers.""" 

358 

359 def __init__( 

360 self, 

361 server_name: str, 

362 server_version: str = "1.0.0", 

363 server_description: str = "", 

364 transport: str = "stdio", 

365 host: str = "localhost", 

366 port: int = 8000, 

367 metadata: Optional[Dict[str, Any]] = None, 

368 ): 

369 """ 

370 Initialize autodiscovery for a server. 

371  

372 Args: 

373 server_name: Name of the server 

374 server_version: Version of the server 

375 server_description: Description of the server 

376 transport: Transport type (stdio, sse, http) 

377 host: Host for network transports 

378 port: Port for network transports 

379 metadata: Additional metadata 

380 """ 

381 self.server_id = str(uuid.uuid4()) 

382 self.server_info = ServerInfo( 

383 id=self.server_id, 

384 name=server_name, 

385 version=server_version, 

386 description=server_description, 

387 host=host if transport != "stdio" else platform.node(), 

388 port=port if transport != "stdio" else 0, 

389 transport=transport, 

390 capabilities={}, 

391 metadata=metadata or {}, 

392 timestamp=time.time() 

393 ) 

394 

395 self.broadcaster = DiscoveryBroadcaster(self.server_info) 

396 logger.info(f"AutoDiscovery initialized for {server_name}") 

397 

398 def update_capabilities(self, tools: List[str], resources: List[str], prompts: List[str]) -> None: 

399 """Update server capabilities.""" 

400 self.server_info.capabilities = { 

401 "tools": tools, 

402 "resources": resources, 

403 "prompts": prompts, 

404 "tool_count": len(tools), 

405 "resource_count": len(resources), 

406 "prompt_count": len(prompts), 

407 } 

408 

409 def start(self) -> None: 

410 """Start autodiscovery broadcasting.""" 

411 self.broadcaster.start() 

412 

413 def stop(self) -> None: 

414 """Stop autodiscovery broadcasting.""" 

415 self.broadcaster.stop() 

416 

417 def __enter__(self): 

418 """Context manager entry.""" 

419 self.start() 

420 return self 

421 

422 def __exit__(self, exc_type, exc_val, exc_tb): 

423 """Context manager exit.""" 

424 self.stop() 

425 

426 

427# Convenience function for discovering servers 

428async def discover_servers(timeout: float = 5.0) -> List[ServerInfo]: 

429 """ 

430 Discover MCP servers on the network. 

431  

432 Args: 

433 timeout: Discovery timeout in seconds 

434  

435 Returns: 

436 List of discovered servers 

437 """ 

438 discovered = [] 

439 

440 def on_discovered(server_info: ServerInfo): 

441 discovered.append(server_info) 

442 

443 listener = DiscoveryListener(callback=on_discovered) 

444 listener.start() 

445 

446 # Wait for discovery 

447 await asyncio.sleep(timeout) 

448 

449 listener.stop() 

450 

451 # Also get any servers that were already discovered 

452 return listener.get_servers() 

453 

454 

455# CLI tool for discovery 

456def discovery_cli(): 

457 """Command-line tool for testing discovery.""" 

458 import argparse 

459 

460 parser = argparse.ArgumentParser(description="MCP Server Discovery Tool") 

461 parser.add_argument("command", choices=["listen", "broadcast"], help="Command to run") 

462 parser.add_argument("--name", default="test-server", help="Server name (for broadcast)") 

463 parser.add_argument("--port", type=int, default=8000, help="Server port (for broadcast)") 

464 parser.add_argument("--transport", default="stdio", help="Transport type (for broadcast)") 

465 parser.add_argument("--timeout", type=float, default=30.0, help="Discovery timeout") 

466 

467 args = parser.parse_args() 

468 

469 logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s") 

470 

471 if args.command == "listen": 

472 print("Listening for MCP servers...") 

473 

474 def on_discovered(server_info: ServerInfo): 

475 print(f"\n=== Discovered Server ===") 

476 print(f"Name: {server_info.name}") 

477 print(f"Version: {server_info.version}") 

478 print(f"Transport: {server_info.transport}") 

479 print(f"Host: {server_info.host}") 

480 print(f"Port: {server_info.port}") 

481 print(f"Capabilities: {server_info.capabilities}") 

482 print("=" * 25) 

483 

484 listener = DiscoveryListener(callback=on_discovered, timeout=args.timeout) 

485 listener.start() 

486 

487 try: 

488 while True: 

489 time.sleep(1) 

490 servers = listener.get_servers() 

491 if servers: 

492 print(f"\rActive servers: {len(servers)}", end="", flush=True) 

493 except KeyboardInterrupt: 

494 print("\nStopping listener...") 

495 listener.stop() 

496 

497 elif args.command == "broadcast": 

498 print(f"Broadcasting server: {args.name}") 

499 

500 server_info = ServerInfo( 

501 id=str(uuid.uuid4()), 

502 name=args.name, 

503 version="1.0.0", 

504 description="Test MCP server", 

505 host=platform.node(), 

506 port=args.port, 

507 transport=args.transport, 

508 capabilities={"tools": ["test_tool"], "resources": [], "prompts": []}, 

509 metadata={}, 

510 timestamp=time.time() 

511 ) 

512 

513 broadcaster = DiscoveryBroadcaster(server_info) 

514 broadcaster.start() 

515 

516 try: 

517 while True: 

518 time.sleep(1) 

519 print(".", end="", flush=True) 

520 except KeyboardInterrupt: 

521 print("\nStopping broadcaster...") 

522 broadcaster.stop() 

523 

524 

525if __name__ == "__main__": 

526 discovery_cli()