use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use crate::attributes::{AttributeBag, AttributeValue};
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
pub struct CandidateConstraint {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_models: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub deny_models: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_regions: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_sites: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_cost_tier: Option<String>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub custom: BTreeMap<String, String>,
#[serde(default)]
pub on_empty: OnEmpty,
}
impl CandidateConstraint {
pub fn is_empty(&self) -> bool {
self.allow_models.is_none()
&& self.deny_models.is_empty()
&& self.allow_regions.is_none()
&& self.allow_sites.is_none()
&& self.max_cost_tier.is_none()
&& self.custom.is_empty()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum OnEmpty {
#[default]
Deny,
Fallback,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum StringSetSpec {
Ref(String),
Literal(Vec<String>),
}
impl StringSetSpec {
fn resolve(&self, bag: &AttributeBag) -> Option<Vec<String>> {
match self {
StringSetSpec::Literal(v) => Some(v.clone()),
StringSetSpec::Ref(path) => {
let key = bag.resolve_key(path)?;
match bag.get(&key) {
Some(AttributeValue::StringSet(s)) => {
let mut v: Vec<String> = s.iter().cloned().collect();
v.sort();
Some(v)
},
Some(AttributeValue::String(s)) => Some(vec![s.clone()]),
_ => None,
}
},
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
#[error(
"restrict: `{field}` reference `{path}` did not resolve to a set — denying \
(a deny-list cannot fail open)"
)]
pub struct RestrictResolveError {
pub field: &'static str,
pub path: String,
}
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
pub struct RestrictSpec {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_models: Option<StringSetSpec>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub deny_models: Option<StringSetSpec>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_regions: Option<StringSetSpec>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_sites: Option<StringSetSpec>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_cost_tier: Option<String>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub custom: BTreeMap<String, String>,
#[serde(default)]
pub on_empty: OnEmpty,
}
impl RestrictSpec {
pub fn is_empty(&self) -> bool {
self.allow_models.is_none()
&& self.deny_models.is_none()
&& self.allow_regions.is_none()
&& self.allow_sites.is_none()
&& self.max_cost_tier.is_none()
&& self.custom.is_empty()
}
pub fn resolve(&self, bag: &AttributeBag) -> Result<CandidateConstraint, RestrictResolveError> {
let allow_models = self
.allow_models
.as_ref()
.map(|s| s.resolve(bag).unwrap_or_default());
let allow_regions = self
.allow_regions
.as_ref()
.map(|s| s.resolve(bag).unwrap_or_default());
let allow_sites = self
.allow_sites
.as_ref()
.map(|s| s.resolve(bag).unwrap_or_default());
let deny_models = match self.deny_models.as_ref() {
None => Vec::new(),
Some(spec) => match spec.resolve(bag) {
Some(v) => v,
None => {
let path = match spec {
StringSetSpec::Ref(p) => p.clone(),
StringSetSpec::Literal(_) => String::new(),
};
return Err(RestrictResolveError {
field: "deny_models",
path,
});
},
},
};
Ok(CandidateConstraint {
allow_models,
deny_models,
allow_regions,
allow_sites,
max_cost_tier: self.max_cost_tier.clone(),
custom: self.custom.clone(),
on_empty: self.on_empty,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::attributes::AttributeBag;
#[test]
fn candidate_constraint_is_empty_covers_every_field() {
let CandidateConstraint {
allow_models,
deny_models,
allow_regions,
allow_sites,
max_cost_tier,
custom,
on_empty: _, } = CandidateConstraint::default();
assert!(allow_models.is_none());
assert!(deny_models.is_empty());
assert!(allow_regions.is_none());
assert!(allow_sites.is_none());
assert!(max_cost_tier.is_none());
assert!(custom.is_empty());
assert!(CandidateConstraint::default().is_empty());
let each: [CandidateConstraint; 6] = [
CandidateConstraint {
allow_models: Some(vec![]),
..Default::default()
},
CandidateConstraint {
deny_models: vec!["x".into()],
..Default::default()
},
CandidateConstraint {
allow_regions: Some(vec![]),
..Default::default()
},
CandidateConstraint {
allow_sites: Some(vec![]),
..Default::default()
},
CandidateConstraint {
max_cost_tier: Some("cheap".into()),
..Default::default()
},
CandidateConstraint {
custom: [("k".to_string(), "v".to_string())].into(),
..Default::default()
},
];
for c in each {
assert!(!c.is_empty(), "field should count toward non-empty: {c:?}");
}
assert!(CandidateConstraint {
on_empty: OnEmpty::Fallback,
..Default::default()
}
.is_empty());
}
#[test]
fn restrict_spec_is_empty_covers_every_field() {
let RestrictSpec {
allow_models,
deny_models,
allow_regions,
allow_sites,
max_cost_tier,
custom,
on_empty: _,
} = RestrictSpec::default();
assert!(allow_models.is_none());
assert!(deny_models.is_none());
assert!(allow_regions.is_none());
assert!(allow_sites.is_none());
assert!(max_cost_tier.is_none());
assert!(custom.is_empty());
assert!(RestrictSpec::default().is_empty());
let lit = |s: &str| Some(StringSetSpec::Literal(vec![s.to_string()]));
let each: [RestrictSpec; 6] = [
RestrictSpec {
allow_models: lit("m"),
..Default::default()
},
RestrictSpec {
deny_models: lit("m"),
..Default::default()
},
RestrictSpec {
allow_regions: lit("eu"),
..Default::default()
},
RestrictSpec {
allow_sites: lit("s"),
..Default::default()
},
RestrictSpec {
max_cost_tier: Some("cheap".into()),
..Default::default()
},
RestrictSpec {
custom: [("k".to_string(), "v".to_string())].into(),
..Default::default()
},
];
for s in each {
assert!(!s.is_empty(), "field should count toward non-empty: {s:?}");
}
assert!(RestrictSpec {
on_empty: OnEmpty::Fallback,
..Default::default()
}
.is_empty());
}
#[test]
fn nonempty_spec_resolves_to_nonempty_constraint() {
let bag = AttributeBag::new();
let lit = |s: &str| Some(StringSetSpec::Literal(vec![s.to_string()]));
let specs: [RestrictSpec; 6] = [
RestrictSpec {
allow_models: lit("m"),
..Default::default()
},
RestrictSpec {
deny_models: lit("m"),
..Default::default()
},
RestrictSpec {
allow_regions: lit("eu"),
..Default::default()
},
RestrictSpec {
allow_sites: lit("s"),
..Default::default()
},
RestrictSpec {
max_cost_tier: Some("cheap".into()),
..Default::default()
},
RestrictSpec {
custom: [("k".to_string(), "v".to_string())].into(),
..Default::default()
},
];
for spec in specs {
assert!(!spec.is_empty(), "spec should be non-empty: {spec:?}");
let c = spec.resolve(&bag).expect("literal spec always resolves");
assert!(
!c.is_empty(),
"non-empty spec resolved to a dropped constraint: {spec:?}"
);
}
}
}