uqa_sql/semantics/
functions.rs1use crate::registry::FunctionKind;
10use crate::{SQLError, ScalarExpr};
11use std::sync::LazyLock;
12use uqa_core::Value;
13
14pub fn merge_action_attribute() -> crate::ast::InternalColumnRef {
16 static ATTRIBUTE: LazyLock<crate::ast::InternalColumnRef> =
17 LazyLock::new(|| crate::ast::InternalRelationId::allocate().column(0));
18 *ATTRIBUTE
19}
20pub fn builtin_function_dispatch_name(name: &str) -> String {
24 let lower = name.to_ascii_lowercase();
25 let Some((schema, local)) = lower.split_once('.') else {
26 return lower;
27 };
28 let is_builtin = match schema {
29 "ag_catalog" => matches!(
30 local,
31 "cypher"
32 | "create_graph"
33 | "drop_graph"
34 | "graph_exists"
35 | "create_vlabel"
36 | "create_elabel"
37 | "drop_label"
38 | "alter_graph"
39 ),
40 "pg_catalog" => {
41 crate::registry::is_registered(local)
42 || matches!(
43 local,
44 "generate_series"
45 | "unnest"
46 | "regexp_split_to_table"
47 | "string_to_table"
48 | "json_array_elements"
49 | "jsonb_array_elements"
50 | "json_array_elements_text"
51 | "jsonb_array_elements_text"
52 | "json_each"
53 | "jsonb_each"
54 | "json_each_text"
55 | "jsonb_each_text"
56 | "json_object_keys"
57 | "jsonb_object_keys"
58 | "upper"
59 | "lower"
60 | "bit_length"
61 | "char_length"
62 | "character_length"
63 | "crc32"
64 | "crc32c"
65 | "gamma"
66 | "json_strip_nulls"
67 | "jsonb_strip_nulls"
68 | "length"
69 | "lgamma"
70 | "md5"
71 | "octet_length"
72 | "reverse"
73 | "random"
74 | "setseed"
75 | "nextval"
76 | "currval"
77 | "lastval"
78 | "setval"
79 | "current_schema"
80 | "current_schemas"
81 | "pg_backend_pid"
82 | "pg_listening_channels"
83 | "pg_notify"
84 | "pg_notification_queue_usage"
85 | "pg_get_expr"
86 | "pg_get_partkeydef"
87 | "pg_get_serial_sequence"
88 | "pg_get_triggerdef"
89 | "pg_get_ruledef"
90 | "pg_get_viewdef"
91 | "pg_get_indexdef"
92 | "format_type"
93 | "pg_has_role"
94 | "has_table_privilege"
95 | "has_column_privilege"
96 | "has_database_privilege"
97 | "has_schema_privilege"
98 | "has_sequence_privilege"
99 )
100 }
101 _ => false,
102 };
103 if is_builtin {
104 local.to_string()
105 } else {
106 lower
107 }
108}
109
110pub fn is_builtin_aggregate(expr: &ScalarExpr) -> bool {
111 matches!(expr, ScalarExpr::Func { name, .. } if matches!(
112 name.to_ascii_lowercase().as_str(),
113 "count"
114 | "sum"
115 | "avg"
116 | "min"
117 | "max"
118 | "string_agg"
119 | "array_agg"
120 | "bool_and"
121 | "bool_or"
122 | "stddev"
123 | "stddev_samp"
124 | "stddev_pop"
125 | "variance"
126 | "var_samp"
127 | "var_pop"
128 | "percentile_cont"
129 | "percentile_disc"
130 | "mode"
131 | "json_agg"
132 | "jsonb_agg"
133 | "json_object_agg"
134 | "jsonb_object_agg"
135 ))
136}
137
138pub enum MultiFieldMatchShape<'a> {
139 FieldsThenQuery {
140 fields: Vec<&'a ScalarExpr>,
141 query_idx: usize,
142 },
143 Pairs {
144 fields: Vec<&'a ScalarExpr>,
145 },
146}
147
148pub fn multi_field_match_shape(args: &[ScalarExpr]) -> Result<MultiFieldMatchShape<'_>, SQLError> {
149 let first_non_column = args.iter().position(|arg| {
150 !matches!(
151 arg,
152 ScalarExpr::Column(_) | ScalarExpr::QualifiedColumn { .. }
153 )
154 });
155 if let Some(query_idx) = first_non_column {
156 if query_idx >= 2 {
157 return Ok(MultiFieldMatchShape::FieldsThenQuery {
158 fields: args[..query_idx].iter().collect(),
159 query_idx,
160 });
161 }
162 }
163 if args.len() < 4 || !args.len().is_multiple_of(2) {
164 if let Some(query_idx) = first_non_column {
165 if query_idx < 2 && args.len() >= 3 {
166 return Err(SQLError::TypeMismatch(format!(
167 "multi_field_match field arguments must be column references, \
168 but argument {} is an expression; store computed text in an \
169 indexed column instead of concatenating at query time",
170 query_idx + 1
171 )));
172 }
173 }
174 return Err(SQLError::BadArity {
175 name: "multi_field_match".into(),
176 expected: ">= 3 (fields..., query[, weights...]) or even >= 4 (field, query pairs)"
177 .into(),
178 actual: args.len(),
179 });
180 }
181 Ok(MultiFieldMatchShape::Pairs {
182 fields: (0..args.len() / 2).map(|i| &args[2 * i]).collect(),
183 })
184}
185
186pub fn is_semantic_field_argument(
188 function: &str,
189 args: &[ScalarExpr],
190 argument_index: usize,
191) -> Result<bool, SQLError> {
192 let dispatch_name = crate::semantics::builtin_function_dispatch_name(function);
193 let Some(kind) = crate::registry::lookup(&dispatch_name) else {
194 return Ok(false);
195 };
196 let is_field = match kind {
197 FunctionKind::TextMatch | FunctionKind::BayesianMatch | FunctionKind::KNNMatch => {
198 argument_index == 0
199 }
200 FunctionKind::FTSMatch => argument_index == 0 && !fts_query_is_jsonpath(args.get(1)),
201 FunctionKind::BayesianMatchWithPrior => matches!(argument_index, 0 | 2),
202 FunctionKind::CalibratedVectorMatch => argument_index == 0,
203 FunctionKind::MultiFieldMatch => match multi_field_match_shape(args)? {
204 MultiFieldMatchShape::FieldsThenQuery { query_idx, .. } => argument_index < query_idx,
205 MultiFieldMatchShape::Pairs { .. } => argument_index.is_multiple_of(2),
206 },
207 FunctionKind::StagedRetrieval => {
208 !matches!(args.first(), Some(ScalarExpr::Func { .. }))
209 && argument_index.is_multiple_of(3)
210 }
211 FunctionKind::UQAFacets => true,
212 FunctionKind::ScoreBM25 | FunctionKind::ScoreBayesianBM25 => {
213 args.len() == 2 && argument_index == 0
214 }
215 FunctionKind::FuseLogOdds
216 | FunctionKind::PositiveEvidencePool
217 | FunctionKind::BayesianEvidenceFusion
218 | FunctionKind::GraphPagerank
219 | FunctionKind::GraphHits
220 | FunctionKind::GraphBetweenness
221 | FunctionKind::GraphTraverse
222 | FunctionKind::GraphNeighbors
223 | FunctionKind::DeepPredict
224 | FunctionKind::UQAHighlight
225 | FunctionKind::TraverseMatch
226 | FunctionKind::TemporalTraverse
227 | FunctionKind::RPQ
228 | FunctionKind::GraphCreate
229 | FunctionKind::GraphDrop
230 | FunctionKind::GraphExists
231 | FunctionKind::GraphLabelCreate
232 | FunctionKind::GraphLabelDrop
233 | FunctionKind::GraphAlter
234 | FunctionKind::GraphEdges
235 | FunctionKind::AttentionFusion
236 | FunctionKind::LearnedFusion
237 | FunctionKind::SparseThreshold
238 | FunctionKind::DeepLearn
239 | FunctionKind::Convolve
240 | FunctionKind::Pool
241 | FunctionKind::Flatten
242 | FunctionKind::Dense
243 | FunctionKind::Softmax
244 | FunctionKind::Layer
245 | FunctionKind::Model => false,
246 };
247 Ok(is_field)
248}
249
250pub fn fts_query_is_jsonpath(query_arg: Option<&ScalarExpr>) -> bool {
254 matches!(
255 query_arg,
256 Some(ScalarExpr::Literal(Value::Str(path))) if path.trim_start().starts_with('$')
257 )
258}
259
260pub fn contains_retrieval(expression: &ScalarExpr) -> bool {
264 match expression {
265 ScalarExpr::Func {
266 name,
267 args,
268 order_by,
269 filter,
270 ..
271 } => {
272 retrieval_function(name)
273 || args.iter().any(contains_retrieval)
274 || order_by.iter().any(|order| contains_retrieval(&order.expr))
275 || filter.as_deref().is_some_and(contains_retrieval)
276 }
277 ScalarExpr::Array(items)
278 | ScalarExpr::Row(items)
279 | ScalarExpr::And(items)
280 | ScalarExpr::Or(items) => items.iter().any(contains_retrieval),
281 ScalarExpr::Binary { lhs, rhs, .. } => contains_retrieval(lhs) || contains_retrieval(rhs),
282 ScalarExpr::UnaryMinus(inner)
283 | ScalarExpr::Not(inner)
284 | ScalarExpr::IsNull { expr: inner, .. }
285 | ScalarExpr::Cast { expr: inner, .. } => contains_retrieval(inner),
286 ScalarExpr::Between { expr, low, high } => {
287 contains_retrieval(expr) || contains_retrieval(low) || contains_retrieval(high)
288 }
289 ScalarExpr::InList { expr, list, .. } => {
290 contains_retrieval(expr) || list.iter().any(contains_retrieval)
291 }
292 ScalarExpr::WindowCall { args, spec, .. } => {
293 args.iter().any(contains_retrieval)
294 || spec.partition_by.iter().any(contains_retrieval)
295 || spec
296 .order_by
297 .iter()
298 .any(|order| contains_retrieval(&order.expr))
299 }
300 ScalarExpr::Case {
301 base,
302 when,
303 else_branch,
304 } => {
305 base.as_deref().is_some_and(contains_retrieval)
306 || when.iter().any(|(condition, result)| {
307 contains_retrieval(condition) || contains_retrieval(result)
308 })
309 || else_branch.as_deref().is_some_and(contains_retrieval)
310 }
311 ScalarExpr::InSubquery { expr, .. } => contains_retrieval(expr),
312 ScalarExpr::Default
313 | ScalarExpr::Star
314 | ScalarExpr::QualifiedStar(_)
315 | ScalarExpr::Column(_)
316 | ScalarExpr::Position(_)
317 | ScalarExpr::InternalColumn(_)
318 | ScalarExpr::QualifiedColumn { .. }
319 | ScalarExpr::Literal(_)
320 | ScalarExpr::TypedLiteral { .. }
321 | ScalarExpr::Param(_)
322 | ScalarExpr::ScalarSubquery(_)
323 | ScalarExpr::Exists { .. } => false,
324 }
325}
326
327pub fn retrieval_function(name: &str) -> bool {
328 matches!(
329 name.to_ascii_lowercase().as_str(),
330 "text_match"
331 | "bayesian_match"
332 | "fts_match"
333 | "bayesian_match_with_prior"
334 | "calibrated_vector_match"
335 | "knn_match"
336 | "fuse_log_odds"
337 | "pool_positive_evidence"
338 | "fuse_bayesian_evidence"
339 | "multi_field_match"
340 | "staged_retrieval"
341 | "attention"
342 | "fuse_attention"
343 | "fuse_multihead"
344 | "learned_fusion"
345 | "fuse_learned"
346 | "sparse_threshold"
347 | "graph_pagerank"
348 | "pagerank"
349 | "graph_hits"
350 | "hits"
351 | "graph_betweenness"
352 | "betweenness"
353 | "graph_traverse"
354 | "traverse_match"
355 | "graph_neighbors"
356 | "graph_edges"
357 | "temporal_traverse"
358 | "rpq"
359 | "deep_predict"
360 )
361}
362
363pub fn expect_column_name(expr: &ScalarExpr, label: &str) -> Result<String, SQLError> {
364 match expr {
365 ScalarExpr::Column(name) => Ok(name.clone()),
366 ScalarExpr::QualifiedColumn { column, .. } => Ok(column.clone()),
367 other => Err(SQLError::TypeMismatch(format!(
368 "{label} must be a column reference, got {other:?}"
369 ))),
370 }
371}