Skip to main content

laminar_connectors/
testing.rs

1//! Testing utilities for connector implementations.
2//!
3//! Provides mock connectors and helper functions for testing
4//! the connector SDK and concrete connector implementations.
5
6use std::sync::atomic::{AtomicU64, Ordering};
7use std::sync::Arc;
8use std::time::Duration;
9
10use arrow_array::{Int64Array, RecordBatch, StringArray};
11use arrow_schema::{DataType, Field, Schema, SchemaRef};
12use async_trait::async_trait;
13use parking_lot::Mutex;
14
15use crate::checkpoint::SourceCheckpoint;
16use crate::config::{ConnectorConfig, ConnectorInfo};
17use crate::connector::{
18    SinkConnector, SinkConsistency, SinkContract, SinkInputMode, SinkTopology, SourceBatch,
19    SourceConnector, WriteResult,
20};
21use crate::error::ConnectorError;
22use crate::registry::ConnectorRegistry;
23
24/// Creates a test schema with `id` (Int64) and `value` (Utf8) columns.
25#[must_use]
26pub fn mock_schema() -> SchemaRef {
27    Arc::new(Schema::new(vec![
28        Field::new("id", DataType::Int64, false),
29        Field::new("value", DataType::Utf8, false),
30    ]))
31}
32
33/// Creates a test `RecordBatch` with `n` rows.
34///
35/// # Panics
36///
37/// Panics if the batch cannot be created (should not happen with valid inputs).
38#[must_use]
39pub fn mock_batch(n: usize) -> RecordBatch {
40    #[allow(clippy::cast_possible_wrap)]
41    let ids: Vec<i64> = (0..n as i64).collect();
42    let values: Vec<String> = (0..n).map(|i| format!("value_{i}")).collect();
43    let value_refs: Vec<&str> = values.iter().map(String::as_str).collect();
44
45    RecordBatch::try_new(
46        mock_schema(),
47        vec![
48            Arc::new(Int64Array::from(ids)),
49            Arc::new(StringArray::from(value_refs)),
50        ],
51    )
52    .unwrap()
53}
54
55/// Mock source connector for testing.
56///
57/// Returns a configurable number of batches, then returns `None`.
58#[derive(Debug)]
59pub struct MockSourceConnector {
60    schema: SchemaRef,
61    initial_batches: u64,
62    batches_remaining: AtomicU64,
63    batch_size: usize,
64    records_produced: AtomicU64,
65    is_open: std::sync::atomic::AtomicBool,
66    committed_epochs: Arc<Mutex<Vec<u64>>>,
67}
68
69impl MockSourceConnector {
70    /// Creates a new mock source that produces 10 batches of 5 records.
71    #[must_use]
72    pub fn new() -> Self {
73        Self {
74            schema: mock_schema(),
75            initial_batches: 10,
76            batches_remaining: AtomicU64::new(10),
77            batch_size: 5,
78            records_produced: AtomicU64::new(0),
79            is_open: std::sync::atomic::AtomicBool::new(false),
80            committed_epochs: Arc::new(Mutex::new(Vec::new())),
81        }
82    }
83
84    /// Creates a mock source with custom batch count and size.
85    #[must_use]
86    pub fn with_batches(count: u64, batch_size: usize) -> Self {
87        Self {
88            schema: mock_schema(),
89            initial_batches: count,
90            batches_remaining: AtomicU64::new(count),
91            batch_size,
92            records_produced: AtomicU64::new(0),
93            is_open: std::sync::atomic::AtomicBool::new(false),
94            committed_epochs: Arc::new(Mutex::new(Vec::new())),
95        }
96    }
97
98    /// Returns the total records produced.
99    #[must_use]
100    pub fn records_produced(&self) -> u64 {
101        self.records_produced.load(Ordering::Relaxed)
102    }
103
104    /// Returns a handle for inspecting epochs reported to
105    /// `notify_epoch_committed` from outside the connector (the
106    /// connector itself is moved into a background task when run via
107    /// the pipeline).
108    #[must_use]
109    pub fn committed_epochs_handle(&self) -> Arc<Mutex<Vec<u64>>> {
110        Arc::clone(&self.committed_epochs)
111    }
112}
113
114impl Default for MockSourceConnector {
115    fn default() -> Self {
116        Self::new()
117    }
118}
119
120#[async_trait]
121impl SourceConnector for MockSourceConnector {
122    async fn start(
123        &mut self,
124        request: crate::connector::SourceStart,
125    ) -> Result<(), ConnectorError> {
126        let records = match request.into_parts().1 {
127            crate::connector::SourcePosition::Initial => 0,
128            crate::connector::SourcePosition::Resume { checkpoint, .. } => checkpoint
129                .get_offset("records")
130                .ok_or_else(|| {
131                    ConnectorError::ConfigurationError(
132                        "mock source checkpoint is missing 'records'".into(),
133                    )
134                })?
135                .parse::<u64>()
136                .map_err(|error| {
137                    ConnectorError::ConfigurationError(format!(
138                        "invalid mock source record cursor: {error}"
139                    ))
140                })?,
141        };
142        let consumed_batches = if self.batch_size == 0 {
143            0
144        } else {
145            records / self.batch_size as u64
146        };
147        self.records_produced.store(records, Ordering::Relaxed);
148        self.batches_remaining.store(
149            self.initial_batches.saturating_sub(consumed_batches),
150            Ordering::Relaxed,
151        );
152        self.is_open
153            .store(true, std::sync::atomic::Ordering::Relaxed);
154        Ok(())
155    }
156
157    async fn poll_batch(
158        &mut self,
159        _max_records: usize,
160    ) -> Result<Option<SourceBatch>, ConnectorError> {
161        let remaining = self.batches_remaining.load(Ordering::Relaxed);
162        if remaining == 0 {
163            return Ok(None);
164        }
165        self.batches_remaining.fetch_sub(1, Ordering::Relaxed);
166
167        let batch = mock_batch(self.batch_size);
168        self.records_produced
169            .fetch_add(self.batch_size as u64, Ordering::Relaxed);
170        Ok(Some(SourceBatch::new(batch)))
171    }
172
173    fn schema(&self) -> SchemaRef {
174        self.schema.clone()
175    }
176
177    fn checkpoint(&self) -> SourceCheckpoint {
178        let mut cp = SourceCheckpoint::new();
179        cp.set_offset(
180            "records",
181            self.records_produced.load(Ordering::Relaxed).to_string(),
182        );
183        cp
184    }
185
186    async fn close(&mut self) -> Result<(), ConnectorError> {
187        self.is_open
188            .store(false, std::sync::atomic::Ordering::Relaxed);
189        Ok(())
190    }
191
192    async fn notify_epoch_committed(
193        &mut self,
194        epoch: u64,
195        _checkpoint: &SourceCheckpoint,
196    ) -> Result<(), ConnectorError> {
197        self.committed_epochs.lock().push(epoch);
198        Ok(())
199    }
200}
201
202/// Mock sink connector for testing.
203///
204/// Stores all written batches in memory for inspection.
205#[derive(Debug)]
206pub struct MockSinkConnector {
207    schema: SchemaRef,
208    written: Arc<Mutex<Vec<RecordBatch>>>,
209    records_written: AtomicU64,
210    is_open: std::sync::atomic::AtomicBool,
211}
212
213impl MockSinkConnector {
214    /// Creates a new mock sink.
215    #[must_use]
216    pub fn new() -> Self {
217        Self {
218            schema: mock_schema(),
219            written: Arc::new(Mutex::new(Vec::new())),
220            records_written: AtomicU64::new(0),
221            is_open: std::sync::atomic::AtomicBool::new(false),
222        }
223    }
224
225    /// Returns the number of batches written.
226    #[must_use]
227    pub fn batch_count(&self) -> usize {
228        self.written.lock().len()
229    }
230
231    /// Returns the total number of records written.
232    #[must_use]
233    pub fn records_written(&self) -> u64 {
234        self.records_written.load(Ordering::Relaxed)
235    }
236
237    /// Returns a clone of all written batches.
238    #[must_use]
239    pub fn written_batches(&self) -> Vec<RecordBatch> {
240        self.written.lock().clone()
241    }
242}
243
244impl Default for MockSinkConnector {
245    fn default() -> Self {
246        Self::new()
247    }
248}
249
250#[async_trait]
251impl SinkConnector for MockSinkConnector {
252    fn contract(&self, _config: &ConnectorConfig) -> Result<SinkContract, ConnectorError> {
253        Ok(SinkContract::new(
254            SinkConsistency::Ephemeral,
255            SinkTopology::NodeLocalEgress,
256            SinkInputMode::AppendOnly,
257        ))
258    }
259
260    async fn open(&mut self, _config: &ConnectorConfig) -> Result<(), ConnectorError> {
261        self.is_open
262            .store(true, std::sync::atomic::Ordering::Relaxed);
263        Ok(())
264    }
265
266    async fn write_batch(&mut self, batch: &RecordBatch) -> Result<WriteResult, ConnectorError> {
267        let num_rows = batch.num_rows();
268        let bytes = batch.get_array_memory_size() as u64;
269        self.written.lock().push(batch.clone());
270        self.records_written
271            .fetch_add(num_rows as u64, Ordering::Relaxed);
272        Ok(WriteResult::new(num_rows, bytes))
273    }
274
275    fn schema(&self) -> SchemaRef {
276        self.schema.clone()
277    }
278
279    fn suggested_write_timeout(&self) -> Duration {
280        Duration::from_secs(60)
281    }
282
283    async fn close(&mut self) -> Result<(), ConnectorError> {
284        self.is_open
285            .store(false, std::sync::atomic::Ordering::Relaxed);
286        Ok(())
287    }
288}
289
290/// Registers a mock source connector with the registry.
291///
292/// # Errors
293/// Returns an error when the mock source name is already registered.
294pub fn register_mock_source(registry: &ConnectorRegistry) -> Result<(), ConnectorError> {
295    registry.register_source(
296        "mock",
297        ConnectorInfo {
298            name: "mock".to_string(),
299            display_name: "Mock Source".to_string(),
300            version: "0.1.0".to_string(),
301            is_source: true,
302            is_sink: false,
303            config_keys: vec![],
304        },
305        Arc::new(|_: Option<&Arc<prometheus::Registry>>| Ok(Box::new(MockSourceConnector::new()))),
306    )
307}
308
309/// Registers a mock sink connector with the registry.
310///
311/// # Errors
312/// Returns an error when the mock sink name is already registered.
313pub fn register_mock_sink(registry: &ConnectorRegistry) -> Result<(), ConnectorError> {
314    registry.register_sink(
315        "mock",
316        ConnectorInfo {
317            name: "mock".to_string(),
318            display_name: "Mock Sink".to_string(),
319            version: "0.1.0".to_string(),
320            is_source: false,
321            is_sink: true,
322            config_keys: vec![],
323        },
324        Arc::new(|_config, _registry| Ok(Box::new(MockSinkConnector::new()))),
325    )
326}
327
328#[cfg(test)]
329mod tests {
330    use super::*;
331
332    #[test]
333    fn test_mock_batch() {
334        let batch = mock_batch(10);
335        assert_eq!(batch.num_rows(), 10);
336        assert_eq!(batch.num_columns(), 2);
337    }
338
339    #[tokio::test]
340    async fn test_mock_source_connector() {
341        let mut source = MockSourceConnector::with_batches(3, 5);
342        source
343            .start(
344                crate::connector::SourceStart::new(
345                    ConnectorConfig::new("mock"),
346                    crate::connector::SourcePosition::Initial,
347                    crate::connector::DeliveryGuarantee::BestEffort,
348                )
349                .unwrap(),
350            )
351            .await
352            .unwrap();
353
354        let b1 = source.poll_batch(100).await.unwrap();
355        assert!(b1.is_some());
356        assert_eq!(b1.unwrap().num_rows(), 5);
357
358        let b2 = source.poll_batch(100).await.unwrap();
359        assert!(b2.is_some());
360
361        let b3 = source.poll_batch(100).await.unwrap();
362        assert!(b3.is_some());
363
364        let b4 = source.poll_batch(100).await.unwrap();
365        assert!(b4.is_none());
366
367        assert_eq!(source.records_produced(), 15);
368
369        let cp = source.checkpoint();
370        assert_eq!(cp.get_offset("records"), Some("15"));
371
372        source.close().await.unwrap();
373    }
374
375    #[tokio::test]
376    async fn test_mock_sink_connector() {
377        let mut sink = MockSinkConnector::new();
378        sink.open(&ConnectorConfig::new("mock")).await.unwrap();
379
380        let batch = mock_batch(10);
381        let result = sink.write_batch(&batch).await.unwrap();
382        assert_eq!(result.records_written, 10);
383
384        assert_eq!(sink.batch_count(), 1);
385        assert_eq!(sink.records_written(), 10);
386
387        sink.write_batch(&mock_batch(5)).await.unwrap();
388
389        assert_eq!(sink.records_written(), 15);
390        assert_eq!(sink.batch_count(), 2);
391
392        let contract = sink.contract(&ConnectorConfig::new("mock")).unwrap();
393        assert_eq!(contract.consistency, SinkConsistency::Ephemeral);
394        assert_eq!(contract.topology, SinkTopology::NodeLocalEgress);
395        assert_eq!(contract.input_mode, SinkInputMode::AppendOnly);
396        assert_eq!(sink.suggested_write_timeout(), Duration::from_secs(60));
397
398        sink.close().await.unwrap();
399    }
400
401    #[tokio::test]
402    async fn test_mock_source_resume() {
403        let mut source = MockSourceConnector::new();
404        let mut checkpoint = SourceCheckpoint::new();
405        checkpoint.set_offset("records", "10");
406        source
407            .start(
408                crate::connector::SourceStart::new(
409                    ConnectorConfig::new("mock"),
410                    crate::connector::SourcePosition::Resume {
411                        attempt: laminar_core::state::CheckpointAttempt::new(5, 5),
412                        checkpoint,
413                    },
414                    crate::connector::DeliveryGuarantee::AtLeastOnce,
415                )
416                .unwrap(),
417            )
418            .await
419            .unwrap();
420        assert_eq!(source.records_produced(), 10);
421    }
422
423    #[test]
424    fn test_register_helpers() {
425        let registry = ConnectorRegistry::new();
426        register_mock_source(&registry).unwrap();
427        register_mock_sink(&registry).unwrap();
428
429        assert!(registry.source_info("mock").is_some());
430        assert!(registry.sink_info("mock").is_some());
431    }
432}