Skip to main content

uqa_sql/retrieval/
binding.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Runtime argument binding and checked logical retrieval selection.
8
9use super::{
10    bind_operator_argument, checked_retrieval_call_tree_present, const_string, const_usize,
11    default_operator_graph, lower_document_boolean, lower_function, lower_where,
12    try_lower_attention_fusion, validate_checked_retrieval_call_tree,
13    validate_operator_function_arity, validate_probability_signal_contract, BindingResult,
14    RetrievalArguments, RetrievalConstants, RetrievalExpr, SQLError, ScalarExpr, Value,
15};
16
17pub fn lower_sql_function_bound(
18    source: &dyn RetrievalArguments,
19    name: &str,
20    args: &[ScalarExpr],
21    constants: &RetrievalConstants<'_>,
22) -> BindingResult<RetrievalExpr> {
23    validate_operator_function_arity(name, args.len())?;
24    validate_probability_signal_contract(name, args)?;
25    let mut bound = args
26        .iter()
27        .map(|argument| bind_operator_argument(source, argument, constants))
28        .collect::<Result<Vec<_>, _>>()?;
29    match name.to_ascii_lowercase().as_str() {
30        "rpq" if bound.len() == 2 => bound.push(ScalarExpr::Literal(Value::Str(
31            default_operator_graph(source, "rpq")?,
32        ))),
33        "graph_pagerank" | "pagerank" | "graph_hits" | "hits" | "graph_betweenness"
34        | "betweenness"
35            if bound.is_empty() =>
36        {
37            bound.push(ScalarExpr::Literal(Value::Str(default_operator_graph(
38                source, name,
39            )?)));
40        }
41        _ => {}
42    }
43    let empty_constants = constants.without_parameters();
44    validate_checked_retrieval_call_tree(name, &bound, &empty_constants)?;
45    if matches!(
46        name.to_ascii_lowercase().as_str(),
47        "attention" | "fuse_attention" | "fuse_multihead"
48    ) {
49        return try_lower_attention_fusion(name, &bound, &empty_constants);
50    }
51    lower_function(name, &bound, &empty_constants).ok_or_else(|| {
52        SQLError::TypeMismatch(format!(
53            "{name} arguments cannot be lowered to the shared operator IR"
54        ))
55    })
56}
57
58fn centrality_kind(name: &str) -> Option<&'static str> {
59    match name {
60        "graph_pagerank" | "pagerank" => Some("pagerank"),
61        "graph_hits" | "hits" => Some("hits"),
62        "graph_betweenness" | "betweenness" => Some("betweenness"),
63        _ => None,
64    }
65}
66
67fn lower_bound_centrality(
68    source: &dyn RetrievalArguments,
69    name: &str,
70    args: &[ScalarExpr],
71    kind: &str,
72) -> BindingResult<RetrievalExpr> {
73    let graph = match args {
74        [] => default_operator_graph(source, name)?,
75        [_] => {
76            return Err(SQLError::TypeMismatch(format!(
77                "{name}.graph must be a constant string"
78            )))
79        }
80        _ => {
81            return Err(SQLError::BadArity {
82                name: name.to_string(),
83                expected: "0..=1".into(),
84                actual: args.len(),
85            })
86        }
87    };
88    Ok(match kind {
89        "pagerank" => RetrievalExpr::PageRank { graph },
90        "hits" => RetrievalExpr::HITS { graph },
91        _ => RetrievalExpr::BetweennessCentrality { graph },
92    })
93}
94
95fn lower_bound_rpq(
96    source: &dyn RetrievalArguments,
97    args: &[ScalarExpr],
98    constants: &RetrievalConstants<'_>,
99) -> BindingResult<RetrievalExpr> {
100    let graph = default_operator_graph(source, "rpq")?;
101    let rpq_source = const_string(&args[0], constants)
102        .ok_or_else(|| SQLError::TypeMismatch("rpq.expr must be a constant string".into()))?;
103    let start_vertex = const_usize(&args[1], constants)
104        .and_then(|value| u64::try_from(value).ok())
105        .ok_or_else(|| SQLError::TypeMismatch("rpq.start must be a non-negative integer".into()))?;
106    Ok(RetrievalExpr::RegularPathQuery {
107        rpq_source,
108        start_vertex,
109        graph,
110    })
111}
112
113fn lower_bound_function(
114    source: &dyn RetrievalArguments,
115    name: &str,
116    args: &[ScalarExpr],
117    constants: &RetrievalConstants<'_>,
118) -> BindingResult<Option<RetrievalExpr>> {
119    validate_operator_function_arity(name, args.len())?;
120    validate_probability_signal_contract(name, args)?;
121
122    let bound;
123    let empty_constants = constants.without_parameters();
124    let (lowering_args, lowering_constants): (&[ScalarExpr], &RetrievalConstants<'_>) =
125        if checked_retrieval_call_tree_present(name, args) {
126            bound = args
127                .iter()
128                .map(|argument| bind_operator_argument(source, argument, constants))
129                .collect::<Result<Vec<_>, _>>()?;
130            (&bound, &empty_constants)
131        } else {
132            (args, constants)
133        };
134    validate_checked_retrieval_call_tree(name, lowering_args, lowering_constants)?;
135
136    if let Some(tree) = lower_function(name, lowering_args, lowering_constants) {
137        return Ok(Some(tree));
138    }
139    if matches!(
140        name.to_ascii_lowercase().as_str(),
141        "attention" | "fuse_attention" | "fuse_multihead"
142    ) {
143        return try_lower_attention_fusion(name, lowering_args, lowering_constants).map(Some);
144    }
145    let lower_name = name.to_ascii_lowercase();
146    if let Some(kind) = centrality_kind(&lower_name) {
147        return lower_bound_centrality(source, name, lowering_args, kind).map(Some);
148    }
149    if lower_name == "rpq" && lowering_args.len() == 2 {
150        return lower_bound_rpq(source, lowering_args, lowering_constants).map(Some);
151    }
152    if matches!(
153        lower_name.as_str(),
154        "graph_traverse"
155            | "traverse_match"
156            | "graph_neighbors"
157            | "graph_edges"
158            | "temporal_traverse"
159            | "rpq"
160            | "deep_predict"
161    ) {
162        return Err(SQLError::TypeMismatch(format!(
163            "{name} arguments must be execution-time constants of the documented types"
164        )));
165    }
166    Ok(None)
167}
168
169pub fn lower_where_bound(
170    source: &dyn RetrievalArguments,
171    expression: &ScalarExpr,
172    constants: &RetrievalConstants<'_>,
173) -> Result<Option<RetrievalExpr>, SQLError> {
174    match expression {
175        ScalarExpr::And(parts) => {
176            let mut children = Vec::with_capacity(parts.len());
177            for part in parts {
178                let Some(child) = lower_where_bound(source, part, constants)? else {
179                    return Ok(None);
180                };
181                children.push(child);
182            }
183            Ok(Some(lower_document_boolean(children, false)))
184        }
185        ScalarExpr::Or(parts) => {
186            let mut children = Vec::with_capacity(parts.len());
187            for part in parts {
188                let Some(child) = lower_where_bound(source, part, constants)? else {
189                    return Ok(None);
190                };
191                children.push(child);
192            }
193            Ok(Some(lower_document_boolean(children, true)))
194        }
195        ScalarExpr::Not(inner) if crate::semantics::expr_is_null_free(inner) => {
196            Ok(lower_where_bound(source, inner, constants)?
197                .map(|child| RetrievalExpr::Complement(Box::new(child))))
198        }
199        ScalarExpr::Func { name, args, .. } => lower_bound_function(source, name, args, constants),
200        _ => Ok(lower_where(expression, constants)),
201    }
202}