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 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}