From a97bfe6148bdb3a32a526f532ffbd035db95d318 Mon Sep 17 00:00:00 2001 From: Eric Coissac Date: Thu, 11 Jun 2026 08:46:11 +0200 Subject: [PATCH] feat: implement Qobuz CMAF streaming and AES decryption Introduces a new CMAF streaming module for Qobuz that handles progressive segment fetching and full-track decryption. The implementation adds HKDF session key derivation, AES-128-CBC key unwrapping, and AES-128-CTR frame decryption, alongside an ISO BMFF parser for extracting FLAC headers and segment metadata. It also integrates automatic session renewal with double-checked locking, a generic async retry mechanism for transient HTTP errors, and centralized error handling. Finally, it exposes `CmafStreamInfo` and async streaming methods in the public API, updating cryptographic dependencies accordingly. --- Cargo.lock | 45 +++++- pmoqobuz/Cargo.toml | 8 ++ pmoqobuz/src/api/cmaf.rs | 261 +++++++++++++++++++++++++++++++++ pmoqobuz/src/api/mod.rs | 56 ++++++++ pmoqobuz/src/client.rs | 103 +++++++++++++ pmoqobuz/src/cmaf/crypto.rs | 154 ++++++++++++++++++++ pmoqobuz/src/cmaf/error.rs | 17 +++ pmoqobuz/src/cmaf/mod.rs | 278 ++++++++++++++++++++++++++++++++++++ pmoqobuz/src/cmaf/parser.rs | 253 ++++++++++++++++++++++++++++++++ pmoqobuz/src/lib.rs | 2 + pmoqobuz/src/models.rs | 25 ++++ pmoqobuz/src/retry.rs | 194 +++++++++++++++++++++++++ 12 files changed, 1395 insertions(+), 1 deletion(-) create mode 100644 pmoqobuz/src/api/cmaf.rs create mode 100644 pmoqobuz/src/cmaf/crypto.rs create mode 100644 pmoqobuz/src/cmaf/error.rs create mode 100644 pmoqobuz/src/cmaf/mod.rs create mode 100644 pmoqobuz/src/cmaf/parser.rs create mode 100644 pmoqobuz/src/retry.rs diff --git a/Cargo.lock b/Cargo.lock index 450d1a02..a8d9a65d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4,7 +4,7 @@ version = 4 [[package]] name = "PMOMusic" -version = "0.3.49" +version = "0.3.50" dependencies = [ "axum 0.8.7", "console-subscriber", @@ -715,6 +715,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-padding" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a8894febbff9f758034a5b8e12d87918f56dfc64a8e1fe757d65e29041538d93" +dependencies = [ + "generic-array", +] + [[package]] name = "block2" version = "0.6.2" @@ -808,6 +817,15 @@ dependencies = [ "rustversion", ] +[[package]] +name = "cbc" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26b52a9543ae338f279b96b0b9fed9c8093744685043739079ce85cd58f289a6" +dependencies = [ + "cipher", +] + [[package]] name = "cc" version = "1.2.46" @@ -1317,6 +1335,7 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer", "crypto-common", + "subtle", ] [[package]] @@ -2025,6 +2044,24 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hkdf" +version = "0.12.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" +dependencies = [ + "hmac", +] + +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest", +] + [[package]] name = "home" version = "0.5.12" @@ -2408,6 +2445,7 @@ version = "0.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" dependencies = [ + "block-padding", "generic-array", ] @@ -4026,12 +4064,16 @@ dependencies = [ name = "pmoqobuz" version = "0.1.0" dependencies = [ + "aes", "anyhow", "async-trait", "axum 0.8.7", "base64 0.22.1", + "cbc", "chrono", + "ctr", "hex", + "hkdf", "indexmap 2.12.0", "md-5", "mockito", @@ -4051,6 +4093,7 @@ dependencies = [ "serde_json", "serde_yaml", "sha1", + "sha2", "tempfile", "thiserror 2.0.17", "tokio", diff --git a/pmoqobuz/Cargo.toml b/pmoqobuz/Cargo.toml index cfd82d24..aec4e01b 100644 --- a/pmoqobuz/Cargo.toml +++ b/pmoqobuz/Cargo.toml @@ -29,9 +29,17 @@ sha1 = "0.10" hex = "0.4" md-5 = "0.10" +# Crypto CMAF (AES-CTR déchiffrement frames, AES-CBC dérobage clé, HKDF dérivation) +aes = "0.8" +cbc = "0.1" +ctr = "0.9" +hkdf = "0.12" +sha2 = "0.10" + # Cache en mémoire avec TTL moka = { version = "0.12", features = ["future"] } + # Logging tracing = { workspace = true } diff --git a/pmoqobuz/src/api/cmaf.rs b/pmoqobuz/src/api/cmaf.rs new file mode 100644 index 00000000..aaa9cf65 --- /dev/null +++ b/pmoqobuz/src/api/cmaf.rs @@ -0,0 +1,261 @@ +//! Endpoints API CMAF de Qobuz : session/start et file/url. +//! +//! Ces endpoints implémentent le nouveau pipeline de streaming Qobuz +//! qui remplace progressivement `/track/getFileUrl`. + +use serde::Deserialize; +use std::time::{SystemTime, UNIX_EPOCH}; +use tokio::sync::RwLock; +use tracing::info; + +use crate::cmaf::crypto::{compute_request_sig, CMAF_SEED}; +use crate::error::{QobuzError, Result}; + +/// État d'une session CMAF active. +pub struct CmafSession { + pub session_id: String, + pub infos: String, + pub expires_at: u64, +} + +/// Réponse de l'endpoint /session/start. +#[derive(Debug, Deserialize)] +struct SessionStartResponse { + session_id: String, + #[serde(default)] + infos: Option, + expires_at: u64, +} + +/// Réponse de l'endpoint /file/url (CMAF). +#[derive(Debug, Deserialize)] +pub struct CmafFileUrlResponse { + /// Modèle d'URL avec le placeholder `$SEGMENT$`. + #[serde(default)] + pub url_template: Option, + /// Clé de contenu enveloppée, format `"qbz-1.wrapped_b64url.iv_b64url"`. + #[serde(default)] + pub key: Option, + /// Nombre de segments audio (hors segment init). + pub n_segments: u8, + #[serde(default)] + pub format_id: Option, + #[serde(default)] + pub mime_type: Option, + #[serde(default)] + pub sampling_rate: Option, + /// Profondeur de bits (champ v1 de l'API). + #[serde(default)] + pub bits_depth: Option, + /// Profondeur de bits (champ v2 de l'API). + #[serde(default)] + pub bit_depth: Option, +} + +impl CmafFileUrlResponse { + /// Retourne la profondeur de bits en préférant `bits_depth` puis `bit_depth`. + pub fn resolved_bit_depth(&self) -> Option { + self.bits_depth.or(self.bit_depth) + } +} + +const BASE_URL: &str = "https://www.qobuz.com/api.json/0.2"; + +fn current_timestamp() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +/// Signe une requête /session/start. +fn sign_session_start(timestamp: u64) -> String { + let mut args = std::collections::BTreeMap::new(); + args.insert("profile", "qbz-1".to_string()); + compute_request_sig("sessionstart", &args, ×tamp.to_string(), CMAF_SEED) +} + +/// Signe une requête /file/url. +fn sign_file_url(track_id: &str, format_id: u32, timestamp: u64) -> String { + let mut args = std::collections::BTreeMap::new(); + args.insert("format_id", format_id.to_string()); + args.insert("intent", "stream".to_string()); + args.insert("track_id", track_id.to_string()); + compute_request_sig("fileurl", &args, ×tamp.to_string(), CMAF_SEED) +} + +/// Gestionnaire de session CMAF avec renouvellement automatique. +/// +/// Utilise le pattern double-checked lock : fast path sous read guard, +/// slow path sous write guard exclusif pour éviter les renouvellements +/// concurrents qui produiraient des `infos` incohérents avec les `key` +/// retournées par `/file/url`. +pub struct CmafSessionManager { + session: RwLock>, +} + +impl CmafSessionManager { + pub fn new() -> Self { + Self { session: RwLock::new(None) } + } + + /// Retourne `(session_id, infos)` d'une session valide, en en démarrant + /// une nouvelle si la session courante est absente ou expire dans < 60s. + pub async fn ensure_session( + &self, + http: &reqwest::Client, + app_id: &str, + auth_token: &str, + ) -> Result<(String, String)> { + let now = current_timestamp(); + + // Fast path : session existante avec > 60s restants. + { + let guard = self.session.read().await; + if let Some(ref cs) = *guard { + if cs.expires_at > now + 60 { + return Ok((cs.session_id.clone(), cs.infos.clone())); + } + } + } + + // Slow path : prend le verrou d'écriture et vérifie à nouveau. + let mut guard = self.session.write().await; + if let Some(ref cs) = *guard { + if cs.expires_at > now + 60 { + return Ok((cs.session_id.clone(), cs.infos.clone())); + } + } + + info!("[CMAF] Démarrage d'une nouvelle session"); + let timestamp = current_timestamp(); + let sig = sign_session_start(timestamp); + + let url = format!("{}/session/start", BASE_URL); + let response = http + .post(&url) + .header("X-App-Id", app_id) + .header("X-User-Auth-Token", auth_token) + .form(&[ + ("profile", "qbz-1"), + ("request_ts", ×tamp.to_string()), + ("request_sig", &sig), + ]) + .send() + .await + .map_err(QobuzError::Http)?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + return Err(QobuzError::ApiError { + code: status.as_u16(), + message: format!("session/start échoué: {}", body), + }); + } + + let resp: SessionStartResponse = response.json().await.map_err(|e| { + QobuzError::Other(format!("parse session/start: {}", e)) + })?; + + let infos = resp.infos.unwrap_or_default(); + info!( + "[CMAF] Session démarrée: id={}..., expires_at={}", + &resp.session_id[..resp.session_id.len().min(8)], + resp.expires_at, + ); + + let session_id = resp.session_id.clone(); + let infos_clone = infos.clone(); + + *guard = Some(CmafSession { + session_id: resp.session_id, + infos, + expires_at: resp.expires_at, + }); + + Ok((session_id, infos_clone)) + } + + /// Invalide la session courante (utile en cas d'erreur de déchiffrement). + pub async fn invalidate(&self) { + *self.session.write().await = None; + } +} + +/// Récupère l'URL de fichier CMAF pour un track. +/// +/// Inclut le retry sur les erreurs transitoires (5xx, 429, erreurs réseau). +/// Un 404 retourne immédiatement une erreur `NotFound`. +pub async fn get_file_url( + http: &reqwest::Client, + app_id: &str, + auth_token: &str, + session_id: &str, + track_id: &str, + format_id: u32, +) -> Result { + use crate::retry::{classify_reqwest, classify_status, retry_transient, DEFAULT_MAX_ATTEMPTS}; + + let url = format!("{}/file/url", BASE_URL); + + let result = retry_transient( + DEFAULT_MAX_ATTEMPTS, + "CMAF file/url", + |e: &QobuzError| matches!(e, QobuzError::RateLimitExceeded | QobuzError::ApiError { code: 500..=599, .. }), + |_attempt| { + let url = url.clone(); + async move { + let timestamp = current_timestamp(); + let sig = sign_file_url(track_id, format_id, timestamp); + + let response = http + .get(&url) + .header("X-App-Id", app_id) + .header("X-User-Auth-Token", auth_token) + .header("X-Session-Id", session_id) + .query(&[ + ("track_id", track_id), + ("format_id", &format_id.to_string()), + ("intent", "stream"), + ("request_ts", ×tamp.to_string()), + ("request_sig", &sig), + ]) + .send() + .await + .map_err(QobuzError::Http)?; + + let status = response.status(); + tracing::info!("[CMAF] file/url track_id={} format_id={} status={}", track_id, format_id, status); + + if !status.is_success() { + let code = status.as_u16(); + return Err(match code { + 404 => QobuzError::NotFound(format!("track {} non disponible", track_id)), + 429 => QobuzError::RateLimitExceeded, + _ => QobuzError::ApiError { + code, + message: format!("file/url status {}", code), + }, + }); + } + + let file_url: CmafFileUrlResponse = response.json().await.map_err(|e| { + QobuzError::Other(format!("parse file/url: {}", e)) + })?; + + tracing::info!( + "[CMAF] file/url: n_segments={}, mime={:?}, sampling_rate={:?}", + file_url.n_segments, + file_url.mime_type, + file_url.sampling_rate, + ); + + Ok(file_url) + } + }, + ) + .await?; + + Ok(result) +} diff --git a/pmoqobuz/src/api/mod.rs b/pmoqobuz/src/api/mod.rs index 88acdd1b..d11ee16a 100644 --- a/pmoqobuz/src/api/mod.rs +++ b/pmoqobuz/src/api/mod.rs @@ -4,10 +4,12 @@ pub mod auth; pub mod catalog; +pub mod cmaf; pub mod signing; pub mod spoofer; pub mod user; +use crate::api::cmaf::CmafSessionManager; use crate::error::{QobuzError, Result}; use crate::models::AudioFormat; use reqwest::{Client, Response}; @@ -95,6 +97,8 @@ pub struct QobuzApi { user_id: RwLock>, /// Format audio par défaut format_id: AudioFormat, + /// Gestionnaire de session CMAF (renouvellement automatique thread-safe) + pub(crate) cmaf_session: CmafSessionManager, } impl QobuzApi { @@ -114,6 +118,7 @@ impl QobuzApi { user_auth_token: RwLock::new(None), user_id: RwLock::new(None), format_id: AudioFormat::default(), + cmaf_session: CmafSessionManager::new(), }) } @@ -316,6 +321,57 @@ impl QobuzApi { self.handle_response(response, endpoint, params).await } + /// Retourne le client HTTP interne (utilisé par les endpoints CMAF). + pub(crate) fn http_client(&self) -> &Client { + &self.client + } + + /// Assure qu'une session CMAF valide existe et retourne `(session_id, infos)`. + /// + /// Crée ou renouvelle automatiquement la session si elle est absente ou expire + /// dans moins de 60 secondes. Les appels concurrents sont sérialisés pour + /// éviter des sessions incohérentes. + pub async fn ensure_cmaf_session(&self) -> Result<(String, String)> { + let app_id = self.app_id.read().unwrap().clone(); + let auth_token = self + .user_auth_token + .read() + .unwrap() + .clone() + .ok_or_else(|| QobuzError::Unauthorized("Token d'authentification manquant pour CMAF".into()))?; + + self.cmaf_session + .ensure_session(&self.client, &app_id, &auth_token) + .await + } + + /// Récupère l'URL de fichier CMAF pour un track avec le format donné. + pub async fn get_cmaf_file_url( + &self, + track_id: &str, + format_id: u32, + ) -> Result { + let app_id = self.app_id.read().unwrap().clone(); + let auth_token = self + .user_auth_token + .read() + .unwrap() + .clone() + .ok_or_else(|| QobuzError::Unauthorized("Token d'authentification manquant pour CMAF".into()))?; + + let (session_id, _infos) = self.ensure_cmaf_session().await?; + + crate::api::cmaf::get_file_url( + &self.client, + &app_id, + &auth_token, + &session_id, + track_id, + format_id, + ) + .await + } + /// Traite la réponse HTTP async fn handle_response( &self, diff --git a/pmoqobuz/src/client.rs b/pmoqobuz/src/client.rs index 4ad8b8e4..fb97c991 100644 --- a/pmoqobuz/src/client.rs +++ b/pmoqobuz/src/client.rs @@ -973,6 +973,109 @@ impl QobuzClient { }) .await } + + // ============ CMAF (streaming moderne) ============ + + /// Prépare le streaming CMAF pour un track : dérive les clés et fetche + /// le segment d'initialisation. Retourne une `CmafStreamInfo` prête à l'emploi. + /// + /// Pour télécharger la totalité du FLAC déchiffré d'un coup, utilisez + /// plutôt `download_cmaf_full`. Pour streamer segment par segment, utilisez + /// les données de `CmafStreamInfo` avec `cmaf::fetch_all_segments`. + /// + /// # Erreurs + /// + /// * `QobuzError::NotFound` — track non disponible sur Qobuz + /// * `QobuzError::Unauthorized` — session expirée (réparée automatiquement) + /// * `QobuzError::Other` — erreur CMAF (clé invalide, segment init corrompu, etc.) + pub async fn get_cmaf_stream_info( + &self, + track_id: &str, + ) -> Result { + let format_id = self.api.format().id() as u32; + + let file_url = self + .call_with_auth_repair("get_cmaf_file_url", || { + self.api.get_cmaf_file_url(track_id, format_id) + }) + .await?; + + let bit_depth = file_url.resolved_bit_depth(); + let url_template = file_url + .url_template + .ok_or_else(|| QobuzError::Other("CMAF file/url: url_template absent".into()))?; + let key_str = file_url + .key + .ok_or_else(|| QobuzError::Other("CMAF file/url: key absent".into()))?; + + let (_session_id, infos) = self + .call_with_auth_repair("ensure_cmaf_session", || { + self.api.ensure_cmaf_session() + }) + .await?; + + let setup = crate::cmaf::setup_streaming( + url_template, + &key_str, + &infos, + file_url.n_segments, + file_url.format_id.unwrap_or(format_id), + file_url.sampling_rate, + bit_depth, + ) + .await?; + + Ok(crate::models::CmafStreamInfo { + url_template: setup.url_template, + n_segments: setup.n_segments, + content_key: setup.content_key, + flac_header: setup.flac_header, + segment_table: setup.segment_table, + format_id: setup.format_id, + sampling_rate: setup.sampling_rate, + bit_depth: setup.bit_depth, + }) + } + + /// Télécharge un track complet via CMAF et retourne les bytes FLAC déchiffrés. + /// + /// Bloquant jusqu'à la fin du téléchargement. Pour les gros fichiers Hi-Res, + /// préférer `get_cmaf_stream_info` + streaming segment par segment. + pub async fn download_cmaf_full(&self, track_id: &str) -> Result> { + let format_id = self.api.format().id() as u32; + + let file_url = self + .call_with_auth_repair("get_cmaf_file_url", || { + self.api.get_cmaf_file_url(track_id, format_id) + }) + .await?; + + let bit_depth = file_url.resolved_bit_depth(); + let url_template = file_url + .url_template + .ok_or_else(|| QobuzError::Other("CMAF file/url: url_template absent".into()))?; + let key_str = file_url + .key + .ok_or_else(|| QobuzError::Other("CMAF file/url: key absent".into()))?; + + let (_session_id, infos) = self + .call_with_auth_repair("ensure_cmaf_session", || { + self.api.ensure_cmaf_session() + }) + .await?; + + crate::cmaf::download_full( + url_template, + &key_str, + &infos, + file_url.n_segments, + file_url.format_id.unwrap_or(format_id), + file_url.sampling_rate, + bit_depth, + None, + ) + .await + } } #[cfg(test)] diff --git a/pmoqobuz/src/cmaf/crypto.rs b/pmoqobuz/src/cmaf/crypto.rs new file mode 100644 index 00000000..dc27ffc9 --- /dev/null +++ b/pmoqobuz/src/cmaf/crypto.rs @@ -0,0 +1,154 @@ +use aes::cipher::{BlockDecryptMut, KeyIvInit, StreamCipher}; +use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; +use hkdf::Hkdf; +use md5::{Digest, Md5}; +use sha2::Sha256; + +use super::error::CmafError; + +type Aes128CbcDec = cbc::Decryptor; +type Aes128Ctr = ctr::Ctr128BE; + +/// Seed publique extraite du bundle web Qobuz. Valeur IKM pour HKDF. +pub const CMAF_SEED: &str = "abb21364945c0583309667d13ca3d93a"; + +fn hex_decode(hex: &str) -> Vec { + (0..hex.len()) + .step_by(2) + .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap_or(0)) + .collect() +} + +/// Dérive la clé de session 16 octets depuis le champ `infos` de session/start. +/// +/// Format infos : `"salt_b64url.info_b64url"` +/// `seed` est le CMAF_SEED hex-encodé utilisé comme IKM HKDF. +pub fn derive_session_key(seed: &str, infos: &str) -> Result<[u8; 16], CmafError> { + let parts: Vec<&str> = infos.split('.').collect(); + if parts.len() < 2 { + return Err(CmafError::InvalidInfos( + "session infos doit avoir au moins 2 parties séparées par des points".into(), + )); + } + + let salt = URL_SAFE_NO_PAD.decode(parts[0])?; + let info = URL_SAFE_NO_PAD.decode(parts[1])?; + let ikm = hex_decode(seed); + + let hk = Hkdf::::new(Some(&salt), &ikm); + let mut okm = [0u8; 16]; + hk.expand(&info, &mut okm).map_err(|_| CmafError::HkdfExpand)?; + + Ok(okm) +} + +/// Déroule la clé de contenu par track avec la clé de session. +/// +/// Format key_str : `"qbz-1.wrapped_key_b64url.iv_b64url"` +pub fn unwrap_content_key(session_key: &[u8; 16], key_str: &str) -> Result<[u8; 16], CmafError> { + let parts: Vec<&str> = key_str.split('.').collect(); + if parts.len() < 3 { + return Err(CmafError::InvalidKey( + "key string doit avoir au moins 3 parties séparées par des points".into(), + )); + } + + let wrapped = URL_SAFE_NO_PAD.decode(parts[1])?; + let iv = URL_SAFE_NO_PAD.decode(parts[2])?; + + if iv.len() != 16 { + return Err(CmafError::InvalidKey(format!( + "IV de dérobage doit faire 16 octets, reçu {}", + iv.len() + ))); + } + + let mut buf = wrapped.clone(); + let decrypted = + Aes128CbcDec::new(session_key.into(), iv.as_slice().into()) + .decrypt_padded_mut::(&mut buf) + .map_err(|e| CmafError::AesDecrypt(format!("AES-CBC unwrap échoué: {e}")))?; + + if decrypted.len() != 16 { + return Err(CmafError::InvalidKey(format!( + "clé déroullée doit faire 16 octets, obtenu {}", + decrypted.len() + ))); + } + + let mut key = [0u8; 16]; + key.copy_from_slice(decrypted); + Ok(key) +} + +/// Déchiffre une frame FLAC en place avec AES-128-CTR. +/// +/// `iv_8` = IV 8 octets du segment UUID box, complété à zéro jusqu'à 16 octets. +pub fn decrypt_frame(content_key: &[u8; 16], iv_8: &[u8; 8], data: &mut [u8]) { + let mut nonce = [0u8; 16]; + nonce[..8].copy_from_slice(iv_8); + Aes128Ctr::new(content_key.into(), &nonce.into()).apply_keystream(data); +} + +/// Calcule la signature MD5 pour les appels API CMAF de Qobuz. +/// +/// Concatène method + paires clé-valeur triées + timestamp + seed, +/// puis retourne le digest MD5 hexadécimal minuscule. +pub fn compute_request_sig( + method: &str, + args: &std::collections::BTreeMap<&str, String>, + timestamp: &str, + seed: &str, +) -> String { + let mut hasher = Md5::new(); + hasher.update(method.as_bytes()); + for (k, v) in args { + hasher.update(k.as_bytes()); + hasher.update(v.as_bytes()); + } + hasher.update(timestamp.as_bytes()); + hasher.update(seed.as_bytes()); + + format!("{:x}", hasher.finalize()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_compute_request_sig_deterministe() { + let mut args = std::collections::BTreeMap::new(); + args.insert("profile", "qbz-1".to_string()); + let sig1 = compute_request_sig("sessionstart", &args, "1775500000", CMAF_SEED); + let sig2 = compute_request_sig("sessionstart", &args, "1775500000", CMAF_SEED); + assert_eq!(sig1.len(), 32); + assert_eq!(sig1, sig2); + } + + #[test] + fn test_decrypt_frame_aller_retour() { + let key = [1u8, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]; + let iv = [1u8, 2, 3, 4, 5, 6, 7, 8]; + let original = b"Hello FLAC frame data here!".to_vec(); + let mut data = original.clone(); + decrypt_frame(&key, &iv, &mut data); + assert_ne!(data, original); + // AES-CTR est son propre inverse + decrypt_frame(&key, &iv, &mut data); + assert_eq!(data, original); + } + + #[test] + fn test_derive_session_key_infos_invalide() { + let result = derive_session_key(CMAF_SEED, "pas_de_point"); + assert!(result.is_err()); + } + + #[test] + fn test_unwrap_content_key_format_invalide() { + let key = [0u8; 16]; + let result = unwrap_content_key(&key, "seulement.deux"); + assert!(result.is_err()); + } +} diff --git a/pmoqobuz/src/cmaf/error.rs b/pmoqobuz/src/cmaf/error.rs new file mode 100644 index 00000000..f8e13f70 --- /dev/null +++ b/pmoqobuz/src/cmaf/error.rs @@ -0,0 +1,17 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum CmafError { + #[error("Invalid infos format: {0}")] + InvalidInfos(String), + #[error("Invalid key format: {0}")] + InvalidKey(String), + #[error("Base64 decode error: {0}")] + Base64(#[from] base64::DecodeError), + #[error("HKDF expand error")] + HkdfExpand, + #[error("AES decrypt error: {0}")] + AesDecrypt(String), + #[error("CMAF parse error: {0}")] + ParseError(String), +} diff --git a/pmoqobuz/src/cmaf/mod.rs b/pmoqobuz/src/cmaf/mod.rs new file mode 100644 index 00000000..1fff9e0f --- /dev/null +++ b/pmoqobuz/src/cmaf/mod.rs @@ -0,0 +1,278 @@ +//! Pipeline CMAF (Common Media Application Format) pour Qobuz. +//! +//! Qobuz utilise CMAF avec chiffrement AES-CTR par frame sur CDN Akamai. +//! C'est le pipeline de l'app Android v9.7+ qui remplace l'endpoint legacy +//! `/track/getFileUrl`. +//! +//! # Pipeline +//! +//! 1. `/session/start` → `{ session_id, infos, expires_at }` +//! 2. `/file/url` → `{ url_template, key (enveloppé), n_segments, ... }` +//! 3. Session key = `HKDF(CMAF_SEED, infos)` +//! 4. Content key = AES-CBC-unwrap(session_key, key) +//! 5. Segment init (s=0) → header FLAC + table des segments +//! 6. Pour chaque s=1..n_segments : fetch → parse crypto boxes → déchiffrement AES-CTR + +pub mod crypto; +pub mod error; +pub mod parser; + +pub use crypto::{compute_request_sig, decrypt_frame, derive_session_key, unwrap_content_key, CMAF_SEED}; +pub use error::CmafError; +pub use parser::{ + parse_init_segment, parse_segment_crypto, FrameEntry, InitInfo, SegmentCrypto, + SegmentTableEntry, +}; + +use std::sync::Arc; +use tokio::sync::Semaphore; +use tracing::{debug, info, warn}; + +use crate::error::{QobuzError, Result}; +use crate::retry::{classify_reqwest, classify_status, retry_transient, FetchError, DEFAULT_MAX_ATTEMPTS}; + +/// Concurrence max pour le fetch de segments CMAF. +/// 3 segments en vol est le compromis optimal — le CDN Akamai rate-limite +/// au-delà de ~5 requêtes parallèles par IP sur des fenêtres de 1s. +pub const CMAF_PREFETCH_CONCURRENCY: usize = 3; + +/// Callback de progression pour les fonctions de téléchargement. +pub type CmafProgressCallback = Arc; + +/// Un tick de progression. `segments_completed` est cumulatif (1..=n). +#[derive(Debug, Clone, Copy)] +pub struct CmafProgressUpdate { + pub segments_completed: u32, + pub n_segments: u32, + pub bytes_this_segment: u64, +} + +/// Info réunies depuis le segment init, suffisantes pour démarrer le streaming. +pub struct CmafStreamingInfo { + pub url_template: String, + pub n_segments: u8, + pub content_key: [u8; 16], + pub flac_header: Vec, + pub segment_table: Vec, + pub format_id: u32, + pub sampling_rate: Option, + pub bit_depth: Option, +} + +/// Construit un client reqwest dédié aux fetches CDN Akamai. +fn build_cdn_client() -> Result { + reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs(10)) + .build() + .map_err(|e| QobuzError::Http(e)) +} + +/// Fetch une URL CDN en bytes avec retry sur les erreurs transitoires. +/// Un 404/403 échoue immédiatement (terminal). 5xx et 429 → retry avec backoff. +async fn fetch_bytes_with_retry( + http: &reqwest::Client, + url: &str, + log_tag: &str, +) -> std::result::Result, FetchError> { + retry_transient( + DEFAULT_MAX_ATTEMPTS, + log_tag, + FetchError::is_transient, + |_attempt| async move { + let response = http + .get(url) + .header("User-Agent", "Mozilla/5.0") + .send() + .await + .map_err(|e| classify_reqwest(&e, "fetch"))?; + let status = response.status(); + if !status.is_success() { + return Err(classify_status(status, "fetch")); + } + response + .bytes() + .await + .map(|b| b.to_vec()) + .map_err(|e| classify_reqwest(&e, "lecture")) + }, + ) + .await +} + +/// Fetch les segments 1..=n_segments avec contrôle de concurrence. +/// Déclenche le callback de progression une fois par segment complété. +/// Les résultats sont retournés triés par index de segment. +async fn fetch_all_segments( + http: &reqwest::Client, + url_template: &str, + n_segments: u8, + log_tag: &str, + on_progress: Option, +) -> Result>> { + let semaphore = Arc::new(Semaphore::new(CMAF_PREFETCH_CONCURRENCY)); + let completed_count = Arc::new(std::sync::atomic::AtomicU32::new(0)); + let mut handles = Vec::with_capacity(n_segments as usize); + + for seg_idx in 1u8..=n_segments { + let sem = semaphore.clone(); + let http = http.clone(); + let seg_url = url_template.replace("$SEGMENT$", &seg_idx.to_string()); + let log_tag = log_tag.to_string(); + let progress = on_progress.clone(); + let counter = completed_count.clone(); + + handles.push(tokio::spawn(async move { + let permit = sem.acquire_owned().await + .map_err(|e| format!("semaphore: {}", e))?; + + let seg_data = fetch_bytes_with_retry(&http, &seg_url, &format!("{} seg {}", log_tag, seg_idx)) + .await + .map_err(|e| format!("[{}] seg {} fetch: {}", log_tag, seg_idx, e))?; + + let bytes_this_segment = seg_data.len() as u64; + if let Some(cb) = progress { + let done = counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1; + cb(CmafProgressUpdate { + segments_completed: done, + n_segments: n_segments as u32, + bytes_this_segment, + }); + } + + // Pause avant de libérer le slot pour respecter les limites CDN + tokio::time::sleep(std::time::Duration::from_millis(500)).await; + drop(permit); + + Ok::<(u8, Vec), String>((seg_idx, seg_data)) + })); + } + + let mut segments: Vec<(u8, Vec)> = Vec::with_capacity(handles.len()); + for handle in handles { + let (idx, data) = handle + .await + .map_err(|e| QobuzError::Other(format!("[{}] panic task: {}", log_tag, e)))? + .map_err(|e| QobuzError::Other(format!("[{}] téléchargement échoué: {}", log_tag, e)))?; + segments.push((idx, data)); + } + segments.sort_by_key(|(idx, _)| *idx); + Ok(segments.into_iter().map(|(_, data)| data).collect()) +} + +/// Déchiffre une séquence de segments CMAF chiffrés et écrit les frames FLAC dans `output`. +/// +/// Optimisation hot-path : extend + decrypt in-place plutôt que copie triple. +pub fn decrypt_segments_into( + segments: &[Vec], + content_key: &[u8; 16], + output: &mut Vec, +) -> Result<()> { + for (seg_idx, seg_data) in segments.iter().enumerate() { + let log_idx = seg_idx + 1; + let crypto = parse_segment_crypto(seg_data) + .map_err(|e| QobuzError::Other(format!("CMAF seg {} parse: {}", log_idx, e)))?; + + let mut data_pos = crypto.data_offset; + for entry in &crypto.entries { + let frame_end = data_pos + entry.size as usize; + if frame_end > seg_data.len() { + return Err(QobuzError::Other(format!("CMAF seg {} débordement frame", log_idx))); + } + let output_start = output.len(); + output.extend_from_slice(&seg_data[data_pos..frame_end]); + if entry.flags != 0 { + decrypt_frame(content_key, &entry.iv, &mut output[output_start..]); + } + data_pos = frame_end; + } + if data_pos < crypto.mdat_end && crypto.mdat_end <= seg_data.len() { + output.extend_from_slice(&seg_data[data_pos..crypto.mdat_end]); + } + } + Ok(()) +} + +/// Prépare le streaming CMAF : dérive les clés, fetche le segment init. +/// Ne télécharge PAS les segments audio — l'appelant les streame en arrière-plan. +pub async fn setup_streaming( + url_template: String, + key_str: &str, + infos: &str, + n_segments: u8, + format_id: u32, + sampling_rate: Option, + bit_depth: Option, +) -> Result { + let session_key = derive_session_key(CMAF_SEED, infos) + .map_err(|e| QobuzError::Other(format!("dérivation clé session: {}", e)))?; + let content_key = unwrap_content_key(&session_key, key_str) + .map_err(|e| QobuzError::Other(format!("dérobage clé contenu: {}", e)))?; + + let http = build_cdn_client()?; + let init_url = url_template.replace("$SEGMENT$", "0"); + + info!("[CMAF] Fetch segment init: {}", &init_url[..init_url.len().min(60)]); + + let init_data = fetch_bytes_with_retry(&http, &init_url, "CMAF init") + .await + .map_err(|e| QobuzError::Other(format!("fetch segment init: {}", e)))?; + + let init_info = parse_init_segment(&init_data) + .map_err(|e| QobuzError::Other(format!("parse segment init: {}", e)))?; + + info!( + "[CMAF] Init: header FLAC {}B, {} segments dans la table, n_segments API={}", + init_info.flac_header.len(), + init_info.segment_table.len(), + n_segments, + ); + if init_info.segment_table.len() != n_segments as usize { + warn!( + "[CMAF] ÉCART: table={} entrées mais API dit n_segments={}", + init_info.segment_table.len(), + n_segments, + ); + } + + Ok(CmafStreamingInfo { + url_template, + n_segments, + content_key, + flac_header: init_info.flac_header, + segment_table: init_info.segment_table, + format_id, + sampling_rate, + bit_depth, + }) +} + +/// Télécharge un track CMAF complet et retourne les bytes FLAC déchiffrés. +pub async fn download_full( + url_template: String, + key_str: &str, + infos: &str, + n_segments: u8, + format_id: u32, + sampling_rate: Option, + bit_depth: Option, + on_progress: Option, +) -> Result> { + let setup = setup_streaming(url_template, key_str, infos, n_segments, format_id, sampling_rate, bit_depth).await?; + let http = build_cdn_client()?; + + let total_size: usize = setup.flac_header.len() + + setup.segment_table.iter().map(|s| s.byte_len as usize).sum::(); + + let segments = fetch_all_segments(&http, &setup.url_template, setup.n_segments, "CMAF-FULL", on_progress).await?; + + let mut output = Vec::with_capacity(total_size); + output.extend_from_slice(&setup.flac_header); + decrypt_segments_into(&segments, &setup.content_key, &mut output)?; + + debug!( + "[CMAF-FULL] Complet: {:.2} MB FLAC, attendu {:.2} MB", + output.len() as f64 / (1024.0 * 1024.0), + total_size as f64 / (1024.0 * 1024.0), + ); + Ok(output) +} diff --git a/pmoqobuz/src/cmaf/parser.rs b/pmoqobuz/src/cmaf/parser.rs new file mode 100644 index 00000000..5a234ab2 --- /dev/null +++ b/pmoqobuz/src/cmaf/parser.rs @@ -0,0 +1,253 @@ +use super::error::CmafError; + +const QBZ_INIT_UUID: [u8; 16] = [ + 0xc7, 0xc7, 0x5d, 0xf0, 0xfd, 0xd9, 0x51, 0xe9, + 0x8f, 0xc2, 0x29, 0x71, 0xe4, 0xac, 0xf8, 0xd2, +]; +const QBZ_SEGMENT_UUID: [u8; 16] = [ + 0x3b, 0x42, 0x12, 0x92, 0x56, 0xf3, 0x5f, 0x75, + 0x92, 0x36, 0x63, 0xb6, 0x9a, 0x1f, 0x52, 0xb2, +]; +const FLAC_MAGIC: &[u8; 4] = b"fLaC"; + +/// Taille en octets et nombre d'échantillons d'un segment. +#[derive(Debug, Clone)] +pub struct SegmentTableEntry { + /// Taille des données FLAC déchiffrées de ce segment. + pub byte_len: u32, + /// Nombre d'échantillons audio dans ce segment. + pub sample_count: u32, +} + +/// Header FLAC et table des segments extraits du segment d'initialisation. +pub struct InitInfo { + pub flac_header: Vec, + /// Tailles par segment (indices 0..n_segments-1 correspondent aux segments 1..n_segments). + pub segment_table: Vec, +} + +/// Une entrée de frame dans le segment UUID box. +pub struct FrameEntry { + pub size: u32, + pub flags: u16, + pub iv: [u8; 8], +} + +/// Informations crypto parsées depuis le UUID box d'un segment audio. +pub struct SegmentCrypto { + /// Offset vers le début des données audio (payload mdat). + pub data_offset: usize, + /// Fin du contenu de la mdat box. + pub mdat_end: usize, + pub entries: Vec, +} + +/// Parcourt les boxes ISO BMFF et trouve le premier UUID box correspondant à `target_uuid`. +/// Retourne `(payload_start, box_end)` où payload_start est après les 16 octets UUID. +fn find_uuid_box(data: &[u8], target_uuid: &[u8; 16]) -> Option<(usize, usize)> { + let mut pos = 0; + while pos + 8 <= data.len() { + let size = read_box_size(data, pos); + if size < 8 || pos + size > data.len() { + break; + } + if &data[pos + 4..pos + 8] == b"uuid" && pos + 24 <= data.len() { + if &data[pos + 8..pos + 24] == target_uuid.as_ref() { + return Some((pos + 24, pos + size)); + } + } + pos += size; + } + None +} + +/// Parse le segment d'initialisation (segment 0) pour extraire le header FLAC et la table des segments. +pub fn parse_init_segment(data: &[u8]) -> Result { + let (payload_start, box_end) = find_uuid_box(data, &QBZ_INIT_UUID) + .ok_or_else(|| CmafError::ParseError("segment init: QBZ_INIT_UUID box non trouvé".into()))?; + + let payload = &data[payload_start..box_end]; + parse_init_uuid_payload(payload) +} + +/// Parse un segment audio pour extraire les informations crypto par frame. +pub fn parse_segment_crypto(data: &[u8]) -> Result { + let mut uuid_box_start: Option = None; + let mut mdat_end = data.len(); + + let mut pos = 0; + while pos + 8 <= data.len() { + let size = read_box_size(data, pos); + if size < 8 || pos + size > data.len() { + break; + } + let box_type = &data[pos + 4..pos + 8]; + if box_type == b"uuid" && pos + 24 <= data.len() { + if &data[pos + 8..pos + 24] == QBZ_SEGMENT_UUID.as_ref() { + uuid_box_start = Some(pos); + } + } else if box_type == b"mdat" { + mdat_end = pos + size; + } + pos += size; + } + + let box_start = uuid_box_start + .ok_or_else(|| CmafError::ParseError("segment audio: QBZ_SEGMENT_UUID box non trouvé".into()))?; + + parse_segment_uuid_payload(data, box_start, mdat_end) +} + +fn parse_init_uuid_payload(payload: &[u8]) -> Result { + // Layout payload: + // [4B padding/version] + // [4B track_id] + // [4B file_id] + // [4B sample_rate] + // [1B bits_per_sample] + // [1B channels + 2B padding] + // [6B total_samples_count] + // [2B raw_data_len] + // [raw_data_len bytes: contient le header FLAC] + // [1B key_id_len] + // [key_id_len bytes: key_id] + // [2B segment_count] + // Par segment: [4B byte_len][4B sample_count] + + if payload.len() < 28 { + return Err(CmafError::ParseError("payload init UUID trop court".into())); + } + + let mut a = 4; // version/padding + a += 4; // track_id + a += 4; // file_id + a += 4; // sample_rate + a += 1; // bits_per_sample + a += 3; // channels + padding + a += 6; // total_samples_count + + if a + 2 > payload.len() { + return Err(CmafError::ParseError("payload init UUID tronqué au raw_len".into())); + } + let raw_len = u16::from_be_bytes([payload[a], payload[a + 1]]) as usize; + a += 2; + + let raw_data = &payload[a..a + raw_len.min(payload.len() - a)]; + a += raw_len; + + let flac_pos = raw_data + .windows(4) + .position(|w| w == FLAC_MAGIC) + .ok_or_else(|| CmafError::ParseError("payload init UUID: magic fLaC non trouvé".into()))?; + + // fLaC (4) + STREAMINFO block header (4) + STREAMINFO data (34) = 42 octets + let header_len = 4 + 4 + 34; + if flac_pos + header_len > raw_data.len() { + return Err(CmafError::ParseError("payload init UUID: STREAMINFO tronqué".into())); + } + + let mut flac_header = raw_data[flac_pos..flac_pos + header_len].to_vec(); + // Marquer le dernier bloc de métadonnées + flac_header[4] |= 0x80; + + if a + 1 > payload.len() { + return Ok(InitInfo { flac_header, segment_table: Vec::new() }); + } + let key_id_len = payload[a] as usize; + a += 1 + key_id_len; + + let mut segment_table = Vec::new(); + if a + 2 <= payload.len() { + let seg_count = u16::from_be_bytes([payload[a], payload[a + 1]]) as usize; + a += 2; + + for _ in 0..seg_count { + if a + 8 > payload.len() { + break; + } + let byte_len = u32::from_be_bytes([payload[a], payload[a + 1], payload[a + 2], payload[a + 3]]); + a += 4; + let sample_count = u32::from_be_bytes([payload[a], payload[a + 1], payload[a + 2], payload[a + 3]]); + a += 4; + segment_table.push(SegmentTableEntry { byte_len, sample_count }); + } + } + + tracing::debug!( + "Init UUID: {} segments dans la table, header FLAC {} octets", + segment_table.len(), + flac_header.len() + ); + + Ok(InitInfo { flac_header, segment_table }) +} + +fn parse_segment_uuid_payload( + data: &[u8], + uuid_box_start: usize, + mdat_end: usize, +) -> Result { + // Layout après box header (8) + UUID (16) = offset 24 depuis uuid_box_start: + // [4B version/padding] + // [4B data_offset_raw] — offset depuis uuid_box_start vers les données audio + // [1B iv_size] + // [3B frame_count (24-bit BE)] + // Par frame: [4B size][2B skip][2B flags][iv_size bytes IV] + + let base = uuid_box_start + 24; + if base + 12 > data.len() { + return Err(CmafError::ParseError( + "payload segment UUID trop court pour le header".into(), + )); + } + + let mut a = base + 4; // skip 4-byte version/padding + + let data_offset_raw = u32::from_be_bytes([data[a], data[a + 1], data[a + 2], data[a + 3]]); + let data_offset = uuid_box_start + data_offset_raw as usize; + a += 4; + + let iv_size = data[a] as usize; + a += 1; + + let frame_count = + ((data[a] as usize) << 16) | ((data[a + 1] as usize) << 8) | (data[a + 2] as usize); + a += 3; + + let entry_size = 4 + 2 + 2 + iv_size; + if a + frame_count * entry_size > data.len() { + return Err(CmafError::ParseError(format!( + "segment UUID: données insuffisantes pour {frame_count} entrées de {entry_size} octets" + ))); + } + + let mut entries = Vec::with_capacity(frame_count); + for _ in 0..frame_count { + let size = u32::from_be_bytes([data[a], data[a + 1], data[a + 2], data[a + 3]]); + a += 4; + a += 2; // 2 octets inconnus + let flags = u16::from_be_bytes([data[a], data[a + 1]]); + a += 2; + + let mut iv = [0u8; 8]; + let copy_len = iv_size.min(8); + iv[..copy_len].copy_from_slice(&data[a..a + copy_len]); + a += iv_size; + + entries.push(FrameEntry { size, flags, iv }); + } + + Ok(SegmentCrypto { data_offset, mdat_end, entries }) +} + +fn read_box_size(data: &[u8], pos: usize) -> usize { + if pos + 8 > data.len() { + return 0; + } + let s = u32::from_be_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]); + match s { + 0 => data.len() - pos, + 1..=7 => 0, + s => s as usize, + } +} diff --git a/pmoqobuz/src/lib.rs b/pmoqobuz/src/lib.rs index ce3b7f5f..e5d82d26 100644 --- a/pmoqobuz/src/lib.rs +++ b/pmoqobuz/src/lib.rs @@ -208,6 +208,7 @@ pub mod api; pub mod cache; +pub mod cmaf; pub mod client; pub mod config_ext; pub mod didl; @@ -216,6 +217,7 @@ pub mod disk_cache; pub mod error; mod lazy_provider; pub mod models; +pub mod retry; pub mod source; // Extension pmoserver (feature-gated) diff --git a/pmoqobuz/src/models.rs b/pmoqobuz/src/models.rs index 50e81400..71932c57 100644 --- a/pmoqobuz/src/models.rs +++ b/pmoqobuz/src/models.rs @@ -208,6 +208,31 @@ pub struct StreamInfo { pub expires_at: DateTime, } +/// Informations pour le streaming CMAF d'un track. +/// +/// Retourné par `QobuzClient::get_cmaf_stream_info`. +/// L'appelant peut ensuite appeler `pmoqobuz::cmaf::download_full` pour +/// obtenir les bytes FLAC déchiffrés, ou streamer les segments manuellement. +#[derive(Debug)] +pub struct CmafStreamInfo { + /// Modèle d'URL des segments, avec le placeholder `$SEGMENT$`. + pub url_template: String, + /// Nombre de segments audio (hors segment init s=0). + pub n_segments: u8, + /// Clé AES-128 déchiffrée pour décoder les frames FLAC. + pub content_key: [u8; 16], + /// Header FLAC extrait du segment init (à placer en tête du flux décodé). + pub flac_header: Vec, + /// Table des segments avec taille et compte d'échantillons. + pub segment_table: Vec, + /// Format ID Qobuz. + pub format_id: u32, + /// Fréquence d'échantillonnage (Hz). + pub sampling_rate: Option, + /// Profondeur de bits. + pub bit_depth: Option, +} + /// Format audio demandé pour le streaming #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[repr(u8)] diff --git a/pmoqobuz/src/retry.rs b/pmoqobuz/src/retry.rs new file mode 100644 index 00000000..660a798d --- /dev/null +++ b/pmoqobuz/src/retry.rs @@ -0,0 +1,194 @@ +//! Helper de retry avec backoff exponentiel pour les fetches réseau transitoires. +//! +//! Un blip réseau transitoire (5xx, timeout, connexion reset, 429) sur le +//! `file/url` du prochain track ou un segment CMAF ne doit pas être fatal. +//! Ce module retente les échecs *transitoires* avec backoff exponentiel +//! et laisse les échecs *terminaux* (404 "disparu définitivement", erreurs auth) +//! se propager immédiatement. + +use std::future::Future; +use std::time::Duration; + +/// Nombre de tentatives : 1 initiale + 2 retrys. +pub const DEFAULT_MAX_ATTEMPTS: u32 = 3; + +/// Erreur de fetch taguée selon son caractère retryable. +#[derive(Debug)] +pub enum FetchError { + /// Vaut la peine de retenter : erreur réseau/timeout/connect/body, 5xx, ou 429. + Transient(String), + /// Ne vaut pas la peine de retenter : 4xx (sauf 429), ou échec définitif. + Terminal(String), +} + +impl FetchError { + pub fn is_transient(&self) -> bool { + matches!(self, FetchError::Transient(_)) + } +} + +impl std::fmt::Display for FetchError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + FetchError::Transient(s) | FetchError::Terminal(s) => write!(f, "{}", s), + } + } +} + +/// Vrai pour les erreurs reqwest qui valent la peine d'être retentées. +pub fn reqwest_is_transient(e: &reqwest::Error) -> bool { + e.is_timeout() || e.is_connect() || e.is_request() || e.is_body() +} + +/// Classifie une erreur reqwest en `FetchError`. +/// Toutes les erreurs transport reqwest sont traitées comme transitoires. +pub fn classify_reqwest(e: &reqwest::Error, context: &str) -> FetchError { + FetchError::Transient(format!("{}: {}", context, e)) +} + +/// Classifie un status HTTP non-succès en `FetchError`. +/// 5xx et 429 → transitoire ; tout le reste (404, 403, ...) → terminal. +pub fn classify_status(status: reqwest::StatusCode, context: &str) -> FetchError { + let msg = format!("{}: HTTP {}", context, status); + if status.is_server_error() || status == reqwest::StatusCode::TOO_MANY_REQUESTS { + FetchError::Transient(msg) + } else { + FetchError::Terminal(msg) + } +} + +/// Backoff exponentiel avec jitter pour la N-ième tentative (base 1) : +/// ~250 ms, ~500 ms, ~1 s, plafonné à 2 s, plus jusqu'à +25% de jitter. +/// Le jitter est dérivé de l'horloge pour éviter une dépendance à `rand`. +fn backoff_delay(attempt: u32) -> Duration { + let exp = attempt.saturating_sub(1).min(3); + let base_ms = 250u64.saturating_mul(1u64 << exp).min(2000); + let jitter_span = base_ms / 4; + let jitter = if jitter_span > 0 { + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.subsec_nanos() as u64) + .unwrap_or(0); + nanos % (jitter_span + 1) + } else { + 0 + }; + Duration::from_millis(base_ms + jitter) +} + +/// Exécute `op` (qui prend le numéro de tentative base 1) et retente +/// tant qu'il retourne une erreur transitoire, avec backoff entre les tentatives. +/// Les erreurs terminales et la dernière tentative retournent immédiatement. +pub async fn retry_transient( + max_attempts: u32, + log_tag: &str, + is_transient: impl Fn(&E) -> bool, + mut op: F, +) -> std::result::Result +where + F: FnMut(u32) -> Fut, + Fut: Future>, + E: std::fmt::Display, +{ + let mut attempt = 1; + loop { + match op(attempt).await { + Ok(value) => return Ok(value), + Err(err) => { + if attempt >= max_attempts || !is_transient(&err) { + return Err(err); + } + let delay = backoff_delay(attempt); + tracing::warn!( + "[{}] erreur transitoire tentative {}/{}: {} — retry dans {}ms", + log_tag, + attempt, + max_attempts, + err, + delay.as_millis() + ); + tokio::time::sleep(delay).await; + attempt += 1; + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicU32, Ordering}; + use std::sync::Arc; + + #[tokio::test] + async fn reussit_au_premier_essai() { + let calls = Arc::new(AtomicU32::new(0)); + let c = calls.clone(); + let r: std::result::Result = + retry_transient(3, "test", FetchError::is_transient, |_| { + let c = c.clone(); + async move { + c.fetch_add(1, Ordering::Relaxed); + Ok(42) + } + }) + .await; + assert_eq!(r.unwrap(), 42); + assert_eq!(calls.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn retente_transitoire_puis_reussit() { + let calls = Arc::new(AtomicU32::new(0)); + let c = calls.clone(); + let r: std::result::Result = + retry_transient(3, "test", FetchError::is_transient, |attempt| { + let c = c.clone(); + async move { + c.fetch_add(1, Ordering::Relaxed); + if attempt < 3 { + Err(FetchError::Transient("503".into())) + } else { + Ok(7) + } + } + }) + .await; + assert_eq!(r.unwrap(), 7); + assert_eq!(calls.load(Ordering::Relaxed), 3); + } + + #[tokio::test] + async fn terminal_ne_retente_pas() { + let calls = Arc::new(AtomicU32::new(0)); + let c = calls.clone(); + let r: std::result::Result = + retry_transient(3, "test", FetchError::is_transient, |_| { + let c = c.clone(); + async move { + c.fetch_add(1, Ordering::Relaxed); + Err(FetchError::Terminal("404".into())) + } + }) + .await; + assert!(r.is_err()); + assert_eq!(calls.load(Ordering::Relaxed), 1); + } + + #[tokio::test] + async fn abandonne_apres_max_tentatives() { + let calls = Arc::new(AtomicU32::new(0)); + let c = calls.clone(); + let r: std::result::Result = + retry_transient(3, "test", FetchError::is_transient, |_| { + let c = c.clone(); + async move { + c.fetch_add(1, Ordering::Relaxed); + Err(FetchError::Transient("timeout".into())) + } + }) + .await; + assert!(r.is_err()); + assert_eq!(calls.load(Ordering::Relaxed), 3); + } +}