use std::sync::Arc;
use reifydb_codec::row::{
bytes::{EncodedBytes, RowBuilder},
series::EncodedSeriesRow,
shape::RowShape,
};
use reifydb_core::{
common::CommitVersion,
error::diagnostic::catalog::{namespace_not_found, series_not_found},
interface::{
catalog::{
config::{ConfigKey, GetConfig},
namespace::Namespace,
object::ObjectId,
policy::{DataOp, PolicyTargetType},
series::Series,
storage::StorageId,
},
change::{Change, ChangeOrigin, Diff},
resolved::{ResolvedNamespace, ResolvedObject, ResolvedSeries},
},
internal_error,
key::{
any::TaggedKey,
series::{PartitionedSeriesRowKey, SeriesRowKey},
},
partition::PartitionError,
value::column::{ColumnWithName, buffer::ColumnBuffer, columns::Columns},
};
use reifydb_evaluate::stack::SymbolTable;
use reifydb_rql::nodes::UpdateSeriesNode;
use reifydb_transaction::{interceptor::series_row::SeriesRowInterceptor, transaction::Transaction};
use reifydb_value::{
fragment::Fragment,
params::Params,
return_error,
value::{
Value, datetime::DateTime, identity::IdentityId, partition::Partition, row_number::RowNumber,
system_columns::SystemColumns,
},
};
use smallvec::smallvec;
use tracing::instrument;
use super::{
context::SeriesTarget,
returning::{decode_returning_dictionaries, decode_rows_to_columns, evaluate_returning, with_pre_image},
};
use crate::{
Result,
error::EngineError,
partition::partition_values,
policy::PolicyEvaluator,
transaction::operation::dictionary::DictionaryOperations,
vm::{
instruction::dml::{shape::get_or_create_series_shape, time::resolve_time_for_update},
services::Services,
volcano::{
compile::compile,
query::{QueryContext, QueryNode, query_budget},
},
},
};
#[instrument(name = "mutate::series::update", level = "trace", skip_all)]
pub(crate) fn update_series(
services: &Arc<Services>,
txn: &mut Transaction<'_>,
plan: UpdateSeriesNode,
params: Params,
symbols: &SymbolTable,
) -> Result<Columns> {
let UpdateSeriesNode {
input,
target,
returning,
} = plan;
let (namespace, series) = resolve_update_series_target(services, txn, &target)?;
let target_data = SeriesTarget {
namespace: &namespace,
series: &series,
};
let context = build_update_series_query_context(services, &target_data, ¶ms, symbols, txn.identity());
let mut input_node = compile(*input, txn, Arc::new(context.clone()));
input_node.initialize(txn, &context)?;
let has_tag = series.tag.is_some();
let mut updated_count = 0u64;
let has_returning = returning.is_some();
let mut returned_rows: Vec<(RowNumber, EncodedBytes)> = Vec::new();
let mut pre_rows: Vec<(RowNumber, EncodedBytes)> = Vec::new();
let mut mutable_context = context.clone();
while let Some(columns) = input_node.next(txn, &mut mutable_context)? {
let row_count = columns.row_count();
if row_count == 0 {
continue;
}
PolicyEvaluator::new(services, symbols).enforce_write_policies(
txn,
namespace.name(),
&series.name,
DataOp::Update,
&columns,
PolicyTargetType::Series,
)?;
let row_numbers = columns.row_numbers();
let updates_to_apply =
build_series_updates_to_apply(services, txn, &series, &columns, row_numbers, has_tag)?;
for (key, row, row_idx) in updates_to_apply {
let pre_values = match txn.get(&key)? {
Some(v) => v.bytes,
None => continue,
};
let old_created_at = EncodedSeriesRow::view(&pre_values).created_at();
let old_time = EncodedSeriesRow::view(&pre_values).time();
let now = services.runtime_context.clock.now();
let update_shape = get_or_create_series_shape(&services.catalog, &series, txn)?;
let mut builder = EncodedSeriesRow::from(row).thaw();
builder.set_timestamps(old_created_at, now);
if let Some(time) = resolve_time_for_update(
&series.name,
&series.columns,
&series.time,
&update_shape,
builder.as_slice(),
old_time,
)? {
builder.set_time(time);
}
let key_value = extract_series_update_key_value(&columns, &series, row_idx);
let row_number = RowNumber::from(u64::from(row_numbers[row_idx]));
let mut rows_buf = [builder];
SeriesRowInterceptor::pre_update(txn, &series, &mut rows_buf)?;
let [row] = rows_buf;
let row = row.freeze_bytes();
if !series.partition_by.is_empty() {
let expected = columns.partitions()[row_idx];
let shape = get_or_create_series_shape(&services.catalog, &series, txn)?;
if series_partition_of_bytes(&series, &shape, &row) != expected {
return Err(PartitionError::ImmutablePartitionColumn {
object: ObjectId::series(series.id),
}
.into());
}
}
if txn.get_committed(&key)?.is_some() {
txn.mark_preexisting(&key)?;
}
txn.set(&key, row.clone())?;
let posts = [row.clone()];
let pres = [pre_values.clone()];
SeriesRowInterceptor::post_update(txn, &series, &posts, &pres)?;
let event = SeriesUpdateEvent {
columns: &columns,
pre: &pre_values,
post: &row,
key_value,
row_number,
row_idx,
};
track_series_update_flow_change(services, txn, &series, &event)?;
if has_returning {
returned_rows.push((row_number, row.clone()));
pre_rows.push((row_number, pre_values.clone()));
}
updated_count += 1;
}
}
if let Some(returning_exprs) = &returning {
let shape = get_or_create_series_shape(&services.catalog, &series, txn)?;
let mut cols = decode_rows_to_columns(&shape, &returned_rows);
decode_returning_dictionaries(services, txn, &series.columns, &mut cols)?;
let mut pre_cols = decode_rows_to_columns(&shape, &pre_rows);
decode_returning_dictionaries(services, txn, &series.columns, &mut pre_cols)?;
let cols = with_pre_image(cols, &pre_cols);
return evaluate_returning(services, symbols, returning_exprs, cols, txn.identity());
}
Ok(update_series_result(namespace.name(), &series.name, updated_count))
}
struct SeriesUpdateEvent<'a> {
columns: &'a Columns,
pre: &'a EncodedBytes,
post: &'a EncodedBytes,
key_value: u64,
row_number: RowNumber,
row_idx: usize,
}
#[inline]
fn resolve_update_series_target(
services: &Arc<Services>,
txn: &mut Transaction<'_>,
target: &ResolvedSeries,
) -> Result<(Namespace, Series)> {
let namespace_name = target.namespace().name();
let Some(namespace) = services.catalog.find_namespace_by_name(txn, namespace_name)? else {
return_error!(namespace_not_found(Fragment::internal(namespace_name), namespace_name));
};
let series_name = target.name();
let Some(series) = services.catalog.find_series_by_name(txn, namespace.id(), series_name)? else {
let fragment = Fragment::internal(target.name());
return_error!(series_not_found(fragment, namespace_name, series_name));
};
Ok((namespace, series))
}
#[inline]
fn build_update_series_query_context(
services: &Arc<Services>,
target: &SeriesTarget<'_>,
params: &Params,
symbols: &SymbolTable,
identity: IdentityId,
) -> QueryContext {
let namespace_ident = Fragment::internal(target.namespace.name());
let resolved_namespace = ResolvedNamespace::new(namespace_ident, target.namespace.clone());
let series_ident = Fragment::internal(target.series.name.clone());
let resolved_series = ResolvedSeries::new(series_ident, resolved_namespace, target.series.clone());
QueryContext {
services: services.clone(),
source: Some(ResolvedObject::Series(resolved_series)),
batch_size: services.catalog.get_config_uint2(ConfigKey::QueryRowBatchSize) as u64,
params: params.clone(),
symbols: symbols.clone(),
identity,
memory: query_budget(services),
}
}
fn build_series_updates_to_apply(
services: &Arc<Services>,
txn: &mut Transaction<'_>,
series: &Series,
columns: &Columns,
row_numbers: &[RowNumber],
has_tag: bool,
) -> Result<Vec<(TaggedKey, EncodedBytes, usize)>> {
let row_count = columns.row_count();
let partitioned = !series.partition_by.is_empty();
if partitioned && columns.partitions().len() != row_count {
return Err(EngineError::MissingPartitionAddress {
object: ObjectId::series(series.id),
operation: "UPDATE",
}
.into());
}
let mut updates_to_apply: Vec<(TaggedKey, EncodedBytes, usize)> = Vec::with_capacity(row_count);
for (row_idx, row_number) in row_numbers.iter().enumerate().take(row_count) {
let sequence = u64::from(*row_number);
let key_value = extract_series_update_key_value(columns, series, row_idx);
let variant_tag = extract_series_update_variant_tag(columns, has_tag, row_idx);
let key: TaggedKey = if partitioned {
let old_partition = columns.partitions()[row_idx];
let new_partition = series_partition_of_columns(series, columns, row_idx)?;
if new_partition != old_partition {
return Err(PartitionError::ImmutablePartitionColumn {
object: ObjectId::series(series.id),
}
.into());
}
PartitionedSeriesRowKey::new(
StorageId::series(series.id),
old_partition,
variant_tag,
key_value,
sequence,
)
.into()
} else {
SeriesRowKey {
storage: StorageId::series(series.id),
variant_tag,
key: key_value,
sequence,
}
.into()
};
let shape = get_or_create_series_shape(&services.catalog, series, txn)?;
let row = build_series_update_bytes(services, txn, series, columns, &shape, row_idx)?;
updates_to_apply.push((key, row, row_idx));
}
Ok(updates_to_apply)
}
#[inline]
fn series_partition_of_bytes(series: &Series, shape: &RowShape, bytes: &EncodedBytes) -> Partition {
let key_column = series.key.column();
let indices: Vec<usize> = series
.partition_by
.iter()
.map(|name| {
if name == key_column {
0
} else {
1 + series
.data_columns()
.position(|c| c.name == *name)
.expect("partition column must exist (validated during planning)")
}
})
.collect();
Partition::of(&partition_values(shape, bytes, &indices))
}
#[inline]
fn series_partition_of_columns(series: &Series, columns: &Columns, row_idx: usize) -> Result<Partition> {
let mut part_values = Vec::with_capacity(series.partition_by.len());
for name in &series.partition_by {
let idx =
columns.names.iter().position(|n| n.text() == name.as_str()).ok_or_else(|| {
internal_error!("partition column {} missing from series update input", name)
})?;
part_values.push(columns[idx].get_value(row_idx));
}
Ok(Partition::of(&part_values))
}
#[inline]
fn extract_series_update_key_value(columns: &Columns, series: &Series, row_idx: usize) -> u64 {
columns.iter()
.find(|c| c.name().text() == series.key.column())
.and_then(|c| series.key_to_u64(c.data().get_value(row_idx)))
.unwrap_or(0)
}
#[inline]
fn extract_series_update_variant_tag(columns: &Columns, has_tag: bool, row_idx: usize) -> Option<u8> {
if !has_tag {
return None;
}
columns.iter().find(|c| c.name().text() == "tag").and_then(|c| match c.data().get_value(row_idx) {
Value::Uint1(v) => Some(v),
_ => None,
})
}
#[inline]
fn build_series_update_bytes(
services: &Arc<Services>,
txn: &mut Transaction<'_>,
series: &Series,
columns: &Columns,
shape: &RowShape,
row_idx: usize,
) -> Result<EncodedBytes> {
let mut row = shape.allocate_series();
let key_col_value = columns
.iter()
.find(|c| c.name().text() == series.key.column())
.map(|c| c.data().get_value(row_idx))
.unwrap_or(Value::Int8(0));
shape.set_value(&mut row, 0, &key_col_value);
let data_columns: Vec<_> = series.data_columns().cloned().collect();
for (i, col_def) in data_columns.iter().enumerate() {
let value = columns
.iter()
.find(|c| c.name().text() == col_def.name)
.map(|c| c.data().get_value(row_idx))
.unwrap_or(Value::none());
let value = match col_def.dictionary_id {
Some(dict_id) => {
let dictionary = services.catalog.find_dictionary(txn, dict_id)?.ok_or_else(|| {
internal_error!(
"Dictionary {:?} not found for column {}",
dict_id,
col_def.name
)
})?;
let entry_id = if matches!(value, Value::None { .. }) {
dictionary.id_type.none()
} else {
txn.insert_into_dictionary(&dictionary, &value)?
};
entry_id.to_value()
}
None => value,
};
shape.set_value(&mut row, i + 1, &value);
}
Ok(row.freeze_bytes())
}
fn track_series_update_flow_change(
services: &Arc<Services>,
txn: &mut Transaction<'_>,
series: &Series,
event: &SeriesUpdateEvent<'_>,
) -> Result<()> {
let read_shape = get_or_create_series_shape(&services.catalog, series, txn)?;
let mut pre_col_vec = Vec::with_capacity(1 + series.columns.len());
pre_col_vec.push(ColumnWithName::new(
Fragment::internal(series.key.column()),
series.key_column_data(vec![event.key_value]),
));
let read_fields = read_shape.fields();
for (i, col_def) in series.data_columns().enumerate() {
let val = read_shape.get_value(event.pre, i + 1);
let mut data = ColumnBuffer::with_capacity(read_fields[i + 1].constraint.get_type(), 1);
data.push_value(val);
pre_col_vec.push(ColumnWithName {
name: Fragment::internal(&col_def.name),
data,
});
}
let mut post_col_vec = Vec::with_capacity(1 + series.columns.len());
post_col_vec.push(ColumnWithName::new(
Fragment::internal(series.key.column()),
series.key_column_data(vec![event.key_value]),
));
for col in event.columns.iter() {
if col.name().text() != series.key.column() && col.name().text() != "tag" {
let mut data = ColumnBuffer::with_capacity(col.data().get_type(), 1);
data.push_value(col.data().get_value(event.row_idx));
post_col_vec.push(ColumnWithName {
name: col.name().clone(),
data,
});
}
}
let pre = Columns::with_system(
pre_col_vec,
SystemColumns::new(
vec![event.row_number],
Vec::new(),
vec![EncodedSeriesRow::view(event.pre).created_at()],
vec![EncodedSeriesRow::view(event.pre).updated_at()],
EncodedSeriesRow::view(event.pre).time().into_iter().collect(),
),
);
let post = Columns::with_system(
post_col_vec,
SystemColumns::new(
vec![event.row_number],
Vec::new(),
vec![EncodedSeriesRow::view(event.post).created_at()],
vec![EncodedSeriesRow::view(event.post).updated_at()],
EncodedSeriesRow::view(event.post).time().into_iter().collect(),
),
);
txn.track_flow_change(Change {
origin: ChangeOrigin::Object(ObjectId::series(series.id)),
version: CommitVersion(0),
diffs: smallvec![Diff::update(pre, post)],
changed_at: DateTime::default(),
});
Ok(())
}
#[inline]
fn update_series_result(namespace: &str, series: &str, updated: u64) -> Columns {
Columns::single_row([
("namespace", Value::Utf8(namespace.to_string())),
("series", Value::Utf8(series.to_string())),
("updated", Value::Uint8(updated)),
])
}