use std::ffi::CString;
use std::fmt;
use std::io;
use std::net::SocketAddr;
use std::os::raw::{c_char, c_int, c_void};
use std::ptr::NonNull;
use std::time::Duration;
#[repr(C)]
struct NativeTlsStream {
_private: [u8; 0],
}
extern "C" {
fn resolved_tls_connect(
address: *const c_char,
port: u16,
scope_id: u32,
ifindex: c_int,
firewall_mark: u32,
server_name: *const c_char,
strict: c_int,
timeout_msec: u32,
ret: *mut *mut NativeTlsStream,
) -> c_int;
fn resolved_tls_set_timeout(stream: *mut NativeTlsStream, timeout_msec: u32) -> c_int;
fn resolved_tls_read(stream: *mut NativeTlsStream, buffer: *mut c_void, capacity: usize)
-> i64;
fn resolved_tls_write(
stream: *mut NativeTlsStream,
buffer: *const c_void,
length: usize,
) -> i64;
fn resolved_tls_free(stream: *mut NativeTlsStream);
}
pub struct TlsStream {
raw: NonNull<NativeTlsStream>,
}
impl fmt::Debug for TlsStream {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("TlsStream")
.field("raw", &self.raw)
.finish()
}
}
unsafe impl Send for TlsStream {}
impl TlsStream {
pub fn connect(
server: SocketAddr,
ifindex: Option<i32>,
firewall_mark: u32,
server_name: Option<&str>,
strict: bool,
timeout: Duration,
) -> io::Result<Self> {
let address = CString::new(server.ip().to_string()).map_err(|_| {
io::Error::new(io::ErrorKind::InvalidInput, "invalid DNS server address")
})?;
let server_name = server_name
.map(CString::new)
.transpose()
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "server name contains NUL"))?;
let scope_id = match server {
SocketAddr::V4(_) => 0,
SocketAddr::V6(address) => address.scope_id(),
};
let timeout_msec = duration_milliseconds(timeout);
let mut raw = std::ptr::null_mut();
let result = unsafe {
resolved_tls_connect(
address.as_ptr(),
server.port(),
scope_id,
ifindex.unwrap_or(0),
firewall_mark,
server_name
.as_ref()
.map_or(std::ptr::null(), |name| name.as_ptr()),
i32::from(strict),
timeout_msec,
&mut raw,
)
};
if result < 0 {
return Err(io::Error::from_raw_os_error(-result));
}
let raw = NonNull::new(raw).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
"TLS connector returned a null stream",
)
})?;
Ok(Self { raw })
}
pub fn set_timeout(&mut self, timeout: Duration) -> io::Result<()> {
let result =
unsafe { resolved_tls_set_timeout(self.raw.as_ptr(), duration_milliseconds(timeout)) };
native_result(result).map(|_| ())
}
pub fn write_all(&mut self, mut buffer: &[u8]) -> io::Result<()> {
while !buffer.is_empty() {
let written = unsafe {
resolved_tls_write(
self.raw.as_ptr(),
buffer.as_ptr().cast::<c_void>(),
buffer.len(),
)
};
let written = signed_result(written)?;
if written == 0 {
return Err(io::Error::new(
io::ErrorKind::WriteZero,
"TLS stream closed",
));
}
buffer = &buffer[written..];
}
Ok(())
}
pub fn read_exact(&mut self, mut buffer: &mut [u8]) -> io::Result<()> {
while !buffer.is_empty() {
let read = unsafe {
resolved_tls_read(
self.raw.as_ptr(),
buffer.as_mut_ptr().cast::<c_void>(),
buffer.len(),
)
};
let read = signed_result(read)?;
if read == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"TLS stream closed",
));
}
let (_, rest) = buffer.split_at_mut(read);
buffer = rest;
}
Ok(())
}
}
impl Drop for TlsStream {
fn drop(&mut self) {
unsafe { resolved_tls_free(self.raw.as_ptr()) };
}
}
fn duration_milliseconds(duration: Duration) -> u32 {
u32::try_from(duration.as_millis())
.unwrap_or(u32::MAX)
.max(1)
}
fn native_result(result: c_int) -> io::Result<c_int> {
if result < 0 {
Err(io::Error::from_raw_os_error(-result))
} else {
Ok(result)
}
}
fn signed_result(result: i64) -> io::Result<usize> {
if result < 0 {
let errno = i32::try_from(-result).unwrap_or(22);
Err(io::Error::from_raw_os_error(errno))
} else {
usize::try_from(result)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid TLS I/O length"))
}
}