1use clap::{ValueEnum, builder::PossibleValue};
5use serde::{Deserialize, Serialize};
6use std::fmt::{Display, Formatter};
7
8#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct Status {
11 pub version: String,
12 pub status: String,
13 pub uptime: u64,
14}
15
16#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq)]
18pub enum ModelStatus {
19 DOWNLOADING,
21 DOWNLOADED,
23 ERROR,
25}
26
27#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Default)]
29pub enum ModelProvider {
30 #[default]
32 HuggingFace,
33 Ngc,
35 Gcs,
37 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#[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}