use std::sync::Arc;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Eq, PartialEq, Serialize)]
#[serde(untagged)]
pub enum UpstreamAuthority {
Literal(Arc<str>),
Derived {
from: AuthoritySource,
},
}
impl UpstreamAuthority {
pub fn literal(&self) -> Option<&str> {
match self {
Self::Literal(authority) => Some(authority),
Self::Derived { .. } => None,
}
}
pub fn follows_endpoint(&self) -> bool {
match self {
Self::Derived {
from: AuthoritySource::Endpoint,
} => true,
Self::Literal(_) => false,
}
}
}
impl From<&str> for UpstreamAuthority {
fn from(authority: &str) -> Self {
Self::Literal(Arc::from(authority))
}
}
impl<'de> Deserialize<'de> for UpstreamAuthority {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error as _;
match serde_yaml::Value::deserialize(deserializer)? {
serde_yaml::Value::String(authority) => Ok(Self::Literal(Arc::from(authority))),
mapping @ serde_yaml::Value::Mapping(_) => DerivedAuthority::deserialize(mapping)
.map(|derived| Self::Derived { from: derived.from })
.map_err(D::Error::custom),
serde_yaml::Value::Null
| serde_yaml::Value::Bool(_)
| serde_yaml::Value::Number(_)
| serde_yaml::Value::Sequence(_)
| serde_yaml::Value::Tagged(_) => Err(D::Error::custom(
"authority must be a host[:port] string or { from: endpoint }",
)),
}
}
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum AuthoritySource {
Endpoint,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct DerivedAuthority {
from: AuthoritySource,
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, clippy::expect_used, reason = "tests")]
mod tests {
use super::*;
#[test]
fn string_parses_as_literal() {
let authority: UpstreamAuthority = serde_yaml::from_str("\"api.example.com:8443\"").unwrap();
assert_eq!(
authority,
UpstreamAuthority::Literal(Arc::from("api.example.com:8443")),
"a string should parse as a fixed authority"
);
}
#[test]
fn mapping_parses_as_derived_from_endpoint() {
let authority: UpstreamAuthority = serde_yaml::from_str("from: endpoint").unwrap();
assert_eq!(
authority,
UpstreamAuthority::Derived {
from: AuthoritySource::Endpoint
},
"{{from: endpoint}} should parse as the endpoint-derived form"
);
}
#[test]
fn unknown_source_is_named_in_the_error() {
let err = serde_yaml::from_str::<UpstreamAuthority>("from: nonsense").unwrap_err();
let message = err.to_string();
assert!(
message.contains("nonsense"),
"error should name the bad value: {message}"
);
assert!(
message.contains("endpoint"),
"error should list the valid value: {message}"
);
}
#[test]
fn unknown_key_is_named_in_the_error() {
let err = serde_yaml::from_str::<UpstreamAuthority>("frm: endpoint").unwrap_err();
let message = err.to_string();
assert!(message.contains("frm"), "error should name the unknown key: {message}");
assert!(
message.contains("from"),
"error should name the expected key: {message}"
);
}
#[test]
fn mapping_without_source_is_rejected() {
let err = serde_yaml::from_str::<UpstreamAuthority>("{}").unwrap_err();
assert!(
err.to_string().contains("from"),
"an empty mapping should report the missing key: {err}"
);
}
#[test]
fn other_shapes_are_rejected() {
for yaml in ["8080", "true", "[api.example.com]"] {
let err = serde_yaml::from_str::<UpstreamAuthority>(yaml).unwrap_err();
assert!(
err.to_string().contains("host[:port] string or { from: endpoint }"),
"{yaml:?} should be rejected with the expected-shape message: {err}"
);
}
}
#[test]
fn literal_round_trips_as_a_string() {
let authority = UpstreamAuthority::from("api.example.com");
let value = serde_yaml::to_value(&authority).unwrap();
assert_eq!(
value,
serde_yaml::Value::String("api.example.com".to_owned()),
"a fixed authority should serialize as a bare string"
);
let back: UpstreamAuthority = serde_yaml::from_value(value).unwrap();
assert_eq!(back, authority, "a fixed authority should round-trip");
}
#[test]
fn derived_round_trips_as_a_mapping() {
let authority = UpstreamAuthority::Derived {
from: AuthoritySource::Endpoint,
};
let value = serde_yaml::to_value(&authority).unwrap();
assert_eq!(
value,
serde_yaml::from_str::<serde_yaml::Value>("from: endpoint").unwrap(),
"the endpoint-derived form should serialize as {{from: endpoint}}"
);
let back: UpstreamAuthority = serde_yaml::from_value(value).unwrap();
assert_eq!(back, authority, "the endpoint-derived form should round-trip");
}
}