use std::io::ErrorKind as IoErrorKind;
use std::time::Duration;
use std::sync::atomic::{AtomicU64, Ordering};
use async_trait::async_trait;
use futures::StreamExt;
use futures::stream::BoxStream;
use serde_json::Value;
use sqlx::postgres::{PgConnectOptions, PgPool, PgPoolOptions};
use sqlx::{Connection, Either, Executor, Row};
use super::error::{GraphError, json_kind};
use super::value::parse_agtype;
#[derive(Debug, Clone, Copy)]
pub struct QueryBounds {
pub timeout: Duration,
pub max_rows: usize,
}
pub type GraphRow = Vec<Value>;
pub type GraphRowStream = BoxStream<'static, Result<GraphRow, GraphError>>;
#[async_trait]
pub trait GraphClient: Send + Sync + std::fmt::Debug {
async fn execute(
&self,
cypher: &str,
params: &Value,
arity: usize,
bounds: QueryBounds,
limit: Option<usize>,
) -> Result<GraphRowStream, GraphError>;
async fn labels(
&self,
bounds: QueryBounds,
limit: Option<usize>,
) -> Result<Vec<(String, String)>, GraphError>;
}
struct OpenTxnGuard {
conn: Option<sqlx::pool::PoolConnection<sqlx::Postgres>>,
prepared: Option<String>,
}
impl OpenTxnGuard {
fn conn(&mut self) -> &mut sqlx::PgConnection {
self.conn.as_mut().expect("armed until defused")
}
fn defuse(mut self) -> sqlx::pool::PoolConnection<sqlx::Postgres> {
self.conn.take().expect("defused exactly once")
}
}
impl Drop for OpenTxnGuard {
fn drop(&mut self) {
if let Some(mut conn) = self.conn.take() {
let prepared = self.prepared.take();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
let _ = conn.execute("ROLLBACK").await;
if let Some(name) = prepared {
let _ = conn.execute(format!("DEALLOCATE {name}").as_str()).await;
}
});
}
}
}
}
const PREFLIGHT_TIMEOUT_CEILING: Duration = Duration::from_secs(30);
#[derive(Debug)]
pub struct AgeClient {
pool: PgPool,
graph_name: String,
source_name: String,
prepare_seq: AtomicU64,
}
impl AgeClient {
pub async fn connect(
source_name: &str,
connection_string: &str,
graph_name: &str,
username_env: Option<&str>,
password_env: Option<&str>,
max_connections: u32,
acquire_timeout: Duration,
) -> Result<Self, GraphError> {
let options = build_options(source_name, connection_string, username_env, password_env)?;
let preflight = async {
let mut conn = sqlx::postgres::PgConnection::connect_with(&options)
.await
.map_err(|e| connect_error(source_name, &e))?;
per_connection_setup(&mut conn)
.await
.map_err(|e| backend_error(source_name, &e))?;
let exists: Option<i32> =
sqlx::query_scalar("SELECT 1 FROM ag_catalog.ag_graph WHERE name = $1")
.bind(graph_name)
.fetch_optional(&mut conn)
.await
.map_err(|e| backend_error(source_name, &e))?;
if exists.is_none() {
return Err(GraphError::InvalidConfig {
name: source_name.to_string(),
reason: format!(
"graph '{graph_name}' does not exist in this database \
(ag_catalog.ag_graph has no such entry; create it with \
SELECT create_graph(...) or fix graph_name)"
),
});
}
Ok::<_, GraphError>(())
};
let preflight_timeout = acquire_timeout.min(PREFLIGHT_TIMEOUT_CEILING);
tokio::time::timeout(preflight_timeout, preflight)
.await
.map_err(|_| GraphError::Unavailable {
source_name: source_name.to_string(),
reason: format!(
"the registration preflight did not complete within {}s",
preflight_timeout.as_secs()
),
})??;
Ok(Self::from_pool(
build_pool(options, max_connections, acquire_timeout),
source_name,
graph_name,
))
}
pub fn connect_degraded(
source_name: &str,
connection_string: &str,
graph_name: &str,
username_env: Option<&str>,
password_env: Option<&str>,
max_connections: u32,
acquire_timeout: Duration,
) -> Result<Self, GraphError> {
let options = build_options(source_name, connection_string, username_env, password_env)?;
Ok(Self::from_pool(
build_pool(options, max_connections, acquire_timeout),
source_name,
graph_name,
))
}
fn from_pool(pool: PgPool, source_name: &str, graph_name: &str) -> Self {
tracing::debug!(
source = source_name,
graph = graph_name,
"graph source connected"
);
Self {
pool,
graph_name: graph_name.to_string(),
source_name: source_name.to_string(),
prepare_seq: AtomicU64::new(0),
}
}
#[doc(hidden)]
pub fn pool_for_tests(&self) -> &PgPool {
&self.pool
}
}
fn build_options(
source_name: &str,
connection_string: &str,
username_env: Option<&str>,
password_env: Option<&str>,
) -> Result<PgConnectOptions, GraphError> {
let mut options: PgConnectOptions =
connection_string
.parse()
.map_err(|e: sqlx::Error| GraphError::InvalidConfig {
name: source_name.to_string(),
reason: format!("connection_string does not parse: {e}"),
})?;
if let Some(env) = username_env {
let user = read_env(source_name, "username_env", env)?;
options = options.username(&user);
}
if let Some(env) = password_env {
let pass = read_env(source_name, "password_env", env)?;
options = options.password(&pass);
}
Ok(options)
}
fn build_pool(
options: PgConnectOptions,
max_connections: u32,
acquire_timeout: Duration,
) -> PgPool {
PgPoolOptions::new()
.max_connections(max_connections)
.acquire_timeout(acquire_timeout)
.after_connect(|conn, _meta| Box::pin(per_connection_setup(conn)))
.connect_lazy_with(options)
}
async fn per_connection_setup(conn: &mut sqlx::PgConnection) -> Result<(), sqlx::Error> {
let _ = conn.execute("LOAD 'age'").await;
conn.execute("SET search_path = ag_catalog, \"$user\", public")
.await?;
Ok(())
}
impl AgeClient {
fn build_sql(
&self,
cypher: &str,
params: &Value,
arity: usize,
fetch: usize,
) -> (String, Option<String>) {
build_cypher_sql(
&self.graph_name,
self.prepare_seq.fetch_add(1, Ordering::Relaxed),
cypher,
params,
arity,
fetch,
)
}
}
#[async_trait]
impl GraphClient for AgeClient {
async fn execute(
&self,
cypher: &str,
params: &Value,
arity: usize,
bounds: QueryBounds,
limit: Option<usize>,
) -> Result<GraphRowStream, GraphError> {
let fetch = match limit {
Some(l) if l <= bounds.max_rows => l,
_ => bounds.max_rows.saturating_add(1),
};
let (sql, prepared) = self.build_sql(cypher, params, arity, fetch);
let source = self.source_name.clone();
let started = std::time::Instant::now();
let acquired = self
.pool
.acquire()
.await
.map_err(|e| map_query_error(&source, bounds, &e))?;
let run = async {
let mut guard = OpenTxnGuard {
conn: Some(acquired),
prepared: prepared.clone(),
};
guard
.conn()
.execute("BEGIN TRANSACTION READ ONLY")
.await
.map_err(|e| backend_error(&source, &e))?;
let collected: Result<Vec<GraphRow>, GraphError> = async {
let conn = guard.conn();
conn.execute(
format!(
"SET LOCAL statement_timeout = '{}ms'",
bounds.timeout.as_millis()
)
.as_str(),
)
.await
.map_err(|e| backend_error(&source, &e))?;
let mut rows: Vec<GraphRow> = Vec::new();
let mut stream = conn.fetch_many(sqlx::raw_sql(&sql));
while let Some(step) = stream.next().await {
let step = step.map_err(|e| map_query_error(&source, bounds, &e))?;
let Either::Right(row) = step else { continue };
if limit.is_some_and(|l| rows.len() >= l) {
break;
}
if rows.len() >= bounds.max_rows {
return Err(GraphError::RowCapExceeded {
max_rows: bounds.max_rows,
});
}
let row_idx = rows.len();
let mut values = Vec::with_capacity(arity);
for col in 0..arity {
let text: Option<String> = row
.try_get_unchecked(col)
.map_err(|e| backend_error(&source, &e))?;
let value = match text {
None => Value::Null,
Some(t) => {
parse_agtype(&t).map_err(|reason| GraphError::MalformedCell {
row: row_idx,
column: col,
reason,
})?
}
};
values.push(value);
}
rows.push(values);
}
Ok(rows)
}
.await;
let mut conn = guard.defuse();
let rollback = conn.execute("ROLLBACK").await;
let mut dealloc = Ok(Default::default());
if let Some(name) = &prepared {
dealloc = conn.execute(format!("DEALLOCATE {name}").as_str()).await;
}
let rows = collected?;
rollback.map_err(|e| backend_error(&source, &e))?;
dealloc.map_err(|e| backend_error(&source, &e))?;
Ok(rows)
};
let rows = tokio::time::timeout(bounds.timeout.saturating_add(Duration::from_secs(5)), run)
.await
.map_err(|_| GraphError::BackendSilent {
seconds: bounds
.timeout
.saturating_add(Duration::from_secs(5))
.as_secs(),
})??;
tracing::debug!(
source = %self.source_name,
elapsed_ms = started.elapsed().as_millis() as u64,
rows = rows.len(),
"cypher_query executed"
);
Ok(futures::stream::iter(rows.into_iter().map(Ok)).boxed())
}
async fn labels(
&self,
bounds: QueryBounds,
limit: Option<usize>,
) -> Result<Vec<(String, String)>, GraphError> {
let fetch = match limit {
Some(l) if l <= bounds.max_rows => l,
_ => bounds.max_rows.saturating_add(1),
};
let source = self.source_name.clone();
let mut acquired = self
.pool
.acquire()
.await
.map_err(|e| map_query_error(&source, bounds, &e))?;
let run = async {
let mut tx = acquired
.begin()
.await
.map_err(|e| backend_error(&source, &e))?;
(&mut *tx)
.execute("SET TRANSACTION READ ONLY")
.await
.map_err(|e| backend_error(&source, &e))?;
(&mut *tx)
.execute(
format!(
"SET LOCAL statement_timeout = '{}ms'",
bounds.timeout.as_millis()
)
.as_str(),
)
.await
.map_err(|e| backend_error(&source, &e))?;
let rows = sqlx::query(
"SELECT l.name, l.kind::text \
FROM ag_catalog.ag_label l \
JOIN ag_catalog.ag_graph g ON g.graphid = l.graph \
WHERE g.name = $1 AND l.name NOT LIKE '\\_ag\\_label\\_%' \
ORDER BY l.name LIMIT $2",
)
.bind(&self.graph_name)
.bind(i64::try_from(fetch).unwrap_or(i64::MAX))
.fetch_all(&mut *tx)
.await
.map_err(|e| map_query_error(&source, bounds, &e))?;
if limit.is_none_or(|l| l > bounds.max_rows) && rows.len() > bounds.max_rows {
return Err(GraphError::RowCapExceeded {
max_rows: bounds.max_rows,
});
}
tx.rollback()
.await
.map_err(|e| backend_error(&source, &e))?;
tracing::debug!(source = %source, labels = rows.len(), "graph_schema catalog read");
rows.into_iter()
.map(|row| {
let name: String = row.try_get(0).map_err(|e| backend_error(&source, &e))?;
let kind: String = row.try_get(1).map_err(|e| backend_error(&source, &e))?;
let kind = match kind.as_str() {
"v" => "vertex".to_string(),
"e" => "edge".to_string(),
other => other.to_string(),
};
Ok((name, kind))
})
.collect()
};
tokio::time::timeout(bounds.timeout.saturating_add(Duration::from_secs(5)), run)
.await
.map_err(|_| GraphError::BackendSilent {
seconds: bounds
.timeout
.saturating_add(Duration::from_secs(5))
.as_secs(),
})?
}
}
fn build_cypher_sql(
graph_name: &str,
seq: u64,
cypher: &str,
params: &Value,
arity: usize,
fetch: usize,
) -> (String, Option<String>) {
let tag = dollar_tag(cypher);
let graph = graph_name.replace('\'', "''");
let cols: Vec<String> = (0..arity)
.map(|i| format!("c{i} ag_catalog.agtype"))
.collect();
let outs: Vec<String> = (0..arity).map(|i| format!("c{i}")).collect();
let has_params = params.as_object().is_some_and(|m| !m.is_empty());
if has_params {
let name = format!("skq_p_{}_{seq}", std::process::id());
let literal = params.to_string().replace('\'', "''");
let batch = format!(
"PREPARE {name}(ag_catalog.agtype) AS \
SELECT {} FROM ag_catalog.cypher('{graph}', {tag}{cypher}{tag}, $1) \
AS t({}) LIMIT {fetch}; \
EXECUTE {name}('{literal}');",
outs.join(", "),
cols.join(", ")
);
(batch, Some(name))
} else {
(
format!(
"SELECT {} FROM ag_catalog.cypher('{graph}', {tag}{cypher}{tag}) \
AS t({}) LIMIT {fetch}",
outs.join(", "),
cols.join(", ")
),
None,
)
}
}
fn dollar_tag(text: &str) -> String {
let probe = format!("{text}$");
let mut base_used = false;
let mut used: std::collections::HashSet<u32> = std::collections::HashSet::new();
for (i, _) in probe.match_indices("$skq") {
let rest = &probe[i + "$skq".len()..];
if let Some(end) = rest.find('$') {
let digits = &rest[..end];
if digits.is_empty() {
base_used = true;
} else if let Ok(n) = digits.parse::<u32>() {
used.insert(n);
}
}
}
if !base_used {
return "$skq$".to_string();
}
let mut n = 1u32;
while used.contains(&n) {
n += 1;
}
format!("$skq{n}$")
}
fn read_env(source: &str, field: &str, env: &str) -> Result<String, GraphError> {
std::env::var(env).map_err(|_| GraphError::InvalidConfig {
name: source.to_string(),
reason: format!("{field}: ${env} is not set in this environment"),
})
}
fn backend_error(source: &str, e: &sqlx::Error) -> GraphError {
let (code, message) = match e {
sqlx::Error::Database(db) => (db.code().map(|c| c.to_string()), db.message().to_string()),
_ => (None, e.to_string()),
};
GraphError::backend(source, code.as_deref().unwrap_or("io"), &message)
}
fn connect_error(source: &str, e: &sqlx::Error) -> GraphError {
match e {
sqlx::Error::Io(io) if io.kind() != IoErrorKind::InvalidData => GraphError::Unavailable {
source_name: source.to_string(),
reason: io.to_string(),
},
other => backend_error(source, other),
}
}
fn map_query_error(source: &str, bounds: QueryBounds, e: &sqlx::Error) -> GraphError {
if let sqlx::Error::Database(db) = e
&& db.code().as_deref() == Some("57014")
{
return GraphError::Timeout {
seconds: bounds.timeout.as_secs(),
};
}
if matches!(e, sqlx::Error::PoolTimedOut) {
return GraphError::ConnectionAcquireTimeout {
seconds: bounds.timeout.as_secs(),
};
}
backend_error(source, e)
}
pub fn validate_params(params: &Value) -> Result<(), GraphError> {
match params {
Value::Object(_) => Ok(()),
other => Err(GraphError::InvalidParams {
found: json_kind(other).to_string(),
}),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dial_io_errors_split_on_talking_versus_nobody_answered() {
let tls = sqlx::Error::Io(std::io::Error::new(
IoErrorKind::InvalidData,
"invalid peer certificate: UnknownIssuer",
));
assert!(
!matches!(connect_error("kg", &tls), GraphError::Unavailable { .. }),
"a TLS handshake failure is not an outage"
);
for kind in [
IoErrorKind::ConnectionRefused,
IoErrorKind::TimedOut,
IoErrorKind::ConnectionReset,
] {
let e = sqlx::Error::Io(std::io::Error::new(kind, "dial failed"));
assert!(
matches!(connect_error("kg", &e), GraphError::Unavailable { .. }),
"{kind:?} must degrade"
);
}
assert!(!matches!(
connect_error("kg", &sqlx::Error::PoolTimedOut),
GraphError::Unavailable { .. }
));
}
#[test]
fn dollar_tags_never_collide_with_the_text() {
assert_eq!(dollar_tag("MATCH (n) RETURN n"), "$skq$");
let hostile = "RETURN '$skq$ $skq1$'";
let tag = dollar_tag(hostile);
assert!(!hostile.contains(&tag), "{tag}");
let boundary = "RETURN $skq";
let tag = dollar_tag(boundary);
assert_ne!(tag, "$skq$", "boundary composition must bump the tag");
assert!(!format!("{boundary}$").contains(&tag), "{tag}");
let (sql, _) = build_cypher_sql("kg", 0, boundary, &serde_json::json!({}), 1, 11);
let open = sql.find(&tag).expect("opening tag");
let close = sql[open + tag.len()..].find(&tag).expect("closing tag");
assert_eq!(
&sql[open + tag.len()..open + tag.len() + close],
boundary,
"the whole query is the literal body: {sql}"
);
}
#[test]
fn parameterless_sql_is_a_single_select_with_no_prepare() {
let (sql, prepared) = build_cypher_sql(
"kg",
0,
"MATCH (n) RETURN n.a, n.b",
&serde_json::json!({}),
2,
101,
);
assert!(prepared.is_none(), "no params → nothing to DEALLOCATE");
assert!(
sql.starts_with("SELECT c0, c1 FROM ag_catalog.cypher('kg', $skq$"),
"{sql}"
);
assert!(
sql.ends_with("AS t(c0 ag_catalog.agtype, c1 ag_catalog.agtype) LIMIT 101"),
"{sql}"
);
assert!(!sql.contains("PREPARE"), "{sql}");
}
#[test]
fn parameterized_sql_prepares_executes_and_names_the_statement() {
let (sql, prepared) = build_cypher_sql(
"kg",
7,
"MATCH (n) WHERE n.x = $x RETURN n",
&serde_json::json!({"x": 1}),
1,
5,
);
assert!(
sql.contains("ag_catalog.agtype) LIMIT 5;"),
"the PREPARE body carries the fetch LIMIT: {sql}"
);
let name = prepared.expect("params → PREPARE name to DEALLOCATE");
assert!(name.starts_with("skq_p_"), "{name}");
assert!(name.ends_with("_7"), "the seq uniquifies: {name}");
assert!(
sql.contains(&format!("PREPARE {name}(ag_catalog.agtype)")),
"{sql}"
);
assert!(
sql.contains(&format!("EXECUTE {name}('{{\"x\":1}}');")),
"{sql}"
);
}
#[test]
fn hostile_values_stay_inside_their_literals() {
let (sql, _) = build_cypher_sql(
"kg",
0,
"RETURN $s",
&serde_json::json!({"s": "O'Brien '; DROP TABLE x; --"}),
1,
10,
);
assert!(sql.contains("O''Brien ''; DROP TABLE x; --"), "{sql}");
let (sql, _) = build_cypher_sql("g'name", 0, "RETURN 1", &serde_json::json!({}), 1, 10);
assert!(sql.contains("cypher('g''name'"), "{sql}");
}
#[test]
fn params_must_be_an_object() {
assert!(validate_params(&serde_json::json!({"a": 1})).is_ok());
let err = validate_params(&serde_json::json!([1])).unwrap_err();
assert!(err.to_string().contains("an array"), "{err}");
}
#[tokio::test]
async fn a_blackholed_backend_fails_the_preflight_as_unavailable_within_the_bound() {
let started = std::time::Instant::now();
let err = AgeClient::connect(
"kg",
"postgres://10.255.255.1:5432/none",
"knowledge",
None,
None,
1,
Duration::from_secs(1),
)
.await
.expect_err("a blackhole cannot be reached");
assert!(
matches!(err, GraphError::Unavailable { .. }),
"classified as an availability failure: {err}"
);
assert!(
started.elapsed() < Duration::from_secs(10),
"bounded by the configured timeout, not the OS's"
);
}
#[test]
fn only_transport_failures_classify_as_unavailable() {
let io: sqlx::Error =
std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "refused").into();
let err = connect_error("kg", &io);
assert!(matches!(err, GraphError::Unavailable { .. }), "{err}");
assert!(err.to_string().contains("unreachable"), "{err}");
let err = connect_error("kg", &sqlx::Error::PoolTimedOut);
assert!(matches!(err, GraphError::Backend { .. }), "{err}");
}
#[test]
fn a_pool_acquire_timeout_is_not_misreported_as_a_statement_timeout() {
let bounds = QueryBounds {
timeout: std::time::Duration::from_secs(3),
max_rows: 10,
};
let err = map_query_error("kg", bounds, &sqlx::Error::PoolTimedOut);
let msg = err.to_string();
assert!(
matches!(err, GraphError::ConnectionAcquireTimeout { seconds: 3 }),
"{msg}"
);
assert!(msg.contains("never started"), "{msg}");
assert!(!msg.contains("narrow the traversal"), "{msg}");
}
#[tokio::test]
async fn connect_degraded_hard_fails_config_errors_but_never_dials() {
let err = AgeClient::connect_degraded(
"kg",
"not a url",
"g",
None,
None,
4,
Duration::from_secs(1),
)
.unwrap_err();
assert!(err.to_string().contains("does not parse"), "{err}");
let err = AgeClient::connect_degraded(
"kg",
"postgres://127.0.0.1:1/none",
"g",
Some("SKARDI_TEST_GRAPH_DEFINITELY_UNSET"),
None,
4,
Duration::from_secs(1),
)
.unwrap_err();
assert!(err.to_string().contains("is not set"), "{err}");
AgeClient::connect_degraded(
"kg",
"postgres://127.0.0.1:1/none",
"g",
None,
None,
4,
Duration::from_secs(1),
)
.expect("no dial at connect_degraded");
}
}