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.
|
//! formats or enum variants.
|
||||||
|
|
||||||
use std::io::{Read, Write};
|
use std::io::{Read, Write};
|
||||||
use std::net::{Shutdown, TcpListener, TcpStream};
|
use std::net::{TcpListener, TcpStream};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use std::sync::Mutex;
|
use std::sync::Mutex;
|
||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
@ -79,6 +79,50 @@ pub fn recv_framed<R: Read, T: serde::de::DeserializeOwned>(reader: &mut R) -> R
|
||||||
Ok(msg)
|
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
|
/// Thread-safe [`PluginRequestSender`] implementation that forwards requests
|
||||||
/// over a single TCP stream to a request proxy server.
|
/// over a single TCP stream to a request proxy server.
|
||||||
pub struct ProxyPluginRequestSender {
|
pub struct ProxyPluginRequestSender {
|
||||||
|
|
@ -113,12 +157,15 @@ impl ProxyPluginRequestSender {
|
||||||
/// [`POLL_INTERVAL`] while waiting. This lets Python callers run
|
/// [`POLL_INTERVAL`] while waiting. This lets Python callers run
|
||||||
/// `py.check_signals()` so that `Ctrl+C`/`KeyboardInterrupt` is raised
|
/// `py.check_signals()` so that `Ctrl+C`/`KeyboardInterrupt` is raised
|
||||||
/// promptly instead of waiting for the full socket timeout.
|
/// 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(
|
pub fn request_with_poll(
|
||||||
&self,
|
&self,
|
||||||
req: PluginRequest,
|
req: PluginRequest,
|
||||||
poll: &mut dyn FnMut() -> Result<(), PluginRequestError>,
|
poll: &mut dyn FnMut() -> Result<(), PluginRequestError>,
|
||||||
) -> Result<PluginResponse, PluginRequestError> {
|
) -> Result<PluginResponse, PluginRequestError> {
|
||||||
use std::io::ErrorKind;
|
|
||||||
let mut stream = self
|
let mut stream = self
|
||||||
.stream
|
.stream
|
||||||
.lock()
|
.lock()
|
||||||
|
|
@ -130,21 +177,9 @@ impl ProxyPluginRequestSender {
|
||||||
stream
|
stream
|
||||||
.set_read_timeout(Some(POLL_INTERVAL))
|
.set_read_timeout(Some(POLL_INTERVAL))
|
||||||
.map_err(|e| PluginRequestError(e.to_string()))?;
|
.map_err(|e| PluginRequestError(e.to_string()))?;
|
||||||
let result = loop {
|
|
||||||
match recv_framed(&mut *stream) {
|
let result = recv_framed_with_poll(&mut *stream, poll);
|
||||||
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 _ = stream.set_read_timeout(original_timeout);
|
let _ = stream.set_read_timeout(original_timeout);
|
||||||
result
|
result
|
||||||
}
|
}
|
||||||
|
|
@ -331,4 +366,55 @@ mod tests {
|
||||||
.unwrap();
|
.unwrap();
|
||||||
assert!(matches!(resp, PluginResponse::Ok));
|
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