use bytes::Bytes;
use std::collections::HashMap;
use std::sync::Arc;
use velo_ext::{WorkerAddress, WorkerAddressError};
#[derive(Debug, Clone, Default)]
pub(crate) struct WorkerAddressBuilder {
entries: HashMap<String, Bytes>,
}
impl WorkerAddressBuilder {
pub fn new() -> Self {
Self {
entries: HashMap::new(),
}
}
pub fn add_entry(
&mut self,
key: impl Into<String>,
value: impl Into<Bytes>,
) -> Result<(), WorkerAddressError> {
let key = key.into();
if self.entries.contains_key(&key) {
return Err(WorkerAddressError::KeyExists(key));
}
self.entries.insert(key, value.into());
Ok(())
}
#[allow(dead_code)]
pub fn has_entry(&self, key: &str) -> bool {
self.entries.contains_key(key)
}
#[allow(dead_code)]
pub fn get_entry(&self, key: &str) -> Option<&Bytes> {
self.entries.get(key)
}
pub fn merge(&mut self, other: &WorkerAddress) -> Result<(), WorkerAddressError> {
let map = decode_to_map(other.as_bytes())?;
for key in map.keys() {
if self.entries.contains_key(key.as_ref()) {
return Err(WorkerAddressError::KeyExists(key.to_string()));
}
}
for (key, value) in map {
self.entries.insert(key.to_string(), value);
}
Ok(())
}
pub fn build(self) -> Result<WorkerAddress, WorkerAddressError> {
let serializable: HashMap<String, Vec<u8>> = self
.entries
.into_iter()
.map(|(k, v)| (k, v.to_vec()))
.collect();
let encoded = rmp_serde::to_vec(&serializable)?;
Ok(WorkerAddress::from_encoded(encoded))
}
}
fn decode_to_map(bytes: &[u8]) -> Result<HashMap<Arc<str>, Bytes>, WorkerAddressError> {
if bytes.is_empty() {
return Err(WorkerAddressError::InvalidFormat("Empty bytes".to_string()));
}
let decoded: HashMap<String, Vec<u8>> = rmp_serde::from_slice(bytes)?;
Ok(decoded
.into_iter()
.map(|(k, v)| (Arc::from(k.as_str()), Bytes::from(v)))
.collect())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_builder_basic() {
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("endpoint", Bytes::from_static(b"tcp://127.0.0.1:5555"))
.unwrap();
builder
.add_entry("protocol", Bytes::from_static(b"tcp"))
.unwrap();
assert!(builder.has_entry("endpoint"));
assert!(builder.has_entry("protocol"));
assert!(!builder.has_entry("nonexistent"));
let address = builder.build().unwrap();
assert!(!address.as_bytes().is_empty());
let entry = address.get_entry("endpoint").unwrap();
assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
}
#[test]
fn test_builder_add_duplicate_key() {
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("key", Bytes::from_static(b"value1"))
.unwrap();
let result = builder.add_entry("key", Bytes::from_static(b"value2"));
assert!(matches!(result, Err(WorkerAddressError::KeyExists(_))));
}
#[test]
fn test_builder_merge() {
let mut builder1 = WorkerAddressBuilder::new();
builder1
.add_entry("tcp", Bytes::from_static(b"tcp://127.0.0.1:5555"))
.unwrap();
let address1 = builder1.build().unwrap();
let mut builder2 = WorkerAddressBuilder::new();
builder2
.add_entry("rdma", Bytes::from_static(b"rdma://10.0.0.1:6666"))
.unwrap();
let address2 = builder2.build().unwrap();
let mut builder3 = WorkerAddressBuilder::new();
builder3.merge(&address1).unwrap();
builder3.merge(&address2).unwrap();
let final_address = builder3.build().unwrap();
assert_eq!(
final_address.get_entry("tcp").unwrap(),
Some(Bytes::from_static(b"tcp://127.0.0.1:5555"))
);
assert_eq!(
final_address.get_entry("rdma").unwrap(),
Some(Bytes::from_static(b"rdma://10.0.0.1:6666"))
);
}
#[test]
fn test_builder_merge_with_conflict() {
let mut builder1 = WorkerAddressBuilder::new();
builder1
.add_entry("tcp", Bytes::from_static(b"tcp://127.0.0.1:5555"))
.unwrap();
let address1 = builder1.build().unwrap();
let mut builder2 = WorkerAddressBuilder::new();
builder2
.add_entry("tcp", Bytes::from_static(b"tcp://different:5555"))
.unwrap();
let address2 = builder2.build().unwrap();
let mut builder3 = WorkerAddressBuilder::new();
builder3.merge(&address1).unwrap();
let result = builder3.merge(&address2);
assert!(matches!(result, Err(WorkerAddressError::KeyExists(_))));
assert!(builder3.has_entry("tcp"));
assert_eq!(
builder3.get_entry("tcp").unwrap(),
&Bytes::from_static(b"tcp://127.0.0.1:5555")
);
}
#[test]
fn test_empty_builder() {
let builder = WorkerAddressBuilder::new();
let address = builder.build().unwrap();
let transports = address.available_transports().unwrap();
assert_eq!(transports.len(), 0);
}
#[test]
fn test_builder_address_integration_get_entry() {
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("tcp", Bytes::from_static(b"tcp://127.0.0.1:5555"))
.unwrap();
builder
.add_entry("rdma", Bytes::from_static(b"rdma://10.0.0.1:6666"))
.unwrap();
builder
.add_entry("binary_data", Bytes::from_static(&[0x00, 0x01, 0x02, 0xFF]))
.unwrap();
let address = builder.build().unwrap();
assert_eq!(
address.get_entry("tcp").unwrap(),
Some(Bytes::from_static(b"tcp://127.0.0.1:5555"))
);
assert_eq!(
address.get_entry("rdma").unwrap(),
Some(Bytes::from_static(b"rdma://10.0.0.1:6666"))
);
assert_eq!(
address.get_entry("binary_data").unwrap(),
Some(Bytes::from_static(&[0x00, 0x01, 0x02, 0xFF]))
);
assert_eq!(address.get_entry("nonexistent").unwrap(), None);
}
#[test]
fn test_builder_address_integration_available_transports() {
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("tcp", Bytes::from_static(b"tcp://127.0.0.1:5555"))
.unwrap();
builder
.add_entry("rdma", Bytes::from_static(b"rdma://10.0.0.1:6666"))
.unwrap();
builder
.add_entry("grpc", Bytes::from_static(b"grpc://localhost:9000"))
.unwrap();
let address = builder.build().unwrap();
let transports = address.available_transports().unwrap();
assert_eq!(transports.len(), 3);
assert!(transports.contains(&velo_ext::TransportKey::from("tcp")));
assert!(transports.contains(&velo_ext::TransportKey::from("rdma")));
assert!(transports.contains(&velo_ext::TransportKey::from("grpc")));
}
#[test]
fn test_builder_address_integration_checksum_stability() {
let mut builder1 = WorkerAddressBuilder::new();
builder1
.add_entry("key", Bytes::from_static(b"value"))
.unwrap();
let address1 = builder1.build().unwrap();
let mut builder2 = WorkerAddressBuilder::new();
builder2
.add_entry("key", Bytes::from_static(b"value"))
.unwrap();
let address2 = builder2.build().unwrap();
assert_eq!(address1.checksum(), address2.checksum());
let mut builder3 = WorkerAddressBuilder::new();
builder3
.add_entry("key", Bytes::from_static(b"different"))
.unwrap();
let address3 = builder3.build().unwrap();
assert_ne!(address1.checksum(), address3.checksum());
}
#[test]
fn test_builder_address_integration_bytes_roundtrip() {
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("endpoint", Bytes::from_static(b"test://value"))
.unwrap();
let address = builder.build().unwrap();
let raw_bytes = address.to_bytes();
let address2 = WorkerAddress::from_encoded(raw_bytes);
assert_eq!(address, address2);
assert_eq!(address.checksum(), address2.checksum());
assert_eq!(
address.get_entry("endpoint").unwrap(),
address2.get_entry("endpoint").unwrap()
);
}
#[test]
fn test_builder_address_integration_serde_roundtrip() {
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("tcp", Bytes::from_static(b"tcp://127.0.0.1:5555"))
.unwrap();
let address = builder.build().unwrap();
let json = serde_json::to_string(&address).unwrap();
let deserialized: WorkerAddress = serde_json::from_str(&json).unwrap();
assert_eq!(address, deserialized);
assert_eq!(
address.get_entry("tcp").unwrap(),
deserialized.get_entry("tcp").unwrap()
);
}
}