use crate::error::Result;
use crate::service::UriService;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use url::Url;
pub struct InMemoryUriRegister {
uri_to_id: Arc<RwLock<HashMap<String, u64>>>,
id_to_uri: Arc<RwLock<HashMap<u64, String>>>,
next_id: Arc<RwLock<u64>>,
}
impl InMemoryUriRegister {
pub fn new() -> Self {
Self {
uri_to_id: Arc::new(RwLock::new(HashMap::new())),
id_to_uri: Arc::new(RwLock::new(HashMap::new())),
next_id: Arc::new(RwLock::new(1)), }
}
fn validate_uri(uri: &str) -> Result<()> {
Url::parse(uri).map_err(|e| {
crate::error::Error::InvalidUri(format!("Invalid URI '{}': {}", uri, e))
})?;
Ok(())
}
}
impl Default for InMemoryUriRegister {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl UriService for InMemoryUriRegister {
async fn register_uri(&self, uri: &str) -> Result<u64> {
Self::validate_uri(uri)?;
{
let uri_to_id = self.uri_to_id.read().await;
if let Some(&id) = uri_to_id.get(uri) {
return Ok(id);
}
}
let mut uri_to_id = self.uri_to_id.write().await;
let mut id_to_uri = self.id_to_uri.write().await;
let mut next_id = self.next_id.write().await;
if let Some(&id) = uri_to_id.get(uri) {
return Ok(id);
}
let id = *next_id;
*next_id += 1;
uri_to_id.insert(uri.to_string(), id);
id_to_uri.insert(id, uri.to_string());
Ok(id)
}
async fn register_uri_batch(&self, uris: &[String]) -> Result<Vec<u64>> {
for uri in uris {
Self::validate_uri(uri)?;
}
let mut result = Vec::with_capacity(uris.len());
{
let uri_to_id = self.uri_to_id.read().await;
for uri in uris {
if let Some(&id) = uri_to_id.get(uri) {
result.push(Some(id));
} else {
result.push(None);
}
}
}
let missing_indices: Vec<usize> = result
.iter()
.enumerate()
.filter_map(|(idx, id)| if id.is_none() { Some(idx) } else { None })
.collect();
if missing_indices.is_empty() {
return Ok(result.into_iter().map(|id| id.unwrap()).collect());
}
let mut uri_to_id = self.uri_to_id.write().await;
let mut id_to_uri = self.id_to_uri.write().await;
let mut next_id = self.next_id.write().await;
for idx in missing_indices {
let uri = &uris[idx];
if let Some(&id) = uri_to_id.get(uri) {
result[idx] = Some(id);
continue;
}
let id = *next_id;
*next_id += 1;
uri_to_id.insert(uri.clone(), id);
id_to_uri.insert(id, uri.clone());
result[idx] = Some(id);
}
Ok(result.into_iter().map(|id| id.unwrap()).collect())
}
async fn register_uri_batch_hashmap(&self, uris: &[String]) -> Result<HashMap<String, u64>> {
for uri in uris {
Self::validate_uri(uri)?;
}
let mut result = HashMap::new();
let unique_uris: std::collections::HashSet<_> = uris.iter().collect();
{
let uri_to_id = self.uri_to_id.read().await;
for uri in &unique_uris {
if let Some(&id) = uri_to_id.get(*uri) {
result.insert((*uri).clone(), id);
}
}
}
let missing_uris: Vec<String> = unique_uris
.iter()
.filter(|uri| !result.contains_key(**uri))
.map(|s| (*s).clone())
.collect();
if missing_uris.is_empty() {
return Ok(result);
}
let mut uri_to_id = self.uri_to_id.write().await;
let mut id_to_uri = self.id_to_uri.write().await;
let mut next_id = self.next_id.write().await;
for uri in missing_uris {
if let Some(&id) = uri_to_id.get(&uri) {
result.insert(uri, id);
continue;
}
let id = *next_id;
*next_id += 1;
uri_to_id.insert(uri.clone(), id);
id_to_uri.insert(id, uri.clone());
result.insert(uri, id);
}
Ok(result)
}
}
impl Clone for InMemoryUriRegister {
fn clone(&self) -> Self {
Self {
uri_to_id: Arc::clone(&self.uri_to_id),
id_to_uri: Arc::clone(&self.id_to_uri),
next_id: Arc::clone(&self.next_id),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_register_uri() {
let register = InMemoryUriRegister::new();
let id1 = register.register_uri("http://example.org/1").await.unwrap();
let id2 = register.register_uri("http://example.org/2").await.unwrap();
let id1_again = register.register_uri("http://example.org/1").await.unwrap();
assert_eq!(id1, id1_again, "Same URI should return same ID");
assert_ne!(id1, id2, "Different URIs should have different IDs");
}
#[tokio::test]
async fn test_register_uri_batch() {
let register = InMemoryUriRegister::new();
let uris = vec![
"http://example.org/1".to_string(),
"http://example.org/2".to_string(),
"http://example.org/3".to_string(),
];
let ids = register.register_uri_batch(&uris).await.unwrap();
assert_eq!(ids.len(), 3, "Should return 3 IDs");
let unique_ids: std::collections::HashSet<_> = ids.iter().copied().collect();
assert_eq!(unique_ids.len(), 3, "All IDs should be unique");
}
#[tokio::test]
async fn test_register_uri_batch_order_preservation() {
let register = InMemoryUriRegister::new();
let uris = vec![
"http://example.org/1".to_string(),
"http://example.org/2".to_string(),
"http://example.org/3".to_string(),
];
let ids = register.register_uri_batch(&uris).await.unwrap();
for (i, uri) in uris.iter().enumerate() {
let id = register.register_uri(uri).await.unwrap();
assert_eq!(
ids[i], id,
"Batch ID at index {} should match single registration",
i
);
}
}
#[tokio::test]
async fn test_register_uri_batch_with_duplicates() {
let register = InMemoryUriRegister::new();
let uris1 = vec![
"http://example.org/1".to_string(),
"http://example.org/2".to_string(),
];
let ids1 = register.register_uri_batch(&uris1).await.unwrap();
assert_eq!(ids1.len(), 2);
let uris2 = vec![
"http://example.org/2".to_string(), "http://example.org/3".to_string(), ];
let ids2 = register.register_uri_batch(&uris2).await.unwrap();
assert_eq!(ids2.len(), 2);
assert_eq!(ids1[1], ids2[0], "Existing URI should return same ID");
}
#[tokio::test]
async fn test_register_uri_batch_empty() {
let register = InMemoryUriRegister::new();
let ids = register.register_uri_batch(&[]).await.unwrap();
assert_eq!(ids.len(), 0, "Empty input should return empty result");
}
#[tokio::test]
async fn test_concurrent_registration() {
let register = InMemoryUriRegister::new();
let uri = "http://example.org/concurrent";
let mut handles = vec![];
for _ in 0..10 {
let reg = register.clone();
let handle = tokio::spawn(async move { reg.register_uri(uri).await.unwrap() });
handles.push(handle);
}
let mut ids = vec![];
for handle in handles {
ids.push(handle.await.unwrap());
}
let unique_ids: std::collections::HashSet<_> = ids.into_iter().collect();
assert_eq!(
unique_ids.len(),
1,
"Concurrent registration should return same ID"
);
}
#[tokio::test]
async fn test_register_uri_batch_hashmap() {
let register = InMemoryUriRegister::new();
let uris = vec![
"http://example.org/1".to_string(),
"http://example.org/2".to_string(),
"http://example.org/3".to_string(),
];
let map = register.register_uri_batch_hashmap(&uris).await.unwrap();
assert_eq!(map.len(), 3, "Should return 3 mappings");
assert!(map.contains_key("http://example.org/1"));
assert!(map.contains_key("http://example.org/2"));
assert!(map.contains_key("http://example.org/3"));
}
#[tokio::test]
async fn test_register_uri_batch_hashmap_with_duplicates() {
let register = InMemoryUriRegister::new();
let uris = vec![
"http://example.org/1".to_string(),
"http://example.org/2".to_string(),
"http://example.org/1".to_string(), ];
let map = register.register_uri_batch_hashmap(&uris).await.unwrap();
assert_eq!(map.len(), 2, "Duplicates should be removed");
assert!(map.contains_key("http://example.org/1"));
assert!(map.contains_key("http://example.org/2"));
}
#[tokio::test]
async fn test_register_uri_batch_hashmap_with_existing() {
let register = InMemoryUriRegister::new();
let uris1 = vec![
"http://example.org/1".to_string(),
"http://example.org/2".to_string(),
];
let map1 = register.register_uri_batch_hashmap(&uris1).await.unwrap();
let uris2 = vec![
"http://example.org/2".to_string(), "http://example.org/3".to_string(), ];
let map2 = register.register_uri_batch_hashmap(&uris2).await.unwrap();
assert_eq!(
map1.get("http://example.org/2"),
map2.get("http://example.org/2")
);
}
#[tokio::test]
async fn test_register_uri_batch_hashmap_empty() {
let register = InMemoryUriRegister::new();
let map = register.register_uri_batch_hashmap(&[]).await.unwrap();
assert_eq!(map.len(), 0, "Empty input should return empty map");
}
#[tokio::test]
async fn test_invalid_uri_validation() {
let register = InMemoryUriRegister::new();
let invalid_uris = vec![
"not a uri",
"://missing-scheme",
"http://",
"",
"just-a-string",
"ftp://[invalid",
];
for invalid_uri in invalid_uris {
let result = register.register_uri(invalid_uri).await;
assert!(
result.is_err(),
"Invalid URI '{}' should be rejected",
invalid_uri
);
if let Err(e) = result {
assert!(
matches!(e, crate::error::Error::InvalidUri(_)),
"Error should be InvalidUri, got: {:?}",
e
);
}
}
}
#[tokio::test]
async fn test_valid_uri_validation() {
let register = InMemoryUriRegister::new();
let valid_uris = vec![
"http://example.org",
"https://example.org/path",
"ftp://ftp.example.org/file.txt",
"http://example.org:8080/path?query=value",
"https://user:pass@example.org/path#fragment",
"file:///path/to/file",
];
for valid_uri in valid_uris {
let result = register.register_uri(valid_uri).await;
assert!(
result.is_ok(),
"Valid URI '{}' should be accepted, got error: {:?}",
valid_uri,
result.err()
);
}
}
#[tokio::test]
async fn test_invalid_uri_batch_validation() {
let register = InMemoryUriRegister::new();
let uris = vec![
"http://valid.org".to_string(),
"invalid uri".to_string(), "http://another-valid.org".to_string(),
];
let result = register.register_uri_batch(&uris).await;
assert!(result.is_err(), "Batch with invalid URI should fail");
}
#[tokio::test]
async fn test_invalid_uri_batch_hashmap_validation() {
let register = InMemoryUriRegister::new();
let uris = vec![
"http://valid.org".to_string(),
"not-a-valid-uri".to_string(), ];
let result = register.register_uri_batch_hashmap(&uris).await;
assert!(
result.is_err(),
"Batch hashmap with invalid URI should fail"
);
}
}