uqa_sql/semantics/rules/
binding.rs1use crate::{
10 ast::Expr,
11 plpgsql::{ResolvedVariable, VariableResolver},
12 ResultRow as Document, SQLError,
13};
14use std::collections::BTreeMap;
15use uqa_core::{DocId, Value};
16
17pub trait RuleRowValues {
19 fn old_row(&self) -> Option<&Document>;
20 fn new_row(&self) -> Option<&Document>;
21 fn old_doc_id(&self) -> Option<DocId>;
22 fn new_doc_id(&self) -> Option<DocId>;
23}
24
25mod actions;
26pub use actions::{bind_insert_values_action, bind_set_oriented_action, BoundSetOrientedAction};
27
28pub struct RuleColumnMetadata {
29 pub ty: crate::ast::ColumnType,
30 pub uses_document_id: bool,
31 pub position: usize,
32}
33
34pub(super) struct RuntimeRuleResolver<'a> {
35 pub(super) old: Option<&'a Document>,
36 pub(super) new: Option<&'a Document>,
37 pub(super) old_doc_id: Option<DocId>,
38 pub(super) new_doc_id: Option<DocId>,
39 pub(super) columns: &'a BTreeMap<String, RuleColumnMetadata>,
40}
41
42impl RuntimeRuleResolver<'_> {
43 pub(super) fn record_field(
44 &self,
45 record: Option<&Document>,
46 doc_id: Option<DocId>,
47 column: &str,
48 ) -> Result<ResolvedVariable, SQLError> {
49 let metadata = self
50 .columns
51 .get(column)
52 .ok_or_else(|| SQLError::UnknownColumn(column.to_string()))?;
53 let value = if let Some(value) = record.and_then(|record| record.get(column).cloned()) {
54 value
55 } else if metadata.uses_document_id {
56 doc_id
57 .map(i64::try_from)
58 .transpose()
59 .map_err(|_| {
60 SQLError::TypeMismatch("document id exceeds PostgreSQL bigint".into())
61 })?
62 .map_or(Value::Null, Value::Int)
63 } else {
64 Value::Null
65 };
66 Ok(ResolvedVariable {
67 value,
68 declared_type: Some(metadata.ty.sql_name()),
69 })
70 }
71
72 pub(super) fn record(
73 &self,
74 record: Option<&Document>,
75 doc_id: Option<DocId>,
76 ) -> Result<ResolvedVariable, SQLError> {
77 let mut columns = self.columns.iter().collect::<Vec<_>>();
78 columns.sort_by_key(|(_, metadata)| metadata.position);
79 let fields = columns
80 .into_iter()
81 .map(|(column, _)| {
82 self.record_field(record, doc_id, column)
83 .map(|field| (column.clone(), field.value))
84 })
85 .collect::<Result<Vec<_>, _>>()?;
86 Ok(ResolvedVariable::untyped(Value::Record(fields)))
87 }
88}
89
90impl VariableResolver for RuntimeRuleResolver<'_> {
91 fn resolve_name(&mut self, name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
92 if name.eq_ignore_ascii_case("old") {
93 return self.record(self.old, self.old_doc_id).map(Some);
94 }
95 if name.eq_ignore_ascii_case("new") {
96 return self.record(self.new, self.new_doc_id).map(Some);
97 }
98 Ok(None)
99 }
100
101 fn resolve_qualified(
102 &mut self,
103 qualifier: &str,
104 column: &str,
105 ) -> Result<Option<ResolvedVariable>, SQLError> {
106 if qualifier.eq_ignore_ascii_case("old") {
107 return self
108 .record_field(self.old, self.old_doc_id, column)
109 .map(Some);
110 }
111 if qualifier.eq_ignore_ascii_case("new") {
112 return self
113 .record_field(self.new, self.new_doc_id, column)
114 .map(Some);
115 }
116 Ok(None)
117 }
118
119 fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
120 Ok(None)
121 }
122
123 fn rewrite_qualified_whole_row(&mut self, qualifier: &str) -> Result<Option<Expr>, SQLError> {
124 Ok(self
125 .resolve_name(qualifier)?
126 .map(|record| Expr::Literal(record.value)))
127 }
128}