1use std::{collections::HashMap, sync::Arc};
5
6use reifydb_codec::row::{
7 bytes::EncodedBytes,
8 pod::EncodedPodRow,
9 shape::RowShape,
10 table::{EncodedTableRow, EncodedTableRowBuilder},
11};
12use reifydb_core::{
13 error::diagnostic::{
14 catalog::{namespace_not_found, table_not_found},
15 engine,
16 },
17 interface::{
18 catalog::{
19 config::{ConfigKey, GetConfig},
20 id::IndexId,
21 key::PrimaryKey,
22 namespace::Namespace,
23 object::ObjectId,
24 policy::{DataOp, PolicyTargetType},
25 table::Table,
26 },
27 resolved::{ResolvedColumn, ResolvedNamespace, ResolvedObject, ResolvedTable},
28 },
29 internal_error,
30 key::{any::TaggedKey, catalog::IndexEntryKey},
31 partition::PartitionError,
32 value::column::columns::Columns,
33};
34use reifydb_evaluate::stack::SymbolTable;
35use reifydb_rql::nodes::UpdateTableNode;
36use reifydb_transaction::transaction::Transaction;
37use reifydb_value::{
38 fragment::Fragment,
39 params::Params,
40 return_error,
41 value::{Value, identity::IdentityId, partition::Partition, row_number::RowNumber},
42};
43
44use super::{
45 context::{TableTarget, WriteExecCtx},
46 primary_key,
47 returning::{decode_returning_dictionaries, decode_rows_to_columns, evaluate_returning, with_pre_image},
48 shape::get_or_create_table_shape,
49};
50use crate::{
51 Result,
52 error::EngineError,
53 partition::{row_key_from_partition, table_partition_of_row},
54 policy::PolicyEvaluator,
55 transaction::operation::{dictionary::DictionaryOperations, table::TableOperations},
56 vm::{
57 instruction::dml::{coerce::coerce_value_to_column_type, time::resolve_time_for_update},
58 services::Services,
59 volcano::{
60 compile::compile,
61 query::{QueryContext, QueryNode, query_budget},
62 },
63 },
64};
65
66pub(crate) fn update_table(
67 services: &Arc<Services>,
68 txn: &mut Transaction<'_>,
69 plan: UpdateTableNode,
70 params: Params,
71 symbols: &SymbolTable,
72) -> Result<Columns> {
73 let UpdateTableNode {
74 input,
75 target,
76 returning,
77 } = plan;
78 let target = target.expect("Cannot infer target table from pipeline - no table found");
79 let (namespace, table) = resolve_update_table_target(services, txn, &target)?;
80 let shape = get_or_create_table_shape(&services.catalog, &table, txn)?;
81 let target_data = TableTarget {
82 namespace: &namespace,
83 table: &table,
84 fragment: target.identifier(),
85 };
86 let context = build_update_table_query_context(services, &target_data, ¶ms, symbols, txn.identity());
87
88 let mut input_node = compile(*input, txn, Arc::new(context.clone()));
89 input_node.initialize(txn, &context)?;
90
91 let exec = WriteExecCtx {
92 services,
93 symbols,
94 };
95 let (updated_count, returned_rows, pre_rows) =
96 run_table_update(&exec, txn, &mut input_node, &target_data, &shape, &context, returning.is_some())?;
97
98 if let Some(returning_exprs) = &returning {
99 let mut columns = decode_rows_to_columns(&shape, &returned_rows);
100 decode_returning_dictionaries(services, txn, &table.columns, &mut columns)?;
101 let mut pre_columns = decode_rows_to_columns(&shape, &pre_rows);
102 decode_returning_dictionaries(services, txn, &table.columns, &mut pre_columns)?;
103 let columns = with_pre_image(columns, &pre_columns);
104 return evaluate_returning(services, symbols, returning_exprs, columns, txn.identity());
105 }
106 Ok(update_table_result(namespace.name(), &table.name, updated_count))
107}
108
109#[inline]
110fn resolve_update_table_target(
111 services: &Arc<Services>,
112 txn: &mut Transaction<'_>,
113 target: &ResolvedTable,
114) -> Result<(Namespace, Table)> {
115 let namespace_name = target.namespace().name();
116 let Some(namespace) = services.catalog.find_namespace_by_name(txn, namespace_name)? else {
117 return_error!(namespace_not_found(Fragment::internal(namespace_name), namespace_name));
118 };
119 let Some(table) = services.catalog.find_table_by_name(txn, namespace.id(), target.name())? else {
120 let fragment = target.identifier().clone();
121 return_error!(table_not_found(fragment.clone(), namespace_name, target.name(),));
122 };
123 Ok((namespace, table))
124}
125
126#[inline]
127fn build_update_table_query_context(
128 services: &Arc<Services>,
129 target: &TableTarget<'_>,
130 params: &Params,
131 symbols: &SymbolTable,
132 identity: IdentityId,
133) -> QueryContext {
134 let namespace_ident = Fragment::internal(target.namespace.name());
135 let resolved_namespace = ResolvedNamespace::new(namespace_ident, target.namespace.clone());
136 let table_ident = Fragment::internal(target.table.name.clone());
137 let resolved_table = ResolvedTable::new(table_ident, resolved_namespace, target.table.clone());
138 QueryContext {
139 services: services.clone(),
140 source: Some(ResolvedObject::Table(resolved_table)),
141 batch_size: services.catalog.get_config_uint2(ConfigKey::QueryRowBatchSize) as u64,
142 params: params.clone(),
143 symbols: symbols.clone(),
144 identity,
145 memory: query_budget(services),
146 }
147}
148
149type ReturnedRows = Vec<(RowNumber, EncodedBytes)>;
150
151fn run_table_update(
152 exec: &WriteExecCtx<'_>,
153 txn: &mut Transaction<'_>,
154 input_node: &mut Box<dyn QueryNode>,
155 target: &TableTarget<'_>,
156 shape: &RowShape,
157 context: &QueryContext,
158 has_returning: bool,
159) -> Result<(u64, ReturnedRows, ReturnedRows)> {
160 let mut updated_count = 0u64;
161 let mut returned_rows: ReturnedRows = Vec::new();
162 let mut pre_rows: ReturnedRows = Vec::new();
163 let mut mutable_context = context.clone();
164
165 while let Some(columns) = input_node.next(txn, &mut mutable_context)? {
166 PolicyEvaluator::new(exec.services, exec.symbols).enforce_write_policies(
167 txn,
168 target.namespace.name(),
169 &target.table.name,
170 DataOp::Update,
171 &columns,
172 PolicyTargetType::Table,
173 )?;
174
175 if columns.row_numbers().is_empty() {
176 return_error!(engine::missing_row_number_column());
177 }
178
179 let partitioned = !target.table.partition_by.is_empty();
180 if partitioned && columns.partitions().len() != columns.row_count() {
181 return Err(EngineError::MissingPartitionAddress {
182 object: ObjectId::Table(target.table.id),
183 operation: "UPDATE",
184 }
185 .into());
186 }
187
188 let row_numbers: Vec<RowNumber> = columns.row_numbers().to_vec();
189 let sidecar_partitions: Vec<Partition> = columns.partitions().to_vec();
190 let row_count = columns.row_count();
191
192 let mut prepared_rows: Vec<EncodedTableRowBuilder> = Vec::with_capacity(row_count);
193 let mut partitions_out: Vec<Partition> = Vec::with_capacity(row_count);
194 let mut pre_by_row: HashMap<RowNumber, EncodedBytes> = HashMap::new();
195 for (row_idx, &row_number) in row_numbers.iter().enumerate() {
196 let mut row = build_updated_table_row(
197 exec.services,
198 txn,
199 target.table,
200 shape,
201 &columns,
202 context,
203 row_idx,
204 )?;
205 let partition = sidecar_partitions.get(row_idx).copied();
206
207 if let Some(old) = partition {
208 let new_partition = table_partition_of_row(target.table, shape, &row);
209 if new_partition != old {
210 return Err(PartitionError::ImmutablePartitionColumn {
211 object: ObjectId::Table(target.table.id),
212 }
213 .into());
214 }
215 }
216
217 let row_key = row_key_from_partition(target.table.id, partition, row_number);
218
219 if let Some(pk_def) = primary_key::get_primary_key(&exec.services.catalog, txn, target.table)? {
220 rotate_table_pk_index(txn, target.table, shape, &pk_def, &row_key, &row, row_number)?;
221 }
222
223 let old_row = txn.get(&row_key)?.expect("bytes must exist for update").bytes;
224 if has_returning {
225 pre_by_row.insert(row_number, old_row.clone());
226 }
227 let old_row = EncodedTableRow::view(&old_row);
228 let old_created_at = old_row.created_at();
229 let old_time = old_row.time();
230 let now = exec.services.runtime_context.clock.now();
231 row.set_timestamps(old_created_at, now);
232 if let Some(time) = resolve_time_for_update(
233 &target.table.name,
234 &target.table.columns,
235 &target.table.time,
236 shape,
237 &row,
238 old_time,
239 )? {
240 row.set_time(time);
241 }
242
243 prepared_rows.push(row);
244 if let Some(p) = partition {
245 partitions_out.push(p);
246 }
247 }
248
249 let stored = txn.update_table(target.table, &row_numbers, &partitions_out, &mut prepared_rows)?;
250 updated_count += stored.len() as u64;
251 if has_returning {
252 for (row_number, _) in &stored {
253 let pre = pre_by_row
254 .remove(row_number)
255 .expect("every stored row must carry the pre image read before its write");
256 pre_rows.push((*row_number, pre));
257 }
258 returned_rows.extend(stored);
259 }
260 }
261 Ok((updated_count, returned_rows, pre_rows))
262}
263
264#[inline]
265fn build_updated_table_row(
266 services: &Arc<Services>,
267 txn: &mut Transaction<'_>,
268 table: &Table,
269 shape: &RowShape,
270 columns: &Columns,
271 context: &QueryContext,
272 row_idx: usize,
273) -> Result<EncodedTableRowBuilder> {
274 let mut row = shape.allocate_table();
275 for (table_idx, table_column) in table.columns.iter().enumerate() {
276 let mut value = if let Some(input_column) = columns.iter().find(|col| col.name() == table_column.name) {
277 input_column.data().get_value(row_idx)
278 } else {
279 Value::none()
280 };
281
282 let column_ident = columns
283 .iter()
284 .find(|col| col.name() == table_column.name)
285 .map(|col| col.name().clone())
286 .unwrap_or_else(|| Fragment::internal(&table_column.name));
287 let resolved_column = ResolvedColumn::new(
288 column_ident.clone(),
289 context.source.clone().unwrap(),
290 table_column.clone(),
291 );
292
293 value = coerce_value_to_column_type(
294 value,
295 table_column.constraint.get_type(),
296 resolved_column,
297 context,
298 )?;
299 if let Err(mut e) = table_column.constraint.validate(&value) {
300 e.0.fragment = column_ident.clone();
301 return Err(e);
302 }
303
304 let value = if let Some(dict_id) = table_column.dictionary_id {
305 let dictionary = services.catalog.find_dictionary(txn, dict_id)?.ok_or_else(|| {
306 internal_error!("Dictionary {:?} not found for column {}", dict_id, table_column.name)
307 })?;
308 let entry_id = txn.insert_into_dictionary(&dictionary, &value)?;
309 entry_id.to_value()
310 } else {
311 value
312 };
313
314 shape.set_value(&mut row, table_idx, &value);
315 }
316 Ok(row)
317}
318
319#[inline]
320fn rotate_table_pk_index(
321 txn: &mut Transaction<'_>,
322 table: &Table,
323 shape: &RowShape,
324 pk_def: &PrimaryKey,
325 row_key: &TaggedKey,
326 new_row: &[u8],
327 row_number: RowNumber,
328) -> Result<()> {
329 if let Some(pre_row_data) = txn.get(row_key)? {
330 let pre_row = pre_row_data.bytes;
331 let pre_key = primary_key::encode_primary_key(pk_def, &pre_row, table, shape)?;
332 txn.remove(&IndexEntryKey::new(table.id, IndexId::primary(pk_def.id), pre_key))?;
333 }
334
335 let post_key = primary_key::encode_primary_key(pk_def, new_row, table, shape)?;
336 txn.set(
337 &IndexEntryKey::new(table.id, IndexId::primary(pk_def.id), post_key),
338 EncodedPodRow::new(&u64::from(row_number).to_be_bytes()).into_bytes(),
339 )?;
340 Ok(())
341}
342
343#[inline]
344fn update_table_result(namespace: &str, table: &str, updated: u64) -> Columns {
345 Columns::single_row([
346 ("namespace", Value::Utf8(namespace.to_string())),
347 ("table", Value::Utf8(table.to_string())),
348 ("updated", Value::Uint8(updated)),
349 ])
350}