Skip to main content

reifydb_engine/vm/instruction/dml/
table_update.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 ReifyDB
3
4use 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, &params, 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}