use crate::db::pg::mutability;
use crate::db::pg::translate;
use std::collections::{BTreeMap, VecDeque};
use std::ops::{Deref, DerefMut};
use std::sync::{Arc, Condvar, Mutex};
pub use crate::db::value::{DataValue, NamedRows};
pub trait DbBackend: Send + Sync {
fn run_script(
&self,
query: &str,
params: BTreeMap<String, serde_json::Value>,
) -> Result<NamedRows, Box<dyn std::error::Error>>;
fn import_relations(
&self,
data: BTreeMap<String, NamedRows>,
) -> Result<(), Box<dyn std::error::Error>>;
fn redacted_url(&self) -> String;
fn mutability_for(&self, query: &str) -> mutability::ScriptMutability;
}
pub type SharedDb = Arc<dyn DbBackend>;
#[derive(Clone)]
pub struct PostgresBackend {
pub pg_url: String,
pub pool: Arc<ClientPool>,
pub ro_pool: Arc<ClientPool>,
pub read_only: bool,
}
pub struct PooledClient {
client: Option<postgres::Client>,
pool: Arc<ClientPool>,
}
impl PooledClient {
fn new(client: postgres::Client, pool: Arc<ClientPool>) -> Self {
Self {
client: Some(client),
pool,
}
}
}
impl Deref for PooledClient {
type Target = postgres::Client;
fn deref(&self) -> &postgres::Client {
self.client.as_ref().unwrap()
}
}
impl DerefMut for PooledClient {
fn deref_mut(&mut self) -> &mut postgres::Client {
self.client.as_mut().unwrap()
}
}
impl Drop for PooledClient {
fn drop(&mut self) {
if let Some(c) = self.client.take() {
self.pool.release(c);
}
}
}
#[derive(Clone)]
pub struct ClientPool {
inner: Arc<ClientPoolState>,
}
struct ClientPoolState {
max: usize,
has_slot: Condvar,
state: Mutex<PoolState>,
}
#[derive(Default)]
struct PoolState {
idle: VecDeque<postgres::Client>,
live: usize,
}
impl Drop for ClientPoolState {
fn drop(&mut self) {
let state = std::mem::take(&mut self.state);
let mut inner = state.into_inner().unwrap_or_else(|e| e.into_inner());
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(move || inner.idle.clear());
} else {
inner.idle.clear();
}
}
}
impl ClientPool {
pub fn new(max: usize) -> Self {
Self {
inner: Arc::new(ClientPoolState {
max: max.max(1),
has_slot: Condvar::new(),
state: Mutex::new(PoolState {
idle: VecDeque::new(),
live: 0,
}),
}),
}
}
pub fn size_from_env() -> usize {
std::env::var("LEANKG_PG_POOL_SIZE")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|v| *v >= 1)
.unwrap_or(5)
}
pub fn live_count(&self) -> usize {
self.inner.state.lock().unwrap().live
}
pub fn checkout(&self, connect_url: &str) -> Result<PooledClient, Box<dyn std::error::Error>> {
let mut guard = self.inner.state.lock().unwrap();
let pool_arc = Arc::new(self.clone());
loop {
if let Some(c) = guard.idle.pop_front() {
return Ok(PooledClient::new(c, pool_arc.clone()));
}
if guard.live < self.inner.max {
let client = postgres::Client::connect(connect_url, postgres::NoTls)?;
guard.live += 1;
return Ok(PooledClient::new(client, pool_arc.clone()));
}
guard = self
.inner
.has_slot
.wait(guard)
.unwrap_or_else(|e| e.into_inner());
}
}
fn release(&self, client: postgres::Client) {
let mut guard = self.inner.state.lock().unwrap();
guard.idle.push_back(client);
self.inner.has_slot.notify_one();
}
}
impl std::fmt::Debug for PostgresBackend {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PostgresBackend")
.field("pg_url", &redact_url(&self.pg_url))
.field("pool", &"<lazy>")
.field("read_only", &self.read_only)
.finish()
}
}
impl PostgresBackend {
pub fn from_env() -> Result<Self, String> {
let url = std::env::var("LEANKG_PG_URL")
.map_err(|_| "LEANKG_PG_URL is not set; the Postgres backend requires it (run `docker compose up postgres`)")?;
if !url.starts_with("postgres://") && !url.starts_with("postgresql://") {
return Err(format!(
"LEANKG_PG_URL must be a postgres:// URL, got: {}",
redact_url(&url)
));
}
Ok(Self {
pg_url: url,
pool: Arc::new(ClientPool::new(ClientPool::size_from_env())),
ro_pool: Arc::new(ClientPool::new(ClientPool::size_from_env())),
read_only: false,
})
}
pub fn from_env_read_only() -> Result<Self, String> {
Ok(Self::from_env()?.with_read_only())
}
pub fn with_read_only(mut self) -> Self {
self.read_only = true;
self
}
pub fn redacted_url(&self) -> String {
redact_url(&self.pg_url)
}
pub fn read_only_url(&self) -> String {
if !self.read_only {
return self.pg_url.clone();
}
let base = &self.pg_url;
if base.contains("default_transaction_read_only") {
return base.clone();
}
const RO_FLAG: &str = "-cdefault_transaction_read_only%3Don";
if let Some(pos) = base.find("options=") {
let after = &base[pos + "options=".len()..];
let end = after.find('&').unwrap_or(after.len());
let value = &after[..end];
let rest = &after[end..]; return format!(
"{}{}%20{}{}",
&base[..pos + "options=".len()],
value,
RO_FLAG,
rest
);
}
let (before, after) = base.split_once('?').unwrap_or((base, ""));
let sep = if after.is_empty() { "?" } else { "&" };
format!("{before}{sep}{after}options={RO_FLAG}")
}
fn checkout(&self) -> Result<PooledClient, Box<dyn std::error::Error>> {
let url = self.pg_url.clone();
let pool = self.pool.clone();
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(move || pool.checkout(&url))
} else {
pool.checkout(&url)
}
}
fn checkout_read_only(&self) -> Result<PooledClient, Box<dyn std::error::Error>> {
let url = self.read_only_url();
let pool = self.ro_pool.clone();
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(move || pool.checkout(&url))
} else {
pool.checkout(&url)
}
}
pub fn advisory_lock(&self, key: i64) -> Result<AdvisoryLock, Box<dyn std::error::Error>> {
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(|| self.advisory_lock_sync(key))
} else {
self.advisory_lock_sync(key)
}
}
fn advisory_lock_sync(&self, key: i64) -> Result<AdvisoryLock, Box<dyn std::error::Error>> {
let mut client = self.checkout()?;
client.execute("SELECT pg_advisory_lock($1)", &[&key])?;
Ok(AdvisoryLock {
client: Some(client),
key,
})
}
pub const INDEX_LOCK_KEY: i64 = 0x6C65616E6B67;
pub fn try_advisory_lock(
&self,
key: i64,
) -> Result<Option<AdvisoryLock>, Box<dyn std::error::Error>> {
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(|| self.try_advisory_lock_sync(key))
} else {
self.try_advisory_lock_sync(key)
}
}
fn try_advisory_lock_sync(
&self,
key: i64,
) -> Result<Option<AdvisoryLock>, Box<dyn std::error::Error>> {
let mut client = self.checkout()?;
let ok: bool = client
.query_one("SELECT pg_try_advisory_lock($1)", &[&key])?
.get(0);
if ok {
Ok(Some(AdvisoryLock {
client: Some(client),
key,
}))
} else {
let c = client.client.take().unwrap();
client.pool.release(c);
Ok(None)
}
}
pub fn run_script(
&self,
query: &str,
params: BTreeMap<String, serde_json::Value>,
) -> Result<NamedRows, Box<dyn std::error::Error>> {
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(|| self.run_script_sync(query, params))
} else {
self.run_script_sync(query, params)
}
}
pub fn mutability_for(&self, query: &str) -> mutability::ScriptMutability {
mutability::mutability_for(query)
}
pub fn import_relations(
&self,
data: BTreeMap<String, NamedRows>,
) -> Result<(), Box<dyn std::error::Error>> {
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(|| self.import_relations_sync(data))
} else {
self.import_relations_sync(data)
}
}
fn run_script_sync(
&self,
query: &str,
params: BTreeMap<String, serde_json::Value>,
) -> Result<NamedRows, Box<dyn std::error::Error>> {
if query.contains("query_cache") {
if query.trim_start().starts_with("?[") && !query.contains(":put") {
let head: Vec<String> = query
.split_once("?[")
.and_then(|(_, rest)| rest.split_once(']'))
.map(|(inner, _)| {
inner
.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
})
.unwrap_or_default();
return Ok(NamedRows::new(head, Vec::new()));
}
return Ok(NamedRows::new(Vec::new(), Vec::new()));
}
let mut client = if self.read_only {
self.checkout_read_only()?
} else {
self.checkout()?
};
let t = translate::translate(query, params).map_err(|e| -> Box<dyn std::error::Error> {
Box::new(std::io::Error::other(format!(
"translate({}): {e}",
&query[..query.len().min(60)]
)))
})?;
let mut head = t.head.clone();
let mut rows: Vec<Vec<DataValue>> = Vec::new();
match t.kind {
translate::TranslationKind::Read => {
let param_refs: Vec<&(dyn postgres::types::ToSql + Sync)> = t
.params
.iter()
.map(|b| b.as_ref() as &(dyn postgres::types::ToSql + Sync))
.collect();
if t.gucs.is_empty() {
let result = client.query(&t.sql, ¶m_refs)?;
for row in &result {
let mapped = translate::map_row(row, &t.head)?;
rows.push(mapped);
}
} else {
let mut tx = client.transaction()?;
apply_gucs(&mut tx, &t.gucs)?;
let result = tx.query(&t.sql, ¶m_refs)?;
for row in &result {
let mapped = translate::map_row(row, &t.head)?;
rows.push(mapped);
}
tx.commit()?;
}
if head.is_empty() && !rows.is_empty() {
head = (0..rows[0].len()).map(|i| format!("col{i}")).collect();
}
}
translate::TranslationKind::Write => {
let param_refs: Vec<&(dyn postgres::types::ToSql + Sync)> = t
.params
.iter()
.map(|b| b.as_ref() as &(dyn postgres::types::ToSql + Sync))
.collect();
let mut tx = client.transaction()?;
apply_gucs(&mut tx, &t.gucs)?;
tx.execute(&t.sql, ¶m_refs)?;
tx.commit()?;
}
translate::TranslationKind::DdlNoop => {
}
}
Ok(NamedRows::new(head, rows))
}
fn import_relations_sync(
&self,
data: BTreeMap<String, NamedRows>,
) -> Result<(), Box<dyn std::error::Error>> {
let mut client = if self.read_only {
self.checkout_read_only()?
} else {
self.checkout()?
};
let mut tx = client.transaction()?;
for (table, named) in data {
let cols = named.headers.clone();
let cols: Vec<String> = cols
.into_iter()
.map(|c| {
if table == "embedding_vectors" && c == "vector" {
"vec".to_string()
} else {
c
}
})
.collect();
let pk_col = match table.as_str() {
"embedding_state" | "embedding_vectors" => Some("qualified_name"),
"index_inventory" => Some("key"),
"index_hashes" => Some("path"),
"migrations" => Some("id"),
_ => None,
};
match pk_col {
Some(pk) if bulk_copy_enabled() => {
self.copy_upsert(&mut tx, &table, &cols, pk, &named)?;
}
_ => {
self.insert_rows(&mut tx, &table, &cols, pk_col, &named)?;
}
}
}
tx.commit()?;
Ok(())
}
fn copy_upsert(
&self,
tx: &mut postgres::Transaction,
table: &str,
cols: &[String],
pk: &str,
named: &NamedRows,
) -> Result<(), Box<dyn std::error::Error>> {
use std::io::Write;
if named.rows.is_empty() {
return Ok(());
}
let pk_idx = cols.iter().position(|c| c.as_str() == pk);
let rows: Vec<&Vec<DataValue>> = if let Some(idx) = pk_idx {
let mut seen: std::collections::HashMap<String, &Vec<DataValue>> =
std::collections::HashMap::with_capacity(named.rows.len());
for row in &named.rows {
seen.insert(row.get(idx).map(|v| v.to_string()).unwrap_or_default(), row);
}
seen.into_values().collect()
} else {
named.rows.iter().collect()
};
let q_table = crate::db::pg::translate::quote_ident(table);
let q_cols = cols
.iter()
.map(|c| crate::db::pg::translate::quote_ident(c))
.collect::<Vec<_>>()
.join(", ");
let staging = crate::db::pg::translate::quote_ident(&format!("{table}_staging"));
tx.batch_execute(&format!(
"CREATE TEMP TABLE {staging} (LIKE {q_table}) ON COMMIT DROP"
))?;
let copy_sql = format!("COPY {staging} ({q_cols}) FROM STDIN");
let mut writer = tx.copy_in(©_sql)?;
for row in rows {
let mut line = String::new();
for (i, val) in row.iter().enumerate() {
if i > 0 {
line.push('\t');
}
push_copy_text(&mut line, &data_to_copy_text(val, &cols[i]));
}
line.push('\n');
writer.write_all(line.as_bytes())?;
}
writer.finish()?;
let update_set = cols
.iter()
.filter(|c| c.as_str() != pk)
.map(|c| {
format!(
"{} = EXCLUDED.{}",
crate::db::pg::translate::quote_ident(c),
crate::db::pg::translate::quote_ident(c)
)
})
.collect::<Vec<_>>()
.join(", ");
let q_pk = crate::db::pg::translate::quote_ident(pk);
tx.execute(
&format!(
"INSERT INTO {q_table} ({q_cols}) SELECT {q_cols} FROM {staging} \
ON CONFLICT ({q_pk}) DO UPDATE SET {update_set}"
),
&[],
)?;
Ok(())
}
fn insert_rows(
&self,
tx: &mut postgres::Transaction,
table: &str,
cols: &[String],
pk: Option<&str>,
named: &NamedRows,
) -> Result<(), Box<dyn std::error::Error>> {
let col_sql = cols
.iter()
.map(|c| crate::db::pg::translate::quote_ident(c))
.collect::<Vec<_>>()
.join(", ");
for row in &named.rows {
let mut values: Vec<Box<dyn postgres::types::ToSql + Sync + Send>> = Vec::new();
for (i, val) in row.iter().enumerate() {
values.push(cozo_to_pg(val, &cols[i]));
}
let value_refs: Vec<&(dyn postgres::types::ToSql + Sync)> = values
.iter()
.map(|b| b.as_ref() as &(dyn postgres::types::ToSql + Sync))
.collect();
let sql = if let Some(pk) = pk {
let update_set = cols
.iter()
.filter(|c| c.as_str() != pk)
.map(|c| {
format!(
"{} = EXCLUDED.{}",
crate::db::pg::translate::quote_ident(c),
crate::db::pg::translate::quote_ident(c)
)
})
.collect::<Vec<_>>()
.join(", ");
format!(
"INSERT INTO {table} ({col_sql}) VALUES ({vals}) ON CONFLICT ({pk}) DO UPDATE SET {update_set}",
vals = (1..=values.len()).map(|i| format!("${i}")).collect::<Vec<_>>().join(", "),
pk = crate::db::pg::translate::quote_ident(pk),
)
} else {
format!(
"INSERT INTO {table} ({col_sql}) VALUES ({vals})",
vals = (1..=values.len())
.map(|i| format!("${i}"))
.collect::<Vec<_>>()
.join(", ")
)
};
tx.execute(&sql, &value_refs)?;
}
Ok(())
}
}
impl DbBackend for PostgresBackend {
fn run_script(
&self,
query: &str,
params: BTreeMap<String, serde_json::Value>,
) -> Result<NamedRows, Box<dyn std::error::Error>> {
PostgresBackend::run_script(self, query, params)
}
fn import_relations(
&self,
data: BTreeMap<String, NamedRows>,
) -> Result<(), Box<dyn std::error::Error>> {
PostgresBackend::import_relations(self, data)
}
fn redacted_url(&self) -> String {
PostgresBackend::redacted_url(self)
}
fn mutability_for(&self, query: &str) -> mutability::ScriptMutability {
PostgresBackend::mutability_for(self, query)
}
}
pub struct AdvisoryLock {
client: Option<PooledClient>,
key: i64,
}
impl Drop for AdvisoryLock {
fn drop(&mut self) {
if let Some(mut client) = self.client.take() {
let key = self.key;
if tokio::runtime::Handle::try_current().is_ok() {
let _ = tokio::task::block_in_place(move || {
client.execute("SELECT pg_advisory_unlock($1)", &[&key])
});
} else {
let _ = client.execute("SELECT pg_advisory_unlock($1)", &[&key]);
}
}
}
}
fn apply_gucs(
tx: &mut postgres::Transaction,
gucs: &[(String, String)],
) -> Result<(), postgres::Error> {
for (name, value) in gucs {
let escaped_value = value.replace('\'', "''");
let sql = format!("SET LOCAL {name} = '{escaped_value}'");
tx.batch_execute(&sql)?;
}
Ok(())
}
fn bulk_copy_enabled() -> bool {
std::env::var("LEANKG_EMBED_COPY")
.map(|v| !matches!(v.as_str(), "0" | "false" | "off"))
.unwrap_or(true)
}
fn bulk_reindex_enabled(total_rows: usize) -> bool {
if std::env::var("LEANKG_EMBED_COPY")
.map(|v| matches!(v.as_str(), "1" | "true" | "on"))
.unwrap_or(false)
{
return true;
}
let threshold = std::env::var("LEANKG_EMBED_BULK_REINDEX_THRESHOLD")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(100_000);
total_rows >= threshold
}
fn data_to_copy_text(v: &DataValue, col: &str) -> String {
match v {
DataValue::Null => String::new(), DataValue::Bool(b) => b.to_string(),
DataValue::Num(crate::db::value::Num::Int(i)) => i.to_string(),
DataValue::Num(crate::db::value::Num::Float(f)) => f.to_string(),
DataValue::Str(s) => s.as_str().to_string(),
DataValue::Json(j) => j.clone(),
DataValue::List(items) if col == "vec" || col == "vector" => {
let mut s = String::from("[");
for (i, item) in items.iter().enumerate() {
if i > 0 {
s.push(',');
}
match item {
DataValue::Num(crate::db::value::Num::Float(f)) => s.push_str(&format!("{f}")),
DataValue::Num(crate::db::value::Num::Int(i)) => s.push_str(&format!("{i}")),
other => s.push_str(&format!("{other}")),
}
}
s.push(']');
s
}
DataValue::Bytes(b) => {
let mut s = String::with_capacity(b.len() * 2);
for byte in b {
s.push_str(&format!("\\{:03o}", byte));
}
s
}
other => format!("{other}"),
}
}
fn push_copy_text(out: &mut String, s: &str) {
for ch in s.chars() {
match ch {
'\t' => out.push_str("\\t"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
'\\' => out.push_str("\\\\"),
_ => out.push(ch),
}
}
}
fn cozo_to_pg(v: &DataValue, col: &str) -> Box<dyn postgres::types::ToSql + Sync + Send> {
match v {
DataValue::Null => Box::new(Option::<String>::None),
DataValue::Bool(b) => Box::new(*b),
DataValue::Num(crate::db::value::Num::Int(i)) => Box::new(*i),
DataValue::Num(crate::db::value::Num::Float(f)) => Box::new(*f),
DataValue::Str(s) => Box::new(s.clone()),
DataValue::Json(j) => Box::new(j.clone()),
DataValue::List(items) if col == "vec" || col == "vector" => {
let mut s = String::from("[");
for (i, item) in items.iter().enumerate() {
if i > 0 {
s.push(',');
}
match item {
DataValue::Num(crate::db::value::Num::Float(f)) => s.push_str(&format!("{f}")),
DataValue::Num(crate::db::value::Num::Int(i)) => s.push_str(&format!("{i}")),
other => s.push_str(&format!("{other}")),
}
}
s.push(']');
Box::new(s)
}
DataValue::Bytes(b) => Box::new(b.clone()),
other => Box::new(format!("{other}")),
}
}
fn redact_url(url: &str) -> String {
let mut out = String::with_capacity(url.len());
let mut in_userinfo = false;
let mut seen_scheme = false;
let mut seen_at = false;
for ch in url.chars() {
match ch {
'@' => {
seen_at = true;
in_userinfo = false;
out.push('@');
}
':' if !in_userinfo && seen_scheme && !seen_at => {
in_userinfo = true;
out.push(':');
out.push_str("****");
}
'/' | '?' | '#' => {
in_userinfo = false;
out.push(ch);
}
_ if !in_userinfo => out.push(ch),
_ => {}
}
if ch == '/' {
seen_scheme = true;
}
}
out
}
pub fn init_db(_db_path: &std::path::Path) -> Result<SharedDb, Box<dyn std::error::Error>> {
#[cfg(test)]
{
return test_init_db(_db_path);
}
#[allow(unreachable_code)]
{
init_db_pg()
}
}
#[cfg(test)]
fn test_init_db(db_path: &std::path::Path) -> Result<SharedDb, Box<dyn std::error::Error>> {
Ok(Arc::new(crate::db::fake::FakeBackend::for_path(db_path)))
}
#[cfg(test)]
pub(crate) fn test_pg_url() -> String {
std::env::var("LEANKG_PG_URL")
.unwrap_or_else(|_| "postgresql://postgres:postgres@localhost:5433/leankg".to_string())
}
#[cfg(test)]
fn test_schema_url(schema: &str) -> Result<String, Box<dyn std::error::Error>> {
let base = test_pg_url();
let sep = if base.contains('?') { '&' } else { '?' };
Ok(format!(
"{base}{sep}options=-csearch_path%3D{schema}%2Cpublic"
))
}
#[cfg(test)]
fn test_scratch_schema(db_path: &std::path::Path) -> Result<String, Box<dyn std::error::Error>> {
use std::collections::HashMap;
use std::sync::Mutex as StdMutex;
use std::sync::OnceLock;
static MAP: OnceLock<StdMutex<HashMap<std::path::PathBuf, String>>> = OnceLock::new();
let map = MAP.get_or_init(|| StdMutex::new(HashMap::new()));
let mut guard = map.lock().unwrap_or_else(|e| e.into_inner());
let key = db_path.to_path_buf();
if let Some(schema) = guard.get(&key) {
return Ok(schema.clone());
}
let schema = create_scratch_schema()?;
guard.insert(key, schema.clone());
Ok(schema)
}
#[cfg(test)]
fn create_scratch_schema() -> Result<String, Box<dyn std::error::Error>> {
static COUNTER: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
let base = test_pg_url();
let name = format!(
"leankg_libtest_{}_{}",
std::process::id(),
COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
);
let mut client = postgres::Client::connect(&base, postgres::NoTls)?;
client.batch_execute(&format!("DROP SCHEMA IF EXISTS {name} CASCADE"))?;
client.batch_execute(&format!("CREATE SCHEMA {name}"))?;
client.batch_execute(&format!("SET search_path TO {name}, public"))?;
crate::db::pg::migrations::run_migrations(&mut client)?;
std::mem::forget(client);
Ok(name)
}
pub fn init_db_readonly(
_db_path: &std::path::Path,
) -> Result<SharedDb, Box<dyn std::error::Error>> {
#[cfg(test)]
{
return test_init_db(_db_path);
}
#[allow(unreachable_code)]
{
let pg = PostgresBackend::from_env_read_only()?;
tracing::info!(
"DB engine = postgres read-only (default_transaction_read_only = on): {}",
redact_url(&pg.pg_url)
);
Ok(Arc::new(pg))
}
}
pub fn init_db_pg() -> Result<SharedDb, Box<dyn std::error::Error>> {
let pg = PostgresBackend::from_env()?;
tracing::info!("DB engine = postgres: {}", redact_url(&pg.pg_url));
Ok(Arc::new(pg))
}
pub fn index_advisory_lock() -> Result<Option<AdvisoryLock>, Box<dyn std::error::Error>> {
if std::env::var("LEANKG_PG_LOCK")
.ok()
.map(|v| v.eq_ignore_ascii_case("0") || v.eq_ignore_ascii_case("false"))
.unwrap_or(false)
{
tracing::info!("LEANKG_PG_LOCK=0 — index advisory lock disabled");
return Ok(None);
}
let key = PostgresBackend::INDEX_LOCK_KEY;
let mut held = INDEX_LOCK_HELD.lock().unwrap();
if *held {
return Ok(None);
}
let pg = PostgresBackend::from_env()?;
let lock = pg.advisory_lock(key)?;
*held = true;
tracing::info!("index advisory lock held (key {key})");
Ok(Some(lock))
}
static INDEX_LOCK_HELD: std::sync::Mutex<bool> = std::sync::Mutex::new(false);
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn postgres_backend_stub_returns_documented_error() {
let pg = PostgresBackend {
pg_url: "postgres://invalid-host-not-real:1/leankg".into(),
pool: std::sync::Arc::new(ClientPool::new(1)),
ro_pool: std::sync::Arc::new(ClientPool::new(1)),
read_only: false,
};
let err = pg
.run_script("?[a] := *x[a]", Default::default())
.unwrap_err()
.to_string();
assert!(!err.is_empty(), "stub error must not be empty: {err}");
assert!(pg.import_relations(BTreeMap::new()).is_err());
}
#[test]
fn postgres_backend_validates_url_and_redacts() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
assert!(PostgresBackend::from_env().is_err(), "no env -> error");
std::env::set_var("LEANKG_PG_URL", "not-a-url");
let err = PostgresBackend::from_env().unwrap_err();
assert!(err.contains("not-a-url"));
let redacted = redact_url("postgres://user:s3cret@host:5432/db?sslmode=require");
assert!(!redacted.contains("s3cret"));
assert!(redacted.contains("postgres://user:****@host:5432/db?sslmode=require"));
std::env::remove_var("LEANKG_PG_URL");
}
#[test]
fn postgres_backend_requires_url() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
std::env::remove_var("LEANKG_PG_URL");
assert!(init_db_pg().is_err(), "no LEANKG_PG_URL -> error");
assert!(init_db(std::path::Path::new("/tmp/none.db")).is_ok());
}
#[test]
fn init_db_accepts_pg_url() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
std::env::set_var(
"LEANKG_PG_URL",
"postgresql://postgres:postgres@localhost:5433/leankg",
);
let db = PostgresBackend::from_env().unwrap();
assert!(db.pg_url.contains("postgresql://"));
assert!(!db.read_only);
let ro = db.clone().with_read_only();
assert!(ro.read_only);
assert!(ro
.read_only_url()
.contains("default_transaction_read_only%3Don"));
let rw_url = db.read_only_url();
assert!(!rw_url.contains("default_transaction_read_only%3Don"));
std::env::remove_var("LEANKG_PG_URL");
}
#[test]
fn init_db_with_url_produces_pg_backend() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
std::env::set_var(
"LEANKG_PG_URL",
"postgresql://postgres:postgres@localhost:5433/leankg",
);
let db = PostgresBackend::from_env().unwrap();
assert!(
db.run_script("?[a] <- [[1]]", Default::default()).is_err(),
"path-init must produce the PG backend (translator rejects bare lists)"
);
std::env::remove_var("LEANKG_PG_URL");
}
#[test]
fn pool_size_from_env_defaults_and_clamps() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
std::env::remove_var("LEANKG_PG_POOL_SIZE");
assert_eq!(ClientPool::size_from_env(), 5);
std::env::set_var("LEANKG_PG_POOL_SIZE", "0");
assert_eq!(ClientPool::size_from_env(), 5, "0 -> clamp to default");
std::env::set_var("LEANKG_PG_POOL_SIZE", "-3");
assert_eq!(ClientPool::size_from_env(), 5, "negative -> default");
std::env::set_var("LEANKG_PG_POOL_SIZE", "12");
assert_eq!(ClientPool::size_from_env(), 12);
std::env::set_var("LEANKG_PG_POOL_SIZE", "banana");
assert_eq!(ClientPool::size_from_env(), 5, "garbage -> default");
std::env::remove_var("LEANKG_PG_POOL_SIZE");
}
#[test]
fn pool_new_clamps_max_to_one() {
let p = ClientPool::new(0);
assert!(p
.checkout("postgres://invalid-host-not-real:1/leankg")
.is_err());
}
#[test]
fn data_value_roundtrips() {
use crate::db::value::DataValue;
let v = DataValue::from(42i64);
assert_eq!(v.get_int(), Some(42));
let f = DataValue::from(3.5f64);
assert_eq!(f.get_float(), Some(3.5));
let s = DataValue::from("hi");
assert_eq!(s.get_str(), Some("hi"));
let b = DataValue::Bool(true);
assert_eq!(b.get_bool(), Some(true));
}
}