Skip to main content

modelexpress_common/
lib.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4#[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// Generated gRPC code
26#[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/// Defines the shared response format between server and client (legacy HTTP)
50#[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/// Common error types that both client and server can use
58#[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
105// Implement From traits for Box<Error> to work with the Result<T> type
106impl 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
124/// Common result type for the project
125pub type Result<T> = std::result::Result<T, Box<Error>>;
126
127/// Marker struct to use Utils methods
128pub struct Utils;
129
130impl Utils {
131    /// Get home directory from environment variables ([`envs::HOME`], then
132    /// [`envs::USERPROFILE`]).
133    pub fn get_home_dir() -> std::result::Result<String, Box<Error>> {
134        envs::home_dir()
135    }
136}
137
138/// Constants shared between client and server
139pub 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    /// Default port for the server's Prometheus `/metrics` listener.
150    ///
151    /// Deliberately not [`DEFAULT_GRPC_PORT`]: tonic serves HTTP/2 only, so a
152    /// scrape aimed at the gRPC port can never succeed. Chosen clear of the
153    /// ports already in play around a ModelExpress deployment — 8001/8002 (gRPC
154    /// and the client worker service) and 9090 (Dynamo's health endpoint).
155    pub const DEFAULT_METRICS_PORT: NonZeroU16 = NonZeroU16::new(9401).expect("9401 is non-zero");
156
157    /// Default setting for shared storage mode (true = client and server share a network drive)
158    pub const DEFAULT_SHARED_STORAGE: bool = true;
159
160    /// Default chunk size for file transfer streaming in bytes (32 KB)
161    pub const DEFAULT_TRANSFER_CHUNK_SIZE: usize = 32 * 1024;
162}
163
164// Conversion utilities between gRPC and legacy models
165impl 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        // The scrape target must never be the gRPC listener: tonic is HTTP/2
404        // only and Prometheus scrapes with an HTTP/1.1 GET.
405        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            // Check against HOME or USERPROFILE
443            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}