use serde::{Deserialize, Serialize};
use super::ids::AgentId;
pub const RELAY_MAGIC: [u8; 8] = *b"BMRELAY1";
pub const RELAY_PROTO_VER: u32 = 1;
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct RelayHello {
pub relay_proto_ver: u32,
pub root: std::path::PathBuf,
pub view: String,
pub agent: AgentId,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct RelayWelcome {
pub relay_proto_ver: u32,
pub daemon_version: String,
pub accepted: bool,
pub code: Option<String>,
}
pub fn encode<T: serde::Serialize>(msg: &T) -> Result<Vec<u8>, rmp_serde::encode::Error> {
rmp_serde::to_vec_named(msg)
}
pub fn decode<T: serde::de::DeserializeOwned>(bytes: &[u8]) -> Result<T, rmp_serde::decode::Error> {
rmp_serde::from_slice(bytes)
}
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use super::transport::MAX_FRAME_BYTES;
fn encode_io<T: serde::Serialize>(msg: &T) -> std::io::Result<Vec<u8>> {
encode(msg).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))
}
fn decode_io<T: serde::de::DeserializeOwned>(bytes: &[u8]) -> std::io::Result<T> {
decode(bytes).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))
}
async fn write_frame<W: AsyncWrite + Unpin>(writer: &mut W, body: &[u8]) -> std::io::Result<()> {
if body.len() > MAX_FRAME_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("relay frame {} exceeds MAX_FRAME_BYTES {MAX_FRAME_BYTES}", body.len()),
));
}
let len = u32::try_from(body.len()).expect("len <= MAX_FRAME_BYTES fits u32");
writer.write_all(&len.to_be_bytes()).await?;
writer.write_all(body).await?;
writer.flush().await?;
Ok(())
}
async fn read_frame<R: AsyncRead + Unpin>(reader: &mut R) -> std::io::Result<Vec<u8>> {
let mut len_bytes = [0u8; 4];
reader.read_exact(&mut len_bytes).await?;
let len = u32::from_be_bytes(len_bytes) as usize;
if len > MAX_FRAME_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("relay frame length {len} exceeds MAX_FRAME_BYTES {MAX_FRAME_BYTES}"),
));
}
let mut body = vec![0u8; len];
reader.read_exact(&mut body).await?;
Ok(body)
}
pub async fn client_handshake<S: AsyncRead + AsyncWrite + Unpin>(
stream: &mut S,
hello: &RelayHello,
) -> std::io::Result<RelayWelcome> {
stream.write_all(&RELAY_MAGIC).await?;
write_frame(stream, &encode_io(hello)?).await?;
let body = read_frame(stream).await?;
decode_io(&body)
}
pub async fn read_hello<R: AsyncRead + Unpin>(reader: &mut R) -> std::io::Result<RelayHello> {
let body = read_frame(reader).await?;
decode_io(&body)
}
pub async fn write_welcome<W: AsyncWrite + Unpin>(writer: &mut W, welcome: &RelayWelcome) -> std::io::Result<()> {
write_frame(writer, &encode_io(welcome)?).await
}
#[cfg(test)]
mod tests {
use super::super::transport::MAX_FRAME_BYTES;
use super::*;
#[test]
fn relay_hello_round_trips_through_msgpack() {
let hello = RelayHello {
relay_proto_ver: RELAY_PROTO_VER,
root: std::path::PathBuf::from("/repo/root"),
view: "main".to_string(),
agent: AgentId::parse("claude-code").expect("agent"),
};
let bytes = encode(&hello).expect("encode");
let back: RelayHello = decode(&bytes).expect("decode");
assert_eq!(hello, back);
}
#[test]
fn relay_welcome_round_trips_through_msgpack() {
let welcome = RelayWelcome {
relay_proto_ver: RELAY_PROTO_VER,
daemon_version: "0.22.6".to_string(),
accepted: false,
code: Some("relay_proto_skew".to_string()),
};
let bytes = encode(&welcome).expect("encode");
let back: RelayWelcome = decode(&bytes).expect("decode");
assert_eq!(welcome, back);
}
#[test]
fn relay_magic_is_eight_bytes() {
assert_eq!(RELAY_MAGIC.len(), 8);
}
#[test]
fn relay_magic_prefix_exceeds_max_frame_bytes() {
let prefix = [RELAY_MAGIC[0], RELAY_MAGIC[1], RELAY_MAGIC[2], RELAY_MAGIC[3]];
let as_len = u32::from_be_bytes(prefix) as usize;
assert!(
as_len > MAX_FRAME_BYTES,
"preamble prefix {as_len} must exceed MAX_FRAME_BYTES {MAX_FRAME_BYTES} for fallback"
);
}
#[tokio::test]
async fn handshake_round_trips_and_leaves_stream_at_first_rmcp_byte() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut client, mut server) = tokio::io::duplex(4096);
let hello = RelayHello {
relay_proto_ver: RELAY_PROTO_VER,
root: std::path::PathBuf::from("/repo/root"),
view: "main".to_string(),
agent: AgentId::parse("claude-code").expect("agent"),
};
let welcome = RelayWelcome {
relay_proto_ver: RELAY_PROTO_VER,
daemon_version: "9.9.9".to_string(),
accepted: true,
code: None,
};
let server_hello = hello.clone();
let server_welcome = welcome.clone();
let server_task = tokio::spawn(async move {
let mut magic = [0u8; RELAY_MAGIC.len()];
server.read_exact(&mut magic).await.expect("read magic");
assert_eq!(magic, RELAY_MAGIC);
let got = read_hello(&mut server).await.expect("read hello");
assert_eq!(got, server_hello);
write_welcome(&mut server, &server_welcome)
.await
.expect("write welcome");
server.write_all(b"{").await.expect("write first rmcp byte");
server.flush().await.expect("flush");
});
let got_welcome = client_handshake(&mut client, &hello).await.expect("handshake");
assert_eq!(got_welcome, welcome);
let mut first = [0u8; 1];
client.read_exact(&mut first).await.expect("read first rmcp byte");
assert_eq!(
&first, b"{",
"the byte after the welcome must be the untouched rmcp stream"
);
server_task.await.expect("server task");
}
}