use dcp::binary::{
BinaryMessageEnvelope, ChunkFlags, HbtpHeader, MessageType, SignedInvocation, SignedToolDef,
StreamChunk, ToolInvocation,
};
use dcp::DCPError;
use proptest::prelude::*;
fn message_type_strategy() -> impl Strategy<Value = u8> {
prop_oneof![
Just(MessageType::Tool as u8),
Just(MessageType::Resource as u8),
Just(MessageType::Prompt as u8),
Just(MessageType::Response as u8),
Just(MessageType::Error as u8),
Just(MessageType::Stream as u8),
]
}
fn flags_strategy() -> impl Strategy<Value = u8> {
0u8..=7u8 }
fn stream_chunk_flags_strategy() -> impl Strategy<Value = u8> {
prop_oneof![
Just(ChunkFlags::FIRST),
Just(ChunkFlags::CONTINUE),
Just(ChunkFlags::LAST),
Just(ChunkFlags::ERROR),
Just(ChunkFlags::FIRST | ChunkFlags::LAST),
]
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
#[test]
fn prop_envelope_round_trip(
msg_type in message_type_strategy(),
flags in flags_strategy(),
payload_len in any::<u32>(),
) {
let msg_type_enum = MessageType::from_u8(msg_type).unwrap();
let original = BinaryMessageEnvelope::new(msg_type_enum, flags, payload_len);
let bytes = original.as_bytes();
let parsed = BinaryMessageEnvelope::from_bytes(bytes).unwrap();
prop_assert_eq!(parsed.magic, original.magic);
prop_assert_eq!(parsed.message_type, original.message_type);
prop_assert_eq!(parsed.flags, original.flags);
prop_assert_eq!(parsed.payload_len, original.payload_len);
}
#[test]
fn prop_hbtp_header_round_trip(
version in any::<u8>(),
msg_type in any::<u8>(),
flags in any::<u16>(),
stream_id in any::<u32>(),
length in any::<u32>(),
) {
let original = HbtpHeader::new(version, msg_type, flags, stream_id, length);
let bytes = original.as_bytes();
let parsed = HbtpHeader::from_bytes(bytes).unwrap();
prop_assert_eq!(parsed.magic, original.magic);
prop_assert_eq!(parsed.version, original.version);
prop_assert_eq!(parsed.msg_type, original.msg_type);
prop_assert_eq!(parsed.flags, original.flags);
prop_assert_eq!(parsed.stream_id, original.stream_id);
prop_assert_eq!(parsed.length, original.length);
prop_assert_eq!(parsed.checksum, original.checksum);
}
#[test]
fn prop_hbtp_checksum_detects_corruption(
version in any::<u8>(),
msg_type in any::<u8>(),
flags in any::<u16>(),
stream_id in any::<u32>(),
length in any::<u32>(),
bit_to_flip in 0usize..16, ) {
let original = HbtpHeader::new(version, msg_type, flags, stream_id, length);
let mut bytes = original.as_bytes().to_vec();
let byte_idx = bit_to_flip / 8;
let bit_idx = bit_to_flip % 8;
bytes[byte_idx] ^= 1 << bit_idx;
let result = HbtpHeader::from_bytes(&bytes);
if let Ok(parsed) = result {
prop_assert!(!parsed.verify_checksum());
}
}
#[test]
fn prop_tool_invocation_round_trip(
tool_id in any::<u32>(),
arg_layout in any::<u64>(),
args_offset in any::<u32>(),
args_len in any::<u32>(),
) {
let original = ToolInvocation::new(tool_id, arg_layout, args_offset, args_len);
let bytes = original.as_bytes();
let parsed = ToolInvocation::from_bytes(bytes).unwrap();
prop_assert_eq!(parsed.tool_id, original.tool_id);
prop_assert_eq!(parsed.arg_layout, original.arg_layout);
prop_assert_eq!(parsed.args_offset, original.args_offset);
prop_assert_eq!(parsed.args_len, original.args_len);
}
#[test]
fn prop_stream_chunk_round_trip(
sequence in any::<u32>(),
flags in stream_chunk_flags_strategy(),
len in any::<u16>(),
) {
let original = StreamChunk::new(sequence, flags, len);
let bytes = original.as_bytes();
let parsed = StreamChunk::from_bytes(bytes).unwrap();
prop_assert_eq!(parsed.sequence, original.sequence);
prop_assert_eq!(parsed.flags, original.flags);
prop_assert_eq!(parsed.len, original.len);
}
#[test]
fn prop_stream_chunk_rejects_malformed_flags(
sequence in any::<u32>(),
flags in prop_oneof![
Just(0u8),
Just(0x10u8),
Just(0xffu8),
Just(ChunkFlags::FIRST | ChunkFlags::CONTINUE),
Just(ChunkFlags::CONTINUE | ChunkFlags::LAST),
Just(ChunkFlags::ERROR | ChunkFlags::FIRST),
Just(ChunkFlags::ERROR | ChunkFlags::LAST),
],
len in any::<u16>(),
) {
let chunk = StreamChunk::new(sequence, flags, len);
let parsed = StreamChunk::from_bytes(chunk.as_bytes());
prop_assert!(parsed.is_err());
}
#[test]
fn prop_signed_tool_def_round_trip(
tool_id in any::<u32>(),
schema_hash in any::<[u8; 32]>(),
capabilities in any::<u64>(),
signature in any::<[u8; 64]>(),
public_key in any::<[u8; 32]>(),
) {
let original = SignedToolDef {
tool_id,
schema_hash,
capabilities,
signature,
public_key,
};
let bytes = original.as_bytes();
let parsed = SignedToolDef::from_bytes(bytes).unwrap();
prop_assert_eq!(parsed.tool_id, original.tool_id);
prop_assert_eq!(parsed.schema_hash, original.schema_hash);
prop_assert_eq!(parsed.capabilities, original.capabilities);
prop_assert_eq!(parsed.signature, original.signature);
prop_assert_eq!(parsed.public_key, original.public_key);
}
#[test]
fn prop_signed_invocation_round_trip(
tool_id in any::<u32>(),
nonce in any::<u64>(),
timestamp in any::<u64>(),
args_hash in any::<[u8; 32]>(),
signature in any::<[u8; 64]>(),
) {
let original = SignedInvocation {
tool_id,
nonce,
timestamp,
args_hash,
signature,
};
let bytes = original.as_bytes();
let parsed = SignedInvocation::from_bytes(bytes).unwrap();
prop_assert_eq!(parsed.tool_id, original.tool_id);
prop_assert_eq!(parsed.nonce, original.nonce);
prop_assert_eq!(parsed.timestamp, original.timestamp);
prop_assert_eq!(parsed.args_hash, original.args_hash);
prop_assert_eq!(parsed.signature, original.signature);
}
}
#[test]
fn tool_invocation_rejects_nonzero_reserved_padding() {
let invocation = ToolInvocation::new(7, 0x0102_0304_0506_0708, 9, 10);
let mut bytes = invocation.as_bytes().to_vec();
bytes[4] = 0xAA;
assert_eq!(
ToolInvocation::from_bytes(&bytes),
Err(DCPError::ValidationFailed)
);
}
#[test]
fn signed_tool_def_rejects_nonzero_reserved_padding() {
let mut bytes = SignedToolDef {
tool_id: 7,
schema_hash: [0x11; 32],
capabilities: 0x0102_0304_0506_0708,
signature: [0x22; 64],
public_key: [0x33; 32],
}
.as_bytes()
.to_vec();
bytes[36] = 0xAA;
assert_eq!(
SignedToolDef::from_bytes(&bytes),
Err(DCPError::ValidationFailed)
);
}
#[test]
fn signed_invocation_rejects_nonzero_reserved_padding() {
let mut bytes = SignedInvocation {
tool_id: 7,
nonce: 8,
timestamp: 9,
args_hash: [0x11; 32],
signature: [0x22; 64],
}
.as_bytes()
.to_vec();
bytes[4] = 0xAA;
assert_eq!(
SignedInvocation::from_bytes(&bytes),
Err(DCPError::ValidationFailed)
);
}
#[test]
fn fixed_binary_structs_reject_trailing_bytes() {
let mut invocation = ToolInvocation::new(7, 0x0102_0304_0506_0708, 9, 10)
.as_bytes()
.to_vec();
invocation.push(0xAA);
assert_eq!(
ToolInvocation::from_bytes(&invocation),
Err(DCPError::ValidationFailed)
);
let mut def = SignedToolDef {
tool_id: 7,
schema_hash: [0x11; 32],
capabilities: 0x0102_0304_0506_0708,
signature: [0x22; 64],
public_key: [0x33; 32],
}
.as_bytes()
.to_vec();
def.push(0xAA);
assert_eq!(
SignedToolDef::from_bytes(&def),
Err(DCPError::ValidationFailed)
);
let mut signed_invocation = SignedInvocation {
tool_id: 7,
nonce: 8,
timestamp: 9,
args_hash: [0x11; 32],
signature: [0x22; 64],
}
.as_bytes()
.to_vec();
signed_invocation.push(0xAA);
assert_eq!(
SignedInvocation::from_bytes(&signed_invocation),
Err(DCPError::ValidationFailed)
);
}
#[test]
fn signed_payload_bytes_are_compact_field_encodings() {
let def = SignedToolDef {
tool_id: 0x0102_0304,
schema_hash: [0x11; 32],
capabilities: 0x0506_0708_090A_0B0C,
signature: [0x22; 64],
public_key: [0x33; 32],
};
let mut expected_def = Vec::with_capacity(44);
expected_def.extend_from_slice(&def.tool_id.to_le_bytes());
expected_def.extend_from_slice(&def.schema_hash);
expected_def.extend_from_slice(&def.capabilities.to_le_bytes());
assert_eq!(def.signed_bytes().as_slice(), expected_def.as_slice());
let inv = SignedInvocation {
tool_id: 0x0102_0304,
nonce: 0x0506_0708_090A_0B0C,
timestamp: 0x0D0E_0F10_1112_1314,
args_hash: [0x44; 32],
signature: [0x55; 64],
};
let mut expected_inv = Vec::with_capacity(52);
expected_inv.extend_from_slice(&inv.tool_id.to_le_bytes());
expected_inv.extend_from_slice(&inv.nonce.to_le_bytes());
expected_inv.extend_from_slice(&inv.timestamp.to_le_bytes());
expected_inv.extend_from_slice(&inv.args_hash);
assert_eq!(inv.signed_bytes().as_slice(), expected_inv.as_slice());
}