use std::io;
use std::sync::Arc;
use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName};
use rustls::{ClientConfig, RootCertStore, ServerConfig};
use tokio::net::{TcpListener, TcpStream};
use tokio_rustls::{TlsAcceptor, TlsConnector};
use tokio::io::{AsyncWriteExt, split};
use crate::errors::{Error, Result};
use crate::server::{ServerHandle, RegistryEntry};
fn ensure_provider() {
let _ = rustls::crypto::ring::default_provider().install_default();
}
#[derive(Clone, Default)]
pub struct TlsConfig {
pub cert_pem: Vec<u8>,
pub key_pem: Vec<u8>,
pub ca_pem: Option<Vec<u8>>,
}
impl TlsConfig {
pub(crate) fn build_server_config(&self) -> io::Result<Arc<ServerConfig>> {
let certs = load_certs(&self.cert_pem)?;
let key = load_private_key(&self.key_pem)?;
if let Some(ca) = &self.ca_pem {
let ca_certs = load_certs(ca)?;
let mut roots = RootCertStore::empty();
for cert in ca_certs {
roots.add(cert)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
}
let verifier = rustls::server::WebPkiClientVerifier::builder(Arc::new(roots))
.build()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
let cfg = ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_single_cert(certs, key)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
Ok(Arc::new(cfg))
} else {
let cfg = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
Ok(Arc::new(cfg))
}
}
pub(crate) fn build_client_config(&self) -> io::Result<Arc<ClientConfig>> {
let mut roots = RootCertStore::empty();
if let Some(ca) = &self.ca_pem {
for cert in load_certs(ca)? {
roots.add(cert)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
}
}
let builder = ClientConfig::builder().with_root_certificates(roots);
let cfg = if !self.cert_pem.is_empty() && !self.key_pem.is_empty() {
let certs = load_certs(&self.cert_pem)?;
let key = load_private_key(&self.key_pem)?;
builder
.with_client_auth_cert(certs, key)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?
} else {
builder.with_no_client_auth()
};
Ok(Arc::new(cfg))
}
pub async fn connect(self, addr: &str) -> Result<crate::Client> {
crate::client::Client::connect_tls_impl(addr.to_string(), self, false).await
}
pub async fn connect_with_reconnect(self, addr: &str) -> Result<crate::Client> {
crate::client::Client::connect_tls_impl(addr.to_string(), self, true).await
}
}
pub async fn serve_on_tls(addr: &str, cfg: TlsConfig) -> Result<ServerHandle> {
ensure_provider();
let server_cfg = cfg.build_server_config()
.map_err(|e| Error::Internal(e.to_string()))?;
let acceptor = TlsAcceptor::from(server_cfg);
let listener = TcpListener::bind(addr).await?;
let (tx, rx) = tokio::sync::watch::channel(false);
tokio::spawn(run_tls_accept_loop(listener, acceptor, rx));
Ok(ServerHandle { tx })
}
async fn run_tls_accept_loop(
listener: TcpListener,
acceptor: TlsAcceptor,
mut rx: tokio::sync::watch::Receiver<bool>,
) {
loop {
tokio::select! {
res = listener.accept() => {
match res {
Ok((tcp, _)) => {
let acceptor = acceptor.clone();
let shutdown = rx.clone();
tokio::spawn(async move {
match acceptor.accept(tcp).await {
Ok(tls_stream) => {
handle_tls_connection(tls_stream, shutdown).await;
}
Err(e) => {
eprintln!("callwire/tls: TLS handshake error: {e}");
}
}
});
}
Err(_) => break,
}
}
_ = rx.changed() => {
if *rx.borrow() { break; }
}
}
}
}
async fn handle_tls_connection(
stream: tokio_rustls::server::TlsStream<TcpStream>,
mut shutdown_rx: tokio::sync::watch::Receiver<bool>,
) {
let (mut reader, writer) = split(stream);
let writer = Arc::new(tokio::sync::Mutex::new(writer));
loop {
tokio::select! {
res = crate::framing::read_frame(&mut reader) => {
match res {
Ok(payload) => {
match crate::codec::unpack(&payload) {
Ok(msg) => {
let writer_clone = writer.clone();
tokio::spawn(async move {
dispatch_tls(writer_clone, msg).await;
});
}
Err(_) => { }
}
}
Err(_) => break,
}
}
_ = shutdown_rx.changed() => {
if *shutdown_rx.borrow() { break; }
}
}
}
let _ = writer.lock().await.shutdown().await;
}
type TlsWriteHalf = tokio::io::WriteHalf<tokio_rustls::server::TlsStream<TcpStream>>;
async fn dispatch_tls(writer: Arc<tokio::sync::Mutex<TlsWriteHalf>>, msg: crate::codec::WireMessage) {
let func_name = match &msg.func {
Some(f) => f.clone(),
None => {
let payload = crate::codec::pack_error(msg.id, "TypeError", "missing func field").unwrap();
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
return;
}
};
let entry = {
let reg = crate::server::REGISTRY.lock().unwrap();
reg.get(&func_name).cloned()
};
let Some(entry) = entry else {
let payload = crate::codec::pack_error(
msg.id,
"NotFoundError",
&format!("function '{}' not exported", func_name),
).unwrap();
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
return;
};
let args = msg.args.unwrap_or(rmpv::Value::Nil);
match entry {
RegistryEntry::Unary(handler) => {
match handler(args).await {
Ok(res) => {
if let Ok(payload) = crate::codec::pack_response(msg.id, &res) {
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
}
}
Err(err) => {
if let Ok(payload) = crate::codec::pack_error(msg.id, &err.error_type, &err.message) {
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
}
}
}
}
RegistryEntry::Stream(handler) => {
match handler(args).await {
Ok(mut stream) => {
use futures_util::StreamExt;
while let Some(res) = stream.next().await {
match res {
Ok(val) => {
if let Ok(payload) = crate::codec::pack_stream_chunk(msg.id, &val) {
let mut w = writer.lock().await;
if crate::framing::write_frame(&mut *w, &payload).await.is_err() {
return;
}
}
}
Err(err) => {
if let Ok(payload) = crate::codec::pack_error(msg.id, &err.error_type, &err.message) {
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
}
return;
}
}
}
if let Ok(payload) = crate::codec::pack_stream_end(msg.id) {
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
}
}
Err(err) => {
if let Ok(payload) = crate::codec::pack_error(msg.id, &err.error_type, &err.message) {
let mut w = writer.lock().await;
let _ = crate::framing::write_frame(&mut *w, &payload).await;
}
}
}
}
}
}
pub(crate) async fn dial_tls(addr: &str, cfg: &TlsConfig) -> io::Result<tokio_rustls::client::TlsStream<TcpStream>> {
ensure_provider();
let client_cfg = cfg.build_client_config()?;
let connector = TlsConnector::from(client_cfg);
let host = addr.split(':').next().unwrap_or("localhost");
let server_name: ServerName<'static> = host.to_string().try_into()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, format!("invalid server name: {e}")))?;
let tcp = TcpStream::connect(addr).await?;
let tls = connector.connect(server_name, tcp).await?;
Ok(tls)
}
fn load_certs(pem: &[u8]) -> io::Result<Vec<CertificateDer<'static>>> {
rustls_pemfile::certs(&mut io::BufReader::new(pem))
.collect::<io::Result<Vec<_>>>()
}
fn load_private_key(pem: &[u8]) -> io::Result<PrivateKeyDer<'static>> {
rustls_pemfile::private_key(&mut io::BufReader::new(pem))?
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "no private key found in PEM"))
}
#[cfg(test)]
pub fn gen_self_signed(san: &str) -> (Vec<u8>, Vec<u8>, Vec<u8>) {
use rcgen::{generate_simple_self_signed, CertifiedKey};
let CertifiedKey { cert, key_pair } =
generate_simple_self_signed(vec![san.to_owned()]).unwrap();
let cert_pem = cert.pem().into_bytes();
let key_pem = key_pair.serialize_pem().into_bytes();
let ca_pem = cert_pem.clone();
(cert_pem, key_pem, ca_pem)
}