use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use base64::Engine;
use base64::engine::general_purpose::STANDARD;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::mpsc;
use tokio::task::JoinSet;
use tokio::time::{MissedTickBehavior, interval, timeout};
use zeroize::Zeroizing;
use crate::error::{Error, Result};
use crate::frame::{Frame, FrameType};
use crate::handshake::{ClientHandshake, accept_client};
use crate::identity::Identity;
use crate::mux::{MuxHandle, MuxStream};
use crate::outbound::{connect_tcp, connect_udp};
use crate::protection::{ProtectionMode, SESSION_KEY_LEN, Side, StreamProtector};
use crate::wire_io::{read_frame, write_frame};
const RELAY_BUFFER: usize = 16 * 1024;
const CONNECT_TIMEOUT: Duration = Duration::from_secs(15);
#[derive(Clone)]
pub struct ClientSession {
mux: MuxHandle,
master: Arc
>,
mode: ProtectionMode,
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub struct Target {
pub host: TargetHost,
pub port: u16,
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub enum TargetHost {
Ip(IpAddr),
Domain(String),
}
pub struct ProtectedStream {
stream: MuxStream,
protector: StreamProtector,
}
pub struct ServerEstablished {
pub master: Zeroizing<[u8; SESSION_KEY_LEN]>,
pub identity: [u8; 32],
}
enum OpenRequest {
Control,
Tcp(Target),
Udp(Target),
}
impl Target {
pub fn new(host: TargetHost, port: u16) -> Result {
if port == 0 {
return Err(Error::Protocol("target port cannot be zero".to_owned()));
}
if let TargetHost::Domain(domain) = &host {
validate_domain(domain)?;
}
Ok(Self { host, port })
}
pub fn domain(domain: impl Into, port: u16) -> Result {
Self::new(TargetHost::Domain(domain.into()), port)
}
pub fn ip(ip: IpAddr, port: u16) -> Result {
Self::new(TargetHost::Ip(ip), port)
}
pub fn domain_name(&self) -> Option<&str> {
match &self.host {
TargetHost::Domain(domain) => Some(domain),
TargetHost::Ip(_) => None,
}
}
pub fn ip_address(&self) -> Option {
match self.host {
TargetHost::Ip(ip) => Some(ip),
TargetHost::Domain(_) => None,
}
}
}
impl ClientSession {
pub fn new(
mux: MuxHandle,
master: Zeroizing<[u8; SESSION_KEY_LEN]>,
mode: ProtectionMode,
) -> Self {
Self {
mux,
master: Arc::new(master),
mode,
}
}
pub async fn open_tcp(&self, target: &Target) -> Result {
let request = encode_open(&OpenRequest::Tcp(target.clone()))?;
self.open(request).await
}
pub async fn open_udp(&self, target: &Target) -> Result {
let request = encode_open(&OpenRequest::Udp(target.clone()))?;
self.open(request).await
}
async fn open_control(&self) -> Result {
self.open(encode_open(&OpenRequest::Control)?).await
}
async fn open(&self, request: Vec) -> Result {
let mut stream = self.mux.open().await?;
let mut protector = StreamProtector::new(&self.master, stream.id, Side::Client, self.mode)?;
let frame = protector.seal(FrameType::Open, 0, &request)?;
write_frame(&mut stream.io, &frame, self.mode.tag_len()).await?;
let response = read_frame(&mut stream.io, self.mode.tag_len()).await?;
let kind = response.kind;
let payload = protector.open(response)?;
match kind {
FrameType::OpenOk if payload.is_empty() => Ok(ProtectedStream { stream, protector }),
FrameType::Error => Err(Error::Carrier(
"remote endpoint rejected the stream".to_owned(),
)),
_ => Err(Error::Protocol(
"unexpected stream-open response".to_owned(),
)),
}
}
pub async fn monitor_heartbeat(&self, period: Duration) -> Result<()> {
if period.is_zero() {
return Err(Error::Config("heartbeat cannot be zero".to_owned()));
}
let mut control = self.open_control().await?;
let mut ticker = interval(period);
ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
ticker.tick().await;
loop {
ticker.tick().await;
let frame = control.protector.seal(FrameType::Heartbeat, 0, &[])?;
write_frame(&mut control.stream.io, &frame, self.mode.tag_len()).await?;
let response = timeout(
period,
read_frame(&mut control.stream.io, self.mode.tag_len()),
)
.await
.map_err(|_| Error::Carrier("heartbeat timed out".to_owned()))??;
let kind = response.kind;
let payload = control.protector.open(response)?;
if kind != FrameType::Heartbeat || !payload.is_empty() {
return Err(Error::Protocol("invalid heartbeat response".to_owned()));
}
}
}
pub async fn abort(&self) {
self.mux.abort().await;
}
pub async fn wait_closed(&self) {
self.mux.wait_closed().await;
}
}
impl ProtectedStream {
pub async fn send(&mut self, payload: &[u8]) -> Result<()> {
let frame = self.protector.seal(FrameType::Data, 0, payload)?;
write_frame(&mut self.stream.io, &frame, self.protector.tag_len()).await
}
/// Reads the next application payload from the stream. Used by the UDP
/// relay paths, where each `Data` frame carries exactly one datagram
/// rather than an arbitrary byte-stream chunk. `Ok(None)` means the
/// remote side closed the stream cleanly.
pub async fn recv(&mut self) -> Result