use std::io::ErrorKind;
use std::net::SocketAddr;
use clap::builder::PossibleValue;
use clap::ValueEnum;
use env_logger::Target;
use log::*;
use mio::net::UdpSocket;
use mio::Interest;
use mio::Registry;
use mio::Token;
use rustc_hash::FxHashMap;
use slab::Slab;
use tquic::PacketInfo;
use tquic::PacketSendHandler;
pub type Result<T> = std::result::Result<T, Box<dyn std::error::Error>>;
#[derive(Clone, Copy, Default, PartialEq, Debug)]
pub enum ApplicationProto {
Interop,
Http09,
#[default]
H3,
}
impl ApplicationProto {
pub fn from_slice(proto: &[u8]) -> Self {
match proto {
b"hq-interop" => Self::Interop,
b"http/0.9" => Self::Http09,
b"h3" => Self::H3,
_ => unreachable!(),
}
}
pub fn to_slice(&self) -> &[u8] {
match self {
Self::Interop => b"hq-interop",
Self::Http09 => b"http/0.9",
Self::H3 => b"h3",
}
}
pub fn convert_to_vec(protos: &[Self]) -> Vec<Vec<u8>> {
protos
.iter()
.map(|proto| proto.to_slice().to_vec())
.collect()
}
}
impl ValueEnum for ApplicationProto {
fn to_possible_value(&self) -> Option<PossibleValue> {
Some(match self {
Self::Interop => PossibleValue::new("hq-interop"),
Self::Http09 => PossibleValue::new("http/0.9"),
Self::H3 => PossibleValue::new("h3"),
})
}
fn value_variants<'a>() -> &'a [Self] {
&[Self::Interop, Self::Http09, Self::H3]
}
}
pub struct QuicSocket {
socks: Slab<UdpSocket>,
addrs: FxHashMap<SocketAddr, usize>,
local_addr: SocketAddr,
}
impl QuicSocket {
pub fn new(local: &SocketAddr, registry: &Registry) -> Result<Self> {
let mut socks = Slab::new();
let mut addrs = FxHashMap::default();
let socket = UdpSocket::bind(*local)?;
let local_addr = socket.local_addr()?;
let sid = socks.insert(socket);
addrs.insert(local_addr, sid);
let socket = socks.get_mut(sid).unwrap();
registry.register(socket, Token(sid), Interest::READABLE)?;
Ok(Self {
socks,
addrs,
local_addr,
})
}
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub fn add(&mut self, local: &SocketAddr, registry: &Registry) -> Result<SocketAddr> {
let socket = UdpSocket::bind(*local)?;
let local_addr = socket.local_addr()?;
let sid = self.socks.insert(socket);
self.addrs.insert(local_addr, sid);
let socket = self.socks.get_mut(sid).unwrap();
registry.register(socket, Token(sid), Interest::READABLE)?;
Ok(local_addr)
}
pub fn del(&mut self, local: &SocketAddr, registry: &Registry) -> Result<()> {
let sid = match self.addrs.get(local) {
Some(sid) => *sid,
None => return Ok(()),
};
let socket = match self.socks.get_mut(sid) {
Some(socket) => socket,
None => return Ok(()),
};
registry.deregister(socket)?;
self.socks.remove(sid);
Ok(())
}
pub fn recv_from(
&self,
buf: &mut [u8],
token: mio::Token,
) -> std::io::Result<(usize, SocketAddr, SocketAddr)> {
let socket = match self.socks.get(token.0) {
Some(socket) => socket,
None => return Err(std::io::Error::new(ErrorKind::Other, "invalid token")),
};
match socket.recv_from(buf) {
Ok((len, remote)) => Ok((len, socket.local_addr()?, remote)),
Err(e) => Err(e),
}
}
pub fn send_to(&self, buf: &[u8], src: SocketAddr, dst: SocketAddr) -> std::io::Result<usize> {
let sid = match self.addrs.get(&src) {
Some(sid) => sid,
None => {
debug!("send_to drop packet with unknown address {:?}", src);
return Ok(buf.len());
}
};
match self.socks.get(*sid) {
Some(socket) => Ok(socket.send_to(buf, dst)?),
None => {
debug!("send_to drop packet with unknown address {:?}", src);
Ok(buf.len())
}
}
}
}
impl PacketSendHandler for QuicSocket {
fn on_packets_send(&self, pkts: &[(Vec<u8>, PacketInfo)]) -> tquic::Result<usize> {
let mut count = 0;
for (pkt, info) in pkts {
if let Err(e) = self.send_to(pkt, info.src, info.dst) {
if e.kind() == std::io::ErrorKind::WouldBlock {
debug!("socket send would block");
return Ok(count);
}
return Err(tquic::Error::InvalidOperation(format!(
"socket send_to(): {:?}",
e
)));
}
debug!("written {} bytes", pkt.len());
count += 1;
}
Ok(count)
}
}
pub fn log_target(log_file: &Option<String>) -> Result<Target> {
if let Some(log_file) = log_file {
if let Ok(file) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(log_file)
{
return Ok(Target::Pipe(Box::new(file)));
}
return Err(format!("create log file {:?} failed", log_file).into());
}
Ok(Target::Stderr)
}