Skip to main content

lancedb/utils/
mod.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The LanceDB Authors
3
4pub(crate) mod background_cache;
5
6use std::sync::Arc;
7
8use arrow_array::RecordBatch;
9use arrow_schema::{DataType, Field, Schema, SchemaRef};
10use datafusion_common::{DataFusionError, Result as DataFusionResult};
11use datafusion_execution::RecordBatchStream;
12use futures::{FutureExt, Stream};
13use lance::arrow::json::JsonDataType;
14use lance::dataset::{ReadParams, WriteParams};
15use lance::index::vector::utils::infer_vector_dim;
16use lance::io::{ObjectStoreParams, WrappingObjectStore};
17use std::pin::Pin;
18
19use crate::error::{Error, Result};
20use datafusion_physical_plan::SendableRecordBatchStream;
21
22static TABLE_NAME_REGEX: std::sync::LazyLock<regex::Regex> =
23    std::sync::LazyLock::new(|| regex::Regex::new(r"^[a-zA-Z0-9_\-\.]+$").unwrap());
24static NAMESPACE_NAME_REGEX: std::sync::LazyLock<regex::Regex> =
25    std::sync::LazyLock::new(|| regex::Regex::new(r"^[a-zA-Z0-9_\-\.]+$").unwrap());
26
27pub trait PatchStoreParam {
28    fn patch_with_store_wrapper(
29        self,
30        wrapper: Arc<dyn WrappingObjectStore>,
31    ) -> Result<Option<ObjectStoreParams>>;
32}
33
34impl PatchStoreParam for Option<ObjectStoreParams> {
35    fn patch_with_store_wrapper(
36        self,
37        wrapper: Arc<dyn WrappingObjectStore>,
38    ) -> Result<Option<ObjectStoreParams>> {
39        let mut params = self.unwrap_or_default();
40        if params.object_store_wrapper.is_some() {
41            return Err(Error::Other {
42                message: "can not patch param because object store is already set".into(),
43                source: None,
44            });
45        }
46        params.object_store_wrapper = Some(wrapper);
47
48        Ok(Some(params))
49    }
50}
51
52pub trait PatchWriteParam {
53    fn patch_with_store_wrapper(self, wrapper: Arc<dyn WrappingObjectStore>)
54    -> Result<WriteParams>;
55}
56
57impl PatchWriteParam for WriteParams {
58    fn patch_with_store_wrapper(
59        mut self,
60        wrapper: Arc<dyn WrappingObjectStore>,
61    ) -> Result<WriteParams> {
62        self.store_params = self.store_params.patch_with_store_wrapper(wrapper)?;
63        Ok(self)
64    }
65}
66
67// NOTE: we have some API inconsistency here.
68// WriteParam is found in the form of Option<WriteParam> and ReadParam is found in the form of ReadParam
69
70pub trait PatchReadParam {
71    fn patch_with_store_wrapper(self, wrapper: Arc<dyn WrappingObjectStore>) -> Result<ReadParams>;
72}
73
74impl PatchReadParam for ReadParams {
75    fn patch_with_store_wrapper(
76        mut self,
77        wrapper: Arc<dyn WrappingObjectStore>,
78    ) -> Result<ReadParams> {
79        self.store_options = self.store_options.patch_with_store_wrapper(wrapper)?;
80        Ok(self)
81    }
82}
83
84/// Validate table name.
85pub fn validate_table_name(name: &str) -> Result<()> {
86    if name.is_empty() {
87        return Err(Error::InvalidTableName {
88            name: name.to_string(),
89            reason: "Table names cannot be empty strings".to_string(),
90        });
91    }
92    if !TABLE_NAME_REGEX.is_match(name) {
93        return Err(Error::InvalidTableName {
94            name: name.to_string(),
95            reason:
96                "Table names can only contain alphanumeric characters, underscores, hyphens, and periods"
97                    .to_string(),
98        });
99    }
100    Ok(())
101}
102
103/// Validate a namespace name component
104///
105/// Namespace names must:
106/// - Not be empty
107/// - Only contain alphanumeric characters, underscores, hyphens, and periods
108///
109/// # Arguments
110/// * `name` - A single namespace component (not the full path)
111///
112/// # Returns
113/// * `Ok(())` if the namespace name is valid
114/// * `Err(Error)` if the namespace name is invalid
115pub fn validate_namespace_name(name: &str) -> Result<()> {
116    if name.is_empty() {
117        return Err(Error::InvalidInput {
118            message: "Namespace names cannot be empty strings".to_string(),
119        });
120    }
121    if !NAMESPACE_NAME_REGEX.is_match(name) {
122        return Err(Error::InvalidInput {
123            message: format!(
124                "Invalid namespace name '{}': Namespace names can only contain alphanumeric characters, underscores, hyphens, and periods",
125                name
126            ),
127        });
128    }
129    Ok(())
130}
131
132/// Validate all components of a namespace
133///
134/// Iterates through all namespace components and validates each one.
135/// Returns an error if any component is invalid.
136///
137/// # Arguments
138/// * `namespace` - The namespace components to validate
139///
140/// # Returns
141/// * `Ok(())` if all namespace components are valid
142/// * `Err(Error)` if any component is invalid
143pub fn validate_namespace(namespace: &[String]) -> Result<()> {
144    for component in namespace {
145        validate_namespace_name(component)?;
146    }
147    Ok(())
148}
149
150/// Find one default column to create index or perform vector query.
151pub(crate) fn default_vector_column(schema: &Schema, dim: Option<i32>) -> Result<String> {
152    // Try to find a vector column.
153    let mut candidates = Vec::new();
154    for field in schema.fields() {
155        collect_vector_columns(field, &mut Vec::new(), dim, &mut candidates);
156    }
157    if candidates.is_empty() {
158        Err(Error::InvalidInput {
159            message: format!(
160                "No vector column found to match with the query vector dimension: {}",
161                dim.unwrap_or_default()
162            ),
163        })
164    } else if candidates.len() != 1 {
165        Err(Error::Schema {
166            message: format!(
167                "More than one vector columns found, \
168                    please specify which column to create index or query: {:?}",
169                candidates
170            ),
171        })
172    } else {
173        Ok(candidates[0].clone())
174    }
175}
176
177fn collect_vector_columns(
178    field: &Field,
179    path: &mut Vec<String>,
180    dim: Option<i32>,
181    candidates: &mut Vec<String>,
182) {
183    path.push(field.name().clone());
184    match infer_vector_dim(field.data_type()) {
185        Ok(d) if dim.is_none() || dim == Some(d as i32) => {
186            let path_segments = path.iter().map(String::as_str).collect::<Vec<_>>();
187            candidates.push(lance_core::datatypes::format_field_path(&path_segments));
188        }
189        _ => {
190            if let DataType::Struct(fields) = field.data_type() {
191                for child in fields {
192                    collect_vector_columns(child, path, dim, candidates);
193                }
194            }
195        }
196    }
197    path.pop();
198}
199
200pub(crate) fn resolve_arrow_field_path(schema: &Schema, column: &str) -> Result<(String, Field)> {
201    lance_core::datatypes::parse_field_path(column).map_err(|e| Error::InvalidInput {
202        message: format!("Invalid field path `{}`: {}", column, e),
203    })?;
204
205    let lance_schema =
206        lance_core::datatypes::Schema::try_from(schema).map_err(|e| Error::Schema {
207            message: format!("Invalid schema: {}", e),
208        })?;
209    let field_path = lance_schema
210        .resolve_case_insensitive(column)
211        .ok_or_else(|| Error::Schema {
212            message: format!(
213                "Field path `{}` not found in schema. Available field paths: {}",
214                column,
215                lance_schema.field_paths().join(", ")
216            ),
217        })?;
218    let field = field_path.last().expect("field path should be non-empty");
219    let path_segments = field_path
220        .iter()
221        .map(|field| field.name.as_str())
222        .collect::<Vec<_>>();
223    let canonical_path = lance_core::datatypes::format_field_path(&path_segments);
224
225    Ok((canonical_path, Field::from(*field)))
226}
227
228pub fn supported_btree_data_type(dtype: &DataType) -> bool {
229    dtype.is_integer()
230        || dtype.is_floating()
231        || matches!(
232            dtype,
233            DataType::Boolean
234                | DataType::Utf8
235                | DataType::Time32(_)
236                | DataType::Time64(_)
237                | DataType::Date32
238                | DataType::Date64
239                | DataType::Timestamp(_, _)
240                | DataType::FixedSizeBinary(_)
241        )
242}
243
244pub fn supported_bitmap_data_type(dtype: &DataType) -> bool {
245    dtype.is_integer()
246        || matches!(
247            dtype,
248            DataType::Utf8
249                | DataType::LargeUtf8
250                | DataType::Binary
251                | DataType::LargeBinary
252                | DataType::Boolean
253        )
254}
255
256pub fn supported_label_list_data_type(dtype: &DataType) -> bool {
257    match dtype {
258        DataType::List(field) | DataType::LargeList(field) => {
259            supported_bitmap_data_type(field.data_type())
260        }
261        DataType::FixedSizeList(field, _) => supported_bitmap_data_type(field.data_type()),
262        _ => false,
263    }
264}
265
266pub fn supported_fts_data_type(dtype: &DataType) -> bool {
267    supported_fts_data_type_impl(dtype, false)
268}
269
270fn supported_fts_data_type_impl(dtype: &DataType, in_list: bool) -> bool {
271    match (dtype, in_list) {
272        (DataType::Utf8 | DataType::LargeUtf8, _) => true,
273        (DataType::List(field) | DataType::LargeList(field), false) => {
274            supported_fts_data_type_impl(field.data_type(), true)
275        }
276        _ => false,
277    }
278}
279
280/// FM-Index accelerates substring (`contains`) search over raw bytes, so it
281/// applies to string and binary columns.
282pub fn supported_fm_data_type(dtype: &DataType) -> bool {
283    matches!(
284        dtype,
285        DataType::Utf8 | DataType::LargeUtf8 | DataType::Binary | DataType::LargeBinary
286    )
287}
288
289pub fn supported_vector_data_type(dtype: &DataType) -> bool {
290    match dtype {
291        DataType::FixedSizeList(field, _) => {
292            field.data_type().is_floating() || field.data_type() == &DataType::UInt8
293        }
294        DataType::List(field) => supported_vector_data_type(field.data_type()),
295        _ => false,
296    }
297}
298
299/// Note: this is temporary until we get a proper datatype conversion in Lance.
300pub fn string_to_datatype(s: &str) -> Option<DataType> {
301    let data_type: serde_json::Value = {
302        if let Ok(data_type) = serde_json::from_str(s) {
303            data_type
304        } else {
305            serde_json::json!({ "type": s })
306        }
307    };
308    let json_type: JsonDataType = serde_json::from_value(data_type).ok()?;
309    (&json_type).try_into().ok()
310}
311
312enum TimeoutState {
313    NotStarted {
314        timeout: std::time::Duration,
315    },
316    Started {
317        deadline: Pin<Box<tokio::time::Sleep>>,
318        timeout: std::time::Duration,
319    },
320    Completed,
321}
322
323/// A `Stream` wrapper that implements a timeout.
324///
325/// The timeout starts when the first `poll_next` is called. As soon as the timeout
326/// duration has passed, the stream will return an `Err` indicating a timeout error
327/// for the next poll.
328pub struct TimeoutStream {
329    inner: SendableRecordBatchStream,
330    state: TimeoutState,
331}
332
333impl TimeoutStream {
334    pub fn new(inner: SendableRecordBatchStream, timeout: std::time::Duration) -> Self {
335        Self {
336            inner,
337            state: TimeoutState::NotStarted { timeout },
338        }
339    }
340
341    pub fn new_boxed(
342        inner: SendableRecordBatchStream,
343        timeout: std::time::Duration,
344    ) -> SendableRecordBatchStream {
345        Box::pin(Self::new(inner, timeout))
346    }
347
348    fn timeout_error(timeout: &std::time::Duration) -> DataFusionError {
349        DataFusionError::Execution(format!("Query timeout after {} ms", timeout.as_millis()))
350    }
351}
352
353impl RecordBatchStream for TimeoutStream {
354    fn schema(&self) -> SchemaRef {
355        self.inner.schema()
356    }
357}
358
359impl Stream for TimeoutStream {
360    type Item = DataFusionResult<RecordBatch>;
361
362    fn poll_next(
363        mut self: std::pin::Pin<&mut Self>,
364        cx: &mut std::task::Context<'_>,
365    ) -> std::task::Poll<Option<Self::Item>> {
366        match &mut self.state {
367            TimeoutState::NotStarted { timeout } => {
368                if timeout.is_zero() {
369                    return std::task::Poll::Ready(Some(Err(Self::timeout_error(timeout))));
370                }
371                let deadline = Box::pin(tokio::time::sleep(*timeout));
372                self.state = TimeoutState::Started {
373                    deadline,
374                    timeout: *timeout,
375                };
376                self.poll_next(cx)
377            }
378            TimeoutState::Started { deadline, timeout } => match deadline.poll_unpin(cx) {
379                std::task::Poll::Ready(_) => {
380                    let err = Self::timeout_error(timeout);
381                    self.state = TimeoutState::Completed;
382                    std::task::Poll::Ready(Some(Err(err)))
383                }
384                std::task::Poll::Pending => {
385                    let inner = Pin::new(&mut self.inner);
386                    inner.poll_next(cx)
387                }
388            },
389            TimeoutState::Completed => std::task::Poll::Ready(None),
390        }
391    }
392}
393
394/// A `Stream` wrapper that slices oversized batches to enforce a maximum batch length.
395pub struct MaxBatchLengthStream {
396    inner: SendableRecordBatchStream,
397    max_batch_length: Option<usize>,
398    buffered_batch: Option<RecordBatch>,
399    buffered_offset: usize,
400}
401
402impl MaxBatchLengthStream {
403    pub fn new(inner: SendableRecordBatchStream, max_batch_length: usize) -> Self {
404        Self {
405            inner,
406            max_batch_length: (max_batch_length > 0).then_some(max_batch_length),
407            buffered_batch: None,
408            buffered_offset: 0,
409        }
410    }
411
412    pub fn new_boxed(
413        inner: SendableRecordBatchStream,
414        max_batch_length: usize,
415    ) -> SendableRecordBatchStream {
416        if max_batch_length == 0 {
417            inner
418        } else {
419            Box::pin(Self::new(inner, max_batch_length))
420        }
421    }
422}
423
424impl RecordBatchStream for MaxBatchLengthStream {
425    fn schema(&self) -> SchemaRef {
426        self.inner.schema()
427    }
428}
429
430impl Stream for MaxBatchLengthStream {
431    type Item = DataFusionResult<RecordBatch>;
432
433    fn poll_next(
434        mut self: Pin<&mut Self>,
435        cx: &mut std::task::Context<'_>,
436    ) -> std::task::Poll<Option<Self::Item>> {
437        loop {
438            let Some(max_batch_length) = self.max_batch_length else {
439                return Pin::new(&mut self.inner).poll_next(cx);
440            };
441
442            if let Some(batch) = self.buffered_batch.clone() {
443                if self.buffered_offset < batch.num_rows() {
444                    let remaining = batch.num_rows() - self.buffered_offset;
445                    let length = remaining.min(max_batch_length);
446                    let sliced = batch.slice(self.buffered_offset, length);
447                    self.buffered_offset += length;
448                    if self.buffered_offset >= batch.num_rows() {
449                        self.buffered_batch = None;
450                        self.buffered_offset = 0;
451                    }
452                    return std::task::Poll::Ready(Some(Ok(sliced)));
453                }
454
455                self.buffered_batch = None;
456                self.buffered_offset = 0;
457            }
458
459            match Pin::new(&mut self.inner).poll_next(cx) {
460                std::task::Poll::Ready(Some(Ok(batch))) => {
461                    if batch.num_rows() <= max_batch_length {
462                        return std::task::Poll::Ready(Some(Ok(batch)));
463                    }
464                    self.buffered_batch = Some(batch);
465                    self.buffered_offset = 0;
466                }
467                other => return other,
468            }
469        }
470    }
471}
472
473#[cfg(test)]
474mod tests {
475    use arrow_array::Int32Array;
476    use arrow_schema::Field;
477    use datafusion_physical_plan::stream::RecordBatchStreamAdapter;
478    use futures::{StreamExt, stream};
479    use tokio::time::sleep;
480
481    use super::*;
482
483    #[test]
484    fn test_guess_default_column() {
485        let schema_no_vector = Schema::new(vec![
486            Field::new("id", DataType::Int16, true),
487            Field::new("tag", DataType::Utf8, false),
488        ]);
489        assert!(
490            default_vector_column(&schema_no_vector, None)
491                .unwrap_err()
492                .to_string()
493                .contains("No vector column")
494        );
495
496        let schema_with_vec_col = Schema::new(vec![
497            Field::new("id", DataType::Int16, true),
498            Field::new(
499                "vec",
500                DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float64, false)), 10),
501                false,
502            ),
503        ]);
504        assert_eq!(
505            default_vector_column(&schema_with_vec_col, None).unwrap(),
506            "vec"
507        );
508
509        let schema_with_nested_vec_col = Schema::new(vec![
510            Field::new("id", DataType::Int16, true),
511            Field::new(
512                "image",
513                DataType::Struct(
514                    vec![Field::new(
515                        "embedding",
516                        DataType::FixedSizeList(
517                            Arc::new(Field::new("item", DataType::Float32, false)),
518                            10,
519                        ),
520                        false,
521                    )]
522                    .into(),
523                ),
524                false,
525            ),
526        ]);
527        assert_eq!(
528            default_vector_column(&schema_with_nested_vec_col, None).unwrap(),
529            "image.embedding"
530        );
531
532        let schema_with_escaped_nested_vec_col = Schema::new(vec![Field::new(
533            "image-meta",
534            DataType::Struct(
535                vec![Field::new(
536                    "embedding.v1",
537                    DataType::FixedSizeList(
538                        Arc::new(Field::new("item", DataType::Float32, false)),
539                        10,
540                    ),
541                    false,
542                )]
543                .into(),
544            ),
545            false,
546        )]);
547        assert_eq!(
548            default_vector_column(&schema_with_escaped_nested_vec_col, None).unwrap(),
549            "`image-meta`.`embedding.v1`"
550        );
551
552        let multi_vec_col = Schema::new(vec![
553            Field::new("id", DataType::Int16, true),
554            Field::new(
555                "vec",
556                DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float64, false)), 10),
557                false,
558            ),
559            Field::new(
560                "vec2",
561                DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float64, false)), 50),
562                false,
563            ),
564        ]);
565        assert!(
566            default_vector_column(&multi_vec_col, None)
567                .unwrap_err()
568                .to_string()
569                .contains("More than one")
570        );
571
572        let multi_nested_vec_col = Schema::new(vec![
573            Field::new(
574                "image",
575                DataType::Struct(
576                    vec![Field::new(
577                        "embedding",
578                        DataType::FixedSizeList(
579                            Arc::new(Field::new("item", DataType::Float32, false)),
580                            10,
581                        ),
582                        false,
583                    )]
584                    .into(),
585                ),
586                false,
587            ),
588            Field::new(
589                "text",
590                DataType::Struct(
591                    vec![Field::new(
592                        "embedding",
593                        DataType::FixedSizeList(
594                            Arc::new(Field::new("item", DataType::Float32, false)),
595                            50,
596                        ),
597                        false,
598                    )]
599                    .into(),
600                ),
601                false,
602            ),
603        ]);
604        assert_eq!(
605            default_vector_column(&multi_nested_vec_col, Some(50)).unwrap(),
606            "text.embedding"
607        );
608        let err = default_vector_column(&multi_nested_vec_col, None)
609            .unwrap_err()
610            .to_string();
611        assert!(err.contains("image.embedding"));
612        assert!(err.contains("text.embedding"));
613    }
614
615    #[test]
616    fn test_validate_table_name() {
617        assert!(validate_table_name("my_table").is_ok());
618        assert!(validate_table_name("my_table_1").is_ok());
619        assert!(validate_table_name("123mytable").is_ok());
620        assert!(validate_table_name("_12345table").is_ok());
621        assert!(validate_table_name("table.12345").is_ok());
622        assert!(validate_table_name("table.._dot_..12345").is_ok());
623
624        assert!(validate_table_name("").is_err());
625        assert!(validate_table_name("my_table!").is_err());
626        assert!(validate_table_name("my/table").is_err());
627        assert!(validate_table_name("my@table").is_err());
628        assert!(validate_table_name("name with space").is_err());
629    }
630
631    #[test]
632    fn test_validate_namespace_name() {
633        // Valid namespace names
634        assert!(validate_namespace_name("ns1").is_ok());
635        assert!(validate_namespace_name("namespace_123").is_ok());
636        assert!(validate_namespace_name("my-namespace").is_ok());
637        assert!(validate_namespace_name("my.namespace").is_ok());
638        assert!(validate_namespace_name("NS_1.2.3").is_ok());
639        assert!(validate_namespace_name("a").is_ok());
640        assert!(validate_namespace_name("123").is_ok());
641        assert!(validate_namespace_name("_underscore").is_ok());
642        assert!(validate_namespace_name("-hyphen").is_ok());
643        assert!(validate_namespace_name(".period").is_ok());
644
645        // Invalid namespace names
646        assert!(validate_namespace_name("").is_err());
647        assert!(validate_namespace_name("namespace with spaces").is_err());
648        assert!(validate_namespace_name("namespace/with/slashes").is_err());
649        assert!(validate_namespace_name("namespace\\with\\backslashes").is_err());
650        assert!(validate_namespace_name("namespace$with$delimiter").is_err());
651        assert!(validate_namespace_name("namespace@special").is_err());
652        assert!(validate_namespace_name("namespace#hash").is_err());
653    }
654
655    #[test]
656    fn test_validate_namespace() {
657        // Valid namespace with single component
658        assert!(validate_namespace(&["ns1".to_string()]).is_ok());
659
660        // Valid namespace with multiple components
661        assert!(
662            validate_namespace(&["ns1".to_string(), "ns2".to_string(), "ns3".to_string()]).is_ok()
663        );
664
665        // Empty namespace (root) is valid
666        assert!(validate_namespace(&[]).is_ok());
667
668        // Invalid: contains empty component
669        assert!(validate_namespace(&["ns1".to_string(), "".to_string()]).is_err());
670
671        // Invalid: contains component with spaces
672        assert!(validate_namespace(&["ns1".to_string(), "ns 2".to_string()]).is_err());
673
674        // Invalid: contains component with special characters
675        assert!(validate_namespace(&["ns1".to_string(), "ns@2".to_string()]).is_err());
676        assert!(validate_namespace(&["ns1".to_string(), "ns/2".to_string()]).is_err());
677        assert!(validate_namespace(&["ns1".to_string(), "ns$2".to_string()]).is_err());
678
679        // Valid: underscores, hyphens, and periods are allowed
680        assert!(
681            validate_namespace(&["ns_1".to_string(), "ns-2".to_string(), "ns.3".to_string()])
682                .is_ok()
683        );
684    }
685
686    #[test]
687    fn test_string_to_datatype() {
688        let string = "int32";
689        let expected = DataType::Int32;
690        assert_eq!(string_to_datatype(string), Some(expected));
691    }
692
693    fn sample_batch(num_rows: i32) -> RecordBatch {
694        let schema = Arc::new(Schema::new(vec![Field::new(
695            "col1",
696            DataType::Int32,
697            false,
698        )]));
699        RecordBatch::try_new(
700            schema.clone(),
701            vec![Arc::new(Int32Array::from_iter_values(0..num_rows))],
702        )
703        .unwrap()
704    }
705
706    #[tokio::test]
707    async fn test_timeout_stream() {
708        let batch = sample_batch(3);
709        let schema = batch.schema();
710        let mock_stream = stream::iter(vec![Ok(batch.clone()), Ok(batch.clone())]);
711
712        let sendable_stream: SendableRecordBatchStream =
713            Box::pin(RecordBatchStreamAdapter::new(schema.clone(), mock_stream));
714        let timeout_duration = std::time::Duration::from_millis(10);
715        let mut timeout_stream = TimeoutStream::new(sendable_stream, timeout_duration);
716
717        // Poll the stream to get the first batch
718        let first_result = timeout_stream.next().await;
719        assert!(first_result.is_some());
720        assert!(first_result.unwrap().is_ok());
721
722        // Sleep for the timeout duration
723        sleep(timeout_duration).await;
724
725        // Poll the stream again and ensure it returns a timeout error
726        let second_result = timeout_stream.next().await.unwrap();
727        assert!(second_result.is_err());
728        assert!(
729            second_result
730                .unwrap_err()
731                .to_string()
732                .contains("Query timeout")
733        );
734    }
735
736    #[tokio::test]
737    async fn test_timeout_stream_zero_duration() {
738        let batch = sample_batch(3);
739        let schema = batch.schema();
740        let mock_stream = stream::iter(vec![Ok(batch.clone()), Ok(batch.clone())]);
741
742        let sendable_stream: SendableRecordBatchStream =
743            Box::pin(RecordBatchStreamAdapter::new(schema.clone(), mock_stream));
744
745        // Setup similar to test_timeout_stream
746        let timeout_duration = std::time::Duration::from_secs(0);
747        let mut timeout_stream = TimeoutStream::new(sendable_stream, timeout_duration);
748
749        // First poll should immediately return a timeout error
750        let result = timeout_stream.next().await.unwrap();
751        assert!(result.is_err());
752        assert!(result.unwrap_err().to_string().contains("Query timeout"));
753    }
754
755    #[tokio::test]
756    async fn test_timeout_stream_completes_normally() {
757        let batch = sample_batch(3);
758        let schema = batch.schema();
759        let mock_stream = stream::iter(vec![Ok(batch.clone()), Ok(batch.clone())]);
760
761        let sendable_stream: SendableRecordBatchStream =
762            Box::pin(RecordBatchStreamAdapter::new(schema.clone(), mock_stream));
763
764        // Setup a stream with 2 batches
765        // Use a longer timeout that won't trigger
766        let timeout_duration = std::time::Duration::from_secs(1);
767        let mut timeout_stream = TimeoutStream::new(sendable_stream, timeout_duration);
768
769        // Both polls should return data normally
770        assert!(timeout_stream.next().await.unwrap().is_ok());
771        assert!(timeout_stream.next().await.unwrap().is_ok());
772        // Stream should be empty now
773        assert!(timeout_stream.next().await.is_none());
774    }
775
776    async fn collect_batch_sizes(
777        stream: SendableRecordBatchStream,
778        max_batch_length: usize,
779    ) -> Vec<usize> {
780        let mut sliced_stream = MaxBatchLengthStream::new(stream, max_batch_length);
781        sliced_stream
782            .by_ref()
783            .map(|batch| batch.unwrap().num_rows())
784            .collect::<Vec<_>>()
785            .await
786    }
787
788    #[tokio::test]
789    async fn test_max_batch_length_stream_behaviors() {
790        let schema = sample_batch(7).schema();
791        let mock_stream = stream::iter(vec![Ok(sample_batch(2)), Ok(sample_batch(7))]);
792
793        let sendable_stream: SendableRecordBatchStream =
794            Box::pin(RecordBatchStreamAdapter::new(schema.clone(), mock_stream));
795        assert_eq!(
796            collect_batch_sizes(sendable_stream, 3).await,
797            vec![2, 3, 3, 1]
798        );
799
800        let sendable_stream: SendableRecordBatchStream = Box::pin(RecordBatchStreamAdapter::new(
801            schema,
802            stream::iter(vec![Ok(sample_batch(2)), Ok(sample_batch(7))]),
803        ));
804        assert_eq!(collect_batch_sizes(sendable_stream, 0).await, vec![2, 7]);
805    }
806}