uta 0.1.2

Command-line music search and downloader for QQ Music and NetEase Cloud Music, lossless first, shipped as a single static binary. For learning and research only; non-commercial use.
//! 接口请求的重试:传输层错误(连接失败、超时、发送中断)与 5xx / 429 最多重试 2 次。
//!
//! 业务层的"空结果"(如 vkeys 某档无 url)不在这里重试。下载有自己的重试逻辑,不走这里。

use std::time::Duration;

use reqwest::{RequestBuilder, Response, StatusCode};
use tracing::debug;

/// 总尝试次数(与 Python 版 max_retries=3 一致)。
pub const ATTEMPTS: usize = 3;

/// 第 `attempt` 次失败后(从 1 开始)的等待时间。
pub fn backoff(attempt: usize) -> Duration {
    Duration::from_millis(500 * attempt as u64)
}

pub fn retryable_status(status: StatusCode) -> bool {
    status.is_server_error() || status == StatusCode::TOO_MANY_REQUESTS
}

pub fn retryable_error(e: &reqwest::Error) -> bool {
    e.is_timeout() || e.is_connect() || e.is_request()
}

/// 发送请求,必要时重试。返回最后一次的响应(状态码由调用方检查)。
/// 请求体不可克隆(流式 body)时只发一次。
pub async fn send(rb: RequestBuilder) -> reqwest::Result<Response> {
    let mut attempt = 1;
    let mut current = rb;
    loop {
        let next = (attempt < ATTEMPTS).then(|| current.try_clone()).flatten();
        let result = current.send().await;
        let retry = match &result {
            Ok(r) => retryable_status(r.status()),
            Err(e) => retryable_error(e),
        };
        let Some(next_rb) = next.filter(|_| retry) else {
            return result;
        };
        match &result {
            Ok(r) => debug!(attempt, status = %r.status(), url = %r.url(), "服务器错误,重试"),
            Err(e) => debug!(attempt, "请求失败,重试: {e}"),
        }
        tokio::time::sleep(backoff(attempt)).await;
        attempt += 1;
        current = next_rb;
    }
}

/// URL 的主机名(去掉端口与 userinfo),不是 URL 返回 None。
pub fn host_of(url: &str) -> Option<&str> {
    let rest = url.split_once("://")?.1;
    let host = rest.split(['/', '?', '#']).next()?;
    host.rsplit_once('@')
        .map_or(host, |(_, h)| h)
        .split(':')
        .next()
}

/// 链接主机是否属于 `hosts` 之一(含子域名,忽略大小写)。
pub fn host_in(url: &str, hosts: &[&str]) -> bool {
    host_of(url).is_some_and(|h| {
        let h = h.to_ascii_lowercase();
        hosts
            .iter()
            .any(|q| h == *q || h.ends_with(&format!(".{q}")))
    })
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::sync::Arc;
    use std::sync::atomic::{AtomicUsize, Ordering};
    use tokio::io::{AsyncReadExt, AsyncWriteExt};
    use tokio::net::TcpListener;

    #[test]
    fn status_classification() {
        assert!(retryable_status(StatusCode::INTERNAL_SERVER_ERROR));
        assert!(retryable_status(StatusCode::BAD_GATEWAY));
        assert!(retryable_status(StatusCode::TOO_MANY_REQUESTS));
        assert!(!retryable_status(StatusCode::OK));
        assert!(!retryable_status(StatusCode::NOT_FOUND));
        assert!(!retryable_status(StatusCode::FORBIDDEN));
    }

    #[test]
    fn backoff_grows() {
        assert_eq!(backoff(1), Duration::from_millis(500));
        assert_eq!(backoff(2), Duration::from_millis(1000));
    }

    /// 本地服务器:前 `fail` 次连接按 `mode` 失败,之后返回 200 "ok"。
    async fn server(fail: usize, mode: &'static str) -> (String, Arc<AtomicUsize>) {
        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
        let addr = listener.local_addr().unwrap();
        let hits = Arc::new(AtomicUsize::new(0));
        let counter = hits.clone();
        tokio::spawn(async move {
            loop {
                let (mut sock, _) = listener.accept().await.unwrap();
                let n = counter.fetch_add(1, Ordering::SeqCst);
                let mut buf = [0u8; 1024];
                let _ = sock.read(&mut buf).await;
                let resp: &[u8] = if n < fail && mode == "drop" {
                    continue; // 直接断开连接
                } else if n < fail {
                    b"HTTP/1.1 503 Service Unavailable\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
                } else {
                    b"HTTP/1.1 200 OK\r\ncontent-length: 2\r\nconnection: close\r\n\r\nok"
                };
                let _ = sock.write_all(resp).await;
            }
        });
        (format!("http://{addr}/"), hits)
    }

    fn client() -> reqwest::Client {
        reqwest::Client::builder()
            .no_proxy()
            .timeout(Duration::from_secs(5))
            .build()
            .unwrap()
    }

    #[tokio::test]
    async fn retries_dropped_connection() {
        let (url, hits) = server(2, "drop").await;
        let r = send(client().get(&url)).await.unwrap();
        assert_eq!(r.text().await.unwrap(), "ok");
        assert_eq!(hits.load(Ordering::SeqCst), 3);
    }

    #[tokio::test]
    async fn retries_503_then_succeeds() {
        let (url, hits) = server(1, "503").await;
        let r = send(client().post(&url).body("x")).await.unwrap();
        assert_eq!(r.status(), StatusCode::OK);
        assert_eq!(hits.load(Ordering::SeqCst), 2);
    }

    #[tokio::test]
    async fn gives_up_after_attempts() {
        let (url, hits) = server(10, "503").await;
        let r = send(client().get(&url)).await.unwrap();
        assert_eq!(r.status(), StatusCode::SERVICE_UNAVAILABLE);
        assert_eq!(hits.load(Ordering::SeqCst), ATTEMPTS);
    }

    #[tokio::test]
    async fn no_retry_on_404() {
        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
        let addr = listener.local_addr().unwrap();
        let hits = Arc::new(AtomicUsize::new(0));
        let counter = hits.clone();
        tokio::spawn(async move {
            loop {
                let (mut sock, _) = listener.accept().await.unwrap();
                counter.fetch_add(1, Ordering::SeqCst);
                let mut buf = [0u8; 1024];
                let _ = sock.read(&mut buf).await;
                let _ = sock
                    .write_all(
                        b"HTTP/1.1 404 Not Found\r\ncontent-length: 0\r\nconnection: close\r\n\r\n",
                    )
                    .await;
            }
        });
        let r = send(client().get(format!("http://{addr}/"))).await.unwrap();
        assert_eq!(r.status(), StatusCode::NOT_FOUND);
        assert_eq!(hits.load(Ordering::SeqCst), 1);
    }
}