use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use crate::{InklogError, LogDbProvider, LogRecord};
#[cfg(all(
feature = "kit",
any(feature = "sqlite", feature = "postgres", feature = "mysql")
))]
use dbnexus::ConnectionPool;
#[cfg(all(
feature = "kit",
any(feature = "sqlite", feature = "postgres", feature = "mysql")
))]
pub struct DbNexusLogDbAdapter {
pool: Arc<dyn ConnectionPool + Send + Sync>,
table_name: String,
}
#[cfg(all(
feature = "kit",
any(feature = "sqlite", feature = "postgres", feature = "mysql")
))]
impl DbNexusLogDbAdapter {
pub fn new(pool: Arc<dyn ConnectionPool + Send + Sync>, table_name: &str) -> Self {
Self {
pool,
table_name: table_name.to_string(),
}
}
pub fn table_name(&self) -> &str {
&self.table_name
}
}
#[cfg(all(
feature = "kit",
any(feature = "sqlite", feature = "postgres", feature = "mysql")
))]
impl LogDbProvider for DbNexusLogDbAdapter {
fn execute_log<'a>(
&'a self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<(), InklogError>> + Send + 'a>> {
Box::pin(async move {
let session =
self.pool.get_session("admin").await.map_err(|e| {
InklogError::DatabaseError(format!("Failed to get session: {e}"))
})?;
if is_ddl_sql(sql) {
session
.execute_raw_ddl(sql)
.await
.map_err(|e| InklogError::DatabaseError(format!("Execute DDL failed: {e}")))?;
} else {
session
.execute_raw(sql)
.await
.map_err(|e| InklogError::DatabaseError(format!("Execute failed: {e}")))?;
}
Ok(())
})
}
fn batch_insert<'a>(
&'a self,
entries: Vec<LogRecord>,
) -> Pin<Box<dyn Future<Output = Result<(), InklogError>> + Send + 'a>> {
Box::pin(async move {
if entries.is_empty() {
return Ok(());
}
let session =
self.pool.get_session("admin").await.map_err(|e| {
InklogError::DatabaseError(format!("Failed to get session: {e}"))
})?;
let sqls: Vec<String> = entries
.iter()
.map(|record| build_insert_sql(record, &self.table_name))
.collect();
let sql_refs: Vec<&str> = sqls.iter().map(|s| s.as_str()).collect();
session
.batch_execute_in_transaction(sql_refs)
.await
.map_err(|e| InklogError::DatabaseError(format!("Batch insert failed: {e}")))?;
Ok(())
})
}
}
fn is_ddl_sql(sql: &str) -> bool {
let sql_upper = sql.trim().to_uppercase();
const DDL_PREFIXES: &[&str] = &[
"CREATE TABLE",
"DROP TABLE",
"ALTER TABLE",
"TRUNCATE TABLE",
"CREATE INDEX",
"DROP INDEX",
"CREATE VIEW",
"DROP VIEW",
];
DDL_PREFIXES
.iter()
.any(|prefix| sql_upper.starts_with(prefix))
}
#[cfg(all(
feature = "kit",
any(feature = "sqlite", feature = "postgres", feature = "mysql")
))]
fn build_insert_sql(record: &LogRecord, table_name: &str) -> String {
let timestamp = record.timestamp.to_rfc3339();
let level = &record.level;
let target = &record.target;
let message = record.message.replace('\'', "''");
let fields_json = serde_json::to_string(&record.fields).unwrap_or_else(|_| "{}".to_string());
let fields_escaped = fields_json.replace('\'', "''");
let file = record
.file
.as_ref()
.map(|f| format!("'{}'", f.replace('\'', "''")))
.unwrap_or_else(|| "NULL".to_string());
let line = record
.line
.map(|l| l.to_string())
.unwrap_or_else(|| "NULL".to_string());
let thread_id = &record.thread_id;
format!(
"INSERT INTO {} (timestamp, level, target, message, fields, file, line, thread_id) \
VALUES ('{}', '{}', '{}', '{}', '{}', {}, {}, '{}')",
table_name,
timestamp,
level,
target.replace('\'', "''"),
message,
fields_escaped,
file,
line,
thread_id.replace('\'', "''")
)
}
#[cfg(all(
test,
feature = "kit",
any(feature = "sqlite", feature = "postgres", feature = "mysql")
))]
mod tests {
use super::*;
use dbnexus::{DbConfig, DbPoolBuilder};
async fn create_sqlite_pool() -> Arc<dyn ConnectionPool + Send + Sync> {
let config = DbConfig {
url: "sqlite::memory:".to_string(),
max_connections: 5,
min_connections: 1,
..Default::default()
};
let pool = DbPoolBuilder::new()
.config(config)
.build()
.await
.expect("Failed to create sqlite pool");
Arc::new(pool) as Arc<dyn ConnectionPool + Send + Sync>
}
const CREATE_TABLE_SQL: &str = "CREATE TABLE IF NOT EXISTS logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
timestamp TEXT NOT NULL,
level TEXT NOT NULL,
target TEXT NOT NULL,
message TEXT NOT NULL,
fields TEXT,
file TEXT,
line INTEGER,
thread_id TEXT NOT NULL
)";
#[tokio::test]
async fn adapter_new_returns_instance() {
let pool = create_sqlite_pool().await;
let adapter = DbNexusLogDbAdapter::new(pool, "logs");
assert_eq!(adapter.table_name(), "logs");
}
#[tokio::test]
async fn adapter_execute_log_runs_ddl() {
let pool = create_sqlite_pool().await;
let adapter = DbNexusLogDbAdapter::new(pool, "logs");
adapter
.execute_log(CREATE_TABLE_SQL)
.await
.expect("execute_log should succeed for DDL");
}
#[tokio::test]
async fn adapter_batch_insert_inserts_records() {
let pool = create_sqlite_pool().await;
let adapter = DbNexusLogDbAdapter::new(pool, "logs");
adapter
.execute_log(CREATE_TABLE_SQL)
.await
.expect("create table");
let records = vec![
LogRecord::new(
tracing::Level::INFO,
"module_a".to_string(),
"message_a".to_string(),
),
LogRecord::new(
tracing::Level::WARN,
"module_b".to_string(),
"message_b".to_string(),
),
];
adapter
.batch_insert(records)
.await
.expect("batch_insert should succeed");
let session = adapter
.pool
.get_session("admin")
.await
.expect("get session for verification");
let _result = session
.execute_raw("SELECT COUNT(*) FROM logs")
.await
.expect("count query should succeed");
}
#[tokio::test]
async fn adapter_batch_insert_empty_succeeds() {
let pool = create_sqlite_pool().await;
let adapter = DbNexusLogDbAdapter::new(pool, "logs");
adapter
.batch_insert(Vec::new())
.await
.expect("empty batch should succeed");
}
#[tokio::test]
async fn adapter_satisfies_log_db_provider() {
let pool = create_sqlite_pool().await;
let adapter = DbNexusLogDbAdapter::new(pool, "logs");
fn assert_impl<T: LogDbProvider>(_: &T) {}
assert_impl(&adapter);
}
#[tokio::test]
async fn adapter_dyn_dispatch_works() {
let pool = create_sqlite_pool().await;
let adapter: Arc<dyn LogDbProvider + Send + Sync> =
Arc::new(DbNexusLogDbAdapter::new(pool, "logs"));
adapter
.execute_log(CREATE_TABLE_SQL)
.await
.expect("execute_log via dyn should succeed");
adapter
.batch_insert(vec![LogRecord::new(
tracing::Level::INFO,
"dyn_test".to_string(),
"via dyn".to_string(),
)])
.await
.expect("batch_insert via dyn should succeed");
}
#[test]
fn build_insert_sql_escapes_single_quotes() {
let mut record = LogRecord::new(
tracing::Level::INFO,
"module".to_string(),
"it's a test".to_string(),
);
record.thread_id = "thread'1".to_string();
let sql = build_insert_sql(&record, "logs");
assert!(
sql.contains("it''s a test"),
"message single quotes should be escaped, got: {sql}"
);
assert!(
sql.contains("thread''1"),
"thread_id single quotes should be escaped, got: {sql}"
);
}
}