extern crate libc;
use libc::{fcntl, F_GETFL, F_SETFL, O_NONBLOCK};
use std::io::{self, ErrorKind, Read};
use std::os::unix::io::{AsRawFd, RawFd};
pub struct NonBlockingReader<R: AsRawFd + Read> {
eof: bool,
reader: R,
}
impl<R: AsRawFd + Read> NonBlockingReader<R> {
pub fn from_fd(reader: R) -> io::Result<NonBlockingReader<R>> {
let fd = reader.as_raw_fd();
set_blocking(fd, false)?;
Ok(NonBlockingReader { reader, eof: false })
}
pub fn into_blocking(self) -> io::Result<R> {
let fd = self.reader.as_raw_fd();
set_blocking(fd, true)?;
Ok(self.reader)
}
pub fn is_eof(&self) -> bool {
self.eof
}
pub fn read_available(&mut self, buf: &mut Vec<u8>) -> io::Result<usize> {
let mut buf_len = 0;
loop {
let mut bytes = [0u8; 1024];
match self.reader.read(&mut bytes[..]) {
Ok(0) => {
self.eof = true;
break;
}
Err(ref err) if err.kind() == ErrorKind::WouldBlock => {
self.eof = false;
break;
}
Err(ref err) if err.kind() == ErrorKind::Interrupted => {}
Ok(len) => {
buf_len += len;
buf.append(&mut bytes[0..(len)].to_owned())
}
Err(err) => {
return Err(err);
}
}
}
Ok(buf_len)
}
pub fn read_available_to_string(&mut self, buf: &mut String) -> io::Result<usize> {
let mut byte_buf: Vec<u8> = Vec::with_capacity(1024);
let res = self.read_available(&mut byte_buf);
match String::from_utf8(byte_buf) {
Ok(utf8_buf) => {
buf.push_str(&utf8_buf);
res
}
Err(err) => {
let _ = res?;
Err(io::Error::new(ErrorKind::InvalidData, err))
}
}
}
}
fn set_blocking(fd: RawFd, blocking: bool) -> io::Result<()> {
let flags = unsafe { fcntl(fd, F_GETFL, 0) };
if flags < 0 {
return Err(io::Error::last_os_error());
}
let flags = if blocking {
flags & !O_NONBLOCK
} else {
flags | O_NONBLOCK
};
let res = unsafe { fcntl(fd, F_SETFL, flags) };
if res != 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::NonBlockingReader;
use std::io::Write;
use std::net::{TcpListener, TcpStream};
use std::sync::mpsc::channel;
use std::thread;
#[test]
fn it_works() {
let server = TcpListener::bind("127.0.0.1:34567").unwrap();
let (tx, rx) = channel();
thread::spawn(move || {
let (stream, _) = server.accept().unwrap();
tx.send(stream).unwrap();
});
let client = TcpStream::connect("127.0.0.1:34567").unwrap();
let mut stream = rx.recv().unwrap();
let mut nonblocking = NonBlockingReader::from_fd(client).unwrap();
let mut buf = Vec::new();
assert_eq!(nonblocking.read_available(&mut buf).unwrap(), 0);
assert_eq!(buf, b"");
stream.write_all(b"foo").unwrap();
assert_eq!(nonblocking.read_available(&mut buf).unwrap(), 3);
assert_eq!(buf, b"foo");
}
}