Skip to main content

laminar_connectors/nats/
source.rs

1//! NATS source: `JetStream` pull consumer with bounded asynchronous acknowledgement, or core
2//! subscribe (at-most-once). Messages are acknowledged only after successful deserialization;
3//! the source remains explicitly ephemeral and does not couple broker acks to checkpoints.
4
5use std::sync::atomic::{AtomicU32, Ordering};
6use std::sync::Arc;
7use std::time::{Duration, Instant};
8
9use arrow_schema::SchemaRef;
10use async_nats::jetstream::{self, consumer::pull};
11use async_trait::async_trait;
12use bytes::Bytes;
13use crossfire::{mpsc, AsyncRx, MAsyncTx, TryRecvError};
14use futures_util::stream::FuturesUnordered;
15use futures_util::StreamExt;
16use tokio::sync::{mpsc as tokio_mpsc, watch, Notify};
17use tokio::task::JoinHandle;
18use tracing::{debug, warn};
19
20use super::config::{build_connect_options, AckPolicy, DeliverPolicy, Mode, NatsSourceConfig};
21use super::metrics::NatsSourceMetrics;
22use super::setup::{
23    classify_connect_error, classify_create_consumer_error, classify_get_stream_error,
24    classify_subscribe_error, track_connection_tasks,
25};
26use crate::checkpoint::SourceCheckpoint;
27use crate::config::ConnectorConfig;
28use crate::connector::{
29    ConnectorTaskGuard, ConnectorTaskOwner, ConnectorTaskTracker, SourceBatch, SourceConnector,
30    SourceConsistency, SourceContract, SourcePosition, SourceStart, SourceTopology,
31};
32use crate::error::ConnectorError;
33use crate::serde::{self, RecordDeserializer};
34
35const ACK_IO_TIMEOUT: Duration = Duration::from_secs(5);
36const CLOSE_DRAIN_TIMEOUT: Duration = Duration::from_secs(5);
37const MAX_ACK_CONCURRENCY: usize = 64;
38const MAX_ACK_BACKLOG: usize = 16_384;
39
40/// `ack` is `Some` only on the `JetStream` path.
41struct Incoming {
42    payload: Bytes,
43    ack: Option<jetstream::Message>,
44}
45
46struct AckRuntime {
47    tx: Option<tokio_mpsc::Sender<jetstream::Message>>,
48    shutdown: watch::Sender<bool>,
49    task: TrackedTask,
50}
51
52struct Running {
53    deserializer: Box<dyn RecordDeserializer>,
54    rx: Option<AsyncRx<mpsc::Array<Incoming>>>,
55    shutdown: watch::Sender<bool>,
56    reader: TrackedTask,
57    ack_runtime: Option<AckRuntime>,
58}
59
60struct TrackedTask {
61    handle: Option<JoinHandle<()>>,
62    reaper_guard: Option<ConnectorTaskGuard>,
63    name: &'static str,
64}
65
66enum TaskWait {
67    Completed(Result<(), tokio::task::JoinError>),
68    TimedOut,
69}
70
71impl TrackedTask {
72    fn spawn(
73        owner: &ConnectorTaskOwner,
74        name: &'static str,
75        future: impl std::future::Future<Output = ()> + Send + 'static,
76    ) -> Result<Self, ConnectorError> {
77        let task_guard = owner.track().ok_or_else(|| {
78            ConnectorError::Internal("NATS source task generation is already retired".into())
79        })?;
80        let reaper_guard = owner.track().ok_or_else(|| {
81            ConnectorError::Internal("NATS source task generation is already retired".into())
82        })?;
83        let handle = tokio::spawn(async move {
84            let _task_guard = task_guard;
85            future.await;
86        });
87        Ok(Self {
88            handle: Some(handle),
89            reaper_guard: Some(reaper_guard),
90            name,
91        })
92    }
93
94    async fn wait_until(&mut self, deadline: tokio::time::Instant) -> TaskWait {
95        let Some(handle) = self.handle.as_mut() else {
96            return TaskWait::Completed(Ok(()));
97        };
98        match tokio::time::timeout_at(deadline, handle).await {
99            Ok(result) => {
100                self.handle.take();
101                self.reaper_guard.take();
102                TaskWait::Completed(result)
103            }
104            Err(_) => TaskWait::TimedOut,
105        }
106    }
107
108    fn retire(&mut self) {
109        let Some(handle) = self.handle.take() else {
110            self.reaper_guard.take();
111            return;
112        };
113        let reaper_guard = self.reaper_guard.take();
114        let Ok(runtime) = tokio::runtime::Handle::try_current() else {
115            // Runtime destruction drops the task future and its task guard. That guard is the
116            // completion proof; a join task is useful for cleanup but is not itself the proof.
117            drop(handle);
118            drop(reaper_guard);
119            return;
120        };
121        let name = self.name;
122        drop(runtime.spawn(async move {
123            let _reaper_guard = reaper_guard;
124            if let Err(error) = handle.await {
125                debug!(task = name, %error, "retired NATS source task reaped");
126            }
127        }));
128    }
129}
130
131impl Drop for TrackedTask {
132    fn drop(&mut self) {
133        self.retire();
134    }
135}
136
137impl Running {
138    fn request_shutdown(&self) {
139        self.shutdown.send_replace(true);
140    }
141}
142
143impl Drop for Running {
144    fn drop(&mut self) {
145        // Drop is the final backstop for a cancelled startup/close future. Tasks retain their
146        // generation guards until they actually exit; the reapers retain the join handles.
147        self.request_shutdown();
148        self.rx.take();
149        if let Some(ack_runtime) = self.ack_runtime.as_mut() {
150            ack_runtime.request_shutdown();
151            ack_runtime.task.retire();
152        }
153        self.reader.retire();
154    }
155}
156
157impl AckRuntime {
158    fn spawn(
159        cfg: &NatsSourceConfig,
160        metrics: NatsSourceMetrics,
161        owner: &ConnectorTaskOwner,
162    ) -> Result<Self, ConnectorError> {
163        let (backlog, concurrency) = ack_runtime_limits(cfg);
164        let (tx, rx) = tokio_mpsc::channel(backlog);
165        let (shutdown, shutdown_rx) = watch::channel(false);
166        let task = TrackedTask::spawn(
167            owner,
168            "ack-worker",
169            run_ack_worker(rx, shutdown_rx, concurrency, metrics),
170        )?;
171        Ok(Self {
172            tx: Some(tx),
173            shutdown,
174            task,
175        })
176    }
177
178    fn request_shutdown(&mut self) {
179        self.shutdown.send_replace(true);
180        self.tx.take();
181    }
182}
183
184impl Drop for AckRuntime {
185    fn drop(&mut self) {
186        self.request_shutdown();
187    }
188}
189
190/// NATS source — core and `JetStream` modes.
191pub struct NatsSource {
192    schema: SchemaRef,
193    config: Option<NatsSourceConfig>,
194    data_ready: Arc<Notify>,
195    metrics: NatsSourceMetrics,
196    running: Option<Running>,
197    task_owner: ConnectorTaskOwner,
198    task_tracker: ConnectorTaskTracker,
199}
200
201impl NatsSource {
202    /// Metrics register on `registry` if provided.
203    #[must_use]
204    pub fn new(schema: SchemaRef, registry: Option<&prometheus::Registry>) -> Self {
205        let (task_owner, task_tracker) = ConnectorTaskOwner::new();
206        Self {
207            schema,
208            config: None,
209            data_ready: Arc::new(Notify::new()),
210            metrics: NatsSourceMetrics::new(registry),
211            running: None,
212            task_owner,
213            task_tracker,
214        }
215    }
216
217    /// Available after [`SourceConnector::start`].
218    #[must_use]
219    pub fn config(&self) -> Option<&NatsSourceConfig> {
220        self.config.as_ref()
221    }
222
223    /// Snapshot accessor for the prometheus-backed metrics struct.
224    #[must_use]
225    pub fn metrics_handle(&self) -> &NatsSourceMetrics {
226        &self.metrics
227    }
228
229    async fn open_jetstream(
230        &mut self,
231        cfg: &NatsSourceConfig,
232        deserializer: Box<dyn RecordDeserializer>,
233    ) -> Result<(), ConnectorError> {
234        let client = connect(cfg, &self.task_owner).await?;
235        let js = jetstream::new(client);
236
237        let stream_name = cfg
238            .stream
239            .as_deref()
240            .ok_or_else(|| err("stream name missing after validation"))?;
241        let consumer_name = cfg
242            .consumer
243            .as_deref()
244            .ok_or_else(|| err("consumer name missing after validation"))?;
245
246        let pull_cfg = build_pull_config(cfg, consumer_name)?;
247        let stream = js
248            .get_stream(stream_name)
249            .await
250            .map_err(|error| classify_get_stream_error(&error, stream_name))?;
251        let consumer = stream
252            .create_consumer(pull_cfg)
253            .await
254            .map_err(|error| classify_create_consumer_error(&error, consumer_name))?;
255
256        let (tx, rx) = mpsc::bounded_async::<Incoming>(cfg.fetch_batch * 2);
257        let (shutdown, shutdown_rx) = watch::channel(false);
258        let requires_ack = cfg.ack_policy == AckPolicy::Explicit;
259        let ack_runtime = if requires_ack {
260            Some(AckRuntime::spawn(
261                cfg,
262                self.metrics.clone(),
263                &self.task_owner,
264            )?)
265        } else {
266            None
267        };
268
269        let reader = JsReader {
270            consumer,
271            tx,
272            shutdown: shutdown_rx,
273            consecutive_errors: Arc::new(AtomicU32::new(0)),
274            data_ready: Arc::clone(&self.data_ready),
275            metrics: self.metrics.clone(),
276            batch_size: cfg.fetch_batch,
277            max_wait: cfg.fetch_max_wait,
278            lag_poll_interval: cfg.lag_poll_interval,
279            requires_ack,
280        };
281        let reader = TrackedTask::spawn(&self.task_owner, "jetstream-reader", reader.run())?;
282
283        self.running = Some(Running {
284            deserializer,
285            rx: Some(rx),
286            shutdown,
287            reader,
288            ack_runtime,
289        });
290        Ok(())
291    }
292
293    async fn open_core(
294        &mut self,
295        cfg: &NatsSourceConfig,
296        deserializer: Box<dyn RecordDeserializer>,
297    ) -> Result<(), ConnectorError> {
298        let client = connect(cfg, &self.task_owner).await?;
299        let subject = cfg
300            .subject
301            .clone()
302            .ok_or_else(|| err("subject missing after validation"))?;
303        let subscriber = if let Some(group) = cfg.queue_group.as_deref() {
304            client
305                .queue_subscribe(subject, group.to_string())
306                .await
307                .map_err(|error| classify_subscribe_error(&error, "NATS queue subscribe"))?
308        } else {
309            client
310                .subscribe(subject)
311                .await
312                .map_err(|error| classify_subscribe_error(&error, "NATS subscribe"))?
313        };
314
315        let (tx, rx) = mpsc::bounded_async::<Incoming>(cfg.fetch_batch * 2);
316        let (shutdown, shutdown_rx) = watch::channel(false);
317
318        let reader = CoreReader {
319            subscriber,
320            tx,
321            shutdown: shutdown_rx,
322            data_ready: Arc::clone(&self.data_ready),
323        };
324        let reader = TrackedTask::spawn(&self.task_owner, "core-reader", reader.run())?;
325
326        self.running = Some(Running {
327            deserializer,
328            rx: Some(rx),
329            shutdown,
330            reader,
331            ack_runtime: None,
332        });
333        Ok(())
334    }
335}
336
337#[async_trait]
338impl SourceConnector for NatsSource {
339    fn terminal_task_tracker(&self) -> Option<ConnectorTaskTracker> {
340        Some(self.task_tracker.clone())
341    }
342
343    fn contract(&self, _config: &ConnectorConfig) -> Result<SourceContract, ConnectorError> {
344        // Neither Core NATS nor the current JetStream implementation can
345        // rewind an abandoned checkpoint attempt deterministically. Durable
346        // consumers alone are insufficient for LaminarDB replay semantics.
347        Ok(SourceContract::new(
348            SourceConsistency::Ephemeral,
349            SourceTopology::Singleton,
350        ))
351    }
352
353    async fn start(&mut self, request: SourceStart) -> Result<(), ConnectorError> {
354        let (config, position, _) = request.into_parts();
355        if let SourcePosition::Resume { attempt, .. } = position {
356            return Err(ConnectorError::ConfigurationError(format!(
357                "NATS is an ephemeral source and cannot resume checkpoint attempt {attempt:?}"
358            )));
359        }
360        let config = &config;
361
362        let cfg = NatsSourceConfig::from_config(config)?;
363        // Keep the candidate schema local until network admission succeeds so
364        // cancelling start leaves the existing instance unchanged.
365        let candidate_schema = config.arrow_schema();
366        let deserializer = serde::create_deserializer(cfg.format)
367            .map_err(|e| err(&format!("deserializer for format {:?}: {e}", cfg.format)))?;
368        match cfg.mode {
369            Mode::JetStream => self.open_jetstream(&cfg, deserializer).await?,
370            Mode::Core => self.open_core(&cfg, deserializer).await?,
371        }
372        if let Some(schema) = candidate_schema {
373            self.schema = schema;
374        }
375        self.config = Some(cfg);
376        Ok(())
377    }
378
379    async fn poll_batch(
380        &mut self,
381        max_records: usize,
382    ) -> Result<Option<SourceBatch>, ConnectorError> {
383        let Some(running) = self.running.as_mut() else {
384            return Ok(None);
385        };
386
387        let mut payloads: Vec<Bytes> = Vec::new();
388        let mut new_acks: Vec<jetstream::Message> = Vec::new();
389        let mut reader_disconnected = false;
390
391        while payloads.len() < max_records {
392            let incoming = match running
393                .rx
394                .as_mut()
395                .expect("running NATS source owns its receiver")
396                .try_recv()
397            {
398                Ok(m) => m,
399                Err(TryRecvError::Empty) => break,
400                Err(TryRecvError::Disconnected) => {
401                    reader_disconnected = true;
402                    break;
403                }
404            };
405            payloads.push(incoming.payload);
406            if let Some(msg) = incoming.ack {
407                new_acks.push(msg);
408            }
409        }
410
411        if payloads.is_empty() {
412            if reader_disconnected {
413                return Err(ConnectorError::ReadError(
414                    "NATS reader task terminated unexpectedly".into(),
415                ));
416            }
417            return Ok(None);
418        }
419
420        let records: Vec<&[u8]> = payloads.iter().map(Bytes::as_ref).collect();
421        let bytes_total: u64 = records.iter().map(|r| r.len() as u64).sum();
422        // Deserialize before scheduling acks: on failure the handles drop unacked and the broker
423        // redelivers after ack_wait. Ack enqueue is non-blocking to keep this poll path bounded.
424        let batch = running
425            .deserializer
426            .deserialize_batch(&records, &self.schema)
427            .map_err(|e| err(&format!("deserialize batch: {e}")))?;
428
429        enqueue_acks(running.ack_runtime.as_ref(), new_acks, &self.metrics);
430
431        self.metrics
432            .record_poll(batch.num_rows() as u64, bytes_total);
433
434        Ok(Some(SourceBatch::new(batch)))
435    }
436
437    fn schema(&self) -> SchemaRef {
438        self.schema.clone()
439    }
440
441    fn checkpoint(&self) -> SourceCheckpoint {
442        // Ephemeral sources deliberately expose no recovery cursor. JetStream acknowledgements
443        // are delivery progress, not checkpoint-owned state.
444        SourceCheckpoint::new()
445    }
446
447    async fn close(&mut self) -> Result<(), ConnectorError> {
448        let Some(mut running) = self.running.take() else {
449            return Ok(());
450        };
451        let close_deadline = tokio::time::Instant::now() + CLOSE_DRAIN_TIMEOUT;
452        running.request_shutdown();
453        match running.reader.wait_until(close_deadline).await {
454            TaskWait::Completed(Ok(())) => {}
455            TaskWait::Completed(Err(error)) => {
456                warn!(%error, "NATS reader task failed while closing");
457            }
458            TaskWait::TimedOut => warn!(
459                "NATS reader exceeded its close deadline; its tracked reaper retains shutdown ownership"
460            ),
461        }
462        // Drop unread messages only after the reader has stopped. Their unacked JetStream
463        // handles remain eligible for broker redelivery.
464        running.rx.take();
465
466        if let Some(mut ack_runtime) = running.ack_runtime.take() {
467            ack_runtime.request_shutdown();
468            match ack_runtime.task.wait_until(close_deadline).await {
469                TaskWait::Completed(Ok(())) => {}
470                TaskWait::Completed(Err(error)) => {
471                    warn!(%error, "NATS ack worker failed while closing");
472                    self.metrics.record_abandoned_acks();
473                }
474                TaskWait::TimedOut => {
475                    warn!(
476                        "NATS ack worker exceeded its close deadline; its tracked reaper retains shutdown ownership"
477                    );
478                }
479            }
480        }
481        Ok(())
482    }
483
484    fn data_ready_notify(&self) -> Option<Arc<Notify>> {
485        Some(Arc::clone(&self.data_ready))
486    }
487}
488
489// ── helpers ──
490
491fn err(msg: &str) -> ConnectorError {
492    ConnectorError::ConfigurationError(msg.to_string())
493}
494
495fn ack_runtime_limits(cfg: &NatsSourceConfig) -> (usize, usize) {
496    let broker_limit = usize::try_from(cfg.max_ack_pending)
497        .ok()
498        .filter(|limit| *limit > 0);
499    let fallback = cfg.fetch_batch.saturating_mul(2);
500    let backlog = broker_limit.unwrap_or(fallback).clamp(1, MAX_ACK_BACKLOG);
501    let concurrency = cfg.fetch_batch.clamp(1, MAX_ACK_CONCURRENCY).min(backlog);
502    (backlog, concurrency)
503}
504
505fn enqueue_acks(
506    runtime: Option<&AckRuntime>,
507    messages: Vec<jetstream::Message>,
508    metrics: &NatsSourceMetrics,
509) {
510    if messages.is_empty() {
511        return;
512    }
513    let Some(runtime) = runtime else {
514        metrics.record_ack_enqueue_errors(messages.len());
515        warn!(
516            rejected = messages.len(),
517            "JetStream ack worker is unavailable; broker will redeliver"
518        );
519        return;
520    };
521
522    let mut rejected = 0usize;
523    for message in messages {
524        // Increment before publication: a fast worker may complete immediately after try_send.
525        metrics.record_ack_enqueued();
526        if runtime
527            .tx
528            .as_ref()
529            .is_none_or(|tx| tx.try_send(message).is_err())
530        {
531            metrics.record_ack_error();
532            rejected += 1;
533        }
534    }
535    if rejected > 0 {
536        warn!(
537            rejected,
538            "JetStream ack backlog is full or closed; broker will redeliver"
539        );
540    }
541}
542
543async fn run_ack_worker(
544    rx: tokio_mpsc::Receiver<jetstream::Message>,
545    shutdown: watch::Receiver<bool>,
546    concurrency: usize,
547    metrics: NatsSourceMetrics,
548) {
549    // Ack calls stay scoped under the worker. Its single generation guard therefore proves that
550    // the receiver and every in-flight acknowledgement have all been dropped or completed.
551    let task_metrics = metrics.clone();
552    let abandoned = run_bounded_queue(rx, shutdown, concurrency, move |message| {
553        let metrics = task_metrics.clone();
554        async move {
555            acknowledge_message(message, &metrics).await;
556        }
557    })
558    .await;
559    if abandoned > 0 {
560        metrics.record_ack_abandoned(abandoned);
561        warn!(
562            abandoned,
563            "discarded queued JetStream acknowledgements during shutdown; broker will redeliver"
564        );
565    }
566}
567
568async fn run_bounded_queue<T, F, Fut>(
569    mut rx: tokio_mpsc::Receiver<T>,
570    mut shutdown: watch::Receiver<bool>,
571    concurrency: usize,
572    process: F,
573) -> usize
574where
575    T: Send + 'static,
576    F: Fn(T) -> Fut,
577    Fut: std::future::Future<Output = ()> + Send,
578{
579    debug_assert!(concurrency > 0);
580    let mut in_flight = FuturesUnordered::new();
581    'input: loop {
582        if shutdown_requested(&shutdown) || rx.is_closed() {
583            break;
584        }
585        while in_flight.len() >= concurrency {
586            tokio::select! {
587                biased;
588                _ = shutdown.changed() => break 'input,
589                _ = in_flight.next() => {}
590            }
591            if rx.is_closed() {
592                break 'input;
593            }
594        }
595
596        tokio::select! {
597            biased;
598            _ = shutdown.changed() => break,
599            _ = in_flight.next(), if !in_flight.is_empty() => {}
600            message = rx.recv() => {
601                let Some(message) = message else {
602                    break;
603                };
604                in_flight.push(process(message));
605            }
606        }
607    }
608
609    // Closing the receiver makes queued-but-unstarted messages immediately eligible for broker
610    // redelivery. Only work already admitted to `in_flight` is allowed to consume the close budget.
611    rx.close();
612    let mut abandoned = 0usize;
613    while rx.try_recv().is_ok() {
614        abandoned = abandoned.saturating_add(1);
615    }
616    while in_flight.next().await.is_some() {}
617    abandoned
618}
619
620async fn acknowledge_message(message: jetstream::Message, metrics: &NatsSourceMetrics) {
621    match tokio::time::timeout(ACK_IO_TIMEOUT, message.double_ack()).await {
622        Ok(Ok(())) => metrics.record_ack(),
623        Ok(Err(error)) => {
624            metrics.record_ack_error();
625            warn!(%error, "JetStream ack failed; broker will redeliver");
626        }
627        Err(_) => {
628            metrics.record_ack_error();
629            warn!(
630                timeout_ms = ACK_IO_TIMEOUT.as_millis(),
631                "JetStream ack timed out; broker will redeliver"
632            );
633        }
634    }
635}
636
637async fn connect(
638    cfg: &NatsSourceConfig,
639    owner: &ConnectorTaskOwner,
640) -> Result<async_nats::Client, ConnectorError> {
641    track_connection_tasks(build_connect_options(&cfg.auth, &cfg.tls)?, owner, "source")?
642        .connect(&cfg.servers)
643        .await
644        .map_err(|error| classify_connect_error(&error))
645}
646
647fn build_pull_config(
648    cfg: &NatsSourceConfig,
649    consumer_name: &str,
650) -> Result<pull::Config, ConnectorError> {
651    let filter_subjects = if cfg.subject_filters.is_empty() {
652        cfg.subject.iter().cloned().collect()
653    } else {
654        cfg.subject_filters.clone()
655    };
656
657    Ok(pull::Config {
658        durable_name: Some(consumer_name.to_string()),
659        filter_subjects,
660        deliver_policy: map_deliver_policy(cfg)?,
661        ack_policy: map_ack_policy(cfg.ack_policy),
662        ack_wait: cfg.ack_wait,
663        max_deliver: cfg.max_deliver,
664        max_ack_pending: cfg.max_ack_pending,
665        ..Default::default()
666    })
667}
668
669fn map_deliver_policy(
670    cfg: &NatsSourceConfig,
671) -> Result<async_nats::jetstream::consumer::DeliverPolicy, ConnectorError> {
672    use async_nats::jetstream::consumer::DeliverPolicy as Nats;
673    Ok(match cfg.deliver_policy {
674        DeliverPolicy::All => Nats::All,
675        DeliverPolicy::New => Nats::New,
676        DeliverPolicy::ByStartSequence => Nats::ByStartSequence {
677            start_sequence: cfg.start_sequence.unwrap_or(1),
678        },
679        DeliverPolicy::ByStartTime => {
680            let raw = cfg
681                .start_time
682                .as_deref()
683                .ok_or_else(|| err("deliver.policy=by_start_time requires 'start.time'"))?;
684            let start_time =
685                time::OffsetDateTime::parse(raw, &time::format_description::well_known::Rfc3339)
686                    .map_err(|e| err(&format!("start.time '{raw}' is not valid RFC3339: {e}")))?;
687            Nats::ByStartTime { start_time }
688        }
689    })
690}
691
692fn map_ack_policy(p: AckPolicy) -> async_nats::jetstream::consumer::AckPolicy {
693    use async_nats::jetstream::consumer::AckPolicy as Nats;
694    match p {
695        AckPolicy::Explicit => Nats::Explicit,
696        AckPolicy::None => Nats::None,
697    }
698}
699
700/// 500ms, 1s, 2s, 4s, cap 5s.
701fn fetch_backoff_base(consecutive_errors: u32) -> Duration {
702    let exp = consecutive_errors.saturating_sub(1).min(4);
703    let ms = 500u64.saturating_mul(1u64 << exp);
704    Duration::from_millis(ms.min(5000))
705}
706
707/// `base ± 20%`. Tests pass a fixed `entropy` seed.
708fn with_jitter(base: Duration, entropy: u64) -> Duration {
709    let base_ms = u64::try_from(base.as_millis()).unwrap_or(u64::MAX);
710    let range = (base_ms / 5).max(1); // 20%
711    let window = range * 2 + 1;
712    let offset = entropy % window;
713    let jittered = base_ms.saturating_add(offset).saturating_sub(range);
714    Duration::from_millis(jittered)
715}
716
717fn fetch_backoff(consecutive_errors: u32, entropy: u64) -> Duration {
718    with_jitter(fetch_backoff_base(consecutive_errors), entropy)
719}
720
721/// Wall-clock nanos for `with_jitter`. `Instant::now().elapsed()` is ~0
722/// and produces correlated jitter across tasks.
723#[allow(clippy::cast_possible_truncation)]
724fn entropy_now() -> u64 {
725    std::time::SystemTime::now()
726        .duration_since(std::time::UNIX_EPOCH)
727        .unwrap_or_default()
728        .as_nanos() as u64
729}
730
731fn shutdown_requested(shutdown: &watch::Receiver<bool>) -> bool {
732    *shutdown.borrow() || shutdown.has_changed().is_err()
733}
734
735struct JsReader {
736    consumer: jetstream::consumer::Consumer<pull::Config>,
737    tx: MAsyncTx<mpsc::Array<Incoming>>,
738    shutdown: watch::Receiver<bool>,
739    consecutive_errors: Arc<AtomicU32>,
740    data_ready: Arc<Notify>,
741    metrics: NatsSourceMetrics,
742    batch_size: usize,
743    max_wait: Duration,
744    /// `Duration::ZERO` disables the poll.
745    lag_poll_interval: Duration,
746    requires_ack: bool,
747}
748
749impl JsReader {
750    async fn run(self) {
751        let Self {
752            mut consumer,
753            tx,
754            mut shutdown,
755            consecutive_errors,
756            data_ready,
757            metrics,
758            batch_size,
759            max_wait,
760            lag_poll_interval,
761            requires_ack,
762        } = self;
763
764        let mut last_lag_poll = Instant::now();
765        let lag_poll_enabled = !lag_poll_interval.is_zero();
766
767        loop {
768            if shutdown_requested(&shutdown) {
769                break;
770            }
771            let fetch_result = tokio::select! {
772                biased;
773                _ = shutdown.changed() => break,
774                r = consumer.fetch().max_messages(batch_size).expires(max_wait).messages() => r,
775            };
776
777            let mut stream = match fetch_result {
778                Ok(s) => s,
779                Err(e) => {
780                    let errs = consecutive_errors.fetch_add(1, Ordering::AcqRel) + 1;
781                    metrics.record_fetch_error();
782                    warn!(
783                        error = %e,
784                        consecutive_errors = errs,
785                        "nats fetch() errored; backing off",
786                    );
787                    let backoff = fetch_backoff(errs, entropy_now());
788                    tokio::select! {
789                        biased;
790                        _ = shutdown.changed() => break,
791                        () = tokio::time::sleep(backoff) => {}
792                    }
793                    continue;
794                }
795            };
796
797            let mut forwarded = 0usize;
798            let mut stream_errors = 0usize;
799            loop {
800                let msg_result = tokio::select! {
801                    biased;
802                    _ = shutdown.changed() => return,
803                    r = stream.next() => match r {
804                        Some(r) => r,
805                        None => break,
806                    },
807                };
808                let msg = match msg_result {
809                    Ok(m) => m,
810                    Err(e) => {
811                        metrics.record_fetch_error();
812                        stream_errors += 1;
813                        warn!(error = %e, "nats message error");
814                        continue;
815                    }
816                };
817                let payload = msg.payload.clone();
818                let incoming = Incoming {
819                    payload,
820                    ack: requires_ack.then_some(msg),
821                };
822                let send_result = tokio::select! {
823                    biased;
824                    _ = shutdown.changed() => return,
825                    result = tx.send(incoming) => result,
826                };
827                if send_result.is_err() {
828                    debug!("nats reader: downstream channel closed");
829                    return;
830                }
831                forwarded += 1;
832            }
833
834            // Reset on progress; an iteration with only errors counts
835            // as one failure; idle iterations don't bump.
836            if forwarded > 0 {
837                consecutive_errors.store(0, Ordering::Release);
838                data_ready.notify_one();
839            } else if stream_errors > 0 {
840                let errs = consecutive_errors.fetch_add(1, Ordering::AcqRel) + 1;
841                let backoff = fetch_backoff(errs, entropy_now());
842                tokio::select! {
843                    biased;
844                    _ = shutdown.changed() => break,
845                    () = tokio::time::sleep(backoff) => {}
846                }
847            }
848
849            if lag_poll_enabled && last_lag_poll.elapsed() >= lag_poll_interval {
850                last_lag_poll = Instant::now();
851                match consumer.info().await {
852                    Ok(info) => metrics.set_consumer_lag(info.num_pending),
853                    Err(e) => warn!(error = %e, "consumer.info() failed; skipping lag update"),
854                }
855            }
856        }
857    }
858}
859
860struct CoreReader {
861    subscriber: async_nats::Subscriber,
862    tx: MAsyncTx<mpsc::Array<Incoming>>,
863    shutdown: watch::Receiver<bool>,
864    data_ready: Arc<Notify>,
865}
866
867impl CoreReader {
868    async fn run(self) {
869        let Self {
870            mut subscriber,
871            tx,
872            mut shutdown,
873            data_ready,
874        } = self;
875
876        loop {
877            if shutdown_requested(&shutdown) {
878                break;
879            }
880            let msg = tokio::select! {
881                biased;
882                _ = shutdown.changed() => break,
883                m = subscriber.next() => match m {
884                    Some(m) => m,
885                    None => break,
886                },
887            };
888            let incoming = Incoming {
889                payload: msg.payload,
890                ack: None,
891            };
892            let send_result = tokio::select! {
893                biased;
894                _ = shutdown.changed() => break,
895                result = tx.send(incoming) => result,
896            };
897            if send_result.is_err() {
898                break;
899            }
900            data_ready.notify_one();
901        }
902        // Wake the coordinator so it observes the now-disconnected channel immediately instead
903        // of treating a terminated Core subscription as an indefinitely idle source.
904        data_ready.notify_one();
905    }
906}
907
908#[cfg(test)]
909mod tests {
910    use super::*;
911    use arrow_schema::Schema;
912
913    struct DropSignal(Option<tokio::sync::oneshot::Sender<()>>);
914
915    impl Drop for DropSignal {
916        fn drop(&mut self) {
917            if let Some(tx) = self.0.take() {
918                let _ = tx.send(());
919            }
920        }
921    }
922
923    async fn pending_task(
924        started: tokio::sync::oneshot::Sender<()>,
925        dropped: tokio::sync::oneshot::Sender<()>,
926        release: tokio::sync::oneshot::Receiver<()>,
927    ) {
928        let _drop_signal = DropSignal(Some(dropped));
929        let _ = started.send(());
930        let _ = release.await;
931    }
932
933    #[test]
934    fn source_contract_is_ephemeral_even_for_jetstream_config() {
935        let source = NatsSource::new(Arc::new(Schema::empty()), None);
936        let mut config = ConnectorConfig::new("nats");
937        config.set("mode", "jetstream");
938        let contract = source.contract(&config).expect("static NATS contract");
939        assert_eq!(contract.consistency, SourceConsistency::Ephemeral);
940        assert_eq!(contract.topology, SourceTopology::Singleton);
941    }
942
943    #[test]
944    fn ephemeral_source_checkpoint_has_no_protocol_state() {
945        let src = NatsSource::new(Arc::new(Schema::empty()), None);
946        assert!(src.checkpoint().is_empty());
947        assert_eq!(
948            src.cancellation_policy(),
949            crate::connector::ConnectorCancellationPolicy::RetireConnector
950        );
951    }
952
953    #[test]
954    fn task_tracker_notifies_waiters_on_another_runtime() {
955        let (release_tx, release_rx) = tokio::sync::oneshot::channel();
956        let (tracker_tx, tracker_rx) = std::sync::mpsc::sync_channel(1);
957        let owner_thread = std::thread::spawn(move || {
958            tokio::runtime::Builder::new_current_thread()
959                .enable_all()
960                .build()
961                .unwrap()
962                .block_on(async move {
963                    let mut source = NatsSource::new(Arc::new(Schema::empty()), None);
964                    let terminal = source.terminal_task_tracker().unwrap();
965                    let owner_waiter = terminal.clone();
966                    let (_, rx) = mpsc::bounded_async::<Incoming>(1);
967                    let (shutdown, _) = watch::channel(false);
968                    let reader = TrackedTask::spawn(
969                        &source.task_owner,
970                        "cross-runtime-reader",
971                        async move {
972                            let _ = release_rx.await;
973                        },
974                    )
975                    .unwrap();
976                    source.running = Some(Running {
977                        deserializer: serde::create_deserializer(serde::Format::Raw).unwrap(),
978                        rx: Some(rx),
979                        shutdown,
980                        reader,
981                        ack_runtime: None,
982                    });
983                    tracker_tx.send(terminal).unwrap();
984                    drop(source);
985                    owner_waiter.wait_terminated().await;
986                });
987        });
988
989        let terminal = tracker_rx.recv().unwrap();
990        assert!(!terminal.is_terminated());
991        tokio::runtime::Builder::new_current_thread()
992            .enable_all()
993            .build()
994            .unwrap()
995            .block_on(async move {
996                release_tx.send(()).unwrap();
997                tokio::time::timeout(Duration::from_secs(1), terminal.wait_terminated())
998                    .await
999                    .expect("cross-runtime tracker waiter was not notified");
1000            });
1001        owner_thread.join().unwrap();
1002    }
1003
1004    #[tokio::test]
1005    async fn disconnected_reader_drains_queued_payload_before_terminal_error() {
1006        let mut src = NatsSource::new(Arc::new(Schema::empty()), None);
1007        let (tx, rx) = mpsc::bounded_async::<Incoming>(2);
1008        assert!(tx
1009            .try_send(Incoming {
1010                payload: Bytes::from_static(b"one"),
1011                ack: None,
1012            })
1013            .is_ok());
1014        drop(tx);
1015        let (shutdown, _) = watch::channel(false);
1016        let reader = TrackedTask::spawn(&src.task_owner, "test-reader", async {}).unwrap();
1017        src.running = Some(Running {
1018            deserializer: serde::create_deserializer(serde::Format::Raw).unwrap(),
1019            rx: Some(rx),
1020            shutdown,
1021            reader,
1022            ack_runtime: None,
1023        });
1024
1025        let final_batch = src.poll_batch(10).await.unwrap().unwrap();
1026        assert_eq!(final_batch.records.num_rows(), 1);
1027
1028        let error = src.poll_batch(10).await.unwrap_err();
1029        assert!(matches!(error, ConnectorError::ReadError(_)));
1030        assert!(error.to_string().contains("reader task terminated"));
1031    }
1032
1033    #[tokio::test]
1034    async fn dropping_source_signals_and_reaps_the_owned_reader() {
1035        let mut source = NatsSource::new(Arc::new(Schema::empty()), None);
1036        let terminal = source.terminal_task_tracker().unwrap();
1037        let (_, rx) = mpsc::bounded_async::<Incoming>(1);
1038        let (shutdown, mut task_shutdown) = watch::channel(false);
1039        let shutdown_observer = task_shutdown.clone();
1040        let (started_tx, started_rx) = tokio::sync::oneshot::channel();
1041        let (dropped_tx, dropped_rx) = tokio::sync::oneshot::channel();
1042        let reader = TrackedTask::spawn(&source.task_owner, "test-reader", async move {
1043            let _drop_signal = DropSignal(Some(dropped_tx));
1044            let _ = started_tx.send(());
1045            let _ = task_shutdown.changed().await;
1046        })
1047        .unwrap();
1048        source.running = Some(Running {
1049            deserializer: serde::create_deserializer(serde::Format::Raw).unwrap(),
1050            rx: Some(rx),
1051            shutdown,
1052            reader,
1053            ack_runtime: None,
1054        });
1055        started_rx.await.expect("reader task started");
1056
1057        drop(source);
1058
1059        assert!(*shutdown_observer.borrow(), "drop must publish shutdown");
1060        tokio::time::timeout(Duration::from_secs(1), dropped_rx)
1061            .await
1062            .expect("reader must observe shutdown on drop")
1063            .expect("reader drop signal");
1064        tokio::time::timeout(Duration::from_secs(1), terminal.wait_terminated())
1065            .await
1066            .expect("reader and its tracked reaper must terminate");
1067    }
1068
1069    #[tokio::test]
1070    async fn normal_close_joins_reader_and_ack_tasks() {
1071        let mut source = NatsSource::new(Arc::new(Schema::empty()), None);
1072        let (_, rx) = mpsc::bounded_async::<Incoming>(1);
1073        let (shutdown, mut shutdown_rx) = watch::channel(false);
1074        let (reader_started_tx, reader_started_rx) = tokio::sync::oneshot::channel();
1075        let (reader_dropped_tx, reader_dropped_rx) = tokio::sync::oneshot::channel();
1076        let reader = TrackedTask::spawn(&source.task_owner, "test-reader", async move {
1077            let _drop_signal = DropSignal(Some(reader_dropped_tx));
1078            let _ = reader_started_tx.send(());
1079            let _ = shutdown_rx.changed().await;
1080        })
1081        .unwrap();
1082
1083        let (ack_tx, mut ack_rx) = tokio_mpsc::channel::<jetstream::Message>(1);
1084        let (ack_shutdown, _) = watch::channel(false);
1085        let (ack_started_tx, ack_started_rx) = tokio::sync::oneshot::channel();
1086        let (ack_dropped_tx, ack_dropped_rx) = tokio::sync::oneshot::channel();
1087        let ack_task = TrackedTask::spawn(&source.task_owner, "test-ack", async move {
1088            let _drop_signal = DropSignal(Some(ack_dropped_tx));
1089            let _ = ack_started_tx.send(());
1090            while ack_rx.recv().await.is_some() {}
1091        })
1092        .unwrap();
1093
1094        source.running = Some(Running {
1095            deserializer: serde::create_deserializer(serde::Format::Raw).unwrap(),
1096            rx: Some(rx),
1097            shutdown,
1098            reader,
1099            ack_runtime: Some(AckRuntime {
1100                tx: Some(ack_tx),
1101                shutdown: ack_shutdown,
1102                task: ack_task,
1103            }),
1104        });
1105        reader_started_rx.await.expect("reader task started");
1106        ack_started_rx.await.expect("ack task started");
1107
1108        tokio::time::timeout(Duration::from_secs(1), source.close())
1109            .await
1110            .expect("normal close must join owned tasks")
1111            .unwrap();
1112
1113        for (name, dropped) in [("reader", reader_dropped_rx), ("ack", ack_dropped_rx)] {
1114            tokio::time::timeout(Duration::from_secs(1), dropped)
1115                .await
1116                .unwrap_or_else(|_| panic!("{name} task was not joined"))
1117                .unwrap_or_else(|_| panic!("{name} drop signal closed"));
1118        }
1119        assert!(source.running.is_none());
1120    }
1121
1122    #[tokio::test]
1123    async fn cancelling_close_does_not_detach_reader_or_ack_tasks() {
1124        let mut source = NatsSource::new(Arc::new(Schema::empty()), None);
1125        let terminal = source.terminal_task_tracker().unwrap();
1126        let (_, rx) = mpsc::bounded_async::<Incoming>(1);
1127        let (shutdown, shutdown_rx) = watch::channel(false);
1128        let (reader_started_tx, reader_started_rx) = tokio::sync::oneshot::channel();
1129        let (reader_dropped_tx, reader_dropped_rx) = tokio::sync::oneshot::channel();
1130        let (reader_release_tx, reader_release_rx) = tokio::sync::oneshot::channel();
1131        let reader = TrackedTask::spawn(
1132            &source.task_owner,
1133            "test-reader",
1134            pending_task(reader_started_tx, reader_dropped_tx, reader_release_rx),
1135        )
1136        .unwrap();
1137
1138        let (ack_tx, ack_rx) = tokio_mpsc::channel::<jetstream::Message>(1);
1139        let (ack_shutdown, _) = watch::channel(false);
1140        let (ack_started_tx, ack_started_rx) = tokio::sync::oneshot::channel();
1141        let (ack_dropped_tx, ack_dropped_rx) = tokio::sync::oneshot::channel();
1142        let (ack_release_tx, ack_release_rx) = tokio::sync::oneshot::channel();
1143        let ack_task = TrackedTask::spawn(&source.task_owner, "test-ack", async move {
1144            let _ack_rx = ack_rx;
1145            pending_task(ack_started_tx, ack_dropped_tx, ack_release_rx).await;
1146        })
1147        .unwrap();
1148        source.running = Some(Running {
1149            deserializer: serde::create_deserializer(serde::Format::Raw).unwrap(),
1150            rx: Some(rx),
1151            shutdown,
1152            reader,
1153            ack_runtime: Some(AckRuntime {
1154                tx: Some(ack_tx),
1155                shutdown: ack_shutdown,
1156                task: ack_task,
1157            }),
1158        });
1159        reader_started_rx.await.expect("reader task started");
1160        ack_started_rx.await.expect("ack task started");
1161
1162        let close = tokio::spawn(async move { source.close().await });
1163        tokio::task::yield_now().await;
1164        assert!(!close.is_finished(), "close must be waiting for the reader");
1165        close.abort();
1166        assert!(close
1167            .await
1168            .expect_err("close waiter cancelled")
1169            .is_cancelled());
1170
1171        assert!(
1172            *shutdown_rx.borrow(),
1173            "cancelling close must publish shutdown"
1174        );
1175        assert!(
1176            !terminal.is_terminated(),
1177            "task guards must keep a cancelled generation non-terminal"
1178        );
1179        reader_release_tx.send(()).expect("release reader");
1180        ack_release_tx.send(()).expect("release ack worker");
1181
1182        for (name, dropped) in [("reader", reader_dropped_rx), ("ack", ack_dropped_rx)] {
1183            tokio::time::timeout(Duration::from_secs(1), dropped)
1184                .await
1185                .unwrap_or_else(|_| panic!("{name} task remained detached"))
1186                .unwrap_or_else(|_| panic!("{name} drop signal closed"));
1187        }
1188        tokio::time::timeout(Duration::from_secs(1), terminal.wait_terminated())
1189            .await
1190            .expect("generation must become terminal after every owned task exits");
1191    }
1192
1193    #[tokio::test]
1194    async fn ack_shutdown_discards_queued_but_unstarted_work() {
1195        let (tx, rx) = tokio_mpsc::channel(8);
1196        let (shutdown_tx, shutdown_rx) = watch::channel(false);
1197        let (admitted_tx, mut admitted_rx) = tokio_mpsc::unbounded_channel();
1198        let release = Arc::new(Notify::new());
1199        let worker_release = Arc::clone(&release);
1200        let worker = tokio::spawn(run_bounded_queue(rx, shutdown_rx, 1, move |message| {
1201            let admitted_tx = admitted_tx.clone();
1202            let release = Arc::clone(&worker_release);
1203            async move {
1204                admitted_tx.send(message).unwrap();
1205                release.notified().await;
1206            }
1207        }));
1208
1209        for message in 1..=3 {
1210            tx.send(message).await.unwrap();
1211        }
1212        assert_eq!(admitted_rx.recv().await, Some(1));
1213
1214        shutdown_tx.send_replace(true);
1215        drop(tx);
1216        tokio::task::yield_now().await;
1217        release.notify_one();
1218
1219        let abandoned = tokio::time::timeout(Duration::from_secs(1), worker)
1220            .await
1221            .expect("worker shutdown must be bounded by admitted work")
1222            .unwrap();
1223        assert_eq!(abandoned, 2);
1224        assert!(admitted_rx.try_recv().is_err());
1225    }
1226
1227    #[test]
1228    fn backoff_base_grows_then_caps_at_5s() {
1229        assert_eq!(fetch_backoff_base(1), Duration::from_millis(500));
1230        assert_eq!(fetch_backoff_base(2), Duration::from_millis(1000));
1231        assert_eq!(fetch_backoff_base(3), Duration::from_millis(2000));
1232        assert_eq!(fetch_backoff_base(4), Duration::from_millis(4000));
1233        assert_eq!(fetch_backoff_base(5), Duration::from_millis(5000));
1234        assert_eq!(fetch_backoff_base(100), Duration::from_millis(5000));
1235    }
1236
1237    #[test]
1238    fn jitter_stays_within_plus_minus_20_percent() {
1239        let base = Duration::from_millis(1000);
1240        for entropy in [0u64, 1, 99, 12345, u64::MAX] {
1241            let j = with_jitter(base, entropy);
1242            assert!(
1243                j >= Duration::from_millis(800) && j <= Duration::from_millis(1200),
1244                "entropy {entropy}: jittered = {j:?} outside ±20% of {base:?}"
1245            );
1246        }
1247    }
1248}