Skip to main content

rns_core/resource/
advertisement.rs

1use alloc::vec;
2use alloc::vec::Vec;
3
4use super::types::{AdvFlags, ResourceError};
5use crate::constants::{RESOURCE_HASHMAP_MAX_LEN, RESOURCE_MAPHASH_LEN};
6use crate::msgpack::{self, Value};
7
8/// Resource advertisement data, corresponding to Python's ResourceAdvertisement.
9#[derive(Debug, Clone)]
10pub struct ResourceAdvertisement {
11    /// Transfer size (encrypted data size)
12    pub transfer_size: u64,
13    /// Total uncompressed data size (including metadata overhead)
14    pub data_size: u64,
15    /// Number of parts
16    pub num_parts: u64,
17    /// Resource hash (full 32 bytes)
18    pub resource_hash: Vec<u8>,
19    /// Random hash (4 bytes)
20    pub random_hash: Vec<u8>,
21    /// Original hash (first segment, 32 bytes)
22    pub original_hash: Vec<u8>,
23    /// Hashmap segment (concatenated 4-byte part hashes)
24    pub hashmap: Vec<u8>,
25    /// Flags byte
26    pub flags: AdvFlags,
27    /// Segment index (1-based)
28    pub segment_index: u64,
29    /// Total segments
30    pub total_segments: u64,
31    /// Request ID (optional)
32    pub request_id: Option<Vec<u8>>,
33}
34
35impl ResourceAdvertisement {
36    /// Pack the advertisement to msgpack bytes.
37    /// `segment` controls which hashmap segment to include (0-based).
38    pub fn pack(&self, segment: usize) -> Vec<u8> {
39        let hashmap_start = segment * RESOURCE_HASHMAP_MAX_LEN * RESOURCE_MAPHASH_LEN;
40        let max_end = (segment + 1) * RESOURCE_HASHMAP_MAX_LEN * RESOURCE_MAPHASH_LEN;
41        let hashmap_end = core::cmp::min(max_end, self.hashmap.len());
42        let hashmap_segment = if hashmap_start < self.hashmap.len() {
43            &self.hashmap[hashmap_start..hashmap_end]
44        } else {
45            &[]
46        };
47
48        let q_value = match &self.request_id {
49            Some(id) => Value::Bin(id.clone()),
50            None => Value::Nil,
51        };
52
53        // Match Python's key order: t, d, n, h, r, o, i, l, q, f, m
54        let entries: Vec<(&str, Value)> = vec![
55            ("t", Value::UInt(self.transfer_size)),
56            ("d", Value::UInt(self.data_size)),
57            ("n", Value::UInt(self.num_parts)),
58            ("h", Value::Bin(self.resource_hash.clone())),
59            ("r", Value::Bin(self.random_hash.clone())),
60            ("o", Value::Bin(self.original_hash.clone())),
61            ("i", Value::UInt(self.segment_index)),
62            ("l", Value::UInt(self.total_segments)),
63            ("q", q_value),
64            ("f", Value::UInt(self.flags.to_byte() as u64)),
65            ("m", Value::Bin(hashmap_segment.to_vec())),
66        ];
67
68        msgpack::pack_str_map(&entries)
69    }
70
71    /// Unpack an advertisement from msgpack bytes.
72    pub fn unpack(data: &[u8]) -> Result<Self, ResourceError> {
73        let value = msgpack::unpack_exact(data).map_err(|_| ResourceError::InvalidAdvertisement)?;
74
75        let t = value
76            .map_get("t")
77            .and_then(|v| v.as_uint())
78            .ok_or(ResourceError::InvalidAdvertisement)?;
79        let d = value
80            .map_get("d")
81            .and_then(|v| v.as_uint())
82            .ok_or(ResourceError::InvalidAdvertisement)?;
83        let n = value
84            .map_get("n")
85            .and_then(|v| v.as_uint())
86            .ok_or(ResourceError::InvalidAdvertisement)?;
87        let h = value
88            .map_get("h")
89            .and_then(|v| v.as_bin())
90            .ok_or(ResourceError::InvalidAdvertisement)?
91            .to_vec();
92        let r = value
93            .map_get("r")
94            .and_then(|v| v.as_bin())
95            .ok_or(ResourceError::InvalidAdvertisement)?
96            .to_vec();
97        let o = value
98            .map_get("o")
99            .and_then(|v| v.as_bin())
100            .ok_or(ResourceError::InvalidAdvertisement)?
101            .to_vec();
102        let m = value
103            .map_get("m")
104            .and_then(|v| v.as_bin())
105            .ok_or(ResourceError::InvalidAdvertisement)?
106            .to_vec();
107        let f = value
108            .map_get("f")
109            .and_then(|v| v.as_uint())
110            .ok_or(ResourceError::InvalidAdvertisement)? as u8;
111        let i = value
112            .map_get("i")
113            .and_then(|v| v.as_uint())
114            .ok_or(ResourceError::InvalidAdvertisement)?;
115        let l = value
116            .map_get("l")
117            .and_then(|v| v.as_uint())
118            .ok_or(ResourceError::InvalidAdvertisement)?;
119
120        let q_val = value
121            .map_get("q")
122            .ok_or(ResourceError::InvalidAdvertisement)?;
123        let request_id = if q_val.is_nil() {
124            None
125        } else {
126            Some(
127                q_val
128                    .as_bin()
129                    .ok_or(ResourceError::InvalidAdvertisement)?
130                    .to_vec(),
131            )
132        };
133
134        if t > (crate::constants::RESOURCE_MAX_EFFICIENT_SIZE * 3) as u64 {
135            return Err(ResourceError::InvalidAdvertisement);
136        }
137
138        Ok(ResourceAdvertisement {
139            transfer_size: t,
140            data_size: d,
141            num_parts: n,
142            resource_hash: h,
143            random_hash: r,
144            original_hash: o,
145            hashmap: m,
146            flags: AdvFlags::from_byte(f),
147            segment_index: i,
148            total_segments: l,
149            request_id,
150        })
151    }
152
153    /// Check if this advertisement is a request.
154    pub fn is_request(&self) -> bool {
155        self.request_id.is_some() && self.flags.is_request
156    }
157
158    /// Check if this advertisement is a response.
159    pub fn is_response(&self) -> bool {
160        self.request_id.is_some() && self.flags.is_response
161    }
162
163    /// Get the number of hashmap segments needed.
164    pub fn hashmap_segments(&self) -> usize {
165        let total_hashes = self.num_parts as usize;
166        if total_hashes == 0 {
167            return 1;
168        }
169        total_hashes.div_ceil(RESOURCE_HASHMAP_MAX_LEN)
170    }
171}
172
173#[cfg(test)]
174mod tests {
175    use super::*;
176
177    fn make_adv(flags: AdvFlags) -> ResourceAdvertisement {
178        ResourceAdvertisement {
179            transfer_size: 1000,
180            data_size: 950,
181            num_parts: 3,
182            resource_hash: vec![0x11; 32],
183            random_hash: vec![0xAA, 0xBB, 0xCC, 0xDD],
184            original_hash: vec![0x22; 32],
185            hashmap: vec![
186                0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x0B, 0x0C,
187            ],
188            flags,
189            segment_index: 1,
190            total_segments: 1,
191            request_id: None,
192        }
193    }
194
195    #[test]
196    fn test_pack_unpack_roundtrip() {
197        let flags = AdvFlags {
198            encrypted: true,
199            compressed: false,
200            split: false,
201            is_request: false,
202            is_response: false,
203            has_metadata: false,
204        };
205        let adv = make_adv(flags);
206        let packed = adv.pack(0);
207        let unpacked = ResourceAdvertisement::unpack(&packed).unwrap();
208
209        assert_eq!(unpacked.transfer_size, 1000);
210        assert_eq!(unpacked.data_size, 950);
211        assert_eq!(unpacked.num_parts, 3);
212        assert_eq!(unpacked.resource_hash, vec![0x11; 32]);
213        assert_eq!(unpacked.random_hash, vec![0xAA, 0xBB, 0xCC, 0xDD]);
214        assert_eq!(unpacked.original_hash, vec![0x22; 32]);
215        assert_eq!(unpacked.flags, flags);
216        assert_eq!(unpacked.segment_index, 1);
217        assert_eq!(unpacked.total_segments, 1);
218        assert!(unpacked.request_id.is_none());
219    }
220
221    #[test]
222    fn test_flags_encrypted_compressed() {
223        let flags = AdvFlags {
224            encrypted: true,
225            compressed: true,
226            split: false,
227            is_request: false,
228            is_response: false,
229            has_metadata: false,
230        };
231        let adv = make_adv(flags);
232        let packed = adv.pack(0);
233        let unpacked = ResourceAdvertisement::unpack(&packed).unwrap();
234        assert!(unpacked.flags.encrypted);
235        assert!(unpacked.flags.compressed);
236        assert!(!unpacked.flags.split);
237    }
238
239    #[test]
240    fn test_flags_with_metadata() {
241        let flags = AdvFlags {
242            encrypted: true,
243            compressed: false,
244            split: false,
245            is_request: false,
246            is_response: false,
247            has_metadata: true,
248        };
249        let adv = make_adv(flags);
250        let packed = adv.pack(0);
251        let unpacked = ResourceAdvertisement::unpack(&packed).unwrap();
252        assert!(unpacked.flags.has_metadata);
253    }
254
255    #[test]
256    fn test_multi_segment() {
257        let flags = AdvFlags {
258            encrypted: true,
259            compressed: false,
260            split: true,
261            is_request: false,
262            is_response: false,
263            has_metadata: false,
264        };
265        let mut adv = make_adv(flags);
266        adv.segment_index = 2;
267        adv.total_segments = 5;
268        let packed = adv.pack(0);
269        let unpacked = ResourceAdvertisement::unpack(&packed).unwrap();
270        assert!(unpacked.flags.split);
271        assert_eq!(unpacked.segment_index, 2);
272        assert_eq!(unpacked.total_segments, 5);
273    }
274
275    #[test]
276    fn test_with_request_id() {
277        let flags = AdvFlags {
278            encrypted: true,
279            compressed: false,
280            split: false,
281            is_request: true,
282            is_response: false,
283            has_metadata: false,
284        };
285        let mut adv = make_adv(flags);
286        adv.request_id = Some(vec![0xDE, 0xAD, 0xBE, 0xEF]);
287        let packed = adv.pack(0);
288        let unpacked = ResourceAdvertisement::unpack(&packed).unwrap();
289        assert!(unpacked.is_request());
290        assert!(!unpacked.is_response());
291        assert_eq!(unpacked.request_id, Some(vec![0xDE, 0xAD, 0xBE, 0xEF]));
292    }
293
294    #[test]
295    fn test_is_response() {
296        let flags = AdvFlags {
297            encrypted: true,
298            compressed: false,
299            split: false,
300            is_request: false,
301            is_response: true,
302            has_metadata: false,
303        };
304        let mut adv = make_adv(flags);
305        adv.request_id = Some(vec![0x42; 16]);
306        assert!(adv.is_response());
307        assert!(!adv.is_request());
308    }
309
310    #[test]
311    fn test_nil_request_id() {
312        let flags = AdvFlags {
313            encrypted: true,
314            compressed: false,
315            split: false,
316            is_request: false,
317            is_response: false,
318            has_metadata: false,
319        };
320        let adv = make_adv(flags);
321        let packed = adv.pack(0);
322        let unpacked = ResourceAdvertisement::unpack(&packed).unwrap();
323        assert!(unpacked.request_id.is_none());
324        assert!(!unpacked.is_request());
325        assert!(!unpacked.is_response());
326    }
327
328    #[test]
329    fn test_hashmap_segmentation() {
330        // Create a large hashmap with > HASHMAP_MAX_LEN(74) hashes
331        let num_hashes = 100;
332        let hashmap: Vec<u8> = (0..num_hashes).flat_map(|i| vec![i as u8; 4]).collect();
333
334        let flags = AdvFlags {
335            encrypted: true,
336            compressed: false,
337            split: false,
338            is_request: false,
339            is_response: false,
340            has_metadata: false,
341        };
342        let adv = ResourceAdvertisement {
343            transfer_size: 50000,
344            data_size: 48000,
345            num_parts: num_hashes,
346            resource_hash: vec![0x11; 32],
347            random_hash: vec![0xAA; 4],
348            original_hash: vec![0x22; 32],
349            hashmap: hashmap.clone(),
350            flags,
351            segment_index: 1,
352            total_segments: 1,
353            request_id: None,
354        };
355
356        // Segment 0: first 74 hashes = 296 bytes
357        let packed0 = adv.pack(0);
358        let unpacked0 = ResourceAdvertisement::unpack(&packed0).unwrap();
359        assert_eq!(unpacked0.hashmap.len(), 74 * 4);
360
361        // Segment 1: remaining 26 hashes = 104 bytes
362        let packed1 = adv.pack(1);
363        let unpacked1 = ResourceAdvertisement::unpack(&packed1).unwrap();
364        assert_eq!(unpacked1.hashmap.len(), 26 * 4);
365    }
366
367    #[test]
368    fn test_hashmap_segments_count() {
369        let flags = AdvFlags {
370            encrypted: true,
371            compressed: false,
372            split: false,
373            is_request: false,
374            is_response: false,
375            has_metadata: false,
376        };
377        let mut adv = make_adv(flags);
378
379        adv.num_parts = 74; // exactly HASHMAP_MAX_LEN
380        assert_eq!(adv.hashmap_segments(), 1);
381
382        adv.num_parts = 75;
383        assert_eq!(adv.hashmap_segments(), 2);
384
385        adv.num_parts = 148;
386        assert_eq!(adv.hashmap_segments(), 2);
387
388        adv.num_parts = 149;
389        assert_eq!(adv.hashmap_segments(), 3);
390    }
391
392    #[test]
393    fn test_unpack_invalid_data() {
394        assert!(ResourceAdvertisement::unpack(&[]).is_err());
395        assert!(ResourceAdvertisement::unpack(&[0xc0]).is_err()); // nil
396        assert!(ResourceAdvertisement::unpack(&[0x01, 0x02]).is_err()); // not a map
397    }
398
399    #[test]
400    fn test_unpack_accepts_exact_upstream_transfer_limit() {
401        let flags = AdvFlags {
402            encrypted: true,
403            compressed: false,
404            split: false,
405            is_request: false,
406            is_response: false,
407            has_metadata: false,
408        };
409        let mut adv = make_adv(flags);
410        adv.transfer_size = (crate::constants::RESOURCE_MAX_EFFICIENT_SIZE * 3) as u64;
411
412        let unpacked = ResourceAdvertisement::unpack(&adv.pack(0)).unwrap();
413        assert_eq!(unpacked.transfer_size, adv.transfer_size);
414    }
415
416    #[test]
417    fn test_unpack_rejects_transfer_above_upstream_limit() {
418        let flags = AdvFlags {
419            encrypted: true,
420            compressed: false,
421            split: false,
422            is_request: false,
423            is_response: false,
424            has_metadata: false,
425        };
426        let mut adv = make_adv(flags);
427        adv.transfer_size = (crate::constants::RESOURCE_MAX_EFFICIENT_SIZE * 3 + 1) as u64;
428
429        assert_eq!(
430            ResourceAdvertisement::unpack(&adv.pack(0)).unwrap_err(),
431            ResourceError::InvalidAdvertisement
432        );
433    }
434}