use crate::Result;
use crate::slirp::tcp_stream::{ConnState, Endpoints};
use crate::vtcp::segment::Segment;
use crate::vtcp::{Conn, ConnConfig};
use std::io::{Read, Write};
use std::net::{IpAddr, Shutdown, SocketAddr, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread;
use std::time::Duration;
const MSS_V4: u16 = 1460;
const MSS_V6: u16 = 1440;
pub(crate) struct TcpOutConn {
state: Arc<ConnState>,
remote: Mutex<Option<TcpStream>>,
closed: Arc<AtomicBool>,
synack: Mutex<Vec<Vec<u8>>>,
}
impl TcpOutConn {
pub(crate) fn accept_syn(
endpoints: Endpoints,
syn: &Segment,
remote: TcpStream,
sink: Arc<dyn Fn(&[u8]) + Send + Sync>,
) -> Result<Arc<TcpOutConn>> {
let remote_read = remote.try_clone()?;
let (local_addr, remote_addr, local_port, remote_port, mss) = match endpoints {
Endpoints::V4 {
local_ip,
local_port,
remote_ip,
remote_port,
} => (
SocketAddr::new(IpAddr::V4(local_ip), local_port),
SocketAddr::new(IpAddr::V4(remote_ip), remote_port),
local_port,
remote_port,
MSS_V4,
),
Endpoints::V6 {
local_ip,
local_port,
remote_ip,
remote_port,
} => (
SocketAddr::new(IpAddr::V6(local_ip), local_port),
SocketAddr::new(IpAddr::V6(remote_ip), remote_port),
local_port,
remote_port,
MSS_V6,
),
};
let cfg = ConnConfig {
local_addr: Some(local_addr),
remote_addr: Some(remote_addr),
local_port,
remote_port,
mss,
keepalive: true,
..Default::default()
};
let mut conn = Conn::new(cfg);
let synack = conn.accept_syn(syn);
let state = Arc::new(ConnState {
endpoints,
conn: Mutex::new(conn),
signal: Condvar::new(),
sink,
});
let bridge = Arc::new(TcpOutConn {
state: state.clone(),
remote: Mutex::new(Some(remote)),
closed: Arc::new(AtomicBool::new(false)),
synack: Mutex::new(synack),
});
let b_r = bridge.clone();
thread::spawn(move || b_r.pump_remote_to_client(remote_read));
let b_w = bridge.clone();
thread::spawn(move || b_w.pump_client_to_remote());
Ok(bridge)
}
pub(crate) fn send_synack(&self) {
let synack = std::mem::take(&mut *self.synack.lock().expect("poisoned"));
if !synack.is_empty() {
self.state.wrap_and_send(synack);
}
}
pub(crate) fn handle_segment(&self, tcp: &[u8]) -> Result<()> {
if let Ok(seg) = Segment::parse(tcp) {
self.state.deliver(&seg);
}
Ok(())
}
pub(crate) fn state(&self) -> &Arc<ConnState> {
&self.state
}
pub(crate) fn is_closed(&self) -> bool {
self.closed.load(Ordering::Acquire) || self.state.conn.lock().expect("poisoned").is_closed()
}
pub(crate) fn close(&self) {
if self.closed.swap(true, Ordering::AcqRel) {
return;
}
let segs = self.state.conn.lock().expect("poisoned").abort();
self.state.wrap_and_send(segs);
self.state.signal.notify_all();
self.shutdown_remote(Shutdown::Both);
}
fn shutdown_remote(&self, how: Shutdown) {
if let Ok(g) = self.remote.lock()
&& let Some(s) = g.as_ref()
{
let _ = s.shutdown(how);
}
}
fn pump_remote_to_client(self: Arc<Self>, mut remote_read: TcpStream) {
let mut buf = vec![0u8; 32 * 1024];
loop {
if self.closed.load(Ordering::Acquire) {
return;
}
let n = match remote_read.read(&mut buf) {
Ok(0) => break, Ok(n) => n,
Err(_) => break, };
if !self.write_all_to_engine(&buf[..n]) {
self.teardown_remote_read();
return;
}
}
let segs = {
let mut conn = self.state.conn.lock().expect("poisoned");
conn.close()
};
self.state.wrap_and_send(segs);
self.state.signal.notify_all();
}
fn write_all_to_engine(&self, mut data: &[u8]) -> bool {
while !data.is_empty() {
if self.closed.load(Ordering::Acquire) {
return false;
}
let (n, segs) = {
let mut conn = self.state.conn.lock().expect("poisoned");
if conn.is_closed() {
return false;
}
conn.write(data)
};
if n > 0 {
self.state.wrap_and_send(segs);
self.state.signal.notify_all();
data = &data[n..];
} else {
let conn = self.state.conn.lock().expect("poisoned");
if conn.is_closed() {
return false;
}
let _ = self
.state
.signal
.wait_timeout(conn, Duration::from_millis(100))
.expect("poisoned");
}
}
true
}
fn pump_client_to_remote(self: Arc<Self>) {
let mut buf = vec![0u8; 32 * 1024];
loop {
let (n, eof) = {
let mut conn = self.state.conn.lock().expect("poisoned");
let n = conn.read(&mut buf);
if n > 0 {
(n, false)
} else if conn.fin_received() || conn.is_closed() {
(0, true)
} else {
let _ = self
.state
.signal
.wait_timeout(conn, Duration::from_millis(100))
.expect("poisoned");
(0, false)
}
};
if self.closed.load(Ordering::Acquire) {
return;
}
if n > 0 {
let res = {
let mut g = self.remote.lock().expect("poisoned");
match g.as_mut() {
Some(s) => s.write_all(&buf[..n]),
None => Ok(()),
}
};
if res.is_err() {
self.close();
return;
}
} else if eof {
self.shutdown_remote(Shutdown::Write);
return;
}
}
}
fn teardown_remote_read(&self) {
self.shutdown_remote(Shutdown::Both);
self.closed.store(true, Ordering::Release);
self.state.signal.notify_all();
}
}
impl Drop for TcpOutConn {
fn drop(&mut self) {
self.closed.store(true, Ordering::Release);
if let Ok(g) = self.remote.lock()
&& let Some(s) = g.as_ref()
{
let _ = s.shutdown(Shutdown::Both);
}
}
}
pub(crate) fn build_rst_for_stray(tcp: &[u8], dst_port: u16, src_port: u16) -> Option<Vec<u8>> {
let seg = Segment::parse(tcp).ok()?;
let rst = if seg.has_flag(crate::vtcp::segment::flags::ACK) {
Segment {
src_port: dst_port,
dst_port: src_port,
seq: seg.ack,
flags: crate::vtcp::segment::flags::RST,
..Default::default()
}
} else {
let mut data_len = seg.data_len();
if seg.has_flag(crate::vtcp::segment::flags::FIN) {
data_len = data_len.wrapping_add(1);
}
Segment {
src_port: dst_port,
dst_port: src_port,
seq: 0,
ack: seg.seq.wrapping_add(data_len),
flags: crate::vtcp::segment::flags::RST | crate::vtcp::segment::flags::ACK,
..Default::default()
}
};
Some(rst.marshal())
}
pub(crate) fn build_refused_rst(src_port: u16, dst_port: u16, client_seq: u32) -> Vec<u8> {
Segment {
src_port: dst_port,
dst_port: src_port,
ack: client_seq.wrapping_add(1),
flags: crate::vtcp::segment::flags::RST | crate::vtcp::segment::flags::ACK,
..Default::default()
}
.marshal()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vtcp::segment::flags;
#[test]
fn build_rst_with_ack() {
let seg = Segment {
src_port: 5000,
dst_port: 80,
ack: 12345,
flags: flags::ACK,
..Default::default()
};
let rst = build_rst_for_stray(&seg.marshal(), 80, 5000).unwrap();
let parsed = Segment::parse(&rst).unwrap();
assert_eq!(parsed.flags, flags::RST);
assert_eq!(parsed.seq, 12345);
}
#[test]
fn build_rst_without_ack() {
let seg = Segment {
src_port: 5000,
dst_port: 80,
seq: 100,
flags: flags::SYN,
..Default::default()
};
let rst = build_rst_for_stray(&seg.marshal(), 80, 5000).unwrap();
let parsed = Segment::parse(&rst).unwrap();
assert_eq!(parsed.flags, flags::RST | flags::ACK);
assert_eq!(parsed.ack, 100);
}
#[test]
fn refused_rst_acks_syn() {
let rst = build_refused_rst(5000, 80, 1000);
let parsed = Segment::parse(&rst).unwrap();
assert_eq!(parsed.flags, flags::RST | flags::ACK);
assert_eq!(parsed.ack, 1001);
assert_eq!(parsed.src_port, 80);
assert_eq!(parsed.dst_port, 5000);
}
}