use axum::{
extract::{Path, State},
response::IntoResponse,
Json,
};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use utoipa::ToSchema;
use velesdb_core::api_types::serde_id;
use velesdb_core::Error;
use crate::handlers::helpers::auto_core_error_response;
use crate::types::{ErrorResponse, VELESQL_CONTRACT_VERSION};
use crate::AppState;
#[derive(Debug, Deserialize, ToSchema)]
pub struct MatchQueryRequest {
pub query: String,
#[serde(default)]
pub params: HashMap<String, serde_json::Value>,
#[serde(default)]
pub vector: Option<Vec<f32>>,
#[serde(default)]
pub threshold: Option<f32>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct MatchQueryResultItem {
#[serde(serialize_with = "serde_id::serialize_id_map_as_strings")]
#[cfg_attr(feature = "openapi", schema(schema_with = serde_id::id_map_schema))]
pub bindings: HashMap<String, u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub score: Option<f32>,
pub depth: u32,
#[serde(skip_serializing_if = "HashMap::is_empty")]
pub projected: HashMap<String, serde_json::Value>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct MatchQueryResponse {
pub results: Vec<MatchQueryResultItem>,
pub took_ms: u64,
pub count: usize,
pub meta: MatchQueryMeta,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct MatchQueryMeta {
pub velesql_contract_version: String,
}
#[utoipa::path(
post,
path = "/collections/{name}/match",
tag = "graph",
params(("name" = String, Path, description = "Collection name")),
request_body = MatchQueryRequest,
responses(
(status = 200, description = "Match query results", body = MatchQueryResponse),
(status = 400, description = "Parse error or invalid query", body = ErrorResponse),
(status = 404, description = "Collection not found", body = ErrorResponse),
(status = 500, description = "Internal server error", body = ErrorResponse)
)
)]
pub async fn match_query(
Path(collection_name): Path<String>,
State(state): State<Arc<AppState>>,
Json(request): Json<MatchQueryRequest>,
) -> axum::response::Response {
let state_clone = Arc::clone(&state);
let outcome = crate::handlers::helpers::run_blocking(move || {
run_match(&state_clone, &collection_name, &request)
})
.await;
match outcome {
Ok(Ok(response)) => Json(response).into_response(),
Ok(Err(e)) => auto_core_error_response(&e),
Err(resp) => resp,
}
}
fn run_match(
state: &AppState,
collection_name: &str,
request: &MatchQueryRequest,
) -> Result<MatchQueryResponse, Error> {
let start = std::time::Instant::now();
let collection = resolve_match_collection(state, collection_name)
.ok_or_else(|| Error::CollectionNotFound(collection_name.to_string()))?;
let match_clause = parse_match_clause(&request.query)?;
validate_threshold(request.threshold)?;
if state
.db
.authorize_read(
collection_name,
velesdb_core::observer::QueryOperationKind::GraphTraversal,
None,
None,
)?
.is_some()
{
return Err(Error::Config(
"scope narrowing is not supported for MATCH queries".to_string(),
));
}
let results = execute_match(&collection, &match_clause, request)?;
let count = results.len();
#[allow(clippy::cast_possible_truncation)]
let took_ms = start.elapsed().as_millis() as u64;
Ok(MatchQueryResponse {
results,
took_ms,
count,
meta: MatchQueryMeta {
velesql_contract_version: VELESQL_CONTRACT_VERSION.to_string(),
},
})
}
fn parse_match_clause(query_str: &str) -> Result<velesdb_core::velesql::MatchClause, Error> {
let query = velesdb_core::velesql::Parser::parse(query_str)?;
query.match_clause.ok_or_else(|| {
Error::Query(
"Query is not a MATCH query. Use MATCH (...) RETURN ... \
or call /query for SELECT statements."
.to_string(),
)
})
}
fn validate_threshold(threshold: Option<f32>) -> Result<(), Error> {
if let Some(t) = threshold {
if !(0.0..=1.0).contains(&t) {
return Err(Error::Query(format!(
"Invalid threshold: {t}. Must be between 0.0 and 1.0"
)));
}
}
Ok(())
}
enum MatchCollection {
Vector(velesdb_core::collection::VectorCollection),
Graph(velesdb_core::collection::GraphCollection),
}
fn resolve_match_collection(state: &AppState, name: &str) -> Option<MatchCollection> {
state
.db
.get_vector_collection(name)
.map(MatchCollection::Vector)
.or_else(|| {
state
.db
.get_graph_collection(name)
.map(MatchCollection::Graph)
})
}
fn execute_match(
collection: &MatchCollection,
match_clause: &velesdb_core::velesql::MatchClause,
request: &MatchQueryRequest,
) -> Result<Vec<MatchQueryResultItem>, Error> {
let raw_results = if let Some(ref vector) = request.vector {
let threshold = request.threshold.unwrap_or(0.0);
match collection {
MatchCollection::Vector(coll) => {
coll.execute_match_with_similarity(match_clause, vector, threshold, &request.params)
}
MatchCollection::Graph(coll) => {
coll.execute_match_with_similarity(match_clause, vector, threshold, &request.params)
}
}
} else {
match collection {
MatchCollection::Vector(coll) => coll.execute_match(match_clause, &request.params),
MatchCollection::Graph(coll) => coll.execute_match(match_clause, &request.params),
}
};
raw_results.map(|results| {
results
.into_iter()
.map(|r| MatchQueryResultItem {
bindings: r.bindings,
score: r.score,
depth: r.depth,
projected: r.projected,
})
.collect()
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_match_query_request_deserialize() {
let json = r#"{
"query": "MATCH (a:Person)-[:KNOWS]->(b) RETURN a.name",
"params": {}
}"#;
let request: MatchQueryRequest = serde_json::from_str(json).unwrap();
assert!(request.query.contains("MATCH"));
assert!(request.params.is_empty());
}
#[test]
fn test_match_query_response_serialize() {
let response = MatchQueryResponse {
results: vec![MatchQueryResultItem {
bindings: HashMap::from([("a".to_string(), 123)]),
score: Some(0.95),
depth: 1,
projected: HashMap::new(),
}],
took_ms: 15,
count: 1,
meta: MatchQueryMeta {
velesql_contract_version: VELESQL_CONTRACT_VERSION.to_string(),
},
};
let json = serde_json::to_string(&response).unwrap();
assert!(json.contains("bindings"));
assert!(json.contains("0.95"));
}
#[test]
fn test_match_query_bindings_serialized_as_strings() {
let above_safe = (1_u64 << 53) + 1; let response = MatchQueryResponse {
results: vec![MatchQueryResultItem {
bindings: HashMap::from([("a".to_string(), above_safe)]),
score: None,
depth: 0,
projected: HashMap::new(),
}],
took_ms: 0,
count: 1,
meta: MatchQueryMeta {
velesql_contract_version: VELESQL_CONTRACT_VERSION.to_string(),
},
};
let json = serde_json::to_value(&response).unwrap();
assert_eq!(
json["results"][0]["bindings"]["a"],
serde_json::json!("9007199254740993"),
"binding IDs must serialize as JSON strings for JS precision safety"
);
}
#[test]
fn test_match_query_response_with_projected_properties() {
let mut projected = HashMap::new();
projected.insert("author.name".to_string(), serde_json::json!("John Doe"));
let response = MatchQueryResponse {
results: vec![MatchQueryResultItem {
bindings: HashMap::from([("author".to_string(), 42)]),
score: Some(0.92),
depth: 1,
projected,
}],
took_ms: 10,
count: 1,
meta: MatchQueryMeta {
velesql_contract_version: VELESQL_CONTRACT_VERSION.to_string(),
},
};
let json = serde_json::to_string(&response).unwrap();
assert!(json.contains("John Doe"));
assert!(json.contains("author.name"));
}
#[test]
fn test_match_handler_applies_return_order_by() {
use velesdb_core::collection::VectorCollection;
use velesdb_core::{DistanceMetric, Point, StorageMode};
let temp = tempfile::tempdir().expect("temp dir");
let coll = VectorCollection::create(
temp.path().to_path_buf(),
"people",
4,
DistanceMetric::Cosine,
StorageMode::default(),
)
.expect("create collection");
let ages = [(1_u64, 30), (2, 10), (3, 50), (4, 20), (5, 40)];
let points: Vec<Point> = ages
.iter()
.map(|(id, age)| {
Point::new(
*id,
vec![1.0, 0.0, 0.0, 0.0],
Some(serde_json::json!({"_labels": ["Person"], "age": age})),
)
})
.collect();
coll.upsert(points).expect("upsert Person nodes");
let collection = MatchCollection::Vector(coll);
let request = MatchQueryRequest {
query: "MATCH (n:Person) RETURN n ORDER BY n.age DESC LIMIT 10".to_string(),
params: HashMap::new(),
vector: None,
threshold: None,
};
let clause = parse_match_clause(&request.query).expect("parse MATCH clause");
let results = execute_match(&collection, &clause, &request).expect("execute_match");
let ids: Vec<u64> = results
.iter()
.map(|r| *r.bindings.get("n").expect("binding 'n'"))
.collect();
assert_eq!(
ids,
vec![3, 5, 1, 4, 2],
"/match must honor RETURN ORDER BY n.age DESC (ages 50,40,30,20,10)"
);
}
}