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.
This commit is contained in:
2026-06-11 08:46:11 +02:00
parent 7c10b22204
commit a97bfe6148
12 changed files with 1395 additions and 1 deletions

45
Cargo.lock generated
View File

@@ -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",

View File

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

261
pmoqobuz/src/api/cmaf.rs Normal file
View File

@@ -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<String>,
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<String>,
/// Clé de contenu enveloppée, format `"qbz-1.wrapped_b64url.iv_b64url"`.
#[serde(default)]
pub key: Option<String>,
/// Nombre de segments audio (hors segment init).
pub n_segments: u8,
#[serde(default)]
pub format_id: Option<u32>,
#[serde(default)]
pub mime_type: Option<String>,
#[serde(default)]
pub sampling_rate: Option<u32>,
/// Profondeur de bits (champ v1 de l'API).
#[serde(default)]
pub bits_depth: Option<u32>,
/// Profondeur de bits (champ v2 de l'API).
#[serde(default)]
pub bit_depth: Option<u32>,
}
impl CmafFileUrlResponse {
/// Retourne la profondeur de bits en préférant `bits_depth` puis `bit_depth`.
pub fn resolved_bit_depth(&self) -> Option<u32> {
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, &timestamp.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, &timestamp.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<Option<CmafSession>>,
}
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", &timestamp.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<CmafFileUrlResponse> {
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", &timestamp.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)
}

View File

@@ -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<Option<String>>,
/// 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<crate::api::cmaf::CmafFileUrlResponse> {
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<T: DeserializeOwned>(
&self,

View File

@@ -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<crate::models::CmafStreamInfo> {
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<Vec<u8>> {
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)]

154
pmoqobuz/src/cmaf/crypto.rs Normal file
View File

@@ -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<aes::Aes128>;
type Aes128Ctr = ctr::Ctr128BE<aes::Aes128>;
/// Seed publique extraite du bundle web Qobuz. Valeur IKM pour HKDF.
pub const CMAF_SEED: &str = "abb21364945c0583309667d13ca3d93a";
fn hex_decode(hex: &str) -> Vec<u8> {
(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::<Sha256>::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::<aes::cipher::block_padding::Pkcs7>(&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());
}
}

View File

@@ -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),
}

278
pmoqobuz/src/cmaf/mod.rs Normal file
View File

@@ -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<dyn Fn(CmafProgressUpdate) + Send + Sync>;
/// 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<u8>,
pub segment_table: Vec<SegmentTableEntry>,
pub format_id: u32,
pub sampling_rate: Option<u32>,
pub bit_depth: Option<u32>,
}
/// Construit un client reqwest dédié aux fetches CDN Akamai.
fn build_cdn_client() -> Result<reqwest::Client> {
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<Vec<u8>, 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<CmafProgressCallback>,
) -> Result<Vec<Vec<u8>>> {
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<u8>), String>((seg_idx, seg_data))
}));
}
let mut segments: Vec<(u8, Vec<u8>)> = 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<u8>],
content_key: &[u8; 16],
output: &mut Vec<u8>,
) -> 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<u32>,
bit_depth: Option<u32>,
) -> Result<CmafStreamingInfo> {
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<u32>,
bit_depth: Option<u32>,
on_progress: Option<CmafProgressCallback>,
) -> Result<Vec<u8>> {
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::<usize>();
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)
}

253
pmoqobuz/src/cmaf/parser.rs Normal file
View File

@@ -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<u8>,
/// Tailles par segment (indices 0..n_segments-1 correspondent aux segments 1..n_segments).
pub segment_table: Vec<SegmentTableEntry>,
}
/// 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<FrameEntry>,
}
/// 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<InitInfo, CmafError> {
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<SegmentCrypto, CmafError> {
let mut uuid_box_start: Option<usize> = 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<InitInfo, CmafError> {
// 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<SegmentCrypto, CmafError> {
// 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,
}
}

View File

@@ -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)

View File

@@ -208,6 +208,31 @@ pub struct StreamInfo {
pub expires_at: DateTime<Utc>,
}
/// 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<u8>,
/// Table des segments avec taille et compte d'échantillons.
pub segment_table: Vec<crate::cmaf::SegmentTableEntry>,
/// Format ID Qobuz.
pub format_id: u32,
/// Fréquence d'échantillonnage (Hz).
pub sampling_rate: Option<u32>,
/// Profondeur de bits.
pub bit_depth: Option<u32>,
}
/// Format audio demandé pour le streaming
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[repr(u8)]

194
pmoqobuz/src/retry.rs Normal file
View File

@@ -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<F, Fut, T, E>(
max_attempts: u32,
log_tag: &str,
is_transient: impl Fn(&E) -> bool,
mut op: F,
) -> std::result::Result<T, E>
where
F: FnMut(u32) -> Fut,
Fut: Future<Output = std::result::Result<T, E>>,
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<u32, FetchError> =
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<u32, FetchError> =
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<u32, FetchError> =
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<u32, FetchError> =
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);
}
}