use std::borrow::Cow;
use std::collections::HashSet;
use saphyr_parser::{Event, Parser, ScanError};
fn fallback_budget_input(input: &str) -> Option<&str> {
let mut offset = 0;
for line in input.split_inclusive('\n') {
let line_end = offset + line.len();
let trimmed = line.trim_end_matches(['\r', '\n']);
if let Some(rest) = trimmed.strip_prefix("...") {
let rest = rest.trim_start_matches([' ', '\t']);
if rest.is_empty() || rest.starts_with('#') {
let tail = &input[line_end..];
let next = tail.trim_start_matches([' ', '\t', '\r', '\n']);
if next.is_empty() || next.starts_with("---") {
return None;
}
return Some(&input[..line_end]);
}
}
offset = line_end;
}
None
}
#[derive(Clone, Debug)]
pub struct Budget {
pub max_events: usize,
pub max_aliases: usize,
pub max_anchors: usize,
pub max_depth: usize,
pub max_documents: usize,
pub max_nodes: usize,
pub max_total_scalar_bytes: usize,
pub enforce_alias_anchor_ratio: bool,
pub alias_anchor_min_aliases: usize,
pub alias_anchor_ratio_multiplier: usize,
}
impl Default for Budget {
fn default() -> Self {
Self {
max_events: 1_000_000, max_aliases: 50_000, max_anchors: 50_000,
max_depth: 2_000, max_documents: 1_024, max_nodes: 250_000, max_total_scalar_bytes: 64 * 1024 * 1024, enforce_alias_anchor_ratio: true,
alias_anchor_min_aliases: 100,
alias_anchor_ratio_multiplier: 10,
}
}
}
#[derive(Clone, Debug)]
pub enum BudgetBreach {
Events {
events: usize,
},
Aliases {
aliases: usize,
},
Anchors {
anchors: usize,
},
Depth {
depth: usize,
},
Documents {
documents: usize,
},
Nodes {
nodes: usize,
},
ScalarBytes {
total_scalar_bytes: usize,
},
AliasAnchorRatio {
aliases: usize,
anchors: usize,
},
SequenceUnbalanced,
}
#[derive(Clone, Debug, Default)]
pub struct BudgetReport {
pub breached: Option<BudgetBreach>,
pub events: usize,
pub aliases: usize,
pub anchors: usize,
pub documents: usize,
pub nodes: usize,
pub max_depth: usize,
pub total_scalar_bytes: usize,
}
#[derive(Clone, Copy, Debug)]
pub(crate) enum BudgetEvent<'a, A> {
StreamStart,
StreamEnd,
DocumentStart,
DocumentEnd,
Alias(A),
Scalar {
anchor: Option<A>,
scalar_bytes: usize,
},
SequenceStart {
anchor: Option<A>,
},
SequenceEnd,
MappingStart {
anchor: Option<A>,
},
MappingEnd,
Nothing,
#[allow(dead_code)]
_Marker(std::marker::PhantomData<&'a A>),
}
#[derive(Clone, Debug)]
struct BudgetScope<A> {
aliases: usize,
nodes: usize,
max_depth: usize,
total_scalar_bytes: usize,
anchors: HashSet<A>,
}
impl<A> Default for BudgetScope<A> {
fn default() -> Self {
Self {
aliases: 0,
nodes: 0,
max_depth: 0,
total_scalar_bytes: 0,
anchors: HashSet::new(),
}
}
}
#[derive(Clone, Debug)]
pub(crate) struct BudgetTracker<A> {
budget: Budget,
report: BudgetReport,
depth: usize,
scope: BudgetScope<A>,
}
impl<A> BudgetTracker<A>
where
A: Eq + std::hash::Hash,
{
pub(crate) fn new(budget: &Budget) -> Self {
Self {
budget: budget.clone(),
report: BudgetReport::default(),
depth: 0,
scope: BudgetScope {
anchors: HashSet::with_capacity(256),
..BudgetScope::default()
},
}
}
fn breach(&mut self, breach: BudgetBreach) -> BudgetBreach {
self.report.breached = Some(breach.clone());
breach
}
fn check_alias_anchor_ratio(&mut self) -> Result<(), BudgetBreach> {
if self.budget.enforce_alias_anchor_ratio
&& self.scope.aliases >= self.budget.alias_anchor_min_aliases
{
let anchors = self.scope.anchors.len();
if anchors == 0 || self.scope.aliases > self.budget.alias_anchor_ratio_multiplier * anchors
{
return Err(self.breach(BudgetBreach::AliasAnchorRatio {
aliases: self.scope.aliases,
anchors,
}));
}
}
Ok(())
}
fn finish_document(&mut self) -> Result<(), BudgetBreach> {
self.check_alias_anchor_ratio()?;
self.report.aliases = self.scope.aliases;
self.report.anchors = self.scope.anchors.len();
self.report.nodes = self.scope.nodes;
self.report.max_depth = self.scope.max_depth;
self.report.total_scalar_bytes = self.scope.total_scalar_bytes;
Ok(())
}
pub(crate) fn observe<'a>(&mut self, event: BudgetEvent<'a, A>) -> Result<(), BudgetBreach> {
self.report.events += 1;
if self.report.events > self.budget.max_events {
return Err(self.breach(BudgetBreach::Events {
events: self.report.events,
}));
}
match event {
BudgetEvent::StreamStart | BudgetEvent::StreamEnd | BudgetEvent::Nothing => {}
BudgetEvent::DocumentStart => {
self.report.documents += 1;
if self.report.documents > self.budget.max_documents {
return Err(self.breach(BudgetBreach::Documents {
documents: self.report.documents,
}));
}
}
BudgetEvent::DocumentEnd => {
self.finish_document()?;
}
BudgetEvent::Alias(_anchor) => {
self.scope.aliases += 1;
self.report.aliases = self.scope.aliases;
if self.scope.aliases > self.budget.max_aliases {
return Err(self.breach(BudgetBreach::Aliases {
aliases: self.scope.aliases,
}));
}
}
BudgetEvent::Scalar {
anchor,
scalar_bytes,
} => {
self.scope.nodes += 1;
self.report.nodes = self.scope.nodes;
if self.scope.nodes > self.budget.max_nodes {
return Err(self.breach(BudgetBreach::Nodes {
nodes: self.scope.nodes,
}));
}
self.scope.total_scalar_bytes =
self.scope.total_scalar_bytes.saturating_add(scalar_bytes);
self.report.total_scalar_bytes = self.scope.total_scalar_bytes;
if self.scope.total_scalar_bytes > self.budget.max_total_scalar_bytes {
return Err(self.breach(BudgetBreach::ScalarBytes {
total_scalar_bytes: self.scope.total_scalar_bytes,
}));
}
if let Some(anchor) = anchor {
if self.scope.anchors.insert(anchor) {
self.report.anchors = self.scope.anchors.len();
if self.scope.anchors.len() > self.budget.max_anchors {
return Err(self.breach(BudgetBreach::Anchors {
anchors: self.scope.anchors.len(),
}));
}
}
}
}
BudgetEvent::SequenceStart { anchor } | BudgetEvent::MappingStart { anchor } => {
self.scope.nodes += 1;
self.report.nodes = self.scope.nodes;
if self.scope.nodes > self.budget.max_nodes {
return Err(self.breach(BudgetBreach::Nodes {
nodes: self.scope.nodes,
}));
}
self.depth += 1;
if self.depth > self.scope.max_depth {
self.scope.max_depth = self.depth;
self.report.max_depth = self.scope.max_depth;
}
if self.scope.max_depth > self.budget.max_depth {
return Err(self.breach(BudgetBreach::Depth {
depth: self.scope.max_depth,
}));
}
if let Some(anchor) = anchor {
if self.scope.anchors.insert(anchor) {
self.report.anchors = self.scope.anchors.len();
if self.scope.anchors.len() > self.budget.max_anchors {
return Err(self.breach(BudgetBreach::Anchors {
anchors: self.scope.anchors.len(),
}));
}
}
}
}
BudgetEvent::SequenceEnd | BudgetEvent::MappingEnd => {
if let Some(new_depth) = self.depth.checked_sub(1) {
self.depth = new_depth;
} else {
return Err(self.breach(BudgetBreach::SequenceUnbalanced));
}
}
BudgetEvent::_Marker(_) => unreachable!(),
}
Ok(())
}
pub(crate) fn finish(mut self) -> Result<BudgetReport, BudgetBreach> {
if self.report.breached.is_none() && self.report.documents > 0 && self.depth == 0 {
if self.scope.aliases > 0
|| self.scope.nodes > 0
|| self.scope.total_scalar_bytes > 0
|| !self.scope.anchors.is_empty()
|| self.scope.max_depth > 0
{
self.finish_document()?;
}
}
Ok(self.report)
}
}
pub fn check_yaml_budget(input: &str, budget: &Budget) -> Result<BudgetReport, ScanError> {
let mut parser = Parser::new_from_str_with_options(
input,
saphyr_parser::options! {
emit_comments: false,
},
);
let mut tracker = BudgetTracker::<usize>::new(budget);
while let Some(item) = parser.next() {
let (ev, _span) = match item {
Ok(item) => item,
Err(err) => {
if let Some(prefix) = fallback_budget_input(input) {
return check_yaml_budget(prefix, budget);
}
return Err(err);
}
};
let budget_event = match ev {
Event::StreamStart => BudgetEvent::StreamStart,
Event::StreamEnd => BudgetEvent::StreamEnd,
Event::DocumentStart(_explicit, _version) => BudgetEvent::DocumentStart,
Event::DocumentEnd => BudgetEvent::DocumentEnd,
Event::Alias(anchor_id) => BudgetEvent::Alias(anchor_id),
Event::Scalar(value, _style, anchor_id, _tag_opt) => BudgetEvent::Scalar {
anchor: (anchor_id != 0).then_some(anchor_id),
scalar_bytes: match value {
Cow::Borrowed(s) => s.len(),
Cow::Owned(s) => s.len(),
},
},
Event::SequenceStart(_style, anchor_id, _tag_opt) => BudgetEvent::SequenceStart {
anchor: (anchor_id != 0).then_some(anchor_id),
},
Event::SequenceEnd => BudgetEvent::SequenceEnd,
Event::MappingStart(_style, anchor_id, _tag_opt) => BudgetEvent::MappingStart {
anchor: (anchor_id != 0).then_some(anchor_id),
},
Event::MappingEnd => BudgetEvent::MappingEnd,
_ => BudgetEvent::Nothing,
};
if let Err(breach) = tracker.observe(budget_event) {
let mut report = tracker.finish().unwrap_or_else(|_| BudgetReport::default());
report.breached = Some(breach);
return Ok(report);
}
}
match tracker.finish() {
Ok(report) => Ok(report),
Err(breach) => Ok(BudgetReport {
breached: Some(breach),
..BudgetReport::default()
}),
}
}
pub fn exceeds_yaml_budget(input: &str, budget: &Budget) -> Result<bool, ScanError> {
let report = check_yaml_budget(input, budget)?;
Ok(report.breached.is_some())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tiny_yaml_ok() {
let b = Budget::default();
let y = "a: [1, 2, 3]\n";
let r = check_yaml_budget(y, &b).unwrap();
assert!(r.breached.is_none());
assert_eq!(r.documents, 1);
assert_eq!(r.nodes > 0, true);
}
#[test]
fn comments_are_not_emitted_as_budget_events() {
let budget = Budget::default();
let without_comment = check_yaml_budget("a: 1\n", &budget).unwrap();
let with_comment = check_yaml_budget("# ignored\na: 1 # ignored\n", &budget).unwrap();
assert_eq!(with_comment.events, without_comment.events);
}
#[test]
fn alias_bomb_trips_alias_limit() {
let y = r#"root: &A [1, 2]
a: *A
b: *A
c: *A
d: *A
e: *A
"#;
let mut b = Budget::default();
b.max_aliases = 3;
let rep = check_yaml_budget(y, &b).unwrap();
assert!(matches!(rep.breached, Some(BudgetBreach::Aliases{ .. })));
}
#[test]
fn deep_nesting_trips_depth() {
let mut y = String::new();
for _ in 0..200 {
y.push('[');
}
for _ in 0..200 {
y.push(']');
}
let mut b = Budget::default();
b.max_depth = 150;
let rep = check_yaml_budget(&y, &b).unwrap();
assert!(matches!(rep.breached, Some(BudgetBreach::Depth{ .. })));
}
#[test]
fn anchors_limit_trips() {
let y = "a: &A 1\nb: &B 2\nc: &C 3\n";
let mut b = Budget::default();
b.max_anchors = 2;
let rep = check_yaml_budget(y, &b).unwrap();
assert!(matches!(rep.breached, Some(BudgetBreach::Anchors { anchors: 3 })));
}
#[test]
fn explicit_end_marker_allows_ignoring_trailing_garbage() {
let yaml = "---\na: 1\n...\n!!! trailing garbage\n";
let report = check_yaml_budget(yaml, &Budget::default()).unwrap();
assert!(report.breached.is_none());
assert_eq!(report.documents, 1);
}
#[test]
fn explicit_end_marker_keeps_following_documents_visible() {
let yaml = "---\na: 1\n...\n---\nb: 2\n";
let report = check_yaml_budget(yaml, &Budget::default()).unwrap();
assert!(report.breached.is_none());
assert_eq!(report.documents, 2);
}
#[test]
fn non_document_limits_apply_globally_across_documents() {
let yaml = "---\na: 1\nb: 2\n...\n---\nc: 3\nd: 4\n";
let mut budget = Budget::default();
budget.max_nodes = 8;
let report = check_yaml_budget(yaml, &budget).unwrap();
assert!(matches!(report.breached, Some(BudgetBreach::Nodes { nodes: 9 })));
assert_eq!(report.documents, 2);
assert_eq!(report.nodes, 9);
}
#[test]
fn document_count_remains_stream_wide() {
let yaml = "---\na: 1\n---\nb: 2\n";
let mut budget = Budget::default();
budget.max_documents = 1;
let report = check_yaml_budget(yaml, &budget).unwrap();
assert!(matches!(
report.breached,
Some(BudgetBreach::Documents { documents: 2 })
));
}
}