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