use super::error::SchemaError;
use super::value::JsonValue;
use std::collections::HashMap;
pub(super) fn collect(root: &JsonValue) -> Result<HashMap<String, JsonValue>, SchemaError> {
let mut refs = HashMap::new();
walk(root, root, &mut refs)?;
Ok(refs)
}
fn walk(
root: &JsonValue,
node: &JsonValue,
refs: &mut HashMap<String, JsonValue>,
) -> Result<(), SchemaError> {
match node {
JsonValue::Array(items) => {
for item in items {
walk(root, item, refs)?;
}
}
JsonValue::Object(entries) => {
if let Some(reference) = node.get("$ref") {
let reference = reference
.as_str()
.ok_or_else(|| SchemaError::UnsupportedRef {
reference: reference.dump(),
})?;
if !refs.contains_key(reference) {
let target = resolve_pointer(root, reference)?;
refs.insert(reference.to_string(), target);
}
} else {
for (_, value) in entries {
walk(root, value, refs)?;
}
}
}
_ => {}
}
Ok(())
}
fn resolve_pointer(root: &JsonValue, reference: &str) -> Result<JsonValue, SchemaError> {
if !reference.starts_with("#/") {
return Err(SchemaError::UnsupportedRef {
reference: reference.to_string(),
});
}
let mut target = root;
for raw in reference[1..].split('/').skip(1) {
let token = unescape_pointer_token(raw);
let next = match target {
JsonValue::Object(_) => target.get(&token),
JsonValue::Array(items) => token.parse::<usize>().ok().and_then(|i| items.get(i)),
_ => None,
};
target = next.ok_or_else(|| SchemaError::RefNotFound {
reference: reference.to_string(),
token,
})?;
}
Ok(target.clone())
}
fn unescape_pointer_token(token: &str) -> String {
if token.contains('~') {
token.replace("~1", "/").replace("~0", "~")
} else {
token.to_string()
}
}
pub(super) fn ref_rule_name(reference: &str) -> String {
let fragment = match reference.find('#') {
Some(i) => &reference[i + 1..],
None => reference,
};
let mut out = String::from("ref");
out.push_str(&super::converter::collapse_invalid(fragment));
out
}
#[cfg(test)]
mod tests {
use super::*;
fn schema(text: &str) -> JsonValue {
JsonValue::parse(text).expect("test schema parses")
}
#[test]
fn ref_names_match_upstream() {
assert_eq!(ref_rule_name("#/definitions/foo"), "ref-definitions-foo");
assert_eq!(ref_rule_name("#/$defs/Node"), "ref-defs-Node");
assert_eq!(
ref_rule_name("#/properties/a/anyOf/0"),
"ref-properties-a-anyOf-0"
);
}
#[test]
fn collects_nested_and_recursive_refs() {
let s = schema(
r##"{"$ref":"#/$defs/a","$defs":{"a":{"properties":{"n":{"$ref":"#/$defs/a"}}}}}"##,
);
let refs = collect(&s).expect("refs resolve");
assert_eq!(refs.len(), 1);
assert!(refs["#/$defs/a"].contains_key("properties"));
}
#[test]
fn refuses_remote_and_dangling_refs() {
let remote = schema(r##"{"$ref":"https://example.com/s.json#/x"}"##);
assert!(matches!(
collect(&remote),
Err(SchemaError::UnsupportedRef { .. })
));
let dangling = schema(r##"{"$ref":"#/$defs/missing","$defs":{}}"##);
assert!(matches!(
collect(&dangling),
Err(SchemaError::RefNotFound { .. })
));
}
#[test]
fn pointer_escapes_are_decoded() {
let s = schema(r##"{"$ref":"#/$defs/a~1b","$defs":{"a/b":{"type":"string"}}}"##);
let refs = collect(&s).expect("escaped pointer resolves");
assert_eq!(
refs["#/$defs/a~1b"].get("type").unwrap().as_str(),
Some("string")
);
}
}