use crate::identity::{InstanceId, WorkerId};
use crate::transport::TransportKey;
use bytes::Bytes;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use xxhash_rust::xxh3::xxh3_64;
#[derive(Debug, thiserror::Error)]
pub enum WorkerAddressError {
#[error("Key already exists: {0}")]
KeyExists(String),
#[error("Key not found: {0}")]
KeyNotFound(String),
#[error("Encoding error: {0}")]
EncodingError(#[from] rmp_serde::encode::Error),
#[error("Decoding error: {0}")]
DecodingError(#[from] rmp_serde::decode::Error),
#[error("Unsupported format version: {0}")]
UnsupportedVersion(u8),
#[error("Invalid format: {0}")]
InvalidFormat(String),
}
#[derive(Clone, PartialEq, Eq, Hash)]
pub struct WorkerAddress(Bytes);
impl Serialize for WorkerAddress {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serde_bytes::serialize(self.0.as_ref(), serializer)
}
}
impl<'de> Deserialize<'de> for WorkerAddress {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let bytes: Vec<u8> = serde_bytes::deserialize(deserializer)?;
Ok(WorkerAddress(Bytes::from(bytes)))
}
}
impl WorkerAddress {
pub fn from_encoded(bytes: impl Into<Bytes>) -> Self {
Self(bytes.into())
}
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn to_bytes(&self) -> Bytes {
self.0.clone()
}
pub fn checksum(&self) -> u64 {
xxh3_64(self.as_bytes())
}
pub fn available_transports(&self) -> Result<Vec<TransportKey>, WorkerAddressError> {
let map = decode_to_map(self.as_bytes())?;
Ok(map.keys().cloned().map(TransportKey::from).collect())
}
pub fn get_entry(&self, key: impl AsRef<str>) -> Result<Option<Bytes>, WorkerAddressError> {
let map = decode_to_map(self.as_bytes())?;
Ok(map.get(key.as_ref()).cloned())
}
}
impl fmt::Debug for WorkerAddress {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("WorkerAddress")
.field(&format_args!(
"len={}, xxh3_64=0x{:016x}",
self.0.len(),
self.checksum()
))
.finish()
}
}
impl fmt::Display for WorkerAddress {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "WorkerAddress(xxh3_64=0x{:016x})", self.checksum())
}
}
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())
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct PeerInfo {
pub instance_id: InstanceId,
pub worker_address: WorkerAddress,
}
impl PeerInfo {
pub fn new(instance_id: InstanceId, worker_address: WorkerAddress) -> Self {
Self {
instance_id,
worker_address,
}
}
pub fn instance_id(&self) -> InstanceId {
self.instance_id
}
pub fn worker_id(&self) -> WorkerId {
self.instance_id.worker_id()
}
pub fn worker_address(&self) -> &WorkerAddress {
&self.worker_address
}
pub fn address_checksum(&self) -> u64 {
self.worker_address.checksum()
}
pub fn into_address(self) -> WorkerAddress {
self.worker_address
}
pub fn into_parts(self) -> (InstanceId, WorkerAddress) {
(self.instance_id, self.worker_address)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_test_address(entries: &[(&str, &[u8])]) -> WorkerAddress {
let map: HashMap<String, Vec<u8>> = entries
.iter()
.map(|(k, v)| (k.to_string(), v.to_vec()))
.collect();
let encoded = rmp_serde::to_vec(&map).unwrap();
WorkerAddress::from_encoded(encoded)
}
#[test]
fn test_worker_address_from_encoded() {
let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let entry = address.get_entry("endpoint").unwrap();
assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
}
#[test]
fn test_worker_address_checksum() {
let address1 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let address2 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let address3 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:6666")]);
assert_eq!(address1.checksum(), address2.checksum());
assert_ne!(address1.checksum(), address3.checksum());
}
#[test]
fn test_worker_address_equality() {
let address1 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let address2 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let address3 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:6666")]);
assert_eq!(address1, address2);
assert_ne!(address1, address3);
}
#[test]
fn test_worker_address_debug() {
let address = make_test_address(&[("test", b"value")]);
let debug_str = format!("{:?}", address);
assert!(debug_str.contains("WorkerAddress"));
assert!(debug_str.contains("len="));
assert!(debug_str.contains("xxh3_64="));
}
#[test]
fn test_available_transports() {
let address = make_test_address(&[
("tcp", b"tcp://127.0.0.1:5555"),
("rdma", b"rdma://10.0.0.1:6666"),
("udp", b"udp://127.0.0.1:7777"),
]);
let transports = address.available_transports().unwrap();
assert_eq!(transports.len(), 3);
assert!(transports.contains(&TransportKey::from("tcp")));
assert!(transports.contains(&TransportKey::from("rdma")));
assert!(transports.contains(&TransportKey::from("udp")));
}
#[test]
fn test_available_transports_empty() {
let address = make_test_address(&[]);
let transports = address.available_transports().unwrap();
assert_eq!(transports.len(), 0);
}
#[test]
fn test_get_entry() {
let address =
make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555"), ("protocol", b"tcp")]);
assert_eq!(
address.get_entry("endpoint").unwrap().unwrap(),
Bytes::from_static(b"tcp://127.0.0.1:5555")
);
assert!(address.get_entry("nonexistent").unwrap().is_none());
}
#[test]
fn test_get_entry_with_transport_key() {
let address = make_test_address(&[
("tcp", b"tcp://127.0.0.1:5555"),
("rdma", b"rdma://10.0.0.1:6666"),
]);
let tcp_key = TransportKey::from("tcp");
let result = address.get_entry(tcp_key).unwrap();
assert_eq!(result, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
let result = address.get_entry(String::from("rdma")).unwrap();
assert_eq!(result, Some(Bytes::from_static(b"rdma://10.0.0.1:6666")));
}
#[test]
fn test_peer_info_creation() {
let instance_id = InstanceId::new_v4();
let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let peer_info = PeerInfo::new(instance_id, address.clone());
assert_eq!(peer_info.instance_id(), instance_id);
assert_eq!(peer_info.worker_id(), instance_id.worker_id());
assert_eq!(peer_info.worker_address(), &address);
}
#[test]
fn test_peer_info_checksum() {
let instance_id = InstanceId::new_v4();
let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let peer_info = PeerInfo::new(instance_id, address.clone());
assert_eq!(peer_info.address_checksum(), address.checksum());
}
#[test]
fn test_peer_info_into_address() {
let instance_id = InstanceId::new_v4();
let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let peer_info = PeerInfo::new(instance_id, address.clone());
let extracted_address = peer_info.into_address();
assert_eq!(extracted_address, address);
}
#[test]
fn test_peer_info_into_parts() {
let instance_id = InstanceId::new_v4();
let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let peer_info = PeerInfo::new(instance_id, address.clone());
let (extracted_id, extracted_address) = peer_info.into_parts();
assert_eq!(extracted_id, instance_id);
assert_eq!(extracted_address, address);
}
#[test]
fn test_peer_info_serde() {
let instance_id = InstanceId::new_v4();
let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
let peer_info = PeerInfo::new(instance_id, address);
let json = serde_json::to_string(&peer_info).unwrap();
let deserialized: PeerInfo = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.instance_id(), instance_id);
assert_eq!(deserialized.worker_id(), instance_id.worker_id());
let entry = deserialized.worker_address().get_entry("endpoint").unwrap();
assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
}
}