Skip to main content

stac_duckdb/
client.rs

1use crate::{Error, Extension, Result};
2use arrow_array::{RecordBatch, RecordBatchIterator};
3use arrow_schema::{ArrowError, SchemaRef};
4use chrono::DateTime;
5use cql2::{Expr, ToDuckSQL};
6use duckdb::{Connection, Statement, types::Value};
7use geo::BoundingRect;
8use geojson::GeometryValue;
9#[cfg(feature = "async")]
10use stac::api::StreamItemsClient;
11use stac::api::{
12    ArrowItemsClient, CollectionsClient, Direction, ItemsClient, RecordBatchReaderAdapter, Search,
13};
14use stac::{Collection, SpatialExtent, TemporalExtent, geoarrow::DATETIME_COLUMNS};
15use std::ops::{Deref, DerefMut};
16use std::sync::Mutex;
17
18/// Default hive partitioning value
19pub const DEFAULT_USE_HIVE_PARTITIONING: bool = false;
20
21/// Default convert wkb value.
22pub const DEFAULT_CONVERT_WKB: bool = true;
23
24/// The default collection description.
25pub const DEFAULT_COLLECTION_DESCRIPTION: &str =
26    "Auto-generated collection from stac-geoparquet extents";
27
28/// The default union by name value.
29pub const DEFAULT_UNION_BY_NAME: bool = true;
30
31/// Whether to remove the filename column by default.
32pub const DEFAULT_REMOVE_FILENAME_COLUMN: bool = true;
33
34/// The default source format for the data being queried.
35pub const DEFAULT_SOURCE_FORMAT: SourceFormat = SourceFormat::Parquet;
36
37/// A client for making DuckDB requests for STAC objects.
38#[derive(Debug)]
39pub struct Client {
40    connection: Connection,
41
42    /// Whether to use hive partitioning
43    pub use_hive_partitioning: bool,
44
45    /// Whether to convert WKB to native geometries.
46    ///
47    /// If False, WKB metadata will be added.
48    pub convert_wkb: bool,
49
50    /// Whether to use `union_by_name` when querying.
51    ///
52    /// Defaults to true.
53    pub union_by_name: bool,
54
55    /// Whether to remove the `filename` column that DuckDB adds automatically.
56    ///
57    /// Defaults to true.
58    pub remove_filename_column: bool,
59
60    source_format: SourceFormat,
61}
62
63impl Client {
64    /// Creates a new client with an in-memory DuckDB connection.
65    ///
66    /// Installs the `spatial`, `icu`, and `httpfs` extensions. The `iceberg`
67    /// extension is installed lazily by [`Client::set_source_format`]. If
68    /// you'd like to manage your own extensions (e.g. if your extensions are
69    /// stored in a different location), set things up then use
70    /// `connection.into()` to get a new `Client`.
71    ///
72    /// # Examples
73    ///
74    /// ```
75    /// use stac_duckdb::Client;
76    ///
77    /// let client = Client::new().unwrap();
78    /// ```
79    pub fn new() -> Result<Client> {
80        let connection = Connection::open_in_memory()?;
81        connection.execute("INSTALL spatial", [])?;
82        connection.execute("LOAD spatial", [])?;
83        connection.execute("INSTALL icu", [])?;
84        connection.execute("LOAD icu", [])?;
85        connection.execute("INSTALL httpfs", [])?;
86        connection.execute("LOAD httpfs", [])?;
87        Ok(connection.into())
88    }
89
90    /// Returns a vector of all extensions.
91    ///
92    /// # Examples
93    ///
94    /// ```
95    /// use stac_duckdb::Client;
96    ///
97    /// let client = Client::new().unwrap();
98    /// let extensions = client.extensions().unwrap();
99    /// ```
100    pub fn extensions(&self) -> Result<Vec<Extension>> {
101        let mut statement = self.prepare(
102            "SELECT extension_name, loaded, installed, install_path, description, extension_version, install_mode, installed_from FROM duckdb_extensions();",
103        )?;
104        let extensions = statement
105            .query_map([], |row| {
106                Ok(Extension {
107                    name: row.get("extension_name")?,
108                    loaded: row.get("loaded")?,
109                    installed: row.get("installed")?,
110                    install_path: row.get("install_path")?,
111                    description: row.get("description")?,
112                    version: row.get("extension_version")?,
113                    install_mode: row.get("install_mode")?,
114                    installed_from: row.get("installed_from")?,
115                })
116            })?
117            .collect::<std::result::Result<Vec<_>, duckdb::Error>>()?;
118        Ok(extensions)
119    }
120
121    /// Returns one or more [stac::Collection] from the items in the stac-geoparquet file.
122    ///
123    /// # Examples
124    ///
125    /// ```
126    /// use stac_duckdb::Client;
127    ///
128    /// let client = Client::new().unwrap();
129    /// let collections = client.collections("data/100-sentinel-2-items.parquet").unwrap();
130    /// ```
131    pub fn collections(&self, href: &str) -> Result<Vec<Collection>> {
132        let start_datetime= if self.prepare(&format!(
133            "SELECT column_name FROM (DESCRIBE SELECT * from {}) where column_name = 'start_datetime'",
134            self.format_source_href(href)
135        ))?.query([])?.next()?.is_some() {
136            "strftime(min(coalesce(start_datetime, datetime)), '%xT%X%z')"
137        } else {
138            "strftime(min(datetime), '%xT%X%z')"
139        };
140        let end_datetime = if self
141            .prepare(&format!(
142            "SELECT column_name FROM (DESCRIBE SELECT * from {}) where column_name = 'end_datetime'",
143            self.format_source_href(href)
144        ))?
145            .query([])?
146            .next()?
147            .is_some()
148        {
149            "strftime(max(coalesce(end_datetime, datetime)), '%xT%X%z')"
150        } else {
151            "strftime(max(datetime), '%xT%X%z')"
152        };
153        let mut statement = self.prepare(&format!(
154            "SELECT DISTINCT collection FROM {}",
155            self.format_source_href(href)
156        ))?;
157        let mut collections = Vec::new();
158        for row in statement.query_map([], |row| row.get::<_, String>(0))? {
159            let collection_id = row?;
160            let mut statement = self.connection.prepare(&format!(
161                "SELECT ST_AsGeoJSON(ST_Extent_Agg({})), {}, {} FROM {} WHERE collection = $1",
162                self.geometry_expr(),
163                start_datetime,
164                end_datetime,
165                self.format_source_href(href)
166            ))?;
167            let row = statement.query_row([&collection_id], |row| {
168                Ok((
169                    row.get::<_, String>(0)?,
170                    row.get::<_, String>(1)?,
171                    row.get::<_, String>(2)?,
172                ))
173            })?;
174            let mut collection = Collection::new(collection_id, DEFAULT_COLLECTION_DESCRIPTION);
175            let geometry: geo::Geometry = serde_json::from_str::<GeometryValue>(&row.0)?
176                .try_into()
177                .map_err(Box::new)?;
178            if let Some(bbox) = geometry.bounding_rect() {
179                collection.extent.spatial = SpatialExtent {
180                    bbox: vec![bbox.into()],
181                };
182            }
183            collection.extent.temporal = TemporalExtent {
184                interval: vec![[
185                    Some(DateTime::parse_from_str(&row.1, "%FT%T%#z")?.into()),
186                    Some(DateTime::parse_from_str(&row.2, "%FT%T%#z")?.into()),
187                ]],
188            };
189            collections.push(collection);
190        }
191        Ok(collections)
192    }
193
194    /// Searches a single stac-geoparquet file.
195    ///
196    /// # Examples
197    ///
198    /// ```
199    /// use stac_duckdb::Client;
200    ///
201    /// let client = Client::new().unwrap();
202    /// let item_collection = client.search("data/100-sentinel-2-items.parquet", Default::default()).unwrap();
203    /// ```
204    pub fn search(&self, href: &str, search: Search) -> Result<stac::api::ItemCollection> {
205        let mut arrow_iter = self.search_to_arrow(href, search)?;
206        let Some(schema) = arrow_iter.schema() else {
207            return Ok(Default::default());
208        };
209
210        let first_batch = match arrow_iter.next() {
211            Some(batch) => batch?,
212            None => return Ok(Default::default()),
213        };
214
215        let batches = std::iter::once(Ok(first_batch))
216            .chain(arrow_iter)
217            .map(|batch| batch.map_err(|err| ArrowError::ExternalError(Box::new(err))));
218
219        let item_collection = stac::geoarrow::json::from_record_batch_reader(
220            RecordBatchIterator::new(batches, schema),
221        )?;
222        Ok(item_collection.into())
223    }
224
225    /// Searches to an iterator of record batches.
226    ///
227    /// # Examples
228    ///
229    /// ```
230    /// use stac_duckdb::Client;
231    ///
232    /// let client = Client::new().unwrap();
233    /// let mut total = 0;
234    /// for batch in client
235    ///     .search_to_arrow("data/100-sentinel-2-items.parquet", Default::default())
236    ///     .unwrap()
237    /// {
238    ///     let batch = batch.unwrap();
239    ///     total += batch.num_rows();
240    /// }
241    /// assert_eq!(total, 100);
242    /// ```
243    pub fn search_to_arrow<'conn>(
244        &'conn self,
245        href: &str,
246        search: Search,
247    ) -> Result<SearchArrowBatchIter<'conn>> {
248        if let Some((sql, params)) = self.build_query(href, search)? {
249            log::debug!("duckdb sql: {sql}");
250            let mut statement = self.prepare(&sql)?;
251            statement.execute(duckdb::params_from_iter(params))?;
252            log::debug!("query complete");
253            Ok(SearchArrowBatchIter::new(
254                statement,
255                self.convert_wkb,
256                self.remove_filename_column,
257            ))
258        } else {
259            Ok(SearchArrowBatchIter::empty(
260                self.convert_wkb,
261                self.remove_filename_column,
262            ))
263        }
264    }
265
266    /// Returns the SQL query string and parameters for this href and search object.
267    ///
268    /// Returns `None` if we can _know_ that the query will return nothing.
269    ///
270    /// # Examples
271    ///
272    /// ```
273    /// use stac_duckdb::Client;
274    ///
275    /// let client = Client::new().unwrap();
276    /// let (sql, params) = client.build_query("data/100-sentinel-2-items.parquet", Default::default()).unwrap().unwrap();
277    /// ```
278    pub fn build_query(&self, href: &str, search: Search) -> Result<Option<(String, Vec<Value>)>> {
279        // Note that we pull out some fields early so we can avoid closing some search strings below.
280
281        if search.items.query.is_some() {
282            return Err(Error::QueryNotImplemented);
283        }
284
285        // Check which columns we'll be selecting
286        let mut statement = self.prepare(&format!(
287            "SELECT column_name FROM (DESCRIBE SELECT * from {})",
288            self.format_source_href(href)
289        ))?;
290        let mut has_start_datetime = false;
291        let mut has_end_datetime = false;
292        let mut column_names = Vec::new();
293        let mut columns = Vec::new();
294        for row in statement.query_map([], |row| row.get::<_, String>(0))? {
295            let column = row?;
296            if column == "start_datetime" {
297                has_start_datetime = true;
298            }
299            if column == "end_datetime" {
300                has_end_datetime = true;
301            }
302
303            if let Some(fields) = search.fields.as_ref()
304                && (fields.exclude.contains(&column)
305                    || !(fields.include.is_empty() || fields.include.contains(&column)))
306            {
307                continue;
308            }
309
310            if column == "geometry" {
311                match self.source_format {
312                    SourceFormat::Parquet => {
313                        columns.push("ST_AsWKB(geometry) geometry".to_string())
314                    }
315                    SourceFormat::Iceberg => columns.push("geometry".to_string()),
316                }
317            } else if DATETIME_COLUMNS.contains(&column.as_str()) {
318                columns.push(format!("\"{column}\"::TIMESTAMPTZ {column}"))
319            } else {
320                columns.push(format!("\"{column}\""));
321            }
322            column_names.push(column);
323        }
324
325        // Get limit and offset
326        let limit = search.items.limit;
327        let offset = search
328            .items
329            .additional_fields
330            .get("offset")
331            .and_then(|v| v.as_i64());
332
333        // Build order_by
334        let mut order_by = Vec::with_capacity(search.sortby.len());
335        for sortby in &search.sortby {
336            order_by.push(format!(
337                "\"{}\" {}",
338                sortby.field,
339                match sortby.direction {
340                    Direction::Ascending => "ASC",
341                    Direction::Descending => "DESC",
342                }
343            ));
344        }
345
346        // Build wheres and params
347        let mut wheres = Vec::new();
348        let mut params = Vec::new();
349        if !search.ids.is_empty() {
350            wheres.push(format!(
351                "id IN ({})",
352                (0..search.ids.len())
353                    .map(|_| "?")
354                    .collect::<Vec<_>>()
355                    .join(",")
356            ));
357            params.extend(search.ids.into_iter().map(Value::Text));
358        }
359        if let Some(intersects) = search.intersects {
360            wheres.push(format!(
361                "ST_Intersects({}, ST_GeomFromGeoJSON(?))",
362                self.geometry_expr()
363            ));
364            params.push(Value::Text(intersects.to_string()));
365        }
366        if !search.collections.is_empty() {
367            wheres.push(format!(
368                "collection IN ({})",
369                (0..search.collections.len())
370                    .map(|_| "?")
371                    .collect::<Vec<_>>()
372                    .join(",")
373            ));
374            params.extend(search.collections.into_iter().map(Value::Text));
375        }
376        if let Some(bbox) = search.items.bbox {
377            wheres.push(format!(
378                "ST_Intersects({}, ST_GeomFromGeoJSON(?))",
379                self.geometry_expr()
380            ));
381            params.push(Value::Text(bbox.to_geometry().to_string()));
382        }
383        if let Some(datetime) = search.items.datetime {
384            let interval = stac::datetime::parse(&datetime)?;
385            if let Some(start) = interval.0 {
386                wheres.push(format!(
387                    "?::TIMESTAMPTZ <= {}",
388                    if has_end_datetime {
389                        "coalesce(end_datetime, datetime)"
390                    } else {
391                        "datetime"
392                    }
393                ));
394                params.push(Value::Text(start.to_rfc3339()));
395            }
396            if let Some(end) = interval.1 {
397                wheres.push(format!(
398                    "?::TIMESTAMPTZ >= {}", // Inclusive, https://github.com/radiantearth/stac-spec/pull/1280
399                    if has_start_datetime {
400                        "coalesce(start_datetime, datetime)"
401                    } else {
402                        "datetime"
403                    }
404                ));
405                params.push(Value::Text(end.to_rfc3339()));
406            }
407        }
408        if let Some(filter) = search.items.filter {
409            let expr: Expr = filter.try_into()?;
410            if expr_properties_match(&expr, &column_names) {
411                let sql = expr.to_ducksql().map_err(Box::new)?;
412                wheres.push(sql);
413            } else {
414                return Ok(None);
415            }
416        }
417
418        let mut suffix = String::new();
419        if !wheres.is_empty() {
420            suffix.push_str(&format!(" WHERE {}", wheres.join(" AND ")));
421        }
422        if !order_by.is_empty() {
423            suffix.push_str(&format!(" ORDER BY {}", order_by.join(", ")));
424        }
425        if let Some(limit) = limit {
426            suffix.push_str(&format!(" LIMIT {limit}"));
427        }
428        if let Some(offset) = offset {
429            suffix.push_str(&format!(" OFFSET {offset}"));
430        }
431
432        let sql = format!(
433            "SELECT {} FROM {}{}",
434            columns.join(","),
435            self.format_source_href(href),
436            suffix,
437        );
438        Ok(Some((sql, params)))
439    }
440
441    fn format_source_href(&self, href: &str) -> String {
442        match self.source_format {
443            SourceFormat::Parquet => format!(
444                "read_parquet('{}', hive_partitioning={}, union_by_name={})",
445                href, self.use_hive_partitioning, self.union_by_name
446            ),
447            SourceFormat::Iceberg => format!("iceberg_scan('{href}')"),
448        }
449    }
450
451    fn geometry_expr(&self) -> &'static str {
452        match self.source_format {
453            SourceFormat::Parquet => "geometry",
454            SourceFormat::Iceberg => "ST_GeomFromWKB(geometry)",
455        }
456    }
457
458    /// Returns the source format of the data being queried.
459    ///
460    /// Defaults to [`SourceFormat::Parquet`].
461    pub fn source_format(&self) -> SourceFormat {
462        self.source_format
463    }
464
465    /// Sets the source format of the data being queried.
466    ///
467    /// # Examples
468    ///
469    /// ```
470    /// use stac_duckdb::{Client, SourceFormat};
471    ///
472    /// let mut client = Client::new().unwrap();
473    /// client.set_source_format(SourceFormat::Iceberg).unwrap();
474    /// ```
475    pub fn set_source_format(&mut self, source_format: SourceFormat) -> Result<()> {
476        match source_format {
477            SourceFormat::Iceberg => {
478                self.connection.execute("INSTALL iceberg", [])?;
479                self.connection.execute("LOAD iceberg", [])?;
480                self.connection
481                    .execute("SET enable_geoparquet_conversion = false", [])?;
482            }
483            SourceFormat::Parquet => {
484                self.connection
485                    .execute("RESET enable_geoparquet_conversion", [])?;
486            }
487        }
488        self.source_format = source_format;
489        Ok(())
490    }
491}
492
493fn expr_properties_match(expr: &Expr, properties: &[String]) -> bool {
494    use Expr::*;
495
496    match expr {
497        Property { property } => properties.contains(property),
498        Float(_) | Literal(_) | Bool(_) | Geometry(_) => true,
499        Operation { args, .. } => args
500            .iter()
501            .all(|expr| expr_properties_match(expr, properties)),
502        Interval { interval } => interval
503            .iter()
504            .all(|expr| expr_properties_match(expr, properties)),
505        Timestamp { timestamp } => expr_properties_match(timestamp, properties),
506        Date { date } => expr_properties_match(date, properties),
507        Array(exprs) => exprs
508            .iter()
509            .all(|expr| expr_properties_match(expr, properties)),
510        BBox { bbox } => bbox
511            .iter()
512            .all(|expr| expr_properties_match(expr, properties)),
513        Null => expr_properties_match(expr, properties),
514    }
515}
516
517impl Deref for Client {
518    type Target = Connection;
519
520    fn deref(&self) -> &Self::Target {
521        &self.connection
522    }
523}
524
525impl DerefMut for Client {
526    fn deref_mut(&mut self) -> &mut Self::Target {
527        &mut self.connection
528    }
529}
530
531impl From<Connection> for Client {
532    fn from(connection: Connection) -> Self {
533        Client {
534            connection,
535            use_hive_partitioning: DEFAULT_USE_HIVE_PARTITIONING,
536            convert_wkb: DEFAULT_CONVERT_WKB,
537            union_by_name: DEFAULT_UNION_BY_NAME,
538            remove_filename_column: DEFAULT_REMOVE_FILENAME_COLUMN,
539            source_format: DEFAULT_SOURCE_FORMAT,
540        }
541    }
542}
543
544/// The source format of the data being queried.
545#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
546pub enum SourceFormat {
547    #[default]
548    Parquet,
549    Iceberg,
550}
551
552/// A DuckDB client bound to a specific stac-geoparquet or iceberg href.
553///
554/// This wraps a [`Client`] with a specific href, implementing the
555/// [`ArrowItemsClient`] trait. Because [`duckdb::Connection`] is not
556/// [`Sync`], use [`Mutex<HrefClient>`](std::sync::Mutex) for the async client
557/// traits ([`ItemsClient`] and [`CollectionsClient`]).
558///
559/// # Examples
560///
561/// ```
562/// use stac::api::ArrowItemsClient;
563/// use stac_duckdb::HrefClient;
564///
565/// let client = HrefClient::new("data/100-sentinel-2-items.parquet").unwrap();
566/// let record_batch_reader = client.search_to_arrow(Default::default()).unwrap();
567/// ```
568#[derive(Debug)]
569pub struct HrefClient {
570    client: Client,
571    href: String,
572}
573
574impl HrefClient {
575    /// Creates a new `HrefClient` for the given href.
576    pub fn new(href: impl ToString) -> Result<HrefClient> {
577        let client = Client::new()?;
578        Ok(HrefClient {
579            client,
580            href: href.to_string(),
581        })
582    }
583
584    /// Creates a new `HrefClient` from an existing [`Client`] and href.
585    pub fn from_client(client: Client, href: impl ToString) -> HrefClient {
586        HrefClient {
587            client,
588            href: href.to_string(),
589        }
590    }
591
592    /// Returns a reference to the underlying [`Client`].
593    pub fn client(&self) -> &Client {
594        &self.client
595    }
596
597    /// Returns a mutable reference to the underlying [`Client`].
598    pub fn client_mut(&mut self) -> &mut Client {
599        &mut self.client
600    }
601
602    /// Returns the href.
603    pub fn href(&self) -> &str {
604        &self.href
605    }
606}
607
608impl ArrowItemsClient for HrefClient {
609    type Error = Error;
610    type RecordBatchStream<'a> = ArrowBatchReader<'a>;
611
612    fn search_to_arrow(&self, search: Search) -> std::result::Result<ArrowBatchReader<'_>, Error> {
613        let iter = self.client.search_to_arrow(&self.href, search)?;
614        Ok(make_arrow_batch_reader(iter))
615    }
616}
617
618/// A thread-safe wrapper around [`HrefClient`] that implements
619/// [`ItemsClient`] and [`CollectionsClient`].
620///
621/// Use this when you need the async client traits. For [`ArrowItemsClient`],
622/// use [`HrefClient`] directly.
623///
624/// # Examples
625///
626/// ```
627/// use stac::api::ItemsClient;
628/// use stac_duckdb::SyncHrefClient;
629///
630/// let client = SyncHrefClient::new("data/100-sentinel-2-items.parquet").unwrap();
631/// # tokio_test::block_on(async {
632/// let item_collection = client.search(Default::default()).await.unwrap();
633/// # })
634/// ```
635#[derive(Debug)]
636pub struct SyncHrefClient {
637    inner: Mutex<HrefClient>,
638}
639
640impl SyncHrefClient {
641    /// Creates a new `SyncHrefClient` for the given href.
642    pub fn new(href: impl ToString) -> Result<SyncHrefClient> {
643        Ok(SyncHrefClient {
644            inner: Mutex::new(HrefClient::new(href)?),
645        })
646    }
647
648    /// Creates a new `SyncHrefClient` from an existing [`Client`] and href.
649    pub fn from_client(client: Client, href: impl ToString) -> SyncHrefClient {
650        SyncHrefClient {
651            inner: Mutex::new(HrefClient::from_client(client, href)),
652        }
653    }
654}
655
656impl ItemsClient for SyncHrefClient {
657    type Error = Error;
658
659    async fn search(
660        &self,
661        search: Search,
662    ) -> std::result::Result<stac::api::ItemCollection, Error> {
663        let guard = self.inner.lock().expect("SyncHrefClient mutex is poisoned");
664        guard.client.search(&guard.href, search)
665    }
666}
667
668impl CollectionsClient for SyncHrefClient {
669    type Error = Error;
670
671    async fn collections(&self) -> std::result::Result<Vec<Collection>, Error> {
672        let guard = self.inner.lock().expect("SyncHrefClient mutex is poisoned");
673        guard.client.collections(&guard.href)
674    }
675}
676
677#[cfg(feature = "async")]
678impl StreamItemsClient for SyncHrefClient {
679    type Error = Error;
680
681    async fn search_stream(
682        &self,
683        search: Search,
684    ) -> std::result::Result<
685        impl futures_core::Stream<Item = std::result::Result<stac::api::Item, Error>> + Send,
686        Error,
687    > {
688        // DuckDB queries run synchronously and return the full result in one go.
689        // We collect eagerly here (holding the mutex only for the duration of
690        // the underlying synchronous query) and then stream the items.
691        let item_collection = ItemsClient::search(self, search).await?;
692        Ok(futures::stream::iter(
693            item_collection.items.into_iter().map(Ok::<_, Error>),
694        ))
695    }
696}
697
698/// A wrapper around [`SearchArrowBatchIter`] that implements
699/// [`arrow_array::RecordBatchReader`].
700///
701/// This is a type alias for [`stac::api::RecordBatchReaderAdapter`], which
702/// provides the generic `Iterator<Item = Result<RecordBatch, E>>` →
703/// [`arrow_array::RecordBatchReader`] bridge.
704pub type ArrowBatchReader<'a> = RecordBatchReaderAdapter<SearchArrowBatchIter<'a>>;
705
706/// Constructs an [`ArrowBatchReader`] from a [`SearchArrowBatchIter`].
707fn make_arrow_batch_reader(inner: SearchArrowBatchIter<'_>) -> ArrowBatchReader<'_> {
708    let schema = inner
709        .schema()
710        .unwrap_or_else(|| arrow_schema::Schema::empty().into());
711    RecordBatchReaderAdapter::new(inner, schema)
712}
713
714/// Iterator returned by [`Client::search_to_arrow`].
715pub struct SearchArrowBatchIter<'conn> {
716    statement: Option<Statement<'conn>>,
717    convert_wkb: bool,
718    remove_filename_column: bool,
719    schema: Option<SchemaRef>,
720}
721
722impl<'conn> SearchArrowBatchIter<'conn> {
723    fn new(statement: Statement<'conn>, convert_wkb: bool, remove_filename_column: bool) -> Self {
724        let schema = Some(statement.schema());
725        Self {
726            statement: Some(statement),
727            convert_wkb,
728            remove_filename_column,
729            schema,
730        }
731    }
732
733    fn empty(convert_wkb: bool, remove_filename_column: bool) -> Self {
734        Self {
735            statement: None,
736            convert_wkb,
737            remove_filename_column,
738            schema: None,
739        }
740    }
741
742    pub fn schema(&self) -> Option<SchemaRef> {
743        self.schema.clone()
744    }
745
746    fn finalize_batch(&self, record_batch: RecordBatch) -> Result<RecordBatch> {
747        let mut record_batch = if self.convert_wkb {
748            stac::geoarrow::with_native_geometry(record_batch, "geometry")?
749        } else {
750            stac::geoarrow::add_wkb_metadata(record_batch, "geometry")?
751        };
752        if self.remove_filename_column {
753            record_batch = remove_column(record_batch, "filename");
754        }
755        Ok(record_batch)
756    }
757}
758
759impl<'conn> Iterator for SearchArrowBatchIter<'conn> {
760    type Item = Result<RecordBatch>;
761
762    fn next(&mut self) -> Option<Self::Item> {
763        let statement = self.statement.as_ref()?;
764
765        match statement.step() {
766            Ok(Some(struct_array)) => {
767                let record_batch = RecordBatch::from(&struct_array);
768                match self.finalize_batch(record_batch) {
769                    Ok(batch) => Some(Ok(batch)),
770                    Err(err) => {
771                        self.statement = None;
772                        Some(Err(err))
773                    }
774                }
775            }
776            Ok(None) => {
777                self.statement = None;
778                None
779            }
780            Err(err) => {
781                self.statement = None;
782                Some(Err(err.into()))
783            }
784        }
785    }
786}
787
788fn remove_column(mut record_batch: RecordBatch, name: &str) -> RecordBatch {
789    if let Some((index, _)) = record_batch.schema().column_with_name(name) {
790        record_batch.remove_column(index);
791    }
792    record_batch
793}
794
795#[cfg(test)]
796mod tests {
797    use super::{Client, SourceFormat};
798    use duckdb::Connection;
799    use geo::Geometry;
800    use rstest::{fixture, rstest};
801    use stac::Bbox;
802    use stac::api::{Items, Search, Sortby};
803    use stac_validate::Validate;
804
805    #[fixture]
806    #[once]
807    fn install_extensions() {
808        let connection = Connection::open_in_memory().unwrap();
809        connection.execute("INSTALL icu", []).unwrap();
810        connection.execute("INSTALL spatial", []).unwrap();
811    }
812
813    #[allow(unused_variables)]
814    #[fixture]
815    fn client(install_extensions: ()) -> Client {
816        Client::new().unwrap()
817    }
818
819    #[rstest]
820    fn extensions(client: Client) {
821        let _ = client.extensions().unwrap();
822    }
823
824    #[rstest]
825    #[tokio::test]
826    async fn search(client: Client) {
827        let item_collection = client
828            .search("data/100-sentinel-2-items.parquet", Search::default())
829            .unwrap();
830        assert_eq!(item_collection.items.len(), 100);
831        item_collection.items[0].validate().await.unwrap();
832    }
833
834    #[rstest]
835    fn search_to_arrow(client: Client) {
836        let record_batches = client
837            .search_to_arrow("data/100-sentinel-2-items.parquet", Search::default())
838            .unwrap()
839            .collect::<std::result::Result<Vec<_>, _>>()
840            .unwrap();
841        assert_eq!(record_batches.len(), 1);
842    }
843
844    #[rstest]
845    fn search_ids(client: Client) {
846        let item_collection = client
847            .search(
848                "data/100-sentinel-2-items.parquet",
849                Search::default().ids(vec![
850                    "S2A_MSIL2A_20240326T174951_R141_T13TDE_20240329T224429".to_string(),
851                ]),
852            )
853            .unwrap();
854        assert_eq!(item_collection.items.len(), 1);
855        assert_eq!(
856            item_collection.items[0]["id"],
857            "S2A_MSIL2A_20240326T174951_R141_T13TDE_20240329T224429"
858        );
859    }
860
861    #[rstest]
862    fn search_intersects(client: Client) {
863        let item_collection = client
864            .search(
865                "data/100-sentinel-2-items.parquet",
866                Search::default().intersects(&Geometry::Point(geo::point! { x: -106., y: 40.5 })),
867            )
868            .unwrap();
869        assert_eq!(item_collection.items.len(), 50);
870    }
871
872    #[rstest]
873    fn search_collections(client: Client) {
874        let item_collection = client
875            .search(
876                "data/100-sentinel-2-items.parquet",
877                Search::default().collections(vec!["sentinel-2-l2a".to_string()]),
878            )
879            .unwrap();
880        assert_eq!(item_collection.items.len(), 100);
881
882        let item_collection = client
883            .search(
884                "data/100-sentinel-2-items.parquet",
885                Search::default().collections(vec!["foobar".to_string()]),
886            )
887            .unwrap();
888        assert_eq!(item_collection.items.len(), 0);
889    }
890
891    #[rstest]
892    fn search_bbox(client: Client) {
893        let item_collection = client
894            .search(
895                "data/100-sentinel-2-items.parquet",
896                Search::default().bbox(Bbox::new(-106.1, 40.5, -106.0, 40.6)),
897            )
898            .unwrap();
899        assert_eq!(item_collection.items.len(), 50);
900    }
901
902    #[rstest]
903    fn search_datetime(client: Client) {
904        let item_collection = client
905            .search(
906                "data/100-sentinel-2-items.parquet",
907                Search::default().datetime("2024-12-02T00:00:00Z/.."),
908            )
909            .unwrap();
910        assert_eq!(item_collection.items.len(), 1);
911        let item_collection = client
912            .search(
913                "data/100-sentinel-2-items.parquet",
914                Search::default().datetime("../2024-12-02T00:00:00Z"),
915            )
916            .unwrap();
917        assert_eq!(item_collection.items.len(), 99);
918    }
919
920    #[rstest]
921    fn search_datetime_empty_interval(client: Client) {
922        let item_collection = client
923            .search(
924                "data/100-sentinel-2-items.parquet",
925                Search::default().datetime("2024-12-02T00:00:00Z/"),
926            )
927            .unwrap();
928        assert_eq!(item_collection.items.len(), 1);
929    }
930
931    #[rstest]
932    fn search_datetime_start_end_datetime(client: Client) {
933        let item_collection = client
934            .search(
935                "data/sentinel-1-global-mosaics.parquet",
936                Search::default().datetime("2026-04-15T00:00:00Z"),
937            )
938            .unwrap();
939        assert_eq!(item_collection.items.len(), 1);
940    }
941
942    #[rstest]
943    fn search_limit(client: Client) {
944        let item_collection = client
945            .search(
946                "data/100-sentinel-2-items.parquet",
947                Search::default().limit(42),
948            )
949            .unwrap();
950        assert_eq!(item_collection.items.len(), 42);
951    }
952
953    #[rstest]
954    fn search_offset(client: Client) {
955        let mut search = Search::default().limit(1);
956        search
957            .items
958            .additional_fields
959            .insert("offset".to_string(), 1.into());
960        let item_collection = client
961            .search("data/100-sentinel-2-items.parquet", search)
962            .unwrap();
963        assert_eq!(
964            item_collection.items[0]["id"],
965            "S2A_MSIL2A_20241201T175721_R141_T13TDE_20241201T213150"
966        );
967    }
968
969    #[rstest]
970    fn search_sortby(client: Client) {
971        let item_collection = client
972            .search(
973                "data/100-sentinel-2-items.parquet",
974                Search::default()
975                    .sortby(vec![Sortby::asc("datetime")])
976                    .limit(1),
977            )
978            .unwrap();
979        assert_eq!(
980            item_collection.items[0]["id"],
981            "S2A_MSIL2A_20240326T174951_R141_T13TDE_20240329T224429"
982        );
983
984        let item_collection = client
985            .search(
986                "data/100-sentinel-2-items.parquet",
987                Search::default()
988                    .sortby(vec![Sortby::desc("datetime")])
989                    .limit(1),
990            )
991            .unwrap();
992        assert_eq!(
993            item_collection.items[0]["id"],
994            "S2B_MSIL2A_20241203T174629_R098_T13TDE_20241203T211406"
995        );
996    }
997
998    #[rstest]
999    fn search_fields(client: Client) {
1000        let item_collection = client
1001            .search(
1002                "data/100-sentinel-2-items.parquet",
1003                Search::default().fields("+id".parse().unwrap()).limit(1),
1004            )
1005            .unwrap();
1006        assert_eq!(item_collection.items[0].len(), 1);
1007    }
1008
1009    #[rstest]
1010    fn collections(client: Client) {
1011        let collections = client
1012            .collections("data/100-sentinel-2-items.parquet")
1013            .unwrap();
1014        assert_eq!(collections.len(), 1);
1015    }
1016
1017    #[rstest]
1018    fn no_convert_wkb(mut client: Client) {
1019        client.convert_wkb = false;
1020        let record_batches = client
1021            .search_to_arrow("data/100-sentinel-2-items.parquet", Search::default())
1022            .unwrap()
1023            .collect::<std::result::Result<Vec<_>, _>>()
1024            .unwrap();
1025        let schema = record_batches[0].schema();
1026        assert_eq!(
1027            schema.field_with_name("geometry").unwrap().metadata()["ARROW:extension:name"],
1028            "geoarrow.wkb"
1029        );
1030    }
1031
1032    #[rstest]
1033    fn filter(client: Client) {
1034        let search = Search {
1035            items: Items {
1036                filter: Some("sat:relative_orbit = 98".parse().unwrap()),
1037                ..Default::default()
1038            },
1039            ..Default::default()
1040        };
1041        let item_collection = client
1042            .search("data/100-sentinel-2-items.parquet", search)
1043            .unwrap();
1044        assert_eq!(item_collection.items.len(), 49);
1045    }
1046
1047    #[rstest]
1048    fn filter_no_column(client: Client) {
1049        let search = Search {
1050            items: Items {
1051                filter: Some("foo:bar = 42".parse().unwrap()),
1052                ..Default::default()
1053            },
1054            ..Default::default()
1055        };
1056        let item_collection = client
1057            .search("data/100-sentinel-2-items.parquet", search)
1058            .unwrap();
1059        assert_eq!(item_collection.items.len(), 0);
1060    }
1061
1062    #[rstest]
1063    fn sortby_property(client: Client) {
1064        let search = Search {
1065            items: Items {
1066                sortby: vec!["eo:cloud_cover".parse().unwrap()],
1067                ..Default::default()
1068            },
1069            ..Default::default()
1070        };
1071        let item_collection = client
1072            .search("data/100-sentinel-2-items.parquet", search)
1073            .unwrap();
1074        assert_eq!(item_collection.items.len(), 100);
1075    }
1076
1077    #[rstest]
1078    fn union_by_name(client: Client) {
1079        let _ = client.search("data/*.parquet", Default::default()).unwrap();
1080    }
1081
1082    #[rstest]
1083    fn no_union_by_name(mut client: Client) {
1084        client.union_by_name = false;
1085        let _ = client
1086            .search("data/*.parquet", Default::default())
1087            .unwrap_err();
1088    }
1089
1090    #[rstest]
1091    fn remove_filename_column(client: Client) {
1092        let item_collection = client
1093            .search("data/100-sentinel-2-items.parquet", Default::default())
1094            .unwrap();
1095        for item in item_collection.items {
1096            assert!(
1097                !item["properties"]
1098                    .as_object()
1099                    .as_ref()
1100                    .unwrap()
1101                    .contains_key("filename")
1102            );
1103        }
1104    }
1105
1106    #[rstest]
1107    fn iceberg_local(mut client: Client) {
1108        client.set_source_format(SourceFormat::Iceberg).unwrap();
1109        let item_collection = client
1110            .search(
1111                "data/hls-iceberg/metadata/v1.metadata.json",
1112                Search::default(),
1113            )
1114            .unwrap();
1115        assert_eq!(item_collection.items.len(), 5);
1116        assert!(item_collection.items[0].get("geometry").is_some());
1117    }
1118}