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, 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
40struct 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 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 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
190pub 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 #[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 #[must_use]
219 pub fn config(&self) -> Option<&NatsSourceConfig> {
220 self.config.as_ref()
221 }
222
223 #[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 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 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 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 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 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
489fn 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 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 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 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
700fn 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
707fn 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); 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#[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 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 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 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}