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.
//! 直链验证(移植自 misc.py `AudioLinkTester.test`):
//! HEAD 推断格式与大小,失败时 GET 前 8KB 按文件头嗅探。

use futures::StreamExt;
use reqwest::header::{
    CONTENT_DISPOSITION, CONTENT_LENGTH, CONTENT_RANGE, CONTENT_TYPE, HeaderMap,
};
use serde::Serialize;
use tracing::debug;

const SNIFF_BYTES: usize = 8192;

/// 无损格式(`--lossless` 判定用)。
pub const LOSSLESS_EXTS: &[&str] = &["flac", "wav", "alac", "ape"];

/// 认可的音频扩展名(取自 Python 版 VALID_AUDIO_EXTS 的常见子集,含加密格式)。
const VALID_AUDIO_EXTS: &[&str] = &[
    "aac", "ac3", "aif", "aifc", "aiff", "alac", "amr", "ape", "au", "caf", "dff", "dsf", "dts",
    "ec3", "flac", "m3u8", "m4a", "m4b", "m4p", "m4r", "m4s", "mid", "midi", "mka", "mp1", "mp2",
    "mp3", "mpa", "mpc", "oga", "ogg", "opus", "ra", "spx", "tak", "tta", "wav", "wave", "weba",
    "wma", "wv", "mflac", "mgg", "qmcflac", "qmc0", "qmc3", "qmcogg", "tkm", "kgm", "kwm", "ncm",
];

#[derive(Debug, Clone, Serialize)]
pub struct ProbeResult {
    /// 跟随重定向后的最终 URL
    pub url: String,
    /// 推断出的扩展名(不含点)
    pub ext: String,
    /// 文件大小(字节),未知为 0
    pub size: u64,
}

impl ProbeResult {
    pub fn is_lossless(&self) -> bool {
        LOSSLESS_EXTS.contains(&self.ext.as_str())
    }
}

pub fn normalize_ext(ext: &str) -> Option<String> {
    let e = ext.trim().trim_start_matches('.').to_ascii_lowercase();
    if e.is_empty() {
        return None;
    }
    let e = match e.as_str() {
        "mpeg" | "mpga" | "x-mp3" | "x-mpeg" => "mp3",
        "wave" | "x-wav" => "wav",
        "oga" | "x-ogg" => "ogg",
        "x-flac" => "flac",
        "x-aac" => "aac",
        "x-m4a" | "mp4a" => "m4a",
        other => other,
    };
    Some(e.to_string())
}

pub fn is_valid_audio_ext(ext: &str) -> bool {
    normalize_ext(ext).is_some_and(|e| VALID_AUDIO_EXTS.contains(&e.as_str()))
}

/// URL 路径最后一段的扩展名。
pub fn ext_from_url(url: &str) -> Option<String> {
    let path = url.split(['?', '#']).next()?;
    let path = path.split_once("://").map_or(path, |(_, rest)| rest);
    let name = path.rsplit('/').next()?;
    if !path.contains('/') {
        return None; // 只有主机名
    }
    let (_, ext) = name.rsplit_once('.')?;
    (!ext.is_empty()).then(|| ext.to_ascii_lowercase())
}

fn ext_from_content_disposition(cd: &str) -> Option<String> {
    cd.split(';').map(str::trim).find_map(|part| {
        let lower = part.to_ascii_lowercase();
        let name = if lower.starts_with("filename*=") {
            let v = &part["filename*=".len()..];
            v.split_once("''").map_or(v, |(_, n)| n)
        } else if lower.starts_with("filename=") {
            &part["filename=".len()..]
        } else {
            return None;
        };
        let name = name.trim_matches('"');
        name.rsplit_once('.').map(|(_, e)| e.to_ascii_lowercase())
    })
}

fn ext_from_mime(ctype: &str) -> Option<&'static str> {
    let ct = ctype.split(';').next()?.trim().to_ascii_lowercase();
    Some(match ct.as_str() {
        "audio/mpeg" | "audio/mp3" => "mp3",
        "audio/wav" | "audio/wave" | "audio/x-wav" => "wav",
        "audio/flac" | "audio/x-flac" | "application/flac" | "application/x-flac" => "flac",
        "audio/aac" | "audio/x-aac" => "aac",
        "audio/ogg" | "audio/x-ogg" | "application/ogg" => "ogg",
        "audio/opus" => "opus",
        "audio/mp4" | "audio/x-m4a" | "audio/x-m4p" | "video/mp4" => "m4a",
        "application/x-mpegurl" | "application/vnd.apple.mpegurl" => "m3u8",
        _ => return None,
    })
}

/// 按文件头嗅探音频格式。
pub fn sniff(bytes: &[u8]) -> Option<&'static str> {
    let b = bytes;
    if b.starts_with(b"fLaC") {
        return Some("flac");
    }
    if b.starts_with(b"OggS") {
        return Some("ogg");
    }
    if b.starts_with(b"ID3") {
        // ID3 后可能跟 FLAC(少见)或 MP3;跳过标签再看
        if b.len() >= 10 {
            let sz = ((b[6] as usize & 0x7f) << 21)
                | ((b[7] as usize & 0x7f) << 14)
                | ((b[8] as usize & 0x7f) << 7)
                | (b[9] as usize & 0x7f);
            if let Some(rest) = b.get(10 + sz..)
                && rest.starts_with(b"fLaC")
            {
                return Some("flac");
            }
        }
        return Some("mp3");
    }
    if b.len() >= 12 && &b[0..4] == b"RIFF" && &b[8..12] == b"WAVE" {
        return Some("wav");
    }
    if b.len() >= 8 && &b[4..8] == b"ftyp" {
        return Some("m4a");
    }
    if b.starts_with(b"MAC ") {
        return Some("ape");
    }
    if b.starts_with(b"wvpk") {
        return Some("wv");
    }
    if b.len() >= 2 && b[0] == 0xFF {
        // ADTS AAC:0xFFF1 / 0xFFF9(layer 位为 00)
        if b[1] & 0xF6 == 0xF0 {
            return Some("aac");
        }
        // MPEG 音频帧同步:11 位 1
        if b[1] & 0xE0 == 0xE0 {
            return Some("mp3");
        }
    }
    None
}

/// 用 URL / 响应头推断扩展名(不含字节嗅探)。
fn infer_ext(original: &str, final_url: &str, headers: &HeaderMap) -> Option<String> {
    let header = |k| headers.get(k).and_then(|v| v.to_str().ok());
    let candidates = [
        ext_from_url(original),
        ext_from_url(final_url),
        header(CONTENT_DISPOSITION).and_then(ext_from_content_disposition),
        header(CONTENT_TYPE)
            .and_then(ext_from_mime)
            .map(str::to_string),
    ];
    candidates
        .into_iter()
        .flatten()
        .filter_map(|e| normalize_ext(&e))
        .find(|e| is_valid_audio_ext(e))
}

fn size_from_headers(headers: &HeaderMap) -> Option<u64> {
    let h = |k| headers.get(k).and_then(|v| v.to_str().ok());
    if let Some(n) = h(CONTENT_LENGTH).and_then(|v| v.trim().parse().ok()) {
        return Some(n);
    }
    h(CONTENT_RANGE)
        .and_then(|v| v.rsplit_once('/'))
        .and_then(|(_, total)| total.trim().parse().ok())
}

/// 验证直链;无效(非 2xx 或推断不出音频格式)返回 None。
pub async fn probe(client: &reqwest::Client, url: &str) -> Option<ProbeResult> {
    // 1) HEAD
    match crate::net::send(client.head(url)).await {
        Ok(resp) if resp.status().is_success() => {
            let final_url = resp.url().to_string();
            let size = size_from_headers(resp.headers()).unwrap_or(0);
            if let Some(ext) = infer_ext(url, &final_url, resp.headers()) {
                return Some(ProbeResult {
                    url: final_url,
                    ext,
                    size,
                });
            }
            debug!(url, "HEAD 成功但推断不出格式,改用 GET 嗅探");
        }
        Ok(resp) => {
            debug!(url, status = %resp.status(), "HEAD 非 2xx");
            return None;
        }
        Err(e) => debug!(url, "HEAD 出错,改用 GET: {e}"),
    }
    // 2) GET 前 8KB
    let resp = match crate::net::send(client.get(url)).await {
        Ok(r) if r.status().is_success() => r,
        Ok(r) => {
            debug!(url, status = %r.status(), "GET 非 2xx");
            return None;
        }
        Err(e) => {
            debug!(url, "GET 出错: {e}");
            return None;
        }
    };
    let final_url = resp.url().to_string();
    let headers = resp.headers().clone();
    let size = size_from_headers(&headers).unwrap_or(0);
    if let Some(ext) = infer_ext(url, &final_url, &headers) {
        return Some(ProbeResult {
            url: final_url,
            ext,
            size,
        });
    }
    let mut sample = Vec::with_capacity(SNIFF_BYTES);
    let mut stream = resp.bytes_stream();
    while sample.len() < SNIFF_BYTES {
        match stream.next().await {
            Some(Ok(chunk)) => {
                let take = (SNIFF_BYTES - sample.len()).min(chunk.len());
                sample.extend_from_slice(&chunk[..take]);
            }
            Some(Err(e)) => {
                debug!(url, "读取样本出错: {e}");
                break;
            }
            None => break,
        }
    }
    let ext = sniff(&sample)?;
    Some(ProbeResult {
        url: final_url,
        ext: ext.to_string(),
        size,
    })
}

#[cfg(test)]
mod tests {
    use super::*;
    use reqwest::header::HeaderValue;

    #[test]
    fn url_ext() {
        assert_eq!(
            ext_from_url("http://ws.stream.qqmusic.qq.com/F0000024jrso28p8VA.flac?guid=0&vkey=AB")
                .as_deref(),
            Some("flac")
        );
        assert_eq!(
            ext_from_url("https://a.com/x/y.MP3#frag").as_deref(),
            Some("mp3")
        );
        assert_eq!(ext_from_url("https://a.com/x/noext"), None);
        assert_eq!(ext_from_url("https://a.com"), None);
        assert_eq!(
            ext_from_url("https://a.com/x.php?u=b.flac").as_deref(),
            Some("php")
        );
    }

    #[test]
    fn ext_normalize_and_valid() {
        assert_eq!(normalize_ext(".MPEG").as_deref(), Some("mp3"));
        assert_eq!(normalize_ext("x-ogg").as_deref(), Some("ogg"));
        assert!(is_valid_audio_ext("flac"));
        assert!(is_valid_audio_ext("mflac"));
        assert!(!is_valid_audio_ext("php"));
        assert!(!is_valid_audio_ext("html"));
    }

    #[test]
    fn infer_prefers_url_over_mime() {
        // 实测 QQ 流媒体对 .flac 也返回 audio/x-ogg,URL 后缀优先
        let mut h = HeaderMap::new();
        h.insert(CONTENT_TYPE, HeaderValue::from_static("audio/x-ogg"));
        assert_eq!(
            infer_ext("http://x/a.flac?v=1", "http://x/a.flac?v=1", &h).as_deref(),
            Some("flac")
        );
        assert_eq!(
            infer_ext("http://x/api.php", "http://x/api.php", &h).as_deref(),
            Some("ogg")
        );
        let mut h = HeaderMap::new();
        h.insert(
            CONTENT_DISPOSITION,
            HeaderValue::from_static("attachment; filename=\"a b.mp3\""),
        );
        assert_eq!(
            infer_ext("http://x/dl", "http://x/dl", &h).as_deref(),
            Some("mp3")
        );
        let mut h = HeaderMap::new();
        h.insert(
            CONTENT_TYPE,
            HeaderValue::from_static("text/html; charset=utf-8"),
        );
        assert_eq!(infer_ext("http://x/dl", "http://x/dl", &h), None);
    }

    #[test]
    fn content_disposition_rfc5987() {
        assert_eq!(
            ext_from_content_disposition("attachment; filename*=UTF-8''%E5%A4%9C.FLAC").as_deref(),
            Some("flac")
        );
    }

    #[test]
    fn size_headers() {
        let mut h = HeaderMap::new();
        h.insert(CONTENT_LENGTH, HeaderValue::from_static("26691277"));
        assert_eq!(size_from_headers(&h), Some(26691277));
        let mut h = HeaderMap::new();
        h.insert(
            CONTENT_RANGE,
            HeaderValue::from_static("bytes 0-8191/155620200"),
        );
        assert_eq!(size_from_headers(&h), Some(155620200));
        assert_eq!(size_from_headers(&HeaderMap::new()), None);
    }

    #[test]
    fn sniff_magic() {
        assert_eq!(sniff(b"fLaC\0\0\0\x22"), Some("flac"));
        assert_eq!(sniff(b"OggS\0\x02"), Some("ogg"));
        assert_eq!(sniff(b"ID3\x04\0\0\0\0\0\x00\xFF\xFB"), Some("mp3"));
        assert_eq!(sniff(b"ID3\x04\0\0\0\0\0\x00fLaC"), Some("flac"));
        assert_eq!(sniff(&[0xFF, 0xFB, 0x90, 0x00]), Some("mp3"));
        assert_eq!(sniff(&[0xFF, 0xF1, 0x50, 0x80]), Some("aac"));
        assert_eq!(sniff(b"\0\0\0\x20ftypM4A "), Some("m4a"));
        assert_eq!(sniff(b"RIFF\0\0\0\0WAVEfmt "), Some("wav"));
        assert_eq!(sniff(b"MAC \x96\x0f"), Some("ape"));
        assert_eq!(sniff(b"<html>"), None);
        assert_eq!(sniff(b""), None);
    }

    #[test]
    fn lossless() {
        let p = |e: &str| ProbeResult {
            url: String::new(),
            ext: e.into(),
            size: 0,
        };
        assert!(p("flac").is_lossless());
        assert!(!p("ogg").is_lossless());
        assert!(!p("mp3").is_lossless());
    }
}