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}