Skip to main content

spvirit_client/
search.rs

1use std::collections::HashSet;
2use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
3use std::sync::Arc;
4use std::time::Duration;
5
6use dns_lookup::lookup_host;
7use get_if_addrs::{IfAddr, get_if_addrs};
8use socket2::{Domain, Protocol, Socket, Type};
9use tokio::io::AsyncWriteExt;
10use tokio::net::UdpSocket;
11use tracing::debug;
12
13use crate::auth::{default_authnz_host, default_authnz_user};
14use crate::transport::read_packet;
15use crate::types::{PvGetError, PvGetOptions};
16use spvirit_codec::SegmentReassembler;
17use spvirit_codec::epics_decode::{PvaPacket, PvaPacketCommand};
18use spvirit_codec::spvirit_encode::{
19    encode_client_connection_validation, encode_search_request, ip_to_bytes,
20    socket_addr_from_pva_bytes,
21};
22
23#[derive(Clone, Copy, Debug)]
24pub struct SearchTarget {
25    pub target: IpAddr,
26    pub bind: IpAddr,
27}
28
29#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
30pub struct DiscoveredServer {
31    pub guid: [u8; 12],
32    pub tcp_addr: SocketAddr,
33}
34
35pub fn parse_addr_list(env: &str) -> Vec<IpAddr> {
36    env.split(|c| c == ',' || c == ' ' || c == '\t')
37        .filter(|s| !s.trim().is_empty())
38        .filter_map(|s| parse_search_target_ip(s.trim()))
39        .collect()
40}
41
42fn parse_search_target_ip(token: &str) -> Option<IpAddr> {
43    if token.is_empty() {
44        return None;
45    }
46
47    if let Ok(ip) = token.parse::<IpAddr>() {
48        return Some(ip);
49    }
50    if let Ok(sock) = token.parse::<SocketAddr>() {
51        return Some(sock.ip());
52    }
53
54    // Accept host:port where host may be a name or an IP literal.
55    // For IPv6 bracket notation [::1]:port, SocketAddr::parse above already handles it.
56    if let Some((host, port_str)) = token.rsplit_once(':') {
57        if !host.is_empty()
58            && !port_str.is_empty()
59            && port_str.chars().all(|c| c.is_ascii_digit())
60            && !host.contains(']')
61        {
62            if let Ok(ip) = host.parse::<IpAddr>() {
63                return Some(ip);
64            }
65            if let Ok(addrs) = lookup_host(host) {
66                // Prefer IPv4 for backward compat, fall back to first IPv6
67                let addrs: Vec<IpAddr> = addrs.collect();
68                if let Some(ip) = addrs
69                    .iter()
70                    .find(|ip| ip.is_ipv4())
71                    .copied()
72                    .or_else(|| addrs.into_iter().next())
73                {
74                    return Some(ip);
75                }
76            }
77        }
78    }
79
80    if let Ok(addrs) = lookup_host(token) {
81        // Prefer IPv4, fall back to first IPv6
82        let addrs: Vec<IpAddr> = addrs.collect();
83        if let Some(ip) = addrs
84            .iter()
85            .find(|ip| ip.is_ipv4())
86            .copied()
87            .or_else(|| addrs.into_iter().next())
88        {
89            return Some(ip);
90        }
91    }
92
93    None
94}
95
96/// Return a default unspecified bind address matching the target's address family.
97fn unspecified_for(ip: IpAddr) -> IpAddr {
98    match ip {
99        IpAddr::V4(_) => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
100        IpAddr::V6(_) => IpAddr::V6(Ipv6Addr::UNSPECIFIED),
101    }
102}
103
104pub fn build_search_targets(
105    search_addr: Option<IpAddr>,
106    bind_addr: Option<IpAddr>,
107) -> Vec<SearchTarget> {
108    // Explicit --search-addr overrides everything (single target).
109    if let Some(ip) = search_addr {
110        return vec![SearchTarget {
111            target: ip,
112            bind: bind_addr.unwrap_or_else(|| unspecified_for(ip)),
113        }];
114    }
115
116    let mut targets = Vec::new();
117    let mut seen = HashSet::new();
118
119    // Addresses from EPICS_PVA_ADDR_LIST.
120    if let Ok(env) = std::env::var("EPICS_PVA_ADDR_LIST") {
121        for ip in parse_addr_list(&env) {
122            if seen.insert(ip) {
123                targets.push(SearchTarget {
124                    target: ip,
125                    bind: bind_addr.unwrap_or_else(|| unspecified_for(ip)),
126                });
127            }
128        }
129    }
130
131    // Merge auto-discovered broadcast addresses unless explicitly disabled.
132    // This matches EPICS Base behaviour: ADDR_LIST + auto-broadcast combined.
133    if is_auto_addr_list_enabled() {
134        for t in build_auto_broadcast_targets() {
135            if seen.insert(t.target) {
136                targets.push(SearchTarget {
137                    target: t.target,
138                    bind: bind_addr.unwrap_or(t.bind),
139                });
140            }
141        }
142    }
143
144    targets
145}
146
147pub fn is_auto_addr_list_enabled() -> bool {
148    match std::env::var("EPICS_PVA_AUTO_ADDR_LIST") {
149        Ok(v) => {
150            let v = v.trim().to_ascii_uppercase();
151            v == "YES" || v == "Y" || v == "1" || v == "TRUE"
152        }
153        Err(_) => true,
154    }
155}
156
157fn ipv4_is_link_local(ip: Ipv4Addr) -> bool {
158    let octets = ip.octets();
159    octets[0] == 169 && octets[1] == 254
160}
161
162fn choose_default_bind_v4() -> Option<Ipv4Addr> {
163    let ifaces = get_if_addrs().ok()?;
164    for iface in ifaces {
165        if let IfAddr::V4(v4) = iface.addr {
166            let ip = v4.ip;
167            if ip.is_loopback() || ipv4_is_link_local(ip) {
168                continue;
169            }
170            return Some(ip);
171        }
172    }
173    None
174}
175
176fn choose_default_bind_v6() -> Option<Ipv6Addr> {
177    let ifaces = get_if_addrs().ok()?;
178    for iface in ifaces {
179        if let IfAddr::V6(v6) = iface.addr {
180            let ip = v6.ip;
181            if ip.is_loopback() {
182                continue;
183            }
184            // Skip link-local (fe80::/10) — not routable without scope id
185            let segs = ip.segments();
186            if segs[0] & 0xffc0 == 0xfe80 {
187                continue;
188            }
189            return Some(ip);
190        }
191    }
192    None
193}
194
195fn broadcast_for(ip: Ipv4Addr, netmask: Ipv4Addr) -> Ipv4Addr {
196    let ip_u = u32::from(ip);
197    let mask_u = u32::from(netmask);
198    Ipv4Addr::from(ip_u | !mask_u)
199}
200
201fn discovery_target_for(ip: Ipv4Addr, netmask: Ipv4Addr) -> Ipv4Addr {
202    let limited_broadcast = Ipv4Addr::new(255, 255, 255, 255);
203    if netmask == Ipv4Addr::new(255, 255, 255, 255) || netmask.is_unspecified() {
204        return limited_broadcast;
205    }
206    let directed = broadcast_for(ip, netmask);
207    if directed == ip {
208        limited_broadcast
209    } else {
210        directed
211    }
212}
213
214pub fn build_auto_broadcast_targets() -> Vec<SearchTarget> {
215    let mut targets = Vec::new();
216    let mut fallback_targets = Vec::new();
217    let mut fallback_seen = HashSet::new();
218    let mut added_v4_multicast = false;
219    let mut added_v6_multicast = false;
220    let ifaces = match get_if_addrs() {
221        Ok(v) => v,
222        Err(_) => return targets,
223    };
224    for iface in &ifaces {
225        if let IfAddr::V4(v4) = &iface.addr {
226            let ip = v4.ip;
227            if ip.is_loopback() || ipv4_is_link_local(ip) {
228                continue;
229            }
230            let bcast = discovery_target_for(ip, v4.netmask);
231            targets.push(SearchTarget {
232                target: IpAddr::V4(bcast),
233                bind: IpAddr::V4(ip),
234            });
235            // Also send to IPv4 multicast group (matching PVXS behaviour).
236            // Docker overlay networks may block broadcast but allow multicast.
237            targets.push(SearchTarget {
238                target: IpAddr::V4(PVA_MULTICAST_V4),
239                bind: IpAddr::V4(ip),
240            });
241            if fallback_seen.insert(IpAddr::V4(bcast)) {
242                fallback_targets.push(SearchTarget {
243                    target: IpAddr::V4(bcast),
244                    bind: IpAddr::V4(Ipv4Addr::UNSPECIFIED),
245                });
246            }
247            if !added_v4_multicast {
248                added_v4_multicast = true;
249                fallback_targets.push(SearchTarget {
250                    target: IpAddr::V4(PVA_MULTICAST_V4),
251                    bind: IpAddr::V4(Ipv4Addr::UNSPECIFIED),
252                });
253            }
254        }
255    }
256    // Add IPv6 multicast targets for each non-loopback, non-link-local v6 iface.
257    for iface in &ifaces {
258        if let IfAddr::V6(v6) = &iface.addr {
259            let ip = v6.ip;
260            if ip.is_loopback() {
261                continue;
262            }
263            let segs = ip.segments();
264            if segs[0] & 0xffc0 == 0xfe80 {
265                continue; // skip link-local
266            }
267            let multicast_target = IpAddr::V6(PVA_MULTICAST_V6);
268            targets.push(SearchTarget {
269                target: multicast_target,
270                bind: IpAddr::V6(ip),
271            });
272            if !added_v6_multicast {
273                added_v6_multicast = true;
274                fallback_targets.push(SearchTarget {
275                    target: multicast_target,
276                    bind: IpAddr::V6(Ipv6Addr::UNSPECIFIED),
277                });
278            }
279        }
280    }
281    targets.extend(fallback_targets);
282    targets
283}
284
285/// PVA multicast group (IPv4).
286const PVA_MULTICAST_V4: Ipv4Addr = Ipv4Addr::new(224, 0, 0, 128);
287
288/// PVA multicast group (IPv6 link-local, ff02::42:1).
289const PVA_MULTICAST_V6: Ipv6Addr = Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 0x42, 1);
290
291/// Best-effort join the PVA multicast group appropriate for the bind address.
292fn join_multicast_any(socket: &std::net::UdpSocket, bind: IpAddr) {
293    match bind {
294        IpAddr::V4(iface) => {
295            let _ = socket.join_multicast_v4(&PVA_MULTICAST_V4, &iface);
296        }
297        IpAddr::V6(_) => {
298            // interface index 0 = OS picks the default interface
299            let _ = socket.join_multicast_v6(&PVA_MULTICAST_V6, 0);
300        }
301    }
302}
303
304fn decode_search_response_addr(addr: [u8; 16], port: u16, src: SocketAddr) -> SocketAddr {
305    socket_addr_from_pva_bytes(addr, port)
306        .filter(|a| !a.ip().is_unspecified())
307        .unwrap_or_else(|| SocketAddr::new(src.ip(), port))
308}
309
310fn normalize_discovered_servers(items: Vec<DiscoveredServer>) -> Vec<DiscoveredServer> {
311    let mut seen = HashSet::new();
312    let mut out = Vec::new();
313    for item in items {
314        if seen.insert((item.guid, item.tcp_addr)) {
315            out.push(item);
316        }
317    }
318    out.sort_by(|a, b| a.tcp_addr.to_string().cmp(&b.tcp_addr.to_string()));
319    out
320}
321
322/// Create a UDP socket with SO_REUSEADDR set (matching PVXS behaviour),
323/// allowing multiple processes to share the search port.
324///
325/// On Windows SO_REUSEADDR has different (unsafe) semantics — it allows
326/// a second socket to steal an actively-used port — so we only enable it
327/// on Unix where it merely permits rebinding during TIME_WAIT.
328fn bind_udp_reuse(addr: SocketAddr) -> std::io::Result<std::net::UdpSocket> {
329    let domain = if addr.is_ipv4() {
330        Domain::IPV4
331    } else {
332        Domain::IPV6
333    };
334    let sock = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP))?;
335    #[cfg(unix)]
336    sock.set_reuse_address(true)?;
337    sock.set_nonblocking(true)?;
338    sock.bind(&addr.into())?;
339    Ok(sock.into())
340}
341
342pub async fn search_pv(
343    pv_name: &str,
344    udp_port: u16,
345    timeout_dur: Duration,
346    targets: &[SearchTarget],
347    debug_enabled: bool,
348) -> Result<SocketAddr, PvGetError> {
349    if targets.is_empty() {
350        return Err(PvGetError::Search("no search targets"));
351    }
352
353    let now = std::time::SystemTime::now()
354        .duration_since(std::time::UNIX_EPOCH)
355        .unwrap_or_default();
356    let seq = (now.as_nanos() as u32).wrapping_add(std::process::id());
357    let cid = seq ^ 0x9E37_79B9;
358
359    let mut last_io_error: Option<std::io::Error> = None;
360    let deadline = tokio::time::Instant::now() + timeout_dur;
361
362    // Group targets by bind address so we can share a socket per bind.
363    let mut bind_groups: Vec<(IpAddr, Vec<IpAddr>)> = Vec::new();
364    for t in targets {
365        if let Some(group) = bind_groups.iter_mut().find(|(b, _)| *b == t.bind) {
366            group.1.push(t.target);
367        } else {
368            bind_groups.push((t.bind, vec![t.target]));
369        }
370    }
371
372    // Open sockets and send to all targets first, then collect responses.
373    // Store (socket, message, destinations) for retransmission.
374    let mut socket_info: Vec<(Arc<UdpSocket>, Vec<u8>, Vec<SocketAddr>)> = Vec::new();
375
376    for (bind_ip, group_targets) in &bind_groups {
377        // Always use an ephemeral port for the search client socket.
378        // We only receive unicast replies, so sharing the server's search
379        // port is unnecessary — and on Linux with SO_REUSEPORT the kernel
380        // would route our own outbound packet back to us instead of the
381        // server.
382        let bind_addr = SocketAddr::new(*bind_ip, 0);
383        let (std_sock, actual_bind_addr) = match bind_udp_reuse(bind_addr) {
384            Ok(sock) => {
385                let actual = sock.local_addr().unwrap_or(bind_addr);
386                (sock, actual)
387            }
388            Err(err) => {
389                if debug_enabled {
390                    debug!(
391                        "pva search skipping bind={} step=bind kind={:?} err={}",
392                        bind_addr,
393                        err.kind(),
394                        err
395                    );
396                }
397                last_io_error = Some(err);
398                continue;
399            }
400        };
401        if let Err(err) = std_sock.set_broadcast(true) {
402            if debug_enabled {
403                debug!(
404                    "pva search skipping bind={} step=set_broadcast kind={:?} err={}",
405                    bind_addr,
406                    err.kind(),
407                    err
408                );
409            }
410            last_io_error = Some(err);
411            continue;
412        }
413
414        join_multicast_any(&std_sock, *bind_ip);
415
416        let reply_addr = ip_to_bytes(*bind_ip);
417        let reply_port = match std_sock.local_addr() {
418            Ok(addr) => addr.port(),
419            Err(err) => {
420                if debug_enabled {
421                    debug!(
422                        "pva search skipping bind={} step=local_addr kind={:?} err={}",
423                        bind_addr,
424                        err.kind(),
425                        err
426                    );
427                }
428                last_io_error = Some(err);
429                continue;
430            }
431        };
432        let requests = [(cid, pv_name)];
433        let msg = encode_search_request(seq, 0x81, reply_port, reply_addr, &requests, 2, false);
434
435        let socket = match UdpSocket::from_std(std_sock) {
436            Ok(socket) => socket,
437            Err(err) => {
438                if debug_enabled {
439                    debug!(
440                        "pva search skipping bind={} step=from_std kind={:?} err={}",
441                        bind_addr,
442                        err.kind(),
443                        err
444                    );
445                }
446                last_io_error = Some(err);
447                continue;
448            }
449        };
450
451        let dests: Vec<SocketAddr> = group_targets
452            .iter()
453            .map(|ip| SocketAddr::new(*ip, udp_port))
454            .collect();
455
456        // Send to every target in this bind group immediately.
457        for dest in &dests {
458            if debug_enabled {
459                debug!(
460                    "pva search bind={} target={} server_port={} reply_port={}",
461                    actual_bind_addr,
462                    dest.ip(),
463                    udp_port,
464                    reply_port
465                );
466                debug!("pva search seq={} cid={}", seq, cid);
467                debug!("pva search send {} bytes to {}", msg.len(), dest);
468            }
469            if let Err(err) = socket.send_to(&msg, dest).await {
470                if debug_enabled {
471                    debug!(
472                        "pva search send_to target={} kind={:?} err={}",
473                        dest,
474                        err.kind(),
475                        err
476                    );
477                }
478                last_io_error = Some(err);
479            }
480        }
481
482        socket_info.push((Arc::new(socket), msg, dests));
483    }
484
485    if socket_info.is_empty() {
486        if let Some(err) = last_io_error {
487            return Err(PvGetError::Io(err));
488        }
489        return Err(PvGetError::Timeout("search response"));
490    }
491
492    // Spawn a receiver task per socket that forwards packets into a shared channel.
493    let (tx, mut rx) = tokio::sync::mpsc::channel::<(Vec<u8>, SocketAddr)>(64);
494    for (sock, _, _) in &socket_info {
495        let sock = Arc::clone(sock);
496        let tx = tx.clone();
497        tokio::spawn(async move {
498            loop {
499                let mut buf = vec![0u8; 2048];
500                match sock.recv_from(&mut buf).await {
501                    Ok((len, src)) => {
502                        buf.truncate(len);
503                        if tx.send((buf, src)).await.is_err() {
504                            break;
505                        }
506                    }
507                    Err(_) => break,
508                }
509            }
510        });
511    }
512    drop(tx); // Only spawned tasks hold senders; channel closes when they exit.
513
514    // Retransmit schedule: exponential backoff from start.
515    let retransmit_offsets = [100u64, 500, 1000, 2000];
516    let start = tokio::time::Instant::now();
517    let mut next_retransmit = 0usize;
518
519    loop {
520        // Compute the next wake-up: either the next retransmit or the deadline.
521        let next_retransmit_at = if next_retransmit < retransmit_offsets.len() {
522            start + Duration::from_millis(retransmit_offsets[next_retransmit])
523        } else {
524            deadline
525        };
526        let wake_at = next_retransmit_at.min(deadline);
527
528        tokio::select! {
529            recv = rx.recv() => {
530                let Some((buf, src)) = recv else { break };
531                let mut pkt = PvaPacket::new(&buf);
532                let cmd = pkt
533                    .decode_payload()
534                    .ok_or(PvGetError::Search("failed to decode search response"))?;
535                if let PvaPacketCommand::SearchResponse(payload) = cmd {
536                    if debug_enabled {
537                        debug!(
538                            "pva search response found={} cids={:?} addr={:?} port={}",
539                            payload.found, payload.cids, payload.addr, payload.port
540                        );
541                    }
542                    if payload.seq != seq {
543                        continue;
544                    }
545                    if !payload.protocol.is_empty() && !payload.protocol.eq_ignore_ascii_case("tcp") {
546                        continue;
547                    }
548                    if !payload.found {
549                        continue;
550                    }
551                    if !payload.cids.is_empty() && !payload.cids.contains(&cid) {
552                        continue;
553                    }
554
555                    let addr = decode_search_response_addr(payload.addr, payload.port, src);
556                    if debug_enabled {
557                        debug!("pva search response from {}", addr);
558                    }
559                    return Ok(addr);
560                }
561            }
562            _ = tokio::time::sleep_until(wake_at) => {
563                if tokio::time::Instant::now() >= deadline {
564                    break;
565                }
566                // Retransmit to all targets on all sockets.
567                if next_retransmit < retransmit_offsets.len() {
568                    if debug_enabled {
569                        debug!("pva search retransmit round {}", next_retransmit + 1);
570                    }
571                    for (sock, msg, dests) in &socket_info {
572                        for dest in dests {
573                            let _ = sock.send_to(msg, dest).await;
574                        }
575                    }
576                    next_retransmit += 1;
577                }
578            }
579        }
580    }
581
582    Err(PvGetError::Timeout("search response"))
583}
584
585pub fn default_bind_ip() -> Option<IpAddr> {
586    choose_default_bind_v4()
587        .map(IpAddr::V4)
588        .or_else(|| choose_default_bind_v6().map(IpAddr::V6))
589}
590
591/// Parse `EPICS_PVA_NAME_SERVERS` value into socket addresses.
592/// Accepts space/comma separated entries: `host:port`, `ip`, `hostname`
593/// (port defaults to 5075).
594pub fn parse_name_servers(env_val: &str) -> Vec<SocketAddr> {
595    let mut out = Vec::new();
596    for token in env_val.split(|c| c == ',' || c == ' ' || c == '\t') {
597        let token = token.trim();
598        if token.is_empty() {
599            continue;
600        }
601        if let Ok(addr) = token.parse::<SocketAddr>() {
602            out.push(addr);
603            continue;
604        }
605        if let Ok(ip) = token.parse::<IpAddr>() {
606            out.push(SocketAddr::new(ip, 5075));
607            continue;
608        }
609        use std::net::ToSocketAddrs;
610        if let Ok(mut addrs) = token.to_socket_addrs() {
611            if let Some(addr) = addrs.next() {
612                out.push(addr);
613                continue;
614            }
615        }
616        let with_port = format!("{}:5075", token);
617        if let Ok(mut addrs) = with_port.to_socket_addrs() {
618            if let Some(addr) = addrs.next() {
619                out.push(addr);
620            }
621        }
622    }
623    out
624}
625
626/// Build a minimal PVA ConnectionValidation response for name server search.
627fn encode_search_validation(version: u8, is_be: bool) -> Vec<u8> {
628    let user = default_authnz_user();
629    let host = default_authnz_host();
630    encode_client_connection_validation(87_040, 32_767, 0, "ca", &user, &host, version, is_be)
631}
632
633/// Search for a PV via a TCP connection to a PVA name server.
634///
635/// Connects to the name server, performs the PVA handshake, sends a search
636/// request over TCP, and returns the server address from the search response.
637pub async fn search_pv_tcp(
638    pv_name: &str,
639    name_server: SocketAddr,
640    timeout_dur: Duration,
641    debug_enabled: bool,
642) -> Result<SocketAddr, PvGetError> {
643    let deadline = tokio::time::Instant::now() + timeout_dur;
644
645    let mut stream = tokio::time::timeout(timeout_dur, tokio::net::TcpStream::connect(name_server))
646        .await
647        .map_err(|_| PvGetError::Timeout("name server connect"))??;
648
649    // One reassembler for this connection, shared by every read below.
650    let mut reassembler = SegmentReassembler::new();
651
652    let mut version = 2u8;
653    let mut is_be = false;
654
655    // Read SET_BYTE_ORDER + ConnectionValidation from name server.
656    for _ in 0..2 {
657        let now = tokio::time::Instant::now();
658        if now >= deadline {
659            return Err(PvGetError::Timeout("name server handshake"));
660        }
661        let remaining = deadline - now;
662        if let Ok(bytes) = read_packet(&mut stream, remaining, &mut reassembler).await {
663            let mut pkt = PvaPacket::new(&bytes);
664            if let Some(cmd) = pkt.decode_payload() {
665                match cmd {
666                    PvaPacketCommand::Control(payload) => {
667                        if payload.command == 2 {
668                            is_be = pkt.header.flags.is_msb;
669                        }
670                    }
671                    PvaPacketCommand::ConnectionValidation(_) => {
672                        version = pkt.header.version;
673                        is_be = pkt.header.flags.is_msb;
674                    }
675                    _ => {}
676                }
677            }
678        }
679    }
680
681    let validation = encode_search_validation(version, is_be);
682    stream.write_all(&validation).await?;
683
684    // Wait for ConnectionValidated.
685    loop {
686        let now = tokio::time::Instant::now();
687        if now >= deadline {
688            return Err(PvGetError::Timeout("name server validated"));
689        }
690        let remaining = deadline - now;
691        let bytes = read_packet(&mut stream, remaining, &mut reassembler).await?;
692        let mut pkt = PvaPacket::new(&bytes);
693        if let Some(cmd) = pkt.decode_payload() {
694            if matches!(cmd, PvaPacketCommand::ConnectionValidated(_)) {
695                break;
696            }
697        }
698    }
699
700    // Send search request over TCP.
701    let now_ts = std::time::SystemTime::now()
702        .duration_since(std::time::UNIX_EPOCH)
703        .unwrap_or_default();
704    let seq = (now_ts.as_nanos() as u32).wrapping_add(std::process::id());
705    let cid = seq ^ 0x9E37_79B9;
706    let requests = [(cid, pv_name)];
707    let msg = encode_search_request(seq, 0x80, 0, [0u8; 16], &requests, version, is_be);
708    stream.write_all(&msg).await?;
709
710    if debug_enabled {
711        debug!(
712            "pva tcp search sent to name_server={} pv={}",
713            name_server, pv_name
714        );
715    }
716
717    // Read search response.
718    loop {
719        let now = tokio::time::Instant::now();
720        if now >= deadline {
721            return Err(PvGetError::Timeout("name server search response"));
722        }
723        let remaining = deadline - now;
724        let bytes = read_packet(&mut stream, remaining, &mut reassembler).await?;
725        let mut pkt = PvaPacket::new(&bytes);
726        if let Some(cmd) = pkt.decode_payload() {
727            if let PvaPacketCommand::SearchResponse(payload) = cmd {
728                if !payload.found {
729                    continue;
730                }
731                if !payload.cids.is_empty() && !payload.cids.contains(&cid) {
732                    continue;
733                }
734                let addr = decode_search_response_addr(payload.addr, payload.port, name_server);
735                if debug_enabled {
736                    debug!(
737                        "pva tcp search response from name_server={}: {}",
738                        name_server, addr
739                    );
740                }
741                return Ok(addr);
742            }
743        }
744    }
745}
746
747/// Resolve the PVA server for a PV using name servers (TCP) and/or UDP search.
748///
749/// - If `opts.server_addr` is set, returns it directly.
750/// - Tries each name server from `opts.name_servers` and `EPICS_PVA_NAME_SERVERS`
751///   via TCP search.
752/// - Falls back to UDP search using `build_search_targets()`.
753pub async fn resolve_pv_server(opts: &PvGetOptions) -> Result<SocketAddr, PvGetError> {
754    if let Some(addr) = opts.server_addr {
755        return Ok(addr);
756    }
757
758    let mut name_servers = opts.name_servers.clone();
759    if let Ok(env) = std::env::var("EPICS_PVA_NAME_SERVERS") {
760        name_servers.extend(parse_name_servers(&env));
761    }
762
763    let no_broadcast = opts.no_broadcast;
764
765    // Fail fast when no search strategy is available.
766    if no_broadcast && name_servers.is_empty() {
767        return Err(PvGetError::Search(
768            "no search strategy: specify --name-server or --server when using --no-broadcast",
769        ));
770    }
771
772    // Launch all search strategies concurrently — TCP name servers + UDP broadcast.
773    // Return the first successful result.
774    let targets = build_search_targets(opts.search_addr, opts.bind_addr);
775
776    let pv = opts.pv_name.clone();
777    let timeout_dur = opts.timeout;
778    let debug_enabled = opts.debug;
779    let udp_port = opts.udp_port;
780
781    let mut set = tokio::task::JoinSet::new();
782
783    for ns in name_servers {
784        let pv = pv.clone();
785        set.spawn(async move {
786            let addr = search_pv_tcp(&pv, ns, timeout_dur, debug_enabled).await?;
787            Ok::<SocketAddr, PvGetError>(addr)
788        });
789    }
790
791    if !no_broadcast {
792        let pv = pv.clone();
793        let targets = targets.clone();
794        set.spawn(async move {
795            let addr = search_pv(&pv, udp_port, timeout_dur, &targets, debug_enabled).await?;
796            Ok(addr)
797        });
798    }
799
800    let mut last_err = None;
801    while let Some(result) = set.join_next().await {
802        match result {
803            Ok(Ok(addr)) => {
804                set.abort_all();
805                return Ok(addr);
806            }
807            Ok(Err(e)) => {
808                if debug_enabled {
809                    debug!("pva search strategy failed: {}", e);
810                }
811                last_err = Some(e);
812            }
813            Err(join_err) => {
814                if debug_enabled {
815                    debug!("pva search task panicked: {}", join_err);
816                }
817            }
818        }
819    }
820
821    Err(last_err.unwrap_or(PvGetError::Timeout("search response")))
822}
823
824pub async fn discover_servers(
825    udp_port: u16,
826    timeout_dur: Duration,
827    targets: &[SearchTarget],
828    debug_enabled: bool,
829) -> Result<Vec<DiscoveredServer>, PvGetError> {
830    if targets.is_empty() {
831        return Err(PvGetError::Search("no search targets"));
832    }
833
834    let now = std::time::SystemTime::now()
835        .duration_since(std::time::UNIX_EPOCH)
836        .unwrap_or_default();
837    let seq = (now.as_nanos() as u32).wrapping_add(std::process::id());
838
839    let mut found: Vec<DiscoveredServer> = Vec::new();
840    let mut last_io_error: Option<std::io::Error> = None;
841    let deadline = tokio::time::Instant::now() + timeout_dur;
842
843    // Group targets by bind address so we can share a socket per bind.
844    let mut bind_groups: Vec<(IpAddr, Vec<IpAddr>)> = Vec::new();
845    for t in targets {
846        if let Some(group) = bind_groups.iter_mut().find(|(b, _)| *b == t.bind) {
847            group.1.push(t.target);
848        } else {
849            bind_groups.push((t.bind, vec![t.target]));
850        }
851    }
852
853    // Open sockets and send to all targets first, then collect responses.
854    // Store (socket, message, destinations) for retransmission.
855    let mut socket_info: Vec<(Arc<UdpSocket>, Vec<u8>, Vec<SocketAddr>)> = Vec::new();
856
857    for (bind_ip, group_targets) in &bind_groups {
858        // Always use an ephemeral port for the discovery client socket.
859        // We only receive unicast replies, so sharing the server's search
860        // port is unnecessary — and on Linux with SO_REUSEPORT the kernel
861        // would route our own outbound packet back to us instead of the
862        // server.
863        let bind_addr = SocketAddr::new(*bind_ip, 0);
864        let (std_sock, actual_bind_addr) = match bind_udp_reuse(bind_addr) {
865            Ok(sock) => {
866                let actual = sock.local_addr().unwrap_or(bind_addr);
867                (sock, actual)
868            }
869            Err(err) => {
870                if debug_enabled {
871                    debug!(
872                        "pva discover skipping bind={} step=bind kind={:?} err={}",
873                        bind_addr,
874                        err.kind(),
875                        err
876                    );
877                }
878                last_io_error = Some(err);
879                continue;
880            }
881        };
882        if let Err(err) = std_sock.set_broadcast(true) {
883            if debug_enabled {
884                debug!(
885                    "pva discover skipping bind={} step=set_broadcast kind={:?} err={}",
886                    bind_addr,
887                    err.kind(),
888                    err
889                );
890            }
891            last_io_error = Some(err);
892            continue;
893        }
894
895        join_multicast_any(&std_sock, *bind_ip);
896
897        let reply_addr = ip_to_bytes(*bind_ip);
898        let reply_port = match std_sock.local_addr() {
899            Ok(addr) => addr.port(),
900            Err(err) => {
901                if debug_enabled {
902                    debug!(
903                        "pva discover skipping bind={} step=local_addr kind={:?} err={}",
904                        bind_addr,
905                        err.kind(),
906                        err
907                    );
908                }
909                last_io_error = Some(err);
910                continue;
911            }
912        };
913        let msg = encode_search_request(seq, 0x81, reply_port, reply_addr, &[], 2, false);
914
915        let socket = match UdpSocket::from_std(std_sock) {
916            Ok(socket) => socket,
917            Err(err) => {
918                if debug_enabled {
919                    debug!(
920                        "pva discover skipping bind={} step=from_std kind={:?} err={}",
921                        bind_addr,
922                        err.kind(),
923                        err
924                    );
925                }
926                last_io_error = Some(err);
927                continue;
928            }
929        };
930
931        let dests: Vec<SocketAddr> = group_targets
932            .iter()
933            .map(|ip| SocketAddr::new(*ip, udp_port))
934            .collect();
935
936        // Send to every target in this bind group immediately.
937        for dest in &dests {
938            if debug_enabled {
939                debug!(
940                    "pva discover bind={} target={} server_port={} reply_port={} seq={}",
941                    actual_bind_addr,
942                    dest.ip(),
943                    udp_port,
944                    reply_port,
945                    seq
946                );
947            }
948            if let Err(err) = socket.send_to(&msg, dest).await {
949                if debug_enabled {
950                    debug!(
951                        "pva discover send_to target={} kind={:?} err={}",
952                        dest,
953                        err.kind(),
954                        err
955                    );
956                }
957                last_io_error = Some(err);
958            }
959        }
960
961        socket_info.push((Arc::new(socket), msg, dests));
962    }
963
964    if socket_info.is_empty() {
965        if let Some(err) = last_io_error {
966            return Err(PvGetError::Io(err));
967        }
968        return Err(PvGetError::Search("no search targets"));
969    }
970
971    // Spawn a receiver task per socket that forwards packets into a shared channel.
972    let (tx, mut rx) = tokio::sync::mpsc::channel::<(Vec<u8>, SocketAddr)>(64);
973    for (sock, _, _) in &socket_info {
974        let sock = Arc::clone(sock);
975        let tx = tx.clone();
976        tokio::spawn(async move {
977            loop {
978                let mut buf = vec![0u8; 2048];
979                match sock.recv_from(&mut buf).await {
980                    Ok((len, src)) => {
981                        buf.truncate(len);
982                        if tx.send((buf, src)).await.is_err() {
983                            break;
984                        }
985                    }
986                    Err(_) => break,
987                }
988            }
989        });
990    }
991    drop(tx); // Only spawned tasks hold senders; channel closes when they exit.
992
993    // Retransmit schedule: exponential backoff from start.
994    let retransmit_offsets = [100u64, 500, 1000, 2000];
995    let start = tokio::time::Instant::now();
996    let mut next_retransmit = 0usize;
997
998    loop {
999        // Compute the next wake-up: either the next retransmit or the deadline.
1000        let next_retransmit_at = if next_retransmit < retransmit_offsets.len() {
1001            start + Duration::from_millis(retransmit_offsets[next_retransmit])
1002        } else {
1003            deadline
1004        };
1005        let wake_at = next_retransmit_at.min(deadline);
1006
1007        tokio::select! {
1008            recv = rx.recv() => {
1009                let Some((buf, src)) = recv else { break };
1010                let mut pkt = PvaPacket::new(&buf);
1011                let Some(cmd) = pkt.decode_payload() else {
1012                    continue;
1013                };
1014                if let PvaPacketCommand::SearchResponse(payload) = cmd {
1015                    if payload.seq != seq {
1016                        continue;
1017                    }
1018                    if !payload.protocol.is_empty() && !payload.protocol.eq_ignore_ascii_case("tcp") {
1019                        continue;
1020                    }
1021                    let tcp_addr = decode_search_response_addr(payload.addr, payload.port, src);
1022                    found.push(DiscoveredServer {
1023                        guid: payload.guid,
1024                        tcp_addr,
1025                    });
1026                }
1027            }
1028            _ = tokio::time::sleep_until(wake_at) => {
1029                if tokio::time::Instant::now() >= deadline {
1030                    break;
1031                }
1032                // Retransmit to all targets on all sockets.
1033                if next_retransmit < retransmit_offsets.len() {
1034                    if debug_enabled {
1035                        debug!("pva discover retransmit round {}", next_retransmit + 1);
1036                    }
1037                    for (sock, msg, dests) in &socket_info {
1038                        for dest in dests {
1039                            let _ = sock.send_to(msg, dest).await;
1040                        }
1041                    }
1042                    next_retransmit += 1;
1043                }
1044            }
1045        }
1046    }
1047
1048    Ok(normalize_discovered_servers(found))
1049}
1050
1051#[cfg(test)]
1052mod tests {
1053    use super::*;
1054    use spvirit_codec::epics_decode::{PvaPacket, PvaPacketCommand};
1055
1056    #[test]
1057    fn encode_decode_search_request_roundtrip() {
1058        let seq = 1234;
1059        let cid = 42;
1060        let port = 5076;
1061        let pv_name = "TEST:PV";
1062        let reply_addr = ip_to_bytes(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 20)));
1063        let requests = [(cid, pv_name)];
1064        let msg = encode_search_request(seq, 0x81, port, reply_addr, &requests, 2, false);
1065        let mut pkt = PvaPacket::new(&msg);
1066        let cmd = pkt.decode_payload().expect("decoded");
1067        match cmd {
1068            PvaPacketCommand::Search(payload) => {
1069                assert_eq!(payload.seq, seq);
1070                assert_eq!(payload.mask, 0x81);
1071                assert_eq!(payload.addr, reply_addr);
1072                assert_eq!(payload.port, port);
1073                assert_eq!(payload.protocols, vec!["tcp".to_string()]);
1074                assert_eq!(payload.pv_requests.len(), 1);
1075                assert_eq!(payload.pv_requests[0].0, cid);
1076                assert_eq!(payload.pv_requests[0].1, pv_name.to_string());
1077            }
1078            other => panic!("unexpected decode: {:?}", other),
1079        }
1080    }
1081
1082    #[test]
1083    fn encode_decode_server_discovery_request_roundtrip() {
1084        let seq = 4321;
1085        let port = 5076;
1086        let reply_addr = ip_to_bytes(IpAddr::V4(Ipv4Addr::new(10, 20, 30, 40)));
1087        let msg = encode_search_request(seq, 0x81, port, reply_addr, &[], 2, false);
1088        let mut pkt = PvaPacket::new(&msg);
1089        let cmd = pkt.decode_payload().expect("decoded");
1090        match cmd {
1091            PvaPacketCommand::Search(payload) => {
1092                assert_eq!(payload.seq, seq);
1093                assert_eq!(payload.pv_requests.len(), 0);
1094                assert_eq!(payload.protocols, vec!["tcp".to_string()]);
1095            }
1096            other => panic!("unexpected decode: {:?}", other),
1097        }
1098    }
1099
1100    #[test]
1101    fn normalize_discovered_servers_deduplicates_by_guid_and_addr() {
1102        let guid = [1u8; 12];
1103        let s1 = DiscoveredServer {
1104            guid,
1105            tcp_addr: "127.0.0.1:5075".parse().unwrap(),
1106        };
1107        let s2 = DiscoveredServer {
1108            guid,
1109            tcp_addr: "127.0.0.1:5075".parse().unwrap(),
1110        };
1111        let s3 = DiscoveredServer {
1112            guid: [2u8; 12],
1113            tcp_addr: "127.0.0.1:5075".parse().unwrap(),
1114        };
1115        let normalized = normalize_discovered_servers(vec![s1, s2, s3]);
1116        assert_eq!(normalized.len(), 2);
1117    }
1118
1119    #[test]
1120    fn parse_addr_list_accepts_ip_and_ip_port() {
1121        let items = parse_addr_list("192.168.1.10 10.0.0.1:5076");
1122        assert!(items.contains(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 10))));
1123        assert!(items.contains(&IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))));
1124    }
1125
1126    #[test]
1127    fn discovery_target_falls_back_to_limited_broadcast_for_invalid_netmask() {
1128        let ip = Ipv4Addr::new(130, 246, 90, 92);
1129        assert_eq!(
1130            discovery_target_for(ip, Ipv4Addr::new(255, 255, 255, 255)),
1131            Ipv4Addr::new(255, 255, 255, 255)
1132        );
1133        assert_eq!(
1134            discovery_target_for(ip, Ipv4Addr::new(0, 0, 0, 0)),
1135            Ipv4Addr::new(255, 255, 255, 255)
1136        );
1137    }
1138
1139    #[test]
1140    fn discovery_target_uses_directed_broadcast_for_normal_subnet() {
1141        let ip = Ipv4Addr::new(192, 168, 56, 1);
1142        let netmask = Ipv4Addr::new(255, 255, 255, 0);
1143        assert_eq!(
1144            discovery_target_for(ip, netmask),
1145            Ipv4Addr::new(192, 168, 56, 255)
1146        );
1147    }
1148
1149    #[test]
1150    fn parse_name_servers_ip_with_port() {
1151        let addrs = parse_name_servers("192.168.1.10:5075");
1152        assert_eq!(
1153            addrs,
1154            vec!["192.168.1.10:5075".parse::<SocketAddr>().unwrap()]
1155        );
1156    }
1157
1158    #[test]
1159    fn parse_name_servers_ip_without_port_defaults_to_5075() {
1160        let addrs = parse_name_servers("10.0.0.1");
1161        assert_eq!(
1162            addrs,
1163            vec![SocketAddr::new(
1164                IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)),
1165                5075
1166            )]
1167        );
1168    }
1169
1170    #[test]
1171    fn parse_name_servers_multiple_comma_separated() {
1172        let addrs = parse_name_servers("10.0.0.1:5075,10.0.0.2:9876");
1173        assert_eq!(addrs.len(), 2);
1174        assert_eq!(addrs[0], "10.0.0.1:5075".parse::<SocketAddr>().unwrap());
1175        assert_eq!(addrs[1], "10.0.0.2:9876".parse::<SocketAddr>().unwrap());
1176    }
1177
1178    #[test]
1179    fn parse_name_servers_multiple_space_separated() {
1180        let addrs = parse_name_servers("10.0.0.1 10.0.0.2:5075");
1181        assert_eq!(addrs.len(), 2);
1182        assert_eq!(
1183            addrs[0],
1184            SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)), 5075)
1185        );
1186        assert_eq!(addrs[1], "10.0.0.2:5075".parse::<SocketAddr>().unwrap());
1187    }
1188
1189    #[test]
1190    fn parse_name_servers_empty_string() {
1191        let addrs = parse_name_servers("");
1192        assert!(addrs.is_empty());
1193    }
1194
1195    #[test]
1196    fn parse_name_servers_whitespace_only() {
1197        let addrs = parse_name_servers("  \t  ");
1198        assert!(addrs.is_empty());
1199    }
1200
1201    #[test]
1202    fn parse_name_servers_mixed_separators() {
1203        let addrs = parse_name_servers("10.0.0.1:5075, 10.0.0.2  ,  10.0.0.3:9999");
1204        assert_eq!(addrs.len(), 3);
1205        assert_eq!(addrs[0], "10.0.0.1:5075".parse::<SocketAddr>().unwrap());
1206        assert_eq!(
1207            addrs[1],
1208            SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2)), 5075)
1209        );
1210        assert_eq!(addrs[2], "10.0.0.3:9999".parse::<SocketAddr>().unwrap());
1211    }
1212
1213    #[test]
1214    fn parse_name_servers_ipv6_with_port() {
1215        let addrs = parse_name_servers("[::1]:5075");
1216        assert_eq!(
1217            addrs,
1218            vec![SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 5075)]
1219        );
1220    }
1221
1222    #[test]
1223    fn parse_name_servers_ipv6_without_port() {
1224        let addrs = parse_name_servers("::1");
1225        assert_eq!(
1226            addrs,
1227            vec![SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 5075)]
1228        );
1229    }
1230
1231    #[test]
1232    fn decode_search_response_addr_falls_back_to_udp_source_when_unspecified() {
1233        let src: SocketAddr = "192.168.1.20:5076".parse().unwrap();
1234        let decoded = decode_search_response_addr([0u8; 16], 5075, src);
1235        assert_eq!(decoded, "192.168.1.20:5075".parse().unwrap());
1236    }
1237}