//! Length-delimited protobuf codec for cTrader Open API. //! //! Wire format: `[4-byte BE length][ProtoMessage bytes]` //! //! `ProtoMessage` is the envelope defined in `OpenApiCommonMessages.proto`: //! ```text //! message ProtoMessage { //! required uint32 payloadType = 1; //! optional bytes payload = 2; //! optional string clientMsgId = 3; //! } //! ``` use bytes::{Buf, BufMut, BytesMut}; use prost::Message; use tokio_util::codec::{Decoder, Encoder}; use crate::error::{CTraderError, Result}; use crate::proto::ProtoMessage; /// Maximum allowed frame size (16 MiB). Protects against unbounded allocations. const MAX_FRAME_SIZE: u32 = 16 * 1024 * 1024; /// Codec for encoding/decoding `ProtoMessage` with 4-byte BE length framing. #[derive(Debug, Default, Clone)] pub struct CTraderCodec; impl CTraderCodec { pub const fn new() -> Self { Self } } impl Decoder for CTraderCodec { type Item = ProtoMessage; type Error = CTraderError; fn decode(&mut self, src: &mut BytesMut) -> Result> { // Need at least 4 bytes for the length prefix. if src.len() < 4 { return Ok(None); } // Peek at the length without consuming. let len = u32::from_be_bytes([src[0], src[1], src[2], src[3]]); if len > MAX_FRAME_SIZE { return Err(CTraderError::CodecError(format!( "frame too large: {len} bytes (max {MAX_FRAME_SIZE})" ))); } let total = 4 + len as usize; // Not enough data yet — tell tokio to read more. if src.len() < total { src.reserve(total - src.len()); return Ok(None); } // Consume the length prefix. src.advance(4); // Consume the payload bytes. let frame = src.split_to(len as usize); let msg = ProtoMessage::decode(frame.as_ref()).map_err(|e| { CTraderError::ProtocolError(format!("failed to decode ProtoMessage: {e}")) })?; Ok(Some(msg)) } } impl Encoder for CTraderCodec { type Error = CTraderError; fn encode(&mut self, item: ProtoMessage, dst: &mut BytesMut) -> Result<()> { let encoded_len = item.encoded_len(); if encoded_len > MAX_FRAME_SIZE as usize { return Err(CTraderError::CodecError(format!( "message too large to encode: {encoded_len} bytes (max {MAX_FRAME_SIZE})" ))); } dst.reserve(4 + encoded_len); dst.put_u32(encoded_len as u32); item.encode(dst).map_err(|e| { CTraderError::ProtocolError(format!("failed to encode ProtoMessage: {e}")) })?; Ok(()) } } #[cfg(test)] mod tests { use super::*; fn heartbeat_msg() -> ProtoMessage { ProtoMessage { payload_type: 51, // HEARTBEAT_EVENT payload: None, client_msg_id: None, } } fn order_msg_with_id() -> ProtoMessage { ProtoMessage { payload_type: 2106, // NEW_ORDER_REQ payload: Some(vec![1, 2, 3, 4]), client_msg_id: Some("req-abc-123".into()), } } #[test] fn roundtrip_heartbeat() { let mut codec = CTraderCodec::new(); let original = heartbeat_msg(); let mut buf = BytesMut::new(); codec.encode(original.clone(), &mut buf).expect("encode"); let decoded = codec.decode(&mut buf).expect("decode").expect("some"); assert_eq!(decoded.payload_type, original.payload_type); assert_eq!(decoded.payload, original.payload); assert_eq!(decoded.client_msg_id, original.client_msg_id); } #[test] fn roundtrip_with_client_msg_id() { let mut codec = CTraderCodec::new(); let original = order_msg_with_id(); let mut buf = BytesMut::new(); codec.encode(original.clone(), &mut buf).expect("encode"); let decoded = codec.decode(&mut buf).expect("decode").expect("some"); assert_eq!(decoded.payload_type, original.payload_type); assert_eq!(decoded.payload, original.payload); assert_eq!(decoded.client_msg_id, original.client_msg_id); } #[test] fn partial_frame_returns_none() { let mut codec = CTraderCodec::new(); let msg = heartbeat_msg(); let mut full_buf = BytesMut::new(); codec.encode(msg, &mut full_buf).expect("encode"); // Feed only partial data (first 3 bytes — not enough for length prefix). let mut partial = full_buf.split_to(3); assert!(codec.decode(&mut partial).expect("decode").is_none()); // Feed length prefix but incomplete payload. let mut partial2 = BytesMut::new(); partial2.extend_from_slice(&full_buf[..]); // Restore the first 3 bytes let mut combined = BytesMut::new(); combined.extend_from_slice(&partial); combined.extend_from_slice(&partial2); // Remove last byte so the frame is incomplete combined.truncate(combined.len() - 1); assert!(codec.decode(&mut combined).expect("decode").is_none()); } #[test] fn multiple_frames_in_buffer() { let mut codec = CTraderCodec::new(); let msg1 = heartbeat_msg(); let msg2 = order_msg_with_id(); let mut buf = BytesMut::new(); codec.encode(msg1.clone(), &mut buf).expect("encode msg1"); codec.encode(msg2.clone(), &mut buf).expect("encode msg2"); let decoded1 = codec.decode(&mut buf).expect("decode1").expect("some1"); assert_eq!(decoded1.payload_type, msg1.payload_type); let decoded2 = codec.decode(&mut buf).expect("decode2").expect("some2"); assert_eq!(decoded2.payload_type, msg2.payload_type); assert_eq!(decoded2.client_msg_id, msg2.client_msg_id); // Buffer should be empty now. assert!(codec.decode(&mut buf).expect("decode3").is_none()); } #[test] fn rejects_oversized_frame() { let mut codec = CTraderCodec::new(); let mut buf = BytesMut::new(); // Write a length that exceeds MAX_FRAME_SIZE. buf.put_u32(MAX_FRAME_SIZE + 1); buf.extend_from_slice(&[0u8; 64]); let result = codec.decode(&mut buf); assert!(result.is_err()); let err_msg = format!("{}", result.unwrap_err()); assert!(err_msg.contains("frame too large")); } }