Skip to main content

modelexpress_common/
models.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use clap::{ValueEnum, builder::PossibleValue};
5use serde::{Deserialize, Serialize};
6use std::fmt::{Display, Formatter};
7
8/// Status model for server health checks
9#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct Status {
11    pub version: String,
12    pub status: String,
13    pub uptime: u64,
14}
15
16/// Status of a model download
17#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq)]
18pub enum ModelStatus {
19    /// Model is currently being downloaded
20    DOWNLOADING,
21    /// Model has been successfully downloaded
22    DOWNLOADED,
23    /// Model download failed with an error
24    ERROR,
25}
26
27/// Supported model providers
28#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Default)]
29pub enum ModelProvider {
30    /// Hugging Face model hub
31    #[default]
32    HuggingFace,
33    /// NVIDIA NGC catalog
34    Ngc,
35    /// Google Cloud Storage
36    Gcs,
37    /// S3-compatible object storage
38    S3,
39}
40
41impl ModelProvider {
42    #[must_use]
43    pub const fn as_str(self) -> &'static str {
44        match self {
45            Self::HuggingFace => "hugging-face",
46            Self::Ngc => "ngc",
47            Self::Gcs => "gcs",
48            Self::S3 => "s3",
49        }
50    }
51
52    #[must_use]
53    pub fn resolve_provider_for_model_name(model_name: &str, default_provider: Self) -> Self {
54        let model_name = model_name.trim_start();
55        if model_name
56            .get(.."s3://".len())
57            .is_some_and(|prefix| prefix.eq_ignore_ascii_case("s3://"))
58        {
59            Self::S3
60        } else if model_name
61            .get(.."gs://".len())
62            .is_some_and(|prefix| prefix.eq_ignore_ascii_case("gs://"))
63        {
64            Self::Gcs
65        } else if model_name
66            .get(.."ngc://".len())
67            .is_some_and(|prefix| prefix.eq_ignore_ascii_case("ngc://"))
68        {
69            Self::Ngc
70        } else {
71            default_provider
72        }
73    }
74}
75
76impl Display for ModelProvider {
77    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
78        f.write_str(self.as_str())
79    }
80}
81
82impl ValueEnum for ModelProvider {
83    fn value_variants<'a>() -> &'a [Self] {
84        &[Self::HuggingFace, Self::Ngc, Self::Gcs, Self::S3]
85    }
86
87    fn to_possible_value(&self) -> Option<PossibleValue> {
88        Some(PossibleValue::new(self.as_str()))
89    }
90}
91
92/// Response for model status request
93#[derive(Debug, Clone, Serialize, Deserialize)]
94pub struct ModelStatusResponse {
95    pub model_name: String,
96    pub status: ModelStatus,
97    pub provider: ModelProvider,
98}
99
100#[cfg(test)]
101#[allow(clippy::expect_used)]
102mod tests {
103    use super::*;
104
105    #[test]
106    fn test_model_status_serialization() {
107        let status = ModelStatus::DOWNLOADING;
108        let serialized = serde_json::to_string(&status).expect("Failed to serialize ModelStatus");
109        let deserialized: ModelStatus =
110            serde_json::from_str(&serialized).expect("Failed to deserialize ModelStatus");
111        assert_eq!(status, deserialized);
112    }
113
114    #[test]
115    fn test_model_provider_serialization() {
116        for provider in [
117            ModelProvider::HuggingFace,
118            ModelProvider::Ngc,
119            ModelProvider::Gcs,
120            ModelProvider::S3,
121        ] {
122            let serialized =
123                serde_json::to_string(&provider).expect("Failed to serialize ModelProvider");
124            let deserialized: ModelProvider =
125                serde_json::from_str(&serialized).expect("Failed to deserialize ModelProvider");
126            assert_eq!(provider, deserialized);
127        }
128    }
129
130    #[test]
131    fn test_model_provider_default() {
132        let provider = ModelProvider::default();
133        assert_eq!(provider, ModelProvider::HuggingFace);
134    }
135
136    #[test]
137    fn test_model_provider_display() {
138        assert_eq!(ModelProvider::HuggingFace.to_string(), "hugging-face");
139        assert_eq!(ModelProvider::Ngc.to_string(), "ngc");
140        assert_eq!(ModelProvider::Gcs.to_string(), "gcs");
141        assert_eq!(ModelProvider::S3.to_string(), "s3");
142    }
143
144    #[test]
145    fn test_model_provider_resolve_provider_for_model_name() {
146        assert_eq!(
147            ModelProvider::resolve_provider_for_model_name(
148                "s3://bucket/model",
149                ModelProvider::HuggingFace,
150            ),
151            ModelProvider::S3
152        );
153        assert_eq!(
154            ModelProvider::resolve_provider_for_model_name(
155                " gs://bucket/model",
156                ModelProvider::HuggingFace,
157            ),
158            ModelProvider::Gcs
159        );
160        assert_eq!(
161            ModelProvider::resolve_provider_for_model_name(
162                "NGC://org/model",
163                ModelProvider::HuggingFace,
164            ),
165            ModelProvider::Ngc
166        );
167        assert_eq!(
168            ModelProvider::resolve_provider_for_model_name("org/model", ModelProvider::Ngc),
169            ModelProvider::Ngc
170        );
171    }
172
173    #[test]
174    fn test_model_provider_value_enum_matches_display() {
175        for provider in [
176            ModelProvider::HuggingFace,
177            ModelProvider::Ngc,
178            ModelProvider::Gcs,
179            ModelProvider::S3,
180        ] {
181            let parsed = ModelProvider::from_str(provider.as_str(), false)
182                .expect("Failed to parse ModelProvider from clap value");
183            assert_eq!(parsed, provider);
184        }
185    }
186
187    #[test]
188    fn test_status_serialization() {
189        let status = Status {
190            version: "1.0.0".to_string(),
191            status: "ok".to_string(),
192            uptime: 3600,
193        };
194
195        let serialized = serde_json::to_string(&status).expect("Failed to serialize Status");
196        let deserialized: Status =
197            serde_json::from_str(&serialized).expect("Failed to deserialize Status");
198
199        assert_eq!(status.version, deserialized.version);
200        assert_eq!(status.status, deserialized.status);
201        assert_eq!(status.uptime, deserialized.uptime);
202    }
203
204    #[test]
205    fn test_model_status_response_serialization() {
206        let response = ModelStatusResponse {
207            model_name: "test-model".to_string(),
208            status: ModelStatus::DOWNLOADED,
209            provider: ModelProvider::HuggingFace,
210        };
211
212        let serialized =
213            serde_json::to_string(&response).expect("Failed to serialize ModelStatusResponse");
214        let deserialized: ModelStatusResponse =
215            serde_json::from_str(&serialized).expect("Failed to deserialize ModelStatusResponse");
216
217        assert_eq!(response.model_name, deserialized.model_name);
218        assert_eq!(response.status, deserialized.status);
219        assert_eq!(response.provider, deserialized.provider);
220    }
221
222    #[test]
223    fn test_model_status_all_variants() {
224        assert_eq!(ModelStatus::DOWNLOADING, ModelStatus::DOWNLOADING);
225        assert_eq!(ModelStatus::DOWNLOADED, ModelStatus::DOWNLOADED);
226        assert_eq!(ModelStatus::ERROR, ModelStatus::ERROR);
227
228        assert_ne!(ModelStatus::DOWNLOADING, ModelStatus::DOWNLOADED);
229        assert_ne!(ModelStatus::DOWNLOADED, ModelStatus::ERROR);
230        assert_ne!(ModelStatus::ERROR, ModelStatus::DOWNLOADING);
231    }
232}