Skip to main content

modelexpress_common/providers/gcs/
provider_cache.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use super::{
5    model_dir::ModelDir,
6    model_name::{BucketName, CACHE_ROOT_DIR_NAME, ModelName},
7};
8use crate::cache::{ModelInfo, ProviderCache};
9use crate::models::ModelProvider;
10use anyhow::{Context, Result};
11use std::{
12    fs,
13    path::{Path, PathBuf},
14};
15use tracing::{info, warn};
16
17pub struct GcsProviderCache;
18
19impl GcsProviderCache {
20    fn collect_cached_models(
21        cache_dir: &Path,
22        current_dir: &Path,
23        models: &mut Vec<ModelInfo>,
24    ) -> Result<()> {
25        let current_model_dir = ModelDir::new(cache_dir, current_dir);
26        if current_model_dir.has_manifest_file() {
27            match current_model_dir.model_name() {
28                Ok(model) => {
29                    if current_model_dir.cache_satisfies_request(false)? {
30                        models.push(ModelInfo {
31                            provider: ModelProvider::Gcs,
32                            name: model.to_string(),
33                            size: current_model_dir.size()?,
34                            path: current_dir.to_path_buf(),
35                        });
36                    }
37                    return Ok(());
38                }
39                Err(err) => {
40                    warn!(
41                        "Skipping invalid GCS cache entry '{}': {}",
42                        current_dir.display(),
43                        err
44                    );
45                }
46            }
47        }
48
49        for entry in fs::read_dir(current_dir)
50            .with_context(|| format!("Failed to read directory '{}'", current_dir.display()))?
51        {
52            let entry = entry?;
53            let path = entry.path();
54            if path.is_dir() {
55                Self::collect_cached_models(cache_dir, &path, models)?;
56            }
57        }
58
59        Ok(())
60    }
61}
62
63impl ProviderCache for GcsProviderCache {
64    fn clear_model(&self, cache_dir: &Path, model_name: &str) -> Result<()> {
65        let model = ModelName::parse(model_name)?;
66        let model_dir = model.model_dir(cache_dir);
67        let model_dir_state = ModelDir::new(cache_dir, &model_dir);
68
69        if !model_dir_state.is_removable()? {
70            info!(
71                "Model not found in cache: {} ({:?})",
72                model_name,
73                ModelProvider::Gcs
74            );
75            return Ok(());
76        }
77
78        model_dir_state.remove()?;
79        info!("Cleared model: {} ({:?})", model_name, ModelProvider::Gcs);
80
81        Ok(())
82    }
83
84    fn resolve_model_path(
85        &self,
86        cache_dir: &Path,
87        model_name: &str,
88        _revision: Option<&str>,
89    ) -> Result<PathBuf> {
90        Ok(ModelName::parse(model_name)?.model_dir(cache_dir))
91    }
92
93    fn list_models(&self, cache_dir: &Path) -> Result<Vec<ModelInfo>> {
94        let mut models = Vec::new();
95        let root = cache_dir.join(CACHE_ROOT_DIR_NAME);
96
97        if !root.exists() {
98            return Ok(models);
99        }
100
101        for bucket_entry in fs::read_dir(&root)? {
102            let bucket_entry = bucket_entry?;
103            let bucket_path = bucket_entry.path();
104            if !bucket_path.is_dir() {
105                continue;
106            }
107
108            let Some(bucket_name) = bucket_path.file_name().and_then(|name| name.to_str()) else {
109                continue;
110            };
111
112            if BucketName::parse(bucket_name).is_err() {
113                warn!(
114                    "Skipping invalid GCS bucket cache entry '{}'",
115                    bucket_path.display()
116                );
117                continue;
118            }
119
120            Self::collect_cached_models(cache_dir, &bucket_path, &mut models)?;
121        }
122
123        Ok(models)
124    }
125}
126
127#[cfg(test)]
128#[allow(clippy::expect_used)]
129mod tests {
130    use super::super::test_support::{
131        expected_model_dir, manifest_entry, write_cached_model, write_incomplete_cached_model,
132        write_manifest_with_payloads,
133    };
134    use super::*;
135    use crate::cache::ProviderCache;
136    use std::fs;
137    use tempfile::TempDir;
138
139    #[test]
140    fn test_list_models_scenarios() {
141        let cache = GcsProviderCache;
142        for scenario in [
143            "partial_manifest",
144            "recursive_siblings",
145            "model_path_contains_internal_segment",
146        ] {
147            let temp_dir = TempDir::new().expect("Failed to create temp dir");
148            match scenario {
149                "partial_manifest" => write_manifest_with_payloads(
150                    temp_dir.path(),
151                    "gs://bucket/foo/bar/baz",
152                    &[("tokenizer.json", b"{}")],
153                    vec![
154                        manifest_entry("tokenizer.json", b"{}"),
155                        manifest_entry("weights/model.bin", b"weights"),
156                    ],
157                ),
158                "recursive_siblings" => {
159                    write_cached_model(
160                        temp_dir.path(),
161                        "gs://bucket/foo/bar/baz",
162                        "tokenizer.json",
163                        b"{}",
164                    );
165                    write_cached_model(
166                        temp_dir.path(),
167                        "gs://bucket/foo/bar/buz",
168                        "weights/model.bin",
169                        b"abcd",
170                    );
171                }
172                "model_path_contains_internal_segment" => {
173                    write_cached_model(
174                        temp_dir.path(),
175                        "gs://bucket/foo/.mx/bar",
176                        "tokenizer.json",
177                        b"{}",
178                    );
179                }
180                _ => unreachable!("unexpected scenario"),
181            }
182
183            let mut models = cache
184                .list_models(temp_dir.path())
185                .expect("Expected model listing");
186            models.sort_by(|left, right| left.name.cmp(&right.name));
187
188            match scenario {
189                "partial_manifest" => {
190                    assert!(models.is_empty(), "partial manifest should not be listed");
191                }
192                "recursive_siblings" => {
193                    assert_eq!(models.len(), 2);
194                    assert_eq!(models[0].name, "gs://bucket/foo/bar/baz");
195                    assert_eq!(models[0].size, 2);
196                    assert_eq!(models[1].name, "gs://bucket/foo/bar/buz");
197                    assert_eq!(models[1].size, 4);
198                }
199                "model_path_contains_internal_segment" => {
200                    assert_eq!(models.len(), 1);
201                    assert_eq!(models[0].name, "gs://bucket/foo/.mx/bar");
202                    assert_eq!(models[0].size, 2);
203                }
204                _ => unreachable!("unexpected scenario"),
205            }
206        }
207    }
208
209    #[test]
210    fn test_clear_model_scenarios() {
211        let cache = GcsProviderCache;
212        for scenario in [
213            "ancestor_keeps_descendant",
214            "incomplete_removed",
215            "descendant_keeps_cached_ancestor",
216        ] {
217            let temp_dir = TempDir::new().expect("Failed to create temp dir");
218            match scenario {
219                "ancestor_keeps_descendant" => {
220                    let ancestor_name = "gs://bucket/foo/bar";
221                    let descendant_name = "gs://bucket/foo/bar/baz";
222                    write_cached_model(temp_dir.path(), descendant_name, "tokenizer.json", b"{}");
223                    cache
224                        .clear_model(temp_dir.path(), ancestor_name)
225                        .expect("Expected clear to succeed");
226                    assert!(expected_model_dir(temp_dir.path(), descendant_name).exists());
227                }
228                "incomplete_removed" => {
229                    let model_name = "gs://bucket/foo/bar";
230                    let model_dir = expected_model_dir(temp_dir.path(), model_name);
231                    write_incomplete_cached_model(
232                        temp_dir.path(),
233                        model_name,
234                        "tokenizer.json",
235                        b"{}",
236                    );
237                    cache
238                        .clear_model(temp_dir.path(), model_name)
239                        .expect("Expected clear to succeed");
240                    assert!(!model_dir.exists());
241                }
242                "descendant_keeps_cached_ancestor" => {
243                    let ancestor_name = "gs://bucket/foo/bar";
244                    let descendant_name = "gs://bucket/foo/bar/baz";
245                    let ancestor_dir = expected_model_dir(temp_dir.path(), ancestor_name);
246                    let descendant_dir = expected_model_dir(temp_dir.path(), descendant_name);
247                    write_cached_model(temp_dir.path(), ancestor_name, "tokenizer.json", b"{}");
248                    fs::create_dir_all(&descendant_dir).expect("Failed to create descendant dir");
249                    fs::write(descendant_dir.join("partial.bin"), b"partial")
250                        .expect("Failed to create descendant payload");
251                    cache
252                        .clear_model(temp_dir.path(), descendant_name)
253                        .expect("Expected clear to succeed");
254                    assert!(ancestor_dir.exists());
255                    assert!(ancestor_dir.join("tokenizer.json").exists());
256                    assert!(descendant_dir.exists());
257                }
258                _ => unreachable!("unexpected scenario"),
259            }
260        }
261    }
262}