use crate::egress::decoder::{DecodedBatch, ZstdScratch, decode_result_batch};
use crate::egress::schema::Schema;
use crate::egress::symbol_dict::SymbolDict;
use crate::egress::wire::ByteReader;
use crate::egress::wire::cache_reset::resets_dict;
use crate::egress::wire::capabilities::has_zone;
use crate::egress::wire::header::FrameHeader;
use crate::egress::wire::msg_kind::{MsgKind, StatusCode};
use crate::egress::wire::roles;
use crate::error::{Result, fmt};
use bytes::Bytes;
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum ServerRole {
Standalone,
Primary,
Replica,
PrimaryCatchup,
Other(u8),
}
impl ServerRole {
pub fn from_u8(byte: u8) -> Self {
match byte {
roles::STANDALONE => ServerRole::Standalone,
roles::PRIMARY => ServerRole::Primary,
roles::REPLICA => ServerRole::Replica,
roles::PRIMARY_CATCHUP => ServerRole::PrimaryCatchup,
other => ServerRole::Other(other),
}
}
pub fn as_str(self) -> String {
match self {
ServerRole::Standalone => roles::NAME_STANDALONE.to_string(),
ServerRole::Primary => roles::NAME_PRIMARY.to_string(),
ServerRole::Replica => roles::NAME_REPLICA.to_string(),
ServerRole::PrimaryCatchup => roles::NAME_PRIMARY_CATCHUP.to_string(),
ServerRole::Other(b) => format!("UNKNOWN({})", b),
}
}
pub fn as_u8(self) -> u8 {
match self {
ServerRole::Standalone => roles::STANDALONE,
ServerRole::Primary => roles::PRIMARY,
ServerRole::Replica => roles::REPLICA,
ServerRole::PrimaryCatchup => roles::PRIMARY_CATCHUP,
ServerRole::Other(b) => b,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ServerInfo {
pub role: ServerRole,
pub epoch: u64,
pub capabilities: u32,
pub server_wall_ns: i64,
pub cluster_id: String,
pub node_id: String,
pub zone_id: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct UpgradeReject {
pub role_byte: u8,
pub role_name: String,
pub zone: Option<String>,
}
impl UpgradeReject {
pub fn new(role_byte: u8, role_name: impl Into<String>, zone: Option<String>) -> Self {
Self {
role_byte,
role_name: role_name.into(),
zone,
}
}
pub fn is_transient(&self) -> bool {
self.role_byte == roles::PRIMARY_CATCHUP
|| self
.role_name
.eq_ignore_ascii_case(roles::NAME_PRIMARY_CATCHUP)
}
}
#[derive(Debug, Clone)]
pub enum ServerEvent {
Batch(DecodedBatch),
End {
request_id: i64,
final_seq: u64,
total_rows: u64,
},
Error {
request_id: i64,
status: StatusCode,
message: String,
},
ExecDone {
request_id: i64,
op_type: u8,
rows_affected: u64,
},
CacheReset {
#[allow(dead_code)]
mask: u8,
},
ServerInfo(ServerInfo),
}
pub fn decode_frame(
header: FrameHeader,
payload: &Bytes,
dict: &mut SymbolDict,
query_schema: &mut Option<Schema>,
zstd_scratch: &mut ZstdScratch,
) -> Result<ServerEvent> {
if payload.is_empty() {
return Err(fmt!(ProtocolError, "frame payload is empty"));
}
let kind_byte = payload[0];
let kind = MsgKind::from_u8(kind_byte)?;
let expected_tc = if matches!(kind, MsgKind::ResultBatch) {
1
} else {
0
};
if header.table_count != expected_tc {
return Err(fmt!(
ProtocolError,
"frame for msg_kind 0x{:02X} has table_count {} (expected {})",
kind_byte,
header.table_count,
expected_tc
));
}
match kind {
MsgKind::ResultBatch => Ok(ServerEvent::Batch(decode_result_batch(
payload,
header.flags,
dict,
query_schema,
zstd_scratch,
)?)),
MsgKind::ResultEnd => decode_result_end(payload),
MsgKind::QueryError => decode_query_error(payload),
MsgKind::ExecDone => decode_exec_done(payload),
MsgKind::CacheReset => decode_cache_reset(payload, dict),
MsgKind::ServerInfo => decode_server_info(payload),
MsgKind::QueryRequest | MsgKind::Cancel | MsgKind::Credit => Err(fmt!(
ProtocolError,
"server sent client-only message kind 0x{:02X}",
kind_byte
)),
}
}
fn decode_result_end(payload: &[u8]) -> Result<ServerEvent> {
let mut r = ByteReader::new(payload);
expect_kind(&mut r, MsgKind::ResultEnd)?;
let request_id = r.read_i64_le()?;
let final_seq = r.read_varint_u64()?;
let total_rows = r.read_varint_u64()?;
expect_eof(&r, "RESULT_END")?;
Ok(ServerEvent::End {
request_id,
final_seq,
total_rows,
})
}
fn decode_query_error(payload: &[u8]) -> Result<ServerEvent> {
let mut r = ByteReader::new(payload);
expect_kind(&mut r, MsgKind::QueryError)?;
let request_id = r.read_i64_le()?;
let status = StatusCode::from_u8(r.read_u8()?)?;
let msg_len = r.read_u16_le()? as usize;
let bytes = r.read_bytes(msg_len)?;
let message = std::str::from_utf8(bytes)
.map_err(|e| fmt!(InvalidUtf8, "QUERY_ERROR message not valid UTF-8: {}", e))?
.to_string();
expect_eof(&r, "QUERY_ERROR")?;
Ok(ServerEvent::Error {
request_id,
status,
message,
})
}
fn decode_exec_done(payload: &[u8]) -> Result<ServerEvent> {
let mut r = ByteReader::new(payload);
expect_kind(&mut r, MsgKind::ExecDone)?;
let request_id = r.read_i64_le()?;
let op_type = r.read_u8()?;
let rows_affected = r.read_varint_u64()?;
expect_eof(&r, "EXEC_DONE")?;
Ok(ServerEvent::ExecDone {
request_id,
op_type,
rows_affected,
})
}
fn decode_cache_reset(payload: &[u8], dict: &mut SymbolDict) -> Result<ServerEvent> {
let mut r = ByteReader::new(payload);
expect_kind(&mut r, MsgKind::CacheReset)?;
let mask = r.read_u8()?;
expect_eof(&r, "CACHE_RESET")?;
if resets_dict(mask) {
dict.reset();
}
Ok(ServerEvent::CacheReset { mask })
}
fn decode_server_info(payload: &[u8]) -> Result<ServerEvent> {
let mut r = ByteReader::new(payload);
expect_kind(&mut r, MsgKind::ServerInfo)?;
let role = ServerRole::from_u8(r.read_u8()?);
let epoch = r.read_u64_le()?;
let capabilities = r.read_u32_le()?;
let server_wall_ns = r.read_i64_le()?;
let cluster_id = read_u16_string(&mut r, "cluster_id")?;
let node_id = read_u16_string(&mut r, "node_id")?;
let zone_id = if has_zone(capabilities) {
Some(read_u16_string(&mut r, "zone_id")?)
} else {
None
};
expect_eof(&r, "SERVER_INFO")?;
Ok(ServerEvent::ServerInfo(ServerInfo {
role,
epoch,
capabilities,
server_wall_ns,
cluster_id,
node_id,
zone_id,
}))
}
fn expect_kind(r: &mut ByteReader<'_>, expected: MsgKind) -> Result<()> {
let got = r.read_u8()?;
if got != expected.as_u8() {
return Err(fmt!(
ProtocolError,
"expected msg_kind 0x{:02X}, got 0x{:02X}",
expected.as_u8(),
got
));
}
Ok(())
}
fn expect_eof(r: &ByteReader<'_>, msg_name: &str) -> Result<()> {
if !r.is_empty() {
return Err(fmt!(
ProtocolError,
"{} has {} trailing bytes",
msg_name,
r.remaining().len()
));
}
Ok(())
}
fn read_u16_string(r: &mut ByteReader<'_>, field: &str) -> Result<String> {
let len = r.read_u16_le()? as usize;
let bytes = r.read_bytes(len)?;
std::str::from_utf8(bytes)
.map_err(|e| fmt!(InvalidUtf8, "{} not valid UTF-8: {}", field, e))
.map(|s| s.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::egress::wire::header::{HEADER_LEN, PROTOCOL_VERSION};
use crate::egress::wire::varint::encode_u64;
use crate::error::ErrorCode;
fn header(payload_len: usize) -> FrameHeader {
FrameHeader {
version: PROTOCOL_VERSION,
flags: 0,
table_count: 0,
payload_length: payload_len as u32,
}
}
fn build_result_end(rid: i64, final_seq: u64, total_rows: u64) -> Bytes {
let mut p = vec![MsgKind::ResultEnd.as_u8()];
p.extend_from_slice(&rid.to_le_bytes());
encode_u64(final_seq, &mut p);
encode_u64(total_rows, &mut p);
Bytes::from(p)
}
#[test]
fn decode_result_end_ok() {
let payload = build_result_end(42, 7, 1_000);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
match event {
ServerEvent::End {
request_id,
final_seq,
total_rows,
} => {
assert_eq!(request_id, 42);
assert_eq!(final_seq, 7);
assert_eq!(total_rows, 1000);
}
_ => panic!("wrong event"),
}
}
fn build_query_error(rid: i64, status: StatusCode, msg: &str) -> Bytes {
let mut p = vec![MsgKind::QueryError.as_u8()];
p.extend_from_slice(&rid.to_le_bytes());
p.push(status.as_u8());
p.extend_from_slice(&(msg.len() as u16).to_le_bytes());
p.extend_from_slice(msg.as_bytes());
Bytes::from(p)
}
#[test]
fn decode_query_error_ok() {
let payload = build_query_error(9, StatusCode::ParseError, "bad SQL");
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
match event {
ServerEvent::Error {
request_id,
status,
message,
} => {
assert_eq!(request_id, 9);
assert_eq!(status, StatusCode::ParseError);
assert_eq!(message, "bad SQL");
}
_ => panic!("wrong event"),
}
}
#[test]
fn query_error_truncated_message_rejected() {
let payload = build_query_error(1, StatusCode::InternalError, "details");
let truncated = payload.slice(..payload.len() - 3);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_frame(
header(truncated.len()),
&truncated,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn query_error_invalid_utf8_rejected() {
let mut p = vec![MsgKind::QueryError.as_u8()];
p.extend_from_slice(&1i64.to_le_bytes());
p.push(StatusCode::InternalError.as_u8());
p.extend_from_slice(&2u16.to_le_bytes());
p.extend_from_slice(&[0xFF, 0xFE]);
let p = Bytes::from(p);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_frame(
header(p.len()),
&p,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::InvalidUtf8);
}
#[test]
fn decode_exec_done_ok() {
let mut p = vec![MsgKind::ExecDone.as_u8()];
p.extend_from_slice(&5i64.to_le_bytes());
p.push(0xAB); encode_u64(0, &mut p); let p = Bytes::from(p);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let event = decode_frame(
header(p.len()),
&p,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
match event {
ServerEvent::ExecDone {
request_id,
op_type,
rows_affected,
} => {
assert_eq!(request_id, 5);
assert_eq!(op_type, 0xAB);
assert_eq!(rows_affected, 0);
}
_ => panic!("wrong event"),
}
}
fn build_cache_reset(mask: u8) -> Bytes {
Bytes::from(vec![MsgKind::CacheReset.as_u8(), mask])
}
#[test]
fn cache_reset_clears_dict() {
let mut dict = SymbolDict::new();
dict.apply_delta(0, [b"x".as_slice()]).unwrap();
let mut query_schema: Option<Schema> = None;
let payload = build_cache_reset(0x01);
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut query_schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert!(matches!(event, ServerEvent::CacheReset { mask: 0x01 }));
assert_eq!(dict.len(), 0);
}
#[test]
fn cache_reset_ignores_reserved_bits() {
let mut dict = SymbolDict::new();
dict.apply_delta(0, [b"x".as_slice()]).unwrap();
let mut query_schema: Option<Schema> = None;
let payload = build_cache_reset(0x83);
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut query_schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert!(matches!(event, ServerEvent::CacheReset { mask: 0x83 }));
assert_eq!(
dict.len(),
0,
"DICT bit must apply even with reserved bits set"
);
}
fn build_server_info(role: u8, cluster: &str, node: &str) -> Bytes {
build_server_info_with(role, 0, cluster, node, None)
}
fn build_server_info_with(
role: u8,
capabilities: u32,
cluster: &str,
node: &str,
zone: Option<&str>,
) -> Bytes {
let mut p = vec![MsgKind::ServerInfo.as_u8()];
p.push(role);
p.extend_from_slice(&7u64.to_le_bytes()); p.extend_from_slice(&capabilities.to_le_bytes());
p.extend_from_slice(&123_456_789i64.to_le_bytes()); p.extend_from_slice(&(cluster.len() as u16).to_le_bytes());
p.extend_from_slice(cluster.as_bytes());
p.extend_from_slice(&(node.len() as u16).to_le_bytes());
p.extend_from_slice(node.as_bytes());
if let Some(z) = zone {
p.extend_from_slice(&(z.len() as u16).to_le_bytes());
p.extend_from_slice(z.as_bytes());
}
Bytes::from(p)
}
#[test]
fn decode_server_info_primary() {
let payload = build_server_info(0x01, "cluster-A", "node-1");
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ServerEvent::ServerInfo(info) = event else {
panic!()
};
assert_eq!(info.role, ServerRole::Primary);
assert_eq!(info.epoch, 7);
assert_eq!(info.capabilities, 0);
assert_eq!(info.server_wall_ns, 123_456_789);
assert_eq!(info.cluster_id, "cluster-A");
assert_eq!(info.node_id, "node-1");
assert_eq!(info.zone_id, None, "CAP_ZONE=0 leaves zone_id absent");
}
#[test]
fn unknown_role_byte_is_other_variant() {
let payload = build_server_info(0x55, "c", "n");
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ServerEvent::ServerInfo(info) = event else {
panic!()
};
assert_eq!(info.role, ServerRole::Other(0x55));
}
#[test]
fn decode_server_info_with_cap_zone_reads_zone_id() {
let payload = build_server_info_with(
0x01,
crate::egress::wire::CAP_ZONE,
"cluster-A",
"node-1",
Some("eu-west-1a"),
);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let event = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ServerEvent::ServerInfo(info) = event else {
panic!()
};
assert_eq!(
info.capabilities & crate::egress::wire::CAP_ZONE,
crate::egress::wire::CAP_ZONE
);
assert_eq!(info.zone_id.as_deref(), Some("eu-west-1a"));
}
#[test]
fn cap_zone_set_but_zone_id_missing_is_protocol_error() {
let payload = build_server_info_with(
0x01,
crate::egress::wire::CAP_ZONE,
"c",
"n",
None, );
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn unknown_capabilities_bit_with_trailing_bytes_is_protocol_error() {
let mut payload = build_server_info(0x01, "c", "n").to_vec();
payload.extend_from_slice(&[0xDE, 0xAD]);
let payload = Bytes::from(payload);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn empty_payload_rejected() {
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let empty = Bytes::new();
let err = decode_frame(
header(0),
&empty,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn unknown_msg_kind_rejected() {
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let p = Bytes::from(vec![0xAA]);
let err = decode_frame(
header(1),
&p,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn client_only_kinds_rejected_from_server() {
for k in [
MsgKind::QueryRequest.as_u8(),
MsgKind::Cancel.as_u8(),
MsgKind::Credit.as_u8(),
] {
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let p = Bytes::from(vec![k]);
let err = decode_frame(
header(1),
&p,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(err.msg().contains("client-only"));
}
}
#[test]
fn trailing_bytes_rejected_for_simple_messages() {
let payload = build_result_end(1, 0, 0);
let mut bytes_vec: Vec<u8> = payload.to_vec();
bytes_vec.push(0xFF);
let payload = Bytes::from(bytes_vec);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_frame(
header(payload.len()),
&payload,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn header_len_is_12() {
assert_eq!(HEADER_LEN, 12);
}
#[test]
fn upgrade_reject_round_trips() {
let r = UpgradeReject::new(
roles::PRIMARY_CATCHUP,
roles::NAME_PRIMARY_CATCHUP,
Some("eu-west-1a".into()),
);
assert_eq!(r.role_byte, roles::PRIMARY_CATCHUP);
assert_eq!(r.zone.as_deref(), Some("eu-west-1a"));
assert!(r.is_transient());
}
#[test]
fn upgrade_reject_topological_for_non_catchup_roles() {
for (byte, name) in [
(roles::STANDALONE, roles::NAME_STANDALONE),
(roles::PRIMARY, roles::NAME_PRIMARY),
(roles::REPLICA, roles::NAME_REPLICA),
(0x99, "FUTURE_ROLE"),
] {
let r = UpgradeReject::new(byte, name, None);
assert!(!r.is_transient(), "role {} should be topological", name);
}
}
#[test]
fn upgrade_reject_is_transient_case_insensitive() {
let r = UpgradeReject::new(0x99, "primary_catchup", None);
assert!(r.is_transient(), "case-insensitive match per spec §5");
}
}