1use 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 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}