use super::stream::MaybeTlsStream;
use super::TRANSPORT_BUFFER_SIZE;
use super::{ExaRowReader, ExaRowWriter, HttpTransportConfig, TransportResult};
use crate::error::HttpTransportError;
use crossbeam::channel::{Receiver, Sender};
use log::debug;
use std::io::{BufRead, BufReader, Read, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
pub mod reader;
pub mod writer;
const SPECIAL_PACKET: [u8; 12] = [2, 33, 33, 2, 1, 0, 0, 0, 1, 0, 0, 0];
const END_PACKET: &[u8; 5] = b"0\r\n\r\n";
const SUCCESS_HEADERS: &[u8; 66] = b"HTTP/1.1 200 OK\r\n\
Connection: close\r\n\
Transfer-Encoding: chunked\r\n\
\r\n";
const ERROR_HEADERS: &[u8; 57] = b"HTTP/1.1 500 Internal Server Error\r\n\
Connection: close\r\n\
\r\n";
pub struct HttpExportThread {
data_handler: Sender<Vec<u8>>,
}
impl HttpTransportWorker for HttpExportThread {
type Channel = Sender<Vec<u8>>;
fn new(data_handler: Self::Channel) -> Self {
Self { data_handler }
}
fn process_data(
&mut self,
stream: &mut MaybeTlsStream,
run: &Arc<AtomicBool>,
compression: bool,
) -> TransportResult<()> {
let mut reader = BufReader::new(stream);
Self::skip_headers(&mut reader, run)?;
let mut processor = ExaRowReader::new(reader, compression);
let mut num_bytes = 1;
while num_bytes > 0 && run.load(Ordering::Acquire) {
let mut buf = Vec::with_capacity(TRANSPORT_BUFFER_SIZE);
num_bytes = processor.read_chunk(&mut buf)?;
self.data_handler
.send(buf)
.map_err(|_| HttpTransportError::SendError)?;
}
Ok(())
}
fn success(stream: &mut MaybeTlsStream) -> TransportResult<()> {
stream
.write_all(SUCCESS_HEADERS)
.and(stream.write_all(END_PACKET))
.map_err(HttpTransportError::IoError)
}
fn error(
stream: &mut MaybeTlsStream,
error: HttpTransportError,
run: &Arc<AtomicBool>,
) -> TransportResult<()> {
Self::stop(run);
stream.write_all(ERROR_HEADERS).ok();
Err(error)
}
}
pub struct HttpImportThread {
data_handler: Receiver<Vec<u8>>,
}
impl HttpTransportWorker for HttpImportThread {
type Channel = Receiver<Vec<u8>>;
fn new(data_handler: Self::Channel) -> Self {
Self { data_handler }
}
fn process_data(
&mut self,
stream: &mut MaybeTlsStream,
run: &Arc<AtomicBool>,
compression: bool,
) -> TransportResult<()> {
let mut reader = BufReader::new(stream);
Self::skip_headers(&mut reader, run)?;
let stream = reader.into_inner();
stream.write_all(SUCCESS_HEADERS)?;
let mut processor = ExaRowWriter::new(stream, compression);
for chunk in &self.data_handler {
processor.write_chunk(chunk)?
}
Ok(())
}
fn success(stream: &mut MaybeTlsStream) -> TransportResult<()> {
stream
.write_all(END_PACKET)
.map_err(HttpTransportError::IoError)
}
fn error(
_stream: &mut MaybeTlsStream,
error: HttpTransportError,
run: &Arc<AtomicBool>,
) -> TransportResult<()> {
Self::stop(run);
Err(error)
}
}
pub trait HttpTransportWorker {
type Channel: Clone + Send;
fn new(channel: Self::Channel) -> Self;
fn process_data(
&mut self,
stream: &mut MaybeTlsStream,
run: &Arc<AtomicBool>,
compression: bool,
) -> TransportResult<()>;
fn success(stream: &mut MaybeTlsStream) -> TransportResult<()>;
fn error(
stream: &mut MaybeTlsStream,
error: HttpTransportError,
run: &Arc<AtomicBool>,
) -> TransportResult<()>;
fn stop(run: &Arc<AtomicBool>) {
run.store(false, Ordering::Release);
}
fn start(&mut self, mut config: HttpTransportConfig) -> TransportResult<()> {
debug!("Worker starting...");
let timeout = config.take_timeout();
let socket = Self::initialize(
config.server_addr.as_str(),
&mut config.addr_sender,
timeout,
);
config.barrier.wait();
let socket = socket
.and_then(|stream| Self::promote(stream, config.encryption))
.and_then(|stream| self.transport(stream, config.run, config.compression));
config.barrier.wait();
socket.map(|_| ())
}
fn initialize<A>(
server_addr: A,
addr_sender: &mut Sender<String>,
timeout: Option<Duration>,
) -> TransportResult<TcpStream>
where
A: ToSocketAddrs,
{
let mut stream = TcpStream::connect(server_addr)?;
stream.set_read_timeout(timeout)?;
stream.set_write_timeout(timeout)?;
stream.write_all(&SPECIAL_PACKET)?;
stream.flush()?;
let mut buf = [0; 24];
stream.read_exact(&mut buf)?;
addr_sender
.send(Self::parse_address(buf)?)
.map_err(|_| HttpTransportError::SendError)?;
Ok(stream)
}
fn transport(
&mut self,
mut stream: MaybeTlsStream,
run: Arc<AtomicBool>,
compression: bool,
) -> TransportResult<MaybeTlsStream> {
let res = match run.load(Ordering::Acquire) {
false => Err(HttpTransportError::ThreadError),
true => self.process_data(&mut stream, &run, compression),
};
match res {
Ok(_) => Self::success(&mut stream),
Err(e) => Self::error(&mut stream, e, &run),
}?;
stream.flush()?;
Ok(stream)
}
fn skip_headers<R>(
mut reader: R,
run: &Arc<AtomicBool>,
) -> std::result::Result<(), std::io::Error>
where
R: BufRead,
{
let mut line = String::new();
while run.load(Ordering::Acquire) && line != "\r\n" {
line.clear();
reader.read_line(&mut line)?;
}
Ok(())
}
fn promote(stream: TcpStream, encryption: bool) -> TransportResult<MaybeTlsStream> {
MaybeTlsStream::wrap(stream, encryption)
}
fn parse_address(buf: [u8; 24]) -> TransportResult<String> {
let port_bytes = <[u8; 4]>::try_from(&buf[4..8])?;
let port = u32::from_le_bytes(port_bytes);
let mut ipaddr = String::with_capacity(16);
buf[8..]
.iter()
.take_while(|b| **b != b'\0')
.for_each(|b| ipaddr.push(char::from(*b)));
Ok(format!("{}:{}", ipaddr, port))
}
}