use std::collections::HashMap;
use graphql_parser::query::{
Definition, Document, Field, OperationDefinition, Selection, SelectionSet, Value,
};
use crate::error::{NodeCountError, PageSizeArgument, Position};
use crate::Variables;
const PAGE_SIZE_ARGUMENTS: [PageSizeArgument; 2] =
[PageSizeArgument::First, PageSizeArgument::Last];
const PAGE_SIZE_RANGE: std::ops::RangeInclusive<i64> = 1..=100;
const ANONYMOUS: &str = "<anonymous>";
type Text<'a> = &'a str;
fn operation_parts<'a>(
operation: &'a OperationDefinition<'a, Text<'a>>,
) -> (Option<&'a str>, &'a SelectionSet<'a, Text<'a>>) {
match operation {
OperationDefinition::SelectionSet(set) => (None, set),
OperationDefinition::Query(query) => (query.name, &query.selection_set),
OperationDefinition::Mutation(mutation) => (mutation.name, &mutation.selection_set),
OperationDefinition::Subscription(sub) => (sub.name, &sub.selection_set),
}
}
fn field_label<'a>(field: &Field<'a, Text<'a>>) -> String {
match field.alias {
Some(alias) => format!("{alias}:{}", field.name),
None => field.name.to_string(),
}
}
fn selection_label<'a>(selection: &Selection<'a, Text<'a>>) -> (String, Position) {
match selection {
Selection::Field(field) => (field_label(field), field.position.into()),
Selection::FragmentSpread(spread) => (
format!("...{}", spread.fragment_name),
spread.position.into(),
),
Selection::InlineFragment(inline) => {
("... (inline fragment)".to_string(), inline.position.into())
}
}
}
fn checked(
value: Option<u64>,
attribution: impl FnOnce() -> (String, Position),
) -> Result<u64, NodeCountError> {
value.ok_or_else(|| {
let (field, position) = attribution();
NodeCountError::Overflow { field, position }
})
}
struct Counter<'a> {
fragments: HashMap<&'a str, &'a SelectionSet<'a, Text<'a>>>,
variables: &'a Variables,
open_fragments: Vec<&'a str>,
}
impl<'a> Counter<'a> {
fn selection_set(
&mut self,
set: &'a SelectionSet<'a, Text<'a>>,
multiplier: u64,
) -> Result<u64, NodeCountError> {
let mut total: u64 = 0;
for selection in &set.items {
let contribution = self.selection(selection, multiplier)?;
total = checked(total.checked_add(contribution), || {
selection_label(selection)
})?;
}
Ok(total)
}
fn selection(
&mut self,
selection: &'a Selection<'a, Text<'a>>,
multiplier: u64,
) -> Result<u64, NodeCountError> {
match selection {
Selection::Field(field) => self.field(field, multiplier),
Selection::InlineFragment(inline) => {
self.selection_set(&inline.selection_set, multiplier)
}
Selection::FragmentSpread(spread) => {
let name = spread.fragment_name;
let position = Position::from(spread.position);
let set = *self
.fragments
.get(name)
.ok_or(NodeCountError::UndefinedFragment {
name: name.to_string(),
position,
})?;
if self.open_fragments.contains(&name) {
return Err(NodeCountError::FragmentCycle {
name: name.to_string(),
position,
});
}
self.open_fragments.push(name);
let total = self.selection_set(set, multiplier);
self.open_fragments.pop();
total
}
}
}
fn field(
&mut self,
field: &'a Field<'a, Text<'a>>,
multiplier: u64,
) -> Result<u64, NodeCountError> {
let Some(page_size) = self.page_size(field)? else {
return self.selection_set(&field.selection_set, multiplier);
};
let label = || (field_label(field), field.position.into());
let nodes = checked(multiplier.checked_mul(u64::from(page_size)), label)?;
let nested = self.selection_set(&field.selection_set, nodes)?;
checked(nodes.checked_add(nested), label)
}
fn page_size(&self, field: &Field<'a, Text<'a>>) -> Result<Option<u32>, NodeCountError> {
let mut largest: Option<u32> = None;
for (name, value) in &field.arguments {
let Some(argument) = PAGE_SIZE_ARGUMENTS
.into_iter()
.find(|candidate| candidate.as_str() == *name)
else {
continue;
};
let resolved = self.resolve_page_size(field, argument, value)?;
largest = Some(largest.map_or(resolved, |seen: u32| seen.max(resolved)));
}
Ok(largest)
}
fn resolve_page_size(
&self,
field: &Field<'a, Text<'a>>,
argument: PageSizeArgument,
value: &Value<'a, Text<'a>>,
) -> Result<u32, NodeCountError> {
let position = Position::from(field.position);
let raw: i64 = match value {
Value::Int(number) => number.as_i64().unwrap_or(i64::MAX),
Value::Variable(name) => i64::from(*self.variables.get(*name).ok_or_else(|| {
NodeCountError::UnboundVariable {
field: field_label(field),
argument,
variable: (*name).to_string(),
position,
}
})?),
other => {
return Err(NodeCountError::PageSizeNotAnInteger {
field: field_label(field),
argument,
found: other.to_string(),
position,
})
}
};
if !PAGE_SIZE_RANGE.contains(&raw) {
return Err(NodeCountError::PageSizeOutOfRange {
field: field_label(field),
argument,
value: raw,
position,
});
}
Ok(raw as u32)
}
}
pub(crate) fn count(document: &str, variables: &Variables) -> Result<u64, NodeCountError> {
let parsed: Document<'_, Text<'_>> =
graphql_parser::parse_query(document).map_err(|error| NodeCountError::Parse {
message: error.to_string(),
})?;
let mut fragments = HashMap::new();
let mut operations = Vec::new();
for definition in &parsed.definitions {
match definition {
Definition::Operation(operation) => operations.push(operation),
Definition::Fragment(fragment) => {
fragments.insert(fragment.name, &fragment.selection_set);
}
}
}
let operation = match operations.as_slice() {
[] => return Err(NodeCountError::NoOperation),
[only] => *only,
many => {
return Err(NodeCountError::MultipleOperations {
names: many
.iter()
.map(|operation| {
operation_parts(operation)
.0
.unwrap_or(ANONYMOUS)
.to_string()
})
.collect(),
})
}
};
let mut counter = Counter {
fragments,
variables,
open_fragments: Vec::new(),
};
counter.selection_set(operation_parts(operation).1, 1)
}