use std::collections::BTreeMap;
use std::net::SocketAddr;
use crate::bridge::envelope::PhysicalPlan;
use crate::types::{TenantId, VShardId};
use nodedb_physical::physical_plan::MetaOp;
use nodedb_physical::physical_task::{PhysicalTask, PostSetOp};
use super::outcome::TxnDataPlane;
use super::state::TransactionState;
use super::store::SessionStore;
#[derive(Debug)]
pub enum SavepointError {
NoActiveTransaction,
NotFound { message: String },
}
fn require_active_txn(sessions: &SessionStore, addr: &SocketAddr) -> Result<(), SavepointError> {
if sessions.transaction_state(addr) == TransactionState::Idle {
return Err(SavepointError::NoActiveTransaction);
}
Ok(())
}
async fn dispatch_overlay_savepoint(
tenant_id: TenantId,
vshard_id: VShardId,
dp: &impl TxnDataPlane,
op: MetaOp,
) -> Option<Vec<u8>> {
let task = PhysicalTask {
tenant_id,
vshard_id,
database_id: crate::types::DatabaseId::DEFAULT,
plan: PhysicalPlan::Meta(op),
post_set_op: PostSetOp::None,
txn_id: None,
};
match dp.dispatch_no_wal(task, None).await {
Ok(resp) => Some(resp.payload.to_vec()),
Err(e) => {
tracing::warn!(error = %e, "savepoint overlay meta-op dispatch failed");
None
}
}
}
fn decode_markers(payload: Option<Vec<u8>>) -> (usize, usize) {
payload
.filter(|bytes| bytes.len() == 16)
.map(|bytes| {
let mut value = [0u8; 8];
value.copy_from_slice(&bytes[..8]);
let mut graph = [0u8; 8];
graph.copy_from_slice(&bytes[8..16]);
(
u64::from_le_bytes(value) as usize,
u64::from_le_bytes(graph) as usize,
)
})
.unwrap_or((0, 0))
}
pub async fn run_savepoint(
sessions: &SessionStore,
addr: &SocketAddr,
tenant_id: TenantId,
dp: &impl TxnDataPlane,
name: &str,
) -> Result<(), SavepointError> {
require_active_txn(sessions, addr)?;
let (txn_id, vshards) = sessions.txn_identity(addr);
let mut markers: BTreeMap<VShardId, (usize, usize)> = BTreeMap::new();
if let Some(txn_id) = txn_id {
for vshard_id in vshards {
let payload = dispatch_overlay_savepoint(
tenant_id,
vshard_id,
dp,
MetaOp::MarkSavepoint { txn_id },
)
.await;
markers.insert(vshard_id, decode_markers(payload));
}
}
sessions.create_savepoint(addr, name.to_string(), markers);
Ok(())
}
pub fn run_release_savepoint(
sessions: &SessionStore,
addr: &SocketAddr,
name: &str,
) -> Result<(), SavepointError> {
require_active_txn(sessions, addr)?;
sessions
.release_savepoint(addr, name)
.map_err(|e| SavepointError::NotFound {
message: e.to_string(),
})
}
pub async fn run_rollback_to_savepoint(
sessions: &SessionStore,
addr: &SocketAddr,
tenant_id: TenantId,
dp: &impl TxnDataPlane,
name: &str,
) -> Result<(), SavepointError> {
require_active_txn(sessions, addr)?;
let markers =
sessions
.rollback_to_savepoint(addr, name)
.map_err(|e| SavepointError::NotFound {
message: e.to_string(),
})?;
let (txn_id, vshards) = sessions.txn_identity(addr);
if let Some(txn_id) = txn_id {
for vshard_id in vshards {
let (value_marker, graph_marker) = markers.get(&vshard_id).copied().unwrap_or((0, 0));
dispatch_overlay_savepoint(
tenant_id,
vshard_id,
dp,
MetaOp::RollbackToSavepoint {
txn_id,
value_marker: value_marker as u64,
graph_marker: graph_marker as u64,
},
)
.await;
}
}
Ok(())
}
pub enum DeferredOffsetCmd {
Single {
stream: String,
group: String,
partition_id: u32,
lsn: u64,
},
Batch { stream: String, group: String },
}
pub fn parse_deferred_offset(sql: &str, upper: &str) -> Option<DeferredOffsetCmd> {
if !(upper.starts_with("COMMIT OFFSET ") || upper.starts_with("COMMIT OFFSETS ")) {
return None;
}
let parts: Vec<&str> = sql.split_whitespace().collect();
if parts.len() >= 11
&& parts[2].eq_ignore_ascii_case("PARTITION")
&& parts[4].eq_ignore_ascii_case("AT")
&& parts[6].eq_ignore_ascii_case("ON")
{
let partition_id: u32 = parts[3].parse().unwrap_or(0);
let lsn: u64 = parts[5].parse().unwrap_or(0);
return Some(DeferredOffsetCmd::Single {
stream: parts[7].to_lowercase(),
group: parts[10].to_lowercase(),
partition_id,
lsn,
});
}
if parts.len() >= 7
&& parts[1].eq_ignore_ascii_case("OFFSETS")
&& parts[2].eq_ignore_ascii_case("ON")
{
return Some(DeferredOffsetCmd::Batch {
stream: parts[3].to_lowercase(),
group: parts[6].to_lowercase(),
});
}
None
}