use crate::error::{KizzasiError, KizzasiResult};
use crate::{Kizzasi, KizzasiConfig};
use serde::{Deserialize, Serialize};
use std::cmp::Ordering;
use std::collections::HashMap;
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct SemanticVersion {
pub major: u32,
pub minor: u32,
pub patch: u32,
}
impl SemanticVersion {
pub fn new(major: u32, minor: u32, patch: u32) -> Self {
Self {
major,
minor,
patch,
}
}
pub fn parse(s: &str) -> KizzasiResult<Self> {
let parts: Vec<&str> = s.split('.').collect();
if parts.len() != 3 {
return Err(KizzasiError::invalid_state(format!(
"Invalid semantic version format: {}. Expected 'major.minor.patch'",
s
)));
}
let major = parts[0]
.parse()
.map_err(|_| KizzasiError::invalid_state("Invalid major version"))?;
let minor = parts[1]
.parse()
.map_err(|_| KizzasiError::invalid_state("Invalid minor version"))?;
let patch = parts[2]
.parse()
.map_err(|_| KizzasiError::invalid_state("Invalid patch version"))?;
Ok(Self {
major,
minor,
patch,
})
}
pub fn is_compatible_with(&self, other: &Self) -> bool {
self.major == other.major
}
}
impl fmt::Display for SemanticVersion {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}.{}.{}", self.major, self.minor, self.patch)
}
}
impl PartialOrd for SemanticVersion {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for SemanticVersion {
fn cmp(&self, other: &Self) -> Ordering {
match self.major.cmp(&other.major) {
Ordering::Equal => match self.minor.cmp(&other.minor) {
Ordering::Equal => self.patch.cmp(&other.patch),
other => other,
},
other => other,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub enum DeploymentStrategy {
Immediate,
Canary { traffic_percent: u8 },
BlueGreen,
Rolling { batch_size: usize },
}
impl DeploymentStrategy {
pub fn should_use_new_version(&self, request_id: u64) -> bool {
match self {
DeploymentStrategy::Immediate => true,
DeploymentStrategy::Canary { traffic_percent } => {
(request_id % 100) < (*traffic_percent as u64)
}
DeploymentStrategy::BlueGreen => false, DeploymentStrategy::Rolling { .. } => false, }
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelMetadata {
pub description: String,
pub changelog: Vec<String>,
pub author: Option<String>,
pub created_at: String,
pub deployment: DeploymentStrategy,
pub tags: Vec<String>,
pub metrics: HashMap<String, f64>,
}
impl Default for ModelMetadata {
fn default() -> Self {
Self {
description: String::new(),
changelog: Vec::new(),
author: None,
created_at: chrono::Utc::now().to_rfc3339(),
deployment: DeploymentStrategy::Immediate,
tags: Vec::new(),
metrics: HashMap::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelVersion {
version: SemanticVersion,
config: KizzasiConfig,
metadata: ModelMetadata,
}
impl ModelVersion {
pub fn new(version: &str, config: KizzasiConfig, description: &str) -> KizzasiResult<Self> {
let version = SemanticVersion::parse(version)?;
let metadata = ModelMetadata {
description: description.to_string(),
..Default::default()
};
Ok(Self {
version,
config,
metadata,
})
}
pub fn with_metadata(
version: &str,
config: KizzasiConfig,
metadata: ModelMetadata,
) -> KizzasiResult<Self> {
let version = SemanticVersion::parse(version)?;
Ok(Self {
version,
config,
metadata,
})
}
pub fn version(&self) -> &SemanticVersion {
&self.version
}
pub fn config(&self) -> &KizzasiConfig {
&self.config
}
pub fn metadata(&self) -> &ModelMetadata {
&self.metadata
}
pub fn metadata_mut(&mut self) -> &mut ModelMetadata {
&mut self.metadata
}
pub fn create_predictor(&self) -> KizzasiResult<Kizzasi> {
Kizzasi::new(self.config.clone())
}
pub fn is_compatible_with(&self, other: &Self) -> bool {
self.version.is_compatible_with(&other.version)
}
}
#[derive(Debug, Clone, Default)]
pub struct ModelRegistry {
versions: HashMap<SemanticVersion, ModelVersion>,
active_version: Option<SemanticVersion>,
canary_version: Option<SemanticVersion>,
}
impl ModelRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, model: ModelVersion) -> KizzasiResult<()> {
let version = *model.version();
if self.versions.contains_key(&version) {
return Err(KizzasiError::invalid_state(format!(
"Version {} already exists",
version
)));
}
self.versions.insert(version, model);
if self.active_version.is_none() {
self.active_version = Some(version);
}
Ok(())
}
pub fn get(&self, version: &str) -> KizzasiResult<&ModelVersion> {
let version = SemanticVersion::parse(version)?;
self.versions
.get(&version)
.ok_or_else(|| KizzasiError::invalid_state(format!("Version {} not found", version)))
}
pub fn get_mut(&mut self, version: &str) -> KizzasiResult<&mut ModelVersion> {
let version = SemanticVersion::parse(version)?;
self.versions
.get_mut(&version)
.ok_or_else(|| KizzasiError::invalid_state(format!("Version {} not found", version)))
}
pub fn get_latest(&self) -> KizzasiResult<&ModelVersion> {
self.versions
.keys()
.max()
.and_then(|v| self.versions.get(v))
.ok_or_else(|| KizzasiError::invalid_state("No versions registered"))
}
pub fn get_active(&self) -> KizzasiResult<&ModelVersion> {
let version = self
.active_version
.ok_or_else(|| KizzasiError::invalid_state("No active version"))?;
self.versions
.get(&version)
.ok_or_else(|| KizzasiError::invalid_state("Active version not found"))
}
pub fn get_canary(&self) -> Option<&ModelVersion> {
self.canary_version.and_then(|v| self.versions.get(&v))
}
pub fn list_versions(&self) -> Vec<&ModelVersion> {
let mut versions: Vec<_> = self.versions.values().collect();
versions.sort_by_key(|v| v.version());
versions
}
pub fn deploy(&mut self, version: &str, strategy: DeploymentStrategy) -> KizzasiResult<()> {
let version = SemanticVersion::parse(version)?;
if !self.versions.contains_key(&version) {
return Err(KizzasiError::invalid_state(format!(
"Version {} not found",
version
)));
}
match strategy {
DeploymentStrategy::Immediate => {
self.active_version = Some(version);
self.canary_version = None;
}
DeploymentStrategy::Canary { .. } => {
self.canary_version = Some(version);
}
DeploymentStrategy::BlueGreen => {
self.canary_version = Some(version);
}
DeploymentStrategy::Rolling { .. } => {
self.active_version = Some(version);
self.canary_version = None;
}
}
if let Some(model) = self.versions.get_mut(&version) {
model.metadata.deployment = strategy;
}
Ok(())
}
pub fn promote_canary(&mut self) -> KizzasiResult<()> {
let canary = self
.canary_version
.ok_or_else(|| KizzasiError::invalid_state("No canary version deployed"))?;
self.active_version = Some(canary);
self.canary_version = None;
Ok(())
}
pub fn rollback(&mut self, version: &str) -> KizzasiResult<()> {
let version = SemanticVersion::parse(version)?;
if !self.versions.contains_key(&version) {
return Err(KizzasiError::invalid_state(format!(
"Version {} not found",
version
)));
}
self.active_version = Some(version);
self.canary_version = None;
Ok(())
}
pub fn select_for_request(&self, request_id: u64) -> KizzasiResult<&ModelVersion> {
if let Some(canary_version) = self.canary_version {
if let Some(canary) = self.versions.get(&canary_version) {
if canary
.metadata
.deployment
.should_use_new_version(request_id)
{
return Ok(canary);
}
}
}
self.get_active()
}
pub fn remove(&mut self, version: &str) -> KizzasiResult<ModelVersion> {
let version = SemanticVersion::parse(version)?;
if Some(version) == self.active_version {
return Err(KizzasiError::invalid_state("Cannot remove active version"));
}
if Some(version) == self.canary_version {
return Err(KizzasiError::invalid_state("Cannot remove canary version"));
}
self.versions
.remove(&version)
.ok_or_else(|| KizzasiError::invalid_state(format!("Version {} not found", version)))
}
pub fn stats(&self) -> RegistryStats {
RegistryStats {
total_versions: self.versions.len(),
active_version: self.active_version.map(|v| v.to_string()),
canary_version: self.canary_version.map(|v| v.to_string()),
latest_version: self.versions.keys().max().map(|v| v.to_string()),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegistryStats {
pub total_versions: usize,
pub active_version: Option<String>,
pub canary_version: Option<String>,
pub latest_version: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_semantic_version() {
let v1 = SemanticVersion::new(1, 0, 0);
let v2 = SemanticVersion::new(1, 1, 0);
let v3 = SemanticVersion::new(2, 0, 0);
assert!(v1 < v2);
assert!(v2 < v3);
assert!(v1.is_compatible_with(&v2));
assert!(!v1.is_compatible_with(&v3));
}
#[test]
fn test_version_parsing() {
let v = SemanticVersion::parse("1.2.3").unwrap();
assert_eq!(v.major, 1);
assert_eq!(v.minor, 2);
assert_eq!(v.patch, 3);
assert_eq!(v.to_string(), "1.2.3");
assert!(SemanticVersion::parse("invalid").is_err());
assert!(SemanticVersion::parse("1.2").is_err());
}
#[test]
fn test_model_registry() {
let mut registry = ModelRegistry::new();
let config1 = KizzasiConfig::new().context_window(4096);
let v1 = ModelVersion::new("1.0.0", config1, "Initial").unwrap();
registry.register(v1).unwrap();
let config2 = KizzasiConfig::new().context_window(8192);
let v2 = ModelVersion::new("2.0.0", config2, "Upgrade").unwrap();
registry.register(v2).unwrap();
assert_eq!(registry.list_versions().len(), 2);
assert_eq!(
registry.get_latest().unwrap().version().to_string(),
"2.0.0"
);
assert_eq!(
registry.get_active().unwrap().version().to_string(),
"1.0.0"
);
}
#[test]
fn test_deployment_strategies() {
let mut registry = ModelRegistry::new();
let config1 = KizzasiConfig::new();
let v1 = ModelVersion::new("1.0.0", config1, "v1").unwrap();
registry.register(v1).unwrap();
let config2 = KizzasiConfig::new();
let v2 = ModelVersion::new("2.0.0", config2, "v2").unwrap();
registry.register(v2).unwrap();
registry
.deploy(
"2.0.0",
DeploymentStrategy::Canary {
traffic_percent: 10,
},
)
.unwrap();
assert_eq!(
registry.get_active().unwrap().version().to_string(),
"1.0.0"
);
assert_eq!(
registry.get_canary().unwrap().version().to_string(),
"2.0.0"
);
registry.promote_canary().unwrap();
assert_eq!(
registry.get_active().unwrap().version().to_string(),
"2.0.0"
);
assert!(registry.get_canary().is_none());
}
#[test]
fn test_traffic_splitting() {
let canary = DeploymentStrategy::Canary {
traffic_percent: 25,
};
let mut new_version_count = 0;
for i in 0..1000 {
if canary.should_use_new_version(i) {
new_version_count += 1;
}
}
assert!((200..=300).contains(&new_version_count));
}
#[test]
fn test_rollback() {
let mut registry = ModelRegistry::new();
let config1 = KizzasiConfig::new();
let v1 = ModelVersion::new("1.0.0", config1, "v1").unwrap();
registry.register(v1).unwrap();
let config2 = KizzasiConfig::new();
let v2 = ModelVersion::new("2.0.0", config2, "v2").unwrap();
registry.register(v2).unwrap();
registry
.deploy("2.0.0", DeploymentStrategy::Immediate)
.unwrap();
assert_eq!(
registry.get_active().unwrap().version().to_string(),
"2.0.0"
);
registry.rollback("1.0.0").unwrap();
assert_eq!(
registry.get_active().unwrap().version().to_string(),
"1.0.0"
);
}
#[test]
fn test_version_removal() {
let mut registry = ModelRegistry::new();
let config1 = KizzasiConfig::new();
let v1 = ModelVersion::new("1.0.0", config1, "v1").unwrap();
registry.register(v1).unwrap();
let config2 = KizzasiConfig::new();
let v2 = ModelVersion::new("2.0.0", config2, "v2").unwrap();
registry.register(v2).unwrap();
assert!(registry.remove("1.0.0").is_err());
assert!(registry.remove("2.0.0").is_ok());
assert_eq!(registry.list_versions().len(), 1);
}
}