1use objects::{
3 object::{AnnotatedTag, State, StateAttachment},
4 store::ObjectStore,
5};
6
7use crate::{ObjectData, ObjectId, ObjectRequest, ObjectType, ProtocolError, Result};
8
9pub const MAX_RECEIVED_REDACTIONS_BLOB_SIZE: u64 = 64 * 1024 * 1024;
17
18pub const MAX_RECEIVED_STATE_VISIBILITY_BLOB_SIZE: u64 = 64 * 1024 * 1024;
25
26const PULL_DECODE_ENVELOPE_HEADROOM: u64 = 1024 * 1024;
37
38const fn max_u64(a: u64, b: u64) -> u64 {
39 if a > b { a } else { b }
40}
41
42pub const MAX_PULL_FRAME_MESSAGE_SIZE: usize = (max_u64(
58 MAX_RECEIVED_REDACTIONS_BLOB_SIZE,
59 MAX_RECEIVED_STATE_VISIBILITY_BLOB_SIZE,
60) + PULL_DECODE_ENVELOPE_HEADROOM) as usize;
61
62pub fn check_received_transfer_blob_size(
71 blob_len: usize,
72 max_bytes: u64,
73 kind: &str,
74) -> Result<()> {
75 let len = u64::try_from(blob_len).map_err(|_| {
76 ProtocolError::InvalidState(format!("{kind} blob length does not fit in u64"))
77 })?;
78 if len > max_bytes {
79 return Err(ProtocolError::InvalidState(format!(
80 "{kind} blob exceeds receive size limit: {len} bytes (max {max_bytes})"
81 )));
82 }
83 Ok(())
84}
85
86pub fn admit_declared_received_len(declared: u64, max_bytes: u64, kind: &str) -> Result<usize> {
92 if declared > max_bytes {
93 return Err(ProtocolError::InvalidState(format!(
94 "{kind} exceeds receive size limit: {declared} bytes (max {max_bytes})"
95 )));
96 }
97 usize::try_from(declared)
98 .map_err(|_| ProtocolError::InvalidState(format!("{kind} exceeds this platform")))
99}
100
101#[allow(dead_code)]
102pub fn chunk_count(object_size: usize, chunk_size: usize) -> usize {
103 if object_size == 0 || chunk_size == 0 {
104 return 0;
105 }
106 object_size.div_ceil(chunk_size)
107}
108
109#[allow(dead_code)]
110pub fn chunk_bounds(
111 object_size: usize,
112 chunk_size: usize,
113 chunk_index: usize,
114) -> Option<(usize, usize)> {
115 if chunk_size == 0 {
116 return None;
117 }
118
119 let start = chunk_index.checked_mul(chunk_size)?;
120 if start >= object_size {
121 return None;
122 }
123 let end = (start + chunk_size).min(object_size);
124 Some((start, end - start))
125}
126
127#[allow(dead_code)]
128pub fn chunk_offset(chunk_index: usize, chunk_size: usize) -> Option<usize> {
129 chunk_index.checked_mul(chunk_size)
130}
131
132pub fn load_requested_object(store: &impl ObjectStore, req: &ObjectRequest) -> Result<ObjectData> {
133 let (obj_type, data) = match &req.id {
140 ObjectId::Hash(hash) => {
141 if let Some(blob) = store.get_blob(hash)? {
142 (ObjectType::Blob, blob.content().to_vec())
143 } else if let Some(data) = store.get_tree_serialized(hash)? {
144 (ObjectType::Tree, data)
145 } else {
146 return Err(ProtocolError::ObjectNotFound(hash.to_hex()));
147 }
148 }
149 ObjectId::StateId(state_id) => {
150 let state = store
151 .get_state(state_id)?
152 .ok_or_else(|| ProtocolError::ObjectNotFound(state_id.to_string()))?;
153 (
154 ObjectType::State,
155 state.encode_current_msgpack().map_err(object_codec_error)?,
156 )
157 }
158 ObjectId::StateAttachment { state, id, kind: _ } => {
159 let attachment = store
160 .get_state_attachment(state, id)?
161 .ok_or_else(|| ProtocolError::ObjectNotFound(id.to_string()))?;
162 (
163 ObjectType::StateAttachment,
164 attachment
165 .encode_current_msgpack()
166 .map_err(object_codec_error)?,
167 )
168 }
169 };
170
171 Ok(ObjectData {
172 id: req.id.clone(),
173 obj_type,
174 data,
175 is_delta: false,
176 })
177}
178
179pub fn load_object_data(
180 store: &impl ObjectStore,
181 id: &ObjectId,
182 obj_type: ObjectType,
183) -> Result<ObjectData> {
184 let data = match (id, obj_type) {
185 (ObjectId::Hash(hash), ObjectType::Blob) => store
186 .get_blob(hash)?
187 .ok_or_else(|| ProtocolError::ObjectNotFound(hash.to_hex()))?
188 .content()
189 .to_vec(),
190 (ObjectId::Hash(hash), ObjectType::Tree) => store
191 .get_tree_serialized(hash)?
192 .ok_or_else(|| ProtocolError::ObjectNotFound(hash.to_hex()))?,
193 (ObjectId::Hash(hash), ObjectType::AnnotatedTag) => store
194 .get_annotated_tag(hash)?
195 .ok_or_else(|| ProtocolError::ObjectNotFound(hash.to_hex()))?
196 .encode_current_msgpack(),
197 (ObjectId::StateId(state_id), ObjectType::State) => {
198 let state = store
199 .get_state(state_id)?
200 .ok_or_else(|| ProtocolError::ObjectNotFound(state_id.to_string()))?;
201 state.encode_current_msgpack().map_err(object_codec_error)?
202 }
203 (ObjectId::Hash(hash), ObjectType::Redaction) => store
204 .get_redactions_bytes_for_blob(hash)?
205 .ok_or_else(|| ProtocolError::ObjectNotFound(hash.to_hex()))?,
206 (ObjectId::Hash(hash), ObjectType::Purge) => store
207 .get_redactions_bytes_for_blob(hash)?
208 .ok_or_else(|| ProtocolError::ObjectNotFound(hash.to_hex()))?,
209 (ObjectId::StateId(state_id), ObjectType::StateVisibility) => store
210 .get_state_visibility_bytes_for_state(state_id)?
211 .ok_or_else(|| ProtocolError::ObjectNotFound(state_id.to_string_full()))?,
212 (ObjectId::StateAttachment { state, id, kind: _ }, ObjectType::StateAttachment) => {
213 let attachment = store
214 .get_state_attachment(state, id)?
215 .ok_or_else(|| ProtocolError::ObjectNotFound(id.to_string()))?;
216 attachment
217 .encode_current_msgpack()
218 .map_err(object_codec_error)?
219 }
220 (ObjectId::Hash(_), ObjectType::KeyBinding) => {
221 return Err(ProtocolError::InvalidState(
222 "KeyBinding registry objects must be constructed with encode_key_binding_registry"
223 .to_string(),
224 ));
225 }
226 _ => {
227 return Err(ProtocolError::InvalidState(
228 "object id/type mismatch".to_string(),
229 ));
230 }
231 };
232
233 Ok(ObjectData {
234 id: id.clone(),
235 obj_type,
236 data,
237 is_delta: false,
238 })
239}
240
241pub fn store_received_object(store: &impl ObjectStore, data: &ObjectData) -> Result<()> {
242 match (&data.id, data.obj_type) {
243 (ObjectId::Hash(hash), ObjectType::Blob) => {
244 store.put_blob_bytes_with_hash(&data.data, *hash)?;
245 }
246 (ObjectId::Hash(hash), ObjectType::Tree) => {
247 store
248 .put_tree_serialized(&data.data, *hash)
249 .map_err(|error| match error {
250 objects::error::HeddleError::Corruption { .. }
251 | objects::error::HeddleError::TreeStream(
252 objects::object::TreeStreamError::HashMismatch { .. },
253 ) => ProtocolError::InvalidState("tree hash mismatch".to_string()),
254 error => ProtocolError::InvalidState(format!("invalid tree object: {error}")),
255 })?;
256 }
257 (ObjectId::Hash(hash), ObjectType::AnnotatedTag) => {
258 let tag = AnnotatedTag::decode_current_msgpack(&data.data)
259 .map_err(|error| ProtocolError::InvalidState(error.to_string()))?;
260 if tag.hash() != *hash {
261 return Err(ProtocolError::InvalidState(
262 "annotated tag hash mismatch".to_string(),
263 ));
264 }
265 store.put_annotated_tag(&tag)?;
266 }
267 (ObjectId::StateId(state_id), ObjectType::State) => {
268 let state = State::decode_current_msgpack(&data.data).map_err(object_codec_error)?;
269 if state.id() != *state_id {
270 return Err(ProtocolError::InvalidState(format!(
271 "StateId mismatch: expected {state_id}, computed {}",
272 state.id()
273 )));
274 }
275 store.put_state_serialized(&data.data, *state_id)?;
276 }
277 (ObjectId::StateAttachment { state, id, kind }, ObjectType::StateAttachment) => {
278 let attachment =
279 StateAttachment::decode_current_msgpack(&data.data).map_err(object_codec_error)?;
280 if attachment.state_id != *state || attachment.id() != *id {
281 return Err(ProtocolError::InvalidState(
282 "state attachment id mismatch".to_string(),
283 ));
284 }
285 let body_kind = attachment.body.kind();
290 if *kind != body_kind {
291 return Err(ProtocolError::InvalidState(format!(
292 "state attachment kind mismatch: descriptor {kind:?}, body {body_kind:?}"
293 )));
294 }
295 store.put_state_attachment(&attachment)?;
296 }
297 (_, ObjectType::Redaction) => {
298 return Err(ProtocolError::InvalidState(
303 "Redaction objects must be persisted via Repository::accept_wire_redactions, \
304 not store_received_object — signature verification is required"
305 .to_string(),
306 ));
307 }
308 (_, ObjectType::Purge) => {
309 return Err(ProtocolError::InvalidState(
310 "Purge objects must be persisted via Repository::accept_wire_purge, not store_received_object — owner authorization is required"
311 .to_string(),
312 ));
313 }
314 (_, ObjectType::StateVisibility) => {
315 return Err(ProtocolError::InvalidState(
319 "StateVisibility objects must be persisted via Repository::accept_wire_state_visibility, \
320 not store_received_object — sidecar validation is required"
321 .to_string(),
322 ));
323 }
324 (_, ObjectType::KeyBinding) => {
325 return Err(ProtocolError::InvalidState(
326 "KeyBinding registry objects must be decoded and verified with decode_key_binding_registry"
327 .to_string(),
328 ));
329 }
330 _ => {
331 return Err(ProtocolError::InvalidState(
332 "object id/type mismatch".to_string(),
333 ));
334 }
335 }
336
337 Ok(())
338}
339
340fn object_codec_error(error: objects::error::HeddleError) -> ProtocolError {
341 ProtocolError::Serialization(error.to_string())
342}
343
344#[cfg(test)]
345mod tests {
346 use objects::{
347 object::{
348 Attribution, Blob, ContentHash, Principal, State, StateAttachment, StateAttachmentBody,
349 TREE_BLOCK_ENCODING_VERSION, TREE_HEADER_LEN, Tree, TreeEntry,
350 },
351 store::{FsStore, ObjectStore},
352 };
353 use tempfile::TempDir;
354
355 use super::*;
356
357 fn create_test_store() -> (TempDir, FsStore) {
358 let temp = TempDir::new().unwrap();
359 let store = FsStore::new(temp.path().join(".heddle"));
360 store.init().unwrap();
361 (temp, store)
362 }
363
364 fn test_attribution() -> Attribution {
365 Attribution::human(Principal::new("Wire Tester", "wire@example.com"))
366 }
367
368 #[test]
369 fn primary_objects_roundtrip_through_wire_data() {
370 let (_source_temp, source) = create_test_store();
371 let (_dest_temp, dest) = create_test_store();
372
373 let blob = Blob::from("wire transfer blob\n");
374 let blob_hash = source.put_blob(&blob).unwrap();
375 let tree = Tree::from_entries(vec![TreeEntry::file("lib.rs", blob_hash, false).unwrap()]);
376 let tree_hash = source.put_tree(&tree).unwrap();
377 let state = State::new(tree_hash, Vec::new(), test_attribution())
378 .with_intent("exercise wire transfer");
379 source.put_state(&state).unwrap();
380
381 let blob_data = load_requested_object(
382 &source,
383 &ObjectRequest {
384 id: ObjectId::Hash(blob_hash),
385 have_base: None,
386 },
387 )
388 .unwrap();
389 assert_eq!(blob_data.obj_type, ObjectType::Blob);
390 assert_eq!(blob_data.data, blob.content());
391 store_received_object(&dest, &blob_data).unwrap();
392 assert_eq!(
393 dest.get_blob(&blob_hash).unwrap().unwrap().content(),
394 blob.content()
395 );
396
397 let tree_data = load_requested_object(
398 &source,
399 &ObjectRequest {
400 id: ObjectId::Hash(tree_hash),
401 have_base: None,
402 },
403 )
404 .unwrap();
405 assert_eq!(tree_data.obj_type, ObjectType::Tree);
406 assert_eq!(
407 objects::store::codec::decode_tree_serialized_with_key(
408 &tree_data.data,
409 tree_hash,
410 None,
411 )
412 .unwrap(),
413 tree,
414 );
415 store_received_object(&dest, &tree_data).unwrap();
416 assert_eq!(dest.get_tree(&tree_hash).unwrap().unwrap(), tree);
417
418 let state_data = load_requested_object(
419 &source,
420 &ObjectRequest {
421 id: ObjectId::StateId(state.state_id),
422 have_base: None,
423 },
424 )
425 .unwrap();
426 assert_eq!(state_data.obj_type, ObjectType::State);
427 assert_eq!(
428 objects::store::codec::decode_state(&state_data.data).unwrap(),
429 state
430 );
431 store_received_object(&dest, &state_data).unwrap();
432 assert_eq!(
433 dest.get_state(&state.state_id).unwrap().unwrap().state_id,
434 state.state_id
435 );
436 }
437
438 #[test]
439 fn load_object_data_reports_missing_and_id_type_mismatch_errors() {
440 let (_temp, store) = create_test_store();
441 let missing_hash = ContentHash::from_bytes([7; 32]);
442 let missing_state = objects::object::StateId::from_bytes([9; 32]);
443
444 let missing = load_requested_object(
445 &store,
446 &ObjectRequest {
447 id: ObjectId::Hash(missing_hash),
448 have_base: None,
449 },
450 )
451 .unwrap_err();
452 assert!(
453 matches!(missing, ProtocolError::ObjectNotFound(id) if id == missing_hash.to_hex())
454 );
455
456 let missing = load_requested_object(
457 &store,
458 &ObjectRequest {
459 id: ObjectId::StateId(missing_state),
460 have_base: None,
461 },
462 )
463 .unwrap_err();
464 assert!(
465 matches!(missing, ProtocolError::ObjectNotFound(id) if id == missing_state.to_string())
466 );
467
468 let mismatch =
469 load_object_data(&store, &ObjectId::Hash(missing_hash), ObjectType::State).unwrap_err();
470 assert!(
471 matches!(mismatch, ProtocolError::InvalidState(message) if message == "object id/type mismatch")
472 );
473
474 let mismatch =
475 load_object_data(&store, &ObjectId::StateId(missing_state), ObjectType::Blob)
476 .unwrap_err();
477 assert!(
478 matches!(mismatch, ProtocolError::InvalidState(message) if message == "object id/type mismatch")
479 );
480 }
481
482 #[test]
483 fn store_received_object_rejects_mismatched_object_identity() {
484 let (_temp, store) = create_test_store();
485 let blob = Blob::from("tree leaf");
486 let blob_hash = store.put_blob(&blob).unwrap();
487 let tree = Tree::from_entries(vec![TreeEntry::file("leaf.txt", blob_hash, false).unwrap()]);
488 let tree_bytes = tree.encode_canonical().unwrap();
489 let wrong_hash = ContentHash::from_bytes([4; 32]);
490
491 let error = store_received_object(
492 &store,
493 &ObjectData {
494 id: ObjectId::Hash(wrong_hash),
495 obj_type: ObjectType::Tree,
496 data: tree_bytes,
497 is_delta: false,
498 },
499 )
500 .unwrap_err();
501 assert!(
502 matches!(error, ProtocolError::InvalidState(message) if message == "tree hash mismatch")
503 );
504
505 let state = State::new(tree.hash(), Vec::new(), test_attribution());
506 let wrong_state_id = objects::object::StateId::from_bytes([5; 32]);
507 let error = store_received_object(
508 &store,
509 &ObjectData {
510 id: ObjectId::StateId(wrong_state_id),
511 obj_type: ObjectType::State,
512 data: rmp_serde::to_vec_named(&state).unwrap(),
513 is_delta: false,
514 },
515 )
516 .unwrap_err();
517 assert!(
518 matches!(error, ProtocolError::InvalidState(message) if message.contains("StateId mismatch"))
519 );
520 }
521
522 #[test]
523 fn store_received_object_rejects_unbounded_tree_block_raw_length() {
524 const TREE_BLOCK_PREAMBLE_LEN: usize = 16;
525 const TREE_BLOCK_INDEX_RAW_LEN_OFFSET: usize = 20;
526
527 let (_temp, store) = create_test_store();
528 let tree = Tree::from_entries(
529 (0..256)
530 .map(|index| {
531 TreeEntry::file(
532 format!("crates__shared_component_{index:04}.rs"),
533 ContentHash::compute(format!("blob-{index}").as_bytes()),
534 false,
535 )
536 .expect("tree entry")
537 })
538 .collect(),
539 );
540 let mut data = tree.encode_canonical_blocked(3, 0).expect("blocked tree");
541 assert_eq!(data[4], TREE_BLOCK_ENCODING_VERSION);
542 let raw_len_offset =
543 TREE_HEADER_LEN + TREE_BLOCK_PREAMBLE_LEN + TREE_BLOCK_INDEX_RAW_LEN_OFFSET;
544 data[raw_len_offset..raw_len_offset + 4].copy_from_slice(&u32::MAX.to_le_bytes());
545
546 let error = store_received_object(
547 &store,
548 &ObjectData {
549 id: ObjectId::Hash(tree.hash()),
550 obj_type: ObjectType::Tree,
551 data,
552 is_delta: false,
553 },
554 )
555 .expect_err("received tree must reject attacker-controlled raw length");
556 assert!(
557 matches!(error, ProtocolError::InvalidState(ref message) if message.contains("raw length") && message.contains("exceeds maximum")),
558 "unexpected error: {error}"
559 );
560 }
561
562 #[test]
563 fn store_received_object_rejects_raw_sidecar_objects() {
564 let (_temp, store) = create_test_store();
565 let blob_hash = ContentHash::from_bytes([1; 32]);
566 let state_id = objects::object::StateId::from_bytes([2; 32]);
567
568 let redaction_error = store_received_object(
569 &store,
570 &ObjectData {
571 id: ObjectId::Hash(blob_hash),
572 obj_type: ObjectType::Redaction,
573 data: b"unsigned redaction bytes".to_vec(),
574 is_delta: false,
575 },
576 )
577 .unwrap_err();
578 assert!(
579 matches!(redaction_error, ProtocolError::InvalidState(message) if message.contains("signature verification is required"))
580 );
581
582 let visibility_error = store_received_object(
583 &store,
584 &ObjectData {
585 id: ObjectId::StateId(state_id),
586 obj_type: ObjectType::StateVisibility,
587 data: b"raw visibility bytes".to_vec(),
588 is_delta: false,
589 },
590 )
591 .unwrap_err();
592 assert!(
593 matches!(visibility_error, ProtocolError::InvalidState(message) if message.contains("sidecar validation is required"))
594 );
595 }
596
597 #[test]
598 fn test_chunk_count_rounds_up() {
599 assert_eq!(chunk_count(0, 64), 0);
600 assert_eq!(chunk_count(1, 64), 1);
601 assert_eq!(chunk_count(64, 64), 1);
602 assert_eq!(chunk_count(65, 64), 2);
603 }
604
605 #[test]
606 fn test_chunk_bounds_returns_ranges() {
607 assert_eq!(chunk_bounds(100, 32, 0), Some((0, 32)));
608 assert_eq!(chunk_bounds(100, 32, 2), Some((64, 32)));
609 assert_eq!(chunk_bounds(100, 32, 3), Some((96, 4)));
610 assert_eq!(chunk_bounds(100, 32, 4), None);
611 assert_eq!(chunk_bounds(100, 0, 0), None);
612 }
613
614 #[test]
615 fn test_chunk_offset_returns_position() {
616 assert_eq!(chunk_offset(0, 64), Some(0));
617 assert_eq!(chunk_offset(3, 64), Some(192));
618 assert_eq!(chunk_offset(usize::MAX, 2), None);
619 }
620
621 #[test]
622 fn received_transfer_blob_at_limit_is_accepted() {
623 check_received_transfer_blob_size(8, 8, "redactions").unwrap();
624 }
625
626 #[test]
627 fn received_transfer_blob_over_limit_is_rejected() {
628 let error = check_received_transfer_blob_size(9, 8, "redactions").unwrap_err();
629 let message = error.to_string();
630 assert!(
631 message.contains("redactions blob exceeds receive size limit"),
632 "unexpected error: {message}"
633 );
634 assert!(
635 message.contains("9 bytes (max 8)"),
636 "unexpected error: {message}"
637 );
638 }
639
640 #[test]
641 fn received_transfer_blob_caps_are_enforced_against_production_limits() {
642 check_received_transfer_blob_size(
643 MAX_RECEIVED_REDACTIONS_BLOB_SIZE as usize,
644 MAX_RECEIVED_REDACTIONS_BLOB_SIZE,
645 "redactions",
646 )
647 .unwrap();
648 check_received_transfer_blob_size(
649 MAX_RECEIVED_STATE_VISIBILITY_BLOB_SIZE as usize,
650 MAX_RECEIVED_STATE_VISIBILITY_BLOB_SIZE,
651 "state-visibility",
652 )
653 .unwrap();
654 }
655
656 #[test]
657 fn declared_receive_len_above_max_is_rejected_before_any_alloc() {
658 let error = admit_declared_received_len(9, 8, "pull raw body")
659 .expect_err("declared length above max must fail closed");
660 assert!(
661 error.to_string().contains("exceeds receive size limit"),
662 "got {error}"
663 );
664 assert!(error.to_string().contains("9 bytes (max 8)"), "got {error}");
665 }
666
667 #[test]
668 fn declared_receive_len_above_pack_cap_is_rejected_before_any_alloc() {
669 let error = admit_declared_received_len(
670 crate::MAX_RECEIVED_PACK_SIZE + 1,
671 crate::MAX_RECEIVED_PACK_SIZE,
672 "pull raw body",
673 )
674 .expect_err("attacker-chosen pack length must fail closed");
675 assert!(
676 error.to_string().contains("exceeds receive size limit"),
677 "got {error}"
678 );
679 }
680
681 #[test]
682 fn declared_receive_len_at_max_is_admitted() {
683 assert_eq!(
684 admit_declared_received_len(8, 8, "pull raw body").unwrap(),
685 8
686 );
687 }
688
689 #[test]
690 fn state_attachment_roundtrips_through_wire_data() {
691 let (_source_temp, source) = create_test_store();
692 let (_dest_temp, dest) = create_test_store();
693 let tree = source.put_tree(&Tree::new()).unwrap();
694 let state = State::new(tree, vec![], test_attribution());
695 source.put_state(&state).unwrap();
696 dest.put_state(&state).unwrap();
697 let attachment = StateAttachment {
698 state_id: state.id(),
699 body: StateAttachmentBody::RiskSignals(ContentHash::compute(b"signals")),
700 attribution: test_attribution(),
701 created_at: chrono::Utc::now(),
702 supersedes: None,
703 };
704 source.put_state_attachment(&attachment).unwrap();
705 let id = ObjectId::StateAttachment {
706 state: state.id(),
707 id: attachment.id(),
708 kind: attachment.body.kind(),
709 };
710 let data = load_object_data(&source, &id, ObjectType::StateAttachment).unwrap();
711 store_received_object(&dest, &data).unwrap();
712 assert_eq!(
713 dest.get_state_attachment(&state.id(), &attachment.id())
714 .unwrap(),
715 Some(attachment)
716 );
717 }
718}