reifydb_engine/vm/volcano/
patch.rs1use 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}