1use std::collections::HashMap;
11use std::net::IpAddr;
12use std::sync::Arc;
13use std::time::{SystemTime, UNIX_EPOCH};
14
15use mdns_sd::{ServiceDaemon, ServiceEvent, ServiceInfo};
16use parking_lot::RwLock;
17use serde::{Deserialize, Serialize};
18use serde_json::json;
19
20use super::eventbus::{EventBus, WSEventType};
21
22const SERVICE_TYPE: &str = "_agentmux._tcp.local.";
23const LAN_AGENT_CACHE_TTL_SECS: u64 = 60;
24const LAN_PEER_QUERY_TIMEOUT_SECS: u64 = 2;
25
26#[derive(Debug, Clone, Serialize, Deserialize)]
27pub struct LanInstance {
28 pub instance_id: String,
29 pub hostname: String,
30 pub version: String,
31 pub address: String,
32 pub port: u16,
33 pub auth_key: String,
34 pub agents: Vec<String>,
35 pub first_seen: u64,
36 pub last_seen: u64,
37}
38
39struct LanCacheEntry {
40 peer_url: Option<String>,
42 auth_key: String,
43 expires: std::time::Instant,
44}
45
46pub struct LanDiscovery {
47 daemon: ServiceDaemon,
48 instances: Arc<RwLock<HashMap<String, LanInstance>>>,
49 instance_id: String,
50 event_bus: Arc<EventBus>,
51 service_fullname: String,
52 auth_key: String,
53}
54
55fn mdns_hostname(os_hostname: &str) -> String {
60 let trimmed = os_hostname.trim_end_matches('.');
61 if trimmed.ends_with(".local") {
62 format!("{trimmed}.")
63 } else {
64 format!("{trimmed}.local.")
65 }
66}
67
68impl LanDiscovery {
69 pub fn start(
71 instance_id: String,
72 hostname: String,
73 version: String,
74 port: u16,
75 auth_key: String,
76 event_bus: Arc<EventBus>,
77 ) -> Result<Arc<Self>, String> {
78 let daemon = ServiceDaemon::new().map_err(|e| format!("mDNS daemon failed: {e}"))?;
79
80 let service_name = format!("agentmux-{}", &instance_id);
84 let host_name_mdns = mdns_hostname(&hostname);
85 let properties = [
86 ("version", version.as_str()),
87 ("hostname", hostname.as_str()),
88 ("instance_id", instance_id.as_str()),
89 ("auth_key", auth_key.as_str()),
90 ];
91 let service_info = ServiceInfo::new(
92 SERVICE_TYPE,
93 &service_name,
94 &host_name_mdns,
95 "", port,
97 &properties[..],
98 )
99 .map_err(|e| format!("ServiceInfo creation failed: {e}"))?;
100
101 let service_fullname = service_info.get_fullname().to_string();
102
103 daemon
104 .register(service_info)
105 .map_err(|e| format!("mDNS register failed: {e}"))?;
106
107 let browse_receiver = daemon
109 .browse(SERVICE_TYPE)
110 .map_err(|e| format!("mDNS browse failed: {e}"))?;
111
112 let instances = Arc::new(RwLock::new(HashMap::new()));
113
114 let discovery = Arc::new(Self {
115 daemon,
116 instances: instances.clone(),
117 instance_id: instance_id.clone(),
118 event_bus: event_bus.clone(),
119 service_fullname,
120 auth_key,
121 });
122
123 let disc = discovery.clone();
125 tokio::task::spawn_blocking(move || {
126 disc.event_loop(browse_receiver);
127 });
128
129 tracing::info!(
130 instance_id = %instance_id,
131 port = port,
132 "LAN discovery started (mDNS)"
133 );
134
135 Ok(discovery)
136 }
137
138 fn event_loop(&self, receiver: mdns_sd::Receiver<ServiceEvent>) {
139 loop {
140 match receiver.recv() {
141 Ok(event) => self.handle_event(event),
142 Err(_) => {
143 tracing::warn!("mDNS event receiver closed");
144 break;
145 }
146 }
147 }
148 }
149
150 fn handle_event(&self, event: ServiceEvent) {
151 match event {
152 ServiceEvent::ServiceResolved(info) => {
153 let peer_id = info
154 .get_property_val_str("instance_id")
155 .unwrap_or_default()
156 .to_string();
157
158 if peer_id == self.instance_id {
160 return;
161 }
162
163 let now = SystemTime::now()
164 .duration_since(UNIX_EPOCH)
165 .unwrap_or_default()
166 .as_secs();
167
168 let address = info
169 .get_addresses()
170 .iter()
171 .find(|a| matches!(a, IpAddr::V4(_)))
172 .or_else(|| info.get_addresses().iter().next())
173 .map(|a| a.to_string())
174 .unwrap_or_default();
175
176 let hostname = info
177 .get_property_val_str("hostname")
178 .unwrap_or_default()
179 .to_string();
180 let version = info
181 .get_property_val_str("version")
182 .unwrap_or_default()
183 .to_string();
184 let auth_key = info
185 .get_property_val_str("auth_key")
186 .unwrap_or_default()
187 .to_string();
188
189 let fullname = info.get_fullname().to_string();
190 let mut instances = self.instances.write();
191 let entry = instances.entry(fullname).or_insert_with(|| LanInstance {
192 instance_id: peer_id.clone(),
193 hostname: hostname.clone(),
194 version: version.clone(),
195 address: address.clone(),
196 port: info.get_port(),
197 auth_key: auth_key.clone(),
198 agents: Vec::new(),
199 first_seen: now,
200 last_seen: now,
201 });
202 entry.last_seen = now;
203 entry.hostname = hostname;
204 entry.version = version;
205 entry.address = address;
206 entry.port = info.get_port();
207 entry.auth_key = auth_key;
208 drop(instances);
209
210 tracing::info!(
211 peer_id = %peer_id,
212 address = %info.get_addresses().iter().next().map(|a| a.to_string()).unwrap_or_default(),
213 port = info.get_port(),
214 "LAN peer discovered"
215 );
216
217 self.broadcast_instances();
218 }
219 ServiceEvent::ServiceRemoved(_, fullname) => {
220 let removed = {
221 let mut instances = self.instances.write();
222 instances.remove(&fullname).is_some()
223 };
224 if removed {
225 tracing::info!(fullname = %fullname, "LAN peer removed");
226 self.broadcast_instances();
227 }
228 }
229 _ => {}
230 }
231 }
232
233 fn broadcast_instances(&self) {
234 let instances: Vec<LanInstance> = self.instances.read().values().cloned().collect();
235 self.event_bus.broadcast_event(&WSEventType {
236 eventtype: "laninstances".to_string(),
237 oref: String::new(),
238 data: Some(json!(instances)),
239 });
240 }
241
242 pub fn get_instances(&self) -> Vec<LanInstance> {
244 self.instances.read().values().cloned().collect()
245 }
246
247 #[allow(dead_code)]
249 pub fn peer_count(&self) -> usize {
250 self.instances.read().len()
251 }
252
253 pub fn shutdown(&self) {
263 if let Err(e) = self.daemon.unregister(&self.service_fullname) {
264 tracing::debug!("mDNS unregister returned: {e}");
266 }
267 if let Err(e) = self.daemon.shutdown() {
268 tracing::debug!("mDNS daemon shutdown returned: {e}");
269 }
270 }
271}
272
273impl Drop for LanDiscovery {
274 fn drop(&mut self) {
275 self.shutdown();
280 }
281}
282
283pub struct LanDiscoveryController {
292 slot: Arc<RwLock<Option<Arc<LanDiscovery>>>>,
293 instance_id: String,
294 hostname: String,
295 version: String,
296 port: u16,
297 auth_key: String,
298 event_bus: Arc<EventBus>,
299 agent_cache: std::sync::RwLock<HashMap<String, LanCacheEntry>>,
302}
303
304impl LanDiscoveryController {
305 pub fn new(
306 instance_id: String,
307 hostname: String,
308 version: String,
309 port: u16,
310 event_bus: Arc<EventBus>,
311 auth_key: String,
312 ) -> Self {
313 Self {
314 slot: Arc::new(RwLock::new(None)),
315 instance_id,
316 hostname,
317 version,
318 port,
319 auth_key,
320 event_bus,
321 agent_cache: std::sync::RwLock::new(HashMap::new()),
322 }
323 }
324
325 pub async fn find_agent(
337 &self,
338 agent_id: &str,
339 http: &reqwest::Client,
340 ) -> Option<(String, String)> {
341 if let Ok(cache) = self.agent_cache.read() {
343 if let Some(e) = cache.get(agent_id) {
344 if e.expires > std::time::Instant::now() {
345 return e.peer_url.as_ref().map(|url| (url.clone(), e.auth_key.clone()));
346 }
347 }
348 }
349
350 let peers = self.get_instances();
353 for peer in &peers {
354 if peer.address.is_empty() || peer.auth_key.is_empty() {
355 continue;
356 }
357 let peer_url = format!("http://{}:{}", peer.address, peer.port);
358 let result = http
359 .get(format!("{}/agentmux/reactive/agent", peer_url))
360 .query(&[("id", agent_id)])
361 .header("X-AuthKey", &peer.auth_key)
362 .timeout(std::time::Duration::from_secs(LAN_PEER_QUERY_TIMEOUT_SECS))
363 .send()
364 .await;
365 if matches!(result, Ok(ref r) if r.status().is_success()) {
366 tracing::debug!(agent_id, peer_url = %peer_url, "LAN agent found on peer");
367 if let Ok(mut cache) = self.agent_cache.write() {
368 cache.insert(
369 agent_id.to_string(),
370 LanCacheEntry {
371 peer_url: Some(peer_url.clone()),
372 auth_key: peer.auth_key.clone(),
373 expires: std::time::Instant::now()
374 + std::time::Duration::from_secs(LAN_AGENT_CACHE_TTL_SECS),
375 },
376 );
377 }
378 return Some((peer_url, peer.auth_key.clone()));
379 }
380 }
381
382 if let Ok(mut cache) = self.agent_cache.write() {
385 cache.insert(
386 agent_id.to_string(),
387 LanCacheEntry {
388 peer_url: None,
389 auth_key: String::new(),
390 expires: std::time::Instant::now()
391 + std::time::Duration::from_secs(LAN_AGENT_CACHE_TTL_SECS),
392 },
393 );
394 }
395 None
396 }
397
398 pub fn evict_agent(&self, agent_id: &str) {
400 if let Ok(mut cache) = self.agent_cache.write() {
401 cache.remove(agent_id);
402 }
403 }
404
405 pub fn apply(&self, enabled: bool) {
415 let mut slot = self.slot.write();
416 let is_running = slot.is_some();
417 match (enabled, is_running) {
418 (true, false) => {
419 match LanDiscovery::start(
420 self.instance_id.clone(),
421 self.hostname.clone(),
422 self.version.clone(),
423 self.port,
424 self.auth_key.clone(),
425 self.event_bus.clone(),
426 ) {
427 Ok(d) => {
428 *slot = Some(d);
429 tracing::info!("LAN discovery enabled via setting");
430 }
431 Err(e) => {
432 tracing::warn!("LAN discovery start failed: {e}");
433 self.event_bus.broadcast_event(&WSEventType {
439 eventtype: "laninstances:error".to_string(),
440 oref: String::new(),
441 data: Some(json!({ "error": e.to_string() })),
442 });
443 }
444 }
445 }
446 (false, true) => {
447 if let Some(d) = slot.as_ref() {
457 d.shutdown();
458 }
459 *slot = None;
460 tracing::info!("LAN discovery disabled via setting");
461 self.event_bus.broadcast_event(&WSEventType {
462 eventtype: "laninstances".to_string(),
463 oref: String::new(),
464 data: Some(json!([])),
465 });
466 }
467 _ => {}
468 }
469 }
470
471 pub fn get_instances(&self) -> Vec<LanInstance> {
474 self.slot
475 .read()
476 .as_ref()
477 .map(|d| d.get_instances())
478 .unwrap_or_default()
479 }
480}
481
482#[cfg(test)]
483mod tests {
484 use super::mdns_hostname;
485
486 #[test]
487 fn appends_local_dot_to_bare_hostname() {
488 assert_eq!(mdns_hostname("claudius"), "claudius.local.");
489 }
490
491 #[test]
492 fn preserves_already_fully_qualified_name() {
493 assert_eq!(mdns_hostname("claudius.local."), "claudius.local.");
494 }
495
496 #[test]
497 fn appends_trailing_dot_to_local_suffix() {
498 assert_eq!(mdns_hostname("claudius.local"), "claudius.local.");
500 }
501
502 #[test]
503 fn handles_trailing_dot_on_bare_hostname() {
504 assert_eq!(mdns_hostname("claudius."), "claudius.local.");
505 }
506
507 #[test]
508 fn does_not_double_suffix() {
509 let once = mdns_hostname("claudius");
511 let twice = mdns_hostname(&once);
512 assert_eq!(twice, once);
513 }
514}