use std::collections::HashSet;
use serde::{Deserialize, Serialize};
use super::error::GraphError;
use super::value::{ACCEPTED_TYPES, DeclaredColumn, GraphType};
pub const DEFAULT_QUERY_TIMEOUT_SECONDS: u64 = 30;
pub const MAX_QUERY_TIMEOUT_SECONDS: u64 = 86_400;
pub const DEFAULT_MAX_ROWS: usize = 10_000;
pub const MAX_MAX_ROWS: usize = 1_000_000;
pub const DEFAULT_MAX_CONNECTIONS: u32 = 4;
pub const MAX_MAX_CONNECTIONS: u32 = 64;
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct GraphConfig {
pub backend: String,
pub graph_name: String,
#[serde(default)]
pub username_env: Option<String>,
#[serde(default)]
pub password_env: Option<String>,
#[serde(default = "default_timeout")]
pub query_timeout_seconds: u64,
#[serde(default = "default_max_rows")]
pub max_rows: usize,
#[serde(default = "default_max_connections")]
pub max_connections: u32,
#[serde(default, deserialize_with = "views_null_is_empty")]
pub views: Vec<GraphView>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct GraphView {
pub name: String,
pub cypher: String,
pub schema: Vec<GraphViewColumn>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct GraphViewColumn {
pub name: String,
#[serde(rename = "type")]
pub r#type: String,
#[serde(default = "default_nullable")]
pub nullable: bool,
}
fn default_nullable() -> bool {
true
}
fn views_null_is_empty<'de, D>(deserializer: D) -> Result<Vec<GraphView>, D::Error>
where
D: serde::Deserializer<'de>,
{
let views: Option<Vec<GraphView>> = Deserialize::deserialize(deserializer)?;
Ok(views.unwrap_or_default())
}
impl GraphView {
pub fn declared_columns(&self) -> Result<Vec<DeclaredColumn>, GraphError> {
self.schema
.iter()
.map(|c| {
Ok(DeclaredColumn {
name: c.name.clone(),
ty: GraphType::parse(&c.r#type).ok_or_else(|| GraphError::InvalidConfig {
name: self.name.clone(),
reason: format!(
"column '{}' declares unknown type '{}' (accepted types: {})",
c.name, c.r#type, ACCEPTED_TYPES
),
})?,
nullable: c.nullable,
})
})
.collect()
}
}
fn default_timeout() -> u64 {
DEFAULT_QUERY_TIMEOUT_SECONDS
}
fn default_max_rows() -> usize {
DEFAULT_MAX_ROWS
}
fn default_max_connections() -> u32 {
DEFAULT_MAX_CONNECTIONS
}
impl GraphConfig {
pub fn validate(&self, name: &str, connection_string: &str) -> Result<(), GraphError> {
if self.backend != "age" {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"backend '{}' is not supported (milestone 1 supports: age; \
neo4j and kuzu are later milestones)",
self.backend
),
});
}
let scheme_ok = connection_string.starts_with("postgres://")
|| connection_string.starts_with("postgresql://");
if !scheme_ok {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: "the age backend requires a postgres:// or postgresql:// \
connection_string"
.to_string(),
});
}
let parsed = url::Url::parse(connection_string).map_err(|_| GraphError::InvalidConfig {
name: name.to_string(),
reason: "connection_string is not a parseable URL (the embedded-credential \
check could not run, so the value is rejected rather than passed \
through unvetted)"
.to_string(),
})?;
if parsed.password().is_some() {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: "connection_string must not embed a password — set \
password_env to the NAME of an environment variable \
instead"
.to_string(),
});
}
if parsed
.query_pairs()
.any(|(k, _)| k.eq_ignore_ascii_case("password"))
{
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: "connection_string must not carry a password= query \
parameter — set password_env instead"
.to_string(),
});
}
if self.graph_name.is_empty()
|| !self
.graph_name
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_')
{
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"graph_name '{}' must be a bare identifier ([A-Za-z0-9_]+)",
self.graph_name
),
});
}
for (field, value) in [
("username_env", &self.username_env),
("password_env", &self.password_env),
] {
if let Some(v) = value
&& !is_identifier(v)
{
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"{field} '{v}' must be an environment variable NAME \
([A-Za-z_][A-Za-z0-9_]*)"
),
});
}
}
if self.max_rows == 0 || self.max_rows > MAX_MAX_ROWS {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"max_rows must be in 1..={MAX_MAX_ROWS} (got {}) — the milestone-1 \
client buffers the whole result, so this knob is peak memory, and \
it gets the same hard ceiling the timeout got",
self.max_rows
),
});
}
if self.query_timeout_seconds == 0 || self.query_timeout_seconds > MAX_QUERY_TIMEOUT_SECONDS
{
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"query_timeout_seconds must be in 1..={MAX_QUERY_TIMEOUT_SECONDS} \
(got {}) — the value feeds Postgres's statement_timeout and the \
client-side wrap",
self.query_timeout_seconds
),
});
}
if self.max_connections == 0 || self.max_connections > MAX_MAX_CONNECTIONS {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"max_connections must be in 1..={MAX_MAX_CONNECTIONS} (got {})",
self.max_connections
),
});
}
let mut view_names = HashSet::new();
for view in &self.views {
if !is_lowercase_identifier(&view.name) {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"view name '{}' must be a lowercase identifier \
([a-z_][a-z0-9_]*) — it becomes the catalog table name \
{name}.main.{}, and DataFusion folds unquoted SQL \
identifiers to lowercase, so an uppercase name would \
register but be unreachable without quoting",
view.name, view.name
),
});
}
if !view_names.insert(&view.name) {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!("duplicate view name '{}'", view.name),
});
}
if view.cypher.trim().is_empty() {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!("view '{}' declares empty cypher", view.name),
});
}
super::guard::reject_mutations(&view.cypher).map_err(|e| {
GraphError::InvalidConfig {
name: name.to_string(),
reason: format!("view '{}': {e}", view.name),
}
})?;
if view.schema.is_empty() {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"view '{}' declares an empty schema — at least one column is \
required (the declared schema is the planning-time contract)",
view.name
),
});
}
let mut column_names = HashSet::new();
for column in &view.schema {
if column.name.trim().is_empty() {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"view '{}' declares a column with an empty name",
view.name
),
});
}
if !column_names.insert(&column.name) {
return Err(GraphError::InvalidConfig {
name: name.to_string(),
reason: format!(
"view '{}' declares column '{}' twice",
view.name, column.name
),
});
}
}
view.declared_columns().map_err(|e| {
let reason = match e {
GraphError::InvalidConfig { reason, .. } => reason,
other => other.to_string(),
};
GraphError::InvalidConfig {
name: name.to_string(),
reason: format!("view '{}': {reason}", view.name),
}
})?;
}
Ok(())
}
}
fn is_identifier(s: &str) -> bool {
let mut chars = s.chars();
chars
.next()
.is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
&& chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
fn is_lowercase_identifier(s: &str) -> bool {
is_identifier(s) && !s.chars().any(|c| c.is_ascii_uppercase())
}
#[cfg(test)]
mod tests {
use super::*;
fn base() -> GraphConfig {
serde_yaml::from_str(
r#"
backend: age
graph_name: knowledge
username_env: AGE_PG_USER
password_env: AGE_PG_PASS
"#,
)
.expect("parses")
}
#[test]
fn defaults_and_valid_config_pass() {
let c = base();
assert_eq!(c.query_timeout_seconds, DEFAULT_QUERY_TIMEOUT_SECONDS);
assert_eq!(c.max_rows, DEFAULT_MAX_ROWS);
c.validate("kg", "postgres://localhost:5432/graphrag")
.expect("valid");
c.validate("kg", "postgresql://h/db").expect("valid");
}
#[test]
fn explicit_null_views_read_as_the_empty_list() {
for spelling in [
"backend: age\ngraph_name: g\nviews: ~\n",
"backend: age\ngraph_name: g\nviews: null\n",
"backend: age\ngraph_name: g\nviews:\n",
"backend: age\ngraph_name: g\n",
] {
let c: GraphConfig = serde_yaml::from_str(spelling)
.unwrap_or_else(|e| panic!("{spelling:?} must parse: {e}"));
assert!(c.views.is_empty(), "{spelling:?}");
c.validate("kg", "postgres://h/db").expect("valid");
}
}
#[test]
fn non_age_backends_and_wrong_schemes_are_named_errors() {
let mut c = base();
c.backend = "neo4j".into();
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("later milestones"), "{err}");
let c = base();
let err = c.validate("kg", "bolt://localhost:7687").unwrap_err();
assert!(err.to_string().contains("postgres://"), "{err}");
}
#[test]
fn graph_name_and_env_names_are_shape_checked() {
let mut c = base();
c.graph_name = "bad-name".into();
assert!(c.validate("kg", "postgres://h/db").is_err());
let mut c = base();
c.username_env = Some("BAD NAME".into());
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("environment variable NAME"));
}
#[test]
fn url_embedded_passwords_are_rejected_without_echoing_the_url() {
let c = base();
for url in [
"postgres://user:s3cret@localhost:5432/db",
"postgresql://h/db?password=s3cret",
"postgres://h/db?PASSWORD=s3cret",
] {
let err = c.validate("kg", url).unwrap_err();
let msg = err.to_string();
assert!(msg.contains("password_env"), "{msg}");
assert!(!msg.contains("s3cret"), "the secret never echoes: {msg}");
}
c.validate("kg", "postgres://postgres@localhost:5432/db")
.expect("username-only URL is fine");
let err = c
.validate("kg", "postgres://[not-a-host/db?password=s3cret")
.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("not a parseable URL"), "{msg}");
assert!(!msg.contains("s3cret"), "never echoed: {msg}");
}
#[test]
fn timeout_has_a_hard_ceiling() {
let mut c = base();
c.query_timeout_seconds = MAX_QUERY_TIMEOUT_SECONDS + 1;
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("1..=86400"), "{err}");
c.query_timeout_seconds = u64::MAX; assert!(c.validate("kg", "postgres://h/db").is_err());
c.query_timeout_seconds = MAX_QUERY_TIMEOUT_SECONDS;
c.validate("kg", "postgres://h/db")
.expect("the ceiling itself is legal");
}
#[test]
fn remaining_bounds_are_validated() {
let mut c = base();
c.max_rows = 0;
assert!(
c.validate("kg", "postgres://h/db")
.unwrap_err()
.to_string()
.contains("max_rows")
);
let mut c = base();
c.max_connections = 0;
assert!(
c.validate("kg", "postgres://h/db")
.unwrap_err()
.to_string()
.contains("max_connections")
);
let mut c = base();
c.graph_name = String::new();
assert!(
c.validate("kg", "postgres://h/db")
.unwrap_err()
.to_string()
.contains("bare identifier")
);
let mut c = base();
c.password_env = Some("2BAD".into());
assert!(c.validate("kg", "postgres://h/db").is_err());
}
#[test]
fn unknown_fields_are_rejected_at_parse() {
let err = serde_yaml::from_str::<GraphConfig>("backend: age\ngraph_name: g\nviewz: []\n")
.unwrap_err();
assert!(err.to_string().contains("viewz"), "{err}");
}
#[test]
fn view_cypher_is_screened_by_the_keyword_guard() {
for (cypher, keyword) in [
("CREATE (n:X) RETURN n", "'CREATE'"),
("CALL db.labels()", "'CALL'"),
("LOAD CSV FROM 'https://x' AS row RETURN row", "'LOAD'"),
] {
let c: GraphConfig = serde_yaml::from_str(&format!(
"backend: age
graph_name: g
views:
- name: v
cypher: \"{cypher}\"
schema:
- name: x
type: string
"
))
.expect("parses");
let err = c.validate("kg", "postgres://h/db").unwrap_err();
let msg = err.to_string();
assert!(msg.contains("view 'v'"), "{cypher}: {msg}");
assert!(msg.contains(keyword), "{cypher}: {msg}");
}
}
#[test]
fn views_parse_validate_and_default_nullable_to_true() {
let c: GraphConfig = serde_yaml::from_str(
r#"
backend: age
graph_name: g
views:
- name: user_posts
cypher: MATCH (u:User)-[:POSTED]->(p:Post) RETURN u.name AS user_name, p.title AS post_title
schema:
- name: user_name
type: string
- name: post_title
type: string
nullable: false
"#,
)
.expect("views parse");
assert_eq!(c.views.len(), 1);
assert!(c.views[0].schema[0].nullable, "nullable defaults to true");
assert!(!c.views[0].schema[1].nullable);
c.validate("kg", "postgres://h/db").expect("valid");
let columns = c.views[0].declared_columns().expect("types parse");
assert_eq!(columns[1].name, "post_title");
assert!(!columns[1].nullable);
}
#[test]
fn view_names_must_be_unique_identifiers() {
for bad in ["user-posts", "1posts", "", "user posts", "userPosts", "Foo"] {
let mut c = base();
c.views = vec![view(bad, "MATCH (n) RETURN n", vec![column("n", "int")])];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("identifier"), "{bad}: {err}");
}
let mut c = base();
c.views = vec![
view("v", "MATCH (n) RETURN n", vec![column("n", "int")]),
view("v", "MATCH (m) RETURN m", vec![column("m", "int")]),
];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("duplicate view name 'v'"), "{err}");
}
#[test]
fn view_cypher_and_schema_must_be_nonempty() {
let mut c = base();
c.views = vec![view("v", " \n ", vec![column("n", "int")])];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("empty cypher"), "{err}");
let mut c = base();
c.views = vec![view("v", "MATCH (n) RETURN n", vec![])];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("empty schema"), "{err}");
}
#[test]
fn view_columns_must_be_named_unique_and_typed() {
let mut c = base();
c.views = vec![view("v", "MATCH (n) RETURN n", vec![column(" ", "int")])];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("empty name"), "{err}");
let mut c = base();
c.views = vec![view(
"v",
"MATCH (n) RETURN n",
vec![column("n", "int"), column("n", "string")],
)];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("twice"), "{err}");
let mut c = base();
c.views = vec![view("v", "MATCH (n) RETURN n", vec![column("n", "Utf8")])];
let err = c.validate("kg", "postgres://h/db").unwrap_err();
let msg = err.to_string();
assert!(msg.contains("unknown type 'Utf8'"), "{msg}");
assert!(msg.contains("node, relationship, path"), "{msg}");
}
fn view(name: &str, cypher: &str, schema: Vec<GraphViewColumn>) -> GraphView {
GraphView {
name: name.to_string(),
cypher: cypher.to_string(),
schema,
}
}
fn column(name: &str, ty: &str) -> GraphViewColumn {
GraphViewColumn {
name: name.to_string(),
r#type: ty.to_string(),
nullable: true,
}
}
#[test]
fn max_rows_and_max_connections_have_hard_ceilings() {
let mut c = base();
c.max_rows = MAX_MAX_ROWS + 1;
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("1..=1000000"), "{err}");
c.max_rows = MAX_MAX_ROWS;
c.validate("kg", "postgres://h/db")
.expect("the ceiling itself is legal");
let mut c = base();
c.max_connections = MAX_MAX_CONNECTIONS + 1;
let err = c.validate("kg", "postgres://h/db").unwrap_err();
assert!(err.to_string().contains("1..=64"), "{err}");
}
#[test]
fn declared_columns_is_typed_even_called_standalone() {
let view = GraphView {
name: "v".to_string(),
cypher: "MATCH (n) RETURN n.id".to_string(),
schema: vec![GraphViewColumn {
name: "id".to_string(),
r#type: "uuid".to_string(),
nullable: true,
}],
};
let err = view.declared_columns().expect_err("unknown type");
let msg = err.to_string();
assert!(msg.contains("column 'id'"), "{msg}");
assert!(msg.contains("unknown type 'uuid'"), "{msg}");
assert!(
msg.contains("string"),
"the accepted list rides along: {msg}"
);
}
}