Skip to main content

wasmhub/
loader.rs

1use crate::cache::CacheManager;
2use crate::error::{Error, Result};
3use crate::manifest::{GlobalManifest, RuntimeManifest};
4use crate::runtime::{Language, Runtime};
5use reqwest::Client;
6use std::path::PathBuf;
7use std::time::Duration;
8
9#[cfg(feature = "progress")]
10use futures_util::StreamExt;
11
12const GITHUB_RELEASES_BASE: &str = "https://github.com/anistark/wasmhub/releases/latest/download";
13const JSDELIVR_BASE: &str = "https://cdn.jsdelivr.net/gh/anistark/wasmhub@latest";
14
15const DEFAULT_MAX_RETRIES: u32 = 3;
16const DEFAULT_INITIAL_BACKOFF_MS: u64 = 500;
17const DEFAULT_MAX_BACKOFF_MS: u64 = 30_000;
18
19#[derive(Debug, Clone)]
20pub enum CdnSource {
21    GitHubReleases,
22    JsDelivr,
23}
24
25impl CdnSource {
26    fn base_url(&self) -> &'static str {
27        match self {
28            CdnSource::GitHubReleases => GITHUB_RELEASES_BASE,
29            CdnSource::JsDelivr => JSDELIVR_BASE,
30        }
31    }
32}
33
34pub struct RuntimeLoader {
35    cache: CacheManager,
36    client: Client,
37    cdn_sources: Vec<CdnSource>,
38    base_url: Option<String>,
39    max_retries: u32,
40    initial_backoff_ms: u64,
41    max_backoff_ms: u64,
42    #[cfg(feature = "progress")]
43    show_progress: bool,
44}
45
46impl RuntimeLoader {
47    pub fn new() -> Result<Self> {
48        Ok(Self {
49            cache: CacheManager::new()?,
50            client: Client::new(),
51            cdn_sources: vec![CdnSource::GitHubReleases, CdnSource::JsDelivr],
52            base_url: None,
53            max_retries: DEFAULT_MAX_RETRIES,
54            initial_backoff_ms: DEFAULT_INITIAL_BACKOFF_MS,
55            max_backoff_ms: DEFAULT_MAX_BACKOFF_MS,
56            #[cfg(feature = "progress")]
57            show_progress: false,
58        })
59    }
60
61    pub fn builder() -> RuntimeLoaderBuilder {
62        RuntimeLoaderBuilder::default()
63    }
64
65    pub async fn get_runtime(&self, language: Language, version: &str) -> Result<Runtime> {
66        if let Some(runtime) = self.cache.get(language, version) {
67            return Ok(runtime);
68        }
69
70        self.download_runtime(language, version).await
71    }
72
73    pub async fn download_runtime(&self, language: Language, version: &str) -> Result<Runtime> {
74        let manifest = self.fetch_runtime_manifest(language).await?;
75        let version_info = manifest
76            .get_version(version)
77            .ok_or_else(|| Error::VersionNotFound {
78                language: language.to_string(),
79                version: version.to_string(),
80            })?;
81
82        let mut last_error = None;
83        for source in &self.cdn_sources {
84            let url = self.build_download_url(source, language, version);
85            match self.download_from_url(&url).await {
86                Ok(data) => {
87                    let computed_hash = self.compute_hash(&data);
88                    if computed_hash != version_info.sha256 {
89                        return Err(Error::IntegrityCheckFailed {
90                            expected: version_info.sha256.clone(),
91                            actual: computed_hash,
92                        });
93                    }
94
95                    return self.cache.store(language, version, &data);
96                }
97                Err(e) => {
98                    last_error = Some(e);
99                    continue;
100                }
101            }
102        }
103
104        Err(last_error.unwrap_or_else(|| Error::Other("All CDN sources failed".to_string())))
105    }
106
107    fn build_download_url(&self, source: &CdnSource, language: Language, version: &str) -> String {
108        let filename = format!("{}-{}.wasm", language.as_str(), version);
109        if let Some(ref base) = self.base_url {
110            return format!("{base}/{filename}");
111        }
112        match source {
113            CdnSource::GitHubReleases => {
114                format!("{}/{}", source.base_url(), filename)
115            }
116            CdnSource::JsDelivr => {
117                format!(
118                    "{}/runtimes/{}/{}",
119                    source.base_url(),
120                    language.as_str(),
121                    filename
122                )
123            }
124        }
125    }
126
127    async fn download_from_url(&self, url: &str) -> Result<Vec<u8>> {
128        let mut last_error = None;
129
130        for attempt in 0..=self.max_retries {
131            if attempt > 0 {
132                let backoff = std::cmp::min(
133                    self.initial_backoff_ms * 2u64.pow(attempt - 1),
134                    self.max_backoff_ms,
135                );
136                tokio::time::sleep(Duration::from_millis(backoff)).await;
137            }
138
139            match self.attempt_download(url).await {
140                Ok(data) => return Ok(data),
141                Err(e) => {
142                    if !Self::is_retryable(&e) {
143                        return Err(e);
144                    }
145                    last_error = Some(e);
146                }
147            }
148        }
149
150        Err(last_error.unwrap_or_else(|| Error::Other("Download failed after retries".to_string())))
151    }
152
153    async fn attempt_download(&self, url: &str) -> Result<Vec<u8>> {
154        #[cfg(feature = "progress")]
155        if self.show_progress {
156            return self.download_with_progress(url).await;
157        }
158
159        let response = self.client.get(url).send().await?;
160        if !response.status().is_success() {
161            return Err(Error::Network(response.error_for_status().unwrap_err()));
162        }
163
164        let bytes = response.bytes().await?;
165        Ok(bytes.to_vec())
166    }
167
168    fn is_retryable(error: &Error) -> bool {
169        match error {
170            Error::Network(e) => {
171                if let Some(status) = e.status() {
172                    status.is_server_error() || status == reqwest::StatusCode::TOO_MANY_REQUESTS
173                } else {
174                    e.is_timeout() || e.is_connect() || e.is_request()
175                }
176            }
177            Error::Io(_) => true,
178            _ => false,
179        }
180    }
181
182    #[cfg(feature = "progress")]
183    async fn download_with_progress(&self, url: &str) -> Result<Vec<u8>> {
184        use indicatif::{ProgressBar, ProgressStyle};
185
186        let response = self.client.get(url).send().await?;
187        if !response.status().is_success() {
188            return Err(Error::Network(response.error_for_status().unwrap_err()));
189        }
190
191        let total_size = response.content_length().unwrap_or(0);
192        let pb = ProgressBar::new(total_size);
193        pb.set_style(
194            ProgressStyle::default_bar()
195                .template("{msg}\n{spinner:.green} [{elapsed_precise}] [{wide_bar:.cyan/blue}] {bytes}/{total_bytes} ({eta})")
196                .unwrap()
197                .progress_chars("#>-"),
198        );
199        pb.set_message(format!("Downloading {url}"));
200
201        let mut downloaded: u64 = 0;
202        let mut stream = response.bytes_stream();
203        let mut data = Vec::new();
204
205        while let Some(chunk) = stream.next().await {
206            let chunk = chunk?;
207            data.extend_from_slice(&chunk);
208            downloaded += chunk.len() as u64;
209            pb.set_position(downloaded);
210        }
211
212        pb.finish_with_message("Download complete");
213        Ok(data)
214    }
215
216    fn compute_hash(&self, data: &[u8]) -> String {
217        use sha2::{Digest, Sha256};
218        let mut hasher = Sha256::new();
219        hasher.update(data);
220        let hash = hasher.finalize();
221        hash.iter().map(|byte| format!("{byte:02x}")).collect()
222    }
223
224    pub async fn list_available(&self) -> Result<GlobalManifest> {
225        self.fetch_global_manifest().await
226    }
227
228    pub async fn get_latest_version(&self, language: Language) -> Result<String> {
229        let manifest = self.fetch_global_manifest().await?;
230        let runtime_info =
231            manifest
232                .get_language(language.as_str())
233                .ok_or_else(|| Error::ManifestNotFound {
234                    language: language.to_string(),
235                })?;
236        Ok(runtime_info.latest.clone())
237    }
238
239    pub fn clear_cache(&self, language: Language, version: &str) -> Result<()> {
240        self.cache.clear(language, version)
241    }
242
243    pub fn clear_all_cache(&self) -> Result<()> {
244        self.cache.clear_all()
245    }
246
247    pub fn list_cached(&self) -> Result<Vec<Runtime>> {
248        self.cache.list()
249    }
250
251    async fn fetch_global_manifest(&self) -> Result<GlobalManifest> {
252        let mut last_error = None;
253        for source in &self.cdn_sources {
254            let url = if let Some(ref base) = self.base_url {
255                format!("{base}/manifest.json")
256            } else {
257                format!("{}/manifest.json", source.base_url())
258            };
259
260            match self.fetch_json(&url).await {
261                Ok(manifest) => return Ok(manifest),
262                Err(e) => {
263                    last_error = Some(e);
264                    continue;
265                }
266            }
267        }
268
269        Err(last_error.unwrap_or_else(|| Error::Other("Failed to fetch manifest".to_string())))
270    }
271
272    pub async fn fetch_runtime_manifest(&self, language: Language) -> Result<RuntimeManifest> {
273        let mut last_error = None;
274        for source in &self.cdn_sources {
275            let url = if let Some(ref base) = self.base_url {
276                format!("{base}/{}-manifest.json", language.as_str())
277            } else {
278                match source {
279                    CdnSource::GitHubReleases => {
280                        format!("{}/{}-manifest.json", source.base_url(), language.as_str())
281                    }
282                    CdnSource::JsDelivr => {
283                        format!(
284                            "{}/runtimes/{}/manifest.json",
285                            source.base_url(),
286                            language.as_str()
287                        )
288                    }
289                }
290            };
291
292            match self.fetch_json(&url).await {
293                Ok(manifest) => return Ok(manifest),
294                Err(e) => {
295                    last_error = Some(e);
296                    continue;
297                }
298            }
299        }
300
301        Err(last_error.unwrap_or_else(|| Error::ManifestNotFound {
302            language: language.to_string(),
303        }))
304    }
305
306    async fn fetch_json<T: serde::de::DeserializeOwned>(&self, url: &str) -> Result<T> {
307        let response = self.client.get(url).send().await?;
308        if !response.status().is_success() {
309            return Err(Error::Network(response.error_for_status().unwrap_err()));
310        }
311        let json = response.json().await?;
312        Ok(json)
313    }
314}
315
316impl Default for RuntimeLoader {
317    fn default() -> Self {
318        Self::new().expect("Failed to create RuntimeLoader")
319    }
320}
321
322#[derive(Default)]
323pub struct RuntimeLoaderBuilder {
324    cache_dir: Option<PathBuf>,
325    cdn_sources: Option<Vec<CdnSource>>,
326    base_url: Option<String>,
327    max_retries: Option<u32>,
328    initial_backoff_ms: Option<u64>,
329    max_backoff_ms: Option<u64>,
330    #[cfg(feature = "progress")]
331    show_progress: bool,
332}
333
334impl RuntimeLoaderBuilder {
335    pub fn new() -> Self {
336        Self::default()
337    }
338
339    pub fn cache_dir(mut self, path: PathBuf) -> Self {
340        self.cache_dir = Some(path);
341        self
342    }
343
344    pub fn cdn_sources(mut self, sources: Vec<CdnSource>) -> Self {
345        self.cdn_sources = Some(sources);
346        self
347    }
348
349    /// Override the base URL for all CDN requests.
350    /// Useful for testing with mock servers.
351    pub fn base_url(mut self, url: String) -> Self {
352        self.base_url = Some(url);
353        self
354    }
355
356    pub fn max_retries(mut self, retries: u32) -> Self {
357        self.max_retries = Some(retries);
358        self
359    }
360
361    pub fn initial_backoff_ms(mut self, ms: u64) -> Self {
362        self.initial_backoff_ms = Some(ms);
363        self
364    }
365
366    pub fn max_backoff_ms(mut self, ms: u64) -> Self {
367        self.max_backoff_ms = Some(ms);
368        self
369    }
370
371    #[cfg(feature = "progress")]
372    pub fn show_progress(mut self, show: bool) -> Self {
373        self.show_progress = show;
374        self
375    }
376
377    pub fn build(self) -> Result<RuntimeLoader> {
378        let cache = if let Some(cache_dir) = self.cache_dir {
379            CacheManager::with_cache_dir(cache_dir)
380        } else {
381            CacheManager::new()?
382        };
383
384        Ok(RuntimeLoader {
385            cache,
386            client: Client::new(),
387            cdn_sources: self
388                .cdn_sources
389                .unwrap_or_else(|| vec![CdnSource::GitHubReleases, CdnSource::JsDelivr]),
390            base_url: self.base_url,
391            max_retries: self.max_retries.unwrap_or(DEFAULT_MAX_RETRIES),
392            initial_backoff_ms: self
393                .initial_backoff_ms
394                .unwrap_or(DEFAULT_INITIAL_BACKOFF_MS),
395            max_backoff_ms: self.max_backoff_ms.unwrap_or(DEFAULT_MAX_BACKOFF_MS),
396            #[cfg(feature = "progress")]
397            show_progress: self.show_progress,
398        })
399    }
400}
401
402#[cfg(test)]
403mod tests {
404    use super::*;
405
406    #[test]
407    fn test_cdn_source_base_url() {
408        assert_eq!(
409            CdnSource::GitHubReleases.base_url(),
410            "https://github.com/anistark/wasmhub/releases/latest/download"
411        );
412        assert_eq!(
413            CdnSource::JsDelivr.base_url(),
414            "https://cdn.jsdelivr.net/gh/anistark/wasmhub@latest"
415        );
416    }
417
418    #[test]
419    fn test_build_download_url() {
420        let loader = RuntimeLoader::new().unwrap();
421
422        let url = loader.build_download_url(&CdnSource::GitHubReleases, Language::Python, "3.11.7");
423        assert!(url.contains("releases/latest/download"));
424        assert!(url.contains("python-3.11.7.wasm"));
425
426        let url = loader.build_download_url(&CdnSource::JsDelivr, Language::Python, "3.11.7");
427        assert!(url.contains("cdn.jsdelivr.net"));
428        assert!(url.contains("runtimes/python/python-3.11.7.wasm"));
429    }
430
431    #[test]
432    fn test_compute_hash() {
433        let loader = RuntimeLoader::new().unwrap();
434        let data = b"test data";
435        let hash = loader.compute_hash(data);
436        assert_eq!(
437            hash,
438            "916f0027a575074ce72a331777c3478d6513f786a591bd892da1a577bf2335f9"
439        );
440    }
441
442    #[test]
443    fn test_builder() {
444        let loader = RuntimeLoader::builder()
445            .cdn_sources(vec![CdnSource::GitHubReleases])
446            .build()
447            .unwrap();
448
449        assert_eq!(loader.cdn_sources.len(), 1);
450    }
451
452    #[test]
453    fn test_builder_with_retry_config() {
454        let loader = RuntimeLoader::builder()
455            .max_retries(5)
456            .initial_backoff_ms(1000)
457            .max_backoff_ms(60_000)
458            .build()
459            .unwrap();
460
461        assert_eq!(loader.max_retries, 5);
462        assert_eq!(loader.initial_backoff_ms, 1000);
463        assert_eq!(loader.max_backoff_ms, 60_000);
464    }
465
466    #[test]
467    fn test_is_retryable() {
468        let io_err = Error::Io(std::io::Error::new(std::io::ErrorKind::TimedOut, "timeout"));
469        assert!(RuntimeLoader::is_retryable(&io_err));
470
471        let integrity_err = Error::IntegrityCheckFailed {
472            expected: "a".to_string(),
473            actual: "b".to_string(),
474        };
475        assert!(!RuntimeLoader::is_retryable(&integrity_err));
476    }
477
478    #[test]
479    fn test_builder_with_cache_dir() {
480        use tempfile::TempDir;
481        let temp_dir = TempDir::new().unwrap();
482
483        let loader = RuntimeLoader::builder()
484            .cache_dir(temp_dir.path().to_path_buf())
485            .build()
486            .unwrap();
487
488        assert!(loader
489            .cache
490            .get_path(Language::Python, "3.11.7")
491            .starts_with(temp_dir.path()));
492    }
493}