use std::collections::BTreeSet;
use chrono::{DateTime, Utc};
use serde::{de, Deserialize, Deserializer, Serialize};
use sha2::{Digest, Sha256};
use crate::error::WeightsError;
pub const CARD_VERSION_V1: &str = "1";
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize)]
#[serde(transparent)]
pub struct StringSet(BTreeSet<String>);
impl<'de> Deserialize<'de> for StringSet {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let items = Vec::<String>::deserialize(deserializer)?;
let mut set = BTreeSet::new();
for item in items {
if item.trim().is_empty() {
return Err(de::Error::custom(
"string set entries must be non-empty strings",
));
}
if item.trim() != item {
return Err(de::Error::custom(
"string set entries must not contain surrounding whitespace",
));
}
if set.contains(&item) {
return Err(de::Error::custom(format!(
"duplicate string set entry {item:?}"
)));
}
set.insert(item);
}
Ok(Self(set))
}
}
impl StringSet {
pub fn new<I, S>(items: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
Self(items.into_iter().map(Into::into).collect())
}
#[must_use]
pub fn covers(&self, subset: &StringSet) -> bool {
subset.0.is_subset(&self.0)
}
#[must_use]
pub fn intersects(&self, other: &StringSet) -> bool {
self.0.intersection(&other.0).next().is_some()
}
#[must_use]
pub fn contains(&self, item: &str) -> bool {
self.0.contains(item)
}
pub fn iter(&self) -> impl Iterator<Item = &str> {
self.0.iter().map(String::as_str)
}
#[must_use]
pub fn len(&self) -> usize {
self.0.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
#[must_use]
pub fn as_set(&self) -> &BTreeSet<String> {
&self.0
}
fn validate_entries(&self, field: &'static str) -> Result<(), WeightsError> {
for item in &self.0 {
if item.trim().is_empty() {
return Err(WeightsError::SchemaRejected(format!(
"{field} entries must be non-empty strings"
)));
}
if item.trim() != item {
return Err(WeightsError::SchemaRejected(format!(
"{field} entries must not contain surrounding whitespace"
)));
}
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ModelCard {
pub card_version: String,
pub weights_hash: String,
pub allowed_capability_set: StringSet,
pub banned_tools: StringSet,
pub training_data_class: String,
pub issuer: String,
pub issued_at: DateTime<Utc>,
pub expires_at: DateTime<Utc>,
}
impl ModelCard {
pub fn new(
weights_hash: impl Into<String>,
allowed_capability_set: StringSet,
banned_tools: StringSet,
training_data_class: impl Into<String>,
issuer: impl Into<String>,
issued_at: DateTime<Utc>,
expires_at: DateTime<Utc>,
) -> Result<Self, WeightsError> {
let card = Self {
card_version: CARD_VERSION_V1.to_string(),
weights_hash: weights_hash.into(),
allowed_capability_set,
banned_tools,
training_data_class: training_data_class.into(),
issuer: issuer.into(),
issued_at,
expires_at,
};
card.validate()?;
Ok(card)
}
pub fn validate(&self) -> Result<(), WeightsError> {
if self.card_version != CARD_VERSION_V1 {
return Err(WeightsError::SchemaRejected(format!(
"card_version must be {CARD_VERSION_V1:?}, got {:?}",
self.card_version
)));
}
if !is_lowercase_sha256_hex(&self.weights_hash) {
return Err(WeightsError::SchemaRejected(format!(
"weights_hash must be 64 lowercase hex chars, got {:?}",
self.weights_hash
)));
}
self.allowed_capability_set
.validate_entries("allowed_capability_set")?;
self.banned_tools.validate_entries("banned_tools")?;
validate_required_text_field(&self.training_data_class, "training_data_class")?;
validate_required_text_field(&self.issuer, "issuer")?;
if self.expires_at < self.issued_at {
return Err(WeightsError::SchemaRejected(format!(
"expires_at ({}) precedes issued_at ({})",
self.expires_at, self.issued_at,
)));
}
Ok(())
}
pub fn require_live(&self, now: DateTime<Utc>) -> Result<(), WeightsError> {
if now < self.expires_at {
Ok(())
} else {
Err(WeightsError::Expired {
expires_at: self.expires_at,
now,
})
}
}
pub fn to_canonical_json(&self) -> Result<Vec<u8>, WeightsError> {
chio_core_types::canonical::canonical_json_bytes(self)
.map_err(|err| WeightsError::Encoding(format!("canonical-json encode: {err}")))
}
pub fn from_canonical_json(bytes: &[u8]) -> Result<Self, WeightsError> {
let card: Self = serde_json::from_slice(bytes)
.map_err(|err| WeightsError::Encoding(format!("canonical-json decode: {err}")))?;
card.validate()?;
let canonical = card.to_canonical_json()?;
if canonical.as_slice() != bytes {
return Err(WeightsError::Encoding(
"model card bytes are not RFC 8785 canonical JSON".to_string(),
));
}
Ok(card)
}
}
#[must_use]
pub fn weights_hash_of(bytes: &[u8]) -> String {
let digest = Sha256::digest(bytes);
hex::encode(digest)
}
#[inline]
fn is_lowercase_hex_byte(b: u8) -> bool {
matches!(b, b'0'..=b'9' | b'a'..=b'f')
}
#[inline]
fn is_lowercase_sha256_hex(value: &str) -> bool {
value.len() == 64 && value.bytes().all(is_lowercase_hex_byte)
}
fn validate_required_text_field(value: &str, field: &'static str) -> Result<(), WeightsError> {
if value.trim().is_empty() {
return Err(WeightsError::MissingField(field));
}
if value.trim() != value {
return Err(WeightsError::SchemaRejected(format!(
"{field} must not contain surrounding whitespace"
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::TimeZone;
fn fixed_issued_at() -> DateTime<Utc> {
match Utc.with_ymd_and_hms(2026, 4, 30, 12, 0, 0) {
chrono::LocalResult::Single(t) => t,
_ => panic!("fixed_issued_at fixture must construct"),
}
}
fn good_card() -> ModelCard {
let issued = fixed_issued_at();
let expires = issued + chrono::Duration::days(30);
match ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::new(["tool:read", "tool:write"]),
StringSet::new(["tool:exec"]),
"public-internet",
"https://example.com/issuer",
issued,
expires,
) {
Ok(card) => card,
Err(e) => panic!("good_card must construct: {e}"),
}
}
fn good_card_value() -> serde_json::Value {
serde_json::json!({
"allowed_capability_set": ["tool:read", "tool:write"],
"banned_tools": ["tool:exec"],
"card_version": CARD_VERSION_V1,
"expires_at": "2026-05-30T12:00:00Z",
"issued_at": "2026-04-30T12:00:00Z",
"issuer": "https://example.com/issuer",
"training_data_class": "public-internet",
"weights_hash": "0000000000000000000000000000000000000000000000000000000000000001",
})
}
fn value_to_bytes(value: &serde_json::Value) -> Vec<u8> {
match serde_json::to_vec(value) {
Ok(bytes) => bytes,
Err(e) => panic!("serialize fixture value: {e}"),
}
}
#[test]
fn lowercase_sha256_hex_helper_accepts_only_exact_lowercase_digest() {
assert!(is_lowercase_sha256_hex(
"0000000000000000000000000000000000000000000000000000000000000001"
));
assert!(!is_lowercase_sha256_hex(
"ABCDEF0000000000000000000000000000000000000000000000000000000001"
));
assert!(!is_lowercase_sha256_hex("abcdef"));
assert!(!is_lowercase_sha256_hex(
"000000000000000000000000000000000000000000000000000000000000000g"
));
}
#[test]
fn validate_rejects_uppercase_weights_hash() {
let issued = fixed_issued_at();
let res = ModelCard::new(
"ABCDEF0000000000000000000000000000000000000000000000000000000001",
StringSet::new(["tool:read"]),
StringSet::default(),
"public-internet",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(res, Err(WeightsError::SchemaRejected(_))));
}
#[test]
fn validate_rejects_short_weights_hash() {
let issued = fixed_issued_at();
let res = ModelCard::new(
"abcdef",
StringSet::default(),
StringSet::default(),
"public-internet",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(res, Err(WeightsError::SchemaRejected(_))));
}
#[test]
fn validate_rejects_empty_training_data_class() {
let issued = fixed_issued_at();
let res = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::default(),
StringSet::default(),
"",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(res, Err(WeightsError::MissingField(_))));
}
#[test]
fn validate_rejects_blank_or_padded_required_text_fields() {
let issued = fixed_issued_at();
let blank_training = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::default(),
StringSet::default(),
" ",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(
blank_training,
Err(WeightsError::MissingField("training_data_class"))
));
let padded_issuer = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::default(),
StringSet::default(),
"public-internet",
" https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(
padded_issuer,
Err(WeightsError::SchemaRejected(message)) if message.contains("issuer")
));
}
#[test]
fn validate_rejects_expires_before_issued() {
let issued = fixed_issued_at();
let res = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::default(),
StringSet::default(),
"public-internet",
"https://example.com/issuer",
issued,
issued - chrono::Duration::seconds(1),
);
assert!(matches!(res, Err(WeightsError::SchemaRejected(_))));
}
#[test]
fn from_canonical_json_rejects_unknown_fields() {
let mut value = good_card_value();
match value {
serde_json::Value::Object(ref mut map) => {
map.insert(
"unexpected".to_string(),
serde_json::Value::String("field".to_string()),
);
}
_ => panic!("fixture must be a JSON object"),
}
let bytes = value_to_bytes(&value);
let res = ModelCard::from_canonical_json(&bytes);
assert!(matches!(res, Err(WeightsError::Encoding(_))));
}
#[test]
fn from_canonical_json_rejects_duplicate_set_entries() {
let mut value = good_card_value();
match value {
serde_json::Value::Object(ref mut map) => {
map.insert(
"allowed_capability_set".to_string(),
serde_json::json!(["tool:read", "tool:read"]),
);
}
_ => panic!("fixture must be a JSON object"),
}
let bytes = value_to_bytes(&value);
let res = ModelCard::from_canonical_json(&bytes);
assert!(matches!(res, Err(WeightsError::Encoding(_))));
}
#[test]
fn from_canonical_json_rejects_empty_set_entries() {
let mut value = good_card_value();
match value {
serde_json::Value::Object(ref mut map) => {
map.insert(
"banned_tools".to_string(),
serde_json::json!(["tool:exec", ""]),
);
}
_ => panic!("fixture must be a JSON object"),
}
let bytes = value_to_bytes(&value);
let res = ModelCard::from_canonical_json(&bytes);
assert!(matches!(res, Err(WeightsError::Encoding(_))));
}
#[test]
fn new_rejects_empty_set_entries() {
let issued = fixed_issued_at();
let res = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::new([""]),
StringSet::default(),
"public-internet",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(res, Err(WeightsError::SchemaRejected(_))));
}
#[test]
fn new_rejects_blank_or_padded_set_entries() {
let issued = fixed_issued_at();
let blank = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::new([" "]),
StringSet::default(),
"public-internet",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(
blank,
Err(WeightsError::SchemaRejected(message)) if message.contains("allowed_capability_set")
));
let padded = ModelCard::new(
"0000000000000000000000000000000000000000000000000000000000000001",
StringSet::default(),
StringSet::new([" tool:exec"]),
"public-internet",
"https://example.com/issuer",
issued,
issued + chrono::Duration::days(1),
);
assert!(matches!(
padded,
Err(WeightsError::SchemaRejected(message)) if message.contains("banned_tools")
));
}
#[test]
fn require_live_rejects_expired_card() {
let card = good_card();
let later = card.expires_at + chrono::Duration::seconds(1);
let res = card.require_live(later);
assert!(matches!(res, Err(WeightsError::Expired { .. })));
}
#[test]
fn canonical_round_trip_is_byte_stable() {
let card = good_card();
let bytes = match card.to_canonical_json() {
Ok(b) => b,
Err(e) => panic!("encode must succeed: {e}"),
};
let again = match ModelCard::from_canonical_json(&bytes) {
Ok(c) => c,
Err(e) => panic!("decode must succeed: {e}"),
};
let bytes2 = match again.to_canonical_json() {
Ok(b) => b,
Err(e) => panic!("re-encode must succeed: {e}"),
};
assert_eq!(bytes, bytes2, "canonical-json bytes must round-trip stably");
}
#[test]
fn from_canonical_json_rejects_noncanonical_whitespace() {
let card = good_card();
let pretty = match serde_json::to_vec_pretty(&card) {
Ok(bytes) => bytes,
Err(e) => panic!("pretty encode must succeed: {e}"),
};
let res = ModelCard::from_canonical_json(&pretty);
assert!(
matches!(res, Err(WeightsError::Encoding(_))),
"signed model cards must be presented as RFC 8785 canonical bytes"
);
}
#[test]
fn weights_hash_of_matches_known_vector() {
assert_eq!(
weights_hash_of(b""),
"e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
);
}
#[test]
fn allowed_capability_set_covers_subset() {
let card = good_card();
let req = StringSet::new(["tool:read"]);
assert!(card.allowed_capability_set.covers(&req));
let bad = StringSet::new(["tool:read", "tool:admin"]);
assert!(!card.allowed_capability_set.covers(&bad));
}
#[test]
fn banned_tools_intersects_detects_overlap() {
let card = good_card();
let req = StringSet::new(["tool:read", "tool:exec"]);
assert!(card.banned_tools.intersects(&req));
let benign = StringSet::new(["tool:read"]);
assert!(!card.banned_tools.intersects(&benign));
}
}