fix(ipc): read response byte-by-byte in request_with_poll to avoid TCP desync
This commit is contained in:
parent
52432af928
commit
4fd94d324a
1 changed files with 103 additions and 17 deletions
|
|
@ -11,7 +11,7 @@
|
|||
//! formats or enum variants.
|
||||
|
||||
use std::io::{Read, Write};
|
||||
use std::net::{Shutdown, TcpListener, TcpStream};
|
||||
use std::net::{TcpListener, TcpStream};
|
||||
use std::sync::Arc;
|
||||
use std::sync::Mutex;
|
||||
use std::time::Duration;
|
||||
|
|
@ -79,6 +79,50 @@ pub fn recv_framed<R: Read, T: serde::de::DeserializeOwned>(reader: &mut R) -> R
|
|||
Ok(msg)
|
||||
}
|
||||
|
||||
/// Read exactly `buf.len()` bytes from `reader`, calling `poll` whenever no
|
||||
/// data is currently available. Unlike `read_exact`, this preserves partial
|
||||
/// progress across short read timeouts so the stream can never get out of sync.
|
||||
fn read_with_poll<R: Read>(
|
||||
reader: &mut R,
|
||||
buf: &mut [u8],
|
||||
poll: &mut dyn FnMut() -> Result<(), PluginRequestError>,
|
||||
) -> Result<(), PluginRequestError> {
|
||||
use std::io::ErrorKind;
|
||||
let mut off = 0;
|
||||
while off < buf.len() {
|
||||
match reader.read(&mut buf[off..]) {
|
||||
Ok(0) => return Err(PluginRequestError("proxy disconnected".to_string())),
|
||||
Ok(n) => off += n,
|
||||
Err(e) if e.kind() == ErrorKind::WouldBlock || e.kind() == ErrorKind::TimedOut => {
|
||||
poll()?;
|
||||
}
|
||||
Err(e) => return Err(ProxyError::Io(e).into()),
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Receive a length-framed bincode message with polling. The length prefix and
|
||||
/// payload are read incrementally so a timeout never splits a message.
|
||||
fn recv_framed_with_poll<R: Read>(
|
||||
reader: &mut R,
|
||||
poll: &mut dyn FnMut() -> Result<(), PluginRequestError>,
|
||||
) -> Result<PluginResponse, PluginRequestError> {
|
||||
let mut len_buf = [0u8; 8];
|
||||
read_with_poll(reader, &mut len_buf, poll)?;
|
||||
let len = u64::from_le_bytes(len_buf) as usize;
|
||||
if len == 0 {
|
||||
return Err(ProxyError::Empty.into());
|
||||
}
|
||||
if len > MAX_MESSAGE_SIZE {
|
||||
return Err(ProxyError::TooLarge(len).into());
|
||||
}
|
||||
let mut buf = vec![0u8; len];
|
||||
read_with_poll(reader, &mut buf, poll)?;
|
||||
let msg = bincode::deserialize(&buf).map_err(ProxyError::Encode)?;
|
||||
Ok(msg)
|
||||
}
|
||||
|
||||
/// Thread-safe [`PluginRequestSender`] implementation that forwards requests
|
||||
/// over a single TCP stream to a request proxy server.
|
||||
pub struct ProxyPluginRequestSender {
|
||||
|
|
@ -113,12 +157,15 @@ impl ProxyPluginRequestSender {
|
|||
/// [`POLL_INTERVAL`] while waiting. This lets Python callers run
|
||||
/// `py.check_signals()` so that `Ctrl+C`/`KeyboardInterrupt` is raised
|
||||
/// promptly instead of waiting for the full socket timeout.
|
||||
///
|
||||
/// The response is read byte-by-byte with polling so that a short timeout
|
||||
/// can never leave the TCP stream in a partially-read state between the
|
||||
/// length prefix and the payload.
|
||||
pub fn request_with_poll(
|
||||
&self,
|
||||
req: PluginRequest,
|
||||
poll: &mut dyn FnMut() -> Result<(), PluginRequestError>,
|
||||
) -> Result<PluginResponse, PluginRequestError> {
|
||||
use std::io::ErrorKind;
|
||||
let mut stream = self
|
||||
.stream
|
||||
.lock()
|
||||
|
|
@ -130,21 +177,9 @@ impl ProxyPluginRequestSender {
|
|||
stream
|
||||
.set_read_timeout(Some(POLL_INTERVAL))
|
||||
.map_err(|e| PluginRequestError(e.to_string()))?;
|
||||
let result = loop {
|
||||
match recv_framed(&mut *stream) {
|
||||
Ok(resp) => break Ok(resp),
|
||||
Err(ProxyError::Io(e))
|
||||
if e.kind() == ErrorKind::WouldBlock
|
||||
|| e.kind() == ErrorKind::TimedOut =>
|
||||
{
|
||||
if let Err(e) = poll() {
|
||||
let _ = stream.shutdown(Shutdown::Both);
|
||||
break Err(e);
|
||||
}
|
||||
}
|
||||
Err(e) => break Err(e.into()),
|
||||
}
|
||||
};
|
||||
|
||||
let result = recv_framed_with_poll(&mut *stream, poll);
|
||||
|
||||
let _ = stream.set_read_timeout(original_timeout);
|
||||
result
|
||||
}
|
||||
|
|
@ -331,4 +366,55 @@ mod tests {
|
|||
.unwrap();
|
||||
assert!(matches!(resp, PluginResponse::Ok));
|
||||
}
|
||||
|
||||
/// A reader that returns one byte per call and then `WouldBlock` on the
|
||||
/// next call, exercising both the incremental read and the poll path.
|
||||
struct OneByteThenBlock<'a> {
|
||||
data: &'a [u8],
|
||||
block_next: bool,
|
||||
}
|
||||
|
||||
impl<'a> Read for OneByteThenBlock<'a> {
|
||||
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
|
||||
if self.block_next {
|
||||
self.block_next = false;
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::WouldBlock,
|
||||
"simulated block",
|
||||
));
|
||||
}
|
||||
if self.data.is_empty() {
|
||||
return Ok(0);
|
||||
}
|
||||
let n = buf.len().min(1);
|
||||
buf[..n].copy_from_slice(&self.data[..n]);
|
||||
self.data = &self.data[n..];
|
||||
self.block_next = true;
|
||||
Ok(n)
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn recv_framed_with_poll_reads_one_byte_at_a_time() {
|
||||
let msg = PluginResponse::Bool(true);
|
||||
let bytes = bincode::serialize(&msg).unwrap();
|
||||
let len = bytes.len() as u64;
|
||||
let mut framed = Vec::with_capacity(8 + bytes.len());
|
||||
framed.extend_from_slice(&len.to_le_bytes());
|
||||
framed.extend_from_slice(&bytes);
|
||||
|
||||
let mut reader = OneByteThenBlock {
|
||||
data: &framed,
|
||||
block_next: false,
|
||||
};
|
||||
let mut polls = 0;
|
||||
let resp = recv_framed_with_poll(&mut reader, &mut || {
|
||||
polls += 1;
|
||||
Ok(())
|
||||
})
|
||||
.unwrap();
|
||||
assert!(matches!(resp, PluginResponse::Bool(true)));
|
||||
// One byte, one block for every byte except the last (no poll after EOF).
|
||||
assert_eq!(polls, framed.len() - 1);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue