1use std::io;
10use std::pin::Pin;
11use std::task::{Context, Poll};
12use tokio::io::{AsyncRead, ReadBuf};
13
14pub const MAX_MCP_LINE_BYTES: usize = 1024 * 1024;
18pub const MAX_TOOL_ARGUMENT_BYTES: usize = 256 * 1024;
20pub const MAX_REQUEST_STATE_BYTES: usize = 1024;
22pub const MAX_INPUT_RESPONSES_BYTES: usize = 64 * 1024;
24
25const MAX_DROP_NOTICES: u64 = 8;
28
29const CHUNK: usize = 8 * 1024;
30
31pub struct LineLimited<R> {
43 inner: R,
44 buf: Box<[u8]>,
46 fill: usize,
49 scanned: usize,
51 total: usize,
53 phase: Phase,
54 scratch: Box<[u8]>,
56 limit: usize,
57 dropped: u64,
58}
59
60#[derive(Clone, Copy)]
61enum Phase {
62 Filling,
64 Draining { drain: usize },
66 Discarding,
68 Eof,
70}
71
72impl<R: AsyncRead> LineLimited<R> {
73 pub fn new(inner: R) -> Self {
74 Self::with_limit(inner, MAX_MCP_LINE_BYTES)
75 }
76
77 pub fn with_limit(inner: R, limit: usize) -> Self {
78 let scratch_len = CHUNK.min(limit + 1);
81 Self {
82 inner,
83 buf: vec![0u8; limit + 1].into_boxed_slice(),
84 fill: 0,
85 scanned: 0,
86 total: 0,
87 phase: Phase::Filling,
88 scratch: vec![0u8; scratch_len].into_boxed_slice(),
89 limit,
90 dropped: 0,
91 }
92 }
93}
94
95impl<R: AsyncRead + Unpin> AsyncRead for LineLimited<R> {
96 fn poll_read(
97 mut self: Pin<&mut Self>,
98 cx: &mut Context<'_>,
99 out: &mut ReadBuf<'_>,
100 ) -> Poll<io::Result<()>> {
101 if out.remaining() == 0 {
104 return Poll::Ready(Ok(()));
105 }
106 let this = &mut *self;
107 loop {
108 match this.phase {
109 Phase::Eof => return Poll::Ready(Ok(())),
110 Phase::Draining { drain } => {
111 let n = (this.total - drain).min(out.remaining());
112 out.put_slice(&this.buf[drain..drain + n]);
115 if drain + n == this.total {
116 let leftover = this.fill - this.total;
119 if leftover > 0 {
120 this.buf.copy_within(this.total..this.fill, 0);
121 }
122 this.fill = leftover;
123 this.scanned = 0;
124 this.phase = Phase::Filling;
125 } else {
126 this.phase = Phase::Draining { drain: drain + n };
127 }
128 return Poll::Ready(Ok(()));
129 }
130 Phase::Discarding => {
131 let mut rb = ReadBuf::new(this.scratch.as_mut());
132 match Pin::new(&mut this.inner).poll_read(cx, &mut rb) {
133 Poll::Ready(Ok(())) => {
134 let n = rb.filled().len();
135 if n == 0 {
136 this.phase = Phase::Eof;
137 continue;
138 }
139 match rb.filled().iter().position(|b| *b == b'\n') {
140 Some(i) => {
141 let leftover = n - (i + 1);
145 if leftover > 0 {
146 this.scratch.copy_within(i + 1..n, 0);
147 this.buf[..leftover]
148 .copy_from_slice(&this.scratch[..leftover]);
149 }
150 this.fill = leftover;
151 this.scanned = 0;
152 this.phase = Phase::Filling;
153 }
154 None => continue,
155 }
156 }
157 Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
158 Poll::Pending => return Poll::Pending,
159 }
160 }
161 Phase::Filling => {
162 if this.scanned < this.fill {
167 if let Some(i) = this.buf[this.scanned..this.fill]
168 .iter()
169 .position(|b| *b == b'\n')
170 {
171 this.total = this.scanned + i + 1;
172 this.phase = Phase::Draining { drain: 0 };
173 continue;
174 }
175 this.scanned = this.fill;
176 }
177 if this.fill > this.limit {
178 this.dropped += 1;
181 if this.dropped <= MAX_DROP_NOTICES {
182 eprintln!(
183 "[sequel-mcp] dropped oversized stdin line (limit {} bytes)",
184 this.limit
185 );
186 }
187 this.fill = 0;
188 this.scanned = 0;
189 this.phase = Phase::Discarding;
190 continue;
191 }
192 let mut rb = ReadBuf::new(&mut this.buf[this.fill..]);
193 match Pin::new(&mut this.inner).poll_read(cx, &mut rb) {
194 Poll::Ready(Ok(())) => {
195 let n = rb.filled().len();
196 if n == 0 {
197 if this.fill > 0 {
201 this.total = this.fill;
202 this.phase = Phase::Draining { drain: 0 };
203 } else {
204 this.phase = Phase::Eof;
205 }
206 continue;
207 }
208 this.fill += n;
209 }
210 Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
211 Poll::Pending => return Poll::Pending,
212 }
213 }
214 }
215 }
216 }
217}
218
219pub fn limited_stdio() -> (LineLimited<tokio::io::Stdin>, tokio::io::Stdout) {
222 (LineLimited::new(tokio::io::stdin()), tokio::io::stdout())
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228 use std::time::Duration;
229 use tokio::io::{AsyncReadExt, AsyncWriteExt};
230
231 async fn run_adapter(input: Vec<u8>, limit: usize) -> Vec<u8> {
232 let (mut client, server) = tokio::io::duplex(8 * CHUNK);
233 let mut adapter = LineLimited::with_limit(server, limit);
234 let writer = tokio::spawn(async move {
235 client.write_all(&input).await.unwrap();
236 client.shutdown().await.unwrap();
237 });
238 let mut out = Vec::new();
239 adapter.read_to_end(&mut out).await.unwrap();
240 writer.await.unwrap();
241 out
242 }
243
244 #[tokio::test]
245 async fn small_lines_pass_through_verbatim() {
246 let input = b"{\"a\":1}\n{\"b\":2}\n".to_vec();
247 let out = run_adapter(input.clone(), 64).await;
248 assert_eq!(out, input);
249 }
250
251 #[tokio::test]
252 async fn line_at_exact_limit_passes_with_terminator() {
253 let body = "x".repeat(63);
254 let input = format!("{body}\nnext\n").into_bytes();
255 let out = run_adapter(input.clone(), 64).await;
256 assert_eq!(out, input, "line of exactly limit bytes passes");
257 }
258
259 #[tokio::test]
260 async fn oversized_line_is_dropped_and_stream_recovers() {
261 let before = "ok1\n";
263 let big = format!("{}\n", "y".repeat(200));
264 let after = "ok2\n";
265 let input = format!("{before}{big}{after}").into_bytes();
266 let out = run_adapter(input, 64).await;
267 assert_eq!(out, b"ok1\nok2\n".to_vec());
268 }
269
270 #[tokio::test]
271 async fn no_newline_oversized_then_newline_then_valid() {
272 let input = format!("{}\nvalid\n", "z".repeat(500)).into_bytes();
273 let out = run_adapter(input, 64).await;
274 assert_eq!(out, b"valid\n".to_vec());
275 }
276
277 #[tokio::test]
278 async fn many_oversized_lines_in_sequence() {
279 let mut input = String::new();
280 for i in 0..5 {
281 input.push_str(&format!("{}{}\n", "w".repeat(100), i));
282 }
283 input.push_str("done\n");
284 let out = run_adapter(input.into_bytes(), 64).await;
285 assert_eq!(out, b"done\n".to_vec());
286 }
287
288 #[tokio::test]
289 async fn slow_byte_by_byte_oversized_sender_stays_bounded() {
290 let (mut client, server) = tokio::io::duplex(64);
293 let mut adapter = LineLimited::with_limit(server, 32);
294 let writer = tokio::spawn(async move {
295 for _ in 0..200 {
296 client.write_all(b"q").await.unwrap();
297 tokio::time::sleep(std::time::Duration::from_millis(1)).await;
298 }
299 client.write_all(b"\nrecover\n").await.unwrap();
300 client.shutdown().await.unwrap();
301 });
302 let mut out = Vec::new();
303 adapter.read_to_end(&mut out).await.unwrap();
304 writer.await.unwrap();
305 assert_eq!(out, b"recover\n".to_vec());
306 }
307
308 #[tokio::test]
309 async fn carried_over_line_is_forwarded_without_new_input() {
310 let (mut client, server) = tokio::io::duplex(64);
314 let mut adapter = LineLimited::with_limit(server, 64);
315 client.write_all(b"first\nsecond\n").await.unwrap();
316 let mut out1 = [0u8; 64];
317 let n1 = tokio::time::timeout(Duration::from_secs(2), adapter.read(&mut out1))
318 .await
319 .expect("first line must not stall")
320 .unwrap();
321 assert_eq!(&out1[..n1], b"first\n");
322 let mut out2 = [0u8; 64];
323 let n2 = tokio::time::timeout(Duration::from_secs(2), adapter.read(&mut out2))
324 .await
325 .expect("carried-over line must not wait for new input")
326 .unwrap();
327 assert_eq!(&out2[..n2], b"second\n");
328 let _ = &mut client;
329 }
330
331 #[tokio::test]
332 async fn chunk_boundary_lines_are_framed_correctly() {
333 let mut input = String::new();
335 for i in 0..40 {
336 input.push_str(&format!("line-{i:02}-{}\n", "p".repeat(500)));
337 }
338 let out = run_adapter(input.clone().into_bytes(), 1024).await;
339 assert_eq!(out, input.into_bytes());
340 }
341}