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