use melin_app::encoder::ResponseEncoder as ResponseEncoderTrait;
use melin_protocol::codec;
use melin_protocol::message::ResponseKind;
use melin_types::types::{ExecutionReport, QueryResponse};
#[derive(Debug, Clone, Copy)]
pub struct ResponseEncoder;
impl ResponseEncoderTrait for ResponseEncoder {
type Report = ExecutionReport;
type Query = QueryResponse;
fn encode_report(
&self,
report: &ExecutionReport,
buf: &mut [u8],
) -> Result<usize, &'static str> {
codec::encode_response(&ResponseKind::Report(*report), buf).map_err(|_| "encode error")
}
fn encode_query(&self, query: &QueryResponse, buf: &mut [u8]) -> Result<usize, &'static str> {
let kind = match *query {
QueryResponse::Stats {
active_connections,
events_processed,
journal_sequence,
} => ResponseKind::StatsHeader {
active_connections,
events_processed,
journal_sequence,
},
QueryResponse::Position {
account,
balances,
count,
} => ResponseKind::PositionSnapshot {
account,
balances,
count,
},
QueryResponse::RequestSeqHwm { hwm } => ResponseKind::RequestSeqHwm { hwm },
};
codec::encode_response(&kind, buf).map_err(|_| "encode error")
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::num::NonZeroU64;
use melin_types::types::*;
const SCRATCH: usize = 512;
fn round_trip(written: &[u8]) -> ResponseKind {
codec::decode_response(&written[4..]).expect("decode")
}
fn sample_placed() -> ExecutionReport {
ExecutionReport::Placed {
order_id: OrderId(1),
symbol: Symbol(1),
account: AccountId(1),
side: Side::Buy,
price: Price(NonZeroU64::new(100).unwrap()),
quantity: Quantity(NonZeroU64::new(10).unwrap()),
}
}
#[test]
fn encodes_report() {
let mut buf = [0u8; SCRATCH];
let n = ResponseEncoder
.encode_report(&sample_placed(), &mut buf)
.unwrap();
assert!(matches!(
round_trip(&buf[..n]),
ResponseKind::Report(ExecutionReport::Placed { order_id, .. })
if order_id == OrderId(1)
));
}
#[test]
fn encodes_query_stats() {
let q = QueryResponse::Stats {
active_connections: 7,
events_processed: 12345,
journal_sequence: 999,
};
let mut buf = [0u8; SCRATCH];
let n = ResponseEncoder.encode_query(&q, &mut buf).unwrap();
assert!(matches!(
round_trip(&buf[..n]),
ResponseKind::StatsHeader {
active_connections: 7,
events_processed: 12345,
journal_sequence: 999,
}
));
}
#[test]
fn encodes_query_position() {
let mut balances = [AccountBalance::ZERO; 16];
balances[0] = AccountBalance {
currency: CurrencyId(1),
free: 100,
reserved: 0,
};
let q = QueryResponse::Position {
account: AccountId(42),
balances,
count: 1,
};
let mut buf = [0u8; SCRATCH];
let n = ResponseEncoder.encode_query(&q, &mut buf).unwrap();
assert!(matches!(
round_trip(&buf[..n]),
ResponseKind::PositionSnapshot { account, count: 1, .. }
if account == AccountId(42)
));
}
#[test]
fn encodes_query_request_seq_hwm() {
let q = QueryResponse::RequestSeqHwm { hwm: 4242 };
let mut buf = [0u8; SCRATCH];
let n = ResponseEncoder.encode_query(&q, &mut buf).unwrap();
assert!(matches!(
round_trip(&buf[..n]),
ResponseKind::RequestSeqHwm { hwm: 4242 }
));
}
}