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