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 .filter(|doc_id| crate::semantics::key_identity::is_key_document_id(*doc_id))
58 .map(i64::try_from)
59 .transpose()
60 .map_err(|_| {
61 SQLError::TypeMismatch("document id exceeds PostgreSQL bigint".into())
62 })?
63 .map_or(Value::Null, Value::Int)
64 } else {
65 Value::Null
66 };
67 Ok(ResolvedVariable {
68 value,
69 declared_type: Some(metadata.ty.catalog_name()),
70 })
71 }
72
73 pub(super) fn record(
74 &self,
75 record: Option<&Document>,
76 doc_id: Option<DocId>,
77 ) -> Result<ResolvedVariable, SQLError> {
78 let mut columns = self.columns.iter().collect::<Vec<_>>();
79 columns.sort_by_key(|(_, metadata)| metadata.position);
80 let fields = columns
81 .into_iter()
82 .map(|(column, _)| {
83 self.record_field(record, doc_id, column)
84 .map(|field| (column.clone(), field.value))
85 })
86 .collect::<Result<Vec<_>, _>>()?;
87 Ok(ResolvedVariable::untyped(Value::Record(fields)))
88 }
89}
90
91impl VariableResolver for RuntimeRuleResolver<'_> {
92 fn resolve_name(&mut self, name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
93 if name.eq_ignore_ascii_case("old") {
94 return self.record(self.old, self.old_doc_id).map(Some);
95 }
96 if name.eq_ignore_ascii_case("new") {
97 return self.record(self.new, self.new_doc_id).map(Some);
98 }
99 Ok(None)
100 }
101
102 fn resolve_qualified(
103 &mut self,
104 qualifier: &str,
105 column: &str,
106 ) -> Result<Option<ResolvedVariable>, SQLError> {
107 if qualifier.eq_ignore_ascii_case("old") {
108 return self
109 .record_field(self.old, self.old_doc_id, column)
110 .map(Some);
111 }
112 if qualifier.eq_ignore_ascii_case("new") {
113 return self
114 .record_field(self.new, self.new_doc_id, column)
115 .map(Some);
116 }
117 Ok(None)
118 }
119
120 fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
121 Ok(None)
122 }
123
124 fn rewrite_qualified_whole_row(&mut self, qualifier: &str) -> Result<Option<Expr>, SQLError> {
125 Ok(self
126 .resolve_name(qualifier)?
127 .map(|record| Expr::Literal(record.value)))
128 }
129}