use std::{
io,
net::{Ipv4Addr, SocketAddr, SocketAddrV4},
num::Wrapping,
sync::Arc,
};
use thiserror::Error;
use tokio::{
net::UdpSocket,
sync::{mpsc, watch},
time::{self, Duration, Instant, Interval, MissedTickBehavior, interval_at},
};
use tracing::{debug, trace};
pub mod packet;
pub use packet::{ExtensionField, Fec2V1Value as DaemonInfo, Packet};
pub mod socket_ext;
use crate::daemon::{
async_ring_buffer::{self, BufferClosedError, SendError},
event::{self, NtpData},
io::{
ClockDisruptionEvent, ControlRequest,
ntp::socket_ext::SocketExt,
tsc::{read_timestamp_counter_begin, read_timestamp_counter_end},
},
selected_clock::SelectedClockSource,
time::TscCount,
};
use packet::Timestamp;
pub const UNSPECIFIED_SOCKET_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0);
#[derive(Debug, Error)]
pub enum NtpIoError {
#[error("Sample IO failure.")]
SampleIo(#[source] io::Error),
#[error("Socket bind failure")]
Bind(#[source] io::Error),
#[error("Operation timed out.")]
Timeout(#[from] time::error::Elapsed),
#[error("TSC order failure. tsc_pre: {pre}. tsc_post: {post}")]
TscOrder { pre: u64, post: u64 },
#[error("IO failure on socket clear")]
SocketClear(#[source] io::Error),
}
#[derive(Debug)]
pub struct Ntp {
event_sender: async_ring_buffer::Sender<event::Ntp>,
ctrl_receiver: mpsc::Receiver<ControlRequest>,
clock_disruption_receiver: watch::Receiver<ClockDisruptionEvent>,
interval: Interval,
normal_interval: Duration,
buffer: Vec<u8>,
selected_clock: Arc<SelectedClockSource>,
transmit_counter: Wrapping<u64>,
extensions: Vec<ExtensionField>,
socket_address: SocketAddr,
timeout: Duration,
burst: Option<BurstState>,
}
#[bon::bon]
impl Ntp {
#[builder]
pub fn new(
event_sender: async_ring_buffer::Sender<event::Ntp>,
ctrl_receiver: mpsc::Receiver<ControlRequest>,
clock_disruption_receiver: watch::Receiver<ClockDisruptionEvent>,
selected_clock: Arc<SelectedClockSource>,
extensions: Vec<ExtensionField>,
socket_address: SocketAddr,
timeout: Duration,
interval: Duration,
burst_config: Option<BurstConfig>,
) -> Self {
let buffer_size = Packet::MIN_SIZE
+ extensions
.iter()
.map(|ext| ext.length() as usize)
.sum::<usize>();
let initial_interval = burst_config.as_ref().map_or(interval, |bc| bc.interval);
let mut ntp_interval = tokio::time::interval(initial_interval);
ntp_interval.set_missed_tick_behavior(MissedTickBehavior::Delay);
let burst = burst_config.map(|config| BurstState {
config,
mode: Mode::burst(),
});
let mut retval = Ntp {
event_sender,
ctrl_receiver,
clock_disruption_receiver,
interval: ntp_interval,
normal_interval: interval,
buffer: vec![0u8; buffer_size],
selected_clock,
transmit_counter: Wrapping(0),
extensions,
socket_address,
timeout,
burst,
};
retval.decorrelate_poll_timing();
retval
}
}
impl Ntp {
pub fn select_branches(&mut self) -> SelectBranches<'_> {
SelectBranches {
clock_disruption_receiver: &mut self.clock_disruption_receiver,
ctrl_receiver: &mut self.ctrl_receiver,
interval: &mut self.interval,
}
}
fn decorrelate_poll_timing(&mut self) {
let max_delay = if let Some(burst) = &self.burst
&& let Mode::Burst(_) = &burst.mode
{
burst.config.interval.as_nanos() as usize
} else {
self.normal_interval.as_nanos() as usize
};
let delay = rand::random_range(0..=max_delay) as u64;
let delay = Duration::from_nanos(delay);
self.interval.reset_after(delay);
}
pub fn socket_address(&self) -> SocketAddr {
self.socket_address
}
pub async fn sample(&mut self) -> Result<event::Ntp, NtpIoError> {
let counter = self.transmit_counter.0;
self.transmit_counter += 1;
let (source, stratum) = self.selected_clock.get_with_client_stratum();
let packet = Packet::builder()
.transmit_timestamp(Timestamp::new(counter))
.stratum(stratum.into())
.reference_id(source.into())
.extensions(self.extensions.clone())
.build();
packet.emit_bytes(&mut self.buffer);
let socket = UdpSocket::bind(UNSPECIFIED_SOCKET_ADDRESS)
.await
.map_err(NtpIoError::Bind)?;
let ntp_event = sample_packet(
&socket,
self.socket_address,
&mut self.buffer,
self.timeout,
counter,
)
.await?;
debug!(?ntp_event, "Received packet.");
Ok(ntp_event)
}
pub fn handle_disruption(&mut self) -> bool {
let Self {
clock_disruption_receiver,
event_sender,
ctrl_receiver: _,
interval,
normal_interval: _,
buffer: _,
selected_clock: _,
transmit_counter: _,
extensions: _,
socket_address: _,
timeout: _,
burst,
} = self;
let val = clock_disruption_receiver.borrow_and_update().clone();
if val.disruption_marker.is_some() {
event_sender.handle_disruption();
if let Some(burst_state) = burst {
burst_state.mode = Mode::burst();
*interval = tokio::time::interval(burst_state.config.interval);
interval.set_missed_tick_behavior(MissedTickBehavior::Delay);
interval.reset_immediately();
}
return true;
}
false
}
pub fn handle_burst_expiry(&mut self) -> bool {
let Some(burst_state) = &mut self.burst else {
return false;
};
if let Mode::Burst(start_time) = burst_state.mode
&& start_time.elapsed() >= burst_state.config.duration
{
burst_state.mode = Mode::Normal;
let normal = self.normal_interval;
self.interval = interval_at(Instant::now() + normal, normal);
self.interval
.set_missed_tick_behavior(MissedTickBehavior::Delay);
self.decorrelate_poll_timing();
return true;
}
false
}
pub fn send_event(&self, event: &event::Ntp) -> Result<(), BufferClosedError> {
match self.event_sender.send(event.clone()) {
Ok(()) => {
debug!(?event, "Successfully sent IO event.");
Ok(())
}
Err(SendError::Disrupted(_)) => {
debug!("Trying to send when there was a disruption event.");
Ok(())
}
Err(SendError::BufferClosed(e)) => Err(e),
}
}
}
#[derive(Debug, Clone)]
pub struct BurstConfig {
pub duration: Duration,
pub interval: Duration,
}
#[derive(Debug)]
enum Mode {
Normal,
Burst(Instant),
}
impl Mode {
fn burst() -> Mode {
Mode::Burst(Instant::now())
}
}
#[derive(Debug)]
pub struct SelectBranches<'a> {
pub clock_disruption_receiver: &'a mut watch::Receiver<ClockDisruptionEvent>,
pub ctrl_receiver: &'a mut mpsc::Receiver<ControlRequest>,
pub interval: &'a mut Interval,
}
#[derive(Debug)]
struct BurstState {
config: BurstConfig,
mode: Mode,
}
async fn sample_packet(
socket: &UdpSocket,
addr: SocketAddr,
send_recv_buffer: &mut [u8],
timeout: std::time::Duration,
expected_counter: u64,
) -> Result<event::Ntp, NtpIoError> {
socket.clear().map_err(NtpIoError::SocketClear)?;
let fut = tokio::time::timeout(
timeout,
inner_timeout(socket, addr, send_recv_buffer, expected_counter),
);
let (send_timestamp, ntp_data, received_timestamp) =
fut.await?.map_err(NtpIoError::SampleIo)?;
#[cfg(not(test))]
let system_clock_reading = crate::daemon::event::SystemClockMeasurement::now();
#[allow(clippy::cast_possible_wrap)]
let tsc_pre = TscCount::new(send_timestamp as i64);
#[allow(clippy::cast_possible_wrap)]
let tsc_post = TscCount::new(received_timestamp as i64);
let builder = event::Ntp::builder()
.counter_pre(tsc_pre)
.counter_post(tsc_post)
.ntp_data(ntp_data);
let ntp_event = {
#[cfg(not(test))]
{
builder.system_clock(system_clock_reading).build()
}
#[cfg(test)]
{
builder.build()
}
};
let ntp_event = ntp_event.ok_or(NtpIoError::TscOrder {
pre: send_timestamp,
post: received_timestamp,
})?;
Ok(ntp_event)
}
async fn inner_timeout(
socket: &UdpSocket,
addr: SocketAddr,
send_recv_buffer: &mut [u8],
expected_counter: u64,
) -> Result<(u64, NtpData, u64), io::Error> {
let send_timestamp = read_timestamp_counter_begin();
socket.send_to(send_recv_buffer, addr).await?;
loop {
let (len, recv_addr) = socket.recv_from(send_recv_buffer).await?;
let received_timestamp = read_timestamp_counter_end();
if recv_addr != addr {
continue;
}
let Ok((_, ntp_packet)) = Packet::parse_from_bytes(&send_recv_buffer[..len])
.inspect_err(|e| trace!(parse_error = ?e.to_string(), "Parsing error."))
else {
continue;
};
if ntp_packet.origin_timestamp.get() != expected_counter {
trace!(
error = ?InnerSamplePacketError::OriginMismatch {
expected: expected_counter,
received: ntp_packet.origin_timestamp.get(),
},
"Origin timestamp mismatch."
);
continue;
}
let Ok(ntp_data) = NtpData::try_from(ntp_packet)
.inspect_err(|e| trace!(error = ?InnerSamplePacketError::PacketParsing(e.to_string()), "Failed to parse NTP data.")) else {
continue;
};
return Ok((send_timestamp, ntp_data, received_timestamp));
}
}
#[derive(Debug, thiserror::Error)]
enum InnerSamplePacketError {
#[error("IO failure.")]
Io(#[from] io::Error),
#[error("Failed to parse NTP packet.")]
PacketParsing(String),
#[error("Mismatched origin. Expected {expected}, got {received}")]
OriginMismatch { expected: u64, received: u64 },
}