use std::sync::Arc;
use std::time::{Duration, Instant};
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
use std::collections::HashMap;
#[cfg(all(test, any(feature = "ladybug", feature = "permission")))]
use super::DbPool;
#[cfg(feature = "permission")]
use super::audit::audit_admin_bypass;
use super::db_pool::DbPoolInner;
use super::{DatabaseConnection, DbConnection};
#[cfg(all(feature = "sql-parser", feature = "permission"))]
use crate::access::SqlParser;
#[cfg(feature = "sql-parser")]
use crate::access::is_ddl_operation;
#[cfg(feature = "sql-parser")]
use crate::access::{DdlGuard, DdlValidationResult};
#[cfg(feature = "permission")]
use crate::access::{PermissionAction, PermissionContext};
use crate::foundation::{DbError, DbResult};
use crate::i18n;
#[cfg(feature = "metrics")]
use crate::observability::MetricsCollector;
use async_trait::async_trait;
use sea_orm::{ConnectionTrait, DatabaseTransaction, ExecResult, TransactionTrait};
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
use tokio::sync::Mutex;
use tokio::sync::RwLock;
struct SessionState {
transaction: Option<Arc<DatabaseTransaction>>,
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
graph_transaction: Option<Box<dyn crate::database::graph::GraphTransaction + Send>>,
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
graph_txn_poisoned: bool,
last_write: Option<Instant>,
}
pub struct Session {
connection: Option<DbConnection>,
pool_inner: Arc<DbPoolInner>,
role: String,
#[cfg(feature = "permission")]
permission_ctx: PermissionContext,
state: RwLock<SessionState>,
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
graph_op_mutex: Mutex<()>,
#[cfg(feature = "metrics")]
metrics_collector: Option<Arc<MetricsCollector>>,
}
impl Session {
pub(crate) fn new(connection: DbConnection, pool_inner: Arc<DbPoolInner>, role: String) -> Self {
#[cfg(feature = "permission")]
let permission_ctx = PermissionContext::new(role.clone(), pool_inner.policy_cache.clone());
#[cfg(feature = "metrics")]
let metrics = pool_inner.metrics_collector.clone();
Session {
connection: Some(connection),
pool_inner,
role,
#[cfg(feature = "permission")]
permission_ctx,
state: RwLock::new(SessionState {
transaction: None,
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
graph_transaction: None,
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
graph_txn_poisoned: false,
last_write: None,
}),
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
graph_op_mutex: Mutex::new(()),
#[cfg(feature = "metrics")]
metrics_collector: metrics,
}
}
pub fn role(&self) -> &str {
&self.role
}
#[cfg(feature = "permission")]
pub fn permission_ctx(&self) -> &PermissionContext {
&self.permission_ctx
}
pub async fn mark_write(&self) {
let mut state = self.state.write().await;
state.last_write = Some(Instant::now());
}
#[cfg(feature = "permission")]
pub async fn check_permission(&self, table: &str, operation: &PermissionAction) -> Result<(), DbError> {
if self.role == self.pool_inner.admin_role {
audit_admin_bypass(&self.role, table, operation);
return Ok(());
}
if self.permission_ctx.check_table_access(table, operation).await {
Ok(())
} else {
Err(permission_denied(operation, table))
}
}
pub async fn is_in_transaction(&self) -> bool {
let state = self.state.read().await;
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
{
state.graph_transaction.is_some() || state.transaction.is_some()
}
#[cfg(not(any(feature = "ladybug", feature = "neo4j")))]
{
state.transaction.is_some()
}
}
pub async fn begin_transaction(&self) -> Result<(), DbError> {
{
let state = self.state.write().await;
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
if state.graph_transaction.is_some() {
return Err(DbError::Transaction("Already in graph transaction".to_string()));
}
if state.transaction.is_some() {
return Err(DbError::Transaction("Already in transaction".to_string()));
}
}
let conn = self.connection.as_ref().ok_or_else(|| {
DbError::Config("Connection not available - Session may have been invalidated".to_string())
})?;
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
if conn.is_graph() {
let graph = conn.as_graph()?;
let graph_txn = graph.begin_graph_txn().await.map_err(|e| {
DbError::Transaction(i18n::t("session-txn-begin-graph-failed", &[("error", e.to_string())]))
})?;
let has_conflict = {
let state = self.state.write().await;
state.graph_transaction.is_some()
};
if has_conflict {
let _ = graph_txn.rollback().await;
return Err(DbError::Transaction(
"Already in graph transaction (concurrent begin detected)".to_string(),
));
}
let mut state = self.state.write().await;
state.graph_transaction = Some(graph_txn);
return Ok(());
}
let conn = conn.as_sea_orm()?;
let transaction = conn
.begin()
.await
.map_err(|e| DbError::Transaction(i18n::t("session-txn-begin-failed", &[("error", e.to_string())])))?;
let has_conflict = {
let state = self.state.write().await;
state.transaction.is_some()
};
if has_conflict {
let _ = transaction.rollback().await;
return Err(DbError::Transaction(
"Already in transaction (concurrent begin detected)".to_string(),
));
}
let transaction = Arc::new(transaction);
let mut state = self.state.write().await;
state.transaction = Some(transaction);
Ok(())
}
pub async fn commit(&self) -> Result<(), DbError> {
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
{
let graph_txn = {
let mut state = self.state.write().await;
state.graph_transaction.take()
};
if let Some(graph_txn) = graph_txn {
graph_txn.commit().await.map_err(|e| {
DbError::Transaction(i18n::t("session-txn-commit-failed", &[("error", e.to_string())]))
})?;
let mut state = self.state.write().await;
state.last_write = None;
return Ok(());
}
}
let transaction_arc = {
let mut state = self.state.write().await;
state
.transaction
.take()
.ok_or_else(|| DbError::Transaction("No active transaction to commit".to_string()))?
};
let transaction = Arc::try_unwrap(transaction_arc).map_err(|_| {
DbError::Transaction("Cannot commit: transaction is in use by a concurrent query".to_string())
})?;
transaction
.commit()
.await
.map_err(|e| DbError::Transaction(e.to_string()))?;
let mut state = self.state.write().await;
state.last_write = None;
Ok(())
}
pub async fn rollback(&self) -> Result<(), DbError> {
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
{
let graph_txn = {
let mut state = self.state.write().await;
state.graph_transaction.take()
};
if let Some(graph_txn) = graph_txn {
graph_txn.rollback().await.map_err(|e| {
DbError::Transaction(i18n::t(
"session-txn-rollback-graph-failed",
&[("error", e.to_string())],
))
})?;
return Ok(());
}
}
let transaction_arc = {
let mut state = self.state.write().await;
if state.transaction.is_none() {
return Err(DbError::Transaction("Not in transaction".to_string()));
}
state
.transaction
.take()
.ok_or_else(|| DbError::Transaction("No active transaction to rollback".to_string()))?
};
let transaction = Arc::try_unwrap(transaction_arc).map_err(|_| {
DbError::Transaction("Cannot rollback: transaction is in use by a concurrent query".to_string())
})?;
transaction
.rollback()
.await
.map_err(|e| DbError::Transaction(i18n::t("session-txn-rollback-failed", &[("error", e.to_string())])))?;
Ok(())
}
pub async fn should_use_master(&self) -> bool {
let state = self.state.read().await;
if state.transaction.is_some() {
return true;
}
state
.last_write
.map(|t| t.elapsed() < Duration::from_secs(5))
.unwrap_or(false)
}
pub fn connection(&self) -> Result<&DatabaseConnection, DbError> {
self.connection
.as_ref()
.ok_or_else(|| DbError::Config("Connection not available - Session may have been invalidated".to_string()))?
.as_sea_orm()
}
#[allow(dead_code)]
#[cfg(feature = "migration")]
pub fn create_migration_executor(
&self,
db_type: crate::foundation::DatabaseType,
) -> Result<super::MigrationExecutor, DbError> {
let conn = self.connection()?.clone();
Ok(super::MigrationExecutor::new(conn, db_type))
}
pub async fn execute_raw(&self, sql: &str) -> DbResult<ExecResult> {
#[cfg(feature = "sql-parser")]
{
if is_ddl_operation(sql) {
return Err(DbError::Permission(
"DDL operations are not allowed in this context".to_string(),
));
}
}
#[cfg(not(feature = "sql-parser"))]
{
let _ = sql;
Err(DbError::Permission(
"execute_raw requires the sql-parser feature to be enabled".to_string(),
))
}
#[cfg(feature = "sql-parser")]
{
#[cfg(all(feature = "sql-parser", feature = "permission"))]
{
let parser = SqlParser::shared().await;
match parser.parse_operation_async(sql).await {
Ok(Some((table_name, action))) => {
if table_name.is_empty() || is_invalid_table_name(&table_name) {
return Err(DbError::Permission(
"Failed to extract table name for permission checking".to_string(),
));
}
if self.role == self.pool_inner.admin_role {
} else if !self.permission_ctx.check_table_access(&table_name, &action).await {
return Err(permission_denied(&action, &table_name));
}
}
Ok(None) => {
return Err(DbError::Permission(
"SQL statement requires a valid table name for permission checking".to_string(),
));
}
Err(_) => {
return Err(DbError::Permission(
"Failed to parse SQL statement for permission checking".to_string(),
));
}
}
}
let tx_opt: Option<Arc<DatabaseTransaction>> = {
let state = self.state.write().await;
state.transaction.clone()
};
#[cfg(feature = "retry")]
{
if let Some(ref policy) = self.pool_inner.config.retry_policy
&& crate::reliability::is_idempotent_operation(sql)
{
let mut last_error: Option<DbError> = None;
for attempt in 0..=policy.max_retries {
if attempt > 0 {
let backoff = Self::calculate_retry_backoff(policy, attempt - 1);
tokio::time::sleep(backoff).await;
}
let result = if let Some(ref tx) = tx_opt {
tx.execute_unprepared(sql).await.map_err(DbError::Connection)
} else {
let conn = self.connection()?;
conn.execute_unprepared(sql).await.map_err(DbError::Connection)
};
match result {
Ok(exec_result) => return Ok(exec_result),
Err(e) => last_error = Some(e),
}
}
return Err(last_error.unwrap());
}
}
if let Some(tx) = tx_opt {
return tx.execute_unprepared(sql).await.map_err(DbError::Connection);
}
let conn = self.connection()?;
conn.execute_unprepared(sql).await.map_err(DbError::Connection)
}
}
#[cfg(feature = "retry")]
fn calculate_retry_backoff(policy: &crate::reliability::RetryPolicy, attempt: u32) -> std::time::Duration {
use std::time::Duration;
let base_ms = policy.initial_backoff_ms as f64;
let backoff_ms = base_ms * policy.multiplier.powi(attempt as i32);
let capped_ms = backoff_ms.min(policy.max_backoff_ms as f64);
Duration::from_millis(capped_ms as u64)
}
pub async fn execute_raw_ddl(&self, sql: &str) -> DbResult<ExecResult> {
if self.role != self.pool_inner.admin_role {
return Err(DbError::Permission(format!(
"DDL operations are only allowed for admin role. Current role: '{}', Admin role: '{}'",
self.role, self.pool_inner.admin_role
)));
}
#[cfg(feature = "sql-parser")]
{
let guard = DdlGuard::new();
match guard.validate(sql) {
Ok(DdlValidationResult::Allowed) => {
}
Ok(DdlValidationResult::Forbidden(reason)) => {
return Err(DbError::Permission(i18n::t(
"session-ddl-not-allowed",
&[("reason", reason.to_string())],
)));
}
Ok(DdlValidationResult::ParseError(error)) => {
return Err(DbError::Config(i18n::t(
"session-ddl-parse-failed",
&[("error", error.to_string())],
)));
}
Err(error) => {
return Err(DbError::Config(i18n::t(
"session-ddl-validation-error",
&[("error", error.to_string())],
)));
}
}
}
let conn = self.connection()?;
conn.execute_unprepared(sql).await.map_err(DbError::Connection)
}
#[cfg(feature = "duckdb")]
pub async fn execute_duckdb(&self, sql: &str) -> DbResult<Vec<crate::database::DuckDbRow>> {
#[cfg(feature = "sql-parser")]
{
if is_ddl_operation(sql) {
return Err(DbError::Permission(
"DDL operations are not allowed in DuckDB query context".to_string(),
));
}
}
#[cfg(not(feature = "sql-parser"))]
{
let _ = sql;
Err(DbError::Permission(
"execute_duckdb requires the sql-parser feature to be enabled for security checks".to_string(),
))
}
#[cfg(feature = "sql-parser")]
{
#[cfg(all(feature = "sql-parser", feature = "permission"))]
{
let parser = SqlParser::shared().await;
match parser.parse_operation_async(sql).await {
Ok(Some((table_name, action))) => {
if table_name.is_empty() || is_invalid_table_name(&table_name) {
return Err(DbError::Permission(
"Failed to extract table name for permission checking".to_string(),
));
}
if self.role != self.pool_inner.admin_role
&& !self.permission_ctx.check_table_access(&table_name, &action).await
{
return Err(permission_denied(&action, &table_name));
}
}
Ok(None) => {
if self.role != self.pool_inner.admin_role {
return Err(DbError::Permission(
"SQL statement requires a valid table name for permission checking".to_string(),
));
}
}
Err(_) => {
return Err(DbError::Permission(
"Failed to parse SQL statement for permission checking".to_string(),
));
}
}
}
let conn = self
.connection
.as_ref()
.ok_or_else(|| DbError::Config("Connection not available".to_string()))?;
let duck_conn = conn.as_duckdb()?;
duck_conn.query(sql).await
}
}
#[cfg(feature = "duckdb")]
pub async fn execute_duckdb_raw(&self, sql: &str) -> DbResult<crate::database::DuckDbExecResult> {
#[cfg(feature = "sql-parser")]
{
if is_ddl_operation(sql) {
if self.role == self.pool_inner.admin_role {
let guard = DdlGuard::new();
match guard.validate(sql) {
Ok(DdlValidationResult::Allowed) => {
let conn = self
.connection
.as_ref()
.ok_or_else(|| DbError::Config("Connection not available".to_string()))?;
let duck_conn = conn.as_duckdb()?;
return duck_conn.execute(sql).await;
}
Ok(DdlValidationResult::Forbidden(reason)) => {
return Err(DbError::Permission(i18n::t(
"session-ddl-not-allowed",
&[("reason", reason.to_string())],
)));
}
Ok(DdlValidationResult::ParseError(error)) => {
return Err(DbError::Config(i18n::t(
"session-ddl-parse-failed",
&[("error", error.to_string())],
)));
}
Err(error) => {
return Err(DbError::Config(i18n::t(
"session-ddl-validation-error",
&[("error", error.to_string())],
)));
}
}
} else {
return Err(DbError::Permission(format!(
"DDL operations are only allowed for admin role in DuckDB context. Current role: '{}', Admin role: '{}'",
self.role, self.pool_inner.admin_role
)));
}
}
}
#[cfg(not(feature = "sql-parser"))]
{
let _ = sql;
Err(DbError::Permission(
"execute_duckdb_raw requires the sql-parser feature to be enabled for security checks".to_string(),
))
}
#[cfg(feature = "sql-parser")]
{
#[cfg(all(feature = "sql-parser", feature = "permission"))]
{
let parser = SqlParser::shared().await;
match parser.parse_operation_async(sql).await {
Ok(Some((table_name, action))) => {
if table_name.is_empty() || is_invalid_table_name(&table_name) {
return Err(DbError::Permission(
"Failed to extract table name for permission checking".to_string(),
));
}
if self.role != self.pool_inner.admin_role
&& !self.permission_ctx.check_table_access(&table_name, &action).await
{
return Err(permission_denied(&action, &table_name));
}
}
Ok(None) => {
if self.role != self.pool_inner.admin_role {
return Err(DbError::Permission(
"SQL statement requires a valid table name for permission checking".to_string(),
));
}
}
Err(_) => {
return Err(DbError::Permission(
"Failed to parse SQL statement for permission checking".to_string(),
));
}
}
}
let conn = self
.connection
.as_ref()
.ok_or_else(|| DbError::Config("Connection not available".to_string()))?;
let duck_conn = conn.as_duckdb()?;
duck_conn.execute(sql).await
}
}
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
pub async fn execute_cypher_with_params(
&self,
cypher: &str,
params: HashMap<String, serde_json::Value>,
) -> DbResult<crate::database::graph::GraphExecResult> {
self.execute_cypher_in_transaction(cypher, Some(params)).await
}
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
async fn execute_cypher_in_transaction(
&self,
cypher: &str,
mut params: Option<HashMap<String, serde_json::Value>>,
) -> DbResult<crate::database::graph::GraphExecResult> {
validate_cypher_safety(cypher)?;
#[cfg(feature = "permission")]
{
let graph_perm_ctx =
crate::access::permission::GraphPermissionContext::new(&self.role, &self.pool_inner.admin_role);
graph_perm_ctx.check_graph_access(crate::access::permission::PermissionAction::Traverse)?;
}
let conn = self.connection.as_ref().ok_or_else(|| {
DbError::Config("Connection not available - Session may have been invalidated".to_string())
})?;
let _graph_op_guard = self.graph_op_mutex.lock().await;
let graph_txn = {
let mut state = self.state.write().await;
if state.graph_txn_poisoned {
return Err(DbError::Transaction(
"Graph transaction is poisoned due to previous panic; \
Session must be dropped and recreated"
.to_string(),
));
}
state.graph_transaction.take()
};
if let Some(graph_txn) = graph_txn {
struct PoisonGuard<'a> {
state: &'a RwLock<SessionState>,
armed: bool,
}
impl<'a> Drop for PoisonGuard<'a> {
fn drop(&mut self) {
if self.armed
&& let Ok(mut state) = self.state.try_write()
{
state.graph_txn_poisoned = true;
}
}
}
let mut guard = PoisonGuard {
state: &self.state,
armed: true,
};
let result = if let Some(p) = params.take() {
graph_txn.execute_cypher_with_params(cypher, p).await
} else {
graph_txn.execute_cypher(cypher).await
};
guard.armed = false;
let mut state = self.state.write().await;
state.graph_transaction = Some(graph_txn);
return result;
}
let graph = conn.as_graph()?;
if let Some(p) = params.take() {
graph.execute_cypher_with_params(cypher, p).await
} else {
graph.execute_cypher(cypher).await
}
}
pub async fn execute(&self, sql: &str) -> DbResult<ExecResult> {
#[cfg(feature = "sql-parser")]
check_ddl_operation(sql)?;
#[cfg(feature = "permission")]
{
let start = Instant::now();
let parsed = parse_sql_for_permission(sql).await?;
match parsed {
Some((table_name, action)) => {
if table_name.is_empty() || is_invalid_table_name(&table_name) {
return Err(DbError::Permission(
"Failed to extract table name for permission checking".to_string(),
));
}
self.check_permission(&table_name, &action).await?;
let result = self.execute_raw(sql).await?;
self.record_metrics_and_mark_write(&action, start).await;
Ok(result)
}
None => {
let result = self.execute_raw(sql).await?;
Ok(result)
}
}
}
#[cfg(not(feature = "permission"))]
{
let result = self.execute_raw(sql).await?;
Ok(result)
}
}
#[cfg(feature = "permission")]
pub async fn execute_with_operation(&self, sql: &str, operation: &PermissionAction) -> DbResult<ExecResult> {
let start = Instant::now();
#[cfg(feature = "sql-parser")]
{
if is_ddl_operation(sql) {
return Err(DbError::Permission(
"DDL operations are not allowed in this context".to_string(),
));
}
}
let table_name: String = extract_table_name_via_parser(sql).await.unwrap_or_default();
#[cfg(feature = "permission")]
{
if !table_name.is_empty() && !self.permission_ctx.check_table_access(&table_name, operation).await {
return Err(permission_denied(operation, &table_name));
}
}
let result = self.execute_raw(sql).await?;
self.record_metrics_and_mark_write(operation, start).await;
Ok(result)
}
pub async fn batch_execute(&self, sqls: Vec<&str>) -> DbResult<Vec<DbResult<ExecResult>>> {
let mut results = Vec::new();
for sql in sqls {
let result = self.execute(sql).await;
results.push(result);
}
Ok(results)
}
pub async fn batch_execute_in_transaction(&self, sqls: Vec<&str>) -> DbResult<Vec<ExecResult>> {
self.begin_transaction().await?;
let result: DbResult<Vec<ExecResult>> = async {
let mut results = Vec::with_capacity(sqls.len());
for sql in sqls {
results.push(self.execute_raw(sql).await?);
}
Ok(results)
}
.await;
match result {
Ok(results) => {
self.commit().await?;
Ok(results)
}
Err(e) => {
match self.rollback().await {
Ok(()) => Err(e),
Err(rollback_err) => Err(DbError::Transaction(format!(
"batch failed: {}; rollback also failed: {}",
e, rollback_err
))),
}
}
}
}
#[cfg(all(feature = "metrics", feature = "permission"))]
fn record_query_metrics(&self, query_type: &str, duration: Duration, success: bool) {
if let Some(metrics) = &self.metrics_collector {
metrics.record_query(query_type, duration, success, None);
}
}
#[cfg(all(not(feature = "metrics"), feature = "permission"))]
fn record_query_metrics(&self, _query_type: &str, _duration: Duration, _success: bool) {
}
#[cfg(feature = "permission")]
async fn record_metrics_and_mark_write(&self, action: &PermissionAction, start: Instant) {
let duration = start.elapsed();
self.record_query_metrics(&format!("{:?}", action), duration, true);
if is_write_action(action) {
self.mark_write().await;
}
}
pub async fn check_table_permission(&self, _table_name: &str, _operation: &str) -> DbResult<()> {
#[cfg(feature = "permission")]
{
let action = match _operation {
"INSERT" => PermissionAction::Insert,
"SELECT" => PermissionAction::Select,
"UPDATE" => PermissionAction::Update,
"DELETE" => PermissionAction::Delete,
_ => {
return Err(DbError::Permission(i18n::t(
"session-unknown-operation",
&[("operation", _operation.to_string())],
)));
}
};
if self.role == self.pool_inner.admin_role {
audit_admin_bypass(&self.role, _table_name, &action);
} else if !self.permission_ctx.check_table_access(_table_name, &action).await {
return Err(permission_denied(_operation, _table_name));
}
}
Ok(())
}
#[cfg(feature = "metrics")]
pub fn record_metric(&self, operation: &str, table_name: &str, success: bool) {
if let Some(metrics) = &self.metrics_collector {
let bytes = Some(table_name.len() as u64);
metrics.record_query(operation, std::time::Duration::from_millis(0), success, bytes);
}
}
}
#[cfg(feature = "permission")]
fn is_invalid_table_name(table_name: &str) -> bool {
let table_name = table_name.trim();
if table_name.is_empty() {
return true;
}
for part in table_name.split('.') {
let part = part.trim();
if part.is_empty() {
return true;
}
let unquoted = part
.strip_prefix('"')
.and_then(|s| s.strip_suffix('"'))
.or_else(|| part.strip_prefix('`').and_then(|s| s.strip_suffix('`')))
.or_else(|| part.strip_prefix('\'').and_then(|s| s.strip_suffix('\'')))
.unwrap_or(part)
.trim();
if unquoted.is_empty() {
return true;
}
}
false
}
#[cfg(feature = "permission")]
fn permission_denied(action: &(impl std::fmt::Display + ?Sized), table: &(impl std::fmt::Display + ?Sized)) -> DbError {
DbError::Permission(i18n::t(
"session-permission-denied",
&[("action", action.to_string()), ("table", table.to_string())],
))
}
#[cfg(feature = "permission")]
fn is_write_action(action: &PermissionAction) -> bool {
matches!(
action,
PermissionAction::Insert | PermissionAction::Update | PermissionAction::Delete
)
}
#[cfg(any(feature = "ladybug", feature = "neo4j"))]
fn validate_cypher_safety(cypher: &str) -> DbResult<()> {
const MAX_CYPHER_BYTES: usize = 10_240;
if cypher.len() > MAX_CYPHER_BYTES {
return Err(DbError::Permission(format!(
"Cypher query exceeds maximum length ({} bytes, got {} bytes) - potential DoS payload",
MAX_CYPHER_BYTES,
cypher.len()
)));
}
let trimmed = cypher.trim();
let inner = trimmed.trim_end_matches(';').trim();
if inner.contains(';') {
return Err(DbError::Permission(
"Cypher query contains multiple statements (';' inside query) - potential injection".to_string(),
));
}
if let Some(pos) = cypher.find("//") {
let is_url_scheme = pos > 0 && {
let prev = cypher.as_bytes()[pos - 1];
prev.is_ascii_alphabetic()
};
if !is_url_scheme {
return Err(DbError::Permission(
"Cypher query contains line comment '//' - potential injection".to_string(),
));
}
}
if cypher.contains("/*") || cypher.contains("*/") {
return Err(DbError::Permission(
"Cypher query contains block comment '/* */' - potential injection".to_string(),
));
}
let cypher_lower = cypher.to_ascii_lowercase();
const DANGEROUS_CALLS: &[&str] = &["call apoc.", "call dbms.", "call db.", "call tx."];
for &dangerous in DANGEROUS_CALLS {
if cypher_lower.contains(dangerous) {
return Err(DbError::Permission(format!(
"Cypher query calls dangerous procedure ('{}') - potential privilege escalation",
dangerous
)));
}
}
Ok(())
}
#[cfg(feature = "sql-parser")]
fn check_ddl_operation(sql: &str) -> DbResult<()> {
if is_ddl_operation(sql) {
return Err(DbError::Permission(
"DDL operations are not allowed in this context".to_string(),
));
}
Ok(())
}
impl Drop for Session {
fn drop(&mut self) {
if let Some(conn) = self.connection.take() {
DbPoolInner::release_connection(&self.pool_inner, conn);
}
}
}
#[cfg(all(feature = "permission", feature = "sql-parser"))]
async fn extract_table_name_via_parser(sql: &str) -> Option<String> {
let parser = SqlParser::shared().await;
match parser.parse_operation_async(sql).await {
Ok(Some((table, _))) => {
if table.is_empty() || is_invalid_table_name(&table) {
None
} else {
Some(table)
}
}
Ok(None) => None,
Err(_) => None,
}
}
#[cfg(feature = "permission")]
async fn parse_sql_for_permission(sql: &str) -> DbResult<Option<(String, PermissionAction)>> {
let parser = SqlParser::shared().await;
match parser.parse_operation_async(sql).await {
Ok(Some((table, action))) => Ok(Some((table, action))),
Ok(None) => Ok(None),
Err(_) => Ok(None),
}
}
#[async_trait]
impl super::DatabaseSession for Session {
async fn execute(&self, sql: &str) -> crate::DbResult<ExecResult> {
Ok(self.execute(sql).await?)
}
async fn execute_raw(&self, sql: &str) -> crate::DbResult<ExecResult> {
Ok(self.execute_raw(sql).await?)
}
async fn execute_raw_ddl(&self, sql: &str) -> crate::DbResult<ExecResult> {
Ok(self.execute_raw_ddl(sql).await?)
}
async fn begin_transaction(&self) -> crate::DbResult<()> {
Ok(self.begin_transaction().await?)
}
async fn commit(&self) -> crate::DbResult<()> {
Ok(self.commit().await?)
}
async fn rollback(&self) -> crate::DbResult<()> {
Ok(self.rollback().await?)
}
fn role(&self) -> &str {
self.role()
}
async fn is_in_transaction(&self) -> bool {
self.is_in_transaction().await
}
}
#[cfg(all(test, feature = "ladybug"))]
mod graph_tests {
use super::*;
use crate::database::graph::{GraphExecResult, GraphValue};
async fn make_ladybug_pool() -> DbPool {
DbPool::new("ladybug::memory:")
.await
.expect("Failed to create Ladybug pool")
}
#[tokio::test]
async fn test_graph_session_is_in_transaction_initial_false() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
assert!(
!session.is_in_transaction().await,
"initial state should be no transaction"
);
}
#[tokio::test]
async fn test_graph_session_begin_sets_in_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session.begin_transaction().await.expect("begin should succeed");
assert!(
session.is_in_transaction().await,
"should be in transaction after begin"
);
}
#[tokio::test]
async fn test_graph_session_commit_clears_in_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session.begin_transaction().await.expect("begin");
session.commit().await.expect("commit");
assert!(
!session.is_in_transaction().await,
"should not be in transaction after commit"
);
}
#[tokio::test]
async fn test_graph_session_rollback_clears_in_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session.begin_transaction().await.expect("begin");
session.rollback().await.expect("rollback");
assert!(
!session.is_in_transaction().await,
"should not be in transaction after rollback"
);
}
#[tokio::test]
async fn test_graph_transaction_commit_e2e() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Person(name STRING, PRIMARY KEY(name))",
HashMap::new(),
)
.await
.expect("create node table");
session.begin_transaction().await.expect("begin");
session
.execute_cypher_with_params("CREATE (:Person {name: 'Alice'})", HashMap::new())
.await
.expect("create in txn");
session.commit().await.expect("commit");
let result = session
.execute_cypher_with_params("MATCH (p:Person) RETURN p.name AS name", HashMap::new())
.await
.expect("match after commit");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 1, "should see 1 person after commit");
let name = &q.rows[0].columns[0].1;
match name {
GraphValue::Scalar(serde_json::Value::String(s)) => assert_eq!(s, "Alice"),
other => panic!("expected String Scalar, got {other:?}"),
}
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
}
#[tokio::test]
async fn test_graph_transaction_rollback_e2e() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Person(name STRING, PRIMARY KEY(name))",
HashMap::new(),
)
.await
.expect("create node table");
session.begin_transaction().await.expect("begin");
session
.execute_cypher_with_params("CREATE (:Person {name: 'Bob'})", HashMap::new())
.await
.expect("create in txn");
session.rollback().await.expect("rollback");
let result = session
.execute_cypher_with_params("MATCH (p:Person) RETURN p.name AS name", HashMap::new())
.await
.expect("match after rollback");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 0, "should see 0 persons after rollback");
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
}
#[tokio::test]
async fn test_graph_double_begin_fails() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session.begin_transaction().await.expect("first begin");
let result = session.begin_transaction().await;
assert!(result.is_err(), "double begin should fail");
let err = result.unwrap_err();
assert!(
matches!(err, DbError::Transaction(ref msg) if msg.contains("Already in")),
"expected 'Already in' error, got {:?}",
err
);
}
#[tokio::test]
async fn test_graph_commit_without_transaction_fails() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
let result = session.commit().await;
assert!(result.is_err(), "commit without transaction should fail");
}
#[tokio::test]
async fn test_graph_rollback_without_transaction_fails() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
let result = session.rollback().await;
assert!(result.is_err(), "rollback without transaction should fail");
}
#[tokio::test]
async fn test_execute_cypher_without_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
let result = session
.execute_cypher_with_params("RETURN 1", HashMap::new())
.await
.expect("execute_cypher should succeed");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 1, "should return 1 row");
let value = &q.rows[0].columns[0].1;
match value {
GraphValue::Scalar(s) => assert_eq!(s, &serde_json::json!(1)),
other => panic!("expected Scalar, got {other:?}"),
}
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
}
#[tokio::test]
async fn test_execute_cypher_in_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Person(name STRING, age INT64, PRIMARY KEY(name))",
HashMap::new(),
)
.await
.expect("create table");
session.begin_transaction().await.expect("begin");
session
.execute_cypher_with_params("CREATE (:Person {name: 'Alice', age: 25})", HashMap::new())
.await
.expect("create in txn");
let result = session
.execute_cypher_with_params("MATCH (p:Person) RETURN p.name AS name, p.age AS age", HashMap::new())
.await
.expect("match in txn");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 1, "should see 1 person in txn");
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
session.commit().await.expect("commit");
}
#[tokio::test]
async fn test_execute_cypher_e2e_create_match() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Person(name STRING, age INT64, PRIMARY KEY(name))",
HashMap::new(),
)
.await
.expect("create node table");
session
.execute_cypher_with_params("CREATE (:Person {name: 'Alice', age: 25})", HashMap::new())
.await
.expect("create alice");
session
.execute_cypher_with_params("CREATE (:Person {name: 'Bob', age: 30})", HashMap::new())
.await
.expect("create bob");
let result = session
.execute_cypher_with_params(
"MATCH (p:Person) RETURN p.name AS name, p.age AS age ORDER BY name",
HashMap::new(),
)
.await
.expect("match");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 2, "should return 2 persons");
let name0 = &q.rows[0].columns[0].1;
match name0 {
GraphValue::Scalar(serde_json::Value::String(s)) => assert_eq!(s, "Alice"),
other => panic!("expected String Scalar, got {other:?}"),
}
let name1 = &q.rows[1].columns[0].1;
match name1 {
GraphValue::Scalar(serde_json::Value::String(s)) => assert_eq!(s, "Bob"),
other => panic!("expected String Scalar, got {other:?}"),
}
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
}
#[tokio::test]
async fn test_execute_cypher_invalid_returns_error() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
let result = session
.execute_cypher_with_params("INVALID CYPHER", HashMap::new())
.await;
assert!(result.is_err(), "invalid cypher should return error");
}
#[tokio::test]
async fn test_execute_cypher_multiple_in_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Person(name STRING, PRIMARY KEY(name))",
HashMap::new(),
)
.await
.expect("create table");
session.begin_transaction().await.expect("begin");
session
.execute_cypher_with_params("CREATE (:Person {name: 'A'})", HashMap::new())
.await
.expect("create A");
session
.execute_cypher_with_params("CREATE (:Person {name: 'B'})", HashMap::new())
.await
.expect("create B");
session
.execute_cypher_with_params("CREATE (:Person {name: 'C'})", HashMap::new())
.await
.expect("create C");
let result = session
.execute_cypher_with_params("MATCH (p:Person) RETURN count(p) AS cnt", HashMap::new())
.await
.expect("count in txn");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 1);
let cnt = &q.rows[0].columns[0].1;
match cnt {
GraphValue::Scalar(s) => assert_eq!(s, &serde_json::json!(3)),
other => panic!("expected Scalar, got {other:?}"),
}
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
session.commit().await.expect("commit");
}
#[cfg(feature = "permission")]
#[tokio::test]
async fn test_execute_cypher_non_admin_denied() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("system").await.expect("get_session");
let result = session.execute_cypher_with_params("RETURN 1", HashMap::new()).await;
assert!(result.is_err(), "non-admin role should be denied");
let err = result.unwrap_err();
assert!(
matches!(err, DbError::Permission(ref msg) if msg.contains("Graph operation denied")),
"expected Permission error, got {:?}",
err
);
}
#[cfg(feature = "permission")]
#[tokio::test]
async fn test_execute_cypher_admin_allowed() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
let result = session.execute_cypher_with_params("RETURN 42", HashMap::new()).await;
assert!(result.is_ok(), "admin role should be allowed");
}
}
#[cfg(test)]
#[cfg(all(feature = "permission", feature = "sqlite"))]
mod vuln_0001_tests {
use super::*;
#[cfg(all(feature = "permission", feature = "sqlite"))]
#[tokio::test]
async fn test_vuln_0001_admin_bypass_returns_ok_with_audit() {
let pool = DbPool::new("sqlite::memory:").await.expect("Failed to create pool");
let session = pool.get_session("admin").await.expect("get_session");
let result = session.check_permission("any_table", &PermissionAction::Select).await;
assert!(result.is_ok(), "admin bypass should return Ok");
let result = session.check_permission("any_table", &PermissionAction::Insert).await;
assert!(result.is_ok(), "admin bypass should return Ok for Insert");
let result = session.check_permission("any_table", &PermissionAction::Delete).await;
assert!(result.is_ok(), "admin bypass should return Ok for Delete");
}
#[cfg(all(feature = "permission", feature = "sqlite"))]
#[tokio::test]
async fn test_vuln_0001_non_admin_denied() {
let pool = DbPool::new("sqlite::memory:").await.expect("Failed to create pool");
let session = pool.get_session("system").await.expect("get_session");
let result = session.check_permission("any_table", &PermissionAction::Select).await;
assert!(result.is_err(), "non-admin should be denied");
}
#[cfg(all(feature = "permission", feature = "sqlite"))]
#[tokio::test]
async fn test_check_permission_non_admin_allowed() {
use std::io::Write;
let yaml_content = r#"
roles:
admin:
tables:
- name: "*"
operations: ["select", "insert", "update", "delete"]
reader:
tables:
- name: "test_tbl"
operations: ["select"]
"#;
let tmp_dir = std::env::temp_dir();
let yaml_path = tmp_dir.join("test_non_admin_perm.yaml");
{
let mut file = std::fs::File::create(&yaml_path).expect("create temp file");
file.write_all(yaml_content.as_bytes()).expect("write temp file");
}
let config = crate::foundation::DbConfig {
url: "sqlite::memory:".to_string(),
permissions_path: Some(yaml_path.to_string_lossy().to_string()),
..Default::default()
};
let pool = DbPool::with_config(config).await.expect("should create pool");
let session = pool.get_session("reader").await.expect("get_session for reader");
let result = session.check_permission("test_tbl", &PermissionAction::Select).await;
assert!(
result.is_ok(),
"reader should have SELECT on test_tbl: {:?}",
result.err()
);
let _ = std::fs::remove_file(&yaml_path);
}
#[cfg(all(feature = "permission", feature = "sqlite"))]
#[tokio::test]
async fn test_vuln_0001_check_table_permission_admin_bypass() {
let pool = DbPool::new("sqlite::memory:").await.expect("Failed to create pool");
let session = pool.get_session("admin").await.expect("get_session");
let result = session.check_table_permission("users", "SELECT").await;
assert!(result.is_ok(), "admin should bypass check_table_permission");
let result = session.check_table_permission("users", "INSERT").await;
assert!(result.is_ok(), "admin should bypass check_table_permission for INSERT");
}
}
#[cfg(test)]
#[cfg(all(feature = "permission", feature = "sql-parser"))]
mod vuln_0003_tests {
use super::*;
#[tokio::test]
async fn test_vuln_0003_parser_correctly_handles_insert() {
let sql = "INSERT INTO users (name) VALUES ('from into values')";
let parser_result = extract_table_name_via_parser(sql).await;
assert_eq!(
parser_result.as_deref(),
Some("users"),
"SqlParser should correctly extract 'users' for INSERT"
);
}
#[tokio::test]
async fn test_vuln_0003_parser_correctly_handles_update() {
let sql = "UPDATE users SET name = 'from users' WHERE id = 1";
let parser_result = extract_table_name_via_parser(sql).await;
assert_eq!(
parser_result.as_deref(),
Some("users"),
"SqlParser should correctly extract 'users' for UPDATE"
);
}
#[tokio::test]
async fn test_vuln_0003_parser_correctly_handles_delete() {
let sql = "DELETE FROM users WHERE name = 'from deleted'";
let parser_result = extract_table_name_via_parser(sql).await;
assert_eq!(
parser_result.as_deref(),
Some("users"),
"SqlParser should correctly extract 'users' for DELETE"
);
}
#[tokio::test]
async fn test_vuln_0003_parser_handles_quoted_table_name() {
let sql = "SELECT * FROM \"users\" WHERE id = 1";
let parser_result = extract_table_name_via_parser(sql).await;
assert!(
parser_result.is_some(),
"SqlParser should extract table name for quoted identifier, got: {:?}",
parser_result
);
let table = parser_result.unwrap();
assert!(
table.contains("users"),
"extracted table name should contain 'users', got: {}",
table
);
}
}
#[cfg(all(test, feature = "ladybug"))]
mod vuln_0005_tests {
use super::*;
use crate::database::graph::{GraphExecResult, GraphValue};
async fn make_ladybug_pool() -> DbPool {
DbPool::new("ladybug::memory:")
.await
.expect("Failed to create Ladybug pool")
}
#[test]
fn test_validate_cypher_safety_rejects_too_long() {
let long_cypher = format!("MATCH (n) RETURN '{}'", "x".repeat(11_200));
assert!(
long_cypher.len() > 10_240,
"test cypher should exceed 10KB, got {} bytes",
long_cypher.len()
);
let result = validate_cypher_safety(&long_cypher);
assert!(
result.is_err(),
"Cypher exceeding 10KB should be rejected (got {} bytes)",
long_cypher.len()
);
match &result {
Err(DbError::Permission(msg)) => {
assert!(
msg.contains("maximum length") || msg.contains("exceeds"),
"error should mention length, got: {}",
msg
);
}
other => panic!("expected DbError::Permission, got {:?}", other),
}
}
#[test]
fn test_validate_cypher_safety_rejects_multi_statement() {
let cypher = "MATCH (n) RETURN n; MATCH (m) RETURN m";
let result = validate_cypher_safety(cypher);
assert!(result.is_err(), "multi-statement Cypher should be rejected");
match &result {
Err(DbError::Permission(msg)) => {
assert!(
msg.contains("multiple statements") || msg.contains("';'"),
"error should mention multiple statements, got: {}",
msg
);
}
other => panic!("expected DbError::Permission, got {:?}", other),
}
}
#[test]
fn test_validate_cypher_safety_rejects_line_comment() {
let cypher = "MATCH (n) // comment RETURN n";
let result = validate_cypher_safety(cypher);
assert!(result.is_err(), "Cypher with line comment '//' should be rejected");
match &result {
Err(DbError::Permission(msg)) => {
assert!(
msg.contains("line comment") || msg.contains("//"),
"error should mention line comment, got: {}",
msg
);
}
other => panic!("expected DbError::Permission, got {:?}", other),
}
}
#[test]
fn test_validate_cypher_safety_rejects_block_comment() {
let cypher = "MATCH (n) /* comment */ RETURN n";
let result = validate_cypher_safety(cypher);
assert!(result.is_err(), "Cypher with block comment '/* */' should be rejected");
match &result {
Err(DbError::Permission(msg)) => {
assert!(
msg.contains("block comment") || msg.contains("/*"),
"error should mention block comment, got: {}",
msg
);
}
other => panic!("expected DbError::Permission, got {:?}", other),
}
}
#[test]
fn test_validate_cypher_safety_rejects_apoc_call() {
let cypher = "CALL apoc.systemdb.admin('something')";
let result = validate_cypher_safety(cypher);
assert!(result.is_err(), "Cypher calling APOC procedure should be rejected");
match &result {
Err(DbError::Permission(msg)) => {
assert!(
msg.contains("dangerous procedure") || msg.contains("apoc"),
"error should mention dangerous procedure, got: {}",
msg
);
}
other => panic!("expected DbError::Permission, got {:?}", other),
}
}
#[test]
fn test_validate_cypher_safety_allows_normal_query() {
let cypher = "MATCH (n:User) RETURN n";
let result = validate_cypher_safety(cypher);
assert!(
result.is_ok(),
"normal Cypher query should pass safety check, got: {:?}",
result
);
}
#[test]
fn test_validate_cypher_safety_allows_trailing_semicolon() {
let cypher = "MATCH (n) RETURN n;";
let result = validate_cypher_safety(cypher);
assert!(
result.is_ok(),
"Cypher with trailing semicolon should pass safety check, got: {:?}",
result
);
}
#[tokio::test]
async fn test_execute_cypher_with_params_passes_params() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Person(name STRING, age INT64, PRIMARY KEY(name))",
HashMap::new(),
)
.await
.expect("create node table");
let mut params_alice = HashMap::new();
params_alice.insert("name".to_string(), serde_json::json!("Alice"));
params_alice.insert("age".to_string(), serde_json::json!(25));
session
.execute_cypher_with_params("CREATE (:Person {name: $name, age: $age})", params_alice)
.await
.expect("create Alice with params");
let mut params_bob = HashMap::new();
params_bob.insert("name".to_string(), serde_json::json!("Bob"));
params_bob.insert("age".to_string(), serde_json::json!(30));
session
.execute_cypher_with_params("CREATE (:Person {name: $name, age: $age})", params_bob)
.await
.expect("create Bob with params");
let mut params_query = HashMap::new();
params_query.insert("target_name".to_string(), serde_json::json!("Alice"));
let result = session
.execute_cypher_with_params(
"MATCH (p:Person) WHERE p.name = $target_name RETURN p.name AS name, p.age AS age",
params_query,
)
.await
.expect("match with params");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 1, "should return 1 person (Alice)");
let name_val = &q.rows[0].columns[0].1;
match name_val {
GraphValue::Scalar(serde_json::Value::String(s)) => {
assert_eq!(s, "Alice", "name should be Alice");
}
other => panic!("expected String Scalar for name, got {other:?}"),
}
let age_val = &q.rows[0].columns[1].1;
match age_val {
GraphValue::Scalar(serde_json::Value::Number(n)) => {
assert_eq!(n.as_i64(), Some(25), "age should be 25");
}
other => panic!("expected Number Scalar for age, got {other:?}"),
}
}
GraphExecResult::Write { .. } => panic!("expected Query variant, got Write"),
}
}
#[tokio::test]
async fn test_execute_cypher_with_params_in_transaction() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
session
.execute_cypher_with_params(
"CREATE NODE TABLE Account(id INT64, balance INT64, PRIMARY KEY(id))",
HashMap::new(),
)
.await
.expect("create node table");
session.begin_transaction().await.expect("begin transaction");
let mut params1 = HashMap::new();
params1.insert("id".to_string(), serde_json::json!(1));
params1.insert("balance".to_string(), serde_json::json!(100));
session
.execute_cypher_with_params("CREATE (:Account {id: $id, balance: $balance})", params1)
.await
.expect("create account 1 in txn");
let mut params2 = HashMap::new();
params2.insert("id".to_string(), serde_json::json!(2));
params2.insert("balance".to_string(), serde_json::json!(200));
session
.execute_cypher_with_params("CREATE (:Account {id: $id, balance: $balance})", params2)
.await
.expect("create account 2 in txn");
let result = session
.execute_cypher_with_params("MATCH (a:Account) RETURN a.id AS id ORDER BY a.id", HashMap::new())
.await
.expect("match in txn");
match result {
GraphExecResult::Query(q) => {
assert_eq!(q.rows.len(), 2, "should see 2 accounts in txn");
}
GraphExecResult::Write { .. } => panic!("expected Query variant"),
}
session.commit().await.expect("commit");
}
#[tokio::test]
async fn test_execute_cypher_with_params_rejects_injection() {
let pool = make_ladybug_pool().await;
let session = pool.get_session("admin").await.expect("get_session");
let result = session
.execute_cypher_with_params("MATCH (n) RETURN n; DELETE (n)", HashMap::new())
.await;
assert!(
result.is_err(),
"multi-statement Cypher should be rejected even in execute_cypher_with_params"
);
match result {
Err(DbError::Permission(msg)) => {
assert!(
msg.contains("multiple statements") || msg.contains("';'"),
"error should mention multiple statements, got: {}",
msg
);
}
other => panic!("expected DbError::Permission, got {:?}", other),
}
}
}
#[cfg(test)]
#[cfg(feature = "sqlite")]
mod session_basic_tests {
use super::*;
async fn make_test_session(role: &str) -> (super::super::DbPool, Session) {
let pool = super::super::DbPool::new("sqlite::memory:")
.await
.expect("Failed to create pool");
let session = pool.get_session(role).await.expect("get_session failed");
(pool, session)
}
#[tokio::test]
async fn test_session_role() {
let (_pool, session) = make_test_session("admin").await;
assert_eq!(session.role(), "admin");
}
#[tokio::test]
async fn test_session_is_in_transaction_initially_false() {
let (_pool, session) = make_test_session("admin").await;
assert!(!session.is_in_transaction().await);
}
#[tokio::test]
async fn test_session_should_use_master_initially_false() {
let (_pool, session) = make_test_session("admin").await;
assert!(!session.should_use_master().await);
}
#[tokio::test]
async fn test_session_mark_write_enables_master() {
let (_pool, session) = make_test_session("admin").await;
session.mark_write().await;
assert!(session.should_use_master().await);
}
#[tokio::test]
async fn test_session_connection_returns_ok() {
let (_pool, session) = make_test_session("admin").await;
assert!(session.connection().is_ok());
}
#[tokio::test]
async fn test_session_begin_and_commit_transaction() {
let (_pool, session) = make_test_session("admin").await;
session.begin_transaction().await.expect("begin_transaction");
assert!(session.is_in_transaction().await);
assert!(session.should_use_master().await);
session.commit().await.expect("commit");
assert!(!session.is_in_transaction().await);
}
#[tokio::test]
async fn test_session_begin_and_rollback_transaction() {
let (_pool, session) = make_test_session("admin").await;
session.begin_transaction().await.expect("begin_transaction");
assert!(session.is_in_transaction().await);
session.rollback().await.expect("rollback");
assert!(!session.is_in_transaction().await);
}
#[tokio::test]
async fn test_session_double_begin_returns_error() {
let (_pool, session) = make_test_session("admin").await;
session.begin_transaction().await.expect("first begin");
let result = session.begin_transaction().await;
assert!(result.is_err(), "double begin should return error");
session.commit().await.expect("commit");
}
#[tokio::test]
async fn test_session_rollback_without_transaction_returns_error() {
let (_pool, session) = make_test_session("admin").await;
let result = session.rollback().await;
assert!(result.is_err(), "rollback without transaction should error");
}
#[tokio::test]
async fn test_session_commit_without_transaction_returns_error() {
let (_pool, session) = make_test_session("admin").await;
let result = session.commit().await;
assert!(result.is_err(), "commit without transaction should error");
}
#[tokio::test]
async fn test_session_execute_raw_select() {
let (_pool, session) = make_test_session("admin").await;
use sea_orm::ConnectionTrait;
session
.connection()
.unwrap()
.execute_unprepared("CREATE TABLE sel_test (id INTEGER PRIMARY KEY)")
.await
.expect("create table");
let result = session.execute_raw("SELECT * FROM sel_test").await;
assert!(result.is_ok(), "SELECT from table should succeed: {:?}", result.err());
}
#[tokio::test]
async fn test_session_execute_raw_ddl_rejected() {
let (_pool, session) = make_test_session("admin").await;
let result = session.execute_raw("CREATE TABLE test (id INTEGER)").await;
assert!(result.is_err(), "DDL should be rejected");
}
#[tokio::test]
async fn test_session_execute_raw_create_table_and_insert() {
let (_pool, session) = make_test_session("admin").await;
use sea_orm::ConnectionTrait;
session
.connection()
.unwrap()
.execute_unprepared("CREATE TABLE test_tbl (id INTEGER PRIMARY KEY, name TEXT)")
.await
.expect("create table");
let result = session
.execute_raw("INSERT INTO test_tbl (id, name) VALUES (1, 'test')")
.await;
assert!(result.is_ok(), "INSERT should succeed for admin: {:?}", result.err());
}
#[tokio::test]
async fn test_session_database_session_trait_commit() {
let (_pool, session) = make_test_session("admin").await;
use super::super::DatabaseSession;
let result = DatabaseSession::commit(&session).await;
assert!(result.is_err(), "commit without transaction via trait should error");
}
#[tokio::test]
async fn test_session_database_session_trait_rollback() {
let (_pool, session) = make_test_session("admin").await;
use super::super::DatabaseSession;
let result = DatabaseSession::rollback(&session).await;
assert!(result.is_err(), "rollback without transaction via trait should error");
}
#[tokio::test]
async fn test_session_database_session_trait_execute() {
let (_pool, session) = make_test_session("admin").await;
use sea_orm::ConnectionTrait;
session
.connection()
.unwrap()
.execute_unprepared("CREATE TABLE trait_exec_test (id INTEGER PRIMARY KEY)")
.await
.expect("create table");
use super::super::DatabaseSession;
let result = DatabaseSession::execute(&session, "INSERT INTO trait_exec_test (id) VALUES (1)").await;
assert!(result.is_ok(), "execute via trait should succeed: {:?}", result.err());
}
#[tokio::test]
async fn test_extract_table_name_via_parser_invalid_table() {
let result = super::extract_table_name_via_parser("SELECT 1").await;
assert!(result.is_none(), "SELECT without FROM table should return None");
}
#[tokio::test]
async fn test_extract_table_name_via_parser_unsupported() {
let result = super::extract_table_name_via_parser("INVALID SQL GIBBERISH").await;
let _ = result;
}
#[tokio::test]
async fn test_extract_table_name_via_parser_parse_error() {
let result = super::extract_table_name_via_parser("/* comment */").await;
assert!(result.is_none());
}
#[test]
fn test_is_invalid_table_name_empty() {
assert!(super::is_invalid_table_name(""));
assert!(super::is_invalid_table_name(" "));
}
#[test]
fn test_is_invalid_table_name_empty_part() {
assert!(super::is_invalid_table_name("schema..table"));
assert!(super::is_invalid_table_name(".table"));
}
#[test]
fn test_is_invalid_table_name_valid() {
assert!(!super::is_invalid_table_name("users"));
assert!(!super::is_invalid_table_name("public.users"));
assert!(!super::is_invalid_table_name("\"quoted\".\"table\""));
}
#[tokio::test]
async fn test_create_migration_executor() {
let (_pool, session) = make_test_session("admin").await;
let result = session.create_migration_executor(crate::foundation::DatabaseType::Sqlite);
assert!(
result.is_ok(),
"create_migration_executor should succeed: {:?}",
result.err()
);
}
}