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
« 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"""
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
16logger = logging.getLogger(__name__)
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"
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
39 def to_json(self) -> str:
40 """Convert to JSON string."""
41 return json.dumps(asdict(self))
43 @classmethod
44 def from_json(cls, data: str) -> "ServerInfo":
45 """Create from JSON string."""
46 return cls(**json.loads(data))
49class DiscoveryBroadcaster:
50 """Broadcasts MCP server information for autodiscovery."""
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.
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
79 self._running = False
80 self._thread: Optional[threading.Thread] = None
81 self._sockets: List[socket.socket] = []
83 logger.info(f"Discovery broadcaster initialized for {server_info.name}")
85 def start(self) -> None:
86 """Start broadcasting discovery information."""
87 if self._running:
88 logger.warning("Discovery broadcaster already running")
89 return
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}")
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}")
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}")
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}")
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()
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()
153 # Create discovery packet
154 packet = self._create_packet()
156 # Broadcast
157 if self.enable_broadcast:
158 self._send_broadcast(packet)
160 # Multicast
161 if self.enable_multicast:
162 self._send_multicast(packet)
164 except Exception as e:
165 logger.error(f"Error in broadcast loop: {e}")
167 # Wait for next broadcast
168 time.sleep(self.interval)
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
176 def _send_broadcast(self, packet: bytes) -> None:
177 """Send broadcast packet."""
178 if not self._sockets:
179 return
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}")
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
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}")
204class DiscoveryListener:
205 """Listens for MCP server discovery broadcasts."""
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.
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
228 self._running = False
229 self._servers: Dict[str, ServerInfo] = {}
230 self._thread: Optional[threading.Thread] = None
231 self._socket: Optional[socket.socket] = None
233 logger.info("Discovery listener initialized")
235 def start(self) -> None:
236 """Start listening for discovery broadcasts."""
237 if self._running:
238 logger.warning("Discovery listener already running")
239 return
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")
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")
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}")
267 return list(self._servers.values())
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)
278 # Bind to port
279 self._socket.bind(('', self.port))
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)
286 # Set timeout for socket operations
287 self._socket.settimeout(1.0)
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
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
305 try:
306 self._socket.close()
307 except:
308 pass
310 self._socket = None
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)
319 # Process packet
320 self._process_packet(data, addr)
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}")
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
336 # Extract JSON data
337 json_data = data[len(MCP_DISCOVERY_MAGIC) + 1:]
339 # Parse server info
340 server_info = ServerInfo.from_json(json_data.decode('utf-8'))
342 # Update or add server
343 is_new = server_info.id not in self._servers
344 self._servers[server_info.id] = server_info
346 # Call callback for new servers
347 if is_new and self.callback:
348 self.callback(server_info)
350 logger.debug(f"Discovered server: {server_info.name} from {addr[0]}")
352 except Exception as e:
353 logger.error(f"Failed to process discovery packet: {e}")
356class AutoDiscovery:
357 """High-level autodiscovery interface for QuickMCP servers."""
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.
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 )
395 self.broadcaster = DiscoveryBroadcaster(self.server_info)
396 logger.info(f"AutoDiscovery initialized for {server_name}")
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 }
409 def start(self) -> None:
410 """Start autodiscovery broadcasting."""
411 self.broadcaster.start()
413 def stop(self) -> None:
414 """Stop autodiscovery broadcasting."""
415 self.broadcaster.stop()
417 def __enter__(self):
418 """Context manager entry."""
419 self.start()
420 return self
422 def __exit__(self, exc_type, exc_val, exc_tb):
423 """Context manager exit."""
424 self.stop()
427# Convenience function for discovering servers
428async def discover_servers(timeout: float = 5.0) -> List[ServerInfo]:
429 """
430 Discover MCP servers on the network.
432 Args:
433 timeout: Discovery timeout in seconds
435 Returns:
436 List of discovered servers
437 """
438 discovered = []
440 def on_discovered(server_info: ServerInfo):
441 discovered.append(server_info)
443 listener = DiscoveryListener(callback=on_discovered)
444 listener.start()
446 # Wait for discovery
447 await asyncio.sleep(timeout)
449 listener.stop()
451 # Also get any servers that were already discovered
452 return listener.get_servers()
455# CLI tool for discovery
456def discovery_cli():
457 """Command-line tool for testing discovery."""
458 import argparse
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")
467 args = parser.parse_args()
469 logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
471 if args.command == "listen":
472 print("Listening for MCP servers...")
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)
484 listener = DiscoveryListener(callback=on_discovered, timeout=args.timeout)
485 listener.start()
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()
497 elif args.command == "broadcast":
498 print(f"Broadcasting server: {args.name}")
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 )
513 broadcaster = DiscoveryBroadcaster(server_info)
514 broadcaster.start()
516 try:
517 while True:
518 time.sleep(1)
519 print(".", end="", flush=True)
520 except KeyboardInterrupt:
521 print("\nStopping broadcaster...")
522 broadcaster.stop()
525if __name__ == "__main__":
526 discovery_cli()