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 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 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 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
96fn 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 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 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 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 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 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 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; }
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
285const PVA_MULTICAST_V4: Ipv4Addr = Ipv4Addr::new(224, 0, 0, 128);
287
288const PVA_MULTICAST_V6: Ipv6Addr = Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 0x42, 1);
290
291fn 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 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
322fn 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 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 let mut socket_info: Vec<(Arc<UdpSocket>, Vec<u8>, Vec<SocketAddr>)> = Vec::new();
375
376 for (bind_ip, group_targets) in &bind_groups {
377 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 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 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); let retransmit_offsets = [100u64, 500, 1000, 2000];
516 let start = tokio::time::Instant::now();
517 let mut next_retransmit = 0usize;
518
519 loop {
520 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 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
591pub 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
626fn 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
633pub 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 let mut reassembler = SegmentReassembler::new();
651
652 let mut version = 2u8;
653 let mut is_be = false;
654
655 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 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 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 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
747pub 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 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 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 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 let mut socket_info: Vec<(Arc<UdpSocket>, Vec<u8>, Vec<SocketAddr>)> = Vec::new();
856
857 for (bind_ip, group_targets) in &bind_groups {
858 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 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 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); let retransmit_offsets = [100u64, 500, 1000, 2000];
995 let start = tokio::time::Instant::now();
996 let mut next_retransmit = 0usize;
997
998 loop {
999 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 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}