use std::path::PathBuf;
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use tokio::io::{AsyncBufRead, AsyncBufReadExt, AsyncReadExt, AsyncWrite, AsyncWriteExt};
pub const PROTOCOL_VERSION: u32 = 1;
const MAX_LINE_BYTES: u64 = 64 * 1024;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Hello {
pub scryer: u32,
pub version: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub build: Option<String>,
#[serde(flatten)]
pub mode: HelloMode,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "mode", rename_all = "snake_case")]
pub enum HelloMode {
Mcp {
cwd: Option<PathBuf>,
#[serde(default)]
watch: bool,
},
Control { cmd: ControlCommand },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ControlCommand {
Status,
Stop,
StopIfIdle,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct HelloReply {
pub ok: bool,
pub version: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub build: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub status: Option<DaemonStatus>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DaemonStatus {
pub pid: u32,
pub version: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub build: Option<String>,
pub db_path: PathBuf,
pub sessions: usize,
pub uptime_secs: u64,
pub watched: Vec<WatchedProject>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WatchedProject {
pub project_id: u64,
pub root: PathBuf,
#[serde(default)]
pub directories: usize,
}
impl Hello {
pub fn new(mode: HelloMode) -> Self {
Self {
scryer: PROTOCOL_VERSION,
version: crate::mcp_version().to_string(),
build: Some(crate::build_info::build_id().to_string()),
mode,
}
}
}
impl HelloReply {
pub fn ok() -> Self {
Self {
ok: true,
version: crate::mcp_version().to_string(),
build: Some(crate::build_info::build_id().to_string()),
error: None,
status: None,
}
}
pub fn error(message: impl Into<String>) -> Self {
Self {
ok: false,
error: Some(message.into()),
..Self::ok()
}
}
}
pub async fn write_json_line<W, T>(writer: &mut W, value: &T) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
T: Serialize,
{
let mut line = serde_json::to_vec(value)?;
line.push(b'\n');
writer.write_all(&line).await?;
writer.flush().await?;
Ok(())
}
pub async fn read_json_line<R, T>(reader: &mut R) -> anyhow::Result<T>
where
R: AsyncBufRead + Unpin,
T: DeserializeOwned,
{
let mut line = Vec::new();
let read = (&mut *reader)
.take(MAX_LINE_BYTES)
.read_until(b'\n', &mut line)
.await?;
if read == 0 {
anyhow::bail!("connection closed before handshake");
}
if line.last() != Some(&b'\n') {
anyhow::bail!("handshake line is truncated or exceeds {MAX_LINE_BYTES} bytes");
}
Ok(serde_json::from_slice(&line)?)
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::BufReader;
#[tokio::test]
async fn handshake_preserves_trailing_bytes() {
let hello = Hello::new(HelloMode::Mcp {
cwd: Some(PathBuf::from("/repo")),
watch: true,
});
let mut wire = Vec::new();
write_json_line(&mut wire, &hello).await.unwrap();
wire.extend_from_slice(b"{\"jsonrpc\":\"2.0\"}\n");
let mut reader = BufReader::new(wire.as_slice());
let decoded: Hello = read_json_line(&mut reader).await.unwrap();
assert_eq!(decoded, hello);
let mut rest = String::new();
reader.read_to_string(&mut rest).await.unwrap();
assert_eq!(rest, "{\"jsonrpc\":\"2.0\"}\n");
}
#[test]
fn hello_wire_format_is_flat() {
let json = serde_json::to_value(Hello::new(HelloMode::Control {
cmd: ControlCommand::Stop,
}))
.unwrap();
assert_eq!(json["mode"], "control");
assert_eq!(json["cmd"], "stop");
assert_eq!(json["scryer"], PROTOCOL_VERSION);
}
}