use std::time::Duration;
use sea_orm::{ConnectOptions, Database, DatabaseConnection};
use sea_orm_migration::MigratorTrait;
use crate::migrator::Migrator;
use crate::types::DatabaseError;
const POOL_MAX_CONNECTIONS: u32 = 10;
const POOL_MIN_CONNECTIONS: u32 = 0;
const POOL_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(30);
const POOL_IDLE_TIMEOUT: Duration = Duration::from_secs(600);
#[cfg(feature = "sqlite")]
const SQLITE_BUSY_TIMEOUT: Duration = Duration::from_secs(120);
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
struct SqliteUrlKind {
in_memory: bool,
shared_cache: bool,
read_only: bool,
}
fn classify_sqlite_url(url: &str) -> SqliteUrlKind {
let rest = url
.strip_prefix("sqlite://")
.or_else(|| url.strip_prefix("sqlite:"))
.unwrap_or(url);
let (path, query) = rest.split_once('?').unwrap_or((rest, ""));
let path = path.strip_prefix("file:").unwrap_or(path);
let mut kind = SqliteUrlKind {
in_memory: path == ":memory:",
..SqliteUrlKind::default()
};
for param in query.split('&') {
match param {
"mode=memory" => kind.in_memory = true,
"cache=shared" => kind.shared_cache = true,
"mode=ro" | "immutable=1" | "immutable=true" => kind.read_only = true,
_ => {}
}
}
kind
}
pub fn sqlite_url_is_in_memory(url: &str) -> bool {
url.starts_with("sqlite") && classify_sqlite_url(url).in_memory
}
pub async fn connect(url: &str) -> Result<DatabaseConnection, DatabaseError> {
#[cfg(feature = "sqlite")]
if url.starts_with("sqlite") {
return connect_sqlite(url).await;
}
let mut opt = ConnectOptions::new(url.to_owned());
opt.max_connections(POOL_MAX_CONNECTIONS)
.min_connections(POOL_MIN_CONNECTIONS)
.acquire_timeout(POOL_ACQUIRE_TIMEOUT)
.idle_timeout(POOL_IDLE_TIMEOUT);
Database::connect(opt)
.await
.map_err(|e| DatabaseError::ConnectionError(e.to_string()))
}
#[cfg(feature = "sqlite")]
async fn connect_sqlite(url: &str) -> Result<DatabaseConnection, DatabaseError> {
use std::str::FromStr;
use sea_orm::SqlxSqliteConnector;
use sea_orm::sqlx::ConnectOptions as _;
use sea_orm::sqlx::sqlite::{
SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteSynchronous,
};
let kind = classify_sqlite_url(url);
let base_opts = SqliteConnectOptions::from_str(url)
.map_err(|e| DatabaseError::ConnectionError(e.to_string()))?
.log_statements(log::LevelFilter::Info)
.busy_timeout(SQLITE_BUSY_TIMEOUT);
let mut want_wal = false;
let mut conn_opts = base_opts.clone();
if !kind.in_memory {
let probe_path = conn_opts.get_filename().to_path_buf();
let writable = tokio::task::spawn_blocking(move || sqlite_path_is_writable(&probe_path))
.await
.map_err(|e| DatabaseError::ConnectionError(e.to_string()))?;
if kind.read_only {
conn_opts = conn_opts.read_only(true);
} else if !writable {
tracing::warn!(
path = %conn_opts.get_filename().display(),
"SQLite database is not writable; opening read-only. Writes will fail. \
Check file and parent-directory permissions (and, on Windows, other \
processes holding the file) if this database is meant to be written."
);
conn_opts = conn_opts.read_only(true);
} else {
want_wal = true;
conn_opts = conn_opts
.journal_mode(SqliteJournalMode::Wal)
.synchronous(SqliteSynchronous::Full);
}
}
let mut pool_opts = SqlitePoolOptions::new()
.max_connections(POOL_MAX_CONNECTIONS)
.min_connections(POOL_MIN_CONNECTIONS)
.acquire_timeout(POOL_ACQUIRE_TIMEOUT)
.idle_timeout(POOL_IDLE_TIMEOUT);
if kind.in_memory {
pool_opts = pool_opts
.min_connections(1)
.idle_timeout(None)
.max_lifetime(None);
if !kind.shared_cache {
pool_opts = pool_opts.max_connections(1);
}
}
let sqlx_pool = match pool_opts.clone().connect_with(conn_opts).await {
Ok(pool) => pool,
Err(original) if want_wal => {
let fallback_opts = base_opts
.journal_mode(SqliteJournalMode::Delete)
.synchronous(SqliteSynchronous::Full);
match pool_opts.connect_with(fallback_opts).await {
Ok(pool) => {
tracing::warn!(
error = %original,
"Enabling SQLite WAL failed; opened with a rollback journal instead \
(this happens on a network/FUSE filesystem without shared-memory \
support). Reader/writer concurrency is reduced for this database."
);
pool
}
Err(_) => return Err(DatabaseError::ConnectionError(original.to_string())),
}
}
Err(e) => return Err(DatabaseError::ConnectionError(e.to_string())),
};
Ok(SqlxSqliteConnector::from_sqlx_sqlite_pool(sqlx_pool))
}
#[cfg(feature = "sqlite")]
fn sqlite_path_is_writable(path: &std::path::Path) -> bool {
let file_writable = if path.exists() {
std::fs::OpenOptions::new().write(true).open(path).is_ok()
} else {
true
};
file_writable && sqlite_parent_dir_is_writable(path)
}
#[cfg(feature = "sqlite")]
fn sqlite_parent_dir_is_writable(path: &std::path::Path) -> bool {
let parent = match path.parent() {
Some(p) if !p.as_os_str().is_empty() => p,
_ => std::path::Path::new("."),
};
let probe = parent.join(format!(".cognee-wal-probe-{}.tmp", std::process::id()));
match std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.open(&probe)
{
Ok(_) => {
let _ = std::fs::remove_file(&probe);
true
}
Err(e) if e.kind() == std::io::ErrorKind::AlreadyExists => true,
Err(_) => false,
}
}
pub async fn initialize(db: &DatabaseConnection) -> Result<(), DatabaseError> {
Migrator::up(db, None)
.await
.map_err(|e| DatabaseError::QueryError(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::{SqliteUrlKind, classify_sqlite_url, sqlite_url_is_in_memory};
#[test]
fn detects_in_memory_spellings() {
for url in [
"sqlite::memory:",
"sqlite://:memory:",
"sqlite:file:pinned?mode=memory",
"sqlite::memory:?cache=shared",
] {
assert!(classify_sqlite_url(url).in_memory, "{url}");
assert!(sqlite_url_is_in_memory(url), "{url}");
}
}
#[test]
fn detects_shared_cache_only_when_explicit() {
assert!(classify_sqlite_url("sqlite::memory:?cache=shared").shared_cache);
assert!(classify_sqlite_url("sqlite:file:x?mode=memory&cache=shared").shared_cache);
assert!(!classify_sqlite_url("sqlite::memory:").shared_cache);
assert!(!classify_sqlite_url("sqlite:file:x?cache=private").shared_cache);
}
#[test]
fn detects_read_only_opens() {
assert!(classify_sqlite_url("sqlite://./a.db?mode=ro").read_only);
assert!(classify_sqlite_url("sqlite:a.db?immutable=1").read_only);
assert!(classify_sqlite_url("sqlite:a.db?immutable=true").read_only);
assert!(!classify_sqlite_url("sqlite://./a.db?mode=rwc").read_only);
assert!(!classify_sqlite_url("sqlite://./a.db?mode=rw").read_only);
}
#[test]
fn file_paths_are_never_misclassified_by_substring() {
let kind = classify_sqlite_url("sqlite:///tmp/mode=memory/app.db?mode=rwc");
assert_eq!(kind, SqliteUrlKind::default());
assert!(!sqlite_url_is_in_memory(
"sqlite:///tmp/mode=memory/app.db?mode=rwc"
));
}
#[test]
fn non_sqlite_urls_are_not_in_memory() {
assert!(!sqlite_url_is_in_memory("postgres://user:pw@localhost/db"));
}
}