use crate::cacheable::{CacheAble, CacheService};
use crate::common::model::cookies::CookieItem;
use crate::common::model::download_config::DownloadConfig;
use crate::common::model::headers::HeaderItem;
use crate::common::model::{Cookies, Headers, Request, Response};
use crate::downloader::Downloader;
use crate::errors::{DownloadError, RequestError, Result};
use crate::utils::distributed_rate_limit::DistributedSlidingWindowRateLimiter;
use crate::utils::lock::DistributedLockManager;
use dashmap::DashMap;
use futures::StreamExt;
use log::{info, warn};
use rand::Rng;
use reqwest::Client;
use reqwest::Method;
use reqwest::Proxy;
use reqwest::header::HeaderMap;
use semver::Version;
use serde::{Deserialize, Serialize};
use std::str::FromStr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant, UNIX_EPOCH};
use url::Url;
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
struct SessionState {
session_id: String,
module_id: String,
headers: Headers,
cookies: Cookies,
version: u64,
}
impl CacheAble for SessionState {
fn field() -> impl AsRef<str> {
"session_state"
}
}
#[derive(Clone)]
pub struct RequestDownloader {
pub limit: Arc<DistributedSlidingWindowRateLimiter>,
pub locker: Arc<DistributedLockManager>,
cache_service: Arc<CacheService>,
enable_session: Arc<AtomicBool>,
enable_locker: Arc<AtomicBool>,
enable_rate_limit: Arc<AtomicBool>,
proxy_clients: Arc<DashMap<String, (Client, Instant)>>,
default_client: Client,
pool_size: usize,
max_response_size: usize,
}
impl RequestDownloader {
#[inline]
fn is_session_enabled(&self, request: &Request) -> bool {
request.enable_session || self.enable_session.load(Ordering::Relaxed)
}
#[inline]
fn is_response_cache_enabled(&self, request: &Request) -> bool {
request.enable_response_cache
}
#[inline]
fn session_scope_key(&self, request: &Request) -> String {
format!("{}:{}", request.module_id(), request.run_id)
}
pub fn new(
limit: Arc<DistributedSlidingWindowRateLimiter>,
locker: Arc<DistributedLockManager>,
sync: Arc<CacheService>,
pool_size: usize,
max_response_size: usize,
) -> Self {
let default_client = Client::builder()
.pool_idle_timeout(Duration::from_secs(90))
.pool_max_idle_per_host(pool_size)
.tcp_keepalive(Duration::from_secs(60))
.tcp_nodelay(true)
.connect_timeout(Duration::from_secs(10))
.http2_keep_alive_interval(Some(Duration::from_secs(30)))
.build()
.expect("Failed to create default client");
let proxy_clients = Arc::new(DashMap::new());
let proxy_clients_clone = proxy_clients.clone();
tokio::spawn(async move {
loop {
tokio::time::sleep(Duration::from_secs(60)).await;
let now = Instant::now();
proxy_clients_clone.retain(|_, (_, last_access)| {
now.duration_since(*last_access) < Duration::from_secs(3600)
});
}
});
RequestDownloader {
limit,
locker,
cache_service: sync,
enable_session: Arc::new(AtomicBool::new(false)),
enable_locker: Arc::new(AtomicBool::new(false)),
enable_rate_limit: Arc::new(AtomicBool::new(true)),
proxy_clients,
default_client,
pool_size,
max_response_size,
}
}
async fn process_request(&self, request: Request) -> (Request, Option<SessionState>) {
if !self.is_session_enabled(&request) {
return (request, None);
}
let session_key = self.session_scope_key(&request);
let session_state = match SessionState::sync(&session_key, &self.cache_service).await {
Ok(v) => v,
Err(err) => {
warn!(
"Session load failed: session={} account={} platform={} url={} error={:?}",
session_key, request.account, request.platform, request.url, err
);
None
}
};
let mut modified_request = request;
if let Some(ref state) = session_state {
modified_request.headers.merge_if_absent(&state.headers);
modified_request.cookies.merge_if_absent(&state.cookies);
}
(modified_request, session_state)
}
async fn save_session_state(
&self,
session_key: &str,
module_id: String,
request: &Request,
response: &Response,
existing_session: Option<SessionState>,
) {
let mut session_state = existing_session.unwrap_or_else(|| SessionState {
session_id: session_key.to_string(),
module_id,
headers: Headers::default(),
cookies: Cookies::default(),
version: 0,
});
let now_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let before_len = session_state.cookies.cookies.len();
session_state.cookies.cookies.retain(|c| {
if c.max_age == Some(0) {
return false;
}
if let Some(exp) = c.expires {
if exp > 0 && exp < now_secs {
return false;
}
}
true
});
let mut dirty = session_state.cookies.cookies.len() < before_len;
let request_host = Url::parse(&request.url)
.ok()
.and_then(|u| u.host_str().map(|h| h.to_ascii_lowercase()));
let cookies_before = session_state.cookies.cookies.len();
session_state.cookies.merge_if_absent(&request.cookies);
if session_state.cookies.cookies.len() != cookies_before {
dirty = true;
}
for item in &response.cookies.cookies {
let mut normalized_item = item.clone();
if normalized_item.domain.trim().is_empty()
&& let Some(host) = &request_host
{
normalized_item.domain = host.clone();
}
if let Some(existing) = session_state
.cookies
.cookies
.iter_mut()
.find(|c| c.name == normalized_item.name && c.domain == normalized_item.domain)
{
if existing.value != normalized_item.value
|| existing.expires != normalized_item.expires
|| existing.max_age != normalized_item.max_age
{
*existing = normalized_item;
dirty = true;
}
} else {
session_state.cookies.cookies.push(normalized_item);
dirty = true;
}
}
if let Some(cache_headers) = &request.cache_headers
&& !cache_headers.is_empty()
{
let cache_headers_set: std::collections::HashSet<String> =
cache_headers.iter().map(|h| h.to_lowercase()).collect();
for (name, value) in &response.headers {
if !cache_headers_set.contains(&name.to_lowercase()) {
continue;
}
if let Some(existing) = session_state
.headers
.headers
.iter_mut()
.find(|h| h.key.eq_ignore_ascii_case(name))
{
if existing.value != *value {
existing.value = value.clone();
dirty = true;
}
} else {
session_state.headers.headers.push(HeaderItem {
key: name.clone(),
value: value.clone(),
});
dirty = true;
}
}
}
if !dirty {
return;
}
session_state.version = session_state.version.saturating_add(1);
if let Err(err) = session_state
.send_persistent(session_key, &self.cache_service)
.await
{
warn!(
"Failed to cache session state: session={} account={} platform={} url={} error={:?}",
session_key, request.account, request.platform, request.url, err
);
}
}
async fn process_response(
&self,
request: Request,
response: reqwest::Response,
pre_calculated_hash: Option<String>,
) -> Result<Response> {
let response_cache_enabled = self.is_response_cache_enabled(&request);
let request_id = request.id;
let request_hash = if response_cache_enabled {
pre_calculated_hash.or_else(|| Some(request.hash()))
} else {
None
};
let status_code = response.status().as_u16();
let response_headers: Vec<(String, String)> = response
.headers()
.iter()
.map(|(name, value)| (name.to_string(), value.to_str().unwrap_or("").to_string()))
.collect();
let response_cookies: Vec<CookieItem> = response
.cookies()
.map(|cookie| CookieItem {
name: cookie.name().to_string(),
value: cookie.value().to_string(),
domain: cookie.domain().unwrap_or("").to_string(),
path: cookie.path().unwrap_or("/").to_string(),
expires: cookie
.expires()
.and_then(|exp| exp.duration_since(UNIX_EPOCH).ok())
.map(|d| d.as_secs()),
max_age: cookie.max_age().map(|d| d.as_secs()),
secure: cookie.secure(),
http_only: Some(cookie.http_only()),
})
.collect();
let content_length = response.content_length();
if let Some(len) = content_length {
if len > self.max_response_size as u64 {
warn!(
"Response size {} exceeds limit {}, aborting download for {}",
len, self.max_response_size, request.url
);
return Err(DownloadError::DownloadFailed("Response too large".into()).into());
}
}
let body_timeout = Duration::from_secs(request.timeout.max(30));
let limit = self.max_response_size;
let url_for_log = request.url.clone();
let content = tokio::time::timeout(body_timeout, async {
let mut buf = Vec::new();
let mut stream = response.bytes_stream();
while let Some(item) = stream.next().await {
let chunk =
item.map_err(|e: reqwest::Error| DownloadError::DownloadFailed(e.into()))?;
if buf.len() + chunk.len() > limit {
warn!(
"Response size exceeds limit {}, aborting download for {}",
limit, url_for_log
);
return Err(DownloadError::DownloadFailed("Response too large".into()).into());
}
buf.extend_from_slice(&chunk);
}
Ok::<Vec<u8>, crate::errors::Error>(buf)
})
.await
.map_err(|_| {
warn!(
"Body read timed out after {}s for {}",
body_timeout.as_secs(),
url_for_log
);
crate::errors::Error::from(DownloadError::DownloadFailed(
format!("body read timed out after {}s", body_timeout.as_secs()).into(),
))
})??;
Ok(Response {
id: request_id,
platform: request.platform,
account: request.account,
module: request.module,
status_code,
cookies: Cookies {
cookies: response_cookies,
},
content,
storage_path: None,
headers: response_headers,
task_retry_times: request.task_retry_times,
metadata: request.meta,
download_middleware: request.download_middleware,
data_middleware: request.data_middleware,
task_finished: request.task_finished,
context: request.context,
run_id: request.run_id,
prefix_request: request.prefix_request,
request_hash,
priority: request.priority,
})
}
async fn get_client(&self, proxy: Option<&String>) -> Result<Client> {
if let Some(proxy_url) = proxy {
if let Some(mut entry) = self.proxy_clients.get_mut(proxy_url) {
entry.1 = Instant::now();
return Ok(entry.0.clone());
}
let reqwest_proxy =
Proxy::all(proxy_url).map_err(|e| DownloadError::ClientError(e.into()))?;
let client = Client::builder()
.proxy(reqwest_proxy)
.pool_idle_timeout(Duration::from_secs(90))
.pool_max_idle_per_host(self.pool_size) .tcp_keepalive(Duration::from_secs(60))
.tcp_nodelay(true)
.connect_timeout(Duration::from_secs(10))
.http2_keep_alive_interval(Some(Duration::from_secs(30)))
.build()
.map_err(|e| DownloadError::ClientError(e.into()))?;
if self.proxy_clients.len() < 1000 {
self.proxy_clients
.insert(proxy_url.clone(), (client.clone(), Instant::now()));
}
Ok(client)
} else {
Ok(self.default_client.clone())
}
}
async fn do_download(
&self,
request: Request,
pre_calculated_hash: Option<String>,
) -> Result<Response> {
let _request_id = request.id;
let session_enabled = self.is_session_enabled(&request);
if let Some(seconds) = request.time_sleep_secs {
tokio::time::sleep(Duration::from_secs(seconds)).await;
}
let rate_limit_enabled = self.enable_rate_limit.load(Ordering::Relaxed);
info!(
"[do_download] enable_rate_limit={} request_id={}",
rate_limit_enabled, request.id
);
if rate_limit_enabled {
let limit_id = if request.limit_id.is_empty() {
request.module_id()
} else {
request.limit_id.clone()
};
self.limit.wait_for_permit(&limit_id).await?;
}
let (mut request, loaded_session) = self.process_request(request).await;
let proxy_url = request.proxy.as_ref().map(|p| p.to_string());
let client = self.get_client(proxy_url.as_ref()).await?;
let method =
Method::from_str(&request.method).map_err(|e| RequestError::InvalidMethod(e.into()))?;
let url = match &request.params {
Some(params) => Url::parse_with_params(&request.url, params)
.map_err(|e| RequestError::InvalidUrl(e.to_string()))?,
None => {
Url::parse(&request.url).map_err(|e| RequestError::InvalidUrl(e.to_string()))?
}
};
let cookie_header = if request.cookies.cookies.is_empty() {
None
} else {
request.cookies.cookie_header_for_url(&url)
};
let mut request_builder = client.request(method, url);
let headers: HeaderMap = HeaderMap::from(&request.headers);
request_builder = request_builder.headers(headers);
if let Some(cookie_str) = cookie_header {
request_builder = request_builder.header(reqwest::header::COOKIE, cookie_str);
}
request_builder = request_builder.timeout(Duration::from_secs(request.timeout));
if let Some(body) = request.body.take() {
request_builder = request_builder.body(body);
}
if let Some(form) = request.form.take() {
request_builder = request_builder.form(&form);
}
if let Some(json) = request.json.take() {
request_builder = request_builder.json(&json);
}
let start = std::time::Instant::now();
let result = request_builder.send().await;
let response = match result {
Ok(res) => res,
Err(e) => {
crate::common::metrics::inc_error(
"engine",
"downloader",
"network",
"send_failed",
1,
);
if self.enable_rate_limit.load(Ordering::Relaxed) {
let limit_id = if request.limit_id.is_empty() {
request.module_id()
} else {
request.limit_id.clone()
};
if e.is_connect() || e.is_timeout() {
warn!(
"Circuit Breaker triggered for {} due to network error: {}. Suspending for 10s.",
limit_id, e
);
self.limit
.suspend(&limit_id, Duration::from_secs(10))
.await
.ok();
}
}
return Err(DownloadError::DownloadFailed(e.into()).into());
}
};
if self.enable_rate_limit.load(Ordering::Relaxed) {
let limit_id = if request.limit_id.is_empty() {
request.module_id()
} else {
request.limit_id.clone()
};
let status = response.status();
if status == reqwest::StatusCode::TOO_MANY_REQUESTS {
warn!(
"Circuit Breaker triggered for {}, suspending for 60s",
limit_id
);
self.limit
.suspend(&limit_id, Duration::from_secs(60))
.await
.ok();
} else if status == reqwest::StatusCode::SERVICE_UNAVAILABLE {
warn!("Backpressure triggered for {}, decreasing limit", limit_id);
self.limit.decrease_limit(&limit_id, 0.5).await.ok();
} else if status.is_success() && rand::rng().random_bool(0.1) {
self.limit.try_restore_limit(&limit_id, 1.1).await.ok();
}
}
let duration = start.elapsed().as_secs_f64();
crate::common::metrics::observe_latency(
"engine",
"downloader",
"http_request",
"success",
duration,
);
crate::common::metrics::inc_throughput(
"engine",
"downloader",
"http_request",
"success",
1,
);
let session_info = if session_enabled {
Some((
self.session_scope_key(&request),
request.module_id(),
request.clone(),
loaded_session,
))
} else {
None
};
let response_processed = self
.process_response(request, response, pre_calculated_hash.clone())
.await?;
if let Some((session_key, module_id, request_for_cache, existing_session)) = session_info {
self.save_session_state(
&session_key,
module_id,
&request_for_cache,
&response_processed,
existing_session,
)
.await;
}
Ok(response_processed)
}
pub async fn set_limit_config(&self, id: &str, limit: f32) {
self.limit.set_limit(id, limit).await.ok();
}
}
#[async_trait::async_trait]
impl Downloader for RequestDownloader {
async fn set_config(&self, id: &str, config: DownloadConfig) {
if config.enable_session {
}
if self.enable_session.load(Ordering::Relaxed) != config.enable_session {
self.enable_session
.store(config.enable_session, Ordering::Relaxed);
}
if self.enable_locker.load(Ordering::Relaxed) != config.enable_locker {
self.enable_locker
.store(config.enable_locker, Ordering::Relaxed);
}
info!(
"[set_config] id={} config.enable_rate_limit={} current={}",
id,
config.enable_rate_limit,
self.enable_rate_limit.load(Ordering::Relaxed)
);
if self.enable_rate_limit.load(Ordering::Relaxed) != config.enable_rate_limit {
info!(
"[set_config] changing enable_rate_limit from {} to {} for id={}",
self.enable_rate_limit.load(Ordering::Relaxed),
config.enable_rate_limit,
id
);
self.enable_rate_limit
.store(config.enable_rate_limit, Ordering::Relaxed);
}
self.limit.set_limit(id, config.rate_limit).await.ok();
}
async fn set_limit(&self, id: &str, limit: f32) {
self.set_limit_config(id, limit).await;
}
fn name(&self) -> String {
"request_downloader".to_string()
}
fn version(&self) -> Version {
Version::parse("0.1.0").unwrap()
}
async fn download(&self, request: Request) -> Result<Response> {
let session_enabled = self.is_session_enabled(&request);
let response_cache_enabled = self.is_response_cache_enabled(&request);
let request_hash = if response_cache_enabled {
Some(request.hash())
} else {
None
};
if let Some(hash) = request_hash.as_ref()
&& let Ok(Some(response)) = Response::sync(hash, &self.cache_service).await
{
info!(
"Cache hit: request_id={} account={} platform={} url={}",
request.id, request.account, request.platform, request.url
);
if session_enabled {
let session_key = self.session_scope_key(&request);
let module_id = request.module_id();
let existing_session = match SessionState::sync(&session_key, &self.cache_service)
.await
{
Ok(v) => v,
Err(err) => {
warn!(
"Session load failed (cache hit path): session={} account={} platform={} url={} error={:?}",
session_key, request.account, request.platform, request.url, err
);
None
}
};
self.save_session_state(
&session_key,
module_id,
&request,
&response,
existing_session,
)
.await;
}
return Ok(Response {
id: request.id,
platform: request.platform.clone(),
account: request.account.clone(),
module: request.module.clone(),
status_code: response.status_code,
cookies: Cookies {
cookies: response.cookies.cookies,
},
content: response.content,
storage_path: response.storage_path,
headers: response.headers,
task_retry_times: request.task_retry_times,
metadata: request.meta.clone(),
download_middleware: request.download_middleware.clone(),
data_middleware: request.data_middleware.clone(),
task_finished: request.task_finished,
context: request.context.clone(),
run_id: request.run_id,
prefix_request: request.id,
request_hash: request_hash.clone(),
priority: request.priority,
});
}
let locker_enabled = request
.enable_locker
.unwrap_or(self.enable_locker.load(Ordering::Relaxed));
if locker_enabled {
let key = format!("task-download-{}-{}", request.module_id(), request.run_id);
match self
.locker
.acquire_lock(&key, 30, Duration::from_secs(10))
.await
{
Ok(_) => {
let response = self.do_download(request, request_hash.clone()).await;
self.locker.release_lock(&key).await.ok();
response
}
Err(_) => {
warn!(
"Failed to acquire download lock: task_id={} account={} platform={} url={}, proceeding without lock",
request.task_id(),
request.account,
request.platform,
request.url
);
self.do_download(request, request_hash.clone()).await
}
}
} else {
self.do_download(request, request_hash.clone()).await
}
}
async fn health_check(&self) -> Result<()> {
Ok(())
}
async fn close(&self) -> Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::utils::distributed_rate_limit::RateLimitConfig;
use std::sync::Arc;
use uuid::Uuid;
#[tokio::test]
async fn test_downloader_creation() {
let lock_manager = Arc::new(DistributedLockManager::new("test"));
let config = RateLimitConfig::new(10.0);
let limiter = Arc::new(DistributedSlidingWindowRateLimiter::new(
lock_manager.clone(),
"test",
config,
));
let cache_service = Arc::new(CacheService::new("test".to_string(), None, None));
let downloader =
RequestDownloader::new(limiter, lock_manager, cache_service, 200, 1024 * 1024 * 10);
assert_eq!(downloader.name(), "request_downloader");
}
#[tokio::test]
async fn test_downloader_rate_limit_execution() {
let lock_manager = Arc::new(DistributedLockManager::new("test_exec"));
let config = RateLimitConfig::new(10.0); let limiter = Arc::new(DistributedSlidingWindowRateLimiter::new(
lock_manager.clone(),
"test_exec",
config,
));
let cache_service = Arc::new(CacheService::new("test".to_string(), None, None));
let downloader = RequestDownloader::new(
limiter.clone(),
lock_manager,
cache_service,
200,
1024 * 1024 * 10,
);
downloader.enable_rate_limit.store(true, Ordering::Relaxed);
downloader.set_limit("test_exec", 1.0).await;
let mut request = Request::new("http://example.com", "GET");
request.id = Uuid::new_v4();
request.limit_id = "test_exec".to_string();
limiter.record("test_exec").await.unwrap();
let delay = limiter.verify("test_exec").await.unwrap();
assert!(delay.is_some());
assert!(delay.unwrap() > 0);
}
}