Skip to main content

dcp/protocol/
parser.rs

1//! Message parsing and validation for DCP protocol.
2
3use crate::binary::{BinaryMessageEnvelope, Flags, MessageType};
4use crate::DCPError;
5
6const KNOWN_FLAG_BITS: u8 = Flags::STREAMING | Flags::COMPRESSED | Flags::SIGNED;
7
8/// Parsed DCP message with envelope and payload reference
9#[derive(Debug, PartialEq)]
10pub struct ParsedMessage<'a> {
11    /// The message envelope
12    pub envelope: &'a BinaryMessageEnvelope,
13    /// The payload bytes
14    pub payload: &'a [u8],
15}
16
17/// Message parser for DCP protocol
18pub struct MessageParser;
19
20impl MessageParser {
21    /// Parse a complete DCP message from bytes
22    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    /// Validate message type is known
47    pub fn validate_message_type(msg_type: u8) -> Result<MessageType, DCPError> {
48        MessageType::from_u8(msg_type).ok_or(DCPError::UnknownMessageType)
49    }
50
51    /// Validate that only defined flag bits are set.
52    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    /// Extract flags from envelope
61    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/// Extracted message flags
71#[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    /// Convert flags back to u8
80    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
95/// Message dispatcher for routing by type
96pub struct MessageDispatcher;
97
98impl MessageDispatcher {
99    /// Dispatch a parsed message to the appropriate handler
100    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
117/// Trait for handling different message types
118pub 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        // Create envelope
135        let envelope = BinaryMessageEnvelope::new(MessageType::Tool, 0, 8);
136        bytes[..8].copy_from_slice(envelope.as_bytes());
137        // Add payload
138        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}