Skip to main content

uqa_sql/semantics/
retrieval.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Retrieval argument binding and weight semantics, independent of physical search.
8
9use crate::{SQLError, ScalarExpr};
10use uqa_core::Value;
11
12pub type ArgumentEvaluator<'a> = dyn FnMut(&ScalarExpr) -> Result<Value, SQLError> + 'a;
13
14pub fn expect_string(
15    expr: &ScalarExpr,
16    name: &str,
17    evaluate: &mut ArgumentEvaluator<'_>,
18) -> Result<String, SQLError> {
19    expect_evaluated_string(evaluate(expr)?, name)
20}
21
22pub fn expect_evaluated_string(value: Value, name: &str) -> Result<String, SQLError> {
23    match value {
24        Value::Str(s) => Ok(s),
25        other => Err(SQLError::TypeMismatch(format!(
26            "{name} must be a string, got {other:?}"
27        ))),
28    }
29}
30
31use super::expect_column_name;
32
33pub fn expect_field_name_or_string(
34    expr: &ScalarExpr,
35    label: &str,
36    evaluate: &mut ArgumentEvaluator<'_>,
37) -> Result<String, SQLError> {
38    match expr {
39        ScalarExpr::Column(name) => Ok(name.clone()),
40        ScalarExpr::QualifiedColumn { column, .. } => Ok(column.clone()),
41        _ => expect_string(expr, label, evaluate),
42    }
43}
44
45pub fn expect_usize(
46    expr: &ScalarExpr,
47    label: &str,
48    evaluate: &mut ArgumentEvaluator<'_>,
49) -> Result<usize, SQLError> {
50    let v = evaluate(expr)?;
51    match v {
52        Value::Int(n) if n >= 0 => usize::try_from(n).map_err(|_| {
53            SQLError::TypeMismatch(format!("{label} exceeds the platform usize range"))
54        }),
55        Value::Int(_) => Err(SQLError::TypeMismatch(format!("{label} must be >= 0"))),
56        other => Err(SQLError::TypeMismatch(format!(
57            "{label} must be an integer, got {other:?}"
58        ))),
59    }
60}
61
62pub type MultiFieldMatchArgs = (Vec<String>, Vec<String>, Vec<f64>);
63
64use super::MultiFieldMatchShape;
65
66use super::multi_field_match_shape;
67
68pub fn parse_multi_field_match_args(
69    args: &[ScalarExpr],
70    evaluate: &mut ArgumentEvaluator<'_>,
71) -> Result<MultiFieldMatchArgs, SQLError> {
72    if args.len() < 3 {
73        return Err(SQLError::BadArity {
74            name: "multi_field_match".into(),
75            expected: ">= 3 (fields..., query[, weights...])".into(),
76            actual: args.len(),
77        });
78    }
79    match multi_field_match_shape(args)? {
80        MultiFieldMatchShape::FieldsThenQuery {
81            fields: field_args,
82            query_idx,
83        } => {
84            let fields = field_args
85                .into_iter()
86                .map(|arg| expect_column_name(arg, "multi_field_match.field"))
87                .collect::<Result<Vec<_>, _>>()?;
88            let query = expect_string_value(&args[query_idx], "multi_field_match.query", evaluate)?;
89            let weight_args = &args[query_idx + 1..];
90            let weights = if weight_args.is_empty() {
91                uniform_weights(fields.len())
92            } else {
93                if weight_args.len() != fields.len() {
94                    return Err(SQLError::BadArity {
95                        name: "multi_field_match".into(),
96                        expected: "one weight per field".into(),
97                        actual: weight_args.len(),
98                    });
99                }
100                normalize_weights(
101                    weight_args
102                        .iter()
103                        .map(|arg| expect_f64_value(arg, "multi_field_match.weight", evaluate))
104                        .collect::<Result<Vec<_>, _>>()?,
105                )?
106            };
107            let queries = vec![query; fields.len()];
108            Ok((fields, queries, weights))
109        }
110        MultiFieldMatchShape::Pairs { fields: field_args } => {
111            let n_fields = field_args.len();
112            let mut fields = Vec::with_capacity(n_fields);
113            let mut queries = Vec::with_capacity(n_fields);
114            for (i, field_arg) in field_args.into_iter().enumerate() {
115                fields.push(expect_column_name(field_arg, "multi_field_match.field")?);
116                queries.push(expect_string_value(
117                    &args[2 * i + 1],
118                    "multi_field_match.query",
119                    evaluate,
120                )?);
121            }
122            Ok((fields, queries, uniform_weights(n_fields)))
123        }
124    }
125}
126
127fn expect_string_value(
128    expr: &ScalarExpr,
129    label: &str,
130    evaluate: &mut ArgumentEvaluator<'_>,
131) -> Result<String, SQLError> {
132    match evaluate(expr)? {
133        Value::Str(s) => Ok(s),
134        other => Err(SQLError::TypeMismatch(format!(
135            "{label} must be string, got {other:?}"
136        ))),
137    }
138}
139
140pub fn expect_f64_value(
141    expr: &ScalarExpr,
142    label: &str,
143    evaluate: &mut ArgumentEvaluator<'_>,
144) -> Result<f64, SQLError> {
145    match evaluate(expr)? {
146        Value::Float(f) => Ok(f),
147        Value::Int(i) => Ok(i as f64),
148        Value::Decimal(d) => d
149            .to_f64()
150            .ok_or_else(|| SQLError::TypeMismatch(format!("{label} decimal is outside f64 range"))),
151        other => Err(SQLError::TypeMismatch(format!(
152            "{label} must be numeric, got {other:?}"
153        ))),
154    }
155}
156
157fn uniform_weights(n: usize) -> Vec<f64> {
158    vec![1.0 / n.max(1) as f64; n]
159}
160
161fn normalize_weights(weights: Vec<f64>) -> Result<Vec<f64>, SQLError> {
162    if weights
163        .iter()
164        .any(|weight| !weight.is_finite() || *weight < 0.0)
165    {
166        return Err(SQLError::TypeMismatch(
167            "multi_field_match weights must be non-negative and finite".into(),
168        ));
169    }
170    let total: f64 = weights.iter().sum();
171    if total > 0.0 {
172        Ok(weights.into_iter().map(|weight| weight / total).collect())
173    } else {
174        Err(SQLError::TypeMismatch(
175            "multi_field_match weights must have a positive sum".into(),
176        ))
177    }
178}
179
180/// Bind the field before the caller validates its catalog and evaluates the query.
181pub fn text_match_field(args: &[ScalarExpr], function_name: &str) -> Result<String, SQLError> {
182    if args.len() != 2 {
183        return Err(SQLError::BadArity {
184            name: function_name.into(),
185            expected: "2".into(),
186            actual: args.len(),
187        });
188    }
189    match &args[0] {
190        ScalarExpr::Column(name) => Ok(name.clone()),
191        ScalarExpr::QualifiedColumn { column, .. } => Ok(column.clone()),
192        ScalarExpr::Literal(Value::Str(s)) if s.is_empty() || s == "_all" => Ok("_all".into()),
193        other => Err(SQLError::TypeMismatch(format!(
194            "{function_name}.field must be a column reference, got {other:?}"
195        ))),
196    }
197}
198
199pub struct PriorMatchArguments {
200    pub field: String,
201    pub prior_field: String,
202    pub query: String,
203    pub mode: String,
204}
205
206pub fn prior_match_arguments(
207    args: &[ScalarExpr],
208    evaluate: &mut ArgumentEvaluator<'_>,
209) -> Result<PriorMatchArguments, SQLError> {
210    if args.len() != 4 {
211        return Err(SQLError::BadArity {
212            name: "bayesian_match_with_prior".into(),
213            expected: "4".into(),
214            actual: args.len(),
215        });
216    }
217    let field = expect_column_name(&args[0], "bayesian_match_with_prior.field")?;
218    let prior_field = expect_column_name(&args[2], "bayesian_match_with_prior.prior_field")?;
219    let query = expect_string(&args[1], "bayesian_match_with_prior.query", evaluate)?;
220    let mode = expect_string(&args[3], "bayesian_match_with_prior.mode", evaluate)?;
221    Ok(PriorMatchArguments {
222        field,
223        prior_field,
224        query,
225        mode,
226    })
227}
228
229pub struct CalibratedVectorArguments {
230    pub field: String,
231    pub query_vector: Vec<f32>,
232    pub k: usize,
233    pub threshold: Option<f64>,
234}
235
236pub fn calibrated_vector_arguments(
237    args: &[ScalarExpr],
238    evaluate: &mut ArgumentEvaluator<'_>,
239) -> Result<CalibratedVectorArguments, SQLError> {
240    if !(3..=4).contains(&args.len()) {
241        return Err(SQLError::BadArity {
242            name: "calibrated_vector_match".into(),
243            expected: "3..=4".into(),
244            actual: args.len(),
245        });
246    }
247    let field = expect_field_name_or_string(&args[0], "calibrated_vector_match.field", evaluate)?;
248    let query_vector = crate::expr::value_to_vector(&evaluate(&args[1])?)?;
249    let k = expect_usize(&args[2], "calibrated_vector_match.k", evaluate)?;
250    let threshold = args
251        .get(3)
252        .map(|arg| expect_f64_value(arg, "calibrated_vector_match.threshold", evaluate))
253        .transpose()?;
254    Ok(CalibratedVectorArguments {
255        field,
256        query_vector,
257        k,
258        threshold,
259    })
260}