Skip to main content

osdk_core/model/
source.rs

1use std::time::{Duration, Instant};
2
3use futures_util::StreamExt;
4use reqwest::header::{HeaderMap, HeaderName, HeaderValue, RANGE};
5
6use crate::backend::Ctx;
7use crate::error::{Error, Result};
8use crate::model::provider::huggingface::HuggingFace;
9use crate::model::provider::modelscope::ModelScope;
10use crate::model::provider::{ModelProvider, RemoteModelFile};
11use crate::model::{ModelRef, ProviderId};
12use crate::source::{ProbeCache, ProbeResult, Selection, Source};
13
14pub fn default_sources(provider: ProviderId) -> Vec<Source> {
15    match provider {
16        ProviderId::HuggingFace => vec![Source::official("official", "https://huggingface.co")],
17        ProviderId::ModelScope => {
18            let mut international = Source::official("modelscope-ai", "https://www.modelscope.ai");
19            international.priority = 10;
20            vec![
21                Source::official("modelscope-cn", "https://modelscope.cn"),
22                international,
23            ]
24        }
25    }
26}
27
28pub fn effective_sources(ctx: &Ctx, provider: ProviderId) -> Vec<Source> {
29    let mut sources = default_sources(provider);
30    if let Some(config) = ctx.config.tool_sources(provider.as_str()) {
31        sources.retain(|source| !config.disable.iter().any(|id| id == &source.id));
32        for custom in &config.custom {
33            sources.retain(|source| source.id != custom.id);
34            sources.push(custom.clone());
35        }
36    }
37    sources.retain(|source| source.enabled);
38    sources.sort_by_key(|source| source.priority);
39    sources
40}
41
42pub async fn ranked_sources(ctx: &Ctx, reference: &ModelRef, refresh: bool) -> Result<Vec<Source>> {
43    let sources = effective_sources(ctx, reference.provider);
44    if sources.is_empty() {
45        return Err(Error::NoUsableSource {
46            tool: reference.provider.to_string(),
47            tried: 0,
48        });
49    }
50    if let Some(pin) = ctx
51        .config
52        .tool_sources(reference.provider.as_str())
53        .and_then(|config| config.pin.as_deref())
54    {
55        if let Some(index) = sources.iter().position(|source| source.id == pin) {
56            let mut ordered = sources;
57            let pinned = ordered.remove(index);
58            ordered.insert(0, pinned);
59            return Ok(ordered);
60        }
61    }
62    if ctx.config.settings.offline
63        || matches!(
64            ctx.config.sources.selection,
65            Selection::Ordered | Selection::Pinned
66        )
67    {
68        return Ok(sources);
69    }
70
71    let results = if refresh {
72        let results = probe_all(ctx, reference, &sources).await;
73        save_cache(ctx, reference, &sources, &results);
74        results
75    } else if let Some(cache) = fresh_cache(ctx, reference, &sources) {
76        cache.results
77    } else {
78        let results = probe_all(ctx, reference, &sources).await;
79        save_cache(ctx, reference, &sources, &results);
80        results
81    };
82    let mut results = results;
83    results.sort_by(|left, right| right.score().total_cmp(&left.score()));
84    let mut ranked = Vec::with_capacity(sources.len());
85    for result in results.iter().filter(|result| result.ok) {
86        if let Some(source) = sources.iter().find(|source| source.id == result.source_id) {
87            ranked.push(source.clone());
88        }
89    }
90    for source in sources {
91        if !ranked.iter().any(|ranked| ranked.id == source.id) {
92            ranked.push(source);
93        }
94    }
95    Ok(ranked)
96}
97
98pub async fn refresh(ctx: &Ctx, reference: &ModelRef) -> Result<Vec<ProbeResult>> {
99    if ctx.config.settings.offline {
100        return Err(Error::other("cannot refresh model sources while offline"));
101    }
102    let sources = effective_sources(ctx, reference.provider);
103    let results = probe_all(ctx, reference, &sources).await;
104    save_cache(ctx, reference, &sources, &results);
105    Ok(results)
106}
107
108pub async fn probe_all(ctx: &Ctx, reference: &ModelRef, sources: &[Source]) -> Vec<ProbeResult> {
109    let timeout = Duration::from_millis(ctx.config.sources.probe_timeout_ms);
110    let mut handles = Vec::with_capacity(sources.len());
111    for source in sources {
112        let ctx = ProbeContext {
113            client: ctx.client.clone(),
114            dirs: ctx.dirs.clone(),
115            config: ctx.config.clone(),
116            cas: ctx.cas.clone(),
117            platform: ctx.platform,
118        };
119        let reference = reference.clone();
120        let source = source.clone();
121        handles.push(tokio::spawn(async move {
122            match tokio::time::timeout(timeout, probe_one(ctx, reference, source.clone())).await {
123                Ok(Ok(result)) => result,
124                _ => ProbeResult::failed(&source.id),
125            }
126        }));
127    }
128    let mut results = Vec::with_capacity(handles.len());
129    for handle in handles {
130        if let Ok(result) = handle.await {
131            results.push(result);
132        }
133    }
134    results
135}
136
137struct ProbeContext {
138    client: reqwest::Client,
139    dirs: crate::dirs::Dirs,
140    config: crate::config::Config,
141    cas: std::sync::Arc<crate::store::Cas>,
142    platform: crate::platform::Platform,
143}
144
145impl ProbeContext {
146    fn as_ctx(&self) -> Ctx {
147        Ctx {
148            dirs: self.dirs.clone(),
149            platform: self.platform,
150            config: self.config.clone(),
151            client: self.client.clone(),
152            cas: self.cas.clone(),
153            show_progress: false,
154        }
155    }
156}
157
158async fn probe_one(
159    probe: ProbeContext,
160    reference: ModelRef,
161    source: Source,
162) -> Result<ProbeResult> {
163    let provider = provider(reference.provider, source.forward_credentials);
164    let ctx = probe.as_ctx();
165    let snapshot = provider
166        .resolve(&ctx, &reference, &source.download_url)
167        .await?;
168    let file = probe_file(&snapshot.files)
169        .ok_or_else(|| Error::other("model source returned no probeable files"))?;
170    let headers = header_map(&file.headers)?;
171    let start = Instant::now();
172    let response = probe
173        .client
174        .get(&file.url)
175        .headers(headers)
176        .header(RANGE, "bytes=0-1048575")
177        .send()
178        .await
179        .map_err(|error| Error::network(&file.url, error))?
180        .error_for_status()
181        .map_err(|error| Error::network(&file.url, error))?;
182    let ttfb = start.elapsed();
183    let mut stream = response.bytes_stream();
184    let body_start = Instant::now();
185    let mut downloaded = 0u64;
186    while let Some(chunk) = stream.next().await {
187        let chunk = chunk.map_err(|error| Error::network(&file.url, error))?;
188        downloaded += chunk.len() as u64;
189        if downloaded >= 1_048_576 {
190            break;
191        }
192    }
193    if downloaded == 0 {
194        return Err(Error::other("model source probe returned no bytes"));
195    }
196    Ok(ProbeResult {
197        source_id: source.id,
198        throughput: downloaded as f64 / body_start.elapsed().as_secs_f64().max(0.001),
199        ttfb_ms: ttfb.as_millis() as u64,
200        ok: true,
201        measured_at: crate::source::now_secs(),
202    })
203}
204
205fn probe_file(files: &[RemoteModelFile]) -> Option<&RemoteModelFile> {
206    files
207        .iter()
208        .filter(|file| file.size.unwrap_or_default() > 0)
209        .max_by_key(|file| file.size.unwrap_or_default())
210        .or_else(|| files.first())
211}
212
213pub fn provider(provider: ProviderId, allow_auth: bool) -> Box<dyn ModelProvider> {
214    match provider {
215        ProviderId::HuggingFace => Box::new(HuggingFace::new(allow_auth)),
216        ProviderId::ModelScope => Box::new(ModelScope::new(allow_auth)),
217    }
218}
219
220fn header_map(headers: &[(String, String)]) -> Result<HeaderMap> {
221    let mut map = HeaderMap::new();
222    for (key, value) in headers {
223        let key = HeaderName::from_bytes(key.as_bytes())
224            .map_err(|error| Error::config(format!("invalid HTTP header `{key}`: {error}")))?;
225        let value = HeaderValue::from_str(value)
226            .map_err(|error| Error::config(format!("invalid HTTP header value: {error}")))?;
227        map.insert(key, value);
228    }
229    Ok(map)
230}
231
232fn cache_path(ctx: &Ctx, reference: &ModelRef, sources: &[Source]) -> std::path::PathBuf {
233    let mut key = format!(
234        "{}:{}@{}",
235        reference.provider, reference.repository, reference.revision
236    );
237    for source in sources {
238        key.push('\0');
239        key.push_str(&source.id);
240        key.push('\0');
241        key.push_str(&source.download_url);
242        key.push('\0');
243        key.push_str(if source.forward_credentials {
244            "credentials"
245        } else {
246            "anonymous"
247        });
248    }
249    let hash = blake3::hash(key.as_bytes()).to_hex().to_string();
250    ctx.dirs
251        .sources_cache()
252        .join("models")
253        .join(reference.provider.as_str())
254        .join(format!("{hash}.json"))
255}
256
257fn fresh_cache(ctx: &Ctx, reference: &ModelRef, sources: &[Source]) -> Option<ProbeCache> {
258    let bytes = std::fs::read(cache_path(ctx, reference, sources)).ok()?;
259    let cache: ProbeCache = serde_json::from_slice(&bytes).ok()?;
260    let now = crate::source::now_secs();
261    let ttl = ctx.config.sources.cache_ttl_secs();
262    (!cache.results.is_empty()
263        && cache
264            .results
265            .iter()
266            .all(|result| now.saturating_sub(result.measured_at) <= ttl))
267    .then_some(cache)
268}
269
270fn save_cache(ctx: &Ctx, reference: &ModelRef, sources: &[Source], results: &[ProbeResult]) {
271    let path = cache_path(ctx, reference, sources);
272    if let Some(parent) = path.parent() {
273        let _ = std::fs::create_dir_all(parent);
274    }
275    if let Ok(bytes) = serde_json::to_vec_pretty(&ProbeCache {
276        results: results.to_vec(),
277    }) {
278        let _ = std::fs::write(path, bytes);
279    }
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285    use std::io::{Read, Write};
286    use std::net::TcpListener;
287    use std::sync::Arc;
288
289    use crate::config::{Config, Settings, ToolSources};
290    use crate::dirs::Dirs;
291    use crate::platform::Platform;
292    use crate::store::Cas;
293
294    #[test]
295    fn model_sources_keep_credentials_off_custom_endpoints_by_default() {
296        let official = default_sources(ProviderId::HuggingFace);
297        assert!(official[0].forward_credentials);
298        let custom = Source::mirror("custom", "https://mirror.example.test", 1);
299        assert!(!custom.forward_credentials);
300    }
301
302    #[tokio::test]
303    async fn source_probe_uses_target_model_range_and_cached_ranking() {
304        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
305        let address = listener.local_addr().unwrap();
306        let server = std::thread::spawn(move || {
307            for request_number in 0..2 {
308                let (mut stream, _) = listener.accept().unwrap();
309                let mut request = Vec::new();
310                let mut buffer = [0u8; 2048];
311                while !request.ends_with(b"\r\n\r\n") {
312                    let read = stream.read(&mut buffer).unwrap();
313                    if read == 0 {
314                        break;
315                    }
316                    request.extend_from_slice(&buffer[..read]);
317                }
318                let request = String::from_utf8(request).unwrap();
319                assert!(!request
320                    .to_ascii_lowercase()
321                    .contains("authorization: bearer"));
322                if request_number == 0 {
323                    let body =
324                        r#"{"sha":"abc123","siblings":[{"rfilename":"weights.bin","size":4}]}"#;
325                    write!(
326                        stream,
327                        "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
328                        body.len(),
329                        body
330                    )
331                    .unwrap();
332                } else {
333                    assert!(request
334                        .to_ascii_lowercase()
335                        .contains("range: bytes=0-1048575"));
336                    stream
337                        .write_all(
338                            b"HTTP/1.1 206 Partial Content\r\nContent-Length: 4\r\nContent-Range: bytes 0-3/4\r\nConnection: close\r\n\r\ndata",
339                        )
340                        .unwrap();
341                }
342            }
343        });
344
345        let temporary = tempfile::tempdir().unwrap();
346        let mut ctx = test_ctx(temporary.path());
347        let mut source = Source::mirror("fixture", &format!("http://{address}"), 0);
348        source.forward_credentials = false;
349        ctx.config.sources.per_tool.insert(
350            "huggingface".into(),
351            ToolSources {
352                custom: vec![source],
353                disable: vec!["official".into()],
354                ..Default::default()
355            },
356        );
357        let reference = ModelRef::parse("hf:owner/repo@main").unwrap();
358        let first = ranked_sources(&ctx, &reference, false).await.unwrap();
359        server.join().unwrap();
360        assert_eq!(first[0].id, "fixture");
361        let second = ranked_sources(&ctx, &reference, false).await.unwrap();
362        assert_eq!(second[0].id, "fixture");
363    }
364
365    fn test_ctx(root: &std::path::Path) -> Ctx {
366        let dirs = Dirs::resolve_from(|key| match key {
367            "OSDK_DATA_DIR" => Some(root.join("data").display().to_string()),
368            "OSDK_CACHE_DIR" => Some(root.join("cache").display().to_string()),
369            "OSDK_CONFIG_DIR" => Some(root.join("config").display().to_string()),
370            _ => None,
371        })
372        .unwrap();
373        dirs.ensure().unwrap();
374        Ctx {
375            dirs: dirs.clone(),
376            platform: Platform::current(),
377            config: Config {
378                settings: Settings::default(),
379                sources: Default::default(),
380                tools: Default::default(),
381                tool_configs: Default::default(),
382                global_tools: Default::default(),
383                global_tool_configs: Default::default(),
384                tool_origins: Default::default(),
385                aliases: Default::default(),
386                project_config_path: None,
387            },
388            client: reqwest::Client::new(),
389            cas: Arc::new(Cas::new(dirs.store.clone())),
390            show_progress: false,
391        }
392    }
393}