Skip to main content

paimon_datafusion/
hybrid_search.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! `hybrid_search` table-valued function for DataFusion.
19//!
20//! Spark-compatible shape:
21//! ```sql
22//! SELECT * FROM hybrid_search(
23//!   'table_name',
24//!   array(named_struct('field', 'embedding', 'query_vector', array(1.0, 0.0))),
25//!   array(named_struct('column', 'content', 'query', 'paimon')),
26//!   10,
27//!   'rrf')
28//! ```
29
30use std::collections::HashMap;
31use std::fmt::Debug;
32use std::sync::Arc;
33
34use async_trait::async_trait;
35use datafusion::arrow::array::Array;
36use datafusion::arrow::datatypes::{
37    DataType as ArrowDataType, Field, Schema, SchemaRef as ArrowSchemaRef,
38};
39use datafusion::catalog::{Session, TableFunctionImpl};
40use datafusion::common::{project_schema, ScalarValue};
41use datafusion::datasource::{TableProvider, TableType};
42use datafusion::error::{DataFusionError, Result as DFResult};
43use datafusion::logical_expr::{Expr, TableProviderFilterPushDown};
44use datafusion::physical_plan::empty::EmptyExec;
45use datafusion::physical_plan::ExecutionPlan;
46use datafusion::prelude::SessionContext;
47use paimon::catalog::Catalog;
48use paimon::spec::{
49    BigIntType, CoreOptions, DataField, DataType, ROW_ID_FIELD_ID, ROW_ID_FIELD_NAME,
50};
51use paimon::table::{HybridSearchRanker, HybridSearchRoute, Table};
52
53use crate::error::to_datafusion_error;
54use crate::physical_plan::{SearchScoreExec, SearchScoreOutputColumn};
55use crate::runtime::{await_with_runtime, block_on_with_runtime};
56use crate::table::{datafusion_read_fields, PaimonScanBuilder, PaimonTableProvider};
57use crate::table_function_args::{
58    extract_int_literal, extract_string_literal, parse_table_identifier,
59};
60use crate::table_loader::load_data_table_for_read;
61
62const FUNCTION_NAME: &str = "hybrid_search";
63const SEARCH_SCORE_COLUMN: &str = "__paimon_search_score";
64
65pub fn register_hybrid_search(
66    ctx: &SessionContext,
67    catalog: Arc<dyn Catalog>,
68    default_database: &str,
69) {
70    ctx.register_udf(
71        datafusion::functions_nested::make_array::make_array_udf()
72            .as_ref()
73            .clone()
74            .with_aliases(["array"]),
75    );
76    ctx.register_udtf(
77        FUNCTION_NAME,
78        Arc::new(HybridSearchFunction::new(catalog, default_database)),
79    );
80}
81
82pub struct HybridSearchFunction {
83    catalog: Arc<dyn Catalog>,
84    default_database: String,
85}
86
87impl Debug for HybridSearchFunction {
88    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
89        f.debug_struct("HybridSearchFunction")
90            .field("default_database", &self.default_database)
91            .finish()
92    }
93}
94
95impl HybridSearchFunction {
96    pub fn new(catalog: Arc<dyn Catalog>, default_database: &str) -> Self {
97        Self {
98            catalog,
99            default_database: default_database.to_string(),
100        }
101    }
102}
103
104impl TableFunctionImpl for HybridSearchFunction {
105    fn call(&self, args: &[Expr]) -> DFResult<Arc<dyn TableProvider>> {
106        if args.len() != 4 && args.len() != 5 {
107            return Err(DataFusionError::Plan(
108                "hybrid_search requires 4 or 5 arguments: (table_name, vector_routes, full_text_routes, limit[, ranker])".to_string(),
109            ));
110        }
111
112        let table_name = extract_string_literal(FUNCTION_NAME, &args[0], "table_name")?;
113        let limit = extract_int_literal(FUNCTION_NAME, &args[3], "limit")?;
114        if limit <= 0 {
115            return Err(DataFusionError::Plan(
116                "hybrid_search: limit must be positive".to_string(),
117            ));
118        }
119
120        let ranker = if args.len() == 5 {
121            extract_string_literal(FUNCTION_NAME, &args[4], "ranker")?
122        } else {
123            HybridSearchRanker::RRF.to_string()
124        };
125        HybridSearchRanker::parse(&ranker).map_err(to_datafusion_error)?;
126
127        let mut routes = parse_vector_routes(&args[1], limit as usize)?;
128        routes.extend(parse_full_text_routes(&args[2], limit as usize)?);
129
130        let identifier =
131            parse_table_identifier(FUNCTION_NAME, &table_name, &self.default_database)?;
132        let catalog = Arc::clone(&self.catalog);
133        let table = block_on_with_runtime(
134            async move { load_data_table_for_read(&catalog, &identifier, FUNCTION_NAME).await },
135            "hybrid_search: catalog access thread panicked",
136        )?;
137
138        Ok(Arc::new(HybridSearchTableProvider::try_new(
139            PaimonTableProvider::try_new(table)?,
140            routes,
141            limit as usize,
142            ranker,
143        )?))
144    }
145}
146
147#[derive(Debug)]
148struct HybridSearchTableProvider {
149    inner: PaimonTableProvider,
150    schema: ArrowSchemaRef,
151    routes: Vec<HybridSearchRoute>,
152    limit: usize,
153    ranker: String,
154}
155
156impl HybridSearchTableProvider {
157    fn try_new(
158        inner: PaimonTableProvider,
159        routes: Vec<HybridSearchRoute>,
160        limit: usize,
161        ranker: String,
162    ) -> DFResult<Self> {
163        let inner_schema = inner.schema();
164        if inner_schema
165            .fields()
166            .iter()
167            .any(|field| field.name() == SEARCH_SCORE_COLUMN)
168        {
169            return Err(DataFusionError::Plan(format!(
170                "hybrid_search: table already contains reserved column {SEARCH_SCORE_COLUMN}"
171            )));
172        }
173
174        let mut fields = inner_schema
175            .fields()
176            .iter()
177            .map(|field| field.as_ref().clone())
178            .collect::<Vec<_>>();
179        fields.push(Field::new(
180            SEARCH_SCORE_COLUMN,
181            ArrowDataType::Float32,
182            true,
183        ));
184        let schema = Arc::new(Schema::new_with_metadata(
185            fields,
186            inner_schema.metadata().clone(),
187        ));
188
189        Ok(Self {
190            inner,
191            schema,
192            routes,
193            limit,
194            ranker,
195        })
196    }
197}
198
199#[async_trait]
200impl TableProvider for HybridSearchTableProvider {
201    fn schema(&self) -> ArrowSchemaRef {
202        self.schema.clone()
203    }
204
205    fn table_type(&self) -> TableType {
206        TableType::Base
207    }
208
209    async fn scan(
210        &self,
211        state: &dyn Session,
212        projection: Option<&Vec<usize>>,
213        _filters: &[Expr],
214        _limit: Option<usize>,
215    ) -> DFResult<Arc<dyn ExecutionPlan>> {
216        let table = self.inner.table();
217
218        let search_result = await_with_runtime(async {
219            let mut builder = table.new_hybrid_search_builder();
220            for route in self.routes.clone() {
221                builder.add_route(route);
222            }
223            builder
224                .with_limit(self.limit)
225                .with_ranker(&self.ranker)
226                .map_err(to_datafusion_error)?;
227            builder.execute_scored().await.map_err(to_datafusion_error)
228        })
229        .await?;
230
231        let row_ranges = search_result.to_row_ranges().map_err(to_datafusion_error)?;
232
233        if row_ranges.is_empty() {
234            let schema = project_schema(&self.schema, projection)?;
235            return Ok(Arc::new(EmptyExec::new(schema)));
236        }
237
238        // Raw row-range scans can include non-matching rows from selected files,
239        // so the outer limit must remain above the row-ID filter.
240        let scan = table
241            .new_read_builder()
242            .new_scan()
243            .with_row_ranges(row_ranges);
244        let plan = await_with_runtime(scan.plan())
245            .await
246            .map_err(to_datafusion_error)?;
247
248        let inner_schema = self.inner.schema();
249        let score_index = inner_schema.fields().len();
250        let input_read_fields = search_read_fields(table)?;
251        let input_schema = paimon::arrow::build_target_arrow_schema(&input_read_fields)
252            .map_err(to_datafusion_error)?;
253        let row_id_table_index = input_read_fields
254            .iter()
255            .position(|field| field.name() == ROW_ID_FIELD_NAME)
256            .expect("search read fields contain _ROW_ID");
257        let projected_indices = projection
258            .cloned()
259            .unwrap_or_else(|| (0..self.schema.fields().len()).collect());
260        let mut input_projection = projected_indices
261            .iter()
262            .copied()
263            .filter(|index| *index != score_index)
264            .collect::<Vec<_>>();
265        if !input_projection.contains(&row_id_table_index) {
266            input_projection.push(row_id_table_index);
267        }
268        let row_id_input_index = input_projection
269            .iter()
270            .position(|index| *index == row_id_table_index)
271            .expect("row ID was added to the input projection");
272        let output_columns = projected_indices
273            .iter()
274            .map(|index| {
275                if *index == score_index {
276                    SearchScoreOutputColumn::Score
277                } else {
278                    let input_index = input_projection
279                        .iter()
280                        .position(|input_index| input_index == index)
281                        .expect("projected table column exists in the input projection");
282                    SearchScoreOutputColumn::Input(input_index)
283                }
284            })
285            .collect();
286        let output_schema = project_schema(&self.schema, projection)?;
287        let input = PaimonScanBuilder {
288            table,
289            schema: &input_schema,
290            plan: &plan,
291            scan_trace: None,
292            projection: Some(&input_projection),
293            pushed_predicate: None,
294            limit: None,
295            target_partitions: state.config_options().execution.target_partitions,
296            filter_exact: false,
297            case_sensitive: true,
298        }
299        .build_with_read_fields(input_read_fields)?;
300        let scores = search_result
301            .row_ids
302            .into_iter()
303            .zip(search_result.scores)
304            .collect();
305        Ok(Arc::new(SearchScoreExec::new(
306            input,
307            output_schema,
308            row_id_input_index,
309            output_columns,
310            Arc::new(scores),
311        )))
312    }
313
314    fn supports_filters_pushdown(
315        &self,
316        filters: &[&Expr],
317    ) -> DFResult<Vec<TableProviderFilterPushDown>> {
318        Ok(vec![
319            TableProviderFilterPushDown::Unsupported;
320            filters.len()
321        ])
322    }
323}
324
325fn search_read_fields(table: &Table) -> DFResult<Vec<DataField>> {
326    let mut fields = datafusion_read_fields(table);
327    if fields.iter().any(|field| field.name() == ROW_ID_FIELD_NAME) {
328        return Ok(fields);
329    }
330    if !CoreOptions::new(table.schema().options()).row_tracking_enabled() {
331        return Err(DataFusionError::Plan(
332            "hybrid_search: cannot materialize search results because _ROW_ID is not available"
333                .to_string(),
334        ));
335    }
336    fields.push(DataField::new(
337        ROW_ID_FIELD_ID,
338        ROW_ID_FIELD_NAME.to_string(),
339        DataType::BigInt(BigIntType::with_nullable(true)),
340    ));
341    Ok(fields)
342}
343
344fn parse_vector_routes(expr: &Expr, default_limit: usize) -> DFResult<Vec<HybridSearchRoute>> {
345    if let Some(routes) = extract_literal_array_values(expr, "vector_routes")? {
346        return routes
347            .iter()
348            .map(|route| parse_vector_route_scalar(route, default_limit))
349            .collect();
350    }
351
352    extract_array_elements(expr, "vector_routes")?
353        .into_iter()
354        .map(|route| parse_vector_route(route, default_limit))
355        .collect()
356}
357
358fn parse_full_text_routes(expr: &Expr, default_limit: usize) -> DFResult<Vec<HybridSearchRoute>> {
359    if let Some(routes) = extract_literal_array_values(expr, "full_text_routes")? {
360        return routes
361            .iter()
362            .map(|route| parse_full_text_route_scalar(route, default_limit))
363            .collect();
364    }
365
366    extract_array_elements(expr, "full_text_routes")?
367        .into_iter()
368        .map(|route| parse_full_text_route(route, default_limit))
369        .collect()
370}
371
372fn parse_vector_route(expr: &Expr, default_limit: usize) -> DFResult<HybridSearchRoute> {
373    let fields = extract_named_struct_fields(expr, "vector route")?;
374    let field_name = optional_field(&fields, &["field", "vector_column"])
375        .ok_or_else(|| {
376            DataFusionError::Plan(
377                "hybrid_search: vector route must define field or vector_column".to_string(),
378            )
379        })
380        .and_then(|expr| extract_string_literal(FUNCTION_NAME, expr, "vector route field"))?;
381    let vector = required_field(&fields, "query_vector")
382        .and_then(|expr| extract_float_array(expr, "query_vector"))?;
383    let limit = optional_field(&fields, &["limit"])
384        .map(|expr| extract_positive_usize(expr, "vector route limit"))
385        .transpose()?
386        .unwrap_or(default_limit);
387    let weight = optional_field(&fields, &["weight"])
388        .map(|expr| extract_positive_f32(expr, "weight"))
389        .transpose()?
390        .unwrap_or(1.0);
391    let options = optional_field(&fields, &["options"])
392        .map(extract_options)
393        .transpose()?
394        .unwrap_or_default();
395
396    HybridSearchRoute::vector(field_name, vector, limit, weight, options)
397        .map_err(to_datafusion_error)
398}
399
400fn parse_vector_route_scalar(
401    scalar: &ScalarValue,
402    default_limit: usize,
403) -> DFResult<HybridSearchRoute> {
404    let fields = extract_struct_scalar_fields(scalar, "vector route")?;
405    let field_name = optional_scalar_field(&fields, &["field", "vector_column"])
406        .ok_or_else(|| {
407            DataFusionError::Plan(
408                "hybrid_search: vector route must define field or vector_column".to_string(),
409            )
410        })
411        .and_then(|scalar| scalar_to_string(scalar, "vector route field"))?;
412    let vector = required_scalar_field(&fields, "query_vector")
413        .and_then(|scalar| scalar_to_float_array(scalar, "query_vector"))?;
414    let limit = optional_scalar_field(&fields, &["limit"])
415        .map(|scalar| scalar_to_positive_usize(scalar, "vector route limit"))
416        .transpose()?
417        .unwrap_or(default_limit);
418    let weight = optional_scalar_field(&fields, &["weight"])
419        .map(|scalar| scalar_to_positive_f32(scalar, "weight"))
420        .transpose()?
421        .unwrap_or(1.0);
422    let options = optional_scalar_field(&fields, &["options"])
423        .map(scalar_to_options)
424        .transpose()?
425        .unwrap_or_default();
426
427    HybridSearchRoute::vector(field_name, vector, limit, weight, options)
428        .map_err(to_datafusion_error)
429}
430
431fn parse_full_text_route(expr: &Expr, default_limit: usize) -> DFResult<HybridSearchRoute> {
432    let fields = extract_named_struct_fields(expr, "full-text route")?;
433    let column = required_field(&fields, "column")
434        .and_then(|expr| extract_string_literal(FUNCTION_NAME, expr, "full-text route column"))?;
435    let query = required_field(&fields, "query")
436        .and_then(|expr| extract_string_literal(FUNCTION_NAME, expr, "full-text route query"))?;
437    let limit = optional_field(&fields, &["limit"])
438        .map(|expr| extract_positive_usize(expr, "full-text route limit"))
439        .transpose()?
440        .unwrap_or(default_limit);
441    let weight = optional_field(&fields, &["weight"])
442        .map(|expr| extract_positive_f32(expr, "weight"))
443        .transpose()?
444        .unwrap_or(1.0);
445    let options = optional_field(&fields, &["options"])
446        .map(extract_options)
447        .transpose()?
448        .unwrap_or_default();
449
450    HybridSearchRoute::full_text(column, query, limit, weight, options).map_err(to_datafusion_error)
451}
452
453fn parse_full_text_route_scalar(
454    scalar: &ScalarValue,
455    default_limit: usize,
456) -> DFResult<HybridSearchRoute> {
457    let fields = extract_struct_scalar_fields(scalar, "full-text route")?;
458    let column = required_scalar_field(&fields, "column")
459        .and_then(|scalar| scalar_to_string(scalar, "full-text route column"))?;
460    let query = required_scalar_field(&fields, "query")
461        .and_then(|scalar| scalar_to_string(scalar, "full-text route query"))?;
462    let limit = optional_scalar_field(&fields, &["limit"])
463        .map(|scalar| scalar_to_positive_usize(scalar, "full-text route limit"))
464        .transpose()?
465        .unwrap_or(default_limit);
466    let weight = optional_scalar_field(&fields, &["weight"])
467        .map(|scalar| scalar_to_positive_f32(scalar, "weight"))
468        .transpose()?
469        .unwrap_or(1.0);
470    let options = optional_scalar_field(&fields, &["options"])
471        .map(scalar_to_options)
472        .transpose()?
473        .unwrap_or_default();
474
475    HybridSearchRoute::full_text(column, query, limit, weight, options).map_err(to_datafusion_error)
476}
477
478fn extract_array_elements<'a>(expr: &'a Expr, name: &str) -> DFResult<Vec<&'a Expr>> {
479    match expr {
480        Expr::ScalarFunction(function)
481            if is_function(function.name(), &["make_array", "array"]) =>
482        {
483            Ok(function.args.iter().collect())
484        }
485        _ => Err(DataFusionError::Plan(format!(
486            "hybrid_search: {name} must be array(...), got: {expr}"
487        ))),
488    }
489}
490
491fn extract_literal_array_values(expr: &Expr, name: &str) -> DFResult<Option<Vec<ScalarValue>>> {
492    let Expr::Literal(scalar, _) = expr else {
493        return Ok(None);
494    };
495    scalar_array_values(scalar, name).map(Some)
496}
497
498fn scalar_array_values(scalar: &ScalarValue, name: &str) -> DFResult<Vec<ScalarValue>> {
499    let values = match scalar {
500        ScalarValue::List(array) => array.value(0),
501        ScalarValue::LargeList(array) => array.value(0),
502        ScalarValue::ListView(array) => array.value(0),
503        ScalarValue::LargeListView(array) => array.value(0),
504        ScalarValue::FixedSizeList(array) => array.value(0),
505        _ => {
506            return Err(DataFusionError::Plan(format!(
507                "hybrid_search: {name} must be an array, got: {scalar}"
508            )));
509        }
510    };
511
512    (0..values.len())
513        .map(|index| ScalarValue::try_from_array(values.as_ref(), index))
514        .collect()
515}
516
517fn extract_named_struct_fields<'a>(
518    expr: &'a Expr,
519    name: &str,
520) -> DFResult<Vec<(String, &'a Expr)>> {
521    let Expr::ScalarFunction(function) = expr else {
522        return Err(DataFusionError::Plan(format!(
523            "hybrid_search: {name} must be named_struct(...), got: {expr}"
524        )));
525    };
526    if !is_function(function.name(), &["named_struct"]) {
527        return Err(DataFusionError::Plan(format!(
528            "hybrid_search: {name} must be named_struct(...), got: {expr}"
529        )));
530    }
531    if function.args.len() % 2 != 0 {
532        return Err(DataFusionError::Plan(format!(
533            "hybrid_search: {name} must contain key/value pairs"
534        )));
535    }
536
537    let mut fields = Vec::with_capacity(function.args.len() / 2);
538    for pair in function.args.chunks_exact(2) {
539        let key = extract_string_literal(FUNCTION_NAME, &pair[0], "route field name")?;
540        fields.push((key, &pair[1]));
541    }
542    Ok(fields)
543}
544
545fn extract_struct_scalar_fields(
546    scalar: &ScalarValue,
547    name: &str,
548) -> DFResult<Vec<(String, ScalarValue)>> {
549    let ScalarValue::Struct(array) = scalar else {
550        return Err(DataFusionError::Plan(format!(
551            "hybrid_search: {name} must be named_struct(...), got: {scalar}"
552        )));
553    };
554    if array.is_null(0) {
555        return Err(DataFusionError::Plan(format!(
556            "hybrid_search: {name} cannot be null"
557        )));
558    }
559
560    array
561        .fields()
562        .iter()
563        .zip(array.columns())
564        .map(|(field, column)| {
565            Ok((
566                field.name().clone(),
567                ScalarValue::try_from_array(column.as_ref(), 0)?,
568            ))
569        })
570        .collect()
571}
572
573fn required_field<'a>(fields: &'a [(String, &'a Expr)], name: &str) -> DFResult<&'a Expr> {
574    optional_field(fields, &[name])
575        .ok_or_else(|| DataFusionError::Plan(format!("hybrid_search: route must define {name}")))
576}
577
578fn optional_field<'a>(fields: &'a [(String, &'a Expr)], names: &[&str]) -> Option<&'a Expr> {
579    fields
580        .iter()
581        .find(|(field_name, _)| names.iter().any(|name| field_name == name))
582        .map(|(_, expr)| *expr)
583}
584
585fn required_scalar_field<'a>(
586    fields: &'a [(String, ScalarValue)],
587    name: &str,
588) -> DFResult<&'a ScalarValue> {
589    optional_scalar_field(fields, &[name])
590        .ok_or_else(|| DataFusionError::Plan(format!("hybrid_search: route must define {name}")))
591}
592
593fn optional_scalar_field<'a>(
594    fields: &'a [(String, ScalarValue)],
595    names: &[&str],
596) -> Option<&'a ScalarValue> {
597    fields
598        .iter()
599        .find(|(field_name, _)| names.iter().any(|name| field_name == name))
600        .map(|(_, scalar)| scalar)
601}
602
603fn extract_positive_usize(expr: &Expr, name: &str) -> DFResult<usize> {
604    let value = extract_int_literal(FUNCTION_NAME, expr, name)?;
605    if value <= 0 {
606        return Err(DataFusionError::Plan(format!(
607            "hybrid_search: {name} must be positive"
608        )));
609    }
610    Ok(value as usize)
611}
612
613fn extract_float_array(expr: &Expr, name: &str) -> DFResult<Vec<f32>> {
614    if let Ok(json) = extract_string_literal(FUNCTION_NAME, expr, name) {
615        let vector: Vec<f32> = serde_json::from_str(&json).map_err(|e| {
616            DataFusionError::Plan(format!(
617                "hybrid_search: {name} string must be a JSON array of floats: {e}"
618            ))
619        })?;
620        if vector.is_empty() {
621            return Err(DataFusionError::Plan(format!(
622                "hybrid_search: {name} cannot be empty"
623            )));
624        }
625        return Ok(vector);
626    }
627
628    let elements = extract_array_elements(expr, name)?;
629    if elements.is_empty() {
630        return Err(DataFusionError::Plan(format!(
631            "hybrid_search: {name} cannot be empty"
632        )));
633    }
634    elements
635        .into_iter()
636        .map(|expr| scalar_to_f32(expr, name))
637        .collect()
638}
639
640fn extract_positive_f32(expr: &Expr, name: &str) -> DFResult<f32> {
641    let value = scalar_to_f32(expr, name)?;
642    if !value.is_finite() || value <= 0.0 {
643        return Err(DataFusionError::Plan(format!(
644            "hybrid_search: {name} must be finite and positive, got: {value}"
645        )));
646    }
647    Ok(value)
648}
649
650fn scalar_to_f32(expr: &Expr, name: &str) -> DFResult<f32> {
651    let Expr::Literal(scalar, _) = expr else {
652        return Err(DataFusionError::Plan(format!(
653            "hybrid_search: {name} must be a numeric literal, got: {expr}"
654        )));
655    };
656    match scalar {
657        ScalarValue::Float32(Some(value)) => Ok(*value),
658        ScalarValue::Float64(Some(value)) => Ok(*value as f32),
659        ScalarValue::Int8(Some(value)) => Ok(*value as f32),
660        ScalarValue::Int16(Some(value)) => Ok(*value as f32),
661        ScalarValue::Int32(Some(value)) => Ok(*value as f32),
662        ScalarValue::Int64(Some(value)) => Ok(*value as f32),
663        ScalarValue::UInt8(Some(value)) => Ok(*value as f32),
664        ScalarValue::UInt16(Some(value)) => Ok(*value as f32),
665        ScalarValue::UInt32(Some(value)) => Ok(*value as f32),
666        ScalarValue::UInt64(Some(value)) => Ok(*value as f32),
667        ScalarValue::Utf8(Some(value)) => value.parse::<f32>().map_err(|e| {
668            DataFusionError::Plan(format!(
669                "hybrid_search: {name} string must be a float, got '{value}': {e}"
670            ))
671        }),
672        _ => Err(DataFusionError::Plan(format!(
673            "hybrid_search: {name} must be a numeric literal, got: {expr}"
674        ))),
675    }
676}
677
678fn scalar_to_string(scalar: &ScalarValue, name: &str) -> DFResult<String> {
679    match scalar {
680        ScalarValue::Utf8(Some(value))
681        | ScalarValue::Utf8View(Some(value))
682        | ScalarValue::LargeUtf8(Some(value)) => Ok(value.clone()),
683        _ => Err(DataFusionError::Plan(format!(
684            "hybrid_search: {name} must be a string literal, got: {scalar}"
685        ))),
686    }
687}
688
689fn scalar_to_positive_usize(scalar: &ScalarValue, name: &str) -> DFResult<usize> {
690    let value = match scalar {
691        ScalarValue::Int8(Some(value)) => *value as i64,
692        ScalarValue::Int16(Some(value)) => *value as i64,
693        ScalarValue::Int32(Some(value)) => *value as i64,
694        ScalarValue::Int64(Some(value)) => *value,
695        ScalarValue::UInt8(Some(value)) => *value as i64,
696        ScalarValue::UInt16(Some(value)) => *value as i64,
697        ScalarValue::UInt32(Some(value)) => *value as i64,
698        ScalarValue::UInt64(Some(value)) => i64::try_from(*value).map_err(|_| {
699            DataFusionError::Plan(format!("hybrid_search: {name} value exceeds i64 range"))
700        })?,
701        _ => {
702            return Err(DataFusionError::Plan(format!(
703                "hybrid_search: {name} must be an integer literal, got: {scalar}"
704            )));
705        }
706    };
707    if value <= 0 {
708        return Err(DataFusionError::Plan(format!(
709            "hybrid_search: {name} must be positive"
710        )));
711    }
712    Ok(value as usize)
713}
714
715fn scalar_value_to_f32(scalar: &ScalarValue, name: &str) -> DFResult<f32> {
716    match scalar {
717        ScalarValue::Float16(Some(value)) => Ok(value.to_f32()),
718        ScalarValue::Float32(Some(value)) => Ok(*value),
719        ScalarValue::Float64(Some(value)) => Ok(*value as f32),
720        ScalarValue::Int8(Some(value)) => Ok(*value as f32),
721        ScalarValue::Int16(Some(value)) => Ok(*value as f32),
722        ScalarValue::Int32(Some(value)) => Ok(*value as f32),
723        ScalarValue::Int64(Some(value)) => Ok(*value as f32),
724        ScalarValue::UInt8(Some(value)) => Ok(*value as f32),
725        ScalarValue::UInt16(Some(value)) => Ok(*value as f32),
726        ScalarValue::UInt32(Some(value)) => Ok(*value as f32),
727        ScalarValue::UInt64(Some(value)) => Ok(*value as f32),
728        ScalarValue::Utf8(Some(value))
729        | ScalarValue::Utf8View(Some(value))
730        | ScalarValue::LargeUtf8(Some(value)) => value.parse::<f32>().map_err(|e| {
731            DataFusionError::Plan(format!(
732                "hybrid_search: {name} string must be a float, got '{value}': {e}"
733            ))
734        }),
735        _ => Err(DataFusionError::Plan(format!(
736            "hybrid_search: {name} must be a numeric literal, got: {scalar}"
737        ))),
738    }
739}
740
741fn scalar_to_positive_f32(scalar: &ScalarValue, name: &str) -> DFResult<f32> {
742    let value = scalar_value_to_f32(scalar, name)?;
743    if !value.is_finite() || value <= 0.0 {
744        return Err(DataFusionError::Plan(format!(
745            "hybrid_search: {name} must be finite and positive, got: {value}"
746        )));
747    }
748    Ok(value)
749}
750
751fn scalar_to_float_array(scalar: &ScalarValue, name: &str) -> DFResult<Vec<f32>> {
752    let values = scalar_array_values(scalar, name)?;
753    if values.is_empty() {
754        return Err(DataFusionError::Plan(format!(
755            "hybrid_search: {name} cannot be empty"
756        )));
757    }
758    values
759        .iter()
760        .map(|value| scalar_value_to_f32(value, name))
761        .collect()
762}
763
764fn scalar_to_options(scalar: &ScalarValue) -> DFResult<HashMap<String, String>> {
765    if matches!(scalar, ScalarValue::Null) {
766        return Ok(HashMap::new());
767    }
768
769    if let Ok(json) = scalar_to_string(scalar, "options") {
770        if json.trim().is_empty() {
771            return Ok(HashMap::new());
772        }
773        return serde_json::from_str(&json).map_err(|e| {
774            DataFusionError::Plan(format!(
775                "hybrid_search: options string must be a JSON object: {e}"
776            ))
777        });
778    }
779
780    let ScalarValue::Map(array) = scalar else {
781        return Err(DataFusionError::Plan(format!(
782            "hybrid_search: options must be map(...), got: {scalar}"
783        )));
784    };
785    if array.is_null(0) {
786        return Ok(HashMap::new());
787    }
788
789    let entries = array.value(0);
790    let keys = entries.column(0);
791    let values = entries.column(1);
792    (0..entries.len())
793        .map(|index| {
794            Ok((
795                scalar_to_string(
796                    &ScalarValue::try_from_array(keys.as_ref(), index)?,
797                    "options key",
798                )?,
799                scalar_to_string(
800                    &ScalarValue::try_from_array(values.as_ref(), index)?,
801                    "options value",
802                )?,
803            ))
804        })
805        .collect()
806}
807
808fn extract_options(expr: &Expr) -> DFResult<HashMap<String, String>> {
809    if let Ok(json) = extract_string_literal(FUNCTION_NAME, expr, "options") {
810        if json.trim().is_empty() {
811            return Ok(HashMap::new());
812        }
813        return serde_json::from_str(&json).map_err(|e| {
814            DataFusionError::Plan(format!(
815                "hybrid_search: options string must be a JSON object: {e}"
816            ))
817        });
818    }
819
820    let Expr::ScalarFunction(function) = expr else {
821        return Err(DataFusionError::Plan(format!(
822            "hybrid_search: options must be map(...), got: {expr}"
823        )));
824    };
825    if !is_function(function.name(), &["map", "make_map"]) {
826        return Err(DataFusionError::Plan(format!(
827            "hybrid_search: options must be map(...), got: {expr}"
828        )));
829    }
830    if function.args.is_empty() {
831        return Ok(HashMap::new());
832    }
833
834    if function.args.len() == 2
835        && is_array_expr(&function.args[0])
836        && is_array_expr(&function.args[1])
837    {
838        let keys = extract_array_elements(&function.args[0], "options keys")?;
839        let values = extract_array_elements(&function.args[1], "options values")?;
840        if keys.len() != values.len() {
841            return Err(DataFusionError::Plan(
842                "hybrid_search: options keys and values must have the same length".to_string(),
843            ));
844        }
845        return keys
846            .into_iter()
847            .zip(values)
848            .map(|(key, value)| {
849                Ok((
850                    extract_string_literal(FUNCTION_NAME, key, "options key")?,
851                    extract_string_literal(FUNCTION_NAME, value, "options value")?,
852                ))
853            })
854            .collect();
855    }
856
857    if function.args.len() % 2 != 0 {
858        return Err(DataFusionError::Plan(
859            "hybrid_search: options map must contain key/value pairs".to_string(),
860        ));
861    }
862
863    function
864        .args
865        .chunks_exact(2)
866        .map(|pair| {
867            Ok((
868                extract_string_literal(FUNCTION_NAME, &pair[0], "options key")?,
869                extract_string_literal(FUNCTION_NAME, &pair[1], "options value")?,
870            ))
871        })
872        .collect()
873}
874
875fn is_array_expr(expr: &Expr) -> bool {
876    matches!(
877        expr,
878        Expr::ScalarFunction(function) if is_function(function.name(), &["make_array", "array"])
879    )
880}
881
882fn is_function(actual: &str, expected: &[&str]) -> bool {
883    expected
884        .iter()
885        .any(|expected| actual.eq_ignore_ascii_case(expected))
886}