use std::collections::HashMap;
use std::sync::Arc;
use alopex_core::kv::{
AnyKV, OwnedKVTransactionAdapter, OwnedReadOptions, OwnedReadSession,
OwnedSessionFactory as CoreOwnedSessionFactory, OwnedTransactionSession,
};
use alopex_core::vector::hnsw::{HnswIndex, HnswTransactionState};
use alopex_core::TxnMode;
use alopex_sql::catalog::CatalogOverlay;
use alopex_sql::storage::LocalRangeChangeJournal;
use crate::{Database, Error, Result};
#[derive(Clone)]
pub struct EmbeddedOwnedSessionFactory {
store: Arc<AnyKV>,
}
impl EmbeddedOwnedSessionFactory {
pub(crate) fn new(store: Arc<AnyKV>) -> Self {
Self { store }
}
pub fn begin_read(&self, options: OwnedReadOptions) -> Result<OwnedReadSession> {
self.store
.clone()
.begin_owned_read(options)
.map_err(Error::Core)
}
pub fn begin_transaction(&self, mode: TxnMode) -> Result<OwnedTransactionSession> {
self.store
.clone()
.begin_owned_transaction(mode)
.map_err(Error::Core)
}
}
impl Database {
pub fn owned_session_factory(&self) -> EmbeddedOwnedSessionFactory {
EmbeddedOwnedSessionFactory::new(self.store.clone())
}
pub fn begin_owned_read(&self, options: OwnedReadOptions) -> Result<OwnedReadSession> {
self.owned_session_factory().begin_read(options)
}
pub fn begin_owned_transaction(&self, mode: TxnMode) -> Result<OwnedTransactionSession> {
self.owned_session_factory().begin_transaction(mode)
}
pub fn begin_owned_embedded_transaction(
self: Arc<Self>,
mode: TxnMode,
) -> Result<OwnedEmbeddedTransaction> {
let session = self.begin_owned_transaction(mode)?;
let journal = if mode == TxnMode::ReadWrite
&& self.store.range_change_journal_capability()
== alopex_core::kv::RangeChangeJournalCapability::Supported
{
let scope = {
let catalog = self.sql_catalog.read().expect("catalog lock poisoned");
crate::sql_api::local_journal_scope(&*catalog)
};
Some(
session
.with_transaction(|transaction| {
let mut transaction = OwnedKVTransactionAdapter::new(transaction);
LocalRangeChangeJournal::capture(&mut transaction, scope)
})
.map_err(Error::Core)?,
)
} else {
None
};
Ok(OwnedEmbeddedTransaction {
db: self,
session,
overlay: CatalogOverlay::new(),
catalog_modified: false,
journal,
hnsw_indices: HashMap::new(),
vector_cache_invalidated: false,
})
}
}
pub struct OwnedEmbeddedTransaction {
pub(crate) db: Arc<Database>,
pub(crate) session: OwnedTransactionSession,
pub(crate) overlay: CatalogOverlay,
pub(crate) catalog_modified: bool,
pub(crate) journal: Option<LocalRangeChangeJournal>,
pub(crate) hnsw_indices: HashMap<String, (HnswIndex, HnswTransactionState)>,
pub(crate) vector_cache_invalidated: bool,
}
impl OwnedEmbeddedTransaction {
pub fn session(&self) -> OwnedTransactionSession {
self.session.clone()
}
pub fn get(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
self.session
.with_transaction(|transaction| transaction.get(&key.to_vec()))
.map_err(Error::Core)
}
pub fn put(&mut self, key: &[u8], value: &[u8]) -> Result<()> {
self.vector_cache_invalidated = true;
self.session
.with_transaction(|transaction| transaction.put(key.to_vec(), value.to_vec()))
.map_err(Error::Core)
}
pub fn delete(&mut self, key: &[u8]) -> Result<()> {
self.vector_cache_invalidated = true;
self.session
.with_transaction(|transaction| transaction.delete(key.to_vec()))
.map_err(Error::Core)
}
pub fn execute_sql(&mut self, sql: &str) -> Result<crate::SqlResult> {
crate::sql_api::execute_sql_owned(self, sql)
}
pub fn preflight_sql_stream(&self, sql: &str) -> Result<crate::OwnedSqlStreamPlan> {
crate::OwnedSqlStreamPlan::preflight_in_transaction(&self.db, &self.overlay, sql)
}
pub fn commit(&mut self) -> Result<()> {
let mut preparation = Ok(());
let journal = self.journal.take();
self.session
.with_transaction(|transaction| {
let mut transaction = alopex_core::kv::any::AnyKVTransaction::Owned(
OwnedKVTransactionAdapter::new(transaction),
);
for (index, state) in self.hnsw_indices.values_mut() {
if preparation.is_ok() {
preparation = index
.commit_staged(&mut transaction, state)
.map_err(Error::Core);
}
}
let mut catalog = self.db.sql_catalog.write().expect("catalog lock poisoned");
if preparation.is_ok() {
preparation = catalog
.persist_overlay(&mut transaction, &self.overlay)
.map_err(|error| Error::Sql(error.into()));
}
if preparation.is_ok() {
if let Some(journal) = journal {
preparation = journal
.stage(&mut transaction)
.map(|_| ())
.map_err(Error::Core);
}
}
Ok(())
})
.map_err(Error::Core)?;
preparation?;
self.session.commit().map_err(Error::Core)?;
let overlay = std::mem::take(&mut self.overlay);
let mut catalog = self.db.sql_catalog.write().expect("catalog lock poisoned");
catalog.apply_overlay(overlay);
drop(catalog);
if self.catalog_modified {
self.db.invalidate_table_info_cache();
}
if self.vector_cache_invalidated {
let mut cache = self
.db
.vector_cache
.write()
.expect("vector cache lock poisoned");
*cache = None;
}
if !self.hnsw_indices.is_empty() {
let hnsw_indices = std::mem::take(&mut self.hnsw_indices);
let mut cache = self
.db
.hnsw_cache
.write()
.expect("hnsw cache lock poisoned");
for (name, (index, _)) in hnsw_indices {
cache.insert(name, Arc::new(index));
}
}
Ok(())
}
pub fn rollback(&mut self) -> Result<()> {
self.session.rollback().map_err(Error::Core)?;
for (index, state) in self.hnsw_indices.values_mut() {
let _ = index.rollback(state);
}
self.hnsw_indices.clear();
self.overlay = CatalogOverlay::default();
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use alopex_core::kv::OwnedReadOptions;
use alopex_core::txn::OwnedLeaseOutcome;
use alopex_core::TxnMode;
use std::sync::Arc;
#[test]
fn embedded_factory_uses_owned_lifecycle_without_changing_borrowed_transactions() {
let database = Database::new();
let session = database
.begin_owned_transaction(TxnMode::ReadWrite)
.unwrap();
let lease = session.acquire_lease().unwrap();
lease
.with_transaction(|transaction| transaction.put(b"owned".to_vec(), b"value".to_vec()))
.unwrap();
assert!(session.commit().is_err());
lease.finish(OwnedLeaseOutcome::Exhausted).unwrap();
session.commit().unwrap();
let mut borrowed = database.begin(TxnMode::ReadOnly).unwrap();
assert_eq!(
borrowed.get(b"owned".as_ref()).unwrap(),
Some(b"value".to_vec())
);
borrowed.commit().unwrap();
let read = database
.owned_session_factory()
.begin_read(OwnedReadOptions::default())
.unwrap();
let lease = read.acquire_lease().unwrap();
let mut cursor = lease
.with_transaction(|transaction| transaction.scan_prefix(b"own"))
.unwrap();
assert_eq!(
cursor.next_entry().unwrap(),
Some((b"owned".to_vec(), b"value".to_vec()))
);
assert_eq!(cursor.next_entry().unwrap(), None);
cursor.close().unwrap();
drop(cursor);
lease.finish(OwnedLeaseOutcome::Exhausted).unwrap();
}
#[test]
fn dropped_embedded_owned_transaction_rolls_back_staged_writes() {
let database = Database::new();
let session = database
.begin_owned_transaction(TxnMode::ReadWrite)
.unwrap();
let lease = session.acquire_lease().unwrap();
lease
.with_transaction(|transaction| transaction.put(b"discard".to_vec(), b"value".to_vec()))
.unwrap();
lease.finish(OwnedLeaseOutcome::Exhausted).unwrap();
drop(session);
let mut borrowed = database.begin(TxnMode::ReadOnly).unwrap();
assert_eq!(borrowed.get(b"discard".as_ref()).unwrap(), None);
borrowed.commit().unwrap();
}
#[test]
fn owned_embedded_transaction_preserves_sql_visibility_until_explicit_commit() {
let database = Arc::new(Database::new());
database
.execute_sql("CREATE TABLE owned_sql (id INTEGER PRIMARY KEY, value TEXT)")
.unwrap();
let mut transaction = Arc::clone(&database)
.begin_owned_embedded_transaction(TxnMode::ReadWrite)
.unwrap();
transaction
.execute_sql("INSERT INTO owned_sql (id, value) VALUES (1, 'staged')")
.unwrap();
let alopex_sql::ExecutionResult::Query(query) = transaction
.execute_sql("SELECT value FROM owned_sql WHERE id = 1")
.unwrap()
else {
panic!("owned transaction select must return a query")
};
assert_eq!(query.rows.len(), 1);
let alopex_sql::ExecutionResult::Query(before_commit) = database
.execute_sql("SELECT value FROM owned_sql WHERE id = 1")
.unwrap()
else {
panic!("database select must return a query")
};
assert!(before_commit.rows.is_empty());
transaction.commit().unwrap();
let alopex_sql::ExecutionResult::Query(after_commit) = database
.execute_sql("SELECT value FROM owned_sql WHERE id = 1")
.unwrap()
else {
panic!("database select must return a query")
};
assert_eq!(after_commit.rows.len(), 1);
}
#[test]
fn owned_embedded_transaction_preserves_vector_and_similarity_workflows() {
let database = Arc::new(Database::new());
let mut transaction = Arc::clone(&database)
.begin_owned_embedded_transaction(TxnMode::ReadWrite)
.unwrap();
transaction
.upsert_vector(b"first", b"one", &[1.0, 0.0], alopex_core::Metric::Cosine)
.unwrap();
transaction
.upsert_vector(b"second", b"two", &[0.0, 1.0], alopex_core::Metric::Cosine)
.unwrap();
let results = transaction
.search_similar(&[1.0, 0.0], alopex_core::Metric::Cosine, 1, None)
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].key, b"first".to_vec());
transaction.commit().unwrap();
let mut reader = Arc::clone(&database)
.begin_owned_embedded_transaction(TxnMode::ReadOnly)
.unwrap();
assert_eq!(
reader
.get_vector(b"second", alopex_core::Metric::Cosine)
.unwrap(),
Some(vec![0.0, 1.0])
);
reader.rollback().unwrap();
}
}