use async_trait::async_trait;
use neo4rs::{Graph, query};
use repolith_core::cache::{Cache, CacheError, Result};
use repolith_core::types::{ActionId, BuildError, BuildEvent, Sha256};
use std::time::Duration;
pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
pub const ENV_URI: &str = "REPOLITH_NEO4J_URI";
pub const ENV_USER: &str = "REPOLITH_NEO4J_USER";
pub const ENV_PASS: &str = "REPOLITH_NEO4J_PASS";
pub struct Neo4jConfig {
pub uri: String,
pub user: String,
pub pass: String,
}
impl Neo4jConfig {
pub fn from_env() -> Result<Self> {
let get = |k: &str| {
std::env::var(k).map_err(|_| {
CacheError::Backend(format!("neo4j cache backend requires the `{k}` env var"))
})
};
Ok(Self {
uri: get(ENV_URI)?,
user: get(ENV_USER)?,
pass: get(ENV_PASS)?,
})
}
}
#[derive(Clone)]
pub struct Neo4jCache {
graph: Graph,
}
impl Neo4jCache {
pub async fn connect(cfg: &Neo4jConfig) -> Result<Self> {
let connect = Graph::new(&cfg.uri, &cfg.user, &cfg.pass);
let graph = match tokio::time::timeout(CONNECT_TIMEOUT, connect).await {
Ok(Ok(g)) => g,
Ok(Err(e)) => {
return Err(CacheError::Backend(format!(
"neo4j connect failed for {}: {}",
cfg.uri,
redact(&e.to_string(), &cfg.pass)
)));
}
Err(_) => {
return Err(CacheError::Backend(format!(
"neo4j unreachable at {} (timed out after {CONNECT_TIMEOUT:?})",
cfg.uri
)));
}
};
let ping = graph.run(query("RETURN 1"));
match tokio::time::timeout(CONNECT_TIMEOUT, ping).await {
Ok(Ok(())) => Ok(Self { graph }),
Ok(Err(e)) => Err(CacheError::Backend(format!(
"neo4j ping failed for {}: {}",
cfg.uri,
redact(&e.to_string(), &cfg.pass)
))),
Err(_) => Err(CacheError::Backend(format!(
"neo4j unreachable at {} (timed out after {CONNECT_TIMEOUT:?})",
cfg.uri
))),
}
}
fn upsert_query(ev: &BuildEvent) -> neo4rs::Query {
let (id, input, output, error_json, kind, ms) = match ev {
BuildEvent::Success {
id,
input,
output,
ms,
} => (
id.0.clone(),
input.to_string(),
output.to_string(),
String::new(),
"success",
i64::try_from(*ms).unwrap_or(i64::MAX),
),
BuildEvent::Failed {
id,
input,
error,
ms,
} => (
id.0.clone(),
input.to_string(),
String::new(),
serde_json::to_string(error).unwrap_or_default(),
"failed",
i64::try_from(*ms).unwrap_or(i64::MAX),
),
};
query(
"MERGE (a:Action {id: $id})
WITH a
OPTIONAL MATCH (a)-[r:LAST]->(old:BuildEvent)
DELETE r
DETACH DELETE old
WITH a
CREATE (a)-[:LAST]->(:BuildEvent {
kind: $kind, input: $input, output: $output,
error_json: $error_json, ms: $ms, recorded_at: timestamp()
})",
)
.param("id", id)
.param("kind", kind)
.param("input", input)
.param("output", output)
.param("error_json", error_json)
.param("ms", ms)
}
}
fn redact(msg: &str, pass: &str) -> String {
if pass.is_empty() {
msg.to_string()
} else {
msg.replace(pass, "***")
}
}
fn parse_sha256(s: &str) -> Option<Sha256> {
let bytes = hex::decode(s).ok()?;
let arr: [u8; 32] = bytes.try_into().ok()?;
Some(Sha256(arr))
}
#[async_trait]
impl Cache for Neo4jCache {
async fn last_build(&self, id: &ActionId) -> Option<BuildEvent> {
let q = query(
"MATCH (a:Action {id: $id})-[:LAST]->(e:BuildEvent)
RETURN e.kind AS kind, e.input AS input, e.output AS output,
e.error_json AS error_json, e.ms AS ms",
)
.param("id", id.0.clone());
let mut rows = self.graph.execute(q).await.ok()?;
let row = rows.next().await.ok()??;
let kind: String = row.get("kind").ok()?;
let input = parse_sha256(&row.get::<String>("input").ok()?)?;
let ms = u64::try_from(row.get::<i64>("ms").ok()?).unwrap_or(0);
if kind == "success" {
let output = parse_sha256(&row.get::<String>("output").ok()?)?;
Some(BuildEvent::Success {
id: id.clone(),
input,
output,
ms,
})
} else {
let error = serde_json::from_str::<BuildError>(&row.get::<String>("error_json").ok()?)
.unwrap_or(BuildError::Cancelled);
Some(BuildEvent::Failed {
id: id.clone(),
input,
error,
ms,
})
}
}
async fn record(&mut self, event: BuildEvent) -> Result<()> {
self.graph
.run(Self::upsert_query(&event))
.await
.map_err(|e| CacheError::Backend(format!("neo4j record: {e}")))
}
async fn record_batch(&mut self, events: Vec<BuildEvent>) -> Result<()> {
if events.is_empty() {
return Ok(());
}
let mut txn = self
.graph
.start_txn()
.await
.map_err(|e| CacheError::Backend(format!("neo4j begin txn: {e}")))?;
for ev in &events {
if let Err(e) = txn.run(Self::upsert_query(ev)).await {
let _ = txn.rollback().await;
return Err(CacheError::Backend(format!("neo4j batch write: {e}")));
}
}
txn.commit()
.await
.map_err(|e| CacheError::Backend(format!("neo4j commit: {e}")))
}
}