594 lines
17 KiB
Rust
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);
|
|
}
|
|
}
|