diff --git a/ctrader-openapi/src/codec.rs b/ctrader-openapi/src/codec.rs new file mode 100644 index 000000000..0208e4f65 --- /dev/null +++ b/ctrader-openapi/src/codec.rs @@ -0,0 +1,204 @@ +//! 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 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")); + } +} diff --git a/ctrader-openapi/src/error.rs b/ctrader-openapi/src/error.rs index 0a5a47c0d..a43be3ea7 100644 --- a/ctrader-openapi/src/error.rs +++ b/ctrader-openapi/src/error.rs @@ -60,6 +60,16 @@ pub enum CTraderError { /// Symbol not found in the cached symbol list. #[error("unknown symbol: {0}")] UnknownSymbol(String), + + /// I/O error (required by `tokio_util::codec`). + #[error("io error: {0}")] + Io(String), +} + +impl From for CTraderError { + fn from(e: std::io::Error) -> Self { + Self::Io(e.to_string()) + } } /// Rate-limit bucket categories. diff --git a/ctrader-openapi/src/lib.rs b/ctrader-openapi/src/lib.rs index ac82b401f..f9e4c7092 100644 --- a/ctrader-openapi/src/lib.rs +++ b/ctrader-openapi/src/lib.rs @@ -18,6 +18,7 @@ //! - **account** — account info and position reconciliation //! - **client** — high-level `CTraderClient` API +pub mod codec; pub mod config; pub mod error; pub mod proto;