use std::{
collections::HashMap,
net::{IpAddr, SocketAddr},
sync::{Arc, Mutex},
time::{Duration, Instant as StdInstant, SystemTime},
};
use agnostic_lite::RuntimeLite;
use agnostic_net::{Net, UdpSocket};
use async_channel::Sender;
use futures::{FutureExt, pin_mut, select_biased};
use mdns_proto::{
QueryHandle, QuerySpec, ServiceHandle, ServiceSpec, ServiceUpdate, endpoint::WithdrawalSend,
event::RouteEvent,
};
use rand::{SeedableRng, rngs::StdRng};
use slab::Slab;
use crate::{
command::{Command, QueryStarted, ServiceRegistered},
error::{RegisterError, StartQueryError},
options::ServerOptions,
proto::{ProtoEndpoint, ProtoService},
query::{QueryMailbox, new_mailbox},
service::{ServiceMailbox, new_service_mailbox},
};
#[cfg(test)]
mod tests;
pub(crate) struct BoundSockets<N: Net> {
pub(crate) v4: Option<N::UdpSocket>,
pub(crate) v6: Option<N::UdpSocket>,
pub(crate) interface_index: u32,
}
struct Packet {
src: SocketAddr,
data: Vec<u8>,
local_ip: IpAddr,
interface_index: u32,
kernel_rx_time: Option<SystemTime>,
read_time: SystemTime,
hop_limit: Option<u8>,
}
struct ServiceCtx {
proto: ProtoService,
mailbox: Arc<Mutex<ServiceMailbox>>,
doorbell: Sender<()>,
encode_failures: u8,
withdrawing: bool,
}
struct QueryCtx {
mailbox: Arc<Mutex<QueryMailbox>>,
doorbell: Sender<()>,
last_seq: u64,
}
struct DriverState<N: Net> {
endpoint: ProtoEndpoint,
services: HashMap<ServiceHandle, ServiceCtx>,
queries: HashMap<QueryHandle, QueryCtx>,
v4: Option<Arc<N::UdpSocket>>,
v6: Option<Arc<N::UdpSocket>>,
#[cfg(feature = "stats")]
stats: std::sync::Arc<hick_trace::stats::Stats>,
recent_sends: Vec<(u64, SystemTime)>,
completed_withdrawals: Vec<ServiceHandle>,
svc_handle_scratch: Vec<ServiceHandle>,
query_handle_scratch: Vec<QueryHandle>,
local_subnets: Vec<(IpAddr, u8)>,
bound_interface: u32,
}
impl<N: Net> DriverState<N> {
fn new(opts: &ServerOptions, sockets: BoundSockets<N>) -> Self {
let rng = StdRng::from_rng(&mut rand::rng());
let endpoint = ProtoEndpoint::try_new(*opts.endpoint_config(), rng);
let bound_interface = sockets.interface_index;
#[cfg(feature = "stats")]
let stats = endpoint.stats_handle();
Self {
endpoint,
services: HashMap::new(),
queries: HashMap::new(),
recent_sends: Vec::new(),
completed_withdrawals: Vec::new(),
svc_handle_scratch: Vec::new(),
query_handle_scratch: Vec::new(),
local_subnets: collect_local_subnets(bound_interface),
bound_interface,
v4: sockets.v4.map(Arc::new),
v6: sockets.v6.map(Arc::new),
#[cfg(feature = "stats")]
stats,
}
}
fn next_withdrawal_deadline(&self) -> Option<StdInstant> {
self.endpoint.next_withdrawal_deadline()
}
fn next_deadline(&self) -> Option<StdInstant> {
let mut best: Option<StdInstant> = self.endpoint.poll_timeout();
for ctx in self.services.values() {
if ctx.withdrawing {
continue;
}
if let Some(t) = ctx.proto.poll_timeout() {
best = Some(min_opt(best, t));
}
}
for handle in self.queries.keys() {
if let Some(t) = self.endpoint.poll_query_timeout(*handle) {
best = Some(min_opt(best, t));
}
}
best
}
fn handle_command(&mut self, cmd: Command, now: StdInstant) {
match cmd {
Command::RegisterService { spec, reply } => {
let result = self.register_service(spec, now);
if let Ok(ref ok) = result {
let handle = ok.handle;
if let Err(returned) = reply.send(result) {
drop(returned);
self.remove_service(handle, now);
hick_trace::debug!(
?handle,
"RegisterService caller cancelled before reply; rolled back orphan state"
);
}
} else {
let _ = reply.send(result);
}
}
Command::UnregisterService { handle } => {
self.remove_service(handle, now);
}
Command::StartQuery { spec, reply } => {
let result = self.start_query(spec, now);
if let Ok(ref ok) = result {
let handle = ok.handle;
if let Err(returned) = reply.send(result) {
drop(returned);
let _ = self.endpoint.cancel_query(handle);
self.queries.remove(&handle);
hick_trace::debug!(
?handle,
"StartQuery caller cancelled before reply; rolled back orphan state"
);
}
} else {
let _ = reply.send(result);
}
}
Command::CancelQuery { handle } => {
let _ = self.endpoint.cancel_query(handle);
self.queries.remove(&handle);
}
Command::SpawnLookup { task } => {
<N::Runtime as RuntimeLite>::spawn_detach(task);
}
}
}
fn register_service(
&mut self,
spec: ServiceSpec,
now: StdInstant,
) -> Result<ServiceRegistered, RegisterError> {
let (handle, svc) = self
.endpoint
.try_register_service::<Slab<_>, Slab<_>>(spec, now)?;
let (mailbox, doorbell_tx, doorbell_rx) = new_service_mailbox();
self.services.insert(
handle,
ServiceCtx {
proto: svc,
mailbox: Arc::clone(&mailbox),
doorbell: doorbell_tx,
encode_failures: 0,
withdrawing: false,
},
);
Ok(ServiceRegistered {
handle,
mailbox,
doorbell: doorbell_rx,
})
}
fn start_query(
&mut self,
spec: QuerySpec,
now: StdInstant,
) -> Result<QueryStarted, StartQueryError> {
let handle = self
.endpoint
.try_start_query(spec, now)
.map_err(|_| StartQueryError::StorageFull)?;
let (mailbox, doorbell_tx, doorbell_rx) = new_mailbox();
self.queries.insert(
handle,
QueryCtx {
mailbox: Arc::clone(&mailbox),
doorbell: doorbell_tx,
last_seq: 0,
},
);
Ok(QueryStarted {
handle,
mailbox,
doorbell: doorbell_rx,
})
}
fn handle_packet(&mut self, pkt: Packet) {
let on_link = match pkt.hop_limit {
Some(_) => is_on_link(pkt.hop_limit),
None => src_on_local_link(
&self.local_subnets,
self.bound_interface,
pkt.interface_index,
pkt.src.ip(),
),
};
if !on_link {
hick_trace::debug!(
src = %pkt.src,
hop_limit = ?pkt.hop_limit,
"dropping off-link packet (RFC 6762 §11 trust boundary)"
);
#[cfg(feature = "stats")]
{
self.stats.packets_rx(1);
self.stats.bytes_rx(pkt.data.len() as u64);
self.stats.packets_dropped(1);
}
return;
}
if packet_is_response(&pkt.data) && pkt.src.port() != hick_udp::constants::MDNS_PORT {
hick_trace::debug!(
src = %pkt.src,
"dropping untrusted response (source port != 5353) before self-send match"
);
#[cfg(feature = "stats")]
{
self.stats.packets_rx(1);
self.stats.bytes_rx(pkt.data.len() as u64);
self.stats.packets_dropped(1);
}
return;
}
let local_ip = pkt.local_ip;
let interface_index = pkt.interface_index;
let caller_is_self = match pkt.kernel_rx_time {
Some(rx) => take_self_send(&mut self.recent_sends, &pkt.data, rx, MatchMode::Ordered),
None => take_self_send(
&mut self.recent_sends,
&pkt.data,
pkt.read_time,
MatchMode::Degraded,
),
};
let now = StdInstant::now();
let Self {
endpoint, services, ..
} = self;
let route_events = match endpoint.handle(
now,
pkt.src,
local_ip,
interface_index,
&pkt.data,
caller_is_self,
) {
Ok(it) => it,
Err(_e) => {
hick_trace::debug!(error = %_e, src = %pkt.src, "endpoint.handle failed");
return;
}
};
for ev in route_events {
match ev {
Ok(RouteEvent::ToService(ts)) => {
if let Some(ctx) = services.get_mut(&ts.handle())
&& !ctx.withdrawing
{
ctx.proto.handle_event(ts.into_event(), now);
}
}
Ok(_) => {}
Err(_e) => {
hick_trace::debug!(error = %_e, "route event error mid-packet; bailing");
break;
}
}
}
}
fn fire_timeouts(&mut self, now: StdInstant) {
let _ = self.endpoint.handle_timeout(now);
for ctx in self.services.values_mut() {
if ctx.withdrawing {
continue;
}
let _ = ctx.proto.handle_timeout(now);
}
let Self {
endpoint, queries, ..
} = &mut *self;
for &h in queries.keys() {
let _ = endpoint.handle_query_timeout(h, now);
}
}
async fn push_updates(&mut self, now: StdInstant) {
let mut removed_services: Vec<ServiceHandle> = Vec::new();
{
let Self {
endpoint,
services,
queries,
query_handle_scratch,
..
} = self;
for (handle, ctx) in services.iter_mut() {
if ctx.withdrawing {
continue;
}
if ctx.doorbell.is_closed() {
removed_services.push(*handle);
continue;
}
while let Some(upd) = ctx.proto.poll() {
let final_upd = match upd {
ServiceUpdate::Renamed(ref renamed) => {
let rename_result =
endpoint.handle_service_renamed(*handle, renamed.new_name().clone());
if let Some(h) = ctx.proto.take_rename_goodbye_handoff() {
endpoint.enqueue_rename_withdrawal(h, now, rename_result.is_err());
}
match rename_result {
Ok(()) => upd,
Err(_) => {
hick_trace::warn!(
handle = ?handle,
new_name = %renamed.new_name(),
"auto-rename collided with another registered service; emitting Conflict"
);
ServiceUpdate::Conflict
}
}
}
_ => upd,
};
let is_terminal = final_upd.is_conflict() || final_upd.is_host_conflict();
deliver_service_update(ctx, final_upd);
if is_terminal {
removed_services.push(*handle);
break;
}
}
}
let mut terminated: Vec<QueryHandle> = Vec::new();
query_handle_scratch.clear();
query_handle_scratch.extend(queries.keys().copied());
for &h in query_handle_scratch.iter() {
if let Some(ctx) = queries.get(&h)
&& ctx.doorbell.is_closed()
{
terminated.push(h);
continue;
}
let last_seq = match queries.get(&h) {
Some(c) => c.last_seq,
None => continue,
};
let accepted = endpoint.query_accepted_count(h).unwrap_or(last_seq);
let expected = accepted.saturating_sub(last_seq);
if expected > 0
&& let Some(ctx) = queries.get_mut(&h)
{
let mut delivered: u64 = 0;
{
let mut mb = ctx.mailbox.lock().unwrap_or_else(|e| e.into_inner());
for ans in endpoint
.collected_answers(h)
.filter(|a| a.seq() >= last_seq)
.cloned()
{
mb.push_answer(ans);
delivered += 1;
}
mb.record_dropped(expected.saturating_sub(delivered));
}
ctx.last_seq = accepted;
if delivered > 0 {
let _ = ctx.doorbell.try_send(());
}
}
if let Some(terminal) = endpoint.poll_query(h)
&& let Some(ctx) = queries.get_mut(&h)
{
ctx
.mailbox
.lock()
.unwrap_or_else(|e| e.into_inner())
.set_terminal(terminal);
let _ = ctx.doorbell.try_send(());
terminated.push(h);
}
}
for h in terminated {
let _ = endpoint.cancel_query(h);
queries.remove(&h);
}
while let Some(_ev) = endpoint.poll() {}
}
for h in removed_services {
self.remove_service(h, now);
}
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip_all, fields(credits = MAX_SEND_CREDITS_PER_DRAIN))
)]
async fn drain_transmits(&mut self, now: StdInstant, scratch: &mut [u8]) -> bool {
#[cfg(feature = "stats")]
let stats = self.stats.clone();
let Self {
endpoint,
services,
queries,
recent_sends,
v4,
v6,
svc_handle_scratch,
query_handle_scratch,
..
} = self;
let mut credits_remaining = MAX_SEND_CREDITS_PER_DRAIN;
svc_handle_scratch.clear();
svc_handle_scratch.extend(services.keys().copied());
for &h in svc_handle_scratch.iter() {
if credits_remaining == 0 {
return true;
}
let live = services
.get(&h)
.map(|c| !c.doorbell.is_closed() && !c.withdrawing)
.unwrap_or(false);
if !live {
continue;
}
let mut hit_encode_error = false;
loop {
if credits_remaining == 0 {
return true;
}
let tx = match services.get_mut(&h) {
Some(ctx) => match ctx.proto.poll_transmit(now, scratch) {
Ok(Some(t)) => {
ctx.encode_failures = 0;
t
}
Ok(None) => {
ctx.encode_failures = 0;
break;
}
Err(_e) => {
ctx.encode_failures = ctx.encode_failures.saturating_add(1);
hick_trace::warn!(
handle = ?h,
error = ?_e,
scratch_size = scratch.len(),
consecutive_failures = ctx.encode_failures,
"Service::poll_transmit failed"
);
hit_encode_error = true;
break;
}
},
None => break,
};
let body_len = tx.size();
let used = send_via::<N>(
recent_sends,
v4,
v6,
tx.dst(),
&scratch[..body_len],
#[cfg(feature = "stats")]
&stats,
)
.await;
if let Some(ctx) = services.get_mut(&h) {
ctx.proto.note_transmit_result(StdInstant::now(), used > 0);
if used > 0 {
endpoint.note_service_advertised(
h,
ctx.proto.advertised_a_addrs(),
ctx.proto.advertised_aaaa_addrs(),
ctx.proto.advertises_instance(),
);
}
}
credits_remaining = credits_remaining.saturating_sub(used);
}
if hit_encode_error {
let escalate = services
.get(&h)
.map(|c| c.encode_failures >= MAX_CONSECUTIVE_ENCODE_ERRORS)
.unwrap_or(false);
if escalate {
hick_trace::warn!(
handle = ?h,
"Service exceeded MAX_CONSECUTIVE_ENCODE_ERRORS; emitting Conflict and withdrawing"
);
if let Some(ctx) = services.get_mut(&h) {
deliver_service_update(ctx, ServiceUpdate::Conflict);
ctx.withdrawing = true;
if let Some(handoff) = ctx.proto.take_rename_goodbye_handoff() {
endpoint.enqueue_rename_withdrawal(handoff, now, true);
}
let snap = ctx.proto.withdrawal_snapshot();
endpoint.begin_withdrawal(h, snap, now);
}
}
}
}
query_handle_scratch.clear();
query_handle_scratch.extend(queries.keys().copied());
let mut encode_retired: Vec<QueryHandle> = Vec::new();
let mut more_pending = false;
'query_loop: for &h in query_handle_scratch.iter() {
if credits_remaining == 0 {
more_pending = true;
break 'query_loop;
}
let live = queries
.get(&h)
.map(|c| !c.doorbell.is_closed())
.unwrap_or(false);
if !live {
continue;
}
while credits_remaining > 0 {
let tx = match endpoint.poll_query_transmit(h, now, scratch) {
Ok(Some(t)) => t,
Ok(None) => break,
Err(_e) => {
endpoint.retire_query(h);
if let Some(terminal) = endpoint.poll_query(h)
&& let Some(ctx) = queries.get(&h)
{
ctx
.mailbox
.lock()
.unwrap_or_else(|e| e.into_inner())
.set_terminal(terminal);
let _ = ctx.doorbell.try_send(());
}
encode_retired.push(h);
hick_trace::warn!(
handle = ?h,
error = ?_e,
scratch_size = scratch.len(),
"Endpoint::poll_query_transmit failed; retiring proto query (terminal pushed to Query::next)"
);
break;
}
};
let body_len = tx.size();
let used = send_via::<N>(
recent_sends,
v4,
v6,
tx.dst(),
&scratch[..body_len],
#[cfg(feature = "stats")]
&stats,
)
.await;
endpoint.note_query_transmit_result(h, StdInstant::now(), used > 0);
credits_remaining = credits_remaining.saturating_sub(used);
}
}
for h in encode_retired {
let _ = endpoint.cancel_query(h);
queries.remove(&h);
}
more_pending
}
fn sweep_closed_handles(&mut self, now: StdInstant) {
let dead_svc: Vec<ServiceHandle> = self
.services
.iter()
.filter(|(_, ctx)| ctx.doorbell.is_closed() && !ctx.withdrawing)
.map(|(h, _)| *h)
.collect();
for h in dead_svc {
self.remove_service(h, now);
}
let dead_q: Vec<QueryHandle> = self
.queries
.iter()
.filter(|(_, ctx)| ctx.doorbell.is_closed())
.map(|(h, _)| *h)
.collect();
for h in dead_q {
self.queries.remove(&h);
let _ = self.endpoint.cancel_query(h);
}
}
fn remove_service(&mut self, handle: ServiceHandle, now: StdInstant) {
self.begin_service_withdrawal(handle, now);
}
fn begin_service_withdrawal(&mut self, handle: ServiceHandle, now: StdInstant) {
let (snap, handoff) = match self.services.get_mut(&handle) {
Some(ctx) => {
ctx.withdrawing = true;
let handoff = ctx.proto.take_rename_goodbye_handoff();
(ctx.proto.withdrawal_snapshot(), handoff)
}
None => return,
};
if let Some(handoff) = handoff {
self.endpoint.enqueue_rename_withdrawal(handoff, now, true);
}
self.endpoint.begin_withdrawal(handle, snap, now);
}
async fn drain_withdrawals(&mut self, now: StdInstant, scratch: &mut [u8]) {
#[cfg(feature = "stats")]
let stats = self.stats.clone();
let Self {
endpoint,
recent_sends,
v4,
v6,
..
} = self;
while let Some((dst, len, token)) = endpoint.poll_withdrawal_transmit(now, scratch) {
debug_assert!(
matches!(dst, SocketAddr::V4(v4a) if v4a.ip().is_multicast() && v4a.port() == 5353),
"withdrawal dst must be the IPv4 multicast marker"
);
let _ = dst;
let (v4_out, v6_out) = send_withdrawal_via::<N>(
recent_sends,
v4,
v6,
&scratch[..len],
#[cfg(feature = "stats")]
&stats,
)
.await;
#[cfg(feature = "stats")]
if matches!(v4_out, WithdrawalSend::Sent) || matches!(v6_out, WithdrawalSend::Sent) {
stats.goodbyes_tx(1);
}
endpoint.note_withdrawal_result(token, now, v4_out, v6_out);
}
self.completed_withdrawals.clear();
self
.endpoint
.drain_completed_withdrawals(now, &mut self.completed_withdrawals);
while let Some(handle) = self.completed_withdrawals.pop() {
self.services.remove(&handle);
}
}
}
const MAX_SEND_CREDITS_PER_DRAIN: usize = 64;
const MAX_CONSECUTIVE_ENCODE_ERRORS: u8 = 3;
fn deliver_service_update(ctx: &mut ServiceCtx, upd: ServiceUpdate) {
ctx
.mailbox
.lock()
.unwrap_or_else(|e| e.into_inner())
.push_update(upd);
let _ = ctx.doorbell.try_send(());
}
fn min_opt(prev: Option<StdInstant>, t: StdInstant) -> StdInstant {
match prev {
Some(b) if b <= t => b,
_ => t,
}
}
const MDNS_V4_DST: SocketAddr = SocketAddr::V4(std::net::SocketAddrV4::new(
std::net::Ipv4Addr::new(224, 0, 0, 251),
5353,
));
const MDNS_V6_DST: SocketAddr = SocketAddr::V6(std::net::SocketAddrV6::new(
std::net::Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 0, 0x00fb),
5353,
0,
0,
));
const SELF_SEND_TTL: Duration = Duration::from_secs(2);
const MAX_SELF_SEND_ENTRIES: usize = 65536;
fn fnv1a(data: &[u8]) -> u64 {
const OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
const PRIME: u64 = 0x0000_0100_0000_01b3;
let mut h = OFFSET;
for &b in data {
h ^= b as u64;
h = h.wrapping_mul(PRIME);
}
h
}
fn record_self_send(tracker: &mut Vec<(u64, SystemTime)>, body: &[u8], sent: SystemTime) {
tracker.retain(|(_, t)| match sent.duration_since(*t) {
Ok(age) => age <= SELF_SEND_TTL,
Err(_) => true,
});
if tracker.len() < MAX_SELF_SEND_ENTRIES {
tracker.push((fnv1a(body), sent));
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum MatchMode {
Ordered,
Degraded,
}
fn take_self_send(
tracker: &mut Vec<(u64, SystemTime)>,
body: &[u8],
reference: SystemTime,
mode: MatchMode,
) -> bool {
let needle = fnv1a(body);
match tracker
.iter()
.position(|(h, sent)| *h == needle && reference_matches(reference, *sent, mode))
{
Some(pos) => {
tracker.remove(pos);
true
}
None => false,
}
}
fn reference_matches(reference: SystemTime, sent: SystemTime, mode: MatchMode) -> bool {
match reference.duration_since(sent) {
Ok(ahead) => ahead <= SELF_SEND_TTL,
Err(behind) => mode == MatchMode::Ordered && behind.duration() <= hick_udp::RX_TIMESTAMP_GRAIN,
}
}
fn is_on_link(hop_limit: Option<u8>) -> bool {
hop_limit.is_none_or(|hl| hl == 255)
}
fn packet_is_response(data: &[u8]) -> bool {
data.get(2).is_some_and(|b| b & 0x80 != 0)
}
fn src_on_local_link(
local_subnets: &[(IpAddr, u8)],
bound_iface: u32,
recv_iface: u32,
src: IpAddr,
) -> bool {
let (is_loopback, is_link_local) = match src {
IpAddr::V4(v4) => (v4.is_loopback(), v4.is_link_local()),
IpAddr::V6(v6) => (v6.is_loopback(), (v6.segments()[0] & 0xffc0) == 0xfe80),
};
if is_loopback {
return true;
}
if is_link_local {
return recv_iface == 0 || recv_iface == bound_iface;
}
local_subnets
.iter()
.any(|&(net, prefix)| addr_in_subnet(net, prefix, src))
}
fn addr_in_subnet(net: IpAddr, prefix: u8, addr: IpAddr) -> bool {
match (net, addr) {
(IpAddr::V4(n), IpAddr::V4(a)) => {
let p = prefix.min(32);
if p == 0 {
return true;
}
let mask: u32 = u32::MAX.checked_shl(32 - u32::from(p)).unwrap_or(0);
(u32::from(n) & mask) == (u32::from(a) & mask)
}
(IpAddr::V6(n), IpAddr::V6(a)) => {
let p = prefix.min(128);
if p == 0 {
return true;
}
let mask: u128 = u128::MAX.checked_shl(128 - u32::from(p)).unwrap_or(0);
(u128::from(n) & mask) == (u128::from(a) & mask)
}
_ => false,
}
}
fn collect_local_subnets(iface_index: u32) -> Vec<(IpAddr, u8)> {
let mut out: Vec<(IpAddr, u8)> = Vec::new();
if iface_index == 0 {
return out;
}
if let Ok(Some(i)) = getifs::interface_by_index(iface_index) {
if let Ok(v4s) = i.ipv4_addrs() {
for n in v4s.iter() {
out.push((IpAddr::V4(n.addr()), n.prefix_len()));
}
}
if let Ok(v6s) = i.ipv6_addrs() {
for n in v6s.iter() {
out.push((IpAddr::V6(n.addr()), n.prefix_len()));
}
}
}
out
}
async fn send_via<N: Net>(
tracker: &mut Vec<(u64, SystemTime)>,
v4: &Option<Arc<N::UdpSocket>>,
v6: &Option<Arc<N::UdpSocket>>,
dst: SocketAddr,
body: &[u8],
#[cfg(feature = "stats")] stats: &std::sync::Arc<hick_trace::stats::Stats>,
) -> usize {
let is_mdns_multicast = matches!(dst, SocketAddr::V4(v4a) if v4a.ip().is_multicast() && v4a.port() == 5353)
|| matches!(dst, SocketAddr::V6(v6a) if v6a.ip().is_multicast() && v6a.port() == 5353);
let mut credits = 0usize;
if is_mdns_multicast {
if let Some(s) = v4 {
let (res, send_wall) = send_to_at::<N>(s, body, MDNS_V4_DST).await;
match res {
Ok(_) => {
hick_trace::trace!(dst = %MDNS_V4_DST, len = body.len(), "send_to v4");
record_self_send(tracker, body, send_wall);
#[cfg(feature = "stats")]
{
stats.packets_tx(1);
stats.bytes_tx(body.len() as u64);
}
credits += 1;
}
Err(_e) => {
hick_trace::debug!(error = %_e, dst = %MDNS_V4_DST, "send_to v4 failed");
#[cfg(feature = "stats")]
stats.send_errors(1);
}
}
}
if let Some(s) = v6 {
let (res, send_wall) = send_to_at::<N>(s, body, MDNS_V6_DST).await;
match res {
Ok(_) => {
hick_trace::trace!(dst = %MDNS_V6_DST, len = body.len(), "send_to v6");
record_self_send(tracker, body, send_wall);
#[cfg(feature = "stats")]
{
stats.packets_tx(1);
stats.bytes_tx(body.len() as u64);
}
credits += 1;
}
Err(_e) => {
hick_trace::debug!(error = %_e, dst = %MDNS_V6_DST, "send_to v6 failed");
#[cfg(feature = "stats")]
stats.send_errors(1);
}
}
}
return credits;
}
let sock = match dst {
SocketAddr::V4(_) => v4.as_ref(),
SocketAddr::V6(_) => v6.as_ref(),
};
if let Some(s) = sock {
let (res, send_wall) = send_to_at::<N>(s, body, dst).await;
match res {
Ok(_) => {
hick_trace::trace!(dst = %dst, len = body.len(), "send_to");
record_self_send(tracker, body, send_wall);
#[cfg(feature = "stats")]
{
stats.packets_tx(1);
stats.bytes_tx(body.len() as u64);
}
credits += 1;
}
Err(_e) => {
hick_trace::debug!(error = %_e, dst = %dst, "send_to failed");
#[cfg(feature = "stats")]
stats.send_errors(1);
}
}
}
credits
}
async fn send_withdrawal_via<N: Net>(
tracker: &mut Vec<(u64, SystemTime)>,
v4: &Option<Arc<N::UdpSocket>>,
v6: &Option<Arc<N::UdpSocket>>,
body: &[u8],
#[cfg(feature = "stats")] stats: &std::sync::Arc<hick_trace::stats::Stats>,
) -> (WithdrawalSend, WithdrawalSend) {
let mut v4_out = WithdrawalSend::WriteOff;
let mut v6_out = WithdrawalSend::WriteOff;
if let Some(s) = v4 {
let (res, send_wall) = send_to_at::<N>(s, body, MDNS_V4_DST).await;
v4_out = present_socket_send_outcome(&res);
match res {
Ok(_) => {
hick_trace::trace!(dst = %MDNS_V4_DST, len = body.len(), "withdrawal send_to v4");
record_self_send(tracker, body, send_wall);
#[cfg(feature = "stats")]
{
stats.packets_tx(1);
stats.bytes_tx(body.len() as u64);
}
}
Err(_e) => {
hick_trace::debug!(error = %_e, dst = %MDNS_V4_DST, "withdrawal send_to v4 failed");
#[cfg(feature = "stats")]
stats.send_errors(1);
}
}
}
if let Some(s) = v6 {
let (res, send_wall) = send_to_at::<N>(s, body, MDNS_V6_DST).await;
v6_out = present_socket_send_outcome(&res);
match res {
Ok(_) => {
hick_trace::trace!(dst = %MDNS_V6_DST, len = body.len(), "withdrawal send_to v6");
record_self_send(tracker, body, send_wall);
#[cfg(feature = "stats")]
{
stats.packets_tx(1);
stats.bytes_tx(body.len() as u64);
}
}
Err(_e) => {
hick_trace::debug!(error = %_e, dst = %MDNS_V6_DST, "withdrawal send_to v6 failed");
#[cfg(feature = "stats")]
stats.send_errors(1);
}
}
}
(v4_out, v6_out)
}
fn present_socket_send_outcome<T>(res: &std::io::Result<T>) -> WithdrawalSend {
match res {
Ok(_) => WithdrawalSend::Sent,
Err(_) => WithdrawalSend::Retry,
}
}
async fn send_to_at<N: Net>(
sock: &N::UdpSocket,
buf: &[u8],
dst: SocketAddr,
) -> (std::io::Result<usize>, SystemTime) {
let mut stamp = SystemTime::now();
let res = core::future::poll_fn(|cx| {
stamp = SystemTime::now();
sock.poll_send_to(cx, buf, dst)
})
.await;
(res, stamp)
}
pub(crate) fn spawn<N: Net>(
opts: ServerOptions,
sockets: BoundSockets<N>,
cmd_rx: async_channel::Receiver<Command>,
#[cfg(feature = "stats")] stats_out: &mut Option<std::sync::Arc<hick_trace::stats::Stats>>,
) {
let max_send = opts.max_payload_size();
let max_recv = opts.max_recv_packet_size();
let state = DriverState::<N>::new(&opts, sockets);
#[cfg(feature = "stats")]
{
*stats_out = Some(state.stats.clone());
}
<N::Runtime as RuntimeLite>::spawn_detach(driver_task::<N>(state, cmd_rx, max_send, max_recv));
}
#[cfg_attr(feature = "tracing", tracing::instrument(level = "trace", skip_all))]
async fn driver_task<N: Net>(
mut state: DriverState<N>,
cmd_rx: async_channel::Receiver<Command>,
max_send: usize,
max_recv: usize,
) {
let mut scratch: Vec<u8> = vec![0u8; max_send.max(512)];
let (packet_tx, packet_rx) = async_channel::bounded::<Packet>(64);
let (_shutdown_tx, shutdown_rx) = async_channel::bounded::<()>(1);
if let Some(sock) = state.v4.clone() {
let tx = packet_tx.clone();
let sd = shutdown_rx.clone();
#[cfg(feature = "stats")]
let stats = state.stats.clone();
<N::Runtime as RuntimeLite>::spawn_detach(recv_loop::<N>(
sock,
tx,
sd,
true,
max_recv,
#[cfg(feature = "stats")]
stats,
));
}
if let Some(sock) = state.v6.clone() {
let tx = packet_tx.clone();
let sd = shutdown_rx.clone();
#[cfg(feature = "stats")]
let stats = state.stats.clone();
<N::Runtime as RuntimeLite>::spawn_detach(recv_loop::<N>(
sock,
tx,
sd,
false,
max_recv,
#[cfg(feature = "stats")]
stats,
));
}
drop(packet_tx);
drop(shutdown_rx);
loop {
const PACKET_PUMP_BUDGET: usize = 64;
for _ in 0..PACKET_PUMP_BUDGET {
match packet_rx.try_recv() {
Ok(pkt) => state.handle_packet(pkt),
Err(_) => break,
}
}
let now = StdInstant::now();
state.sweep_closed_handles(now);
state.fire_timeouts(now);
let more_transmits_pending = state.drain_transmits(now, &mut scratch).await;
state.push_updates(now).await;
state.drain_withdrawals(now, &mut scratch).await;
if more_transmits_pending {
const COMMAND_DRAIN_BUDGET: usize = 8;
let mut cmd_closed = false;
for _ in 0..COMMAND_DRAIN_BUDGET {
match cmd_rx.try_recv() {
Ok(cmd) => state.handle_command(cmd, StdInstant::now()),
Err(async_channel::TryRecvError::Empty) => break,
Err(async_channel::TryRecvError::Closed) => {
cmd_closed = true;
break;
}
}
}
if cmd_closed {
break;
}
continue;
}
let deadline = state.next_deadline();
let cmd_fut = cmd_rx.recv().fuse();
let pkt_fut = packet_rx.recv().fuse();
pin_mut!(cmd_fut, pkt_fut);
let mut closed = false;
if let Some(at) = deadline {
let dur = at.saturating_duration_since(now);
let sleep = <N::Runtime as RuntimeLite>::sleep(dur).fuse();
pin_mut!(sleep);
select_biased! {
c = cmd_fut => match c {
Ok(cmd) => state.handle_command(cmd, StdInstant::now()),
Err(_) => closed = true,
},
p = pkt_fut => match p {
Ok(pkt) => state.handle_packet(pkt),
Err(_) => closed = true,
},
_ = sleep => { }
}
} else {
select_biased! {
c = cmd_fut => match c {
Ok(cmd) => state.handle_command(cmd, StdInstant::now()),
Err(_) => closed = true,
},
p = pkt_fut => match p {
Ok(pkt) => state.handle_packet(pkt),
Err(_) => closed = true,
},
}
}
if closed {
break;
}
}
let shutdown_now = StdInstant::now();
let live_services: Vec<ServiceHandle> = state.services.keys().copied().collect();
for h in live_services {
state.remove_service(h, shutdown_now);
}
let shutdown_deadline = StdInstant::now() + Duration::from_secs(10);
loop {
let now = StdInstant::now();
state.drain_withdrawals(now, &mut scratch).await;
let Some(next) = state.next_withdrawal_deadline() else {
break;
};
if now >= shutdown_deadline {
hick_trace::debug!("shutdown withdrawal flush hit its wall-clock backstop; exiting");
break;
}
let dur = next
.saturating_duration_since(now)
.min(shutdown_deadline.saturating_duration_since(now));
if dur > Duration::ZERO {
<N::Runtime as RuntimeLite>::sleep(dur).await;
}
}
}
#[cfg(feature = "stats")]
#[inline]
fn count_consumed_oversized(stats: &hick_trace::stats::Stats, buf_len: usize) {
stats.packets_rx(1);
stats.bytes_rx(buf_len as u64);
stats.packets_dropped(1);
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip_all, fields(via_v4))
)]
async fn recv_loop<N: Net>(
sock: Arc<N::UdpSocket>,
tx: async_channel::Sender<Packet>,
shutdown: async_channel::Receiver<()>,
via_v4: bool,
max_recv: usize,
#[cfg(feature = "stats")] stats: std::sync::Arc<hick_trace::stats::Stats>,
) {
let mut buf = vec![0u8; max_recv.max(1500)];
loop {
#[cfg(unix)]
{
let ready = {
let peek_fut = sock.peek_from(&mut buf).fuse();
let shutdown_fut = shutdown.recv().fuse();
pin_mut!(peek_fut, shutdown_fut);
select_biased! {
_ = shutdown_fut => return,
r = peek_fut => r,
}
};
if let Err(_e) = ready {
hick_trace::debug!(error = %_e, via_v4, "peek_from failed");
return;
}
use std::os::fd::AsRawFd;
let fd = sock.as_raw_fd();
match hick_udp::recv_with_meta(fd, &mut buf, via_v4) {
Ok(meta) => {
let n = meta.len();
hick_trace::trace!(src = %meta.peer(), len = n, via_v4, "recv datagram");
let data = buf.get(..n).unwrap_or(&buf).to_vec();
let pkt = Packet {
src: meta.peer(),
data,
local_ip: meta.local_ip(),
interface_index: meta.interface_index(),
kernel_rx_time: meta.rx_time(),
read_time: SystemTime::now(),
hop_limit: meta.hop_limit(),
};
if tx.send(pkt).await.is_err() {
return;
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
continue;
}
Err(e) if e.kind() == std::io::ErrorKind::InvalidData => {
hick_trace::debug!(error = %e, via_v4, "dropping unusable datagram");
#[cfg(feature = "stats")]
count_consumed_oversized(&stats, buf.len());
continue;
}
Err(_e) => {
hick_trace::debug!(error = %_e, via_v4, "recv_with_meta failed");
return;
}
}
}
#[cfg(windows)]
{
use std::os::windows::io::AsRawSocket;
const WSAEMSGSIZE: i32 = 10040;
let ready = {
let peek_fut = sock.peek_from(&mut buf).fuse();
let shutdown_fut = shutdown.recv().fuse();
pin_mut!(peek_fut, shutdown_fut);
select_biased! {
_ = shutdown_fut => return,
r = peek_fut => r,
}
};
match ready {
Ok(_) => {}
Err(ref e) if e.raw_os_error() == Some(WSAEMSGSIZE) => {}
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => continue,
Err(_e) => {
hick_trace::debug!(error = %_e, via_v4, "peek_from failed");
return;
}
}
let raw = sock.as_raw_socket();
match hick_udp::recv_with_meta(raw, &mut buf, via_v4) {
Ok(meta) => {
let n = meta.len();
hick_trace::trace!(src = %meta.peer(), len = n, via_v4, "recv datagram");
let data = buf.get(..n).unwrap_or(&buf).to_vec();
let pkt = Packet {
src: meta.peer(),
data,
local_ip: meta.local_ip(),
interface_index: meta.interface_index(),
kernel_rx_time: meta.rx_time(),
read_time: SystemTime::now(),
hop_limit: meta.hop_limit(),
};
if tx.send(pkt).await.is_err() {
return;
}
}
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
continue;
}
Err(ref e) if e.raw_os_error() == Some(WSAEMSGSIZE) => {
hick_trace::debug!(via_v4, "dropping oversized datagram (WSAEMSGSIZE)");
#[cfg(feature = "stats")]
count_consumed_oversized(&stats, buf.len());
continue;
}
Err(_e) => {
hick_trace::debug!(error = %_e, via_v4, "recv_with_meta (windows) failed");
return;
}
}
}
#[cfg(all(not(unix), not(windows)))]
{
let recv_result = {
let recv_fut = sock.recv_from(&mut buf).fuse();
let shutdown_fut = shutdown.recv().fuse();
pin_mut!(recv_fut, shutdown_fut);
select_biased! {
_ = shutdown_fut => return,
r = recv_fut => r,
}
};
match recv_result {
Ok((n, src)) => {
hick_trace::trace!(src = %src, len = n, via_v4, "recv datagram");
let data = buf.get(..n).unwrap_or(&buf).to_vec();
let local_ip = if via_v4 {
IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)
} else {
IpAddr::V6(std::net::Ipv6Addr::UNSPECIFIED)
};
let pkt = Packet {
src,
data,
local_ip,
interface_index: 0,
kernel_rx_time: None,
read_time: SystemTime::now(),
hop_limit: None,
};
if tx.send(pkt).await.is_err() {
return;
}
}
Err(_e) => {
hick_trace::debug!(error = %_e, via_v4, "recv_from failed");
return;
}
}
}
}
}