diff --git a/hive-c0re/src/socket_server/mod.rs b/hive-c0re/src/socket_server/mod.rs index 3777cc3e..aff0862f 100644 --- a/hive-c0re/src/socket_server/mod.rs +++ b/hive-c0re/src/socket_server/mod.rs @@ -15,7 +15,7 @@ use anyhow::{Context, Result}; use hive_core_agent_sock::{Request, Response}; use hive_sh4re::inbox::Message; use hive_sh4re::manager::MANAGER_AGENT; -use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; +use tokio::io::{AsyncBufRead, AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; use tokio::net::{UnixListener, UnixStream}; use tokio::task::JoinHandle; @@ -128,26 +128,77 @@ pub fn start_manager(coord: Arc) -> Result<()> { Ok(()) } +/// Longest request line accepted on the agent and manager sockets, newline +/// included. The largest request a client sends is an `OperatorMsg` from the +/// agent web UI's `/send` (`hive-agent/src/web_ui/actions.rs:21`), whose body +/// axum's default limit caps at 2 MiB; JSON escaping can roughly double that to +/// 4 MiB, and this is 4× that. +const REQUEST_LINE_MAX_BYTES: usize = 16 * 1024 * 1024; + +enum RequestLine { + Eof, + Complete, + TooLong, +} + +/// Read one request line into `buf`, consuming at most +/// [`REQUEST_LINE_MAX_BYTES`] from `reader`. +async fn read_request_line( + reader: &mut R, + buf: &mut Vec, +) -> std::io::Result { + buf.clear(); + let n = (&mut *reader) + .take(REQUEST_LINE_MAX_BYTES as u64) + .read_until(b'\n', buf) + .await?; + Ok(if n == 0 { + RequestLine::Eof + } else if n == REQUEST_LINE_MAX_BYTES && !buf.ends_with(b"\n") { + RequestLine::TooLong + } else { + RequestLine::Complete + }) +} + +/// Answer each request line on `stream` until EOF. An over-long line gets a +/// `Response::Err` and then the connection is closed, because the rest of the +/// line is still unread. async fn serve(stream: UnixStream, agent: String, coord: Arc) -> Result<()> { let (read, mut write) = stream.into_split(); let mut reader = BufReader::new(read); - let mut line = String::new(); + let mut line = Vec::new(); loop { - line.clear(); - let n = reader.read_line(&mut line).await?; - if n == 0 { - return Ok(()); - } - let resp = match serde_json::from_str::(line.trim()) { - Ok(req) => dispatch(&req, &agent, &coord).await, - Err(e) => Response::Err { - message: format!("parse error: {e}"), + let (resp, too_long) = match read_request_line(&mut reader, &mut line).await? { + RequestLine::Eof => return Ok(()), + RequestLine::TooLong => ( + Response::Err { + message: format!( + "request too long: no newline within {REQUEST_LINE_MAX_BYTES} bytes; \ + connection closed" + ), + }, + true, + ), + RequestLine::Complete => match serde_json::from_slice::(&line) { + Ok(req) => (dispatch(&req, &agent, &coord).await, false), + Err(e) => ( + Response::Err { + message: format!("parse error: {e}"), + }, + false, + ), }, }; let mut payload = serde_json::to_string(&resp)?; payload.push('\n'); write.write_all(payload.as_bytes()).await?; write.flush().await?; + if too_long { + anyhow::bail!( + "`{agent}` sent a request line over {REQUEST_LINE_MAX_BYTES} bytes; closed" + ); + } } } @@ -791,4 +842,55 @@ mod tests { let err = check_can_cancel_approval("nobody-with-no-groups").unwrap_err(); assert!(err.contains("approvals` tool group"), "{err}"); } + + #[tokio::test] + async fn over_long_line_is_rejected_after_reading_only_the_bound() { + let total = 2 * REQUEST_LINE_MAX_BYTES as u64; + let mut reader = BufReader::new(tokio::io::repeat(b'x').take(total)); + let mut buf = Vec::new(); + let outcome = read_request_line(&mut reader, &mut buf) + .await + .expect("read"); + // At most one `BufReader` fill past the bound left the source. + let unread = reader.into_inner().limit(); + let consumed = total - unread; + assert!( + consumed <= (REQUEST_LINE_MAX_BYTES + 8 * 1024) as u64, + "consumed {consumed} of {total} bytes" + ); + assert_eq!(buf.len(), REQUEST_LINE_MAX_BYTES); + assert!(matches!(outcome, RequestLine::TooLong)); + } + + #[tokio::test] + async fn serve_answers_an_over_long_line_with_an_error_and_closes() { + let (_dir, coord) = schedules::tests::coordinator(); + let (client, server) = UnixStream::pair().expect("socketpair"); + let handler = tokio::spawn(serve(server, "iris".to_owned(), coord)); + let (read, mut write) = client.into_split(); + let writer = tokio::spawn(async move { + let chunk = vec![b'x'; 64 * 1024]; + for _ in 0..2 * REQUEST_LINE_MAX_BYTES / chunk.len() { + if write.write_all(&chunk).await.is_err() { + return; + } + } + let _ = write.shutdown().await; + }); + let mut reply = String::new(); + BufReader::new(read) + .read_line(&mut reply) + .await + .expect("reply"); + let resp: Response = serde_json::from_str(&reply).expect("response json"); + assert!( + matches!(&resp, Response::Err { message } if message.starts_with("request too long")), + "{reply}" + ); + assert!( + handler.await.expect("join").is_err(), + "connection not closed" + ); + writer.await.expect("join"); + } } diff --git a/hive-c0re/src/socket_server/schedules.rs b/hive-c0re/src/socket_server/schedules.rs index d41d2dfe..64d018ad 100644 --- a/hive-c0re/src/socket_server/schedules.rs +++ b/hive-c0re/src/socket_server/schedules.rs @@ -302,7 +302,7 @@ fn schedule_to_wire(s: crate::scheduled_prompts::Schedule) -> hive_sh4re::schedu } #[cfg(test)] -mod tests { +pub(super) mod tests { use super::*; fn target(name: &str) -> hive_sh4re::schedule::WireScheduleTarget { @@ -377,10 +377,10 @@ mod tests { /// anybody's ancestor or descendant. const STRANGER: &str = "requester-with-no-relations"; - /// A real `Coordinator` over a throwaway sqlite dir. The schedule + /// A real `Coordinator` over a throwaway sqlite dir. The socket-server /// handlers take `&Arc`, so there is no lighter way in; /// `open` touches nothing outside the db path it is handed. - fn coordinator() -> (tempfile::TempDir, Arc) { + pub(in crate::socket_server) fn coordinator() -> (tempfile::TempDir, Arc) { let dir = tempfile::tempdir().expect("tempdir"); let coord = Coordinator::open( &dir.path().join("broker.sqlite"),