1use 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" => is_age_catalog_function(local),
30 "pg_catalog" => {
31 crate::registry::is_registered(local)
32 || crate::type_resolution::is_fixed_builtin(local)
33 || is_pg_catalog_function(local)
34 }
35 _ => false,
36 };
37 if is_builtin {
38 local.to_string()
39 } else {
40 lower
41 }
42}
43
44fn is_age_catalog_function(local: &str) -> bool {
46 matches!(
47 local,
48 "cypher"
49 | "create_graph"
50 | "drop_graph"
51 | "graph_exists"
52 | "create_vlabel"
53 | "create_elabel"
54 | "drop_label"
55 | "alter_graph"
56 )
57}
58
59fn is_pg_catalog_function(local: &str) -> bool {
61 matches!(
62 local,
63 "generate_series"
64 | "unnest"
65 | "regexp_split_to_table"
66 | "string_to_table"
67 | "json_array_elements"
68 | "jsonb_array_elements"
69 | "json_array_elements_text"
70 | "jsonb_array_elements_text"
71 | "json_each"
72 | "jsonb_each"
73 | "json_each_text"
74 | "jsonb_each_text"
75 | "aclexplode"
76 | "json_object_keys"
77 | "jsonb_object_keys"
78 | "upper"
79 | "lower"
80 | "bit_length"
81 | "char_length"
82 | "character_length"
83 | "crc32"
84 | "crc32c"
85 | "gamma"
86 | "json_strip_nulls"
87 | "jsonb_strip_nulls"
88 | "length"
89 | "lgamma"
90 | "md5"
91 | "octet_length"
92 | "reverse"
93 | "random"
94 | "setseed"
95 | "nextval"
96 | "currval"
97 | "lastval"
98 | "setval"
99 | "current_schema"
100 | "current_schemas"
101 | "current_setting"
102 | "set_config"
103 | "pg_sleep"
104 | "pg_sleep_for"
105 | "pg_sleep_until"
106 | "pg_backend_pid"
107 | "pg_my_temp_schema"
108 | "pg_is_other_temp_schema"
109 | "pg_trigger_depth"
110 | "pg_listening_channels"
111 | "pg_notify"
112 | "pg_notification_queue_usage"
113 | "pg_get_expr"
114 | "pg_get_partkeydef"
115 | "pg_get_serial_sequence"
116 | "pg_get_sequence_data"
117 | "pg_sequence_last_value"
118 | "pg_sequence_parameters"
119 | "pg_get_triggerdef"
120 | "pg_get_ruledef"
121 | "pg_get_viewdef"
122 | "pg_get_indexdef"
123 | "pg_get_constraintdef"
124 | "pg_get_function_arguments"
125 | "pg_get_function_identity_arguments"
126 | "pg_get_function_result"
127 | "pg_get_functiondef"
128 | "pg_get_function_sqlbody"
129 | "format_type"
130 | "pg_describe_object"
131 | "pg_has_role"
132 | "pg_get_userbyid"
133 | "has_table_privilege"
134 | "has_column_privilege"
135 | "has_database_privilege"
136 | "has_schema_privilege"
137 | "has_sequence_privilege"
138 | "has_function_privilege"
139 | "has_type_privilege"
140 )
141}
142
143pub fn is_builtin_aggregate(expr: &ScalarExpr) -> bool {
144 matches!(expr, ScalarExpr::Func { name, binding, .. } if is_builtin_aggregate_call(name, binding.as_ref()))
145}
146
147pub fn is_builtin_aggregate_call(
149 name: &str,
150 binding: Option<&crate::ast::FunctionBinding>,
151) -> bool {
152 if binding.is_some_and(|binding| !binding.builtin || binding.dispatch.is_some()) {
153 return false;
154 }
155 let name = binding
156 .map_or(name, |binding| binding.name.as_str())
157 .to_ascii_lowercase();
158 let local = name.strip_prefix("pg_catalog.").unwrap_or(&name);
159 crate::ast::is_builtin_aggregate_function(local)
160}
161
162pub enum MultiFieldMatchShape<'a> {
163 FieldsThenQuery {
164 fields: Vec<&'a ScalarExpr>,
165 query_idx: usize,
166 },
167 Pairs {
168 fields: Vec<&'a ScalarExpr>,
169 },
170}
171
172pub fn multi_field_match_shape(args: &[ScalarExpr]) -> Result<MultiFieldMatchShape<'_>, SQLError> {
173 let first_non_column = args.iter().position(|arg| {
174 !matches!(
175 arg,
176 ScalarExpr::Column(_) | ScalarExpr::QualifiedColumn { .. }
177 )
178 });
179 if let Some(query_idx) = first_non_column {
180 if query_idx >= 2 {
181 return Ok(MultiFieldMatchShape::FieldsThenQuery {
182 fields: args[..query_idx].iter().collect(),
183 query_idx,
184 });
185 }
186 }
187 if args.len() < 4 || !args.len().is_multiple_of(2) {
188 if let Some(query_idx) = first_non_column {
189 if query_idx < 2 && args.len() >= 3 {
190 return Err(SQLError::TypeMismatch(format!(
191 "multi_field_match field arguments must be column references, \
192 but argument {} is an expression; store computed text in an \
193 indexed column instead of concatenating at query time",
194 query_idx + 1
195 )));
196 }
197 }
198 return Err(SQLError::BadArity {
199 name: "multi_field_match".into(),
200 expected: ">= 3 (fields..., query[, weights...]) or even >= 4 (field, query pairs)"
201 .into(),
202 actual: args.len(),
203 });
204 }
205 Ok(MultiFieldMatchShape::Pairs {
206 fields: (0..args.len() / 2).map(|i| &args[2 * i]).collect(),
207 })
208}
209
210pub fn is_semantic_field_argument(
212 function: &str,
213 args: &[ScalarExpr],
214 argument_index: usize,
215) -> Result<bool, SQLError> {
216 let dispatch_name = crate::semantics::builtin_function_dispatch_name(function);
217 let Some(kind) = crate::registry::lookup(&dispatch_name) else {
218 return Ok(false);
219 };
220 let is_field = match kind {
221 FunctionKind::TextMatch | FunctionKind::BayesianMatch | FunctionKind::KNNMatch => {
222 argument_index == 0
223 }
224 FunctionKind::FTSMatch => argument_index == 0 && !fts_query_is_jsonpath(args.get(1)),
225 FunctionKind::BayesianMatchWithPrior => matches!(argument_index, 0 | 2),
226 FunctionKind::CalibratedVectorMatch => argument_index == 0,
227 FunctionKind::MultiFieldMatch => match multi_field_match_shape(args)? {
228 MultiFieldMatchShape::FieldsThenQuery { query_idx, .. } => argument_index < query_idx,
229 MultiFieldMatchShape::Pairs { .. } => argument_index.is_multiple_of(2),
230 },
231 FunctionKind::StagedRetrieval => {
232 !matches!(args.first(), Some(ScalarExpr::Func { .. }))
233 && argument_index.is_multiple_of(3)
234 }
235 FunctionKind::UQAFacets => true,
236 FunctionKind::ScoreBM25 | FunctionKind::ScoreBayesianBM25 => {
237 args.len() == 2 && argument_index == 0
238 }
239 FunctionKind::FuseLogOdds
240 | FunctionKind::PositiveEvidencePool
241 | FunctionKind::BayesianEvidenceFusion
242 | FunctionKind::GraphPagerank
243 | FunctionKind::GraphHits
244 | FunctionKind::GraphBetweenness
245 | FunctionKind::GraphTraverse
246 | FunctionKind::GraphNeighbors
247 | FunctionKind::DeepPredict
248 | FunctionKind::UQAHighlight
249 | FunctionKind::TraverseMatch
250 | FunctionKind::TemporalTraverse
251 | FunctionKind::RPQ
252 | FunctionKind::GraphCreate
253 | FunctionKind::GraphDrop
254 | FunctionKind::GraphExists
255 | FunctionKind::GraphLabelCreate
256 | FunctionKind::GraphLabelDrop
257 | FunctionKind::GraphAlter
258 | FunctionKind::GraphEdges
259 | FunctionKind::AttentionFusion
260 | FunctionKind::LearnedFusion
261 | FunctionKind::SparseThreshold
262 | FunctionKind::DeepLearn
263 | FunctionKind::Convolve
264 | FunctionKind::Pool
265 | FunctionKind::Flatten
266 | FunctionKind::Dense
267 | FunctionKind::Softmax
268 | FunctionKind::Layer
269 | FunctionKind::Model => false,
270 };
271 Ok(is_field)
272}
273
274pub fn fts_query_is_jsonpath(query_arg: Option<&ScalarExpr>) -> bool {
278 matches!(
279 query_arg,
280 Some(ScalarExpr::Literal(Value::Str(path))) if path.trim_start().starts_with('$')
281 )
282}
283
284pub fn contains_retrieval(expression: &ScalarExpr) -> bool {
288 match expression {
289 ScalarExpr::Func {
290 name,
291 args,
292 order_by,
293 filter,
294 ..
295 } => {
296 retrieval_function(name)
297 || args.iter().any(contains_retrieval)
298 || order_by.iter().any(|order| contains_retrieval(&order.expr))
299 || filter.as_deref().is_some_and(contains_retrieval)
300 }
301 ScalarExpr::Array(items)
302 | ScalarExpr::Row(items)
303 | ScalarExpr::CompositeRow { items, .. }
304 | ScalarExpr::And(items)
305 | ScalarExpr::Or(items) => items.iter().any(contains_retrieval),
306 ScalarExpr::Binary { lhs, rhs, .. } => contains_retrieval(lhs) || contains_retrieval(rhs),
307 ScalarExpr::UnaryMinus(inner)
308 | ScalarExpr::Not(inner)
309 | ScalarExpr::IsNull { expr: inner, .. }
310 | ScalarExpr::Cast { expr: inner, .. } => contains_retrieval(inner),
311 ScalarExpr::Between { expr, low, high } => {
312 contains_retrieval(expr) || contains_retrieval(low) || contains_retrieval(high)
313 }
314 ScalarExpr::InList { expr, list, .. } => {
315 contains_retrieval(expr) || list.iter().any(contains_retrieval)
316 }
317 ScalarExpr::WindowCall {
318 args, spec, filter, ..
319 } => {
320 args.iter().any(contains_retrieval)
321 || filter.as_deref().is_some_and(contains_retrieval)
322 || spec.partition_by.iter().any(contains_retrieval)
323 || spec
324 .order_by
325 .iter()
326 .any(|order| contains_retrieval(&order.expr))
327 }
328 ScalarExpr::Case {
329 base,
330 when,
331 else_branch,
332 } => {
333 base.as_deref().is_some_and(contains_retrieval)
334 || when.iter().any(|(condition, result)| {
335 contains_retrieval(condition) || contains_retrieval(result)
336 })
337 || else_branch.as_deref().is_some_and(contains_retrieval)
338 }
339 ScalarExpr::InSubquery { expr, .. } => contains_retrieval(expr),
340 ScalarExpr::Default
341 | ScalarExpr::Star
342 | ScalarExpr::QualifiedStar(_)
343 | ScalarExpr::Column(_)
344 | ScalarExpr::Position(_)
345 | ScalarExpr::InternalColumn(_)
346 | ScalarExpr::QualifiedColumn { .. }
347 | ScalarExpr::Literal(_)
348 | ScalarExpr::TypedLiteral { .. }
349 | ScalarExpr::Param(_)
350 | ScalarExpr::ScalarSubquery(_)
351 | ScalarExpr::Exists { .. } => false,
352 }
353}
354
355pub fn retrieval_function(name: &str) -> bool {
356 matches!(
357 name.to_ascii_lowercase().as_str(),
358 "text_match"
359 | "bayesian_match"
360 | "fts_match"
361 | "bayesian_match_with_prior"
362 | "calibrated_vector_match"
363 | "knn_match"
364 | "fuse_log_odds"
365 | "pool_positive_evidence"
366 | "fuse_bayesian_evidence"
367 | "multi_field_match"
368 | "staged_retrieval"
369 | "attention"
370 | "fuse_attention"
371 | "fuse_multihead"
372 | "learned_fusion"
373 | "fuse_learned"
374 | "sparse_threshold"
375 | "graph_pagerank"
376 | "pagerank"
377 | "graph_hits"
378 | "hits"
379 | "graph_betweenness"
380 | "betweenness"
381 | "graph_traverse"
382 | "traverse_match"
383 | "graph_neighbors"
384 | "graph_edges"
385 | "temporal_traverse"
386 | "rpq"
387 | "deep_predict"
388 )
389}
390
391pub fn expect_column_name(expr: &ScalarExpr, label: &str) -> Result<String, SQLError> {
392 match expr {
393 ScalarExpr::Column(name) => Ok(name.clone()),
394 ScalarExpr::QualifiedColumn { column, .. } => Ok(column.clone()),
395 other => Err(SQLError::TypeMismatch(format!(
396 "{label} must be a column reference, got {other:?}"
397 ))),
398 }
399}
400
401#[cfg(test)]
402mod tests {
403 use super::builtin_function_dispatch_name;
404
405 #[test]
406 fn builtin_dispatch_preserves_schema_identity() {
407 for name in [
408 "pg_get_sequence_data",
409 "pg_sequence_last_value",
410 "pg_sequence_parameters",
411 "date_trunc",
412 "mod",
413 ] {
414 assert_eq!(builtin_function_dispatch_name(name), name);
415 assert_eq!(
416 builtin_function_dispatch_name(&format!("pg_catalog.{name}")),
417 name
418 );
419 assert_eq!(
420 builtin_function_dispatch_name(&format!("PG_CATALOG.{}", name.to_uppercase())),
421 name
422 );
423 let user_function = format!("public.{name}");
424 assert_eq!(
425 builtin_function_dispatch_name(&user_function),
426 user_function
427 );
428 }
429 assert_eq!(
430 builtin_function_dispatch_name("pg_catalog.custom_function"),
431 "pg_catalog.custom_function"
432 );
433 }
434}