1#[cfg(not(any(feature = "tls-rustls", feature = "tls-native")))]
5compile_error!(
6 "no TLS backend selected: enable `tls-rustls` (default) or `tls-native`/`openssl`. \
7 Without one, reqwest builds with no TLS support and every HTTPS request fails at runtime."
8);
9
10use serde::{Deserialize, Serialize};
11use std::error::Error as StdError;
12
13pub mod artifact_manifest;
14pub mod cache;
15pub mod client_config;
16pub mod config;
17pub mod download;
18pub mod envs;
19pub mod models;
20pub mod providers;
21#[cfg(any(test, feature = "test-support"))]
22#[doc(hidden)]
23pub mod test_support;
24
25#[allow(clippy::similar_names)]
27#[allow(clippy::default_trait_access)]
28#[allow(clippy::doc_markdown)]
29#[allow(clippy::must_use_candidate)]
30#[allow(clippy::result_large_err)]
31pub mod grpc {
32 pub mod health {
33 tonic::include_proto!("model_express.health");
34 }
35 pub mod api {
36 tonic::include_proto!("model_express.api");
37 }
38 pub mod model {
39 tonic::include_proto!("model_express.model");
40 }
41 pub mod p2p {
42 tonic::include_proto!("model_express.p2p");
43 }
44 pub mod refit {
45 tonic::include_proto!("model_express.refit");
46 }
47}
48
49#[derive(Debug, Clone, Serialize, Deserialize)]
51pub struct Response<T> {
52 pub success: bool,
53 pub data: Option<T>,
54 pub error: Option<String>,
55}
56
57#[derive(Debug, thiserror::Error)]
59pub enum Error {
60 #[error("Server returned error: {0}")]
61 Server(String),
62
63 #[error("I/O error: {0}")]
64 Io(String),
65
66 #[error("Validation error: {0}")]
67 Validation(String),
68
69 #[error("Serialization error: {0}")]
70 Serialization(String),
71
72 #[error("gRPC error: {0}")]
73 Grpc(#[from] tonic::Status),
74
75 #[error("Transport error: {0}")]
76 Transport(String),
77
78 #[error("Generic error: {0}")]
79 Generic(String),
80}
81
82fn format_error_chain(err: &(dyn StdError + 'static)) -> String {
83 let mut parts = Vec::new();
84 let mut current = Some(err);
85
86 while let Some(error) = current {
87 let part = error.to_string();
88 if !part.is_empty() && parts.last() != Some(&part) {
89 parts.push(part);
90 }
91 current = error.source();
92 }
93
94 if parts.len() > 1 && parts.first().is_some_and(|part| part == "transport error") {
95 parts.remove(0);
96 }
97
98 if parts.is_empty() {
99 "transport error".to_string()
100 } else {
101 parts.join(": ")
102 }
103}
104
105impl From<tonic::Status> for Box<Error> {
107 fn from(err: tonic::Status) -> Self {
108 Box::new(Error::Grpc(err))
109 }
110}
111
112impl From<tonic::transport::Error> for Error {
113 fn from(err: tonic::transport::Error) -> Self {
114 Error::Transport(format_error_chain(&err))
115 }
116}
117
118impl From<tonic::transport::Error> for Box<Error> {
119 fn from(err: tonic::transport::Error) -> Self {
120 Box::new(Error::from(err))
121 }
122}
123
124pub type Result<T> = std::result::Result<T, Box<Error>>;
126
127pub struct Utils;
129
130impl Utils {
131 pub fn get_home_dir() -> std::result::Result<String, Box<Error>> {
134 envs::home_dir()
135 }
136}
137
138pub mod constants {
140 use std::num::NonZeroU16;
141
142 pub const DEFAULT_CACHE_PATH: &str = ".model-express/cache";
143 pub const DEFAULT_HF_CACHE_PATH: &str = ".cache/huggingface/hub";
144 pub const DEFAULT_CONFIG_PATH: &str = ".model-express/config.yaml";
145
146 pub const DEFAULT_GRPC_PORT: NonZeroU16 = NonZeroU16::new(8001).expect("8001 is non-zero");
147 pub const DEFAULT_TIMEOUT_SECS: u64 = 30;
148
149 pub const DEFAULT_METRICS_PORT: NonZeroU16 = NonZeroU16::new(9401).expect("9401 is non-zero");
156
157 pub const DEFAULT_SHARED_STORAGE: bool = true;
159
160 pub const DEFAULT_TRANSFER_CHUNK_SIZE: usize = 32 * 1024;
162}
163
164impl From<&models::Status> for grpc::health::HealthResponse {
166 fn from(status: &models::Status) -> Self {
167 Self {
168 version: status.version.clone(),
169 status: status.status.clone(),
170 uptime: status.uptime,
171 }
172 }
173}
174
175impl From<grpc::health::HealthResponse> for models::Status {
176 fn from(response: grpc::health::HealthResponse) -> Self {
177 Self {
178 version: response.version,
179 status: response.status,
180 uptime: response.uptime,
181 }
182 }
183}
184
185impl From<models::ModelProvider> for grpc::model::ModelProvider {
186 fn from(provider: models::ModelProvider) -> Self {
187 match provider {
188 models::ModelProvider::HuggingFace => grpc::model::ModelProvider::HuggingFace,
189 models::ModelProvider::Ngc => grpc::model::ModelProvider::Ngc,
190 models::ModelProvider::Gcs => grpc::model::ModelProvider::Gcs,
191 models::ModelProvider::S3 => grpc::model::ModelProvider::S3,
192 }
193 }
194}
195
196impl From<grpc::model::ModelProvider> for models::ModelProvider {
197 fn from(provider: grpc::model::ModelProvider) -> Self {
198 match provider {
199 grpc::model::ModelProvider::HuggingFace => models::ModelProvider::HuggingFace,
200 grpc::model::ModelProvider::Ngc => models::ModelProvider::Ngc,
201 grpc::model::ModelProvider::Gcs => models::ModelProvider::Gcs,
202 grpc::model::ModelProvider::S3 => models::ModelProvider::S3,
203 }
204 }
205}
206
207impl From<models::ModelStatus> for grpc::model::ModelStatus {
208 fn from(status: models::ModelStatus) -> Self {
209 match status {
210 models::ModelStatus::DOWNLOADING => grpc::model::ModelStatus::Downloading,
211 models::ModelStatus::DOWNLOADED => grpc::model::ModelStatus::Downloaded,
212 models::ModelStatus::ERROR => grpc::model::ModelStatus::Error,
213 }
214 }
215}
216
217impl From<grpc::model::ModelStatus> for models::ModelStatus {
218 fn from(status: grpc::model::ModelStatus) -> Self {
219 match status {
220 grpc::model::ModelStatus::Downloading => models::ModelStatus::DOWNLOADING,
221 grpc::model::ModelStatus::Downloaded => models::ModelStatus::DOWNLOADED,
222 grpc::model::ModelStatus::Error => models::ModelStatus::ERROR,
223 }
224 }
225}
226
227impl From<&models::ModelStatusResponse> for grpc::model::ModelStatusUpdate {
228 fn from(response: &models::ModelStatusResponse) -> Self {
229 Self {
230 model_name: response.model_name.clone(),
231 status: grpc::model::ModelStatus::from(response.status) as i32,
232 message: None,
233 provider: grpc::model::ModelProvider::from(response.provider) as i32,
234 resolved_revision: None,
235 }
236 }
237}
238
239impl From<grpc::model::ModelStatusUpdate> for models::ModelStatusResponse {
240 fn from(update: grpc::model::ModelStatusUpdate) -> Self {
241 Self {
242 model_name: update.model_name,
243 status: grpc::model::ModelStatus::try_from(update.status)
244 .unwrap_or(grpc::model::ModelStatus::Error)
245 .into(),
246 provider: grpc::model::ModelProvider::try_from(update.provider)
247 .unwrap_or(grpc::model::ModelProvider::HuggingFace)
248 .into(),
249 }
250 }
251}
252
253#[cfg(test)]
254mod tests {
255 use super::*;
256 use std::env;
257 use std::io;
258
259 #[test]
260 fn test_status_conversion_from_models_to_grpc() {
261 let status = models::Status {
262 version: "1.0.0".to_string(),
263 status: "ok".to_string(),
264 uptime: 3600,
265 };
266
267 let grpc_response: grpc::health::HealthResponse = (&status).into();
268
269 assert_eq!(grpc_response.version, status.version);
270 assert_eq!(grpc_response.status, status.status);
271 assert_eq!(grpc_response.uptime, status.uptime);
272 }
273
274 #[derive(Debug, thiserror::Error)]
275 #[error("outer error")]
276 struct OuterError(#[source] io::Error);
277
278 #[derive(Debug, thiserror::Error)]
279 #[error("transport error")]
280 struct TransportWrapper(#[source] io::Error);
281
282 #[test]
283 fn test_format_error_chain_includes_nested_causes() {
284 let err = OuterError(io::Error::other("connection reset by peer"));
285 assert_eq!(
286 format_error_chain(&err),
287 "outer error: connection reset by peer"
288 );
289 }
290
291 #[test]
292 fn test_format_error_chain_skips_repeated_transport_prefix() {
293 let err = TransportWrapper(io::Error::other("underlying cause"));
294 assert_eq!(format_error_chain(&err), "underlying cause");
295 }
296
297 #[test]
298 fn test_status_conversion_from_grpc_to_models() {
299 let grpc_response = grpc::health::HealthResponse {
300 version: "1.0.0".to_string(),
301 status: "ok".to_string(),
302 uptime: 3600,
303 };
304
305 let status: models::Status = grpc_response.into();
306
307 assert_eq!(status.version, "1.0.0");
308 assert_eq!(status.status, "ok");
309 assert_eq!(status.uptime, 3600);
310 }
311
312 #[test]
313 fn test_model_provider_conversion_both_ways() {
314 for model_provider in [
315 models::ModelProvider::HuggingFace,
316 models::ModelProvider::Ngc,
317 models::ModelProvider::Gcs,
318 models::ModelProvider::S3,
319 ] {
320 let grpc_provider: grpc::model::ModelProvider = model_provider.into();
321 let back_to_model: models::ModelProvider = grpc_provider.into();
322 assert_eq!(model_provider, back_to_model);
323 }
324 }
325
326 #[test]
327 fn test_model_status_conversion_both_ways() {
328 let statuses = vec![
329 models::ModelStatus::DOWNLOADING,
330 models::ModelStatus::DOWNLOADED,
331 models::ModelStatus::ERROR,
332 ];
333
334 for status in statuses {
335 let grpc_status: grpc::model::ModelStatus = status.into();
336 let back_to_model: models::ModelStatus = grpc_status.into();
337 assert_eq!(status, back_to_model);
338 }
339 }
340
341 #[test]
342 fn test_model_status_response_conversion_from_models_to_grpc() {
343 let response = models::ModelStatusResponse {
344 model_name: "test-model".to_string(),
345 status: models::ModelStatus::DOWNLOADED,
346 provider: models::ModelProvider::HuggingFace,
347 };
348
349 let grpc_update: grpc::model::ModelStatusUpdate = (&response).into();
350
351 assert_eq!(grpc_update.model_name, response.model_name);
352 assert_eq!(
353 grpc_update.status,
354 grpc::model::ModelStatus::Downloaded as i32
355 );
356 assert_eq!(
357 grpc_update.provider,
358 grpc::model::ModelProvider::HuggingFace as i32
359 );
360 assert!(grpc_update.message.is_none());
361 }
362
363 #[test]
364 fn test_model_status_response_conversion_from_grpc_to_models() {
365 let grpc_update = grpc::model::ModelStatusUpdate {
366 model_name: "test-model".to_string(),
367 status: grpc::model::ModelStatus::Downloaded as i32,
368 message: Some("Test message".to_string()),
369 provider: grpc::model::ModelProvider::HuggingFace as i32,
370 resolved_revision: None,
371 };
372
373 let response: models::ModelStatusResponse = grpc_update.into();
374
375 assert_eq!(response.model_name, "test-model");
376 assert_eq!(response.status, models::ModelStatus::DOWNLOADED);
377 assert_eq!(response.provider, models::ModelProvider::HuggingFace);
378 }
379
380 #[test]
381 fn test_error_types() {
382 let server_error = Error::Server("Internal error".to_string());
383 assert!(server_error.to_string().contains("Server returned error"));
384
385 let io_error = Error::Io("Permission denied".to_string());
386 assert!(io_error.to_string().contains("I/O error"));
387
388 let validation_error = Error::Validation("Unsafe path".to_string());
389 assert!(validation_error.to_string().contains("Validation error"));
390
391 let serialization_error = Error::Serialization("JSON parse error".to_string());
392 assert!(
393 serialization_error
394 .to_string()
395 .contains("Serialization error")
396 );
397 }
398
399 #[test]
400 fn test_constants() {
401 assert_eq!(constants::DEFAULT_GRPC_PORT.get(), 8001);
402 assert_eq!(constants::DEFAULT_METRICS_PORT.get(), 9401);
403 assert_ne!(
406 constants::DEFAULT_METRICS_PORT,
407 constants::DEFAULT_GRPC_PORT
408 );
409 assert_eq!(constants::DEFAULT_TIMEOUT_SECS, 30);
410 assert_eq!(constants::DEFAULT_TRANSFER_CHUNK_SIZE, 32 * 1024);
411 }
412
413 #[test]
414 fn test_response_creation() {
415 let success_response = Response {
416 success: true,
417 data: Some("test data".to_string()),
418 error: None,
419 };
420
421 assert!(success_response.success);
422 assert!(success_response.data.is_some());
423 assert!(success_response.error.is_none());
424
425 let error_response: Response<String> = Response {
426 success: false,
427 data: None,
428 error: Some("test error".to_string()),
429 };
430
431 assert!(!error_response.success);
432 assert!(error_response.data.is_none());
433 assert!(error_response.error.is_some());
434 }
435
436 #[test]
437 fn test_utils_get_home_dir() {
438 let home_dir = Utils::get_home_dir();
439
440 if let Ok(home_dir) = home_dir {
441 assert!(!home_dir.is_empty());
442 if let Ok(expected_home) = env::var("HOME") {
444 assert_eq!(home_dir, expected_home);
445 } else if let Ok(expected_home) = env::var("USERPROFILE") {
446 assert_eq!(home_dir, expected_home);
447 }
448 }
449 }
450}