Skip to main content

osdk_core/model/
pull.rs

1use std::path::{Path, PathBuf};
2
3use futures_util::stream::{self, StreamExt, TryStreamExt};
4use globset::{Glob, GlobSet, GlobSetBuilder};
5use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
6
7use crate::backend::Ctx;
8use crate::error::{Error, Result};
9use crate::model::provider::{ModelProvider, RemoteModelFile};
10use crate::model::{DownloadedModelFile, InstalledModel, ModelRef, ModelStore, SnapshotIdentity};
11use crate::pipeline::download;
12use crate::pipeline::verify::{hash_file, HashAlgo};
13
14#[derive(Debug, Clone, Default)]
15pub struct PullOptions {
16    pub include: Vec<String>,
17    pub exclude: Vec<String>,
18    pub variant: Option<String>,
19}
20
21pub async fn pull(
22    ctx: &Ctx,
23    provider: &dyn ModelProvider,
24    name: &str,
25    reference: &ModelRef,
26    endpoint: &str,
27    options: &PullOptions,
28) -> Result<InstalledModel> {
29    let remote = provider.resolve(ctx, reference, endpoint).await?;
30    let selector = FileSelector::new(&options.include, &options.exclude)?;
31    let files: Vec<_> = remote
32        .files
33        .into_iter()
34        .filter(|file| selector.matches(&file.path))
35        .collect();
36    if files.is_empty() {
37        return Err(Error::other(format!(
38            "model file selection for {reference} matched no files"
39        )));
40    }
41
42    let jobs = ctx.config.settings.jobs.max(1);
43    let revision = remote.revision.clone();
44    let downloaded = stream::iter(
45        files
46            .into_iter()
47            .map(|file| download_file(ctx, reference, &revision, file)),
48    )
49    .buffer_unordered(jobs)
50    .try_collect::<Vec<_>>()
51    .await?;
52    let store = ModelStore::new(
53        ctx.dirs.clone(),
54        ctx.cas.clone(),
55        ctx.config.settings.link_mode,
56    );
57    store.publish(
58        SnapshotIdentity {
59            name: name.to_string(),
60            provider: reference.provider,
61            repository: reference.repository.clone(),
62            requested_revision: reference.revision.clone(),
63            revision: remote.revision,
64            endpoint: remote.endpoint,
65            variant: options.variant.clone(),
66        },
67        downloaded,
68    )
69}
70
71async fn download_file(
72    ctx: &Ctx,
73    reference: &ModelRef,
74    revision: &str,
75    file: RemoteModelFile,
76) -> Result<DownloadedModelFile> {
77    let relative = crate::model::safe_relative_path(&file.path)?;
78    let destination = download_path(ctx, reference, revision, &relative);
79    if ctx.config.settings.offline && !destination.is_file() {
80        return Err(Error::other(format!(
81            "offline model file cache miss for {} ({})",
82            reference, file.path
83        )));
84    }
85    if !ctx.config.settings.offline {
86        let headers = header_map(&file.headers)?;
87        download::download_with_headers(
88            &ctx.client,
89            &file.url,
90            &destination,
91            &format!("{}:{}", reference.repository, file.path),
92            ctx.show_progress,
93            &headers,
94        )
95        .await?;
96    }
97    let size = std::fs::metadata(&destination)
98        .map_err(|error| Error::io(&destination, error))?
99        .len();
100    if let Some(expected) = file.size {
101        if expected != size {
102            if !ctx.config.settings.offline {
103                let _ = std::fs::remove_file(&destination);
104            }
105            return Err(Error::other(format!(
106                "model file size mismatch for {}: expected {}, got {}",
107                file.path, expected, size
108            )));
109        }
110    }
111    let sha256 = match file.sha256 {
112        Some(expected) => {
113            if let Err(error) = crate::pipeline::verify::verify_file(
114                &destination,
115                &expected,
116                HashAlgo::Sha256,
117                &file.path,
118            ) {
119                if !ctx.config.settings.offline {
120                    let _ = std::fs::remove_file(&destination);
121                }
122                return Err(error);
123            }
124            Some(expected)
125        }
126        None => Some(hash_file(&destination, HashAlgo::Sha256)?),
127    };
128    Ok(DownloadedModelFile {
129        path: file.path,
130        source: destination,
131        size,
132        sha256,
133        etag: file.etag,
134    })
135}
136
137fn download_path(ctx: &Ctx, reference: &ModelRef, revision: &str, path: &Path) -> PathBuf {
138    ctx.dirs
139        .downloads()
140        .join("models")
141        .join(reference.provider.as_str())
142        .join(crate::dirs::sanitize_tool_id(&reference.repository))
143        .join(crate::dirs::sanitize_tool_id(revision))
144        .join(path)
145}
146
147fn header_map(headers: &[(String, String)]) -> Result<HeaderMap> {
148    let mut map = HeaderMap::new();
149    for (key, value) in headers {
150        let key = HeaderName::from_bytes(key.as_bytes())
151            .map_err(|error| Error::config(format!("invalid HTTP header `{key}`: {error}")))?;
152        let value = HeaderValue::from_str(value)
153            .map_err(|error| Error::config(format!("invalid HTTP header value: {error}")))?;
154        map.insert(key, value);
155    }
156    Ok(map)
157}
158
159struct FileSelector {
160    include: Option<GlobSet>,
161    exclude: GlobSet,
162}
163
164impl FileSelector {
165    fn new(include: &[String], exclude: &[String]) -> Result<Self> {
166        Ok(Self {
167            include: if include.is_empty() {
168                None
169            } else {
170                Some(build_globs(include)?)
171            },
172            exclude: build_globs(exclude)?,
173        })
174    }
175
176    fn matches(&self, path: &str) -> bool {
177        self.include
178            .as_ref()
179            .map(|include| include.is_match(path))
180            .unwrap_or(true)
181            && !self.exclude.is_match(path)
182    }
183}
184
185fn build_globs(patterns: &[String]) -> Result<GlobSet> {
186    let mut builder = GlobSetBuilder::new();
187    for pattern in patterns {
188        builder.add(
189            Glob::new(pattern)
190                .map_err(|error| Error::config(format!("invalid model file glob: {error}")))?,
191        );
192    }
193    builder
194        .build()
195        .map_err(|error| Error::config(format!("invalid model file globs: {error}")))
196}
197
198#[cfg(test)]
199mod tests {
200    use super::*;
201    use async_trait::async_trait;
202    use std::io::{Read, Write};
203    use std::net::TcpListener;
204    use std::sync::Arc;
205
206    use crate::config::{Config, Settings};
207    use crate::dirs::Dirs;
208    use crate::model::provider::huggingface::HuggingFace;
209    use crate::model::provider::modelscope::ModelScope;
210    use crate::model::provider::RemoteSnapshot;
211    use crate::platform::Platform;
212    use crate::store::link::LinkMode;
213    use crate::store::Cas;
214
215    #[test]
216    fn selector_supports_includes_and_excludes() {
217        let selector = FileSelector::new(
218            &["*.json".into(), "*.safetensors".into()],
219            &["tokenizer.json".into()],
220        )
221        .unwrap();
222        assert!(selector.matches("config.json"));
223        assert!(selector.matches("model.safetensors"));
224        assert!(!selector.matches("tokenizer.json"));
225        assert!(!selector.matches("README.md"));
226    }
227
228    #[tokio::test]
229    async fn huggingface_pull_locks_revision_downloads_and_rebuilds_offline() {
230        let config_bytes = br#"{"model":"fixture"}"#.to_vec();
231        let digest = crate::pipeline::verify::hash_bytes(&config_bytes, HashAlgo::Sha256);
232        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
233        let address = listener.local_addr().unwrap();
234        let server_digest = digest.clone();
235        let server_bytes = config_bytes.clone();
236        let server = std::thread::spawn(move || {
237            for request_number in 0..2 {
238                let (mut stream, _) = listener.accept().unwrap();
239                let mut request = Vec::new();
240                let mut buffer = [0u8; 2048];
241                while !request.ends_with(b"\r\n\r\n") {
242                    let read = stream.read(&mut buffer).unwrap();
243                    if read == 0 {
244                        break;
245                    }
246                    request.extend_from_slice(&buffer[..read]);
247                }
248                let request = String::from_utf8(request).unwrap();
249                assert!(request
250                    .to_ascii_lowercase()
251                    .contains("authorization: bearer fixture-token"));
252                if request_number == 0 {
253                    let body = format!(
254                        r#"{{"sha":"abc123","siblings":[{{"rfilename":"config.json","lfs":{{"sha256":"{}","size":{}}}}}]}}"#,
255                        server_digest,
256                        server_bytes.len()
257                    );
258                    write!(
259                        stream,
260                        "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
261                        body.len(),
262                        body
263                    )
264                    .unwrap();
265                } else {
266                    write!(
267                        stream,
268                        "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nETag: \"fixture\"\r\nConnection: close\r\n\r\n",
269                        server_bytes.len()
270                    )
271                    .unwrap();
272                    stream.write_all(&server_bytes).unwrap();
273                }
274            }
275        });
276
277        let temporary = tempfile::tempdir().unwrap();
278        let mut ctx = test_ctx(temporary.path(), false);
279        let reference = ModelRef::parse("hf:owner/repo@main").unwrap();
280        let endpoint = format!("http://{address}");
281        let installed = pull(
282            &ctx,
283            &HuggingFace::with_token("fixture-token"),
284            "fixture",
285            &reference,
286            &endpoint,
287            &PullOptions::default(),
288        )
289        .await
290        .unwrap();
291        server.join().unwrap();
292        assert_eq!(installed.manifest.revision, "abc123");
293        assert_eq!(
294            std::fs::read(installed.path.join("config.json")).unwrap(),
295            config_bytes
296        );
297
298        ModelStore::new(
299            ctx.dirs.clone(),
300            ctx.cas.clone(),
301            ctx.config.settings.link_mode,
302        )
303        .remove("fixture")
304        .unwrap();
305        ctx.config.settings.offline = true;
306        let rebuilt = pull(
307            &ctx,
308            &HuggingFace::default(),
309            "fixture",
310            &reference,
311            &endpoint,
312            &PullOptions::default(),
313        )
314        .await
315        .unwrap();
316        assert_eq!(rebuilt.manifest.revision, "abc123");
317    }
318
319    #[tokio::test]
320    async fn modelscope_pull_verifies_manifest_and_rebuilds_offline() {
321        let config_bytes = br#"{"model":"modelscope-fixture"}"#.to_vec();
322        let digest = crate::pipeline::verify::hash_bytes(&config_bytes, HashAlgo::Sha256);
323        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
324        let address = listener.local_addr().unwrap();
325        let server_digest = digest.clone();
326        let server_bytes = config_bytes.clone();
327        let server = std::thread::spawn(move || {
328            for request_number in 0..2 {
329                let (mut stream, _) = listener.accept().unwrap();
330                let mut request = Vec::new();
331                let mut buffer = [0u8; 2048];
332                while !request.ends_with(b"\r\n\r\n") {
333                    let read = stream.read(&mut buffer).unwrap();
334                    if read == 0 {
335                        break;
336                    }
337                    request.extend_from_slice(&buffer[..read]);
338                }
339                let request = String::from_utf8(request).unwrap();
340                let lower = request.to_ascii_lowercase();
341                assert!(lower.contains("authorization: bearer fixture-token"));
342                assert!(lower.contains("cookie: m_session_id=fixture-token"));
343                if request_number == 0 {
344                    assert!(request.contains("/repo/files?Revision=master&Recursive=true"));
345                    let body = format!(
346                        r#"{{"Code":200,"Success":true,"Message":"success","Data":{{"Files":[{{"Path":"config.json","Size":{},"Sha256":"{}","Type":"blob"}}]}}}}"#,
347                        server_bytes.len(),
348                        server_digest
349                    );
350                    write!(
351                        stream,
352                        "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
353                        body.len(),
354                        body
355                    )
356                    .unwrap();
357                } else {
358                    assert!(request.contains("/repo?Revision=master&FilePath=config.json"));
359                    write!(
360                        stream,
361                        "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nETag: \"fixture\"\r\nConnection: close\r\n\r\n",
362                        server_bytes.len()
363                    )
364                    .unwrap();
365                    stream.write_all(&server_bytes).unwrap();
366                }
367            }
368        });
369
370        let temporary = tempfile::tempdir().unwrap();
371        let mut ctx = test_ctx(temporary.path(), false);
372        let reference = ModelRef::parse("ms:owner/repo@master").unwrap();
373        let endpoint = format!("http://{address}");
374        let installed = pull(
375            &ctx,
376            &ModelScope::with_token("fixture-token"),
377            "modelscope-fixture",
378            &reference,
379            &endpoint,
380            &PullOptions::default(),
381        )
382        .await
383        .unwrap();
384        server.join().unwrap();
385        assert!(installed.manifest.revision.starts_with("master+manifest-"));
386        assert_eq!(
387            std::fs::read(installed.path.join("config.json")).unwrap(),
388            config_bytes
389        );
390
391        ModelStore::new(
392            ctx.dirs.clone(),
393            ctx.cas.clone(),
394            ctx.config.settings.link_mode,
395        )
396        .remove("modelscope-fixture")
397        .unwrap();
398        ctx.config.settings.offline = true;
399        let rebuilt = pull(
400            &ctx,
401            &ModelScope::default(),
402            "modelscope-fixture",
403            &reference,
404            &endpoint,
405            &PullOptions::default(),
406        )
407        .await
408        .unwrap();
409        assert_eq!(rebuilt.manifest.revision, installed.manifest.revision);
410    }
411
412    struct EmptyProvider;
413
414    #[async_trait]
415    impl ModelProvider for EmptyProvider {
416        async fn resolve(
417            &self,
418            _ctx: &Ctx,
419            _reference: &ModelRef,
420            _endpoint: &str,
421        ) -> Result<RemoteSnapshot> {
422            Ok(RemoteSnapshot {
423                revision: "abc".into(),
424                endpoint: "https://example.test".into(),
425                files: Vec::new(),
426            })
427        }
428    }
429
430    #[tokio::test]
431    async fn empty_selection_fails_before_publish() {
432        let temporary = tempfile::tempdir().unwrap();
433        let ctx = test_ctx(temporary.path(), false);
434        let error = pull(
435            &ctx,
436            &EmptyProvider,
437            "fixture",
438            &ModelRef::parse("hf:owner/repo").unwrap(),
439            "https://example.test",
440            &PullOptions::default(),
441        )
442        .await
443        .unwrap_err();
444        assert!(error.to_string().contains("matched no files"));
445    }
446
447    fn test_ctx(root: &Path, offline: bool) -> Ctx {
448        let dirs = Dirs::resolve_from(|key| match key {
449            "OSDK_DATA_DIR" => Some(root.join("data").display().to_string()),
450            "OSDK_CACHE_DIR" => Some(root.join("cache").display().to_string()),
451            "OSDK_CONFIG_DIR" => Some(root.join("config").display().to_string()),
452            _ => None,
453        })
454        .unwrap();
455        dirs.ensure().unwrap();
456        Ctx {
457            dirs: dirs.clone(),
458            platform: Platform::current(),
459            config: Config {
460                settings: Settings {
461                    offline,
462                    link_mode: LinkMode::Copy,
463                    ..Default::default()
464                },
465                sources: Default::default(),
466                tools: Default::default(),
467                tool_configs: Default::default(),
468                global_tools: Default::default(),
469                global_tool_configs: Default::default(),
470                tool_origins: Default::default(),
471                aliases: Default::default(),
472                project_config_path: None,
473            },
474            client: reqwest::Client::new(),
475            cas: Arc::new(Cas::new(dirs.store.clone())),
476            show_progress: false,
477        }
478    }
479}