use std::collections::HashMap;
use graphql_tools::{
ast::{OperationVisitor, OperationVisitorContext},
static_graphql::query::{Definition, Document, FragmentDefinition},
validation::{
rules::{ValidationRule, ValidationVisitor},
utils::{ValidationError, ValidationErrorContext},
},
};
use hive_router_config::limits::MaxAliasesRuleConfig;
use crate::pipeline::validation::shared::{CountableNode, VisitedFragment};
pub struct MaxAliasesRule {
pub config: MaxAliasesRuleConfig,
}
impl ValidationRule for MaxAliasesRule {
fn error_code(&self) -> &'static str {
"MAX_ALIASES_EXCEEDED"
}
fn visitor<'doc>(&self) -> ValidationVisitor<'doc> {
Box::new(MaxAliasesVisitor {
config: self.config.clone(),
visited_fragments: HashMap::new(),
})
}
}
struct MaxAliasesVisitor<'doc> {
config: MaxAliasesRuleConfig,
visited_fragments: HashMap<&'doc str, VisitedFragment>,
}
impl<'doc> MaxAliasesVisitor<'doc> {
fn check_limit(&self, count: usize) -> Result<usize, ValidationError> {
if count > self.config.n {
Err(ValidationError {
locations: vec![],
message: "Aliases limit exceeded.".to_string(),
error_code: "MAX_ALIASES_EXCEEDED",
})
} else {
Ok(count)
}
}
fn count_aliases(
&mut self,
known_fragments: &HashMap<&'doc str, &'doc FragmentDefinition>,
countable_node: CountableNode<'doc>,
) -> Result<usize, ValidationError> {
let mut alias_count: usize = 0;
if let CountableNode::Field(field) = countable_node {
if field.alias.is_some() {
alias_count = self.check_limit(alias_count + 1)?;
}
}
if let Some(selection_set) = countable_node.selection_set() {
for selection in &selection_set.items {
let countable_node: CountableNode<'doc> = selection.into();
let child_aliases = self.count_aliases(known_fragments, countable_node)?;
alias_count = self.check_limit(alias_count + child_aliases)?;
}
}
if let CountableNode::FragmentSpread(node) = countable_node {
let fragment_name = node.fragment_name.as_str();
match self.visited_fragments.get(fragment_name) {
Some(VisitedFragment::Counted(num)) => {
return self.check_limit(alias_count + num);
}
Some(VisitedFragment::Visiting) => return Ok(alias_count),
None => {}
}
self.visited_fragments
.insert(fragment_name, VisitedFragment::Visiting);
if let Some(fragment_def) = known_fragments.get(fragment_name).copied() {
let countable_node: CountableNode<'doc> =
CountableNode::FragmentDefinition(fragment_def);
let fragment_alias_count = self.count_aliases(known_fragments, countable_node)?;
self.visited_fragments.insert(
fragment_name,
VisitedFragment::Counted(fragment_alias_count),
);
alias_count = self.check_limit(alias_count + fragment_alias_count)?;
}
}
Ok(alias_count)
}
}
impl<'doc> OperationVisitor<'doc, ValidationErrorContext> for MaxAliasesVisitor<'doc> {
fn enter_document(
&mut self,
context: &mut OperationVisitorContext<'doc>,
user_context: &mut ValidationErrorContext,
document: &'doc Document,
) {
self.visited_fragments = HashMap::with_capacity(context.known_fragments.len());
for definition in &document.definitions {
let Definition::Operation(op) = definition else {
continue;
};
if let Err(err) = self.count_aliases(&context.known_fragments, op.into()) {
user_context.report_error(err);
}
}
}
}
#[cfg(test)]
mod tests {
use std::vec;
use graphql_tools::{
parser::parse_schema,
validation::validate::{validate, ValidationPlan},
};
use hive_router_config::limits::MaxAliasesRuleConfig;
use crate::pipeline::validation::max_aliases_rule::MaxAliasesRule;
const TYPE_DEFS: &'static str = r#"
type Book {
title: String
author: String
}
type Query {
books: [Book]
getBook(title: String): Book
}
"#;
const QUERY: &'static str = r#"
query {
firstBooks: getBook(title: "null") {
author
title
}
secondBooks: getBook(title: "null") {
author
title
}
}
"#;
#[test]
fn should_work_by_default() {
let schema = parse_schema(TYPE_DEFS)
.expect("Failed to parse schema")
.into_static();
let query = graphql_tools::parser::parse_query(QUERY)
.expect("Failed to parse query")
.into_static();
let validation_plan = ValidationPlan::from(vec![Box::new(MaxAliasesRule {
config: MaxAliasesRuleConfig { n: 15 },
})]);
let errors = validate(&schema, &query, &validation_plan);
assert!(errors.is_empty());
}
#[test]
fn rejects_query_exceeding_max_aliases() {
let schema = parse_schema(TYPE_DEFS)
.expect("Failed to parse schema")
.into_static();
let query = graphql_tools::parser::parse_query(QUERY)
.expect("Failed to parse query")
.into_static();
let validation_plan = ValidationPlan::from(vec![Box::new(MaxAliasesRule {
config: MaxAliasesRuleConfig { n: 1 },
})]);
let errors = validate(&schema, &query, &validation_plan);
assert_eq!(errors.len(), 1);
assert_eq!(errors[0].error_code, "MAX_ALIASES_EXCEEDED");
}
#[test]
fn respects_fragment_aliases() {
let schema = parse_schema(TYPE_DEFS)
.expect("Failed to parse schema")
.into_static();
let query = graphql_tools::parser::parse_query(
r#"
query A {
getBook(title: "null") {
firstTitle: title
...BookFragment
}
}
fragment BookFragment on Book {
secondTitle: title
}
"#,
)
.expect("Failed to parse query")
.into_static();
let validation_plan = ValidationPlan::from(vec![Box::new(MaxAliasesRule {
config: MaxAliasesRuleConfig { n: 1 },
})]);
let errors = validate(&schema, &query, &validation_plan);
assert_eq!(errors.len(), 1);
assert_eq!(errors[0].error_code, "MAX_ALIASES_EXCEEDED");
}
#[test]
fn do_not_crash_on_recursive_fragment() {
let schema = parse_schema(TYPE_DEFS)
.expect("Failed to parse schema")
.into_static();
let query = graphql_tools::parser::parse_query(
r#"
query {
...A
}
fragment A on Query {
...B
}
fragment B on Query {
...A
}
"#,
)
.expect("Failed to parse query")
.into_static();
let validation_plan = ValidationPlan::from(vec![Box::new(MaxAliasesRule {
config: MaxAliasesRuleConfig { n: 10 },
})]);
let errors = validate(&schema, &query, &validation_plan);
assert!(errors.is_empty());
}
}