Skip to main content

reifydb_engine/vm/volcano/
patch.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 ReifyDB
3
4use std::{mem, sync::Arc};
5
6use reifydb_core::{
7	interface::{evaluate::TargetColumn, resolved::ResolvedColumn},
8	value::column::{ColumnWithName, cast::cast_column_data, columns::Columns, headers::ColumnHeaders},
9};
10use reifydb_evaluate::expression::{
11	compile::{CompiledExpr, compile_expression},
12	context::{CompileContext, EvalContext},
13};
14use reifydb_extension::transform::{Transform, context::TransformContext};
15use reifydb_rql::expression::{Expression, name::display_label};
16use reifydb_transaction::transaction::Transaction;
17use reifydb_value::{fragment::Fragment, reifydb_assertions};
18use tracing::instrument;
19
20use super::NoopNode;
21use crate::{
22	Result,
23	vm::volcano::{
24		query::{QueryContext, QueryNode, eval_context_from_transform},
25		udf::{UdfEvalNode, strip_udf_columns},
26	},
27};
28
29pub(crate) struct PatchNode {
30	input: Box<dyn QueryNode>,
31	expressions: Vec<Expression>,
32	udf_names: Vec<String>,
33	headers: Option<ColumnHeaders>,
34	context: Option<(Arc<QueryContext>, Vec<CompiledExpr>)>,
35}
36
37impl PatchNode {
38	pub fn new(input: Box<dyn QueryNode>, expressions: Vec<Expression>) -> Self {
39		Self {
40			input,
41			expressions,
42			udf_names: Vec::new(),
43			headers: None,
44			context: None,
45		}
46	}
47}
48
49impl QueryNode for PatchNode {
50	#[instrument(name = "volcano::patch::initialize", level = "trace", skip_all)]
51	fn initialize<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &QueryContext) -> Result<()> {
52		let (input, expressions, udf_names) = UdfEvalNode::wrap_if_needed(
53			mem::replace(&mut self.input, Box::new(NoopNode)),
54			&self.expressions,
55			&ctx.symbols,
56		);
57		self.input = input;
58		self.expressions = expressions;
59		self.udf_names = udf_names;
60
61		let compile_ctx = CompileContext {
62			symbols: &ctx.symbols,
63		};
64		let compiled = self
65			.expressions
66			.iter()
67			.map(|e| compile_expression(&compile_ctx, e).expect("compile"))
68			.collect();
69		self.context = Some((Arc::new(ctx.clone()), compiled));
70		self.input.initialize(rx, ctx)?;
71		Ok(())
72	}
73
74	#[instrument(name = "volcano::patch::next", level = "trace", skip_all)]
75	fn next<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &mut QueryContext) -> Result<Option<Columns>> {
76		reifydb_assertions! {
77			assert!(self.context.is_some(), "PatchNode::next() called before initialize()");
78		}
79
80		if let Some(columns) = self.input.next(rx, ctx)? {
81			let stored_ctx = &self.context.as_ref().unwrap().0;
82			let transform_ctx = TransformContext {
83				routines: &ctx.services.routines,
84				runtime_context: &stored_ctx.services.runtime_context,
85				params: &stored_ctx.params,
86			};
87			let result = self.apply(&transform_ctx, columns)?;
88
89			if self.headers.is_none() {
90				let result_headers: Vec<Fragment> = result.iter().map(|c| c.name().clone()).collect();
91				self.headers = Some(ColumnHeaders {
92					columns: result_headers,
93				});
94			}
95
96			let mut result = result;
97			strip_udf_columns(&mut result, &self.udf_names);
98			Ok(Some(result))
99		} else {
100			Ok(None)
101		}
102	}
103
104	fn headers(&self) -> Option<ColumnHeaders> {
105		if let Some(ref headers) = self.headers {
106			return Some(headers.clone());
107		}
108
109		let input_headers = self.input.headers()?;
110		let patch_names: Vec<Fragment> = self.expressions.iter().map(display_label).collect();
111
112		let mut result = Vec::new();
113		for col in &input_headers.columns {
114			if let Some(patch_idx) = patch_names.iter().position(|n| n.text() == col.text()) {
115				result.push(patch_names[patch_idx].clone());
116			} else {
117				result.push(col.clone());
118			}
119		}
120
121		for patch_name in &patch_names {
122			if !result.iter().any(|h| h.text() == patch_name.text()) {
123				result.push(patch_name.clone());
124			}
125		}
126
127		Some(ColumnHeaders {
128			columns: result,
129		})
130	}
131}
132
133impl Transform for PatchNode {
134	fn apply(&self, ctx: &TransformContext, input: Columns) -> Result<Columns> {
135		let (stored_ctx, compiled) =
136			self.context.as_ref().expect("PatchNode::apply() called before initialize()");
137
138		let row_count = input.row_count();
139
140		let patch_names: Vec<Fragment> = self.expressions.iter().map(display_label).collect();
141
142		let session = eval_context_from_transform(ctx, stored_ctx);
143		let mut patch_columns = Vec::with_capacity(self.expressions.len());
144		for (expr, compiled_expr) in self.expressions.iter().zip(compiled.iter()) {
145			let mut exec_ctx = Self::eval_context(&session, &input, row_count);
146
147			if let (Expression::Alias(alias_expr), Some(source)) = (expr, &stored_ctx.source) {
148				let alias_name = alias_expr.alias.name();
149
150				if let Some(table_column) = source.columns().iter().find(|col| col.name == alias_name) {
151					let column_ident = Fragment::internal(&table_column.name);
152					let resolved_column =
153						ResolvedColumn::new(column_ident, source.clone(), table_column.clone());
154					exec_ctx.target = Some(TargetColumn::Resolved(resolved_column));
155				}
156			}
157
158			let mut column = Self::eval_patch(compiled_expr, &exec_ctx)?;
159
160			if let Some(target_type) = exec_ctx.target.as_ref().map(|t| t.column_type())
161				&& column.data.get_type() != target_type
162			{
163				let data =
164					cast_column_data(&exec_ctx, &column.data, target_type, &expr.lazy_fragment())?;
165				column = ColumnWithName {
166					name: column.name,
167					data,
168				};
169			}
170
171			patch_columns.push(column);
172		}
173
174		Ok(Self::merge(input, &patch_names, patch_columns))
175	}
176}
177
178impl PatchNode {
179	#[instrument(level = "trace", skip_all, name = "volcano::patch::eval_context")]
180	fn eval_context<'e>(session: &EvalContext<'e>, input: &Columns, row_count: usize) -> EvalContext<'e> {
181		session.with_eval(input.clone(), row_count)
182	}
183
184	#[instrument(level = "trace", skip_all, name = "volcano::patch::eval")]
185	fn eval_patch(compiled: &CompiledExpr, exec_ctx: &EvalContext) -> Result<ColumnWithName> {
186		compiled.execute(exec_ctx)
187	}
188
189	#[instrument(level = "trace", skip_all, name = "volcano::patch::merge")]
190	fn merge(input: Columns, patch_names: &[Fragment], patch_columns: Vec<ColumnWithName>) -> Columns {
191		let system = input.system.clone();
192		let mut result_columns: Vec<ColumnWithName> = Vec::new();
193
194		for (original_name, original_data) in input.names.iter().zip(input.columns.iter()) {
195			let original_name_text = original_name.text();
196
197			if let Some(patch_idx) = patch_names.iter().position(|n| n.text() == original_name_text) {
198				result_columns.push(patch_columns[patch_idx].clone());
199			} else {
200				result_columns.push(ColumnWithName::new(original_name.clone(), original_data.clone()));
201			}
202		}
203
204		for (patch_idx, patch_name) in patch_names.iter().enumerate() {
205			if !result_columns.iter().any(|c| c.name().text() == patch_name.text()) {
206				result_columns.push(patch_columns[patch_idx].clone());
207			}
208		}
209
210		let mut names_vec = Vec::with_capacity(result_columns.len());
211		let mut buffers_vec = Vec::with_capacity(result_columns.len());
212		for c in result_columns {
213			names_vec.push(c.name);
214			buffers_vec.push(c.data);
215		}
216		Columns {
217			system,
218			columns: buffers_vec,
219			names: names_vec,
220		}
221	}
222}