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}