1use crate::identity::{InstanceId, WorkerId};
14use crate::transport::TransportKey;
15
16use bytes::Bytes;
17use serde::{Deserialize, Serialize};
18use std::collections::HashMap;
19use std::fmt;
20use std::sync::Arc;
21use xxhash_rust::xxh3::xxh3_64;
22
23#[derive(Debug, thiserror::Error)]
25pub enum WorkerAddressError {
26 #[error("Key already exists: {0}")]
28 KeyExists(String),
29
30 #[error("Key not found: {0}")]
32 KeyNotFound(String),
33
34 #[error("Encoding error: {0}")]
36 EncodingError(#[from] rmp_serde::encode::Error),
37
38 #[error("Decoding error: {0}")]
40 DecodingError(#[from] rmp_serde::decode::Error),
41
42 #[error("Unsupported format version: {0}")]
44 UnsupportedVersion(u8),
45
46 #[error("Invalid format: {0}")]
48 InvalidFormat(String),
49}
50
51#[derive(Clone, PartialEq, Eq, Hash)]
61pub struct WorkerAddress(Bytes);
62
63impl Serialize for WorkerAddress {
65 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
66 where
67 S: serde::Serializer,
68 {
69 serde_bytes::serialize(self.0.as_ref(), serializer)
70 }
71}
72
73impl<'de> Deserialize<'de> for WorkerAddress {
74 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
75 where
76 D: serde::Deserializer<'de>,
77 {
78 let bytes: Vec<u8> = serde_bytes::deserialize(deserializer)?;
79 Ok(WorkerAddress(Bytes::from(bytes)))
80 }
81}
82
83impl WorkerAddress {
84 pub fn from_encoded(bytes: impl Into<Bytes>) -> Self {
89 Self(bytes.into())
90 }
91
92 pub fn as_bytes(&self) -> &[u8] {
94 &self.0
95 }
96
97 pub fn to_bytes(&self) -> Bytes {
99 self.0.clone()
100 }
101
102 pub fn checksum(&self) -> u64 {
106 xxh3_64(self.as_bytes())
107 }
108
109 pub fn available_transports(&self) -> Result<Vec<TransportKey>, WorkerAddressError> {
130 let map = decode_to_map(self.as_bytes())?;
131 Ok(map.keys().cloned().map(TransportKey::from).collect())
132 }
133
134 pub fn get_entry(&self, key: impl AsRef<str>) -> Result<Option<Bytes>, WorkerAddressError> {
145 let map = decode_to_map(self.as_bytes())?;
146 Ok(map.get(key.as_ref()).cloned())
147 }
148}
149
150impl fmt::Debug for WorkerAddress {
151 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
152 f.debug_tuple("WorkerAddress")
153 .field(&format_args!(
154 "len={}, xxh3_64=0x{:016x}",
155 self.0.len(),
156 self.checksum()
157 ))
158 .finish()
159 }
160}
161
162impl fmt::Display for WorkerAddress {
163 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
164 write!(f, "WorkerAddress(xxh3_64=0x{:016x})", self.checksum())
165 }
166}
167
168fn decode_to_map(bytes: &[u8]) -> Result<HashMap<Arc<str>, Bytes>, WorkerAddressError> {
174 if bytes.is_empty() {
175 return Err(WorkerAddressError::InvalidFormat("Empty bytes".to_string()));
176 }
177
178 let decoded: HashMap<String, Vec<u8>> = rmp_serde::from_slice(bytes)?;
180
181 Ok(decoded
183 .into_iter()
184 .map(|(k, v)| (Arc::from(k.as_str()), Bytes::from(v)))
185 .collect())
186}
187
188#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
208pub struct PeerInfo {
209 pub instance_id: InstanceId,
211 pub worker_address: WorkerAddress,
213}
214
215impl PeerInfo {
216 pub fn new(instance_id: InstanceId, worker_address: WorkerAddress) -> Self {
218 Self {
219 instance_id,
220 worker_address,
221 }
222 }
223
224 pub fn instance_id(&self) -> InstanceId {
226 self.instance_id
227 }
228
229 pub fn worker_id(&self) -> WorkerId {
231 self.instance_id.worker_id()
232 }
233
234 pub fn worker_address(&self) -> &WorkerAddress {
236 &self.worker_address
237 }
238
239 pub fn address_checksum(&self) -> u64 {
241 self.worker_address.checksum()
242 }
243
244 pub fn into_address(self) -> WorkerAddress {
246 self.worker_address
247 }
248
249 pub fn into_parts(self) -> (InstanceId, WorkerAddress) {
251 (self.instance_id, self.worker_address)
252 }
253}
254
255#[cfg(test)]
256mod tests {
257 use super::*;
258
259 fn make_test_address(entries: &[(&str, &[u8])]) -> WorkerAddress {
261 let map: HashMap<String, Vec<u8>> = entries
262 .iter()
263 .map(|(k, v)| (k.to_string(), v.to_vec()))
264 .collect();
265 let encoded = rmp_serde::to_vec(&map).unwrap();
266 WorkerAddress::from_encoded(encoded)
267 }
268
269 #[test]
270 fn test_worker_address_from_encoded() {
271 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
272
273 let entry = address.get_entry("endpoint").unwrap();
275 assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
276 }
277
278 #[test]
279 fn test_worker_address_checksum() {
280 let address1 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
281 let address2 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
282 let address3 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:6666")]);
283
284 assert_eq!(address1.checksum(), address2.checksum());
286
287 assert_ne!(address1.checksum(), address3.checksum());
289 }
290
291 #[test]
292 fn test_worker_address_equality() {
293 let address1 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
294 let address2 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
295 let address3 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:6666")]);
296
297 assert_eq!(address1, address2);
298 assert_ne!(address1, address3);
299 }
300
301 #[test]
302 fn test_worker_address_debug() {
303 let address = make_test_address(&[("test", b"value")]);
304 let debug_str = format!("{:?}", address);
305
306 assert!(debug_str.contains("WorkerAddress"));
307 assert!(debug_str.contains("len="));
308 assert!(debug_str.contains("xxh3_64="));
309 }
310
311 #[test]
312 fn test_available_transports() {
313 let address = make_test_address(&[
314 ("tcp", b"tcp://127.0.0.1:5555"),
315 ("rdma", b"rdma://10.0.0.1:6666"),
316 ("udp", b"udp://127.0.0.1:7777"),
317 ]);
318
319 let transports = address.available_transports().unwrap();
320 assert_eq!(transports.len(), 3);
321 assert!(transports.contains(&TransportKey::from("tcp")));
322 assert!(transports.contains(&TransportKey::from("rdma")));
323 assert!(transports.contains(&TransportKey::from("udp")));
324 }
325
326 #[test]
327 fn test_available_transports_empty() {
328 let address = make_test_address(&[]);
329 let transports = address.available_transports().unwrap();
330 assert_eq!(transports.len(), 0);
331 }
332
333 #[test]
334 fn test_get_entry() {
335 let address =
336 make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555"), ("protocol", b"tcp")]);
337
338 assert_eq!(
340 address.get_entry("endpoint").unwrap().unwrap(),
341 Bytes::from_static(b"tcp://127.0.0.1:5555")
342 );
343
344 assert!(address.get_entry("nonexistent").unwrap().is_none());
346 }
347
348 #[test]
349 fn test_get_entry_with_transport_key() {
350 let address = make_test_address(&[
351 ("tcp", b"tcp://127.0.0.1:5555"),
352 ("rdma", b"rdma://10.0.0.1:6666"),
353 ]);
354
355 let tcp_key = TransportKey::from("tcp");
357 let result = address.get_entry(tcp_key).unwrap();
358 assert_eq!(result, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
359
360 let result = address.get_entry(String::from("rdma")).unwrap();
362 assert_eq!(result, Some(Bytes::from_static(b"rdma://10.0.0.1:6666")));
363 }
364
365 #[test]
366 fn test_peer_info_creation() {
367 let instance_id = InstanceId::new_v4();
368 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
369
370 let peer_info = PeerInfo::new(instance_id, address.clone());
371
372 assert_eq!(peer_info.instance_id(), instance_id);
373 assert_eq!(peer_info.worker_id(), instance_id.worker_id());
374 assert_eq!(peer_info.worker_address(), &address);
375 }
376
377 #[test]
378 fn test_peer_info_checksum() {
379 let instance_id = InstanceId::new_v4();
380 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
381
382 let peer_info = PeerInfo::new(instance_id, address.clone());
383
384 assert_eq!(peer_info.address_checksum(), address.checksum());
385 }
386
387 #[test]
388 fn test_peer_info_into_address() {
389 let instance_id = InstanceId::new_v4();
390 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
391
392 let peer_info = PeerInfo::new(instance_id, address.clone());
393 let extracted_address = peer_info.into_address();
394
395 assert_eq!(extracted_address, address);
396 }
397
398 #[test]
399 fn test_peer_info_into_parts() {
400 let instance_id = InstanceId::new_v4();
401 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
402
403 let peer_info = PeerInfo::new(instance_id, address.clone());
404 let (extracted_id, extracted_address) = peer_info.into_parts();
405
406 assert_eq!(extracted_id, instance_id);
407 assert_eq!(extracted_address, address);
408 }
409
410 #[test]
411 fn test_peer_info_serde() {
412 let instance_id = InstanceId::new_v4();
413 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
414 let peer_info = PeerInfo::new(instance_id, address);
415
416 let json = serde_json::to_string(&peer_info).unwrap();
418
419 let deserialized: PeerInfo = serde_json::from_str(&json).unwrap();
421
422 assert_eq!(deserialized.instance_id(), instance_id);
423 assert_eq!(deserialized.worker_id(), instance_id.worker_id());
424
425 let entry = deserialized.worker_address().get_entry("endpoint").unwrap();
427 assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
428 }
429}