1use crate::binary::{BinaryMessageEnvelope, Flags, MessageType};
4use crate::DCPError;
5
6const KNOWN_FLAG_BITS: u8 = Flags::STREAMING | Flags::COMPRESSED | Flags::SIGNED;
7
8#[derive(Debug, PartialEq)]
10pub struct ParsedMessage<'a> {
11 pub envelope: &'a BinaryMessageEnvelope,
13 pub payload: &'a [u8],
15}
16
17pub struct MessageParser;
19
20impl MessageParser {
21 pub fn parse(bytes: &[u8]) -> Result<ParsedMessage<'_>, DCPError> {
23 let envelope = BinaryMessageEnvelope::from_bytes(bytes)?;
24 let payload_start = BinaryMessageEnvelope::SIZE;
25 let payload_end = payload_start
26 .checked_add(envelope.payload_len as usize)
27 .ok_or(DCPError::ValidationFailed)?;
28
29 Self::validate_message_type(envelope.message_type)?;
30 Self::validate_flags(envelope.flags)?;
31
32 if bytes.len() < payload_end {
33 return Err(DCPError::InsufficientData);
34 }
35
36 if bytes.len() != payload_end {
37 return Err(DCPError::ValidationFailed);
38 }
39
40 Ok(ParsedMessage {
41 envelope,
42 payload: &bytes[payload_start..payload_end],
43 })
44 }
45
46 pub fn validate_message_type(msg_type: u8) -> Result<MessageType, DCPError> {
48 MessageType::from_u8(msg_type).ok_or(DCPError::UnknownMessageType)
49 }
50
51 pub fn validate_flags(flags: u8) -> Result<(), DCPError> {
53 if flags & !KNOWN_FLAG_BITS == 0 {
54 Ok(())
55 } else {
56 Err(DCPError::ValidationFailed)
57 }
58 }
59
60 pub fn extract_flags(envelope: &BinaryMessageEnvelope) -> MessageFlags {
62 MessageFlags {
63 streaming: envelope.flags & Flags::STREAMING != 0,
64 compressed: envelope.flags & Flags::COMPRESSED != 0,
65 signed: envelope.flags & Flags::SIGNED != 0,
66 }
67 }
68}
69
70#[derive(Debug, Clone, Copy, PartialEq, Eq)]
72pub struct MessageFlags {
73 pub streaming: bool,
74 pub compressed: bool,
75 pub signed: bool,
76}
77
78impl MessageFlags {
79 pub fn to_u8(&self) -> u8 {
81 let mut flags = 0u8;
82 if self.streaming {
83 flags |= Flags::STREAMING;
84 }
85 if self.compressed {
86 flags |= Flags::COMPRESSED;
87 }
88 if self.signed {
89 flags |= Flags::SIGNED;
90 }
91 flags
92 }
93}
94
95pub struct MessageDispatcher;
97
98impl MessageDispatcher {
99 pub fn dispatch<'a, H: MessageHandler>(
101 message: &ParsedMessage<'a>,
102 handler: &H,
103 ) -> Result<(), DCPError> {
104 let msg_type = MessageParser::validate_message_type(message.envelope.message_type)?;
105
106 match msg_type {
107 MessageType::Tool => handler.handle_tool(message),
108 MessageType::Resource => handler.handle_resource(message),
109 MessageType::Prompt => handler.handle_prompt(message),
110 MessageType::Response => handler.handle_response(message),
111 MessageType::Error => handler.handle_error(message),
112 MessageType::Stream => handler.handle_stream(message),
113 }
114 }
115}
116
117pub trait MessageHandler {
119 fn handle_tool(&self, message: &ParsedMessage<'_>) -> Result<(), DCPError>;
120 fn handle_resource(&self, message: &ParsedMessage<'_>) -> Result<(), DCPError>;
121 fn handle_prompt(&self, message: &ParsedMessage<'_>) -> Result<(), DCPError>;
122 fn handle_response(&self, message: &ParsedMessage<'_>) -> Result<(), DCPError>;
123 fn handle_error(&self, message: &ParsedMessage<'_>) -> Result<(), DCPError>;
124 fn handle_stream(&self, message: &ParsedMessage<'_>) -> Result<(), DCPError>;
125}
126
127#[cfg(test)]
128mod tests {
129 use super::*;
130
131 #[test]
132 fn test_parse_message() {
133 let mut bytes = vec![0u8; 16];
134 let envelope = BinaryMessageEnvelope::new(MessageType::Tool, 0, 8);
136 bytes[..8].copy_from_slice(envelope.as_bytes());
137 bytes[8..16].copy_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]);
139
140 let parsed = MessageParser::parse(&bytes).unwrap();
141 assert_eq!(parsed.envelope.message_type, MessageType::Tool as u8);
142 assert_eq!(parsed.payload, &[1, 2, 3, 4, 5, 6, 7, 8]);
143 }
144
145 #[test]
146 fn test_parse_insufficient_payload() {
147 let mut bytes = vec![0u8; 12];
148 let envelope = BinaryMessageEnvelope::new(MessageType::Tool, 0, 100);
149 bytes[..8].copy_from_slice(envelope.as_bytes());
150
151 assert_eq!(
152 MessageParser::parse(&bytes),
153 Err(DCPError::InsufficientData)
154 );
155 }
156
157 #[test]
158 fn test_validate_message_type() {
159 assert_eq!(
160 MessageParser::validate_message_type(1),
161 Ok(MessageType::Tool)
162 );
163 assert_eq!(
164 MessageParser::validate_message_type(6),
165 Ok(MessageType::Stream)
166 );
167 assert_eq!(
168 MessageParser::validate_message_type(0),
169 Err(DCPError::UnknownMessageType)
170 );
171 assert_eq!(
172 MessageParser::validate_message_type(7),
173 Err(DCPError::UnknownMessageType)
174 );
175 }
176
177 #[test]
178 fn test_extract_flags() {
179 let envelope =
180 BinaryMessageEnvelope::new(MessageType::Tool, Flags::STREAMING | Flags::SIGNED, 0);
181 let flags = MessageParser::extract_flags(&envelope);
182
183 assert!(flags.streaming);
184 assert!(!flags.compressed);
185 assert!(flags.signed);
186 }
187
188 #[test]
189 fn test_flags_round_trip() {
190 let flags = MessageFlags {
191 streaming: true,
192 compressed: false,
193 signed: true,
194 };
195 let byte = flags.to_u8();
196 assert_eq!(byte, Flags::STREAMING | Flags::SIGNED);
197 }
198}