use thiserror::Error;
#[derive(Debug, Clone, PartialEq, Error)]
pub enum ProfileValidationError {
#[error("profile field `{field}` cannot be empty")]
EmptyField {
field: &'static str,
},
#[error("bundle SHA-256 must contain exactly 64 lowercase hexadecimal characters")]
InvalidBundleSha256,
#[error("calibration temperature must be finite and greater than zero")]
InvalidCalibrationTemperature,
#[error("policy threshold must be finite and lie in [0, 1]")]
InvalidPolicyThreshold,
#[error("{field} must be finite and non-negative")]
InvalidTolerance {
field: &'static str,
},
}
fn required(
value: impl Into<String>,
field: &'static str,
) -> Result<String, ProfileValidationError> {
let value = value.into();
if value.trim().is_empty() {
Err(ProfileValidationError::EmptyField { field })
} else {
Ok(value)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ArtifactIdentity {
id: String,
revision: String,
}
impl ArtifactIdentity {
pub fn new(
id: impl Into<String>,
revision: impl Into<String>,
) -> Result<Self, ProfileValidationError> {
Ok(Self {
id: required(id, "artifact.id")?,
revision: required(revision, "artifact.revision")?,
})
}
pub fn id(&self) -> &str {
&self.id
}
pub fn revision(&self) -> &str {
&self.revision
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProfileSource {
profile_id: String,
repository: String,
repository_revision: String,
bundle_sha256: String,
}
impl ProfileSource {
pub fn new(
profile_id: impl Into<String>,
repository: impl Into<String>,
repository_revision: impl Into<String>,
bundle_sha256: impl Into<String>,
) -> Result<Self, ProfileValidationError> {
let bundle_sha256 = bundle_sha256.into();
if bundle_sha256.len() != 64
|| !bundle_sha256
.bytes()
.all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
{
return Err(ProfileValidationError::InvalidBundleSha256);
}
Ok(Self {
profile_id: required(profile_id, "profile_id")?,
repository: required(repository, "source.repository")?,
repository_revision: required(repository_revision, "source.repository_revision")?,
bundle_sha256,
})
}
pub fn profile_id(&self) -> &str {
&self.profile_id
}
pub fn repository(&self) -> &str {
&self.repository
}
pub fn repository_revision(&self) -> &str {
&self.repository_revision
}
pub fn bundle_sha256(&self) -> &str {
&self.bundle_sha256
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProbabilitySpace {
ConditionalOnOfferedOptions,
OfferedOptionsPlusSemanticNone,
}
impl ProbabilitySpace {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::ConditionalOnOfferedOptions => "conditional_on_offered_options",
Self::OfferedOptionsPlusSemanticNone => "offered_options_plus_semantic_none",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ExecutionSemantics {
renderer: String,
head: String,
rejection: String,
probability_space: ProbabilitySpace,
}
impl ExecutionSemantics {
pub fn new(
renderer: impl Into<String>,
head: impl Into<String>,
rejection: impl Into<String>,
probability_space: ProbabilitySpace,
) -> Result<Self, ProfileValidationError> {
Ok(Self {
renderer: required(renderer, "execution.renderer")?,
head: required(head, "execution.head")?,
rejection: required(rejection, "execution.rejection")?,
probability_space,
})
}
pub fn renderer(&self) -> &str {
&self.renderer
}
pub fn head(&self) -> &str {
&self.head
}
pub fn rejection(&self) -> &str {
&self.rejection
}
#[must_use]
pub const fn probability_space(&self) -> ProbabilitySpace {
self.probability_space
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ParityContract {
calibration_temperature: f64,
policy_threshold: f64,
probability_tolerance: f64,
ordering_tolerance: f64,
}
impl ParityContract {
pub fn new(
calibration_temperature: f64,
policy_threshold: f64,
probability_tolerance: f64,
ordering_tolerance: f64,
) -> Result<Self, ProfileValidationError> {
if !calibration_temperature.is_finite() || calibration_temperature <= 0.0 {
return Err(ProfileValidationError::InvalidCalibrationTemperature);
}
if !policy_threshold.is_finite() || !(0.0..=1.0).contains(&policy_threshold) {
return Err(ProfileValidationError::InvalidPolicyThreshold);
}
for (field, value) in [
("probability_tolerance", probability_tolerance),
("ordering_tolerance", ordering_tolerance),
] {
if !value.is_finite() || value < 0.0 {
return Err(ProfileValidationError::InvalidTolerance { field });
}
}
Ok(Self {
calibration_temperature,
policy_threshold,
probability_tolerance,
ordering_tolerance,
})
}
pub fn calibration_temperature(&self) -> f64 {
self.calibration_temperature
}
pub fn policy_threshold(&self) -> f64 {
self.policy_threshold
}
pub fn probability_tolerance(&self) -> f64 {
self.probability_tolerance
}
pub fn ordering_tolerance(&self) -> f64 {
self.ordering_tolerance
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelExecutionProfile {
source: ProfileSource,
backbone: ArtifactIdentity,
execution: ExecutionSemantics,
parity: ParityContract,
}
impl ModelExecutionProfile {
pub fn new(
source: ProfileSource,
backbone: ArtifactIdentity,
execution: ExecutionSemantics,
parity: ParityContract,
) -> Self {
Self {
source,
backbone,
execution,
parity,
}
}
pub fn source(&self) -> &ProfileSource {
&self.source
}
pub fn backbone(&self) -> &ArtifactIdentity {
&self.backbone
}
pub fn execution(&self) -> &ExecutionSemantics {
&self.execution
}
pub fn parity(&self) -> &ParityContract {
&self.parity
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn profile_contract_is_constructed_from_valid_components() {
let source = ProfileSource::new(
"profile",
"owner/repository",
"revision",
"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef",
)
.unwrap();
let backbone = ArtifactIdentity::new("model", "model-revision").unwrap();
let execution = ExecutionSemantics::new(
"state-first",
"head.safetensors",
"score-summary",
ProbabilitySpace::OfferedOptionsPlusSemanticNone,
)
.unwrap();
let parity = ParityContract::new(1.5, 0.98, 0.005, 0.00001).unwrap();
let profile = ModelExecutionProfile::new(source, backbone, execution, parity);
assert_eq!(profile.source().profile_id(), "profile");
assert_eq!(profile.backbone().id(), "model");
assert_eq!(profile.execution().renderer(), "state-first");
assert_eq!(profile.parity().policy_threshold(), 0.98);
assert_eq!(
profile.execution().probability_space(),
ProbabilitySpace::OfferedOptionsPlusSemanticNone
);
assert_eq!(
ProbabilitySpace::ConditionalOnOfferedOptions.as_str(),
"conditional_on_offered_options"
);
assert_eq!(
ProbabilitySpace::OfferedOptionsPlusSemanticNone.as_str(),
"offered_options_plus_semantic_none"
);
}
#[test]
fn profile_components_reject_invalid_values() {
assert!(matches!(
ArtifactIdentity::new("", "revision"),
Err(ProfileValidationError::EmptyField { .. })
));
assert_eq!(
ProfileSource::new("profile", "repository", "revision", "ABC"),
Err(ProfileValidationError::InvalidBundleSha256)
);
assert_eq!(
ParityContract::new(0.0, 0.98, 0.005, 0.00001),
Err(ProfileValidationError::InvalidCalibrationTemperature)
);
assert_eq!(
ParityContract::new(1.0, 1.1, 0.005, 0.00001),
Err(ProfileValidationError::InvalidPolicyThreshold)
);
assert!(matches!(
ParityContract::new(1.0, 0.98, f64::NAN, 0.00001),
Err(ProfileValidationError::InvalidTolerance { .. })
));
}
}