Skip to main content

reifydb_engine/vm/volcano/
assert.rs

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