reifydb_engine/vm/volcano/
assert.rs1use std::sync::Arc;
5
6use reifydb_core::value::column::{ColumnWithName, buffer::ColumnBuffer, columns::Columns, headers::ColumnHeaders};
7use reifydb_evaluate::expression::{context::EvalContext, eval::evaluate};
8use reifydb_rql::expression::{Expression, name::display_label};
9use reifydb_transaction::transaction::Transaction;
10use reifydb_value::reifydb_assertions;
11use tracing::instrument;
12
13use crate::{
14 Result,
15 error::EngineError,
16 vm::volcano::query::{QueryContext, QueryNode, eval_context_from_query},
17};
18
19pub(crate) struct AssertNode {
20 input: Box<dyn QueryNode>,
21 expressions: Vec<Expression>,
22 message: Option<String>,
23 context: Option<Arc<QueryContext>>,
24}
25
26impl AssertNode {
27 pub fn new(input: Box<dyn QueryNode>, expressions: Vec<Expression>, message: Option<String>) -> Self {
28 Self {
29 input,
30 expressions,
31 message,
32 context: None,
33 }
34 }
35
36 #[instrument(level = "trace", skip_all, name = "volcano::assert::eval")]
37 fn eval(
38 session: &EvalContext<'_>,
39 columns: &Columns,
40 row_count: usize,
41 assert_expr: &Expression,
42 ) -> Result<ColumnWithName> {
43 let eval_ctx = session.with_eval(columns.clone(), row_count);
44 evaluate(&eval_ctx, assert_expr)
45 }
46
47 #[instrument(level = "trace", skip_all, name = "volcano::assert::verify")]
48 fn verify(&self, data: &ColumnBuffer, row_count: usize, assert_expr: &Expression) -> Result<()> {
49 let frag = assert_expr.full_fragment_owned();
50 let label = display_label(assert_expr);
51 match data {
52 ColumnBuffer::Bool(container) => {
53 for i in 0..row_count {
54 let valid = container.is_defined(i);
55 let value = container.data().get(i);
56 if !valid || !value {
57 return Err(EngineError::AssertionFailed {
58 fragment: frag.clone(),
59 message: self.message.clone().unwrap_or_default(),
60 expression: Some(label.text().to_string()),
61 }
62 .into());
63 }
64 }
65 }
66 ColumnBuffer::Option {
67 inner,
68 bitvec,
69 } => match inner.as_ref() {
70 ColumnBuffer::Bool(container) => {
71 for i in 0..row_count {
72 let defined = i < bitvec.len() && bitvec.get(i);
73 let valid = defined && container.is_defined(i);
74 let value = valid && container.data().get(i);
75 if !value {
76 return Err(EngineError::AssertionFailed {
77 fragment: frag.clone(),
78 message: self.message.clone().unwrap_or_default(),
79 expression: Some(label.text().to_string()),
80 }
81 .into());
82 }
83 }
84 }
85 _ => {
86 return Err(EngineError::AssertionFailed {
87 fragment: frag.clone(),
88 message: "assert expression must evaluate to a boolean".to_string(),
89 expression: Some(label.text().to_string()),
90 }
91 .into());
92 }
93 },
94 _ => {
95 return Err(EngineError::AssertionFailed {
96 fragment: frag.clone(),
97 message: "assert expression must evaluate to a boolean".to_string(),
98 expression: Some(label.text().to_string()),
99 }
100 .into());
101 }
102 }
103 Ok(())
104 }
105}
106
107impl QueryNode for AssertNode {
108 #[instrument(level = "trace", skip_all, name = "volcano::assert::initialize")]
109 fn initialize<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &QueryContext) -> Result<()> {
110 self.context = Some(Arc::new(ctx.clone()));
111 self.input.initialize(rx, ctx)?;
112 Ok(())
113 }
114
115 #[instrument(level = "trace", skip_all, name = "volcano::assert::next")]
116 fn next<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &mut QueryContext) -> Result<Option<Columns>> {
117 reifydb_assertions! {
118 assert!(self.context.is_some(), "AssertNode::next() called before initialize()");
119 }
120 let stored_ctx = self.context.as_ref().unwrap();
121
122 if let Some(columns) = self.input.next(rx, ctx)? {
123 let row_count = columns.row_count();
124 let session = eval_context_from_query(stored_ctx);
125
126 for assert_expr in &self.expressions {
127 let result = Self::eval(&session, &columns, row_count, assert_expr)?;
128 self.verify(result.data(), row_count, assert_expr)?;
129 }
130
131 Ok(Some(columns))
132 } else {
133 Ok(None)
134 }
135 }
136
137 fn headers(&self) -> Option<ColumnHeaders> {
138 self.input.headers()
139 }
140}
141
142pub(crate) struct AssertWithoutInputNode {
143 expressions: Vec<Expression>,
144 message: Option<String>,
145 context: Option<Arc<QueryContext>>,
146 done: bool,
147}
148
149impl AssertWithoutInputNode {
150 pub fn new(expressions: Vec<Expression>, message: Option<String>) -> Self {
151 Self {
152 expressions,
153 message,
154 context: None,
155 done: false,
156 }
157 }
158
159 #[instrument(level = "trace", skip_all, name = "volcano::assert::noinput::eval")]
160 fn eval(session: &EvalContext<'_>, assert_expr: &Expression) -> Result<ColumnWithName> {
161 let eval_ctx = session.with_eval_empty();
162 evaluate(&eval_ctx, assert_expr)
163 }
164
165 #[instrument(level = "trace", skip_all, name = "volcano::assert::noinput::verify")]
166 fn verify(&self, data: &ColumnBuffer, assert_expr: &Expression) -> Result<()> {
167 let frag = assert_expr.full_fragment_owned();
168 let label = display_label(assert_expr);
169 match data {
170 ColumnBuffer::Bool(container) => {
171 let valid = container.is_defined(0);
172 let value = container.data().get(0);
173 if !valid || !value {
174 return Err(EngineError::AssertionFailed {
175 fragment: frag.clone(),
176 message: self.message.clone().unwrap_or_default(),
177 expression: Some(label.text().to_string()),
178 }
179 .into());
180 }
181 }
182 ColumnBuffer::Option {
183 inner,
184 bitvec,
185 } => match inner.as_ref() {
186 ColumnBuffer::Bool(container) => {
187 let defined = !bitvec.is_empty() && bitvec.get(0);
188 let valid = defined && container.is_defined(0);
189 let value = valid && container.data().get(0);
190 if !value {
191 return Err(EngineError::AssertionFailed {
192 fragment: frag.clone(),
193 message: self.message.clone().unwrap_or_default(),
194 expression: Some(label.text().to_string()),
195 }
196 .into());
197 }
198 }
199 _ => {
200 return Err(EngineError::AssertionFailed {
201 fragment: frag.clone(),
202 message: "assert expression must evaluate to a boolean".to_string(),
203 expression: Some(label.text().to_string()),
204 }
205 .into());
206 }
207 },
208 _ => {
209 return Err(EngineError::AssertionFailed {
210 fragment: frag.clone(),
211 message: "assert expression must evaluate to a boolean".to_string(),
212 expression: Some(label.text().to_string()),
213 }
214 .into());
215 }
216 }
217 Ok(())
218 }
219}
220
221impl QueryNode for AssertWithoutInputNode {
222 #[instrument(level = "trace", skip_all, name = "volcano::assert::noinput::initialize")]
223 fn initialize<'a>(&mut self, _rx: &mut Transaction<'a>, ctx: &QueryContext) -> Result<()> {
224 self.context = Some(Arc::new(ctx.clone()));
225 Ok(())
226 }
227
228 #[instrument(level = "trace", skip_all, name = "volcano::assert::noinput::next")]
229 fn next<'a>(&mut self, _rx: &mut Transaction<'a>, _ctx: &mut QueryContext) -> Result<Option<Columns>> {
230 if self.done {
231 return Ok(None);
232 }
233 self.done = true;
234
235 reifydb_assertions! {
236 assert!(self.context.is_some(), "AssertWithoutInputNode::next() called before initialize()");
237 }
238 let stored_ctx = self.context.as_ref().unwrap();
239 let session = eval_context_from_query(stored_ctx);
240
241 for assert_expr in &self.expressions {
242 let result = Self::eval(&session, assert_expr)?;
243 self.verify(result.data(), assert_expr)?;
244 }
245
246 Ok(None)
247 }
248
249 fn headers(&self) -> Option<ColumnHeaders> {
250 None
251 }
252}