mod endpoint;
mod health_check;
mod load_balancer_strategy;
mod retry_policy;
mod upstream_authority;
use std::{fmt, sync::Arc};
pub use endpoint::Endpoint;
pub use health_check::{HealthCheckConfig, HealthCheckType};
pub use load_balancer_strategy::{
ConsistentHashOpts, HashFunction, LoadBalancerStrategy, MaglevOpts, ParameterisedStrategy, PriorityOpts,
RingHashOpts, SimpleStrategy, SubsetFallbackPolicy, SubsetOpts, ZoneAwareOpts,
};
pub use retry_policy::{
BackoffConfig, BudgetPercent, DEFAULT_MAX_RETRIES, DEFAULT_RETRY_BODY_LIMIT_BYTES, HttpStatusCode,
MAX_EFFECTIVE_RETRIES, MAX_RETRY_BODY_LIMIT_BYTES, RetriableCondition, RetryBodyLimit, RetryBudgetConfig,
RetryPolicy,
};
use serde::{Deserialize, Serialize};
pub use upstream_authority::{AuthoritySource, UpstreamAuthority};
use crate::errors::ProxyError;
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum UpstreamHttpVersion {
#[default]
H1,
H2,
Auto,
}
impl fmt::Display for UpstreamHttpVersion {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::H1 => f.write_str("h1"),
Self::H2 => f.write_str("h2"),
Self::Auto => f.write_str("auto"),
}
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ClusterHttpOptions {
#[serde(default)]
pub authority: Option<UpstreamAuthority>,
#[serde(default)]
pub application_protocol: Option<Arc<str>>,
#[serde(default)]
pub application_provider: Option<Arc<str>>,
#[serde(default)]
pub version: UpstreamHttpVersion,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct Cluster {
pub name: Arc<str>,
#[serde(default)]
pub http: ClusterHttpOptions,
#[serde(default)]
pub connection_timeout_ms: Option<u64>,
pub endpoints: Vec<Endpoint>,
#[serde(default)]
pub health_check: Option<HealthCheckConfig>,
#[serde(default)]
pub idle_timeout_ms: Option<u64>,
#[serde(default)]
pub load_balancer_strategy: LoadBalancerStrategy,
#[serde(default)]
pub max_connections: Option<u32>,
#[serde(default)]
pub read_timeout_ms: Option<u64>,
#[serde(default)]
pub tls: Option<praxis_tls::ClusterTls>,
#[serde(default)]
pub total_connection_timeout_ms: Option<u64>,
#[serde(default)]
pub trusted_private_endpoints: Vec<String>,
#[serde(default)]
pub write_timeout_ms: Option<u64>,
#[serde(default)]
pub retry_policy: Option<RetryPolicy>,
}
impl Cluster {
pub fn validate_authority(&self) -> Result<(), ProxyError> {
match &self.http.authority {
Some(UpstreamAuthority::Literal(authority)) => {
super::validate::cluster::validate_authority(authority, &self.name)
},
Some(UpstreamAuthority::Derived { .. }) | None => Ok(()),
}
}
pub fn with_defaults(name: &str, endpoints: Vec<Endpoint>) -> Self {
Self {
name: Arc::from(name),
http: ClusterHttpOptions::default(),
connection_timeout_ms: None,
endpoints,
health_check: None,
idle_timeout_ms: None,
load_balancer_strategy: LoadBalancerStrategy::default(),
max_connections: None,
read_timeout_ms: None,
tls: None,
total_connection_timeout_ms: None,
trusted_private_endpoints: Vec::new(),
write_timeout_ms: None,
retry_policy: None,
}
}
}
#[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,
reason = "tests use unwrap/expect/indexing/raw strings for brevity"
)]
mod tests {
use super::*;
#[test]
fn parse_cluster_minimal() {
let yaml = r#"
name: "backend"
endpoints: ["10.0.0.1:8080"]
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert_eq!(&*cluster.name, "backend", "cluster name mismatch");
assert_eq!(
cluster.endpoints[0].address(),
"10.0.0.1:8080",
"endpoint address mismatch"
);
assert_eq!(cluster.endpoints[0].weight(), 1, "default weight should be 1");
assert_eq!(
cluster.load_balancer_strategy,
LoadBalancerStrategy::default(),
"strategy should default"
);
assert!(
cluster.connection_timeout_ms.is_none(),
"connection_timeout should default to None"
);
}
#[test]
fn parse_cluster_with_weights() {
let yaml = r#"
name: "backend"
endpoints:
- "10.0.0.1:8080"
- address: "10.0.0.2:8080"
weight: 3
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert_eq!(cluster.endpoints.len(), 2, "should parse two endpoints");
assert_eq!(cluster.endpoints[0].weight(), 1, "simple endpoint weight should be 1");
assert_eq!(cluster.endpoints[1].weight(), 3, "weighted endpoint weight should be 3");
}
#[test]
fn parse_cluster_with_timeouts() {
let yaml = r#"
name: "backend"
endpoints: ["10.0.0.1:8080"]
connection_timeout_ms: 5000
idle_timeout_ms: 30000
read_timeout_ms: 10000
write_timeout_ms: 10000
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert_eq!(
cluster.connection_timeout_ms,
Some(5000),
"connection_timeout_ms mismatch"
);
assert_eq!(cluster.idle_timeout_ms, Some(30000), "idle_timeout_ms mismatch");
assert_eq!(cluster.read_timeout_ms, Some(10000), "read_timeout_ms mismatch");
assert_eq!(cluster.write_timeout_ms, Some(10000), "write_timeout_ms mismatch");
}
#[test]
fn cluster_roundtrips_via_serde() {
let cluster = Cluster {
connection_timeout_ms: Some(1000),
..Cluster::with_defaults("web", vec!["10.0.0.1:80".into()])
};
let value = serde_yaml::to_value(&cluster).unwrap();
let back: Cluster = serde_yaml::from_value(value).unwrap();
assert_eq!(back.name, cluster.name, "name should roundtrip");
assert_eq!(back.endpoints, cluster.endpoints, "endpoints should roundtrip");
assert_eq!(
back.connection_timeout_ms, cluster.connection_timeout_ms,
"timeout should roundtrip"
);
}
#[test]
fn endpoint_authority_roundtrips_via_serde() {
let cluster = Cluster {
http: ClusterHttpOptions {
authority: Some(UpstreamAuthority::Derived {
from: AuthoritySource::Endpoint,
}),
..ClusterHttpOptions::default()
},
..Cluster::with_defaults("web", vec!["api.example.com:443".into()])
};
let value = serde_yaml::to_value(&cluster).unwrap();
let back: Cluster = serde_yaml::from_value(value).unwrap();
assert_eq!(back.http, cluster.http, "endpoint-derived authority should roundtrip");
}
#[test]
fn validate_authority_skips_endpoint_authority() {
let cluster: Cluster = serde_yaml::from_str(
r#"
name: "web"
endpoints: ["api.example.com:443"]
http:
authority: { from: endpoint }
"#,
)
.unwrap();
cluster
.validate_authority()
.expect("an endpoint-derived authority has no fixed value to validate");
}
#[test]
fn validate_authority_still_checks_a_fixed_authority() {
let cluster = Cluster {
http: ClusterHttpOptions {
authority: Some("api.example.com/v1".into()),
..ClusterHttpOptions::default()
},
..Cluster::with_defaults("web", vec!["10.0.0.1:80".into()])
};
let err = cluster.validate_authority().unwrap_err();
assert!(
err.to_string().contains("not a valid HTTP authority"),
"a fixed authority must still be validated: {err}"
);
}
#[test]
fn tls_and_sni_parse_correctly() {
let yaml = r#"
name: "backend"
endpoints: ["10.0.0.1:443"]
tls:
sni: "api.example.com"
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert!(cluster.tls.is_some(), "tls should be present");
assert_eq!(
cluster.tls.as_ref().unwrap().sni.as_deref(),
Some("api.example.com"),
"sni mismatch"
);
}
#[test]
fn tls_verify_defaults_to_true() {
let yaml = r#"
name: "backend"
endpoints: ["10.0.0.1:443"]
tls: {}
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert!(cluster.tls.as_ref().unwrap().verify, "verify should default to true");
}
#[test]
fn tls_verify_can_be_disabled() {
let yaml = r#"
name: "backend"
endpoints: ["10.0.0.1:443"]
tls:
verify: false
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert!(
!cluster.tls.as_ref().unwrap().verify,
"verify should be false when explicitly set"
);
}
#[test]
fn no_tls_by_default() {
let cluster = Cluster::with_defaults("web", vec!["10.0.0.1:80".into()]);
assert!(cluster.tls.is_none(), "tls should be None by default");
}
#[test]
fn parse_cluster_application_metadata() {
let yaml = r#"
name: "backend"
endpoints: ["10.0.0.1:8080"]
http:
application_protocol: openai_chat_completions
application_provider: vllm
"#;
let cluster: Cluster = serde_yaml::from_str(yaml).unwrap();
assert_eq!(
cluster.http.application_protocol.as_deref(),
Some("openai_chat_completions"),
"application_protocol mismatch"
);
assert_eq!(
cluster.http.application_provider.as_deref(),
Some("vllm"),
"application_provider mismatch"
);
}
#[test]
fn application_metadata_defaults_to_none() {
let cluster = Cluster::with_defaults("web", vec!["10.0.0.1:80".into()]);
assert!(
cluster.http.application_protocol.is_none(),
"application_protocol should default to None"
);
assert!(
cluster.http.application_provider.is_none(),
"application_provider should default to None"
);
}
}