1use 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
41struct 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 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 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
191pub 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 #[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 #[must_use]
220 pub fn config(&self) -> Option<&NatsSourceConfig> {
221 self.config.as_ref()
222 }
223
224 #[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 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 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 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 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 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
505fn 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 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 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 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
716fn 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
723fn 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); 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#[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 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 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 data_ready.notify_one();
921 }
922}
923
924#[cfg(test)]
925mod tests;