Files
rustytorch/crates/production/rtx-hub/src/versioning.rs
T
2026-03-04 00:08:42 +00:00

594 lines
17 KiB
Rust

//! Model versioning and dependency resolution system.
use crate::{HubError, HubResult, ModelId};
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
use std::fmt;
/// Semantic version wrapper with additional metadata.
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
pub struct ModelVersion {
/// Semantic version
pub semver: semver::Version,
/// Pre-release metadata
pub prerelease: Option<String>,
/// Build metadata
pub build: Option<String>,
/// Creation timestamp
pub created_at: DateTime<Utc>,
}
impl ModelVersion {
/// Create a new model version.
pub fn new(version: semver::Version) -> Self {
Self {
semver: version,
prerelease: None,
build: None,
created_at: Utc::now(),
}
}
/// Create a new model version with prerelease metadata.
pub fn with_prerelease(mut self, prerelease: String) -> Self {
self.prerelease = Some(prerelease);
self
}
/// Create a new model version with build metadata.
pub fn with_build(mut self, build: String) -> Self {
self.build = Some(build);
self
}
/// Parse a version string.
pub fn parse(s: &str) -> HubResult<Self> {
let version = semver::Version::parse(s)?;
Ok(Self::new(version))
}
/// Check if this version satisfies a constraint.
pub fn satisfies(&self, constraint: &VersionConstraint) -> bool {
constraint.req.matches(&self.semver)
}
/// Get the major version.
pub fn major(&self) -> u64 {
self.semver.major
}
/// Get the minor version.
pub fn minor(&self) -> u64 {
self.semver.minor
}
/// Get the patch version.
pub fn patch(&self) -> u64 {
self.semver.patch
}
/// Check if this is a prerelease version.
pub fn is_prerelease(&self) -> bool {
!self.semver.pre.is_empty() || self.prerelease.is_some()
}
/// Get the next major version.
pub fn next_major(&self) -> Self {
let mut next = self.semver.clone();
next.major += 1;
next.minor = 0;
next.patch = 0;
next.pre = semver::Prerelease::EMPTY;
next.build = semver::BuildMetadata::EMPTY;
Self::new(next)
}
/// Get the next minor version.
pub fn next_minor(&self) -> Self {
let mut next = self.semver.clone();
next.minor += 1;
next.patch = 0;
next.pre = semver::Prerelease::EMPTY;
next.build = semver::BuildMetadata::EMPTY;
Self::new(next)
}
/// Get the next patch version.
pub fn next_patch(&self) -> Self {
let mut next = self.semver.clone();
next.patch += 1;
next.pre = semver::Prerelease::EMPTY;
next.build = semver::BuildMetadata::EMPTY;
Self::new(next)
}
}
impl fmt::Display for ModelVersion {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.semver)?;
if let Some(ref prerelease) = self.prerelease {
write!(f, "-{prerelease}")?;
}
if let Some(ref build) = self.build {
write!(f, "+{build}")?;
}
Ok(())
}
}
impl std::str::FromStr for ModelVersion {
type Err = HubError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Self::parse(s)
}
}
/// Version constraint for dependency resolution.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VersionConstraint {
/// Semver requirement
pub req: semver::VersionReq,
/// Original constraint string
pub constraint: String,
}
impl VersionConstraint {
/// Create a new version constraint.
pub fn new(constraint: &str) -> HubResult<Self> {
let req = semver::VersionReq::parse(constraint)?;
Ok(Self {
req,
constraint: constraint.to_string(),
})
}
/// Check if a version matches this constraint.
pub fn matches(&self, version: &ModelVersion) -> bool {
self.req.matches(&version.semver)
}
/// Get the constraint as a string.
pub fn as_str(&self) -> &str {
&self.constraint
}
}
impl fmt::Display for VersionConstraint {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.constraint)
}
}
/// Dependency specification for version resolution.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DependencySpec {
/// Model ID
pub model_id: ModelId,
/// Version constraint
pub constraint: VersionConstraint,
/// Whether this dependency is optional
pub optional: bool,
/// Dependency features or capabilities required
pub features: Vec<String>,
}
impl DependencySpec {
/// Create a new dependency specification.
pub fn new(model_id: ModelId, constraint: &str) -> HubResult<Self> {
Ok(Self {
model_id,
constraint: VersionConstraint::new(constraint)?,
optional: false,
features: vec![],
})
}
/// Mark this dependency as optional.
pub fn optional(mut self) -> Self {
self.optional = true;
self
}
/// Add required features.
pub fn with_features(mut self, features: Vec<String>) -> Self {
self.features = features;
self
}
}
/// Resolved dependency with specific version.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResolvedDependency {
/// Model ID
pub model_id: ModelId,
/// Resolved version
pub version: ModelVersion,
/// Dependency specification
pub spec: DependencySpec,
}
/// Version resolution result.
#[derive(Debug, Clone)]
pub struct ResolutionResult {
/// Resolved dependencies
pub dependencies: Vec<ResolvedDependency>,
/// Resolution graph
pub graph: DependencyGraph,
/// Any warnings during resolution
pub warnings: Vec<String>,
}
/// Dependency graph for conflict detection.
#[derive(Debug, Clone)]
pub struct DependencyGraph {
/// Nodes in the graph (model_id -> available versions)
pub nodes: HashMap<ModelId, Vec<ModelVersion>>,
/// Edges in the graph (dependent -> dependencies)
pub edges: HashMap<ModelId, Vec<DependencySpec>>,
/// Resolved versions
pub resolved: HashMap<ModelId, ModelVersion>,
}
impl DependencyGraph {
/// Create a new empty dependency graph.
pub fn new() -> Self {
Self {
nodes: HashMap::new(),
edges: HashMap::new(),
resolved: HashMap::new(),
}
}
/// Add a model with its available versions.
pub fn add_model(&mut self, model_id: ModelId, versions: Vec<ModelVersion>) {
self.nodes.insert(model_id, versions);
}
/// Add dependencies for a model.
pub fn add_dependencies(&mut self, model_id: ModelId, dependencies: Vec<DependencySpec>) {
self.edges.insert(model_id, dependencies);
}
/// Check for circular dependencies.
pub fn has_cycles(&self) -> bool {
let mut visited = HashSet::new();
let mut rec_stack = HashSet::new();
for model_id in self.nodes.keys() {
if self.has_cycle_util(model_id, &mut visited, &mut rec_stack) {
return true;
}
}
false
}
fn has_cycle_util(
&self,
model_id: &ModelId,
visited: &mut HashSet<ModelId>,
rec_stack: &mut HashSet<ModelId>,
) -> bool {
if rec_stack.contains(model_id) {
return true;
}
if visited.contains(model_id) {
return false;
}
visited.insert(model_id.clone());
rec_stack.insert(model_id.clone());
if let Some(dependencies) = self.edges.get(model_id) {
for dep in dependencies {
if self.has_cycle_util(&dep.model_id, visited, rec_stack) {
return true;
}
}
}
rec_stack.remove(model_id);
false
}
/// Get topologically sorted order.
pub fn topological_sort(&self) -> Result<Vec<ModelId>, HubError> {
if self.has_cycles() {
return Err(HubError::VersionConflict {
details: "Circular dependency detected".to_string(),
});
}
let mut visited = HashSet::new();
let mut stack = Vec::new();
for model_id in self.nodes.keys() {
if !visited.contains(model_id) {
self.topological_sort_util(model_id, &mut visited, &mut stack);
}
}
// Don't reverse - DFS post-order naturally gives dependencies before dependents
Ok(stack)
}
fn topological_sort_util(
&self,
model_id: &ModelId,
visited: &mut HashSet<ModelId>,
stack: &mut Vec<ModelId>,
) {
visited.insert(model_id.clone());
if let Some(dependencies) = self.edges.get(model_id) {
for dep in dependencies {
if !visited.contains(&dep.model_id) {
self.topological_sort_util(&dep.model_id, visited, stack);
}
}
}
stack.push(model_id.clone());
}
}
/// Version resolver for dependency management.
pub struct VersionResolver {
/// Available versions for each model
available_versions: HashMap<ModelId, Vec<ModelVersion>>,
}
impl VersionResolver {
/// Create a new version resolver.
pub fn new() -> Self {
Self {
available_versions: HashMap::new(),
}
}
/// Get available versions for a model.
pub fn get_versions(&self, model_id: &ModelId) -> Option<&Vec<ModelVersion>> {
self.available_versions.get(model_id)
}
/// Add available versions for a model.
pub fn add_versions(&mut self, model_id: ModelId, mut versions: Vec<ModelVersion>) {
versions.sort_by(|a, b| b.semver.cmp(&a.semver)); // Latest first
self.available_versions.insert(model_id, versions);
}
/// Resolve dependencies for a set of requirements.
pub fn resolve(&self, requirements: Vec<DependencySpec>) -> HubResult<ResolutionResult> {
let mut graph = DependencyGraph::new();
let mut to_process = requirements.clone();
let mut processed = HashSet::new();
// Build dependency graph
while let Some(spec) = to_process.pop() {
if processed.contains(&spec.model_id) {
continue;
}
processed.insert(spec.model_id.clone());
// Get available versions for this model
let versions = self
.available_versions
.get(&spec.model_id)
.ok_or_else(|| HubError::ModelNotFound {
model_id: spec.model_id.to_string(),
})?
.clone();
graph.add_model(spec.model_id.clone(), versions);
// Add dependencies of this model to processing queue
// For now, we'll assume no transitive dependencies
// In a real implementation, you'd query the registry for each model's dependencies
}
// Resolve versions using constraint satisfaction
let resolved = self.resolve_constraints(&graph, &requirements)?;
Ok(ResolutionResult {
dependencies: resolved,
graph,
warnings: vec![],
})
}
fn resolve_constraints(
&self,
graph: &DependencyGraph,
requirements: &[DependencySpec],
) -> HubResult<Vec<ResolvedDependency>> {
let mut resolved = Vec::new();
for spec in requirements {
let versions =
graph
.nodes
.get(&spec.model_id)
.ok_or_else(|| HubError::ModelNotFound {
model_id: spec.model_id.to_string(),
})?;
// Find the latest version that satisfies the constraint
let matching_version = versions
.iter()
.find(|v| spec.constraint.matches(v))
.ok_or_else(|| HubError::VersionConflict {
details: format!(
"No version of {} satisfies constraint {}",
spec.model_id, spec.constraint
),
})?;
resolved.push(ResolvedDependency {
model_id: spec.model_id.clone(),
version: matching_version.clone(),
spec: spec.clone(),
});
}
Ok(resolved)
}
}
impl Default for VersionResolver {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use semver::Version;
#[test]
fn test_model_version_creation() {
let version = ModelVersion::new(Version::parse("1.2.3").unwrap());
assert_eq!(version.major(), 1);
assert_eq!(version.minor(), 2);
assert_eq!(version.patch(), 3);
assert!(!version.is_prerelease());
}
#[test]
fn test_model_version_with_prerelease() {
// Create a base version and add prerelease separately
let version = ModelVersion::new(Version::parse("1.2.3").unwrap())
.with_prerelease("alpha.1".to_string());
assert!(version.is_prerelease());
assert_eq!(version.to_string(), "1.2.3-alpha.1");
}
#[test]
fn test_version_constraint() {
let constraint = VersionConstraint::new(">=1.0.0").unwrap();
let version = ModelVersion::new(Version::parse("1.2.3").unwrap());
assert!(constraint.matches(&version));
let old_version = ModelVersion::new(Version::parse("0.9.0").unwrap());
assert!(!constraint.matches(&old_version));
}
#[test]
fn test_dependency_spec() {
let model_id = ModelId::new("rustytorch", "tokenizer");
let spec = DependencySpec::new(model_id.clone(), "^1.0.0").unwrap();
assert_eq!(spec.model_id, model_id);
assert!(!spec.optional);
assert_eq!(spec.constraint.as_str(), "^1.0.0");
}
#[test]
fn test_dependency_graph_cycle_detection() {
let mut graph = DependencyGraph::new();
let model_a = ModelId::new("test", "a");
let model_b = ModelId::new("test", "b");
let version = vec![ModelVersion::new(Version::parse("1.0.0").unwrap())];
graph.add_model(model_a.clone(), version.clone());
graph.add_model(model_b.clone(), version);
// Create circular dependency: A -> B -> A
let dep_b = DependencySpec::new(model_b.clone(), "1.0.0").unwrap();
let dep_a = DependencySpec::new(model_a.clone(), "1.0.0").unwrap();
graph.add_dependencies(model_a, vec![dep_b]);
graph.add_dependencies(model_b, vec![dep_a]);
assert!(graph.has_cycles());
}
#[test]
fn test_version_resolver() {
let mut resolver = VersionResolver::new();
let model_id = ModelId::new("rustytorch", "tokenizer");
let versions = vec![
ModelVersion::new(Version::parse("1.0.0").unwrap()),
ModelVersion::new(Version::parse("1.1.0").unwrap()),
ModelVersion::new(Version::parse("2.0.0").unwrap()),
];
resolver.add_versions(model_id.clone(), versions);
let requirements = vec![DependencySpec::new(model_id, "^1.0.0").unwrap()];
let result = resolver.resolve(requirements).unwrap();
assert_eq!(result.dependencies.len(), 1);
let resolved = &result.dependencies[0];
assert_eq!(resolved.version.semver, Version::parse("1.1.0").unwrap());
}
#[test]
fn test_version_increment() {
let version = ModelVersion::new(Version::parse("1.2.3").unwrap());
let next_major = version.next_major();
assert_eq!(next_major.semver, Version::parse("2.0.0").unwrap());
let next_minor = version.next_minor();
assert_eq!(next_minor.semver, Version::parse("1.3.0").unwrap());
let next_patch = version.next_patch();
assert_eq!(next_patch.semver, Version::parse("1.2.4").unwrap());
}
#[test]
fn test_version_parsing() {
let version = ModelVersion::parse("1.2.3-alpha+build").unwrap();
assert_eq!(version.major(), 1);
assert_eq!(version.minor(), 2);
assert_eq!(version.patch(), 3);
assert!(ModelVersion::parse("invalid").is_err());
}
#[test]
fn test_topological_sort() {
let mut graph = DependencyGraph::new();
let model_a = ModelId::new("test", "a");
let model_b = ModelId::new("test", "b");
let model_c = ModelId::new("test", "c");
let version = vec![ModelVersion::new(Version::parse("1.0.0").unwrap())];
graph.add_model(model_a.clone(), version.clone());
graph.add_model(model_b.clone(), version.clone());
graph.add_model(model_c.clone(), version);
// A -> B -> C
let dep_b = DependencySpec::new(model_b.clone(), "1.0.0").unwrap();
let dep_c = DependencySpec::new(model_c.clone(), "1.0.0").unwrap();
graph.add_dependencies(model_a.clone(), vec![dep_b]);
graph.add_dependencies(model_b.clone(), vec![dep_c]);
let sorted = graph.topological_sort().unwrap();
// C should come before B, B should come before A
let c_pos = sorted.iter().position(|x| x == &model_c).unwrap();
let b_pos = sorted.iter().position(|x| x == &model_b).unwrap();
let a_pos = sorted.iter().position(|x| x == &model_a).unwrap();
assert!(c_pos < b_pos);
assert!(b_pos < a_pos);
}
}