fix: ci lint

pull/370/head
zijiren233 5 months ago
parent 1116908ceb
commit 04be224d67
No known key found for this signature in database
GPG Key ID: 534E082AAA9B39DC

@ -29,6 +29,13 @@ const LIST_PAGE_SIZE: usize = 50;
const SHUFFLE_MAX_ITEMS: usize = 200;
const RELATED_SUBTITLE_FETCH_LIMIT: usize = 32;
fn alist_headers() -> HashMap<String, String> {
HashMap::from([(
"User-Agent".to_string(),
synctv_media_providers::error::PROVIDER_USER_AGENT.to_string(),
)])
}
fn alist_modified_to_i64(value: u64) -> i64 {
i64::try_from(value).unwrap_or(i64::MAX)
}
@ -137,6 +144,7 @@ fn subtitle_name_from_task(task: &AlistSubtitleTask, index: usize) -> String {
}
fn subtitles_from_video_preview(preview: Option<&AlistVideoPreview>) -> Vec<SubtitleTrack> {
let headers = alist_headers();
preview.map_or_else(Vec::new, |preview| {
preview
.subtitle_tasks
@ -153,7 +161,7 @@ fn subtitles_from_video_preview(preview: Option<&AlistVideoPreview>) -> Vec<Subt
},
name,
url: sub.url.clone(),
headers: HashMap::new(),
headers: headers.clone(),
format: "srt".to_string(),
}
})
@ -162,6 +170,7 @@ fn subtitles_from_video_preview(preview: Option<&AlistVideoPreview>) -> Vec<Subt
}
fn subtitles_from_related_files(related: &[AlistRelatedFile]) -> Vec<SubtitleTrack> {
let headers = alist_headers();
related
.iter()
.filter(|related| {
@ -173,7 +182,7 @@ fn subtitles_from_related_files(related: &[AlistRelatedFile]) -> Vec<SubtitleTra
language: external_subtitle_language(&related.name),
name: related.name.clone(),
url: related.raw_url.clone(),
headers: HashMap::new(),
headers: headers.clone(),
format: subtitle_format_from_name(&related.name),
})
.collect()
@ -642,7 +651,7 @@ impl AlistProvider {
token: config.token.clone(),
path,
password: config.password.clone().unwrap_or_default(),
user_agent: String::new(),
headers: alist_headers(),
};
if let Ok(subtitle_info) = client.fs_get(request).await {
@ -673,7 +682,7 @@ impl AlistProvider {
token: config.token.clone(),
path: config.path.clone(),
password: config.password.clone().unwrap_or_default(),
user_agent: String::new(),
headers: alist_headers(),
};
// Call client (trait method - implementation handles local/remote)
@ -750,6 +759,7 @@ impl AlistProvider {
};
let mut transcoded_modes = Vec::new();
let headers = alist_headers();
if let Some(preview) = video_preview {
// Add transcoding quality options
@ -780,7 +790,7 @@ impl AlistProvider {
PlaybackInfo {
urls: vec![task.url.clone()],
format: "hls".to_string(),
headers: HashMap::new(),
headers: headers.clone(),
subtitles: combined_subtitles.clone(),
expires_at: task_expires_at,
cors_proxy_required: false,
@ -837,7 +847,7 @@ impl AlistProvider {
PlaybackInfo {
urls: vec![file_info.raw_url.clone()],
format: Self::detect_format(&file_info.name),
headers: HashMap::new(),
headers: headers.clone(),
subtitles: combined_subtitles,
expires_at: direct_expires_at,
cors_proxy_required: false,
@ -1753,6 +1763,10 @@ mod tests {
#[async_trait]
impl AlistInterface for FakeAlistSubtitleClient {
async fn fs_get(&self, request: FsGetReq) -> Result<FsGetResp, AlistError> {
assert_eq!(
request.headers.get("User-Agent").map(String::as_str),
Some(synctv_media_providers::error::PROVIDER_USER_AGENT)
);
self.requested_paths
.lock()
.expect("requested_paths mutex should not be poisoned")
@ -2406,7 +2420,15 @@ mod tests {
.get("transcoded_HD")
.expect("transcoded HD mode should exist");
assert_eq!(transcoded.format, "hls");
assert_eq!(
transcoded.headers.get("User-Agent").map(String::as_str),
Some(synctv_media_providers::error::PROVIDER_USER_AGENT)
);
assert_eq!(transcoded.subtitles.len(), 3);
assert!(transcoded.subtitles.iter().all(|sub| {
sub.headers.get("User-Agent").map(String::as_str)
== Some(synctv_media_providers::error::PROVIDER_USER_AGENT)
}));
assert!(transcoded
.subtitles
.iter()
@ -2421,6 +2443,10 @@ mod tests {
.get("direct")
.expect("direct fallback should exist");
assert_eq!(direct.format, "mkv");
assert_eq!(
direct.headers.get("User-Agent").map(String::as_str),
Some(synctv_media_providers::error::PROVIDER_USER_AGENT)
);
assert_eq!(direct.subtitles.len(), 3);
assert_eq!(result.metadata["transcoding_count"], json!(2));
assert_eq!(result.metadata["video_preview_subtitle_count"], json!(1));

@ -817,6 +817,25 @@ impl super::proxy::ProviderProxy for BilibiliProvider {
// Use the shared bilibili_headers() from the parent module.
use super::bilibili_headers;
fn bilibili_subtitle_track(name: String, url: String) -> SubtitleTrack {
SubtitleTrack {
language: name.clone(),
name,
url,
headers: bilibili_headers(),
format: "json".to_string(),
}
}
fn bilibili_live_headers() -> HashMap<String, String> {
let mut headers = bilibili_headers();
headers.insert(
"Referer".to_string(),
"https://live.bilibili.com".to_string(),
);
headers
}
impl BilibiliProvider {
/// Resolve danmaku connection info from a media item's source config.
///
@ -996,13 +1015,7 @@ impl BilibiliProvider {
subtitles = subtitle_resp
.subtitles
.into_iter()
.map(|(name, url)| SubtitleTrack {
language: name.clone(),
name,
url,
headers: HashMap::new(),
format: "json".to_string(),
})
.map(|(name, url)| bilibili_subtitle_track(name, url))
.collect();
}
Err(e) => {
@ -1076,13 +1089,7 @@ impl BilibiliProvider {
subtitles = subtitle_resp
.subtitles
.into_iter()
.map(|(name, url)| SubtitleTrack {
language: name.clone(),
name,
url,
headers: HashMap::new(),
format: "json".to_string(),
})
.map(|(name, url)| bilibili_subtitle_track(name, url))
.collect();
}
Err(e) => {
@ -1155,14 +1162,7 @@ impl BilibiliProvider {
PlaybackInfo {
urls: stream.urls,
format: "hls".to_string(),
headers: {
let mut h = HashMap::new();
h.insert(
"Referer".to_string(),
"https://live.bilibili.com".to_string(),
);
h
},
headers: bilibili_live_headers(),
subtitles: Vec::new(),
expires_at: live_expires_at,
cors_proxy_required: true,
@ -1197,10 +1197,13 @@ impl BilibiliProvider {
mod tests {
use super::*;
use crate::models::UserId;
use crate::provider::{MediaProvider, ProviderContext};
use crate::provider::{MediaProvider, ProviderClientManager, ProviderContext};
use crate::repository::ProviderInstanceRepository;
use crate::service::RemoteProviderManager;
use async_trait::async_trait;
use std::sync::Arc;
use synctv_media_providers::bilibili::{BilibiliError, BilibiliInterface};
use synctv_media_providers::grpc::bilibili as proto;
fn fake_provider_instance_manager() -> Arc<RemoteProviderManager> {
let pool = sqlx::PgPool::connect_lazy("postgresql://fake").expect("lazy pool");
@ -1219,6 +1222,254 @@ mod tests {
})
}
struct MockBilibiliClient;
fn mock_not_implemented() -> BilibiliError {
BilibiliError::NotImplemented("mock bilibili method not implemented".to_string())
}
#[async_trait]
impl BilibiliInterface for MockBilibiliClient {
async fn new_qr_code(
&self,
_request: proto::Empty,
) -> Result<proto::NewQrCodeResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn login_with_qr_code(
&self,
_request: proto::LoginWithQrCodeReq,
) -> Result<proto::LoginWithQrCodeResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn new_captcha(
&self,
_request: proto::Empty,
) -> Result<proto::NewCaptchaResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn new_sms(
&self,
_request: proto::NewSmsReq,
) -> Result<proto::NewSmsResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn login_with_sms(
&self,
_request: proto::LoginWithSmsReq,
) -> Result<proto::LoginWithSmsResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn parse_video_page(
&self,
_request: proto::ParseVideoPageReq,
) -> Result<proto::VideoPageInfo, BilibiliError> {
Err(mock_not_implemented())
}
async fn get_video_url(
&self,
_request: proto::GetVideoUrlReq,
) -> Result<proto::VideoUrl, BilibiliError> {
Err(mock_not_implemented())
}
async fn get_dash_video_url(
&self,
_request: proto::GetDashVideoUrlReq,
) -> Result<proto::GetDashVideoUrlResp, BilibiliError> {
Ok(mock_dash_response("https://upos.example/video.m4s"))
}
async fn get_subtitles(
&self,
_request: proto::GetSubtitlesReq,
) -> Result<proto::GetSubtitlesResp, BilibiliError> {
Ok(proto::GetSubtitlesResp {
subtitles: HashMap::from([(
"zh-CN".to_string(),
"https://subtitle.example/zh.json".to_string(),
)]),
})
}
async fn parse_pgc_page(
&self,
_request: proto::ParsePgcPageReq,
) -> Result<proto::VideoPageInfo, BilibiliError> {
Err(mock_not_implemented())
}
async fn get_pgcurl(
&self,
_request: proto::GetPgcurlReq,
) -> Result<proto::VideoUrl, BilibiliError> {
Err(mock_not_implemented())
}
async fn get_dash_pgcurl(
&self,
_request: proto::GetDashPgcurlReq,
) -> Result<proto::GetDashPgcurlResp, BilibiliError> {
Ok(proto::GetDashPgcurlResp {
dash: mock_dash_response("https://upos.example/pgc.m4s").dash,
hevc_dash: None,
})
}
async fn user_info(
&self,
_request: proto::UserInfoReq,
) -> Result<proto::UserInfoResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn r#match(
&self,
_request: proto::MatchReq,
) -> Result<proto::MatchResp, BilibiliError> {
Err(mock_not_implemented())
}
async fn get_live_streams(
&self,
_request: proto::GetLiveStreamsReq,
) -> Result<proto::GetLiveStreamsResp, BilibiliError> {
Ok(proto::GetLiveStreamsResp {
live_streams: vec![proto::LiveStream {
quality: 10000,
urls: vec!["https://live.example/stream.m3u8".to_string()],
desc: "origin".to_string(),
}],
})
}
async fn parse_live_page(
&self,
_request: proto::ParseLivePageReq,
) -> Result<proto::VideoPageInfo, BilibiliError> {
Err(mock_not_implemented())
}
async fn get_live_danmu_info(
&self,
_request: proto::GetLiveDanmuInfoReq,
) -> Result<proto::GetLiveDanmuInfoResp, BilibiliError> {
Err(mock_not_implemented())
}
}
fn mock_dash_response(url: &str) -> proto::GetDashVideoUrlResp {
proto::GetDashVideoUrlResp {
dash: Some(proto::DashInfo {
duration: 120.0,
min_buffer_time: 1.5,
video_streams: vec![proto::VideoStream {
id: 80,
base_url: url.to_string(),
mime_type: "video/mp4".to_string(),
codecs: "avc1.640028".to_string(),
width: 1920,
height: 1080,
frame_rate: "60".to_string(),
bandwidth: 1_000_000,
start_with_sap: 1,
segment_base: None,
}],
audio_streams: Vec::new(),
}),
hevc_dash: None,
}
}
fn provider_with_mock_bilibili_client() -> BilibiliProvider {
let default_clients = ProviderClientManager::new();
let client_manager = Arc::new(ProviderClientManager::with_custom_clients(
default_clients.local_alist_client(),
Arc::new(MockBilibiliClient),
default_clients.local_emby_client(),
));
BilibiliProvider::with_client_manager(fake_provider_instance_manager(), client_manager)
}
fn assert_bilibili_cdn_headers(headers: &HashMap<String, String>, expected_referer: &str) {
assert_eq!(
headers.get("Referer"),
Some(&expected_referer.to_string()),
"Bilibili direct playback must return the required Referer"
);
assert_eq!(
headers.get("User-Agent"),
Some(&synctv_media_providers::error::PROVIDER_USER_AGENT.to_string()),
"Bilibili direct playback must return the provider User-Agent"
);
}
#[tokio::test]
async fn test_video_direct_playback_returns_stream_and_subtitle_headers() {
let provider = provider_with_mock_bilibili_client();
let result = provider
.generate_playback(
&ProviderContext::new("test").with_user_id(UserId::from(1)),
&json!({
"type": "video",
"bvid": "BV1GJ411x7gL",
"cid": 12345
}),
)
.await
.expect("mock video playback should resolve");
let dash = &result.playback_infos["dash"];
assert_bilibili_cdn_headers(&dash.headers, "https://www.bilibili.com");
assert_eq!(dash.subtitles.len(), 1);
assert_bilibili_cdn_headers(&dash.subtitles[0].headers, "https://www.bilibili.com");
}
#[tokio::test]
async fn test_pgc_direct_playback_returns_stream_and_subtitle_headers() {
let provider = provider_with_mock_bilibili_client();
let result = provider
.generate_playback(
&ProviderContext::new("test").with_user_id(UserId::from(1)),
&json!({
"type": "pgc",
"epid": 98765,
"cid": 12345
}),
)
.await
.expect("mock PGC playback should resolve");
let dash = &result.playback_infos["dash"];
assert_bilibili_cdn_headers(&dash.headers, "https://www.bilibili.com");
assert_eq!(dash.subtitles.len(), 1);
assert_bilibili_cdn_headers(&dash.subtitles[0].headers, "https://www.bilibili.com");
}
#[tokio::test]
async fn test_live_direct_playback_returns_live_headers() {
let provider = provider_with_mock_bilibili_client();
let result = provider
.generate_playback(
&ProviderContext::new("test").with_user_id(UserId::from(1)),
&json!({
"type": "live",
"room_id": 12345
}),
)
.await
.expect("mock live playback should resolve");
let playback = &result.playback_infos[&result.default_mode];
assert_bilibili_cdn_headers(&playback.headers, "https://live.bilibili.com");
}
#[test]
fn test_valid_video_config_with_bvid() {
let config = json!({

@ -246,6 +246,10 @@ fn grpc_playback_request_hints(
)
}
fn emby_auth_headers(token: &str) -> HashMap<String, String> {
HashMap::from([("X-Emby-Token".to_string(), token.to_string())])
}
/// Emby `MediaProvider`
///
/// Holds a reference to `RemoteProviderManager` to select appropriate provider instance.
@ -569,11 +573,7 @@ impl EmbyProvider {
// Auth headers for Emby: use X-Emby-Token header instead of
// embedding api_key in query strings to avoid credential exposure
// in URLs (which end up in logs, browser history, Referer headers).
let emby_auth_headers = {
let mut h = HashMap::new();
h.insert("X-Emby-Token".to_string(), config.token.clone());
h
};
let emby_auth_headers = emby_auth_headers(&config.token);
// Process media sources
for (idx, source) in playback_info.media_source_info.iter().enumerate() {
@ -593,9 +593,8 @@ impl EmbyProvider {
};
// Extract subtitles -- do NOT include api_key in the URL to avoid
// leaking the Emby token to clients. Instead, subtitle URLs are
// fetched through the server-side proxy which injects the
// X-Emby-Token header (same as video streams).
// leaking the Emby token to clients. Direct clients and the server
// proxy both use X-Emby-Token headers, same as video streams.
let subtitles: Vec<SubtitleTrack> = source
.media_stream_info
.iter()
@ -616,7 +615,7 @@ impl EmbyProvider {
language: stream.language.clone(),
name: stream.display_title.clone(),
url: subtitle_url,
headers: HashMap::new(),
headers: emby_auth_headers.clone(),
format: stream.codec.to_lowercase(),
})
})
@ -1552,8 +1551,12 @@ impl DynamicFolder for EmbyProvider {
mod tests {
use super::*;
use crate::models::UserId;
use crate::provider::ProviderClientManager;
use crate::repository::ProviderInstanceRepository;
use async_trait::async_trait;
use std::sync::Arc;
use synctv_media_providers::emby::{EmbyError, EmbyInterface};
use synctv_media_providers::grpc::emby as proto;
fn fake_provider_instance_manager() -> Arc<RemoteProviderManager> {
let pool = sqlx::PgPool::connect_lazy("postgresql://fake").expect("lazy pool");
@ -1579,6 +1582,169 @@ mod tests {
Ok(())
}
struct MockEmbyClient;
fn mock_not_implemented() -> EmbyError {
EmbyError::NotImplemented("mock emby method not implemented".to_string())
}
#[async_trait]
impl EmbyInterface for MockEmbyClient {
async fn login(&self, _request: proto::LoginReq) -> Result<proto::LoginResp, EmbyError> {
Err(mock_not_implemented())
}
async fn me(&self, _request: proto::MeReq) -> Result<proto::MeResp, EmbyError> {
Err(mock_not_implemented())
}
async fn get_items(
&self,
_request: proto::GetItemsReq,
) -> Result<proto::GetItemsResp, EmbyError> {
Err(mock_not_implemented())
}
async fn get_item(&self, _request: proto::GetItemReq) -> Result<proto::Item, EmbyError> {
Ok(proto::Item {
name: "Mock Movie".to_string(),
id: "item-1".to_string(),
r#type: "Movie".to_string(),
parent_id: String::new(),
series_name: String::new(),
series_id: String::new(),
season_name: String::new(),
season_id: String::new(),
is_folder: false,
media_source_info: Vec::new(),
collection_type: String::new(),
})
}
async fn fs_list(
&self,
_request: proto::FsListReq,
) -> Result<proto::FsListResp, EmbyError> {
Err(mock_not_implemented())
}
async fn get_system_info(
&self,
_request: proto::SystemInfoReq,
) -> Result<proto::SystemInfoResp, EmbyError> {
Err(mock_not_implemented())
}
async fn logout(&self, _request: proto::LogoutReq) -> Result<proto::Empty, EmbyError> {
Err(mock_not_implemented())
}
async fn playback_info(
&self,
_request: proto::PlaybackInfoReq,
) -> Result<proto::PlaybackInfoResp, EmbyError> {
Ok(proto::PlaybackInfoResp {
play_session_id: "play-session-1".to_string(),
media_source_info: vec![proto::MediaSourceInfo {
id: "source-1".to_string(),
name: "Main".to_string(),
path: String::new(),
container: "mp4".to_string(),
protocol: "File".to_string(),
default_subtitle_stream_index: 2,
default_audio_stream_index: 1,
media_stream_info: vec![proto::MediaStreamInfo {
codec: "srt".to_string(),
language: "eng".to_string(),
r#type: "Subtitle".to_string(),
title: "English".to_string(),
display_title: "English".to_string(),
display_language: "English".to_string(),
is_default: true,
index: 2,
protocol: "File".to_string(),
}],
direct_play_url: "/Videos/item-1/stream.mp4".to_string(),
transcoding_url: String::new(),
}],
})
}
async fn delete_active_encodings(
&self,
_request: proto::DeleteActiveEncodingsReq,
) -> Result<proto::Empty, EmbyError> {
Err(mock_not_implemented())
}
async fn report_playback_start(
&self,
_request: proto::ReportPlaybackStartReq,
) -> Result<proto::Empty, EmbyError> {
Err(mock_not_implemented())
}
async fn report_playback_stop(
&self,
_request: proto::ReportPlaybackStopReq,
) -> Result<proto::Empty, EmbyError> {
Err(mock_not_implemented())
}
async fn report_playback_progress(
&self,
_request: proto::ReportPlaybackProgressReq,
) -> Result<proto::Empty, EmbyError> {
Err(mock_not_implemented())
}
}
fn provider_with_mock_emby_client() -> EmbyProvider {
let default_clients = ProviderClientManager::new();
let client_manager = Arc::new(ProviderClientManager::with_custom_clients(
default_clients.local_alist_client(),
default_clients.local_bilibili_client(),
Arc::new(MockEmbyClient),
));
EmbyProvider::with_client_manager(fake_provider_instance_manager(), client_manager)
}
#[tokio::test]
async fn test_emby_direct_playback_returns_subtitle_auth_headers() {
let provider = provider_with_mock_emby_client();
let result = provider
.resolve_from_api(
&ResolvedEmbyConfig {
host: "https://emby.example.com".to_string(),
token: "token-123".to_string(),
user_id: "user-1".to_string(),
item_id: "item-1".to_string(),
credential_owner_id: "owner-1".to_string(),
credential_revision: "credential-1:1".to_string(),
provider_instance_name: None,
},
None,
None,
)
.await
.expect("mock Emby playback should resolve");
let playback = &result.playback_infos["Main"];
assert_eq!(
playback.headers.get("X-Emby-Token").map(String::as_str),
Some("token-123")
);
assert_eq!(playback.subtitles.len(), 1);
assert_eq!(
playback.subtitles[0]
.headers
.get("X-Emby-Token")
.map(String::as_str),
Some("token-123"),
"direct subtitle clients need the same Emby auth header as video streams"
);
}
#[test]
fn test_valid_emby_config() {
let config = json!({

@ -213,7 +213,7 @@ pub fn bilibili_headers() -> std::collections::HashMap<String, String> {
);
headers.insert(
"User-Agent".to_string(),
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36".to_string(),
synctv_media_providers::error::PROVIDER_USER_AGENT.to_string(),
);
headers
}
@ -371,6 +371,7 @@ pub fn sign_playback_urls(
user_id,
expires_at,
);
subtitle.headers.clear();
}
}
}
@ -617,7 +618,10 @@ mod tests {
language: "zh-CN".to_string(),
name: "Chinese".to_string(),
url: "https://cdn.example.com/subtitle.json".to_string(),
headers: std::collections::HashMap::new(),
headers: std::collections::HashMap::from([(
"Authorization".to_string(),
"Bearer subtitle-token".to_string(),
)]),
format: "json".to_string(),
}],
expires_at: None,
@ -659,6 +663,10 @@ mod tests {
.starts_with("/api/providers/proxy/bilibili/ver-1/subtitle%2Fdash%2F0?"),
"subtitle URLs may still use the signed proxy contract"
);
assert!(
dash.subtitles[0].headers.is_empty(),
"signed subtitle proxy should not expose upstream headers to clients"
);
}
#[test]

@ -29,7 +29,7 @@ message FsGetReq {
string path = 3;
string password = 4;
string user_agent = 5;
map<string, string> headers = 5;
}
message FsGetResp {

@ -2,10 +2,14 @@
//!
//! Pure HTTP client for Alist API, no dependency on `MediaProvider`
use std::collections::HashMap;
use std::sync::LazyLock;
use reqwest::{
header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE, ORIGIN, REFERER, USER_AGENT},
header::{
HeaderMap, HeaderName, HeaderValue, AUTHORIZATION, CONTENT_TYPE, ORIGIN, REFERER,
USER_AGENT,
},
Client,
};
use serde_json::json;
@ -50,6 +54,18 @@ fn referer_value(url: &url::Url) -> Result<HeaderValue, AlistError> {
.map_err(|e| AlistError::InvalidConfig(format!("Invalid Referer header value: {e}")))
}
fn header_value<'a>(headers: &'a HashMap<String, String>, name: &str) -> Option<&'a str> {
headers
.iter()
.find(|(key, _)| key.eq_ignore_ascii_case(name))
.map(|(_, value)| value.trim())
.filter(|value| !value.is_empty())
}
fn effective_user_agent(headers: &HashMap<String, String>) -> &str {
header_value(headers, USER_AGENT.as_str()).unwrap_or(crate::error::PROVIDER_USER_AGENT)
}
/// Shared HTTP client for all Alist requests (connection pooling).
/// SSRF-safe: uses the common DNS resolver and disables redirects.
static SHARED_CLIENT: LazyLock<Result<Client, reqwest::Error>> =
@ -126,18 +142,34 @@ impl AlistClient {
self.token.is_some()
}
/// Build request headers
fn build_headers(&self) -> Result<HeaderMap, AlistError> {
/// Build request headers.
fn build_headers(
&self,
request_headers: &HashMap<String, String>,
) -> Result<HeaderMap, AlistError> {
let parsed_host = parse_host_url(&self.host)?;
let user_agent = effective_user_agent(request_headers);
let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert(
USER_AGENT,
HeaderValue::from_static(crate::error::PROVIDER_USER_AGENT),
);
headers.insert(USER_AGENT, HeaderValue::from_str(user_agent)?);
headers.insert(ORIGIN, origin_value(&parsed_host)?);
headers.insert(REFERER, referer_value(&parsed_host)?);
for (name, value) in request_headers {
if name.eq_ignore_ascii_case(CONTENT_TYPE.as_str())
|| name.eq_ignore_ascii_case(AUTHORIZATION.as_str())
|| name.eq_ignore_ascii_case(ORIGIN.as_str())
|| name.eq_ignore_ascii_case(REFERER.as_str())
{
continue;
}
let header_name = HeaderName::from_bytes(name.as_bytes())
.map_err(|err| AlistError::InvalidHeader(err.to_string()))?;
let header_value = HeaderValue::from_str(value)?;
headers.insert(header_name, header_value);
}
if let Some(ref token) = self.token {
headers.insert(AUTHORIZATION, HeaderValue::from_str(token)?);
}
@ -181,7 +213,7 @@ impl AlistClient {
"password": password,
"otp_code": otp_code.unwrap_or(""),
});
let headers = self.build_headers()?;
let headers = self.build_headers(&HashMap::new())?;
let client = self.client.clone();
let token = with_retry(|| {
@ -228,14 +260,17 @@ impl AlistClient {
&self,
path: &str,
password: Option<&str>,
request_headers: &HashMap<String, String>,
) -> Result<HttpFsGetResp, AlistError> {
validate_path(path)?;
let user_agent = effective_user_agent(request_headers);
let url = format!("{}/api/fs/get", self.host);
let body = json!({
"path": path,
"password": password.unwrap_or(""),
"user_agent": user_agent,
});
let headers = self.build_headers()?;
let headers = self.build_headers(request_headers)?;
let client = self.client.clone();
let result = with_retry(|| {
@ -307,7 +342,7 @@ impl AlistClient {
"per_page": per_page,
"refresh": refresh,
});
let headers = self.build_headers()?;
let headers = self.build_headers(&HashMap::new())?;
let client = self.client.clone();
let result = with_retry(|| {
@ -371,7 +406,7 @@ impl AlistClient {
"method": method,
"password": password.unwrap_or(""),
});
let headers = self.build_headers()?;
let headers = self.build_headers(&HashMap::new())?;
let client = self.client.clone();
let result = with_retry(|| {
@ -412,7 +447,7 @@ impl AlistClient {
/// Requires authentication token
pub async fn me(&self) -> Result<HttpMeResp, AlistError> {
let url = format!("{}/api/me", self.host);
let headers = self.build_headers()?;
let headers = self.build_headers(&HashMap::new())?;
let client = self.client.clone();
with_retry(|| {
@ -486,7 +521,7 @@ impl AlistClient {
"per_page": per_page,
"password": password.unwrap_or(""),
});
let headers = self.build_headers()?;
let headers = self.build_headers(&HashMap::new())?;
let client = self.client.clone();
with_retry(|| {
@ -635,7 +670,7 @@ mod tests {
#[test]
fn test_build_headers_uses_origin_without_path_or_query() {
let client = AlistClient::new("https://alist.example.com/base?token=secret#frag").unwrap();
let headers = client.build_headers().unwrap();
let headers = client.build_headers(&HashMap::new()).unwrap();
assert_eq!(
headers.get(ORIGIN).and_then(|v| v.to_str().ok()),
@ -651,7 +686,7 @@ mod tests {
fn test_build_headers_rejects_userinfo_in_host() {
let client = AlistClient::new("https://user:pass@alist.example.com").unwrap();
let err = client
.build_headers()
.build_headers(&HashMap::new())
.expect_err("userinfo must not be accepted in provider host");
assert!(
err.to_string().contains("Origin header")

@ -11,7 +11,16 @@
//! # async fn example() -> Result<(), Box<dyn std::error::Error>> {
//! let mut client = AlistClient::new("https://alist.example.com")?;
//! let token = client.login("username", "password", false).await?;
//! let file_info = client.fs_get("/movies/video.mp4", None).await?;
//! let file_info = client
//! .fs_get(
//! "/movies/video.mp4",
//! None,
//! &std::collections::HashMap::from([(
//! "User-Agent".to_string(),
//! synctv_media_providers::error::PROVIDER_USER_AGENT.to_string(),
//! )]),
//! )
//! .await?;
//! # Ok(())
//! # }
//! ```

@ -76,7 +76,9 @@ impl AlistInterface for AlistService {
} else {
Some(request.password.as_str())
};
let http_resp = client.fs_get(&request.path, password).await?;
let http_resp = client
.fs_get(&request.path, password, &request.headers)
.await?;
Ok(http_resp.into())
}

@ -29,7 +29,7 @@ pub struct LoginResp {
pub token: ::prost::alloc::string::String,
}
#[derive(serde::Serialize, serde::Deserialize)]
#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)]
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct FsGetReq {
#[prost(string, tag = "1")]
pub host: ::prost::alloc::string::String,
@ -39,8 +39,11 @@ pub struct FsGetReq {
pub path: ::prost::alloc::string::String,
#[prost(string, tag = "4")]
pub password: ::prost::alloc::string::String,
#[prost(string, tag = "5")]
pub user_agent: ::prost::alloc::string::String,
#[prost(map = "string, string", tag = "5")]
pub headers: ::std::collections::HashMap<
::prost::alloc::string::String,
::prost::alloc::string::String,
>,
}
#[derive(serde::Serialize, serde::Deserialize)]
#[derive(Clone, PartialEq, ::prost::Message)]

@ -9,10 +9,17 @@
//! from the API responses as-is.
#![allow(clippy::unwrap_used)]
use std::collections::HashMap;
use synctv_media_providers::error::PROVIDER_USER_AGENT;
use synctv_media_providers::AlistClient;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn provider_headers() -> HashMap<String, String> {
HashMap::from([("User-Agent".to_string(), PROVIDER_USER_AGENT.to_string())])
}
#[tokio::test]
async fn test_alist_fs_get_preserves_urls() {
let server = MockServer::start().await;
@ -39,7 +46,10 @@ async fn test_alist_fs_get_preserves_urls() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let resp = client.fs_get("/movies/video.mp4", None).await.unwrap();
let resp = client
.fs_get("/movies/video.mp4", None, &provider_headers())
.await
.unwrap();
assert_eq!(resp.raw_url, "https://cdn.example.com/video.mp4");
assert_eq!(resp.thumb, "https://cdn.example.com/thumb.jpg");
@ -71,7 +81,10 @@ async fn test_alist_fs_get_empty_urls_preserved() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let resp = client.fs_get("/movies", None).await.unwrap();
let resp = client
.fs_get("/movies", None, &provider_headers())
.await
.unwrap();
assert_eq!(resp.raw_url, "");
assert_eq!(resp.thumb, "");

@ -3,8 +3,11 @@
//! Tests for path validation, client creation, and HTTP API interactions using wiremock.
#![allow(clippy::unwrap_used)]
use std::collections::HashMap;
use serde_json::json;
use synctv_media_providers::alist::{AlistInterface, AlistService};
use synctv_media_providers::error::PROVIDER_USER_AGENT;
use synctv_media_providers::grpc::alist::{
alist_client::AlistClient as GrpcAlistClient, alist_server::AlistServer, login_req, FsGetReq,
FsListReq, FsOtherReq, FsSearchReq, LoginReq, MeReq,
@ -21,11 +24,17 @@ use wiremock::{Mock, MockServer, ResponseTemplate};
// validate_path is a private function, but we test it indirectly through the public API
// (fs_get, fs_list, etc.) which call validate_path internally.
fn provider_headers() -> HashMap<String, String> {
HashMap::from([("User-Agent".to_string(), PROVIDER_USER_AGENT.to_string())])
}
#[tokio::test]
async fn test_validate_path_url_encoded_dotdot_rejected() {
// "%2e%2e" decodes to ".." which should be rejected
let client = AlistClient::with_token("https://alist.example.com", "token123").unwrap();
let result = client.fs_get("/movies/%2e%2e/etc/passwd", None).await;
let result = client
.fs_get("/movies/%2e%2e/etc/passwd", None, &provider_headers())
.await;
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
@ -38,7 +47,9 @@ async fn test_validate_path_url_encoded_dotdot_rejected() {
async fn test_validate_path_url_encoded_slash_dotdot_rejected() {
// "%2F%2e%2e" decodes to "/.." which should be rejected
let client = AlistClient::with_token("https://alist.example.com", "token123").unwrap();
let result = client.fs_get("/movies%2F%2e%2e/secret", None).await;
let result = client
.fs_get("/movies%2F%2e%2e/secret", None, &provider_headers())
.await;
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
@ -73,7 +84,9 @@ async fn test_validate_path_normal_paths_accepted() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let result = client.fs_get("/movies/video.mp4", None).await;
let result = client
.fs_get("/movies/video.mp4", None, &provider_headers())
.await;
assert!(result.is_ok(), "Normal path should be accepted: {result:?}");
}
@ -102,7 +115,9 @@ async fn test_validate_path_dotfiles_accepted() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let result = client.fs_get("/movies/.hidden", None).await;
let result = client
.fs_get("/movies/.hidden", None, &provider_headers())
.await;
assert!(
result.is_ok(),
"Dotfile path should be accepted: {result:?}"
@ -180,7 +195,10 @@ async fn test_alist_client_fs_get_success() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let resp = client.fs_get("/movies/movie.mkv", None).await.unwrap();
let resp = client
.fs_get("/movies/movie.mkv", None, &provider_headers())
.await
.unwrap();
assert_eq!(resp.name, "movie.mkv");
assert_eq!(resp.size, 2_000_000_000);
assert!(!resp.is_dir);
@ -193,12 +211,6 @@ async fn test_alist_client_fs_get_sends_auth_headers_and_password() {
Mock::given(method("POST"))
.and(path("/api/fs/get"))
.and(header("authorization", "token123"))
.and(header("origin", server.uri()))
.and(body_json(json!({
"path": "/protected/movie.mkv",
"password": "dir-password"
})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"code": 200,
"message": "success",
@ -217,12 +229,53 @@ async fn test_alist_client_fs_get_sends_auth_headers_and_password() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let mut headers = provider_headers();
headers.insert("X-Alist-Test".to_string(), "custom-header".to_string());
let resp = client
.fs_get("/protected/movie.mkv", Some("dir-password"))
.fs_get("/protected/movie.mkv", Some("dir-password"), &headers)
.await
.unwrap();
assert_eq!(resp.name, "movie.mkv");
assert!(resp.related.is_empty());
let requests = server
.received_requests()
.await
.expect("wiremock should record requests");
assert_eq!(requests.len(), 1);
let request = &requests[0];
assert_eq!(
request
.headers
.get("authorization")
.and_then(|value| value.to_str().ok()),
Some("token123")
);
assert_eq!(
request
.headers
.get("origin")
.and_then(|value| value.to_str().ok()),
Some(server.uri().as_str())
);
assert_eq!(
request
.headers
.get("user-agent")
.and_then(|value| value.to_str().ok()),
Some(PROVIDER_USER_AGENT)
);
assert_eq!(
request
.headers
.get("x-alist-test")
.and_then(|value| value.to_str().ok()),
Some("custom-header")
);
let body: serde_json::Value = request.body_json().expect("fs/get body should be JSON");
assert_eq!(body["path"], json!("/protected/movie.mkv"));
assert_eq!(body["password"], json!("dir-password"));
assert_eq!(body["user_agent"], json!(PROVIDER_USER_AGENT));
}
#[tokio::test]
@ -373,7 +426,9 @@ async fn test_alist_client_5xx_retries() {
.await;
let client = AlistClient::with_token(server.uri(), "token123").unwrap();
let result = client.fs_get("/movies/video.mp4", None).await;
let result = client
.fs_get("/movies/video.mp4", None, &provider_headers())
.await;
assert!(result.is_ok(), "Should succeed after retry: {result:?}");
assert_eq!(result.unwrap().name, "video.mp4");
}
@ -1016,7 +1071,10 @@ async fn test_openlist_container_exercises_real_alist_client_api() {
assert_eq!(empty.total, 0);
assert!(empty.content.is_empty());
let file = client.fs_get("/local/video.mp4", None).await.unwrap();
let file = client
.fs_get("/local/video.mp4", None, &provider_headers())
.await
.unwrap();
assert_eq!(file.name, "video.mp4");
assert_eq!(file.size, 15);
assert_eq!(file.provider, "Local");
@ -1121,7 +1179,7 @@ async fn test_openlist_container_exercises_real_alist_grpc_service() {
token: token.clone(),
path: "/local/video.mp4".to_string(),
password: String::new(),
user_agent: String::new(),
headers: HashMap::new(),
})
.await
.unwrap()

@ -132,7 +132,7 @@ async fn test_alist_grpc_fs_get_rejects_missing_token_before_io() {
token: String::new(),
path: "/local/video.mp4".to_string(),
password: String::new(),
user_agent: String::new(),
headers: std::collections::HashMap::new(),
}))
.await
.unwrap_err();

Loading…
Cancel
Save