use super::{
SqlWriteMutationExecution, reject_explicit_sql_write_to_generated_field,
reject_explicit_sql_write_to_managed_field, sql_exact_update_candidate_bounds,
sql_write_candidate_bounds, sql_write_input_for_accepted_field,
sql_write_patch_set_accepted_field, sql_write_patch_set_update_default,
};
use crate::db::query::preparation::PreparationWork;
use crate::{
db::{
DbSession, MissingRowPolicy, QueryError,
data::{AcceptedMutationIntentPatch, DecodedDataStoreKey},
executor::StructuralProjectionScanBudget,
query::intent::StructuralQuery,
schema::AcceptedRowLayoutRuntimeContract,
session::{
AcceptedSchemaCatalogContext, AcceptedStructuralMutationTarget,
sql::{
SqlExactUpdatePolicy, SqlExactUpdatePolicyRejection, SqlPublicBoundedUpdatePlan,
SqlPublicPrimaryKeyUpdatePlan, SqlStatementDispatch, SqlStatementResult,
SqlTrustedExactUpdatePlan, SqlUpdateExposurePolicy, SqlUpdatePolicyRejection,
SqlUpdatePolicyResult, SqlValidatedUpdatePlan,
classify_sql_update_policy_for_entity, sql_statement_dispatch,
with_accepted_sql_update_policy_context,
write_policy::{SqlWriteExecutionBounds, SqlWriteShapePolicyRejection},
},
},
sql::{
lowering::bind_sql_update_selector_query_structural_with_schema,
parser::{SqlUpdateStatement, SqlWriteValue},
},
write_context::{AcceptedWriteContext, MutationMode},
},
traits::CanisterKind,
types::{CurrentTimestamp, Timestamp},
value::Value,
};
use icydb_diagnostic_code::{DiagnosticFactTag, SqlWriteBoundaryCode};
fn sql_exact_update_policy_error(
require_affected_at_most: u32,
rejection: SqlExactUpdatePolicyRejection,
) -> QueryError {
let (boundary, bound_tag, bound) = match rejection {
SqlExactUpdatePolicyRejection::AssertionRequired => (
SqlWriteBoundaryCode::ExactUpdateAssertionRequired,
DiagnosticFactTag::Minimum,
1,
),
SqlExactUpdatePolicyRejection::AssertionTooHigh => (
SqlWriteBoundaryCode::ExactUpdateAssertionTooHigh,
DiagnosticFactTag::Limit,
u64::from(SqlExactUpdatePolicy::max_affected_rows()),
),
};
QueryError::sql_write_boundary_with_facts(
boundary,
vec![
(
DiagnosticFactTag::ActualCount,
u64::from(require_affected_at_most),
),
(bound_tag, bound),
],
)
}
fn require_sql_exact_update_plan(
result: SqlUpdatePolicyResult,
) -> Result<SqlTrustedExactUpdatePlan, QueryError> {
let rejection = match result {
Ok(SqlValidatedUpdatePlan::TrustedExact(plan)) => return Ok(plan),
Err(rejection) => rejection,
Ok(
SqlValidatedUpdatePlan::PublicPrimaryKeyOnly(_)
| SqlValidatedUpdatePlan::PublicBoundedDeterministic(_),
) => return Err(QueryError::unsupported_query()),
};
let boundary = match rejection {
SqlUpdatePolicyRejection::WriteShape(SqlWriteShapePolicyRejection::MissingWhere) => {
SqlWriteBoundaryCode::UpdateMissingWherePredicate
}
SqlUpdatePolicyRejection::PrimaryKeyMutation => {
SqlWriteBoundaryCode::UpdatePrimaryKeyMutation
}
SqlUpdatePolicyRejection::GeneratedFieldMutation => {
SqlWriteBoundaryCode::ExplicitGeneratedField
}
SqlUpdatePolicyRejection::ManagedFieldMutation => {
SqlWriteBoundaryCode::ExplicitManagedField
}
SqlUpdatePolicyRejection::ExactWindowUnsupported => {
SqlWriteBoundaryCode::ExactUpdateWindowUnsupported
}
SqlUpdatePolicyRejection::NotUpdate
| SqlUpdatePolicyRejection::WriteShape(_)
| SqlUpdatePolicyRejection::ResumableWindowUnsupported
| SqlUpdatePolicyRejection::ResumableReturningUnsupported => {
return Err(QueryError::unsupported_query());
}
};
Err(QueryError::sql_write_boundary(boundary))
}
#[derive(Clone, Copy)]
enum SqlUpdateExecutionContract {
Validated(SqlWriteExecutionBounds),
Exact {
policy: SqlExactUpdatePolicy,
bounds: SqlWriteExecutionBounds,
},
}
impl SqlUpdateExecutionContract {
const fn candidate_bounds(self) -> super::SqlWriteCandidateBounds {
match self {
Self::Validated(bounds) => sql_write_candidate_bounds(Some(bounds)),
Self::Exact { policy, .. } => sql_exact_update_candidate_bounds(policy),
}
}
fn selector(self, selector: StructuralQuery) -> StructuralQuery {
match self {
Self::Exact { policy, .. } => selector.limit(policy.selection_limit()),
Self::Validated(_) => selector,
}
}
const fn returning_bounds(
self,
returning_requested: bool,
) -> Option<crate::db::session::sql::write_policy::SqlWriteReturningBounds> {
if !returning_requested {
return None;
}
match self {
Self::Validated(bounds) | Self::Exact { bounds, .. } => Some(bounds.returning),
}
}
fn scan_budget(self) -> Result<Option<StructuralProjectionScanBudget>, QueryError> {
let Self::Exact { .. } = self else {
return Ok(None);
};
StructuralProjectionScanBudget::try_new(SqlExactUpdatePolicy::scan_budget())
.map(Some)
.ok_or_else(QueryError::invariant)
}
}
impl<C: CanisterKind> DbSession<C> {
pub(in crate::db::session::sql) fn sql_structural_patch(
descriptor: &AcceptedRowLayoutRuntimeContract<'_>,
statement: &SqlUpdateStatement,
) -> Result<AcceptedMutationIntentPatch, QueryError> {
let mut patch = AcceptedMutationIntentPatch::new();
for assignment in &statement.assignments {
if descriptor.is_primary_key_field_name(assignment.field.as_str()) {
return Err(QueryError::sql_write_boundary(
SqlWriteBoundaryCode::UpdatePrimaryKeyMutation,
));
}
patch = match &assignment.value {
SqlWriteValue::Literal(value) => {
reject_explicit_sql_write_to_generated_field(
descriptor,
assignment.field.as_str(),
)?;
reject_explicit_sql_write_to_managed_field(
descriptor,
assignment.field.as_str(),
)?;
let input = sql_write_input_for_accepted_field(
descriptor,
assignment.field.as_str(),
value,
)?;
sql_write_patch_set_accepted_field(
descriptor,
patch,
assignment.field.as_str(),
input,
)?
}
SqlWriteValue::Default => sql_write_patch_set_update_default(
descriptor,
patch,
assignment.field.as_str(),
)?,
};
}
Ok(patch)
}
pub(in crate::db::session::sql) fn sql_update_selector_query(
&self,
schema_info: &crate::db::schema::SchemaInfo,
statement: &SqlUpdateStatement,
) -> Result<StructuralQuery, QueryError> {
if schema_info.primary_key_names().is_empty() {
return Err(QueryError::invariant());
}
let primary_key_names = schema_info
.primary_key_names()
.iter()
.map(String::as_str)
.collect::<Vec<_>>();
let selector = PreparationWork::run(
self.db.request_execution_scope(),
icydb_diagnostic_code::DiagnosticExecutionLane::Mutation,
|work| {
bind_sql_update_selector_query_structural_with_schema(
statement,
MissingRowPolicy::Ignore,
schema_info,
work,
)
.map_err(QueryError::from_sql_lowering_error)
},
)?;
Ok(selector.select_fields(primary_key_names))
}
fn sql_write_key_from_projected_row(
entity_tag: crate::types::EntityTag,
descriptor: &AcceptedRowLayoutRuntimeContract<'_>,
row: &[Value],
) -> Result<DecodedDataStoreKey, QueryError> {
let primary_key_names = descriptor.primary_key_names();
if row.len() != primary_key_names.len() {
return Err(QueryError::invariant());
}
DecodedDataStoreKey::try_from_structural_key_values(entity_tag, row)
.map_err(QueryError::execute)
}
fn execute_sql_update_statement_with_contract(
&self,
statement: &SqlUpdateStatement,
catalog: Option<&AcceptedSchemaCatalogContext>,
execution_contract: SqlUpdateExecutionContract,
) -> Result<SqlStatementResult, QueryError> {
self.with_checked_accepted_write_descriptor_for_returning(
catalog,
Some(statement.entity.as_str()),
statement.returning.as_ref(),
|catalog, descriptor| {
let (authority, schema_info) =
Self::accepted_sql_write_authority_schema_info(catalog);
let entity_tag = catalog.identity().entity_tag();
let selector = execution_contract
.selector(self.sql_update_selector_query(&schema_info, statement)?);
let patch = Self::sql_structural_patch(&descriptor, statement)?;
let write_context = AcceptedWriteContext::new(Timestamp::now());
let candidate_bounds = execution_contract.candidate_bounds();
let scan_budget = execution_contract.scan_budget()?;
let rows = self.collect_bounded_sql_write_mutation_batch_from_structural_query(
catalog.snapshot(),
authority,
&selector,
candidate_bounds,
scan_budget,
|row| {
let key =
Self::sql_write_key_from_projected_row(entity_tag, &descriptor, row)?;
Ok((
AcceptedStructuralMutationTarget::expected(key),
patch.clone(),
))
},
)?;
self.execute_sql_write_mutation_batch(
catalog,
&descriptor,
SqlWriteMutationExecution::from_bounded_batch(
rows,
candidate_bounds,
MutationMode::Update,
write_context,
execution_contract.returning_bounds(statement.returning.is_some()),
)?,
statement.returning.as_ref(),
)
},
)
}
fn schema_derived_sql_update_policy_result(
&self,
dispatch: &SqlStatementDispatch<'_>,
policy: SqlUpdateExposurePolicy,
) -> Result<SqlUpdatePolicyResult, QueryError> {
let entity_name = dispatch.entity_name();
self.with_checked_accepted_write_descriptor_for_returning(
None,
entity_name,
None,
|catalog, descriptor| {
with_accepted_sql_update_policy_context(&descriptor, |context| {
PreparationWork::run(
self.db.request_execution_scope(),
icydb_diagnostic_code::DiagnosticExecutionLane::Mutation,
|work| {
classify_sql_update_policy_for_entity(
dispatch,
catalog.snapshot().persisted_snapshot().entity_name(),
policy,
context,
work,
)
},
)
})
},
)
}
fn schema_derived_sql_update_plan(
&self,
dispatch: &SqlStatementDispatch<'_>,
policy: SqlUpdateExposurePolicy,
) -> Result<SqlValidatedUpdatePlan, QueryError> {
let result = self.schema_derived_sql_update_policy_result(dispatch, policy)?;
result.map_err(|_| QueryError::unsupported_query())
}
#[doc(hidden)]
pub(in crate::db) fn execute_validated_sql_public_primary_key_update(
&self,
plan: &SqlPublicPrimaryKeyUpdatePlan,
) -> Result<SqlStatementResult, QueryError> {
self.execute_sql_update_statement_with_contract(
plan.statement(),
None,
SqlUpdateExecutionContract::Validated(plan.execution_bounds()),
)
}
#[doc(hidden)]
pub(in crate::db) fn execute_validated_sql_public_bounded_update(
&self,
plan: &SqlPublicBoundedUpdatePlan,
) -> Result<SqlStatementResult, QueryError> {
self.execute_sql_update_statement_with_contract(
plan.statement(),
None,
SqlUpdateExecutionContract::Validated(plan.execution_bounds()),
)
}
#[doc(hidden)]
pub(in crate::db) fn execute_validated_sql_trusted_exact_update(
&self,
plan: &SqlTrustedExactUpdatePlan,
) -> Result<SqlStatementResult, QueryError> {
self.execute_sql_update_statement_with_contract(
plan.statement(),
None,
SqlUpdateExecutionContract::Exact {
policy: plan.policy(),
bounds: plan.execution_bounds(),
},
)
}
#[doc(hidden)]
pub fn execute_sql_public_primary_key_update(
&self,
dispatch: &SqlStatementDispatch<'_>,
) -> Result<SqlStatementResult, QueryError> {
let plan = self.schema_derived_sql_update_plan(
dispatch,
SqlUpdateExposurePolicy::PublicPrimaryKeyOnly,
)?;
let SqlValidatedUpdatePlan::PublicPrimaryKeyOnly(plan) = plan else {
return Err(QueryError::invariant());
};
self.execute_validated_sql_public_primary_key_update(&plan)
}
#[doc(hidden)]
pub fn execute_sql_public_bounded_update(
&self,
dispatch: &SqlStatementDispatch<'_>,
) -> Result<SqlStatementResult, QueryError> {
let plan = self.schema_derived_sql_update_plan(
dispatch,
SqlUpdateExposurePolicy::PublicBoundedDeterministic,
)?;
let SqlValidatedUpdatePlan::PublicBoundedDeterministic(plan) = plan else {
return Err(QueryError::invariant());
};
self.execute_validated_sql_public_bounded_update(&plan)
}
pub fn execute_trusted_sql_prefix_update(
&self,
sql: &str,
) -> Result<SqlStatementResult, QueryError> {
let dispatch = sql_statement_dispatch(sql)?;
self.execute_trusted_sql_prefix_update_dispatch(&dispatch)
}
#[doc(hidden)]
pub fn execute_trusted_sql_prefix_update_dispatch(
&self,
dispatch: &SqlStatementDispatch<'_>,
) -> Result<SqlStatementResult, QueryError> {
self.execute_sql_public_bounded_update(dispatch)
}
pub fn execute_trusted_sql_exact_update(
&self,
sql: &str,
require_affected_at_most: u32,
) -> Result<SqlStatementResult, QueryError> {
let dispatch = sql_statement_dispatch(sql)?;
self.execute_trusted_sql_exact_update_dispatch(&dispatch, require_affected_at_most)
}
#[doc(hidden)]
pub fn execute_trusted_sql_exact_update_dispatch(
&self,
dispatch: &SqlStatementDispatch<'_>,
require_affected_at_most: u32,
) -> Result<SqlStatementResult, QueryError> {
let policy =
SqlExactUpdatePolicy::try_new(require_affected_at_most).map_err(|rejection| {
sql_exact_update_policy_error(require_affected_at_most, rejection)
})?;
let result = self.schema_derived_sql_update_policy_result(
dispatch,
SqlUpdateExposurePolicy::TrustedExact(policy),
)?;
let plan = require_sql_exact_update_plan(result)?;
self.execute_validated_sql_trusted_exact_update(&plan)
}
}
#[cfg(test)]
mod tests {
use super::{
SqlExactUpdatePolicy, SqlExactUpdatePolicyRejection, sql_exact_update_policy_error,
};
use icydb_diagnostic_code::DiagnosticFactTag;
#[test]
fn exact_update_assertion_errors_retain_requested_and_bound_counts() {
let required =
sql_exact_update_policy_error(0, SqlExactUpdatePolicyRejection::AssertionRequired);
assert_eq!(
required.diagnostic_facts(),
vec![
(DiagnosticFactTag::ActualCount, 0),
(DiagnosticFactTag::Minimum, 1),
],
);
let requested = SqlExactUpdatePolicy::max_affected_rows() + 1;
let too_high = sql_exact_update_policy_error(
requested,
SqlExactUpdatePolicyRejection::AssertionTooHigh,
);
assert_eq!(
too_high.diagnostic_facts(),
vec![
(DiagnosticFactTag::ActualCount, u64::from(requested)),
(
DiagnosticFactTag::Limit,
u64::from(SqlExactUpdatePolicy::max_affected_rows()),
),
],
);
}
}