use std::{borrow::Cow, sync::Arc};
use azure_core::fmt::SafeDebug;
use azure_core::Bytes;
use azure_data_cosmos_driver::models as driver_models;
use serde::{de::DeserializeOwned, Serialize};
use crate::clients::{ClientContext, ContainerClient};
use crate::diagnostics::DiagnosticsContext;
use crate::models::{PartitionKey, PatchInstructions, ResponseHeaders};
use crate::options::{Precondition, SessionToken};
#[derive(Clone, Default)]
#[non_exhaustive]
pub struct DistributedTransactionOperationOptions {
pub session_token: Option<SessionToken>,
pub precondition: Option<Precondition>,
}
impl DistributedTransactionOperationOptions {
pub fn with_session_token(mut self, session_token: impl Into<SessionToken>) -> Self {
self.session_token = Some(session_token.into());
self
}
pub fn with_precondition(mut self, precondition: Precondition) -> Self {
self.precondition = Some(precondition);
self
}
}
#[derive(Clone, Default)]
#[non_exhaustive]
pub struct DistributedTransactionPatchOperationOptions {
pub session_token: Option<SessionToken>,
pub precondition: Option<Precondition>,
pub filter_predicate: Option<Cow<'static, str>>,
}
impl DistributedTransactionPatchOperationOptions {
pub fn with_session_token(mut self, session_token: impl Into<SessionToken>) -> Self {
self.session_token = Some(session_token.into());
self
}
pub fn with_precondition(mut self, precondition: Precondition) -> Self {
self.precondition = Some(precondition);
self
}
pub fn with_filter_predicate(mut self, predicate: impl Into<Cow<'static, str>>) -> Self {
self.filter_predicate = Some(predicate.into());
self
}
}
#[derive(Clone, SafeDebug)]
#[safe(true)]
pub struct DistributedWriteTransaction {
operations: Vec<driver_models::DistributedTransactionOperation>,
}
impl DistributedWriteTransaction {
pub fn new() -> Self {
Self {
operations: Vec::new(),
}
}
pub(crate) fn into_operations(self) -> Vec<driver_models::DistributedTransactionOperation> {
self.operations
}
pub fn create_item<T: Serialize>(
mut self,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
item: T,
options: Option<DistributedTransactionOperationOptions>,
) -> crate::Result<Self> {
let body = serde_json::to_vec(&item)?;
self.operations.push(operation_with_options(
driver_models::DistributedTransactionOperationKind::Create,
container,
partition_key,
item_id,
Some(Bytes::from(body)),
options,
));
Ok(self)
}
pub fn replace_item<T: Serialize>(
mut self,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
item: T,
options: Option<DistributedTransactionOperationOptions>,
) -> crate::Result<Self> {
let body = serde_json::to_vec(&item)?;
self.operations.push(operation_with_options(
driver_models::DistributedTransactionOperationKind::Replace,
container,
partition_key,
item_id,
Some(Bytes::from(body)),
options,
));
Ok(self)
}
pub fn upsert_item<T: Serialize>(
mut self,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
item: T,
options: Option<DistributedTransactionOperationOptions>,
) -> crate::Result<Self> {
let body = serde_json::to_vec(&item)?;
self.operations.push(operation_with_options(
driver_models::DistributedTransactionOperationKind::Upsert,
container,
partition_key,
item_id,
Some(Bytes::from(body)),
options,
));
Ok(self)
}
pub fn delete_item(
mut self,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
options: Option<DistributedTransactionOperationOptions>,
) -> Self {
self.operations.push(operation_with_options(
driver_models::DistributedTransactionOperationKind::Delete,
container,
partition_key,
item_id,
None,
options,
));
self
}
pub fn patch_item(
mut self,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
patch: PatchInstructions,
options: Option<DistributedTransactionPatchOperationOptions>,
) -> crate::Result<Self> {
let body = serde_json::to_vec(&patch)?;
self.operations.push(patch_operation_with_options(
container,
partition_key,
item_id,
Bytes::from(body),
options,
)?);
Ok(self)
}
}
impl Default for DistributedWriteTransaction {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone, SafeDebug)]
#[safe(true)]
pub struct DistributedReadTransaction {
operations: Vec<driver_models::DistributedTransactionOperation>,
}
impl DistributedReadTransaction {
pub fn new() -> Self {
Self {
operations: Vec::new(),
}
}
pub(crate) fn into_operations(self) -> Vec<driver_models::DistributedTransactionOperation> {
self.operations
}
pub fn read_item(
mut self,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
options: Option<DistributedTransactionOperationOptions>,
) -> Self {
self.operations.push(operation_with_options(
driver_models::DistributedTransactionOperationKind::Read,
container,
partition_key,
item_id,
None,
options,
));
self
}
}
impl Default for DistributedReadTransaction {
fn default() -> Self {
Self::new()
}
}
fn operation_with_options(
kind: driver_models::DistributedTransactionOperationKind,
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
body: Option<Bytes>,
options: Option<DistributedTransactionOperationOptions>,
) -> driver_models::DistributedTransactionOperation {
let mut operation = driver_models::DistributedTransactionOperation::new(
kind,
driver_models::DistributedTransactionTarget::new(
container.container_reference().clone(),
partition_key,
item_id,
),
);
if let Some(body) = body {
operation = operation.with_resource_body(body);
}
if let Some(options) = options {
if let Some(session_token) = options.session_token {
operation = operation.with_session_token(session_token);
}
if let Some(precondition) = options.precondition {
operation = operation.with_precondition(precondition);
}
}
operation
}
fn patch_operation_with_options(
container: &ContainerClient,
partition_key: impl Into<PartitionKey>,
item_id: impl Into<std::borrow::Cow<'static, str>>,
body: Bytes,
options: Option<DistributedTransactionPatchOperationOptions>,
) -> crate::Result<driver_models::DistributedTransactionOperation> {
let mut operation = driver_models::DistributedTransactionOperation::new(
driver_models::DistributedTransactionOperationKind::Patch,
driver_models::DistributedTransactionTarget::new(
container.container_reference().clone(),
partition_key,
item_id,
),
)
.with_resource_body(body);
if let Some(options) = options {
if let Some(session_token) = options.session_token {
operation = operation.with_session_token(session_token);
}
if let Some(precondition) = options.precondition {
operation = operation.with_precondition(precondition);
}
if let Some(predicate) = options.filter_predicate {
operation = operation.with_patch_filter_predicate(predicate);
}
}
Ok(operation)
}
pub(crate) async fn commit_distributed_write(
context: &ClientContext,
transaction: DistributedWriteTransaction,
) -> crate::Result<DistributedTransactionResponse> {
let operations = transaction.into_operations();
validate_transaction_account(context, &operations)?;
let request = driver_models::DistributedTransactionRequest::new(
driver_models::DistributedTransactionType::Write,
operations,
);
let response = context
.driver
.execute_distributed_transaction(request, Default::default())
.await?;
Ok(DistributedTransactionResponse::from_driver(response))
}
pub(crate) async fn execute_distributed_read(
context: &ClientContext,
transaction: DistributedReadTransaction,
) -> crate::Result<DistributedTransactionResponse> {
let operations = transaction.into_operations();
validate_transaction_account(context, &operations)?;
let request = driver_models::DistributedTransactionRequest::new(
driver_models::DistributedTransactionType::Read,
operations,
);
let response = context
.driver
.execute_distributed_transaction(request, Default::default())
.await?;
Ok(DistributedTransactionResponse::from_driver(response))
}
fn validate_transaction_account(
context: &ClientContext,
operations: &[driver_models::DistributedTransactionOperation],
) -> crate::Result<()> {
if operations
.iter()
.all(|operation| operation.target.container.account() == context.driver.account())
{
Ok(())
} else {
Err(crate::DriverCosmosError::builder()
.with_status(crate::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(
"distributed transaction operations must target containers from the same Cosmos account as the committing client",
)
.build()
.into())
}
}
#[derive(Clone, SafeDebug)]
#[safe(true)]
pub struct DistributedTransactionResponse {
inner: driver_models::DistributedTransactionResponse,
headers: ResponseHeaders,
}
impl DistributedTransactionResponse {
fn from_driver(inner: driver_models::DistributedTransactionResponse) -> Self {
let headers = inner.headers.clone().into();
Self { inner, headers }
}
pub fn status(&self) -> crate::CosmosStatus {
let status = crate::CosmosStatus::new(self.inner.status_code);
match self.inner.sub_status_code {
Some(sub_status) => status.with_sub_status(sub_status.value()),
None => status,
}
}
pub fn is_success_status_code(&self) -> bool {
self.inner.is_success_status_code()
}
pub fn is_completed_status_code(&self) -> bool {
self.inner.is_completed_status_code()
}
pub fn len(&self) -> usize {
self.inner.len()
}
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
pub fn operation_result(
&self,
index: usize,
) -> Option<DistributedTransactionOperationResult<'_>> {
self.inner
.operation_results
.get(index)
.map(|inner| DistributedTransactionOperationResult { inner })
}
pub fn headers(&self) -> &ResponseHeaders {
&self.headers
}
pub fn diagnostic_string(&self) -> Option<&str> {
self.inner.diagnostic_string.as_deref()
}
pub fn idempotency_token(&self) -> String {
self.inner.idempotency_token.to_string()
}
pub fn is_retriable(&self) -> bool {
self.inner.is_retriable
}
pub fn error_message(&self) -> Option<&str> {
self.inner.error_message.as_deref()
}
pub fn diagnostics(&self) -> Option<Arc<DiagnosticsContext>> {
self.inner.diagnostics.clone()
}
pub fn activity_id(&self) -> Option<&str> {
self.inner.activity_id.as_ref().map(|id| id.as_str())
}
pub fn request_charge(&self) -> Option<f64> {
self.inner.request_charge.map(|charge| charge.value())
}
pub fn retry_after_ms(&self) -> Option<u64> {
self.inner.retry_after_ms
}
}
#[derive(Clone, Copy, SafeDebug)]
#[safe(true)]
pub struct DistributedTransactionOperationResult<'a> {
inner: &'a driver_models::DistributedTransactionOperationResult,
}
impl DistributedTransactionOperationResult<'_> {
pub fn index(&self) -> usize {
self.inner.index
}
pub fn status_code(&self) -> azure_core::http::StatusCode {
self.inner.status_code
}
pub fn sub_status_code(&self) -> Option<crate::SubStatusCode> {
self.inner.sub_status_code
}
pub fn is_success_status_code(&self) -> bool {
self.inner.is_success_status_code()
}
pub fn is_completed_status_code(&self) -> bool {
self.inner.is_completed_status_code()
}
pub fn etag(&self) -> Option<&azure_core::http::Etag> {
self.inner.etag.as_ref()
}
pub fn session_token(&self) -> Option<&SessionToken> {
self.inner.session_token.as_ref()
}
pub fn partition_key_range_id(&self) -> Option<&str> {
self.inner.partition_key_range_id.as_deref()
}
pub fn request_charge(&self) -> Option<f64> {
self.inner.request_charge.map(|charge| charge.value())
}
pub fn resource<T: DeserializeOwned>(&self) -> crate::Result<Option<T>> {
match &self.inner.resource_body {
driver_models::DistributedTransactionResultBody::None => Ok(None),
driver_models::DistributedTransactionResultBody::Bytes(bytes) => {
serde_json::from_slice(bytes).map(Some).map_err(Into::into)
}
_ => Ok(None),
}
}
}