use crate::http_client::conn::Conn;
use crate::stack::{Stack, pop, push};
use crate::value::Value;
use rustls::pki_types::ServerName;
use rustls::{ClientConfig, ClientConnection, RootCertStore, StreamOwned};
#[cfg(test)]
use std::sync::Mutex;
use std::sync::{Arc, LazyLock};
static TLS_CONFIG: LazyLock<Arc<ClientConfig>> = LazyLock::new(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
let mut roots = RootCertStore::empty();
roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let config = ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
Arc::new(config)
});
#[cfg(test)]
static TEST_TLS_CONFIG: LazyLock<Mutex<Option<Arc<ClientConfig>>>> =
LazyLock::new(|| Mutex::new(None));
#[cfg(test)]
pub(crate) fn install_test_tls_config(cfg: Arc<ClientConfig>) {
*TEST_TLS_CONFIG.lock().unwrap() = Some(cfg);
}
#[cfg(test)]
pub(crate) fn clear_test_tls_config() {
*TEST_TLS_CONFIG.lock().unwrap() = None;
}
fn current_tls_config() -> Arc<ClientConfig> {
#[cfg(test)]
if let Some(cfg) = TEST_TLS_CONFIG.lock().unwrap().as_ref() {
return cfg.clone();
}
TLS_CONFIG.clone()
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn patch_seq_tls_client(stack: Stack) -> Stack {
unsafe {
let (stack, host_val) = pop(stack);
let host = match host_val {
Value::String(s) => s,
_ => return push_failure(stack),
};
let (stack, sock_val) = pop(stack);
let socket_id = match sock_val {
Value::Int(id) => id as usize,
_ => return push_failure(stack),
};
let hostname = host.as_str_or_empty().to_string();
if hostname.is_empty() {
return push_failure(stack);
}
let ok = crate::tcp::upgrade_tcp_in_place(socket_id, |tcp| build_tls(tcp, hostname));
if !ok {
return push_failure(stack);
}
let stack = push(stack, Value::Int(socket_id as i64));
push(stack, Value::Bool(true))
}
}
const DEFAULT_TLS_HANDSHAKE_TIMEOUT_MS: u64 = 10_000;
static TLS_HANDSHAKE_TIMEOUT: LazyLock<std::time::Duration> = LazyLock::new(|| {
let ms = std::env::var("SEQ_TLS_HANDSHAKE_TIMEOUT_MS")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.filter(|n| *n > 0)
.unwrap_or(DEFAULT_TLS_HANDSHAKE_TIMEOUT_MS);
std::time::Duration::from_millis(ms)
});
#[cfg(test)]
static TLS_HANDSHAKE_TIMEOUT_OVERRIDE: Mutex<Option<std::time::Duration>> = Mutex::new(None);
#[cfg(test)]
pub(crate) fn set_test_tls_handshake_timeout(dur: Option<std::time::Duration>) {
*TLS_HANDSHAKE_TIMEOUT_OVERRIDE.lock().unwrap() = dur;
}
fn tls_handshake_timeout() -> std::time::Duration {
#[cfg(test)]
if let Some(dur) = *TLS_HANDSHAKE_TIMEOUT_OVERRIDE.lock().unwrap() {
return dur;
}
*TLS_HANDSHAKE_TIMEOUT
}
fn build_tls(
mut tcp: may::net::TcpStream,
hostname: String,
) -> Result<StreamOwned<ClientConnection, may::net::TcpStream>, ()> {
let handshake_timeout = Some(tls_handshake_timeout());
tcp.set_read_timeout(handshake_timeout).map_err(|_| ())?;
tcp.set_write_timeout(handshake_timeout).map_err(|_| ())?;
let server_name = ServerName::try_from(hostname).map_err(|_| ())?;
let mut conn = ClientConnection::new(current_tls_config(), server_name).map_err(|_| ())?;
conn.complete_io(&mut tcp).map_err(|_| ())?;
let _ = tcp.set_read_timeout(None);
let _ = tcp.set_write_timeout(None);
Ok(StreamOwned::new(conn, tcp))
}
pub(crate) fn dial_tls(tcp: may::net::TcpStream, hostname: String) -> Result<Conn, ()> {
let stream = build_tls(tcp, hostname)?;
Ok(Box::new(stream) as Conn)
}
unsafe fn push_failure(stack: Stack) -> Stack {
unsafe {
let stack = push(stack, Value::Int(0));
push(stack, Value::Bool(false))
}
}