use anyctx::AnyCtx;
use anyhow::Context;
use bytes::Bytes;
use futures_concurrency::future::Race as _;
use futures_util::{FutureExt, TryFutureExt, future::Shared, task::noop_waker};
use geph5_broker_protocol::UserInfo;
use geph5_misc_rpc::client_control::{ControlClient, ControlService};
use geph5_rt::Immortal;
use nanorpc::DynRpcTransport;
use sillad::Pipe;
use std::sync::Arc;
#[cfg(unix)]
use std::{
io::{Read, Write},
os::fd::{AsRawFd, FromRawFd},
};
#[cfg(unix)]
use tokio::io::{Interest, unix::AsyncFd};
use crate::{
auth::{auth_loop, get_auth_token},
broker::broker_client,
bw_token::bw_token_refresh_loop,
control_prot::{ControlProtocolImpl, DummyControlProtocolTransport},
http_proxy::http_proxy_serve,
logging,
pac::pac_serve,
port_forward::port_forward,
session::{open_conn, run_session},
socks5::socks5_loop,
vpn::{recv_vpn_packet, send_vpn_packet, vpn_loop},
};
pub use geph5_misc_rpc::client_config::{BrokerKeys, Config};
#[derive(Clone)]
pub struct Client {
task: Shared<geph5_rt::Task<Result<(), Arc<anyhow::Error>>>>,
ctx: AnyCtx<Config>,
}
impl Client {
pub fn start(cfg: Config) -> Self {
Self::start_with_vpn_fd(cfg, None)
}
#[cfg(unix)]
pub fn start_with_vpn_fd(cfg: Config, vpn_fd: Option<i32>) -> Self {
let ctx = AnyCtx::new(cfg.clone());
let _ = logging::init_logging(&ctx);
let ((fd_limit, _), _) = binary_search::binary_search((1, ()), (65536, ()), |lim| {
if rlimit::increase_nofile_limit(lim).unwrap_or_default() >= lim {
binary_search::Direction::Low(())
} else {
binary_search::Direction::High(())
}
});
tracing::info!("raised file descriptor limit to {}", fd_limit);
let client_ctx = ctx.clone();
let combined = async move {
let main_fut = client_main(ctx.clone());
match vpn_fd {
Some(fd) => (main_fut, run_vpn_fd_handler(ctx, fd)).race().await,
None => main_fut.await,
}
};
let task = geph5_rt::spawn(combined.map_err(Arc::new));
Client {
task: task.shared(),
ctx: client_ctx,
}
}
#[cfg(not(unix))]
pub fn start_with_vpn_fd(cfg: Config, _vpn_fd: Option<i32>) -> Self {
let ctx = AnyCtx::new(cfg.clone());
let _ = logging::init_logging(&ctx);
let ((fd_limit, _), _) = binary_search::binary_search((1, ()), (65536, ()), |lim| {
if rlimit::increase_nofile_limit(lim).unwrap_or_default() >= lim {
binary_search::Direction::Low(())
} else {
binary_search::Direction::High(())
}
});
tracing::info!("raised file descriptor limit to {}", fd_limit);
let task = geph5_rt::spawn(client_main(ctx.clone()).map_err(Arc::new));
Client {
task: task.shared(),
ctx,
}
}
pub async fn open_conn(&self, remote: &str) -> anyhow::Result<Box<dyn Pipe>> {
open_conn(&self.ctx, "tcp", remote).await
}
pub async fn wait_until_dead(self) -> anyhow::Result<()> {
self.task.await.map_err(|e| anyhow::anyhow!(e))
}
pub fn check_dead(&self) -> anyhow::Result<()> {
match self
.task
.clone()
.poll_unpin(&mut std::task::Context::from_waker(&noop_waker()))
{
std::task::Poll::Ready(val) => val.map_err(|e| anyhow::anyhow!(e))?,
std::task::Poll::Pending => {}
}
Ok(())
}
pub fn control_client(&self) -> ControlClient {
ControlClient(DynRpcTransport::new(DummyControlProtocolTransport(
ControlService(ControlProtocolImpl {
ctx: self.ctx.clone(),
}),
)))
}
pub async fn user_info(&self) -> anyhow::Result<UserInfo> {
let auth_token = get_auth_token(&self.ctx).await?;
let user_info = broker_client(&self.ctx)?
.get_user_info(auth_token)
.await??
.context("no such user")?;
Ok(user_info)
}
pub async fn send_vpn_packet(&self, bts: Bytes) -> anyhow::Result<()> {
send_vpn_packet(&self.ctx, bts).await;
Ok(())
}
pub async fn recv_vpn_packet(&self) -> anyhow::Result<Bytes> {
let packet = recv_vpn_packet(&self.ctx).await;
Ok(packet)
}
}
pub type CtxField<T> = fn(&AnyCtx<Config>) -> T;
#[cfg(unix)]
async fn run_vpn_fd_handler(ctx: AnyCtx<Config>, fd: i32) -> anyhow::Result<()> {
let file = unsafe { std::fs::File::from_raw_fd(fd) };
let flags = unsafe { libc::fcntl(file.as_raw_fd(), libc::F_GETFL) };
if flags < 0 {
return Err(std::io::Error::last_os_error()).context("could not get VPN fd flags");
}
if unsafe { libc::fcntl(file.as_raw_fd(), libc::F_SETFL, flags | libc::O_NONBLOCK) } < 0 {
return Err(std::io::Error::last_os_error()).context("could not make VPN fd nonblocking");
}
let async_fd = AsyncFd::new(file).context("could not register VPN fd with Tokio")?;
let read_task = async {
let mut buf = vec![0u8; 65535]; loop {
match async_fd
.async_io(Interest::READABLE, |mut file| file.read(&mut buf))
.await
{
Ok(n) if n > 0 => {
#[cfg(target_os = "macos")]
let pkt = {
if n <= 4 {
continue;
}
bytes::Bytes::copy_from_slice(&buf[4..n])
};
#[cfg(not(target_os = "macos"))]
let pkt = bytes::Bytes::copy_from_slice(&buf[..n]);
send_vpn_packet(&ctx, pkt).await;
}
Ok(0) => {
tracing::warn!("VPN fd reached EOF");
break;
}
Err(e) => {
tracing::error!("Error reading from VPN fd: {}", e);
break;
}
_ => break,
}
}
anyhow::Ok(())
};
let write_task = async {
loop {
let packet = recv_vpn_packet(&ctx).await;
#[cfg(target_os = "macos")]
let packet = {
let af: u32 = if packet.first().map(|b| b >> 4) == Some(6) {
30
} else {
2
};
let mut framed = Vec::with_capacity(4 + packet.len());
framed.extend_from_slice(&af.to_be_bytes());
framed.extend_from_slice(&packet);
bytes::Bytes::from(framed)
};
match async_fd
.async_io(Interest::WRITABLE, |mut file| file.write(&packet))
.await
{
Ok(written) if written == packet.len() => {}
Ok(written) => {
tracing::error!(
written,
expected = packet.len(),
"Partial packet write to VPN fd"
);
break;
}
Err(e) => {
tracing::error!("Error writing to VPN fd: {}", e);
break;
}
}
}
anyhow::Ok(())
};
let res = (read_task, write_task).race().await;
tracing::warn!("VPN fd handler exited");
res
}
async fn client_main(ctx: AnyCtx<Config>) -> anyhow::Result<()> {
let tcp_rpc_serve = async {
if let Some(control_listen) = ctx.init().control_listen {
nanorpc_sillad::rpc_serve(
sillad::tcp::TcpListener::bind(control_listen).await?,
ControlService(ControlProtocolImpl { ctx: ctx.clone() }),
)
.await?;
anyhow::Ok(())
} else {
std::future::pending().await
}
};
let unix_rpc_serve = async {
#[cfg(unix)]
if let Some(path) = ctx.init().control_listen_unix.as_ref() {
nanorpc_sillad::rpc_serve(
sillad::unix::UnixListener::bind(path).await?,
ControlService(ControlProtocolImpl { ctx: ctx.clone() }),
)
.await?;
return anyhow::Ok(());
}
std::future::pending().await
};
let pipe_rpc_serve = async {
#[cfg(windows)]
if let Some(name) = ctx.init().control_listen_pipe.as_ref() {
nanorpc_sillad::rpc_serve(
sillad::windows_pipe::NamedPipeListener::bind(name, None)?,
ControlService(ControlProtocolImpl { ctx: ctx.clone() }),
)
.await?;
return anyhow::Ok(());
}
std::future::pending().await
};
let rpc_serve = (tcp_rpc_serve, unix_rpc_serve, pipe_rpc_serve).race();
if ctx.init().dry_run {
rpc_serve.await
} else {
let vpn_loop = vpn_loop(&ctx);
let _client_loop = Immortal::spawn(run_session(ctx.clone()));
(
socks5_loop(&ctx)
.inspect_err(|e| tracing::error!(err = debug(e), "socks5 loop stopped")),
vpn_loop.inspect_err(|e| tracing::error!(err = debug(e), "vpn loop stopped")),
http_proxy_serve(&ctx)
.inspect_err(|e| tracing::error!(err = debug(e), "http proxy stopped")),
auth_loop(&ctx).inspect_err(|e| tracing::error!(err = debug(e), "auth loop stopped")),
bw_token_refresh_loop(&ctx)
.inspect_err(|e| tracing::error!(err = debug(e), "bw token loop stopped")),
rpc_serve,
pac_serve(&ctx),
port_forward(&ctx),
)
.race()
.await
}
}