use dashmap::DashMap;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use tokio::sync::RwLock;
use tracing::{debug, info};
use super::request_group::{DownloadOptions, DownloadStatus, GroupId, RequestGroup};
use crate::error::Result;
pub struct RequestGroupMan {
groups: DashMap<GroupId, Arc<RwLock<RequestGroup>>>,
next_gid: AtomicU64,
global_download_limit: Arc<RwLock<Option<u64>>>,
global_upload_limit: Arc<RwLock<Option<u64>>>,
}
impl RequestGroupMan {
pub fn new() -> Self {
info!("Initializing request group manager");
RequestGroupMan {
groups: DashMap::new(),
next_gid: AtomicU64::new(1),
global_download_limit: Arc::new(RwLock::new(None)),
global_upload_limit: Arc::new(RwLock::new(None)),
}
}
pub async fn add_group(&self, uris: Vec<String>, options: DownloadOptions) -> Result<GroupId> {
let gid = self.generate_gid();
let group = RequestGroup::new(gid, uris, options);
self.groups.insert(gid, Arc::new(RwLock::new(group)));
info!("Adding download task #{}", gid.value());
debug!("Current total tasks: {}", self.groups.len());
Ok(gid)
}
pub async fn add_group_with_gid(
&self,
gid: GroupId,
uris: Vec<String>,
options: DownloadOptions,
) -> Result<()> {
if self.groups.contains_key(&gid) {
return Err(crate::error::Aria2Error::DownloadFailed(format!(
"GID {} already exists",
gid.to_hex_string()
)));
}
let group = RequestGroup::new(gid, uris, options);
self.groups.insert(gid, Arc::new(RwLock::new(group)));
info!("Adding download task (RPC) #{}", gid.to_hex_string());
Ok(())
}
pub fn group_by_hex(&self, hex: &str) -> Option<Arc<RwLock<RequestGroup>>> {
let gid = GroupId::from_hex_string(hex)?;
self.groups.get(&gid).map(|v| v.clone())
}
pub fn group_by_id(&self, gid: GroupId) -> Option<Arc<RwLock<RequestGroup>>> {
self.groups.get(&gid).map(|v| v.clone())
}
pub fn all_groups(&self) -> Vec<(GroupId, Arc<RwLock<RequestGroup>>)> {
self.groups
.iter()
.map(|entry| (*entry.key(), entry.value().clone()))
.collect()
}
pub fn remove_group_by_id(&self, gid: GroupId) -> Option<Arc<RwLock<RequestGroup>>> {
self.groups.remove(&gid).map(|(_, v)| v)
}
pub async fn remove_group(&self, gid: GroupId) -> Result<()> {
if let Some((_, group_lock)) = self.groups.remove(&gid) {
let mut group = group_lock.write().await;
group.remove().await?;
info!("Removing download task #{}", gid.value());
debug!("Remaining tasks: {}", self.groups.len());
}
Ok(())
}
pub async fn pause_group(&self, gid: GroupId) -> Result<()> {
if let Some(group_lock) = self.groups.get(&gid) {
let mut group = group_lock.write().await;
group.pause().await?;
info!("Pausing download task #{}", gid.value());
}
Ok(())
}
pub async fn unpause_group(&self, gid: GroupId) -> Result<()> {
if let Some(group_lock) = self.groups.get(&gid) {
let mut group = group_lock.write().await;
if group.status().await.is_paused() {
group.start().await?;
info!("Resuming download task #{}", gid.value());
}
}
Ok(())
}
pub async fn update_group_options(
&self,
gid_hex: &str,
changes: HashMap<String, serde_json::Value>,
) -> std::result::Result<(), String> {
let group = self
.group_by_hex(gid_hex)
.ok_or_else(|| format!("GID {} not found", gid_hex))?;
let mut g = group.write().await;
for (key, value) in changes {
let applied = g.update_option(&key, value).await;
if !applied {
return Err(format!("Option '{}' cannot be changed at runtime", key));
}
}
Ok(())
}
pub async fn get_group(&self, gid: GroupId) -> Option<Arc<RwLock<RequestGroup>>> {
self.groups.get(&gid).map(|v| v.clone())
}
pub async fn list_groups(&self) -> Vec<Arc<RwLock<RequestGroup>>> {
self.groups.iter().map(|v| v.clone()).collect()
}
pub async fn get_active_groups(&self) -> Vec<Arc<RwLock<RequestGroup>>> {
let mut active = Vec::new();
for entry in self.groups.iter() {
let group = entry.read().await;
if group.status().await.is_active() {
active.push(entry.clone());
}
}
active
}
pub async fn get_waiting_groups(&self) -> Vec<Arc<RwLock<RequestGroup>>> {
let mut waiting = Vec::new();
for entry in self.groups.iter() {
let group = entry.read().await;
if matches!(group.status().await, DownloadStatus::Waiting) {
waiting.push(entry.clone());
}
}
waiting
}
pub async fn count(&self) -> usize {
self.groups.len()
}
pub async fn active_count(&self) -> usize {
self.get_active_groups().await.len()
}
pub async fn set_global_speed_limit(
&self,
download_limit: Option<u64>,
upload_limit: Option<u64>,
) {
*self.global_download_limit.write().await = download_limit;
*self.global_upload_limit.write().await = upload_limit;
debug!(
"Setting global speed limit - download: {:?}, upload: {:?}",
download_limit, upload_limit
);
}
pub async fn global_download_limit(&self) -> Option<u64> {
*self.global_download_limit.read().await
}
pub async fn global_upload_limit(&self) -> Option<u64> {
*self.global_upload_limit.read().await
}
fn generate_gid(&self) -> GroupId {
let gid = self.next_gid.fetch_add(1, Ordering::SeqCst);
GroupId(gid)
}
pub async fn clear_completed(&self) -> Result<usize> {
let to_remove: Vec<GroupId> = self
.groups
.iter()
.filter_map(|entry| {
let group_lock = entry.value();
futures::executor::block_on(async {
let group = group_lock.read().await;
if matches!(
group.status().await,
DownloadStatus::Complete | DownloadStatus::Error(_)
) {
Some(*entry.key())
} else {
None
}
})
})
.collect();
let count = to_remove.len();
for gid in &to_remove {
self.groups.remove(gid);
}
info!("Cleared {} completed tasks", count);
Ok(count)
}
}
impl Default for RequestGroupMan {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Instant;
#[tokio::test]
async fn test_concurrent_add_groups() {
let man = Arc::new(RequestGroupMan::new());
let num_tasks = 100;
let mut handles = vec![];
for i in 0..num_tasks {
let man_clone = man.clone();
let handle = tokio::spawn(async move {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man_clone.add_group(vec![uri], options).await
});
handles.push(handle);
}
let results: Vec<_> = futures::future::join_all(handles).await;
for result in results {
assert!(result.is_ok());
let gid = result.unwrap().unwrap();
assert!(gid.value() > 0);
}
assert_eq!(man.count().await, num_tasks);
}
#[tokio::test]
async fn test_concurrent_get_groups() {
let man = Arc::new(RequestGroupMan::new());
let num_groups = 50;
for i in 0..num_groups {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man.add_group(vec![uri], options).await.unwrap();
}
let mut handles = vec![];
for i in 1..=num_groups {
let man_clone = man.clone();
let handle = tokio::spawn(async move {
let gid = GroupId(i);
man_clone.get_group(gid).await
});
handles.push(handle);
}
let results: Vec<_> = futures::future::join_all(handles).await;
for result in results {
assert!(result.is_ok());
let group = result.unwrap();
assert!(group.is_some());
}
}
#[tokio::test]
async fn test_concurrent_add_remove() {
let man = Arc::new(RequestGroupMan::new());
let num_tasks = 50;
let mut add_handles = vec![];
for i in 0..num_tasks {
let man_clone = man.clone();
let handle = tokio::spawn(async move {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man_clone.add_group(vec![uri], options).await
});
add_handles.push(handle);
}
let add_results: Vec<_> = futures::future::join_all(add_handles).await;
let gids: Vec<_> = add_results
.into_iter()
.map(|r| r.unwrap().unwrap())
.collect();
let mut remove_handles = vec![];
for gid in gids {
let man_clone = man.clone();
let handle = tokio::spawn(async move { man_clone.remove_group(gid).await });
remove_handles.push(handle);
}
let remove_results: Vec<_> = futures::future::join_all(remove_handles).await;
for result in remove_results {
assert!(result.is_ok());
}
assert_eq!(man.count().await, 0);
}
#[tokio::test]
async fn test_lock_contention_reduction() {
let man = Arc::new(RequestGroupMan::new());
let num_operations = 100;
for i in 0..num_operations {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man.add_group(vec![uri], options).await.unwrap();
}
let start = Instant::now();
let mut handles = vec![];
for i in 1..=num_operations {
let man_clone = man.clone();
let handle = tokio::spawn(async move {
for _ in 0..10 {
let gid = GroupId(i);
let _ = man_clone.get_group(gid).await;
}
});
handles.push(handle);
}
futures::future::join_all(handles).await;
let duration = start.elapsed();
println!("Concurrent read operations took: {:?}", duration);
assert!(
duration.as_millis() < 1000,
"Concurrent operations should be fast"
);
}
#[tokio::test]
async fn test_concurrent_list_groups() {
let man = Arc::new(RequestGroupMan::new());
let num_groups = 30;
for i in 0..num_groups {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man.add_group(vec![uri], options).await.unwrap();
}
let man_clone = man.clone();
let list_handle = tokio::spawn(async move {
let mut counts = vec![];
for _ in 0..10 {
let groups = man_clone.list_groups().await;
counts.push(groups.len());
}
counts
});
for i in num_groups..num_groups + 10 {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man.add_group(vec![uri], options).await.unwrap();
}
let counts = list_handle.await.unwrap();
for count in counts {
assert!(count >= num_groups);
}
}
#[tokio::test]
async fn test_atomic_gid_generation() {
let man = Arc::new(RequestGroupMan::new());
let num_tasks = 1000;
let mut handles = vec![];
for _ in 0..num_tasks {
let man_clone = man.clone();
let handle = tokio::spawn(async move {
let uri = "http://example.com/file.bin".to_string();
let options = DownloadOptions::default();
man_clone.add_group(vec![uri], options).await
});
handles.push(handle);
}
let results: Vec<_> = futures::future::join_all(handles).await;
let gids: Vec<_> = results.into_iter().map(|r| r.unwrap().unwrap()).collect();
let mut gid_values: Vec<_> = gids.iter().map(|g| g.value()).collect();
gid_values.sort();
gid_values.dedup();
assert_eq!(gid_values.len(), num_tasks, "All GIDs should be unique");
}
#[tokio::test]
async fn test_dashmap_iteration_non_blocking() {
let man = Arc::new(RequestGroupMan::new());
let num_groups = 50;
for i in 0..num_groups {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man.add_group(vec![uri], options).await.unwrap();
}
let man_clone = man.clone();
let iter_handle = tokio::spawn(async move {
let mut sum = 0;
for _ in 0..10 {
let groups = man_clone.list_groups().await;
sum += groups.len();
}
sum
});
for i in num_groups..num_groups + 20 {
let uri = format!("http://example.com/file{}.bin", i);
let options = DownloadOptions::default();
man.add_group(vec![uri], options).await.unwrap();
}
let sum = iter_handle.await.unwrap();
assert!(sum > 0);
}
}