use std::collections::HashMap;
use serde::Deserialize;
use super::{impl_condition_deserialize, impl_condition_serialize};
#[derive(Clone, Debug)]
pub enum Condition {
When(ConditionMatch),
Unless(ConditionMatch),
}
impl_condition_deserialize!(Condition, ConditionMatch, "condition");
impl_condition_serialize!(Condition, ConditionMatch);
#[derive(Clone, Debug, Deserialize, serde::Serialize)]
#[serde(deny_unknown_fields)]
pub struct ConditionMatch {
#[serde(default)]
pub grpc: Option<bool>,
#[serde(default)]
pub path: Option<String>,
#[serde(default)]
pub path_prefix: Option<String>,
#[serde(default)]
pub methods: Option<Vec<String>>,
#[serde(default)]
pub headers: Option<HashMap<String, String>>,
#[serde(default)]
pub bound_upstream: Option<ApplicationMatch>,
#[serde(default)]
pub selected_upstream: Option<SelectedUpstreamMatch>,
}
#[derive(Clone, Debug, Deserialize, serde::Serialize)]
#[serde(deny_unknown_fields)]
pub struct ApplicationMatch {
#[serde(default)]
pub application_protocol: Option<String>,
#[serde(default)]
pub application_provider: Option<String>,
}
impl ApplicationMatch {
#[must_use]
pub fn is_empty(&self) -> bool {
self.application_protocol.is_none() && self.application_provider.is_none()
}
}
pub type SelectedUpstreamMatch = ApplicationMatch;
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::needless_raw_strings,
clippy::needless_raw_string_hashes,
clippy::min_ident_chars,
reason = "tests use unwrap/expect/indexing/raw strings for brevity"
)]
mod tests {
use super::*;
#[test]
fn parse_condition_match_all_fields() {
let yaml = r#"
path_prefix: "/api"
methods: ["GET", "POST"]
headers:
x-tenant: "acme"
x-debug: "true"
"#;
let m: ConditionMatch = serde_yaml::from_str(yaml).unwrap();
assert_eq!(m.path_prefix.as_deref(), Some("/api"), "path_prefix mismatch");
let methods = m.methods.unwrap();
assert_eq!(methods, vec!["GET", "POST"], "methods mismatch");
let headers = m.headers.unwrap();
assert_eq!(headers.get("x-tenant").unwrap(), "acme", "x-tenant header mismatch");
assert_eq!(headers.get("x-debug").unwrap(), "true", "x-debug header mismatch");
}
#[test]
fn parse_condition_match_partial() {
let yaml = r#"
path_prefix: "/health"
"#;
let m: ConditionMatch = serde_yaml::from_str(yaml).unwrap();
assert_eq!(m.path_prefix.as_deref(), Some("/health"), "path_prefix mismatch");
assert!(m.methods.is_none(), "methods should be None when omitted");
assert!(m.headers.is_none(), "headers should be None when omitted");
}
#[test]
fn parse_when_condition() {
let yaml = r#"
- when:
path_prefix: "/api"
"#;
let conditions: Vec<Condition> = serde_yaml::from_str(yaml).unwrap();
assert_eq!(conditions.len(), 1, "should parse 1 condition");
assert!(
matches!(&conditions[0], Condition::When(m) if m.path_prefix.as_deref() == Some("/api")),
"should be When with /api prefix"
);
}
#[test]
fn parse_unless_condition() {
let yaml = r#"
- unless:
methods: ["OPTIONS"]
"#;
let conditions: Vec<Condition> = serde_yaml::from_str(yaml).unwrap();
assert_eq!(conditions.len(), 1, "should parse 1 condition");
assert!(
matches!(&conditions[0], Condition::Unless(m) if m.methods.as_ref().unwrap() == &["OPTIONS"]),
"should be Unless with OPTIONS method"
);
}
#[test]
fn parse_mixed_conditions() {
let yaml = r#"
- when:
path_prefix: "/api"
- unless:
headers:
x-internal: "true"
- when:
methods: ["POST", "PUT", "DELETE"]
"#;
let conditions: Vec<Condition> = serde_yaml::from_str(yaml).unwrap();
assert_eq!(conditions.len(), 3, "should parse 3 conditions");
assert!(matches!(&conditions[0], Condition::When(_)), "first should be When");
assert!(
matches!(&conditions[1], Condition::Unless(_)),
"second should be Unless"
);
assert!(matches!(&conditions[2], Condition::When(_)), "third should be When");
}
#[test]
fn parse_empty_conditions() {
let conditions: Vec<Condition> = serde_yaml::from_str("[]").unwrap();
assert!(conditions.is_empty(), "empty array should parse to empty vec");
}
#[test]
fn reject_both_when_and_unless() {
let yaml = r#"
- when:
path_prefix: "/api"
unless:
methods: ["GET"]
"#;
let err = serde_yaml::from_str::<Vec<Condition>>(yaml).unwrap_err();
assert!(err.to_string().contains("exactly one"));
}
#[test]
fn reject_neither_when_nor_unless() {
let yaml = "- {}";
let err = serde_yaml::from_str::<Vec<Condition>>(yaml).unwrap_err();
assert!(err.to_string().contains("either"));
}
#[test]
fn parse_grpc_predicate_true() {
let m: ConditionMatch = serde_yaml::from_str("grpc: true\n").unwrap();
assert_eq!(m.grpc, Some(true), "grpc: true should parse");
assert!(m.path.is_none(), "path should stay unset");
}
#[test]
fn parse_grpc_predicate_false() {
let m: ConditionMatch = serde_yaml::from_str("grpc: false\n").unwrap();
assert_eq!(m.grpc, Some(false), "grpc: false should parse");
}
#[test]
fn grpc_predicate_defaults_to_unset() {
let m: ConditionMatch = serde_yaml::from_str("path: \"/\"\n").unwrap();
assert!(m.grpc.is_none(), "grpc should be None when omitted");
}
#[test]
fn reject_non_boolean_grpc_predicate() {
let err = serde_yaml::from_str::<ConditionMatch>("grpc: \"yes\"\n").unwrap_err();
assert!(
err.to_string().contains("bool"),
"a non-boolean grpc value should be rejected: {err}"
);
}
#[test]
fn grpc_predicate_round_trips_through_serialization() {
let m: ConditionMatch = serde_yaml::from_str("grpc: true\npath_prefix: \"/pkg.Svc\"\n").unwrap();
let yaml = serde_yaml::to_string(&m).unwrap();
let back: ConditionMatch = serde_yaml::from_str(&yaml).unwrap();
assert_eq!(back.grpc, Some(true), "grpc should survive a serialize round trip");
assert_eq!(
back.path_prefix.as_deref(),
Some("/pkg.Svc"),
"path_prefix should survive a serialize round trip"
);
}
#[test]
fn parse_exact_path_condition() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
path: "/"
"#,
)
.unwrap();
assert_eq!(m.path.as_deref(), Some("/"), "exact path should be /");
assert!(
m.path_prefix.is_none(),
"path_prefix should be None for exact path match"
);
}
#[test]
fn parse_bound_upstream_condition() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
bound_upstream:
application_protocol: "openai_responses"
application_provider: "openai"
"#,
)
.unwrap();
let bound = m.bound_upstream.expect("bound_upstream should parse");
assert_eq!(
bound.application_protocol.as_deref(),
Some("openai_responses"),
"application_protocol mismatch"
);
assert_eq!(
bound.application_provider.as_deref(),
Some("openai"),
"application_provider mismatch"
);
}
#[test]
fn parse_bound_upstream_partial() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
bound_upstream:
application_protocol: "openai_responses"
"#,
)
.unwrap();
let bound = m.bound_upstream.expect("bound_upstream should parse");
assert_eq!(
bound.application_protocol.as_deref(),
Some("openai_responses"),
"application_protocol mismatch"
);
assert!(
bound.application_provider.is_none(),
"application_provider should be None when omitted"
);
}
#[test]
fn bound_upstream_defaults_to_unset() {
let m: ConditionMatch = serde_yaml::from_str("path: \"/\"\n").unwrap();
assert!(m.bound_upstream.is_none(), "bound_upstream should be None when omitted");
}
#[test]
fn reject_unknown_bound_upstream_field() {
let err = serde_yaml::from_str::<ConditionMatch>(
r#"
bound_upstream:
application_flavour: "openai"
"#,
)
.unwrap_err();
assert!(
err.to_string().contains("unknown field"),
"an unknown bound_upstream field should be rejected: {err}"
);
}
#[test]
fn bound_upstream_round_trips_through_serialization() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
bound_upstream:
application_protocol: "openai_responses"
application_provider: "openai"
"#,
)
.unwrap();
let yaml = serde_yaml::to_string(&m).unwrap();
let back: ConditionMatch = serde_yaml::from_str(&yaml).unwrap();
let bound = back.bound_upstream.expect("bound_upstream should survive a round trip");
assert_eq!(
bound.application_protocol.as_deref(),
Some("openai_responses"),
"application_protocol should survive a round trip"
);
assert_eq!(
bound.application_provider.as_deref(),
Some("openai"),
"application_provider should survive a round trip"
);
}
#[test]
fn parse_selected_upstream_both_fields() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
selected_upstream:
application_protocol: openai_chat_completions
application_provider: vllm
"#,
)
.unwrap();
let su = m.selected_upstream.expect("selected_upstream should parse");
assert_eq!(
su.application_protocol.as_deref(),
Some("openai_chat_completions"),
"application_protocol mismatch"
);
assert_eq!(
su.application_provider.as_deref(),
Some("vllm"),
"application_provider mismatch"
);
assert!(!su.is_empty(), "a populated selected_upstream is not empty");
}
#[test]
fn parse_selected_upstream_protocol_only() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
selected_upstream:
application_protocol: openai_responses
"#,
)
.unwrap();
let su = m.selected_upstream.expect("selected_upstream should parse");
assert_eq!(su.application_protocol.as_deref(), Some("openai_responses"));
assert!(
su.application_provider.is_none(),
"provider should be None when omitted"
);
}
#[test]
fn selected_upstream_provider_only_round_trips_through_serialization() {
let m: ConditionMatch = serde_yaml::from_str(
r#"
selected_upstream:
application_provider: vllm
"#,
)
.unwrap();
let yaml = serde_yaml::to_string(&m).unwrap();
let back: ConditionMatch = serde_yaml::from_str(&yaml).unwrap();
let su = back.selected_upstream.expect("selected_upstream should round-trip");
assert_eq!(su.application_provider.as_deref(), Some("vllm"));
assert!(
su.application_protocol.is_none(),
"protocol should remain None after round-trip"
);
}
#[test]
fn parse_selected_upstream_empty_is_empty() {
let m: ConditionMatch = serde_yaml::from_str("selected_upstream: {}").unwrap();
let su = m.selected_upstream.expect("empty map still parses");
assert!(su.is_empty(), "an all-absent selected_upstream reports empty");
}
#[test]
fn reject_unknown_selected_upstream_field() {
let err = serde_yaml::from_str::<ConditionMatch>(
r#"
selected_upstream:
application_flavor: spicy
"#,
)
.unwrap_err();
assert!(
err.to_string().contains("application_flavor"),
"unknown selected_upstream field should be rejected: {err}"
);
}
#[test]
fn parse_when_selected_upstream_condition() {
let conditions: Vec<Condition> = serde_yaml::from_str(
r#"
- when:
selected_upstream:
application_provider: vllm
"#,
)
.unwrap();
assert_eq!(conditions.len(), 1, "should parse 1 condition");
assert!(
matches!(
&conditions[0],
Condition::When(m)
if m.selected_upstream.as_ref().and_then(|su| su.application_provider.as_deref()) == Some("vllm")
),
"should be When gating on selected_upstream provider"
);
}
}