use crate::error::{Error, Result};
use crate::parser::{self};
use crate::prelude::*;
#[cfg(feature = "std")]
use crate::span_context;
use crate::value::Value;
#[cfg(feature = "std")]
use std::io;
mod config;
mod deserializer;
pub use config::{
DuplicateKeyPolicy, MergeKeyPolicy, NonScalarKeyPolicy, ParserConfig, ParserLimits,
ParserProfile, RequireIndent, YamlVersion,
};
pub use deserializer::Deserializer;
pub(crate) use deserializer::{EmptyMapAccess, SpannedMapAccess, is_binary_tag};
pub fn from_str<T>(s: &str) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
from_str_with_config(s, &ParserConfig::default())
}
pub fn from_str_borrowing<'de, T>(s: &'de str) -> Result<T>
where
T: serde_core::Deserialize<'de>,
{
from_str_borrowing_with_config(s, &ParserConfig::default())
}
pub fn from_str_borrowing_with_config<'de, T>(s: &'de str, config: &ParserConfig) -> Result<T>
where
T: serde_core::Deserialize<'de>,
{
let parse_config = parser::ParseConfig::from(config);
if s.len() > parse_config.max_document_length {
return Err(Error::Parse(format!(
"document exceeds maximum length of {} bytes",
parse_config.max_document_length
)));
}
if !config.policies.is_empty() {
let document = parser::parse_exactly_one_value(s, &parse_config)?;
crate::policy::check_document(&config.policies, &document)?;
}
let mut de = crate::streaming::StreamingDeserializer::with_config(s, parse_config);
if let Some(registry) = config.tag_registry.as_ref() {
de = de.with_tag_registry(Arc::clone(registry));
}
let value = T::deserialize(&mut de)?;
de.end()?;
Ok(value)
}
fn is_value_target<T: 'static + ?Sized>() -> bool {
use core::any::TypeId;
TypeId::of::<T>() == TypeId::of::<Value>()
}
#[cfg(all(feature = "std", feature = "figment"))]
pub(crate) fn from_str_typed_no_tag_preserve<T>(s: &str, config: &ParserConfig) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de>,
{
if s.len() > config.max_document_length {
return Err(Error::Parse(format!(
"document exceeds maximum length of {} bytes",
config.max_document_length
)));
}
let stream_eligible = config.merge_key_policy == MergeKeyPolicy::Auto
&& !config.ignore_binary_tag_for_string
&& config.policies.is_empty();
let mut streaming_err = None;
if stream_eligible {
if let Some(res) = crate::streaming::from_str_streaming(s, config) {
match res {
Err(err) if err.location().is_none() => streaming_err = Some(err),
other => return other,
}
}
}
let parse_config = parser::ParseConfig::from(config);
let (value, span_tree) = parser::parse_exactly_one(s, &parse_config)?;
for p in &config.policies {
p.check_value(&value)?;
}
let spans = span_context::build_span_map(&value, &span_tree);
let ctx = span_context::SpanContext::new(spans, s.into());
let _guard = span_context::set_span_context(ctx);
let de = Deserializer::with_options(
&value,
Some(_guard.as_ref()),
config.ignore_binary_tag_for_string,
config.plain_scalar_strings,
);
attach_field_path(
locate_streaming_error(T::deserialize(de), streaming_err),
&value,
)
}
#[cfg(all(feature = "std", feature = "strict-deserialise"))]
pub fn from_str_strict<T>(s: &str) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let unknown = std::sync::Mutex::new(Vec::<String>::new());
let value: Value = from_str_with_config(s, &ParserConfig::default())?;
let result: Result<T> = serde_ignored::deserialize(&value, |path| {
unknown
.lock()
.expect("from_str_strict: ignored-paths lock poisoned")
.push(path.to_string());
});
let extras = unknown
.into_inner()
.expect("from_str_strict: ignored-paths lock poisoned");
let typed = result?;
if !extras.is_empty() {
let msg = if extras.len() == 1 {
format!("unknown field at `{}`", extras[0])
} else {
let joined = extras
.iter()
.map(|p| format!("`{p}`"))
.collect::<Vec<_>>()
.join(", ");
format!("unknown fields: {joined}")
};
return Err(Error::UnknownField(msg));
}
Ok(typed)
}
#[cfg(all(feature = "std", feature = "strict-deserialise"))]
pub fn from_slice_strict<T>(b: &[u8]) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let s = core::str::from_utf8(b).map_err(|e| Error::Deserialize(e.to_string()))?;
from_str_strict(s)
}
#[cfg(all(feature = "std", feature = "strict-deserialise"))]
pub fn from_reader_strict<R, T>(reader: R) -> Result<T>
where
R: io::Read,
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let s = read_to_string_bounded(reader, &ParserConfig::default())?;
from_str_strict(&s)
}
#[cfg(feature = "std")]
pub(crate) fn read_to_string_bounded<R>(mut reader: R, config: &ParserConfig) -> Result<String>
where
R: io::Read,
{
use io::Read as _;
let max_len = parser::ParseConfig::from(config).max_document_length;
let limit = u64::try_from(max_len).unwrap_or(u64::MAX).saturating_add(1);
let mut buf = Vec::new();
let _ = reader
.by_ref()
.take(limit)
.read_to_end(&mut buf)
.map_err(Error::Io)?;
if buf.len() > max_len {
return Err(Error::Parse(format!(
"document exceeds maximum length of {max_len} bytes"
)));
}
String::from_utf8(buf).map_err(|_| {
Error::Io(io::Error::new(
io::ErrorKind::InvalidData,
"stream did not contain valid UTF-8",
))
})
}
pub fn from_str_with_config<T>(s: &str, config: &ParserConfig) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let value_target_bypass = is_value_target::<T>();
let stream_eligible = config.merge_key_policy == MergeKeyPolicy::Auto
&& !config.ignore_binary_tag_for_string
&& config.policies.is_empty()
&& properties_inactive(config)
&& includes_inactive(config)
&& !value_target_bypass;
let mut streaming_err = None;
if stream_eligible {
if let Some(res) = crate::streaming::from_str_streaming(s, config) {
match res {
Err(err) if err.location().is_none() => streaming_err = Some(err),
other => return other,
}
}
}
let parse_config = parser::ParseConfig::from(config);
if s.len() > parse_config.max_document_length {
return Err(Error::Parse(format!(
"document exceeds maximum length of {} bytes",
parse_config.max_document_length
)));
}
if is_value_target::<T>() {
let mut value = parser::parse_exactly_one_value(s, &parse_config)?;
apply_includes(&mut value, config)?;
apply_properties(&mut value, config)?;
crate::policy::check_document(&config.policies, &value)?;
let boxed: Box<dyn core::any::Any> = Box::new(value);
let downcast: Box<T> = boxed
.downcast::<T>()
.expect("is_value_target proved T == Value");
return Ok(*downcast);
}
#[cfg(feature = "std")]
{
let (mut value, span_tree) = parser::parse_exactly_one(s, &parse_config)?;
apply_includes(&mut value, config)?;
apply_properties(&mut value, config)?;
crate::policy::check_document(&config.policies, &value)?;
let spans = span_context::build_span_map(&value, &span_tree);
let ctx = span_context::SpanContext::new(spans, s.into());
let _guard = span_context::set_span_context(ctx);
let de = Deserializer::with_options(
&value,
Some(_guard.as_ref()),
config.ignore_binary_tag_for_string,
config.plain_scalar_strings,
);
attach_field_path(
locate_streaming_error(T::deserialize(de), streaming_err),
&value,
)
}
#[cfg(not(feature = "std"))]
{
let value = parser::parse_exactly_one_value(s, &parse_config)?;
crate::policy::check_document(&config.policies, &value)?;
let de = Deserializer::with_options(
&value,
None,
config.ignore_binary_tag_for_string,
config.plain_scalar_strings,
);
locate_streaming_error(T::deserialize(de), streaming_err)
}
}
fn locate_streaming_error<T>(ast_result: Result<T>, streaming_err: Option<Error>) -> Result<T> {
match (ast_result, streaming_err) {
(Err(ast_err), Some(_))
if matches!(
ast_err.kind(),
crate::error::ErrorKind::DuplicateKey | crate::error::ErrorKind::KeyCollision
) && ast_err.location().is_some() =>
{
Err(ast_err)
}
(Err(ast_err), Some(streaming_err)) => Err(match ast_err.location() {
Some(location) => {
let message = match streaming_err {
Error::Custom(message) | Error::Deserialize(message) => message,
other => other.to_string(),
};
Error::DeserializeWithLocation { message, location }
}
None => streaming_err,
}),
(result, _) => result,
}
}
#[cfg(feature = "std")]
fn attach_field_path<T>(result: Result<T>, root: &Value) -> Result<T> {
let recorded = span_context::take_error_node();
let err = match result {
Ok(v) => return Ok(v),
Err(err) => err,
};
let Some(addr) = recorded else {
return Err(err);
};
let mut segments = Vec::new();
if !find_value_path(root, addr, &mut segments) || segments.is_empty() {
return Err(err);
}
let path = format_value_path(&segments);
Err(match err {
Error::DeserializeWithLocation { message, location } => Error::DeserializeWithLocation {
message: format!("{path}: {message}"),
location,
},
Error::Deserialize(message) => Error::Deserialize(format!("{path}: {message}")),
Error::Custom(message) => Error::Custom(format!("{path}: {message}")),
other => other,
})
}
#[cfg(feature = "std")]
enum PathSegment<'a> {
Key(&'a str),
Index(usize),
}
#[cfg(feature = "std")]
fn find_value_path<'a>(value: &'a Value, addr: usize, segments: &mut Vec<PathSegment<'a>>) -> bool {
if core::ptr::from_ref(value) as usize == addr {
return true;
}
match value {
Value::Mapping(map) => {
for (key, child) in map {
segments.push(PathSegment::Key(key));
if find_value_path(child, addr, segments) {
return true;
}
let _ = segments.pop();
}
}
Value::Sequence(seq) => {
for (index, child) in seq.iter().enumerate() {
segments.push(PathSegment::Index(index));
if find_value_path(child, addr, segments) {
return true;
}
let _ = segments.pop();
}
}
Value::Tagged(tagged) => return find_value_path(tagged.value(), addr, segments),
_ => {}
}
false
}
#[cfg(feature = "std")]
fn format_value_path(segments: &[PathSegment<'_>]) -> String {
use core::fmt::Write as _;
let mut out = String::new();
for segment in segments {
match segment {
PathSegment::Key(key) => {
if !out.is_empty() {
out.push('.');
}
out.push_str(key);
}
PathSegment::Index(index) => {
let _ = write!(out, "[{index}]");
}
}
}
out
}
#[cfg(feature = "std")]
#[inline]
fn properties_inactive(config: &ParserConfig) -> bool {
config.properties.is_none()
}
#[cfg(not(feature = "std"))]
#[inline]
fn properties_inactive(_config: &ParserConfig) -> bool {
true
}
#[cfg(feature = "std")]
fn apply_properties(value: &mut Value, config: &ParserConfig) -> Result<()> {
if let Some(props) = config.properties.as_ref() {
let action = if config.strict_properties {
crate::value::MissingAction::Error(false)
} else {
crate::value::MissingAction::Empty
};
value.interpolate_inner(
&|name| match props.get(name) {
Some(v) => crate::value::ResolveOutcome::Found(v.clone()),
None => crate::value::ResolveOutcome::Missing,
},
action,
)?;
}
Ok(())
}
#[cfg(not(feature = "std"))]
#[inline]
fn apply_properties(_value: &mut Value, _config: &ParserConfig) -> Result<()> {
Ok(())
}
#[cfg(feature = "include")]
#[inline]
fn includes_inactive(config: &ParserConfig) -> bool {
config.include_resolver.is_none()
}
#[cfg(not(feature = "include"))]
#[inline]
fn includes_inactive(_config: &ParserConfig) -> bool {
true
}
#[cfg(all(feature = "include", feature = "std"))]
type IncludeVisited = std::collections::HashSet<String>;
#[cfg(all(feature = "include", not(feature = "std")))]
type IncludeVisited = FxHashSet<String>;
#[cfg(feature = "include")]
fn apply_includes(value: &mut Value, config: &ParserConfig) -> Result<()> {
if let Some(resolver) = config.include_resolver.as_ref() {
let mut walk = IncludeWalk {
resolver,
parse_config: parser::ParseConfig::from(config),
max_include_depth: config.max_include_depth,
max_include_sources: config.max_include_sources,
max_total_include_bytes: config.max_total_include_bytes,
visited: IncludeVisited::default(),
next_id: 1,
sources: 0,
bytes: 0,
};
walk.walk(value, IncludeAt::default())?;
enforce_expanded_node_budget(value, config.max_nodes)?;
}
Ok(())
}
#[cfg(feature = "include")]
fn enforce_expanded_node_budget(value: &Value, max_nodes: usize) -> Result<()> {
let mut pending = vec![value];
let mut observed = 0usize;
while let Some(value) = pending.pop() {
let charge = match value {
Value::Tagged(tagged) => {
pending.push(tagged.value());
continue;
}
Value::Sequence(sequence) => {
pending.extend(sequence);
1
}
Value::Mapping(mapping) => {
pending.extend(mapping.values());
1usize.saturating_add(mapping.len())
}
Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => 1,
};
observed = observed.saturating_add(charge);
if observed > max_nodes {
return Err(Error::Budget(crate::BudgetBreach::MaxNodes {
limit: max_nodes,
observed,
}));
}
}
Ok(())
}
#[cfg(not(feature = "include"))]
#[inline]
fn apply_includes(_value: &mut Value, _config: &ParserConfig) -> Result<()> {
Ok(())
}
#[cfg(feature = "include")]
#[derive(Debug, Clone, Copy, Default)]
struct IncludeAt {
include_depth: usize,
from_id: usize,
tree_depth: usize,
}
#[cfg(feature = "include")]
struct IncludeWalk<'a> {
resolver: &'a crate::include::IncludeResolver,
parse_config: parser::ParseConfig,
max_include_depth: usize,
max_include_sources: usize,
max_total_include_bytes: usize,
visited: IncludeVisited,
next_id: usize,
sources: usize,
bytes: usize,
}
#[cfg(feature = "include")]
impl IncludeWalk<'_> {
fn walk(&mut self, value: &mut Value, at: IncludeAt) -> Result<()> {
if let Value::Tagged(tagged) = value {
if tagged.tag().as_str() == "!include" {
let spec = tagged.value().as_str().map(str::to_owned).ok_or_else(|| {
Error::Custom("!include directive expects a scalar string spec".into())
})?;
*value = self.load(&spec, at)?;
return Ok(());
}
}
let inner = IncludeAt {
tree_depth: at.tree_depth + 1,
..at
};
match value {
Value::Tagged(tagged) => self.walk(tagged.value_mut(), at),
Value::Sequence(seq) => seq.iter_mut().try_for_each(|v| self.walk(v, inner)),
Value::Mapping(map) => map.values_mut().try_for_each(|v| self.walk(v, inner)),
Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => Ok(()),
}
}
fn load(&mut self, spec: &str, at: IncludeAt) -> Result<Value> {
let include_depth = at.include_depth + 1;
if include_depth > self.max_include_depth {
return Err(Error::RecursionLimitExceeded {
depth: include_depth,
});
}
let remaining = self.max_total_include_bytes.saturating_sub(self.bytes);
let req = crate::include::IncludeRequest {
spec,
from_id: at.from_id,
depth: at.include_depth,
max_bytes: remaining.min(self.parse_config.max_document_length),
};
let source = self.resolver.resolve(req)?;
self.charge(&source)?;
let identity = source.name;
if !self.visited.insert(identity.clone()) {
return Err(Error::Custom(format!(
"!include cycle detected: canonical source `{identity}` is already in the resolution chain"
)));
}
let id = self.next_id;
self.next_id += 1;
let mut source_config = self.parse_config.clone();
source_config.max_depth = source_config.max_depth.saturating_sub(at.tree_depth);
let mut included = parser::parse_exactly_one_value(&source.bytes, &source_config)?;
let nested = IncludeAt {
include_depth,
from_id: id,
tree_depth: at.tree_depth,
};
self.walk(&mut included, nested)?;
let _ = self.visited.remove(&identity);
select_include_fragment(included, spec)
}
fn charge(&mut self, source: &crate::include::InputSource) -> Result<()> {
self.sources = self.sources.saturating_add(1);
if self.sources > self.max_include_sources {
return Err(Error::Budget(crate::BudgetBreach::MaxIncludeSources {
limit: self.max_include_sources,
observed: self.sources,
}));
}
self.bytes = self.bytes.saturating_add(source.bytes.len());
if self.bytes > self.max_total_include_bytes {
return Err(Error::Budget(crate::BudgetBreach::MaxIncludeBytes {
limit: self.max_total_include_bytes,
observed: self.bytes,
}));
}
let max_len = self.parse_config.max_document_length;
if source.bytes.len() > max_len {
return Err(Error::Parse(format!(
"included document `{}` exceeds maximum length of {max_len} bytes",
source.name
)));
}
Ok(())
}
}
#[cfg(feature = "include")]
fn select_include_fragment(included: Value, spec: &str) -> Result<Value> {
let (path, fragment) = crate::include::split_fragment(spec);
let Some(frag) = fragment else {
return Ok(included);
};
let Some(map) = included.as_mapping() else {
return Err(Error::Custom(format!(
"!include fragment `#{frag}` requires a mapping-shaped \
included document; `{path}` is not a mapping"
)));
};
map.get(frag)
.cloned()
.ok_or_else(|| Error::Custom(format!("!include fragment `#{frag}` not found in `{path}`")))
}
pub fn from_slice<T>(b: &[u8]) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let s = core::str::from_utf8(b).map_err(|e| Error::Deserialize(e.to_string()))?;
from_str(s)
}
pub fn from_slice_with_config<T>(b: &[u8], config: &ParserConfig) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let s = core::str::from_utf8(b).map_err(|e| Error::Deserialize(e.to_string()))?;
from_str_with_config(s, config)
}
#[cfg(feature = "std")]
#[cfg_attr(docsrs, doc(cfg(feature = "std")))]
pub fn from_reader<R, T>(reader: R) -> Result<T>
where
R: io::Read,
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
from_reader_with_config(reader, &ParserConfig::default())
}
#[cfg(feature = "std")]
#[cfg_attr(docsrs, doc(cfg(feature = "std")))]
pub fn from_reader_with_config<R, T>(reader: R, config: &ParserConfig) -> Result<T>
where
R: io::Read,
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
let s = read_to_string_bounded(reader, config)?;
from_str_with_config(&s, config)
}
pub fn from_value<T>(value: &Value) -> Result<T>
where
T: for<'de> serde_core::Deserialize<'de> + 'static,
{
if is_value_target::<T>() {
let cloned = value.clone();
let boxed: Box<dyn core::any::Any> = Box::new(cloned);
let downcast: Box<T> = boxed
.downcast::<T>()
.expect("is_value_target proved T == Value");
return Ok(*downcast);
}
T::deserialize(Deserializer::new(value))
}
#[cfg(all(test, feature = "std", feature = "figment"))]
mod figment_entry_tests {
use super::*;
#[derive(Debug, serde::Deserialize, PartialEq)]
struct Endpoint {
name: String,
port: u16,
}
#[test]
fn no_tag_preserve_takes_streaming_fast_path() {
let cfg = ParserConfig::default();
let got: Endpoint = from_str_typed_no_tag_preserve("name: alpha\nport: 7\n", &cfg).unwrap();
assert_eq!(
got,
Endpoint {
name: "alpha".into(),
port: 7
}
);
}
#[test]
fn no_tag_preserve_falls_back_to_ast_on_merge_key() {
let cfg = ParserConfig::default();
let yaml = "\
_defaults: &d\n name: beta\n port: 9\n<<: *d\n";
let got: Endpoint = from_str_typed_no_tag_preserve(yaml, &cfg).unwrap();
assert_eq!(
got,
Endpoint {
name: "beta".into(),
port: 9
}
);
}
#[test]
fn no_tag_preserve_surfaces_parse_error() {
let cfg = ParserConfig::default().max_document_length(4);
let res: Result<Endpoint> = from_str_typed_no_tag_preserve("name: alpha\n", &cfg);
assert!(matches!(res, Err(Error::Parse(_))), "got {res:?}");
}
#[test]
fn no_tag_preserve_runs_value_policies() {
#[derive(Debug)]
struct RejectAll;
impl crate::policy::Policy for RejectAll {
fn check_value(&self, _v: &Value) -> Result<()> {
Err(Error::Deserialize("rejected by policy".into()))
}
}
let cfg = ParserConfig::default().with_policy(RejectAll);
let res: Result<Endpoint> = from_str_typed_no_tag_preserve("name: alpha\nport: 7\n", &cfg);
assert!(
matches!(res, Err(Error::Deserialize(_))),
"policy rejection must surface: {res:?}"
);
}
}