1use super::identity::{InstanceId, WorkerId};
14use super::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 empty() -> Self {
98 let empty: HashMap<String, Vec<u8>> = HashMap::new();
99 let encoded = rmp_serde::to_vec(&empty).expect("encoding empty HashMap cannot fail");
100 Self(Bytes::from(encoded))
101 }
102
103 pub fn as_bytes(&self) -> &[u8] {
105 &self.0
106 }
107
108 pub fn to_bytes(&self) -> Bytes {
110 self.0.clone()
111 }
112
113 pub fn checksum(&self) -> u64 {
117 xxh3_64(self.as_bytes())
118 }
119
120 pub fn available_transports(&self) -> Result<Vec<TransportKey>, WorkerAddressError> {
141 let map = decode_to_map(self.as_bytes())?;
142 Ok(map.keys().cloned().map(TransportKey::from).collect())
143 }
144
145 pub fn get_entry(&self, key: impl AsRef<str>) -> Result<Option<Bytes>, WorkerAddressError> {
156 let map = decode_to_map(self.as_bytes())?;
157 Ok(map.get(key.as_ref()).cloned())
158 }
159}
160
161impl fmt::Debug for WorkerAddress {
162 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
163 f.debug_tuple("WorkerAddress")
164 .field(&format_args!(
165 "len={}, xxh3_64=0x{:016x}",
166 self.0.len(),
167 self.checksum()
168 ))
169 .finish()
170 }
171}
172
173impl fmt::Display for WorkerAddress {
174 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
175 write!(f, "WorkerAddress(xxh3_64=0x{:016x})", self.checksum())
176 }
177}
178
179fn decode_to_map(bytes: &[u8]) -> Result<HashMap<Arc<str>, Bytes>, WorkerAddressError> {
185 if bytes.is_empty() {
186 return Err(WorkerAddressError::InvalidFormat("Empty bytes".to_string()));
187 }
188
189 let decoded: HashMap<String, Vec<u8>> = rmp_serde::from_slice(bytes)?;
191
192 Ok(decoded
194 .into_iter()
195 .map(|(k, v)| (Arc::from(k.as_str()), Bytes::from(v)))
196 .collect())
197}
198
199#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
219pub struct PeerInfo {
220 pub instance_id: InstanceId,
222 pub worker_address: WorkerAddress,
224}
225
226impl PeerInfo {
227 pub fn new(instance_id: InstanceId, worker_address: WorkerAddress) -> Self {
229 Self {
230 instance_id,
231 worker_address,
232 }
233 }
234
235 pub fn instance_id(&self) -> InstanceId {
237 self.instance_id
238 }
239
240 pub fn worker_id(&self) -> WorkerId {
242 self.instance_id.worker_id()
243 }
244
245 pub fn worker_address(&self) -> &WorkerAddress {
247 &self.worker_address
248 }
249
250 pub fn address_checksum(&self) -> u64 {
252 self.worker_address.checksum()
253 }
254
255 pub fn into_address(self) -> WorkerAddress {
257 self.worker_address
258 }
259
260 pub fn into_parts(self) -> (InstanceId, WorkerAddress) {
262 (self.instance_id, self.worker_address)
263 }
264}
265
266#[cfg(test)]
267mod tests {
268 use super::*;
269
270 fn make_test_address(entries: &[(&str, &[u8])]) -> WorkerAddress {
272 let map: HashMap<String, Vec<u8>> = entries
273 .iter()
274 .map(|(k, v)| (k.to_string(), v.to_vec()))
275 .collect();
276 let encoded = rmp_serde::to_vec(&map).unwrap();
277 WorkerAddress::from_encoded(encoded)
278 }
279
280 #[test]
281 fn test_worker_address_from_encoded() {
282 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
283
284 let entry = address.get_entry("endpoint").unwrap();
286 assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
287 }
288
289 #[test]
290 fn test_worker_address_checksum() {
291 let address1 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
292 let address2 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
293 let address3 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:6666")]);
294
295 assert_eq!(address1.checksum(), address2.checksum());
297
298 assert_ne!(address1.checksum(), address3.checksum());
300 }
301
302 #[test]
303 fn test_worker_address_equality() {
304 let address1 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
305 let address2 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
306 let address3 = make_test_address(&[("endpoint", b"tcp://127.0.0.1:6666")]);
307
308 assert_eq!(address1, address2);
309 assert_ne!(address1, address3);
310 }
311
312 #[test]
313 fn test_worker_address_debug() {
314 let address = make_test_address(&[("test", b"value")]);
315 let debug_str = format!("{:?}", address);
316
317 assert!(debug_str.contains("WorkerAddress"));
318 assert!(debug_str.contains("len="));
319 assert!(debug_str.contains("xxh3_64="));
320 }
321
322 #[test]
323 fn test_available_transports() {
324 let address = make_test_address(&[
325 ("tcp", b"tcp://127.0.0.1:5555"),
326 ("rdma", b"rdma://10.0.0.1:6666"),
327 ("udp", b"udp://127.0.0.1:7777"),
328 ]);
329
330 let transports = address.available_transports().unwrap();
331 assert_eq!(transports.len(), 3);
332 assert!(transports.contains(&TransportKey::from("tcp")));
333 assert!(transports.contains(&TransportKey::from("rdma")));
334 assert!(transports.contains(&TransportKey::from("udp")));
335 }
336
337 #[test]
338 fn test_available_transports_empty() {
339 let address = make_test_address(&[]);
340 let transports = address.available_transports().unwrap();
341 assert_eq!(transports.len(), 0);
342 }
343
344 #[test]
345 fn test_get_entry() {
346 let address =
347 make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555"), ("protocol", b"tcp")]);
348
349 assert_eq!(
351 address.get_entry("endpoint").unwrap().unwrap(),
352 Bytes::from_static(b"tcp://127.0.0.1:5555")
353 );
354
355 assert!(address.get_entry("nonexistent").unwrap().is_none());
357 }
358
359 #[test]
360 fn test_get_entry_with_transport_key() {
361 let address = make_test_address(&[
362 ("tcp", b"tcp://127.0.0.1:5555"),
363 ("rdma", b"rdma://10.0.0.1:6666"),
364 ]);
365
366 let tcp_key = TransportKey::from("tcp");
368 let result = address.get_entry(tcp_key).unwrap();
369 assert_eq!(result, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
370
371 let result = address.get_entry(String::from("rdma")).unwrap();
373 assert_eq!(result, Some(Bytes::from_static(b"rdma://10.0.0.1:6666")));
374 }
375
376 #[test]
377 fn test_peer_info_creation() {
378 let instance_id = InstanceId::new_v4();
379 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
380
381 let peer_info = PeerInfo::new(instance_id, address.clone());
382
383 assert_eq!(peer_info.instance_id(), instance_id);
384 assert_eq!(peer_info.worker_id(), instance_id.worker_id());
385 assert_eq!(peer_info.worker_address(), &address);
386 }
387
388 #[test]
389 fn test_peer_info_checksum() {
390 let instance_id = InstanceId::new_v4();
391 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
392
393 let peer_info = PeerInfo::new(instance_id, address.clone());
394
395 assert_eq!(peer_info.address_checksum(), address.checksum());
396 }
397
398 #[test]
399 fn test_peer_info_into_address() {
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_address = peer_info.into_address();
405
406 assert_eq!(extracted_address, address);
407 }
408
409 #[test]
410 fn test_peer_info_into_parts() {
411 let instance_id = InstanceId::new_v4();
412 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
413
414 let peer_info = PeerInfo::new(instance_id, address.clone());
415 let (extracted_id, extracted_address) = peer_info.into_parts();
416
417 assert_eq!(extracted_id, instance_id);
418 assert_eq!(extracted_address, address);
419 }
420
421 #[test]
422 fn test_peer_info_serde() {
423 let instance_id = InstanceId::new_v4();
424 let address = make_test_address(&[("endpoint", b"tcp://127.0.0.1:5555")]);
425 let peer_info = PeerInfo::new(instance_id, address);
426
427 let json = serde_json::to_string(&peer_info).unwrap();
429
430 let deserialized: PeerInfo = serde_json::from_str(&json).unwrap();
432
433 assert_eq!(deserialized.instance_id(), instance_id);
434 assert_eq!(deserialized.worker_id(), instance_id.worker_id());
435
436 let entry = deserialized.worker_address().get_entry("endpoint").unwrap();
438 assert_eq!(entry, Some(Bytes::from_static(b"tcp://127.0.0.1:5555")));
439 }
440
441 #[test]
442 fn test_worker_address_empty() {
443 let empty = WorkerAddress::empty();
444 assert!(empty.get_entry("anything").unwrap().is_none());
446 assert!(empty.available_transports().unwrap().is_empty());
448 assert_eq!(empty, WorkerAddress::empty());
450 assert_eq!(empty.checksum(), WorkerAddress::empty().checksum());
451 }
452}