use std::any::Any;
use std::collections::{HashMap, HashSet};
use std::fmt;
use std::sync::{Arc, Mutex as StdMutex, RwLock};
use std::time::SystemTime;
use tokio::sync::Mutex as AsyncMutex;
use arrow::datatypes::{DataType, Field, Schema, SchemaRef};
use arrow::record_batch::RecordBatch;
use async_trait::async_trait;
use datafusion::catalog::{Session, TableFunctionImpl, TableProvider};
use datafusion::common::plan_err;
use datafusion::datasource::TableType;
use datafusion::error::{DataFusionError, Result as DFResult};
use datafusion::execution::{SendableRecordBatchStream, TaskContext};
use datafusion::logical_expr::Expr;
use datafusion::physical_expr::EquivalenceProperties;
use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
use datafusion::physical_plan::{
DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, PlanProperties,
};
use datafusion::prelude::SessionContext;
use futures::StreamExt;
use futures::stream::{self, TryStreamExt};
use serde_json::{Map, Value};
use super::client::{GraphClient, QueryBounds, validate_params};
use super::error::{GraphError, json_kind};
use super::guard::reject_mutations;
use super::value::{ACCEPTED_TYPES, DeclaredColumn, GraphType, build_batch, declared_schema};
use crate::sources::providers::udtf_args::{strict_string_arg, string_arg};
const CONVERSION_BATCH_ROWS: usize = 1024;
#[derive(Debug)]
pub struct GraphSourceHandle {
pub client: Arc<dyn GraphClient>,
pub bounds: QueryBounds,
pub health: Arc<RwLock<GraphSourceHealth>>,
pub view_contracts: Arc<Vec<super::view::ViewContract>>,
pub recovery_gate: Arc<AsyncMutex<()>>,
pub last_failed_recovery: Arc<StdMutex<Option<tokio::time::Instant>>>,
pub validation_limit: usize,
pub health_changed_at: Arc<StdMutex<SystemTime>>,
}
impl GraphSourceHandle {
pub fn new(
client: Arc<dyn GraphClient>,
bounds: QueryBounds,
health: GraphSourceHealth,
view_contracts: Arc<Vec<super::view::ViewContract>>,
validation_limit: usize,
) -> Self {
Self {
client,
bounds,
health: Arc::new(RwLock::new(health)),
view_contracts,
recovery_gate: Arc::new(AsyncMutex::new(())),
last_failed_recovery: Arc::new(StdMutex::new(None)),
validation_limit,
health_changed_at: Arc::new(StdMutex::new(SystemTime::now())),
}
}
}
#[derive(Debug, Clone)]
pub enum GraphSourceHealth {
Healthy,
Degraded(String),
}
impl GraphSourceHealth {
pub fn is_healthy(&self) -> bool {
matches!(self, Self::Healthy)
}
}
pub type GraphSources = Arc<RwLock<HashMap<String, Arc<GraphSourceHandle>>>>;
pub fn register_graph_udtfs(ctx: &SessionContext, sources: GraphSources) -> DFResult<()> {
ctx.register_udtf(
"cypher_query",
Arc::new(CypherQueryFunction {
sources: Arc::clone(&sources),
}),
);
ctx.register_udtf("graph_schema", Arc::new(GraphSchemaFunction { sources }));
Ok(())
}
fn lookup(sources: &GraphSources, name: &str) -> DFResult<Arc<GraphSourceHandle>> {
let map = sources.read().unwrap_or_else(|p| p.into_inner());
map.get(name).cloned().ok_or_else(|| {
let mut known: Vec<&str> = map.keys().map(String::as_str).collect();
known.sort_unstable();
plan_error(GraphError::ConnectionNotFound {
name: name.to_string(),
known: if known.is_empty() {
"none".to_string()
} else {
known.join(", ")
},
})
})
}
fn plan_error(e: GraphError) -> DataFusionError {
DataFusionError::Plan(e.to_string())
}
#[derive(Debug)]
pub struct CypherQueryFunction {
sources: GraphSources,
}
impl TableFunctionImpl for CypherQueryFunction {
fn call(&self, exprs: &[Expr]) -> DFResult<Arc<dyn TableProvider>> {
if exprs.len() < 2 || exprs.len() > 4 {
return plan_err!(
"cypher_query(connection, cypher, [params_json], [columns_json]) \
expects 2-4 arguments, got {}",
exprs.len()
);
}
let connection = strict_string_arg(&exprs[0], "cypher_query", "connection")?;
let cypher = strict_string_arg(&exprs[1], "cypher_query", "cypher")?;
let params_json = exprs
.get(2)
.map(|e| string_arg(e, "cypher_query", "params_json"))
.transpose()?;
let columns_json = exprs
.get(3)
.map(|e| strict_string_arg(e, "cypher_query", "columns_json"))
.transpose()?;
reject_mutations(&cypher).map_err(plan_error)?;
let params: Value = match params_json.as_deref() {
None | Some("") => Value::Object(Map::new()),
Some(text) => {
let parsed: Value = serde_json::from_str(text).map_err(|e| {
plan_error(GraphError::InvalidParams {
found: format!("unparseable JSON ({e})"),
})
})?;
validate_params(&parsed).map_err(plan_error)?;
parsed
}
};
let Some(columns_json) = columns_json else {
return plan_err!(
"cypher_query: 'columns' is required on the age backend — declare the \
output columns IN THE SAME ORDER AS YOUR RETURN CLAUSE (the binding \
is positional; two same-typed columns declared out of order swap \
silently), e.g. '{{\"name\": \"string\", \"n\": \"node\"}}' \
(accepted types: {})",
ACCEPTED_TYPES
);
};
let columns = parse_columns(&columns_json).map_err(plan_error)?;
let handle = lookup(&self.sources, &connection)?;
Ok(Arc::new(CypherQueryProvider {
handle,
cypher,
params,
columns: Arc::new(columns),
}))
}
}
fn parse_columns(text: &str) -> Result<Vec<DeclaredColumn>, GraphError> {
let parsed: Value = serde_json::from_str(text).map_err(|e| GraphError::InvalidColumns {
reason: format!("unparseable JSON ({e})"),
accepted: ACCEPTED_TYPES,
})?;
let Value::Object(map) = parsed else {
return Err(GraphError::InvalidColumns {
reason: format!("expected a JSON object, got {}", json_kind(&parsed)),
accepted: ACCEPTED_TYPES,
});
};
if map.is_empty() {
return Err(GraphError::InvalidColumns {
reason: "at least one column must be declared".to_string(),
accepted: ACCEPTED_TYPES,
});
}
let pairs = object_pairs(text).map_err(|e| GraphError::InvalidColumns {
reason: format!("unparseable JSON ({e})"),
accepted: ACCEPTED_TYPES,
})?;
let mut seen = HashSet::with_capacity(pairs.len());
for (name, _) in &pairs {
if !seen.insert(name.as_str()) {
return Err(GraphError::InvalidColumns {
reason: format!("column '{name}' is declared twice"),
accepted: ACCEPTED_TYPES,
});
}
}
map.into_iter()
.map(|(name, ty)| {
let Value::String(ty_name) = &ty else {
return Err(GraphError::InvalidColumns {
reason: format!(
"column '{name}': type must be a string, got {}",
json_kind(&ty)
),
accepted: ACCEPTED_TYPES,
});
};
let ty = GraphType::parse(ty_name).ok_or_else(|| GraphError::InvalidColumns {
reason: format!("column '{name}': unknown type '{ty_name}'"),
accepted: ACCEPTED_TYPES,
})?;
Ok(DeclaredColumn {
name,
ty,
nullable: true,
})
})
.collect()
}
fn object_pairs(text: &str) -> Result<Vec<(String, Value)>, serde_json::Error> {
struct Pairs;
impl<'de> serde::de::Visitor<'de> for Pairs {
type Value = Vec<(String, Value)>;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("a JSON object")
}
fn visit_map<A: serde::de::MapAccess<'de>>(
self,
mut access: A,
) -> Result<Self::Value, A::Error> {
let mut pairs = Vec::new();
while let Some(entry) = access.next_entry::<String, Value>()? {
pairs.push(entry);
}
Ok(pairs)
}
}
let mut de = serde_json::Deserializer::from_str(text);
let pairs = serde::de::Deserializer::deserialize_map(&mut de, Pairs)?;
de.end()?;
Ok(pairs)
}
#[derive(Debug)]
struct CypherQueryProvider {
handle: Arc<GraphSourceHandle>,
cypher: String,
params: Value,
columns: Arc<Vec<DeclaredColumn>>,
}
#[async_trait]
impl TableProvider for CypherQueryProvider {
fn as_any(&self) -> &dyn Any {
self
}
fn schema(&self) -> SchemaRef {
declared_schema(&self.columns)
}
fn table_type(&self) -> TableType {
TableType::Base
}
async fn scan(
&self,
_state: &dyn Session,
projection: Option<&Vec<usize>>,
_filters: &[Expr],
limit: Option<usize>,
) -> DFResult<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(GraphScanExec::new(
GraphScanKind::Cypher {
handle: Arc::clone(&self.handle),
cypher: self.cypher.clone(),
params: self.params.clone(),
columns: Arc::clone(&self.columns),
limit,
},
self.schema(),
projection.cloned(),
)?))
}
}
#[derive(Debug)]
pub struct GraphSchemaFunction {
sources: GraphSources,
}
impl TableFunctionImpl for GraphSchemaFunction {
fn call(&self, exprs: &[Expr]) -> DFResult<Arc<dyn TableProvider>> {
if exprs.len() != 1 {
return plan_err!(
"graph_schema(connection) expects exactly 1 argument, got {}",
exprs.len()
);
}
let connection = strict_string_arg(&exprs[0], "graph_schema", "connection")?;
let handle = lookup(&self.sources, &connection)?;
Ok(Arc::new(GraphSchemaProvider { handle }))
}
}
fn graph_schema_schema() -> SchemaRef {
Arc::new(Schema::new(vec![
Field::new("label", DataType::Utf8, false),
Field::new("kind", DataType::Utf8, false),
]))
}
#[derive(Debug)]
struct GraphSchemaProvider {
handle: Arc<GraphSourceHandle>,
}
#[async_trait]
impl TableProvider for GraphSchemaProvider {
fn as_any(&self) -> &dyn Any {
self
}
fn schema(&self) -> SchemaRef {
graph_schema_schema()
}
fn table_type(&self) -> TableType {
TableType::Base
}
async fn scan(
&self,
_state: &dyn Session,
projection: Option<&Vec<usize>>,
_filters: &[Expr],
limit: Option<usize>,
) -> DFResult<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(GraphScanExec::new(
GraphScanKind::Labels {
handle: Arc::clone(&self.handle),
limit,
},
self.schema(),
projection.cloned(),
)?))
}
}
pub(crate) enum GraphScanKind {
Cypher {
handle: Arc<GraphSourceHandle>,
cypher: String,
params: Value,
columns: Arc<Vec<DeclaredColumn>>,
limit: Option<usize>,
},
View {
handle: Arc<GraphSourceHandle>,
view_name: String,
cypher: String,
columns: Arc<Vec<DeclaredColumn>>,
limit: Option<usize>,
},
Labels {
handle: Arc<GraphSourceHandle>,
limit: Option<usize>,
},
}
impl fmt::Debug for GraphScanKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Cypher { columns, .. } => f
.debug_struct("Cypher")
.field("columns", &columns.len())
.finish_non_exhaustive(),
Self::View { view_name, .. } => f
.debug_struct("View")
.field("view", view_name)
.finish_non_exhaustive(),
Self::Labels { .. } => f.debug_struct("Labels").finish_non_exhaustive(),
}
}
}
#[derive(Debug)]
pub(crate) struct GraphScanExec {
kind: GraphScanKind,
projection: Option<Vec<usize>>,
properties: PlanProperties,
}
impl GraphScanExec {
pub(crate) fn new(
kind: GraphScanKind,
schema: SchemaRef,
projection: Option<Vec<usize>>,
) -> DFResult<Self> {
let projected = match &projection {
Some(indices) => Arc::new(schema.project(indices)?),
None => Arc::clone(&schema),
};
let properties = PlanProperties::new(
EquivalenceProperties::new(Arc::clone(&projected)),
Partitioning::UnknownPartitioning(1),
EmissionType::Final,
Boundedness::Bounded,
);
Ok(Self {
kind,
projection,
properties,
})
}
fn projected_schema(&self) -> SchemaRef {
Arc::clone(self.properties.eq_properties.schema())
}
}
impl DisplayAs for GraphScanExec {
fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
match &self.kind {
GraphScanKind::Cypher { columns, .. } => {
write!(f, "GraphScanExec: cypher_query columns={}", columns.len())
}
GraphScanKind::View { view_name, .. } => {
write!(f, "GraphScanExec: view {view_name}")
}
GraphScanKind::Labels { .. } => write!(f, "GraphScanExec: graph_schema"),
}
}
}
impl ExecutionPlan for GraphScanExec {
fn name(&self) -> &str {
"GraphScanExec"
}
fn as_any(&self) -> &dyn Any {
self
}
fn properties(&self) -> &PlanProperties {
&self.properties
}
fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
vec![]
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn ExecutionPlan>>,
) -> DFResult<Arc<dyn ExecutionPlan>> {
if children.is_empty() {
Ok(self)
} else {
Err(DataFusionError::Internal(
"GraphScanExec is a leaf plan and takes no children".to_string(),
))
}
}
fn execute(
&self,
partition: usize,
_context: Arc<TaskContext>,
) -> DFResult<SendableRecordBatchStream> {
if partition != 0 {
return Err(DataFusionError::Internal(format!(
"GraphScanExec has 1 partition, got partition {partition}"
)));
}
let projected = self.projected_schema();
let projection = self.projection.clone();
let batches = match &self.kind {
GraphScanKind::Cypher {
handle,
cypher,
params,
columns,
limit,
} => cypher_batches(
Arc::clone(handle),
cypher.clone(),
params.clone(),
Arc::clone(columns),
*limit,
),
GraphScanKind::View {
handle,
view_name,
cypher,
columns,
limit,
} => super::view::view_batches(
Arc::clone(handle),
view_name.clone(),
cypher.clone(),
Arc::clone(columns),
*limit,
),
GraphScanKind::Labels { handle, limit } => labels_batch(Arc::clone(handle), *limit),
};
let stream = batches
.map(move |batch| {
let batch = batch?;
match &projection {
Some(indices) => batch
.project(indices)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None)),
None => Ok(batch),
}
})
.boxed();
Ok(Box::pin(RecordBatchStreamAdapter::new(projected, stream)))
}
}
pub(crate) fn degraded_reason(handle: &GraphSourceHandle) -> Option<String> {
let health = handle.health.read().unwrap_or_else(|p| p.into_inner());
match &*health {
GraphSourceHealth::Healthy => None,
GraphSourceHealth::Degraded(reason) => Some(reason.clone()),
}
}
pub(crate) fn mark_healthy(handle: &GraphSourceHandle) {
*handle.health.write().unwrap_or_else(|p| p.into_inner()) = GraphSourceHealth::Healthy;
*handle
.health_changed_at
.lock()
.unwrap_or_else(|p| p.into_inner()) = SystemTime::now();
}
fn degraded_execution_error(degraded: Option<String>, e: GraphError) -> DataFusionError {
match degraded {
Some(reason) if super::is_availability_artifact(&e) => DataFusionError::Execution(format!(
"graph source is registered DEGRADED (registration error: {reason}); \
the query was retried against the backend and failed: {e}"
)),
_ => execution_error(e),
}
}
fn recover_if_degraded(handle: &Arc<GraphSourceHandle>, degraded: bool) {
if !degraded {
return;
}
if degraded_reason(handle).is_none() {
return; }
tracing::info!(
"graph source recovered: an ad-hoc query answered on a degraded source — \
marking healthy (view contracts re-prove on their own scans)"
);
mark_healthy(handle);
}
pub(crate) fn cypher_batches(
handle: Arc<GraphSourceHandle>,
cypher: String,
params: Value,
columns: Arc<Vec<DeclaredColumn>>,
limit: Option<usize>,
) -> futures::stream::BoxStream<'static, DFResult<RecordBatch>> {
stream::once(async move {
let degraded_at_start = degraded_reason(&handle).is_some();
let rows = match handle
.client
.execute(&cypher, ¶ms, columns.len(), handle.bounds, limit)
.await
{
Ok(stream) => stream.try_collect::<Vec<_>>().await,
Err(e) => Err(e),
};
let rows = match rows {
Ok(rows) => {
recover_if_degraded(&handle, degraded_at_start);
rows
}
Err(e) => return Err(degraded_execution_error(degraded_reason(&handle), e)),
};
let mut batches = Vec::with_capacity(rows.len() / CONVERSION_BATCH_ROWS + 1);
for (chunk_idx, chunk) in rows.chunks(CONVERSION_BATCH_ROWS).enumerate() {
batches.push(
build_batch(&columns, chunk, chunk_idx * CONVERSION_BATCH_ROWS)
.map_err(execution_error)?,
);
}
if batches.is_empty() {
batches.push(RecordBatch::new_empty(declared_schema(&columns)));
}
Ok::<_, DataFusionError>(batches)
})
.map_ok(|batches| stream::iter(batches.into_iter().map(Ok)))
.try_flatten()
.boxed()
}
fn labels_batch(
handle: Arc<GraphSourceHandle>,
limit: Option<usize>,
) -> futures::stream::BoxStream<'static, DFResult<RecordBatch>> {
stream::once(async move {
let degraded_at_start = degraded_reason(&handle).is_some();
let labels = match handle.client.labels(handle.bounds, limit).await {
Ok(labels) => {
recover_if_degraded(&handle, degraded_at_start);
labels
}
Err(e) => return Err(degraded_execution_error(degraded_reason(&handle), e)),
};
let mut names = arrow::array::StringBuilder::new();
let mut kinds = arrow::array::StringBuilder::new();
for (name, kind) in &labels {
names.append_value(name);
kinds.append_value(kind);
}
RecordBatch::try_new(
graph_schema_schema(),
vec![Arc::new(names.finish()), Arc::new(kinds.finish())],
)
.map_err(|e| DataFusionError::ArrowError(Box::new(e), None))
})
.boxed()
}
fn execution_error(e: GraphError) -> DataFusionError {
DataFusionError::Execution(e.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use futures::stream::BoxStream;
#[derive(Debug)]
struct MockClient {
rows: Vec<Vec<Value>>,
labels: Vec<(String, String)>,
}
#[async_trait]
impl GraphClient for MockClient {
async fn execute(
&self,
_cypher: &str,
_params: &Value,
_arity: usize,
_bounds: QueryBounds,
limit: Option<usize>,
) -> Result<BoxStream<'static, Result<Vec<Value>, GraphError>>, GraphError> {
let mut rows = self.rows.clone();
if let Some(l) = limit {
rows.truncate(l);
}
Ok(stream::iter(rows.into_iter().map(Ok)).boxed())
}
async fn labels(
&self,
_bounds: QueryBounds,
limit: Option<usize>,
) -> Result<Vec<(String, String)>, GraphError> {
let mut labels = self.labels.clone();
if let Some(l) = limit {
labels.truncate(l);
}
Ok(labels)
}
}
fn sources_with(rows: Vec<Vec<Value>>) -> GraphSources {
let handle = Arc::new(GraphSourceHandle::new(
Arc::new(MockClient {
rows,
labels: vec![
("Person".to_string(), "vertex".to_string()),
("KNOWS".to_string(), "edge".to_string()),
],
}),
QueryBounds {
timeout: std::time::Duration::from_secs(5),
max_rows: 100,
},
GraphSourceHealth::Healthy,
Arc::new(vec![]),
4,
));
Arc::new(RwLock::new(HashMap::from([("kg".to_string(), handle)])))
}
async fn ctx_with(rows: Vec<Vec<Value>>) -> SessionContext {
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, sources_with(rows)).expect("registration");
ctx.register_udf((*datafusion_functions_json::udfs::json_get_str_udf()).clone());
ctx
}
async fn collect(ctx: &SessionContext, sql: &str) -> Vec<RecordBatch> {
ctx.sql(sql)
.await
.expect("plan")
.collect()
.await
.expect("collect")
}
#[derive(Debug)]
struct FailingClient {
message: String,
availability: bool,
}
impl FailingClient {
fn error(&self) -> GraphError {
if self.availability {
GraphError::Unavailable {
source_name: "kg".to_string(),
reason: self.message.clone(),
}
} else {
GraphError::backend("kg", "42601", &self.message)
}
}
}
#[async_trait]
impl GraphClient for FailingClient {
async fn execute(
&self,
_cypher: &str,
_params: &Value,
_arity: usize,
_bounds: QueryBounds,
_limit: Option<usize>,
) -> Result<BoxStream<'static, Result<Vec<Value>, GraphError>>, GraphError> {
Err(self.error())
}
async fn labels(
&self,
_bounds: QueryBounds,
_limit: Option<usize>,
) -> Result<Vec<(String, String)>, GraphError> {
Err(self.error())
}
}
fn sources_with_health(
health: GraphSourceHealth,
client: Arc<dyn GraphClient>,
) -> GraphSources {
let handle = Arc::new(GraphSourceHandle::new(
client,
QueryBounds {
timeout: std::time::Duration::from_secs(5),
max_rows: 100,
},
health,
Arc::new(vec![]),
4,
));
Arc::new(RwLock::new(HashMap::from([("kg".to_string(), handle)])))
}
fn health_of(sources: &GraphSources) -> GraphSourceHealth {
sources
.read()
.unwrap_or_else(|p| p.into_inner())
.get("kg")
.expect("kg registered")
.health
.read()
.unwrap_or_else(|p| p.into_inner())
.clone()
}
#[tokio::test]
async fn a_degraded_source_reports_the_registration_reason_not_a_bare_timeout() {
let sources = sources_with_health(
GraphSourceHealth::Degraded(
"graph backend error on 'kg' [io]: Connection refused".to_string(),
),
Arc::new(FailingClient {
message: "could not acquire a connection".to_string(),
availability: true,
}),
);
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, Arc::clone(&sources)).expect("registration");
let err = ctx
.sql(
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}')",
)
.await
.expect("plans")
.collect()
.await
.expect_err("the retried query fails");
let msg = err.to_string();
assert!(msg.contains("DEGRADED"), "{msg}");
assert!(
msg.contains("Connection refused"),
"the registration error survives: {msg}"
);
assert!(
msg.contains("could not acquire a connection"),
"the fresh failure rides along: {msg}"
);
assert!(
!msg.contains("narrow the traversal"),
"no misleading advice: {msg}"
);
assert!(!health_of(&sources).is_healthy());
let err = ctx
.sql("SELECT * FROM graph_schema('kg')")
.await
.expect("plans")
.collect()
.await
.expect_err("labels fail too");
assert!(err.to_string().contains("DEGRADED"), "{err}");
}
#[tokio::test]
async fn a_caller_error_on_a_degraded_source_keeps_its_own_headline() {
let sources = sources_with_health(
GraphSourceHealth::Degraded(
"graph backend error on 'kg' [io]: Connection refused".to_string(),
),
Arc::new(FailingClient {
message: "syntax error at or near \"RETRUN\"".to_string(),
availability: false,
}),
);
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, Arc::clone(&sources)).expect("registration");
let err = ctx
.sql(
"SELECT name FROM cypher_query('kg', 'RETRUN 1', '{}', \
'{\"name\": \"string\"}')",
)
.await
.expect("plans")
.collect()
.await
.expect_err("the backend refuses the Cypher");
let msg = err.to_string();
assert!(msg.contains("syntax error"), "{msg}");
assert!(
!msg.contains("DEGRADED"),
"a caller-caused failure is not wrapped in connectivity framing: {msg}"
);
}
#[tokio::test]
async fn a_successful_query_on_a_degraded_source_flips_it_healthy() {
let sources = sources_with_health(
GraphSourceHealth::Degraded("connection refused at startup".to_string()),
Arc::new(MockClient {
rows: vec![vec![serde_json::json!("ada")]],
labels: vec![("Person".to_string(), "vertex".to_string())],
}),
);
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, Arc::clone(&sources)).expect("registration");
let batches = collect(
&ctx,
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}')",
)
.await;
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
assert!(
health_of(&sources).is_healthy(),
"a successful retry flips the source healthy"
);
}
#[tokio::test]
async fn a_recovered_query_keeps_its_rows_when_a_sibling_view_fails_revalidation() {
let contract = super::super::view::ViewContract {
name: "people".to_string(),
cypher: "MATCH (p:Person) RETURN p.name, p.age".to_string(),
columns: vec![
DeclaredColumn {
name: "name".to_string(),
ty: GraphType::String,
nullable: true,
},
DeclaredColumn {
name: "age".to_string(),
ty: GraphType::Int,
nullable: true,
},
],
};
let handle = Arc::new(GraphSourceHandle::new(
Arc::new(MockClient {
rows: vec![vec![serde_json::json!("ada")]],
labels: vec![("Person".to_string(), "vertex".to_string())],
}),
QueryBounds {
timeout: std::time::Duration::from_secs(5),
max_rows: 100,
},
GraphSourceHealth::Degraded("connection refused at startup".to_string()),
Arc::new(vec![contract]),
4,
));
let sources: GraphSources =
Arc::new(RwLock::new(HashMap::from([("kg".to_string(), handle)])));
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, Arc::clone(&sources)).expect("registration");
let batches = collect(
&ctx,
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}')",
)
.await;
assert_eq!(
batches.iter().map(|b| b.num_rows()).sum::<usize>(),
1,
"the successful query's rows are emitted, not discarded"
);
assert!(
health_of(&sources).is_healthy(),
"the backend answered (a contract violation is an answer) — the source \
recovers; the broken view's own scans report its failure"
);
}
#[tokio::test]
async fn a_healthy_source_failure_is_not_wrapped_in_degraded_context() {
let sources = sources_with_health(
GraphSourceHealth::Healthy,
Arc::new(FailingClient {
message: "syntax error at or near".to_string(),
availability: false,
}),
);
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, sources).expect("registration");
let err = ctx
.sql(
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}')",
)
.await
.expect("plans")
.collect()
.await
.expect_err("the backend error passes through");
let msg = err.to_string();
assert!(msg.contains("syntax error at or near"), "{msg}");
assert!(
!msg.contains("DEGRADED"),
"a healthy source never had a registration error to cite: {msg}"
);
}
#[tokio::test]
async fn declared_columns_scan_end_to_end() {
let ctx = ctx_with(vec![
vec![serde_json::json!("ada"), serde_json::json!(1)],
vec![serde_json::json!("bob"), Value::Null],
])
.await;
let batches = collect(
&ctx,
"SELECT name, n FROM cypher_query('kg', \
'MATCH (p:Person) RETURN p.name, p.n', '{}', \
'{\"name\": \"string\", \"n\": \"int\"}') ORDER BY name",
)
.await;
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total, 2);
}
#[tokio::test]
async fn a_duplicate_column_declaration_is_rejected_not_silently_collapsed() {
let ctx = ctx_with(vec![]).await;
let err = ctx
.sql(
"SELECT * FROM cypher_query('kg', 'MATCH (p) RETURN p.n', '{}', \
'{\"n\": \"int\", \"n\": \"string\"}')",
)
.await
.expect_err("duplicate columns fail at plan time");
let msg = err.to_string();
assert!(msg.contains("column 'n' is declared twice"), "{msg}");
}
#[tokio::test]
async fn mutating_cypher_fails_at_plan_time_with_the_keyword() {
let ctx = ctx_with(vec![]).await;
let err = ctx
.sql(
"SELECT * FROM cypher_query('kg', 'CREATE (n) RETURN n', '{}', \
'{\"n\": \"node\"}')",
)
.await
.expect_err("plans must fail");
let msg = err.to_string();
assert!(msg.contains("'CREATE'"), "{msg}");
assert!(msg.contains("read-only"), "{msg}");
}
#[tokio::test]
async fn omitted_columns_is_a_targeted_age_error() {
let ctx = ctx_with(vec![]).await;
let err = ctx
.sql("SELECT * FROM cypher_query('kg', 'MATCH (n) RETURN n', '{}')")
.await
.expect_err("columns required on age");
let msg = err.to_string();
assert!(msg.contains("'columns' is required"), "{msg}");
assert!(msg.contains("age backend"), "{msg}");
}
#[tokio::test]
async fn unknown_connection_lists_the_known_roster() {
let ctx = ctx_with(vec![]).await;
let err = ctx
.sql(
"SELECT * FROM cypher_query('nope', 'MATCH (n) RETURN n', '{}', \
'{\"n\": \"node\"}')",
)
.await
.expect_err("unknown connection");
let msg = err.to_string();
assert!(msg.contains("'nope'"), "{msg}");
assert!(msg.contains("kg"), "the roster names what exists: {msg}");
}
#[tokio::test]
async fn unknown_type_names_the_accepted_set() {
let ctx = ctx_with(vec![]).await;
let err = ctx
.sql(
"SELECT * FROM cypher_query('kg', 'MATCH (n) RETURN n', '{}', \
'{\"n\": \"Utf8\"}')",
)
.await
.expect_err("PascalCase is not the vocabulary");
let msg = err.to_string();
assert!(msg.contains("unknown type 'Utf8'"), "{msg}");
assert!(msg.contains("node, relationship, path"), "{msg}");
}
#[tokio::test]
async fn sql_limit_pushes_to_the_consumption_side() {
let ctx = ctx_with(vec![
vec![serde_json::json!("a")],
vec![serde_json::json!("b")],
vec![serde_json::json!("c")],
])
.await;
let batches = collect(
&ctx,
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}') LIMIT 1",
)
.await;
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
}
#[tokio::test]
async fn empty_results_keep_the_declared_schema() {
let ctx = ctx_with(vec![]).await;
let batches = collect(
&ctx,
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}')",
)
.await;
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 0);
assert_eq!(batches[0].schema().field(0).name(), "name");
}
#[tokio::test]
async fn mid_scan_type_mismatch_is_typed_and_value_free() {
let ctx = ctx_with(vec![
vec![serde_json::json!(1)],
vec![serde_json::json!("secret")],
])
.await;
let err = ctx
.sql(
"SELECT n FROM cypher_query('kg', 'MATCH (x) RETURN x.n', '{}', \
'{\"n\": \"int\"}')",
)
.await
.expect("plans")
.collect()
.await
.expect_err("second row is a string");
let msg = err.to_string();
assert!(msg.contains("declared 'int'"), "{msg}");
assert!(!msg.contains("secret"), "values never leak: {msg}");
}
#[tokio::test]
async fn wrong_arg_counts_fail_with_the_signature_named() {
let ctx = ctx_with(vec![]).await;
for sql in [
"SELECT * FROM cypher_query('kg')",
"SELECT * FROM cypher_query('kg', 'RETURN 1', '{}', '{\"a\":\"int\"}', 'extra')",
"SELECT * FROM graph_schema()",
"SELECT * FROM graph_schema('kg', 'extra')",
] {
let err = ctx.sql(sql).await.expect_err(sql);
let msg = err.to_string();
assert!(
msg.contains("expects") || msg.contains("exactly"),
"{sql}: {msg}"
);
}
}
#[tokio::test]
async fn malformed_params_and_columns_fail_at_planning() {
let ctx = ctx_with(vec![]).await;
for (params, needle) in [("{not json", "unparseable JSON"), ("[1]", "an array")] {
let err = ctx
.sql(&format!(
"SELECT * FROM cypher_query('kg', 'MATCH (n) RETURN n', '{params}', \
'{{\"n\": \"node\"}}')"
))
.await
.expect_err(params);
assert!(err.to_string().contains(needle), "{params}: {err}");
}
for (columns, needle) in [
("{not json", "unparseable JSON"),
("[1]", "expected a JSON object"),
("{}", "at least one column"),
("{\"n\": 7}", "type must be a string"),
] {
let err = ctx
.sql(&format!(
"SELECT * FROM cypher_query('kg', 'MATCH (n) RETURN n', '{{}}', '{columns}')"
))
.await
.expect_err(columns);
assert!(err.to_string().contains(needle), "{columns}: {err}");
}
}
#[tokio::test]
async fn projection_prunes_columns_and_explain_renders_the_plan() {
let ctx = ctx_with(vec![vec![serde_json::json!("ada"), serde_json::json!(1)]]).await;
let batches = collect(
&ctx,
"SELECT n FROM cypher_query('kg', 'MATCH (p) RETURN p.name, p.n', '{}', \
'{\"name\": \"string\", \"n\": \"int\"}')",
)
.await;
assert_eq!(batches[0].num_columns(), 1);
assert_eq!(batches[0].schema().field(0).name(), "n");
let plan = collect(
&ctx,
"EXPLAIN SELECT n FROM cypher_query('kg', 'MATCH (p) RETURN p.n', '{}', \
'{\"n\": \"int\"}')",
)
.await;
assert!(!plan.is_empty());
let plan = collect(&ctx, "EXPLAIN SELECT * FROM graph_schema('kg')").await;
assert!(!plan.is_empty());
}
#[tokio::test]
async fn an_unprojected_declared_column_still_fails_its_type_contract() {
let ctx = ctx_with(vec![vec![
serde_json::json!("ada"),
serde_json::json!("not-an-int"),
]])
.await;
let err = ctx
.sql(
"SELECT a FROM cypher_query('kg', 'MATCH (x) RETURN x.a, x.b', '{}', \
'{\"a\": \"string\", \"b\": \"int\"}')",
)
.await
.expect("plans")
.collect()
.await
.expect_err("the unprojected column's violated declaration is still loud");
let msg = err.to_string();
assert!(msg.contains("'b'"), "{msg}");
assert!(msg.contains("declared 'int'"), "{msg}");
}
#[tokio::test]
async fn graph_schema_respects_a_sql_limit() {
let ctx = ctx_with(vec![]).await;
let batches = collect(&ctx, "SELECT label FROM graph_schema('kg') LIMIT 1").await;
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
}
#[tokio::test]
async fn graph_schema_lists_labels_and_kinds() {
let ctx = ctx_with(vec![]).await;
let batches = collect(
&ctx,
"SELECT label, kind FROM graph_schema('kg') ORDER BY label",
)
.await;
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total, 2);
}
#[tokio::test]
async fn json_get_str_extracts_properties_when_the_session_registers_it() {
let node = serde_json::json!({
"id": 1, "label": "Person", "properties": {"name": "ada"}
});
let ctx = ctx_with(vec![vec![node]]).await;
let batches = collect(
&ctx,
"SELECT json_get_str(v.properties, 'name') AS name FROM \
cypher_query('kg', 'MATCH (v) RETURN v', '{}', '{\"v\": \"node\"}')",
)
.await;
let col = batches[0]
.column(0)
.as_any()
.downcast_ref::<arrow::array::StringArray>()
.unwrap();
assert_eq!(col.value(0), "ada");
}
#[tokio::test]
async fn a_null_params_placeholder_plans_as_no_params() {
let ctx = ctx_with(vec![vec![serde_json::json!("ada")]]).await;
let batches = collect(
&ctx,
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', NULL, \
'{\"name\": \"string\"}')",
)
.await;
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
for sql in [
"SELECT * FROM cypher_query(NULL, 'MATCH (n) RETURN n', '{}', '{\"n\": \"node\"}')",
"SELECT * FROM cypher_query('kg', NULL, '{}', '{\"n\": \"node\"}')",
"SELECT * FROM cypher_query('kg', 'MATCH (n) RETURN n', '{}', NULL)",
] {
let err = ctx.sql(sql).await.expect_err(sql);
assert!(err.to_string().contains("not NULL"), "{sql}: {err}");
}
}
#[tokio::test]
async fn an_unknown_source_on_an_empty_registry_says_none() {
let sources: GraphSources = Arc::new(RwLock::new(HashMap::new()));
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, sources).expect("registration");
let err = ctx
.sql(
"SELECT * FROM cypher_query('nope', 'RETURN 1', '{}', \
'{\"n\": \"int\"}')",
)
.await
.expect_err("an unknown source fails at plan time");
let msg = err.to_string();
assert!(msg.contains("nope"), "{msg}");
assert!(msg.contains("none"), "the empty registry says so: {msg}");
}
#[tokio::test]
async fn a_stale_degraded_flag_rechecks_and_returns() {
let sources = sources_with(vec![]);
let handle = Arc::clone(
sources
.read()
.unwrap_or_else(|p| p.into_inner())
.get("kg")
.expect("kg registered"),
);
let stamped = *handle
.health_changed_at
.lock()
.unwrap_or_else(|p| p.into_inner());
recover_if_degraded(&handle, true);
assert!(
health_of(&sources).is_healthy(),
"a healthy source stays healthy through the stale-flag path"
);
assert_eq!(
stamped,
*handle
.health_changed_at
.lock()
.unwrap_or_else(|p| p.into_inner()),
"no transition was re-stamped"
);
}
#[tokio::test]
async fn an_armed_backoff_does_not_outweigh_a_successful_answer() {
let sources = sources_with_health(
GraphSourceHealth::Degraded("connection refused at startup".to_string()),
Arc::new(MockClient {
rows: vec![vec![serde_json::json!("ada")]],
labels: vec![("Person".to_string(), "vertex".to_string())],
}),
);
let handle = Arc::clone(
sources
.read()
.unwrap_or_else(|p| p.into_inner())
.get("kg")
.expect("kg registered"),
);
super::super::view::arm_recovery_backoff(&handle);
let ctx = SessionContext::new();
register_graph_udtfs(&ctx, Arc::clone(&sources)).expect("registration");
let batches = collect(
&ctx,
"SELECT name FROM cypher_query('kg', 'MATCH (p) RETURN p.name', '{}', \
'{\"name\": \"string\"}')",
)
.await;
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
assert!(
health_of(&sources).is_healthy(),
"the successful answer beats the armed backoff"
);
}
#[tokio::test]
async fn the_leaf_plan_pins_its_contract_and_redacts_cypher() {
let sources = sources_with(vec![vec![serde_json::json!("ada")]]);
let handle = Arc::clone(
sources
.read()
.unwrap_or_else(|p| p.into_inner())
.get("kg")
.expect("kg registered"),
);
let columns = Arc::new(vec![DeclaredColumn {
name: "name".to_string(),
ty: GraphType::String,
nullable: true,
}]);
let secret = "MATCH (creds {token: 'hunter2'}) RETURN creds.name";
let kinds = [
GraphScanKind::Cypher {
handle: Arc::clone(&handle),
cypher: secret.to_string(),
params: serde_json::json!({}),
columns: Arc::clone(&columns),
limit: None,
},
GraphScanKind::View {
handle: Arc::clone(&handle),
view_name: "people".to_string(),
cypher: secret.to_string(),
columns: Arc::clone(&columns),
limit: None,
},
GraphScanKind::Labels {
handle: Arc::clone(&handle),
limit: None,
},
];
for kind in &kinds {
let dbg = format!("{kind:?}");
assert!(
!dbg.contains("hunter2") && !dbg.contains("MATCH"),
"the Cypher text never appears in Debug: {dbg}"
);
}
let [_, view_kind, _] = kinds;
assert!(format!("{view_kind:?}").contains("people"));
let schema = declared_schema(&columns);
let exec = Arc::new(
GraphScanExec::new(view_kind, Arc::clone(&schema), None).expect("plan builds"),
);
assert_eq!(exec.schema(), schema);
let display = format!(
"{}",
datafusion::physical_plan::displayable(exec.as_ref()).one_line()
);
assert!(display.contains("view people"), "{display}");
assert!(!display.contains("hunter2"), "{display}");
let err = Arc::clone(&exec)
.with_new_children(vec![Arc::clone(&exec) as Arc<dyn ExecutionPlan>])
.expect_err("a leaf takes no children");
assert!(err.to_string().contains("leaf"), "{err}");
let err = match exec.execute(1, Arc::new(TaskContext::default())) {
Ok(_) => panic!("only partition 0 exists"),
Err(e) => e,
};
assert!(err.to_string().contains("partition 1"), "{err}");
let stream = exec
.execute(0, Arc::new(TaskContext::default()))
.expect("partition 0 executes");
let batches: Vec<RecordBatch> = stream.try_collect().await.expect("collects");
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
}
}