Skip to main content

laminar_connectors/nats/source/
mod.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, SourceInputMode, SourcePosition, SourceStart,
31    SourceTopology,
32};
33use crate::error::ConnectorError;
34use crate::serde::{self, RecordDeserializer};
35
36const ACK_IO_TIMEOUT: Duration = Duration::from_secs(5);
37const CLOSE_DRAIN_TIMEOUT: Duration = Duration::from_secs(5);
38const MAX_ACK_CONCURRENCY: usize = 64;
39const MAX_ACK_BACKLOG: usize = 16_384;
40
41/// `ack` is `Some` only on the `JetStream` path.
42struct Incoming {
43    payload: Bytes,
44    ack: Option<jetstream::Message>,
45}
46
47struct AckRuntime {
48    tx: Option<tokio_mpsc::Sender<jetstream::Message>>,
49    shutdown: watch::Sender<bool>,
50    task: TrackedTask,
51}
52
53struct Running {
54    deserializer: Box<dyn RecordDeserializer>,
55    rx: Option<AsyncRx<mpsc::Array<Incoming>>>,
56    shutdown: watch::Sender<bool>,
57    reader: TrackedTask,
58    ack_runtime: Option<AckRuntime>,
59}
60
61struct TrackedTask {
62    handle: Option<JoinHandle<()>>,
63    reaper_guard: Option<ConnectorTaskGuard>,
64    name: &'static str,
65}
66
67enum TaskWait {
68    Completed(Result<(), tokio::task::JoinError>),
69    TimedOut,
70}
71
72impl TrackedTask {
73    fn spawn(
74        owner: &ConnectorTaskOwner,
75        name: &'static str,
76        future: impl std::future::Future<Output = ()> + Send + 'static,
77    ) -> Result<Self, ConnectorError> {
78        let task_guard = owner.track().ok_or_else(|| {
79            ConnectorError::Internal("NATS source task generation is already retired".into())
80        })?;
81        let reaper_guard = owner.track().ok_or_else(|| {
82            ConnectorError::Internal("NATS source task generation is already retired".into())
83        })?;
84        let handle = tokio::spawn(async move {
85            let _task_guard = task_guard;
86            future.await;
87        });
88        Ok(Self {
89            handle: Some(handle),
90            reaper_guard: Some(reaper_guard),
91            name,
92        })
93    }
94
95    async fn wait_until(&mut self, deadline: tokio::time::Instant) -> TaskWait {
96        let Some(handle) = self.handle.as_mut() else {
97            return TaskWait::Completed(Ok(()));
98        };
99        match tokio::time::timeout_at(deadline, handle).await {
100            Ok(result) => {
101                self.handle.take();
102                self.reaper_guard.take();
103                TaskWait::Completed(result)
104            }
105            Err(_) => TaskWait::TimedOut,
106        }
107    }
108
109    fn retire(&mut self) {
110        let Some(handle) = self.handle.take() else {
111            self.reaper_guard.take();
112            return;
113        };
114        let reaper_guard = self.reaper_guard.take();
115        let Ok(runtime) = tokio::runtime::Handle::try_current() else {
116            // Runtime destruction drops the task future and its task guard. That guard is the
117            // completion proof; a join task is useful for cleanup but is not itself the proof.
118            drop(handle);
119            drop(reaper_guard);
120            return;
121        };
122        let name = self.name;
123        drop(runtime.spawn(async move {
124            let _reaper_guard = reaper_guard;
125            if let Err(error) = handle.await {
126                debug!(task = name, %error, "retired NATS source task reaped");
127            }
128        }));
129    }
130}
131
132impl Drop for TrackedTask {
133    fn drop(&mut self) {
134        self.retire();
135    }
136}
137
138impl Running {
139    fn request_shutdown(&self) {
140        self.shutdown.send_replace(true);
141    }
142}
143
144impl Drop for Running {
145    fn drop(&mut self) {
146        // Drop is the final backstop for a cancelled startup/close future. Tasks retain their
147        // generation guards until they actually exit; the reapers retain the join handles.
148        self.request_shutdown();
149        self.rx.take();
150        if let Some(ack_runtime) = self.ack_runtime.as_mut() {
151            ack_runtime.request_shutdown();
152            ack_runtime.task.retire();
153        }
154        self.reader.retire();
155    }
156}
157
158impl AckRuntime {
159    fn spawn(
160        cfg: &NatsSourceConfig,
161        metrics: NatsSourceMetrics,
162        owner: &ConnectorTaskOwner,
163    ) -> Result<Self, ConnectorError> {
164        let (backlog, concurrency) = ack_runtime_limits(cfg);
165        let (tx, rx) = tokio_mpsc::channel(backlog);
166        let (shutdown, shutdown_rx) = watch::channel(false);
167        let task = TrackedTask::spawn(
168            owner,
169            "ack-worker",
170            run_ack_worker(rx, shutdown_rx, concurrency, metrics),
171        )?;
172        Ok(Self {
173            tx: Some(tx),
174            shutdown,
175            task,
176        })
177    }
178
179    fn request_shutdown(&mut self) {
180        self.shutdown.send_replace(true);
181        self.tx.take();
182    }
183}
184
185impl Drop for AckRuntime {
186    fn drop(&mut self) {
187        self.request_shutdown();
188    }
189}
190
191/// NATS source — core and `JetStream` modes.
192pub struct NatsSource {
193    schema: SchemaRef,
194    config: Option<NatsSourceConfig>,
195    data_ready: Arc<Notify>,
196    metrics: NatsSourceMetrics,
197    running: Option<Running>,
198    task_owner: ConnectorTaskOwner,
199    task_tracker: ConnectorTaskTracker,
200}
201
202impl NatsSource {
203    /// Metrics register on `registry` if provided.
204    #[must_use]
205    pub fn new(schema: SchemaRef, registry: Option<&prometheus::Registry>) -> Self {
206        let (task_owner, task_tracker) = ConnectorTaskOwner::new();
207        Self {
208            schema,
209            config: None,
210            data_ready: Arc::new(Notify::new()),
211            metrics: NatsSourceMetrics::new(registry),
212            running: None,
213            task_owner,
214            task_tracker,
215        }
216    }
217
218    /// Available after [`SourceConnector::start`].
219    #[must_use]
220    pub fn config(&self) -> Option<&NatsSourceConfig> {
221        self.config.as_ref()
222    }
223
224    /// Snapshot accessor for the prometheus-backed metrics struct.
225    #[must_use]
226    pub fn metrics_handle(&self) -> &NatsSourceMetrics {
227        &self.metrics
228    }
229
230    async fn open_jetstream(
231        &mut self,
232        cfg: &NatsSourceConfig,
233        deserializer: Box<dyn RecordDeserializer>,
234    ) -> Result<(), ConnectorError> {
235        let client = connect(cfg, &self.task_owner).await?;
236        let js = jetstream::new(client);
237
238        let stream_name = cfg
239            .stream
240            .as_deref()
241            .ok_or_else(|| err("stream name missing after validation"))?;
242        let consumer_name = cfg
243            .consumer
244            .as_deref()
245            .ok_or_else(|| err("consumer name missing after validation"))?;
246
247        let pull_cfg = build_pull_config(cfg, consumer_name)?;
248        let stream = js
249            .get_stream(stream_name)
250            .await
251            .map_err(|error| classify_get_stream_error(&error, stream_name))?;
252        let consumer = stream
253            .create_consumer(pull_cfg)
254            .await
255            .map_err(|error| classify_create_consumer_error(&error, consumer_name))?;
256
257        let (tx, rx) = mpsc::bounded_async::<Incoming>(cfg.fetch_batch * 2);
258        let (shutdown, shutdown_rx) = watch::channel(false);
259        let requires_ack = cfg.ack_policy == AckPolicy::Explicit;
260        let ack_runtime = if requires_ack {
261            Some(AckRuntime::spawn(
262                cfg,
263                self.metrics.clone(),
264                &self.task_owner,
265            )?)
266        } else {
267            None
268        };
269
270        let reader = JsReader {
271            consumer,
272            tx,
273            shutdown: shutdown_rx,
274            consecutive_errors: Arc::new(AtomicU32::new(0)),
275            data_ready: Arc::clone(&self.data_ready),
276            metrics: self.metrics.clone(),
277            batch_size: cfg.fetch_batch,
278            max_wait: cfg.fetch_max_wait,
279            lag_poll_interval: cfg.lag_poll_interval,
280            requires_ack,
281        };
282        let reader = TrackedTask::spawn(&self.task_owner, "jetstream-reader", reader.run())?;
283
284        self.running = Some(Running {
285            deserializer,
286            rx: Some(rx),
287            shutdown,
288            reader,
289            ack_runtime,
290        });
291        Ok(())
292    }
293
294    async fn open_core(
295        &mut self,
296        cfg: &NatsSourceConfig,
297        deserializer: Box<dyn RecordDeserializer>,
298    ) -> Result<(), ConnectorError> {
299        let client = connect(cfg, &self.task_owner).await?;
300        let subject = cfg
301            .subject
302            .clone()
303            .ok_or_else(|| err("subject missing after validation"))?;
304        let subscriber = if let Some(group) = cfg.queue_group.as_deref() {
305            client
306                .queue_subscribe(subject, group.to_string())
307                .await
308                .map_err(|error| classify_subscribe_error(&error, "NATS queue subscribe"))?
309        } else {
310            client
311                .subscribe(subject)
312                .await
313                .map_err(|error| classify_subscribe_error(&error, "NATS subscribe"))?
314        };
315
316        let (tx, rx) = mpsc::bounded_async::<Incoming>(cfg.fetch_batch * 2);
317        let (shutdown, shutdown_rx) = watch::channel(false);
318
319        let reader = CoreReader {
320            subscriber,
321            tx,
322            shutdown: shutdown_rx,
323            data_ready: Arc::clone(&self.data_ready),
324        };
325        let reader = TrackedTask::spawn(&self.task_owner, "core-reader", reader.run())?;
326
327        self.running = Some(Running {
328            deserializer,
329            rx: Some(rx),
330            shutdown,
331            reader,
332            ack_runtime: None,
333        });
334        Ok(())
335    }
336}
337
338#[async_trait]
339impl SourceConnector for NatsSource {
340    fn terminal_task_tracker(&self) -> Option<ConnectorTaskTracker> {
341        Some(self.task_tracker.clone())
342    }
343
344    fn contract(&self, config: &ConnectorConfig) -> Result<SourceContract, ConnectorError> {
345        let format = match config.get("format") {
346            Some(value) => serde::Format::parse(value)
347                .map_err(|error| ConnectorError::ConfigurationError(error.to_string()))?,
348            None => self
349                .config
350                .as_ref()
351                .map_or(serde::Format::Json, |config| config.format),
352        };
353        let input_mode = if format == serde::Format::Debezium {
354            SourceInputMode::KeyedUpsert
355        } else {
356            SourceInputMode::AppendOnly
357        };
358
359        // Neither Core NATS nor the current JetStream implementation can
360        // rewind an abandoned checkpoint attempt deterministically. Durable
361        // consumers alone are insufficient for LaminarDB replay semantics.
362        Ok(SourceContract::new(
363            SourceConsistency::Ephemeral,
364            SourceTopology::Singleton,
365            input_mode,
366        ))
367    }
368
369    async fn start(&mut self, request: SourceStart) -> Result<(), ConnectorError> {
370        let (config, position, _) = request.into_parts();
371        if let SourcePosition::Resume { attempt, .. } = position {
372            return Err(ConnectorError::ConfigurationError(format!(
373                "NATS is an ephemeral source and cannot resume checkpoint attempt {attempt:?}"
374            )));
375        }
376        let config = &config;
377
378        let cfg = NatsSourceConfig::from_config(config)?;
379        // Keep the candidate schema local until network admission succeeds so
380        // cancelling start leaves the existing instance unchanged.
381        let candidate_schema = config.arrow_schema();
382        let deserializer = serde::create_deserializer(cfg.format)
383            .map_err(|e| err(&format!("deserializer for format {:?}: {e}", cfg.format)))?;
384        match cfg.mode {
385            Mode::JetStream => self.open_jetstream(&cfg, deserializer).await?,
386            Mode::Core => self.open_core(&cfg, deserializer).await?,
387        }
388        if let Some(schema) = candidate_schema {
389            self.schema = schema;
390        }
391        self.config = Some(cfg);
392        Ok(())
393    }
394
395    async fn poll_batch(
396        &mut self,
397        max_records: usize,
398    ) -> Result<Option<SourceBatch>, ConnectorError> {
399        let Some(running) = self.running.as_mut() else {
400            return Ok(None);
401        };
402
403        let mut payloads: Vec<Bytes> = Vec::new();
404        let mut new_acks: Vec<jetstream::Message> = Vec::new();
405        let mut reader_disconnected = false;
406
407        while payloads.len() < max_records {
408            let incoming = match running
409                .rx
410                .as_mut()
411                .expect("running NATS source owns its receiver")
412                .try_recv()
413            {
414                Ok(m) => m,
415                Err(TryRecvError::Empty) => break,
416                Err(TryRecvError::Disconnected) => {
417                    reader_disconnected = true;
418                    break;
419                }
420            };
421            payloads.push(incoming.payload);
422            if let Some(msg) = incoming.ack {
423                new_acks.push(msg);
424            }
425        }
426
427        if payloads.is_empty() {
428            if reader_disconnected {
429                return Err(ConnectorError::ReadError(
430                    "NATS reader task terminated unexpectedly".into(),
431                ));
432            }
433            return Ok(None);
434        }
435
436        let records: Vec<&[u8]> = payloads.iter().map(Bytes::as_ref).collect();
437        let bytes_total: u64 = records.iter().map(|r| r.len() as u64).sum();
438        // Deserialize before scheduling acks: on failure the handles drop unacked and the broker
439        // redelivers after ack_wait. Ack enqueue is non-blocking to keep this poll path bounded.
440        let batch = running
441            .deserializer
442            .deserialize_batch(&records, &self.schema)
443            .map_err(|e| err(&format!("deserialize batch: {e}")))?;
444
445        enqueue_acks(running.ack_runtime.as_ref(), new_acks, &self.metrics);
446
447        self.metrics
448            .record_poll(batch.num_rows() as u64, bytes_total);
449
450        Ok(Some(SourceBatch::new(batch)))
451    }
452
453    fn schema(&self) -> SchemaRef {
454        self.schema.clone()
455    }
456
457    fn checkpoint(&self) -> SourceCheckpoint {
458        // Ephemeral sources deliberately expose no recovery cursor. JetStream acknowledgements
459        // are delivery progress, not checkpoint-owned state.
460        SourceCheckpoint::new()
461    }
462
463    async fn close(&mut self) -> Result<(), ConnectorError> {
464        let Some(mut running) = self.running.take() else {
465            return Ok(());
466        };
467        let close_deadline = tokio::time::Instant::now() + CLOSE_DRAIN_TIMEOUT;
468        running.request_shutdown();
469        match running.reader.wait_until(close_deadline).await {
470            TaskWait::Completed(Ok(())) => {}
471            TaskWait::Completed(Err(error)) => {
472                warn!(%error, "NATS reader task failed while closing");
473            }
474            TaskWait::TimedOut => warn!(
475                "NATS reader exceeded its close deadline; its tracked reaper retains shutdown ownership"
476            ),
477        }
478        // Drop unread messages only after the reader has stopped. Their unacked JetStream
479        // handles remain eligible for broker redelivery.
480        running.rx.take();
481
482        if let Some(mut ack_runtime) = running.ack_runtime.take() {
483            ack_runtime.request_shutdown();
484            match ack_runtime.task.wait_until(close_deadline).await {
485                TaskWait::Completed(Ok(())) => {}
486                TaskWait::Completed(Err(error)) => {
487                    warn!(%error, "NATS ack worker failed while closing");
488                    self.metrics.record_abandoned_acks();
489                }
490                TaskWait::TimedOut => {
491                    warn!(
492                        "NATS ack worker exceeded its close deadline; its tracked reaper retains shutdown ownership"
493                    );
494                }
495            }
496        }
497        Ok(())
498    }
499
500    fn data_ready_notify(&self) -> Option<Arc<Notify>> {
501        Some(Arc::clone(&self.data_ready))
502    }
503}
504
505// ── helpers ──
506
507fn err(msg: &str) -> ConnectorError {
508    ConnectorError::ConfigurationError(msg.to_string())
509}
510
511fn ack_runtime_limits(cfg: &NatsSourceConfig) -> (usize, usize) {
512    let broker_limit = usize::try_from(cfg.max_ack_pending)
513        .ok()
514        .filter(|limit| *limit > 0);
515    let fallback = cfg.fetch_batch.saturating_mul(2);
516    let backlog = broker_limit.unwrap_or(fallback).clamp(1, MAX_ACK_BACKLOG);
517    let concurrency = cfg.fetch_batch.clamp(1, MAX_ACK_CONCURRENCY).min(backlog);
518    (backlog, concurrency)
519}
520
521fn enqueue_acks(
522    runtime: Option<&AckRuntime>,
523    messages: Vec<jetstream::Message>,
524    metrics: &NatsSourceMetrics,
525) {
526    if messages.is_empty() {
527        return;
528    }
529    let Some(runtime) = runtime else {
530        metrics.record_ack_enqueue_errors(messages.len());
531        warn!(
532            rejected = messages.len(),
533            "JetStream ack worker is unavailable; broker will redeliver"
534        );
535        return;
536    };
537
538    let mut rejected = 0usize;
539    for message in messages {
540        // Increment before publication: a fast worker may complete immediately after try_send.
541        metrics.record_ack_enqueued();
542        if runtime
543            .tx
544            .as_ref()
545            .is_none_or(|tx| tx.try_send(message).is_err())
546        {
547            metrics.record_ack_error();
548            rejected += 1;
549        }
550    }
551    if rejected > 0 {
552        warn!(
553            rejected,
554            "JetStream ack backlog is full or closed; broker will redeliver"
555        );
556    }
557}
558
559async fn run_ack_worker(
560    rx: tokio_mpsc::Receiver<jetstream::Message>,
561    shutdown: watch::Receiver<bool>,
562    concurrency: usize,
563    metrics: NatsSourceMetrics,
564) {
565    // Ack calls stay scoped under the worker. Its single generation guard therefore proves that
566    // the receiver and every in-flight acknowledgement have all been dropped or completed.
567    let task_metrics = metrics.clone();
568    let abandoned = run_bounded_queue(rx, shutdown, concurrency, move |message| {
569        let metrics = task_metrics.clone();
570        async move {
571            acknowledge_message(message, &metrics).await;
572        }
573    })
574    .await;
575    if abandoned > 0 {
576        metrics.record_ack_abandoned(abandoned);
577        warn!(
578            abandoned,
579            "discarded queued JetStream acknowledgements during shutdown; broker will redeliver"
580        );
581    }
582}
583
584async fn run_bounded_queue<T, F, Fut>(
585    mut rx: tokio_mpsc::Receiver<T>,
586    mut shutdown: watch::Receiver<bool>,
587    concurrency: usize,
588    process: F,
589) -> usize
590where
591    T: Send + 'static,
592    F: Fn(T) -> Fut,
593    Fut: std::future::Future<Output = ()> + Send,
594{
595    debug_assert!(concurrency > 0);
596    let mut in_flight = FuturesUnordered::new();
597    'input: loop {
598        if shutdown_requested(&shutdown) || rx.is_closed() {
599            break;
600        }
601        while in_flight.len() >= concurrency {
602            tokio::select! {
603                biased;
604                _ = shutdown.changed() => break 'input,
605                _ = in_flight.next() => {}
606            }
607            if rx.is_closed() {
608                break 'input;
609            }
610        }
611
612        tokio::select! {
613            biased;
614            _ = shutdown.changed() => break,
615            _ = in_flight.next(), if !in_flight.is_empty() => {}
616            message = rx.recv() => {
617                let Some(message) = message else {
618                    break;
619                };
620                in_flight.push(process(message));
621            }
622        }
623    }
624
625    // Closing the receiver makes queued-but-unstarted messages immediately eligible for broker
626    // redelivery. Only work already admitted to `in_flight` is allowed to consume the close budget.
627    rx.close();
628    let mut abandoned = 0usize;
629    while rx.try_recv().is_ok() {
630        abandoned = abandoned.saturating_add(1);
631    }
632    while in_flight.next().await.is_some() {}
633    abandoned
634}
635
636async fn acknowledge_message(message: jetstream::Message, metrics: &NatsSourceMetrics) {
637    match tokio::time::timeout(ACK_IO_TIMEOUT, message.double_ack()).await {
638        Ok(Ok(())) => metrics.record_ack(),
639        Ok(Err(error)) => {
640            metrics.record_ack_error();
641            warn!(%error, "JetStream ack failed; broker will redeliver");
642        }
643        Err(_) => {
644            metrics.record_ack_error();
645            warn!(
646                timeout_ms = ACK_IO_TIMEOUT.as_millis(),
647                "JetStream ack timed out; broker will redeliver"
648            );
649        }
650    }
651}
652
653async fn connect(
654    cfg: &NatsSourceConfig,
655    owner: &ConnectorTaskOwner,
656) -> Result<async_nats::Client, ConnectorError> {
657    track_connection_tasks(build_connect_options(&cfg.auth, &cfg.tls)?, owner, "source")?
658        .connect(&cfg.servers)
659        .await
660        .map_err(|error| classify_connect_error(&error))
661}
662
663fn build_pull_config(
664    cfg: &NatsSourceConfig,
665    consumer_name: &str,
666) -> Result<pull::Config, ConnectorError> {
667    let filter_subjects = if cfg.subject_filters.is_empty() {
668        cfg.subject.iter().cloned().collect()
669    } else {
670        cfg.subject_filters.clone()
671    };
672
673    Ok(pull::Config {
674        durable_name: Some(consumer_name.to_string()),
675        filter_subjects,
676        deliver_policy: map_deliver_policy(cfg)?,
677        ack_policy: map_ack_policy(cfg.ack_policy),
678        ack_wait: cfg.ack_wait,
679        max_deliver: cfg.max_deliver,
680        max_ack_pending: cfg.max_ack_pending,
681        ..Default::default()
682    })
683}
684
685fn map_deliver_policy(
686    cfg: &NatsSourceConfig,
687) -> Result<async_nats::jetstream::consumer::DeliverPolicy, ConnectorError> {
688    use async_nats::jetstream::consumer::DeliverPolicy as Nats;
689    Ok(match cfg.deliver_policy {
690        DeliverPolicy::All => Nats::All,
691        DeliverPolicy::New => Nats::New,
692        DeliverPolicy::ByStartSequence => Nats::ByStartSequence {
693            start_sequence: cfg.start_sequence.unwrap_or(1),
694        },
695        DeliverPolicy::ByStartTime => {
696            let raw = cfg
697                .start_time
698                .as_deref()
699                .ok_or_else(|| err("deliver.policy=by_start_time requires 'start.time'"))?;
700            let start_time =
701                time::OffsetDateTime::parse(raw, &time::format_description::well_known::Rfc3339)
702                    .map_err(|e| err(&format!("start.time '{raw}' is not valid RFC3339: {e}")))?;
703            Nats::ByStartTime { start_time }
704        }
705    })
706}
707
708fn map_ack_policy(p: AckPolicy) -> async_nats::jetstream::consumer::AckPolicy {
709    use async_nats::jetstream::consumer::AckPolicy as Nats;
710    match p {
711        AckPolicy::Explicit => Nats::Explicit,
712        AckPolicy::None => Nats::None,
713    }
714}
715
716/// 500ms, 1s, 2s, 4s, cap 5s.
717fn fetch_backoff_base(consecutive_errors: u32) -> Duration {
718    let exp = consecutive_errors.saturating_sub(1).min(4);
719    let ms = 500u64.saturating_mul(1u64 << exp);
720    Duration::from_millis(ms.min(5000))
721}
722
723/// `base ± 20%`. Tests pass a fixed `entropy` seed.
724fn with_jitter(base: Duration, entropy: u64) -> Duration {
725    let base_ms = u64::try_from(base.as_millis()).unwrap_or(u64::MAX);
726    let range = (base_ms / 5).max(1); // 20%
727    let window = range * 2 + 1;
728    let offset = entropy % window;
729    let jittered = base_ms.saturating_add(offset).saturating_sub(range);
730    Duration::from_millis(jittered)
731}
732
733fn fetch_backoff(consecutive_errors: u32, entropy: u64) -> Duration {
734    with_jitter(fetch_backoff_base(consecutive_errors), entropy)
735}
736
737/// Wall-clock nanos for `with_jitter`. `Instant::now().elapsed()` is ~0
738/// and produces correlated jitter across tasks.
739#[allow(clippy::cast_possible_truncation)]
740fn entropy_now() -> u64 {
741    std::time::SystemTime::now()
742        .duration_since(std::time::UNIX_EPOCH)
743        .unwrap_or_default()
744        .as_nanos() as u64
745}
746
747fn shutdown_requested(shutdown: &watch::Receiver<bool>) -> bool {
748    *shutdown.borrow() || shutdown.has_changed().is_err()
749}
750
751struct JsReader {
752    consumer: jetstream::consumer::Consumer<pull::Config>,
753    tx: MAsyncTx<mpsc::Array<Incoming>>,
754    shutdown: watch::Receiver<bool>,
755    consecutive_errors: Arc<AtomicU32>,
756    data_ready: Arc<Notify>,
757    metrics: NatsSourceMetrics,
758    batch_size: usize,
759    max_wait: Duration,
760    /// `Duration::ZERO` disables the poll.
761    lag_poll_interval: Duration,
762    requires_ack: bool,
763}
764
765impl JsReader {
766    async fn run(self) {
767        let Self {
768            mut consumer,
769            tx,
770            mut shutdown,
771            consecutive_errors,
772            data_ready,
773            metrics,
774            batch_size,
775            max_wait,
776            lag_poll_interval,
777            requires_ack,
778        } = self;
779
780        let mut last_lag_poll = Instant::now();
781        let lag_poll_enabled = !lag_poll_interval.is_zero();
782
783        loop {
784            if shutdown_requested(&shutdown) {
785                break;
786            }
787            let fetch_result = tokio::select! {
788                biased;
789                _ = shutdown.changed() => break,
790                r = consumer.fetch().max_messages(batch_size).expires(max_wait).messages() => r,
791            };
792
793            let mut stream = match fetch_result {
794                Ok(s) => s,
795                Err(e) => {
796                    let errs = consecutive_errors.fetch_add(1, Ordering::AcqRel) + 1;
797                    metrics.record_fetch_error();
798                    warn!(
799                        error = %e,
800                        consecutive_errors = errs,
801                        "nats fetch() errored; backing off",
802                    );
803                    let backoff = fetch_backoff(errs, entropy_now());
804                    tokio::select! {
805                        biased;
806                        _ = shutdown.changed() => break,
807                        () = tokio::time::sleep(backoff) => {}
808                    }
809                    continue;
810                }
811            };
812
813            let mut forwarded = 0usize;
814            let mut stream_errors = 0usize;
815            loop {
816                let msg_result = tokio::select! {
817                    biased;
818                    _ = shutdown.changed() => return,
819                    r = stream.next() => match r {
820                        Some(r) => r,
821                        None => break,
822                    },
823                };
824                let msg = match msg_result {
825                    Ok(m) => m,
826                    Err(e) => {
827                        metrics.record_fetch_error();
828                        stream_errors += 1;
829                        warn!(error = %e, "nats message error");
830                        continue;
831                    }
832                };
833                let payload = msg.payload.clone();
834                let incoming = Incoming {
835                    payload,
836                    ack: requires_ack.then_some(msg),
837                };
838                let send_result = tokio::select! {
839                    biased;
840                    _ = shutdown.changed() => return,
841                    result = tx.send(incoming) => result,
842                };
843                if send_result.is_err() {
844                    debug!("nats reader: downstream channel closed");
845                    return;
846                }
847                forwarded += 1;
848            }
849
850            // Reset on progress; an iteration with only errors counts
851            // as one failure; idle iterations don't bump.
852            if forwarded > 0 {
853                consecutive_errors.store(0, Ordering::Release);
854                data_ready.notify_one();
855            } else if stream_errors > 0 {
856                let errs = consecutive_errors.fetch_add(1, Ordering::AcqRel) + 1;
857                let backoff = fetch_backoff(errs, entropy_now());
858                tokio::select! {
859                    biased;
860                    _ = shutdown.changed() => break,
861                    () = tokio::time::sleep(backoff) => {}
862                }
863            }
864
865            if lag_poll_enabled && last_lag_poll.elapsed() >= lag_poll_interval {
866                last_lag_poll = Instant::now();
867                match consumer.info().await {
868                    Ok(info) => metrics.set_consumer_lag(info.num_pending),
869                    Err(e) => warn!(error = %e, "consumer.info() failed; skipping lag update"),
870                }
871            }
872        }
873    }
874}
875
876struct CoreReader {
877    subscriber: async_nats::Subscriber,
878    tx: MAsyncTx<mpsc::Array<Incoming>>,
879    shutdown: watch::Receiver<bool>,
880    data_ready: Arc<Notify>,
881}
882
883impl CoreReader {
884    async fn run(self) {
885        let Self {
886            mut subscriber,
887            tx,
888            mut shutdown,
889            data_ready,
890        } = self;
891
892        loop {
893            if shutdown_requested(&shutdown) {
894                break;
895            }
896            let msg = tokio::select! {
897                biased;
898                _ = shutdown.changed() => break,
899                m = subscriber.next() => match m {
900                    Some(m) => m,
901                    None => break,
902                },
903            };
904            let incoming = Incoming {
905                payload: msg.payload,
906                ack: None,
907            };
908            let send_result = tokio::select! {
909                biased;
910                _ = shutdown.changed() => break,
911                result = tx.send(incoming) => result,
912            };
913            if send_result.is_err() {
914                break;
915            }
916            data_ready.notify_one();
917        }
918        // Wake the coordinator so it observes the now-disconnected channel immediately instead
919        // of treating a terminated Core subscription as an indefinitely idle source.
920        data_ready.notify_one();
921    }
922}
923
924#[cfg(test)]
925mod tests;