use std::convert::TryInto;
use std::ffi::c_void;
use std::ptr::null;
use std::slice;
use std::sync::Arc;
use libc::size_t;
use rustls::server::{
AllowAnyAnonymousOrAuthenticatedClient, AllowAnyAuthenticatedClient, ClientCertVerifier,
ClientHello, NoClientAuth, ResolvesServerCert, ServerConfig, ServerConnection,
StoresServerSessions,
};
use rustls::sign::CertifiedKey;
use rustls::{
ProtocolVersion, SignatureScheme, SupportedCipherSuite, WantsVerifier, ALL_CIPHER_SUITES,
};
use crate::cipher::{
rustls_certified_key, rustls_client_cert_verifier, rustls_client_cert_verifier_optional,
rustls_supported_ciphersuite,
};
use crate::connection::{rustls_connection, Connection};
use crate::error::rustls_result::{InvalidParameter, NullParameter};
use crate::error::{map_error, rustls_result};
use crate::rslice::{rustls_slice_bytes, rustls_slice_slice_bytes, rustls_slice_u16, rustls_str};
use crate::session::{
rustls_session_store_get_callback, rustls_session_store_put_callback, SessionStoreBroker,
SessionStoreGetCallback, SessionStorePutCallback,
};
use crate::{
ffi_panic_boundary, try_arc_from_ptr, try_box_from_ptr, try_mut_from_ptr, try_ref_from_ptr,
try_slice, userdata_get, ArcCastPtr, BoxCastPtr, CastConstPtr, CastPtr,
};
pub struct rustls_server_config_builder {
_private: [u8; 0],
}
pub(crate) struct ServerConfigBuilder {
base: rustls::ConfigBuilder<ServerConfig, WantsVerifier>,
verifier: Arc<dyn ClientCertVerifier>,
cert_resolver: Option<Arc<dyn ResolvesServerCert>>,
session_storage: Option<Arc<dyn StoresServerSessions + Send + Sync>>,
alpn_protocols: Vec<Vec<u8>>,
ignore_client_order: Option<bool>,
}
impl CastPtr for rustls_server_config_builder {
type RustType = ServerConfigBuilder;
}
impl BoxCastPtr for rustls_server_config_builder {}
pub struct rustls_server_config {
_private: [u8; 0],
}
impl CastConstPtr for rustls_server_config {
type RustType = ServerConfig;
}
impl ArcCastPtr for rustls_server_config {}
impl rustls_server_config_builder {
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_new() -> *mut rustls_server_config_builder {
ffi_panic_boundary! {
let builder = ServerConfigBuilder {
base: rustls::ServerConfig::builder().with_safe_defaults(),
verifier: NoClientAuth::new(),
cert_resolver: None,
session_storage: None,
alpn_protocols: vec![],
ignore_client_order: None,
};
BoxCastPtr::to_mut_ptr(builder)
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_new_custom(
cipher_suites: *const *const rustls_supported_ciphersuite,
cipher_suites_len: size_t,
tls_versions: *const u16,
tls_versions_len: size_t,
builder_out: *mut *mut rustls_server_config_builder,
) -> rustls_result {
ffi_panic_boundary! {
let cipher_suites: &[*const rustls_supported_ciphersuite] = try_slice!(cipher_suites, cipher_suites_len);
let mut cs_vec: Vec<SupportedCipherSuite> = Vec::new();
for &cs in cipher_suites.iter() {
let cs = try_ref_from_ptr!(cs);
match ALL_CIPHER_SUITES.iter().find(|&acs| cs.eq(acs)) {
Some(scs) => cs_vec.push(*scs),
None => return InvalidParameter,
}
}
let tls_versions: &[u16] = try_slice!(tls_versions, tls_versions_len);
let mut versions = vec![];
for version_number in tls_versions {
let proto = ProtocolVersion::from(*version_number);
if proto == rustls::version::TLS12.version {
versions.push(&rustls::version::TLS12);
} else if proto == rustls::version::TLS13.version {
versions.push(&rustls::version::TLS13);
}
}
let result = rustls::ServerConfig::builder().with_cipher_suites(&cs_vec).with_safe_default_kx_groups().with_protocol_versions(&versions);
let base = match result {
Ok(new) => new,
Err(_) => return rustls_result::InvalidParameter,
};
let builder = ServerConfigBuilder {
base,
verifier: NoClientAuth::new(),
cert_resolver: None,
session_storage: None,
alpn_protocols: vec![],
ignore_client_order: None,
};
BoxCastPtr::set_mut_ptr(builder_out, builder);
rustls_result::Ok
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_client_verifier(
builder: *mut rustls_server_config_builder,
verifier: *const rustls_client_cert_verifier,
) {
ffi_panic_boundary! {
let builder: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
let verifier: Arc<AllowAnyAuthenticatedClient> = try_arc_from_ptr!(verifier);
builder.verifier = verifier;
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_client_verifier_optional(
builder: *mut rustls_server_config_builder,
verifier: *const rustls_client_cert_verifier_optional,
) {
ffi_panic_boundary! {
let builder: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
let verifier: Arc<AllowAnyAnonymousOrAuthenticatedClient> = try_arc_from_ptr!(verifier);
builder.verifier = verifier;
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_free(config: *mut rustls_server_config_builder) {
ffi_panic_boundary! {
BoxCastPtr::to_box(config);
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_ignore_client_order(
builder: *mut rustls_server_config_builder,
ignore: bool,
) -> rustls_result {
ffi_panic_boundary! {
let config: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
config.ignore_client_order = Some(ignore);
rustls_result::Ok
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_alpn_protocols(
builder: *mut rustls_server_config_builder,
protocols: *const rustls_slice_bytes,
len: size_t,
) -> rustls_result {
ffi_panic_boundary! {
let config: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
let protocols: &[rustls_slice_bytes] = try_slice!(protocols, len);
let mut vv: Vec<Vec<u8>> = Vec::new();
for p in protocols {
let v: &[u8] = try_slice!(p.data, p.len);
vv.push(v.to_vec());
}
config.alpn_protocols = vv;
rustls_result::Ok
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_certified_keys(
builder: *mut rustls_server_config_builder,
certified_keys: *const *const rustls_certified_key,
certified_keys_len: size_t,
) -> rustls_result {
ffi_panic_boundary! {
let builder: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
let keys_ptrs: &[*const rustls_certified_key] = try_slice!(certified_keys, certified_keys_len);
let mut keys: Vec<Arc<CertifiedKey>> = Vec::new();
for &key_ptr in keys_ptrs {
let certified_key: Arc<CertifiedKey> = try_arc_from_ptr!(key_ptr);
keys.push(certified_key);
}
builder.cert_resolver = Some(Arc::new(ResolvesServerCertFromChoices::new(&keys)));
rustls_result::Ok
}
}
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_build(
builder: *mut rustls_server_config_builder,
) -> *const rustls_server_config {
ffi_panic_boundary! {
let builder = try_box_from_ptr!(builder);
let base = builder.base.with_client_cert_verifier(builder.verifier);
let mut config = if let Some(r) = builder.cert_resolver {
base.with_cert_resolver(r)
} else {
return null();
};
if let Some(ss) = builder.session_storage {
config.session_storage = ss;
}
config.alpn_protocols = builder.alpn_protocols;
if let Some(ignore_client_order) = builder.ignore_client_order {
config.ignore_client_order = ignore_client_order;
}
ArcCastPtr::to_const_ptr(config)
}
}
}
impl rustls_server_config {
#[no_mangle]
pub extern "C" fn rustls_server_config_free(config: *const rustls_server_config) {
ffi_panic_boundary! {
rustls_server_config::free(config);
}
}
#[no_mangle]
pub extern "C" fn rustls_server_connection_new(
config: *const rustls_server_config,
conn_out: *mut *mut rustls_connection,
) -> rustls_result {
ffi_panic_boundary! {
let config: Arc<ServerConfig> = try_arc_from_ptr!(config);
let server_connection = match ServerConnection::new(config) {
Ok(sc) => sc,
Err(e) => return map_error(e),
};
let c = Connection::from_server(server_connection);
BoxCastPtr::set_mut_ptr(conn_out, c);
rustls_result::Ok
}
}
}
#[no_mangle]
pub extern "C" fn rustls_server_connection_get_sni_hostname(
conn: *const rustls_connection,
buf: *mut u8,
count: size_t,
out_n: *mut size_t,
) -> rustls_result {
ffi_panic_boundary! {
let conn: &Connection = try_ref_from_ptr!(conn);
if buf.is_null() {
return NullParameter
}
if out_n.is_null() {
return NullParameter
}
let server_connection = match conn.as_server() {
Some(s) => s,
_ => return rustls_result::InvalidParameter,
};
let sni_hostname = match server_connection.sni_hostname() {
Some(sni_hostname) => sni_hostname,
None => {
unsafe {
*out_n = 0;
}
return rustls_result::Ok
},
};
let len: usize = sni_hostname.len();
if len > count {
unsafe {
*out_n = 0
}
return rustls_result::InsufficientSize;
}
unsafe {
std::ptr::copy_nonoverlapping(sni_hostname.as_ptr(), buf, len);
*out_n = len;
}
rustls_result::Ok
}
}
struct ResolvesServerCertFromChoices {
choices: Vec<Arc<CertifiedKey>>,
}
impl ResolvesServerCertFromChoices {
pub fn new(choices: &[Arc<CertifiedKey>]) -> Self {
ResolvesServerCertFromChoices {
choices: Vec::from(choices),
}
}
}
impl ResolvesServerCert for ResolvesServerCertFromChoices {
fn resolve(&self, client_hello: ClientHello) -> Option<Arc<CertifiedKey>> {
for key in self.choices.iter() {
if key
.key
.choose_scheme(client_hello.signature_schemes())
.is_some()
{
return Some(key.clone());
}
}
None
}
}
#[repr(C)]
pub struct rustls_client_hello<'a> {
sni_name: rustls_str<'a>,
signature_schemes: rustls_slice_u16<'a>,
alpn: *const rustls_slice_slice_bytes<'a>,
}
impl<'a> CastPtr for rustls_client_hello<'a> {
type RustType = rustls_client_hello<'a>;
}
pub type rustls_client_hello_userdata = *mut c_void;
pub type rustls_client_hello_callback = Option<
unsafe extern "C" fn(
userdata: rustls_client_hello_userdata,
hello: *const rustls_client_hello,
) -> *const rustls_certified_key,
>;
type ClientHelloCallback = unsafe extern "C" fn(
userdata: rustls_client_hello_userdata,
hello: *const rustls_client_hello,
) -> *const rustls_certified_key;
struct ClientHelloResolver {
pub callback: ClientHelloCallback,
}
impl ClientHelloResolver {
pub fn new(callback: ClientHelloCallback) -> ClientHelloResolver {
ClientHelloResolver { callback }
}
}
impl ResolvesServerCert for ClientHelloResolver {
fn resolve(&self, client_hello: ClientHello) -> Option<Arc<CertifiedKey>> {
let sni_name: &str = {
match client_hello.server_name() {
Some(c) => c,
None => "",
}
};
let sni_name: rustls_str = match sni_name.try_into() {
Ok(r) => r,
Err(_) => return None,
};
let mapped_sigs: Vec<u16> = client_hello
.signature_schemes()
.iter()
.map(|s| s.get_u16())
.collect();
let alpn = match client_hello.alpn() {
Some(iter) => iter.collect(),
None => vec![],
};
let alpn = rustls_slice_slice_bytes { inner: &alpn };
let signature_schemes: rustls_slice_u16 = (&*mapped_sigs).into();
let hello = rustls_client_hello {
sni_name,
signature_schemes,
alpn: &alpn,
};
let cb = self.callback;
let userdata = match userdata_get() {
Ok(u) => u,
Err(_) => return None,
};
let key_ptr: *const rustls_certified_key = unsafe { cb(userdata, &hello) };
let certified_key: &CertifiedKey = try_ref_from_ptr!(key_ptr);
Some(Arc::new(certified_key.clone()))
}
}
unsafe impl Sync for ClientHelloResolver {}
unsafe impl Send for ClientHelloResolver {}
impl rustls_server_config_builder {
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_hello_callback(
builder: *mut rustls_server_config_builder,
callback: rustls_client_hello_callback,
) -> rustls_result {
ffi_panic_boundary! {
let callback: ClientHelloCallback = match callback {
Some(cb) => cb,
None => return rustls_result::NullParameter,
};
let builder: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
builder.cert_resolver = Some(Arc::new(ClientHelloResolver::new(
callback
)));
rustls_result::Ok
}
}
}
fn sigschemes(input: &[u16]) -> Vec<SignatureScheme> {
use rustls::SignatureScheme::*;
input
.iter()
.map(|n| match n {
0x0201 => RSA_PKCS1_SHA1,
0x0203 => ECDSA_SHA1_Legacy,
0x0401 => RSA_PKCS1_SHA256,
0x0403 => ECDSA_NISTP256_SHA256,
0x0501 => RSA_PKCS1_SHA384,
0x0503 => ECDSA_NISTP384_SHA384,
0x0601 => RSA_PKCS1_SHA512,
0x0603 => ECDSA_NISTP521_SHA512,
0x0804 => RSA_PSS_SHA256,
0x0805 => RSA_PSS_SHA384,
0x0806 => RSA_PSS_SHA512,
0x0807 => ED25519,
0x0808 => ED448,
n => SignatureScheme::Unknown(*n),
})
.collect()
}
#[no_mangle]
pub extern "C" fn rustls_client_hello_select_certified_key(
hello: *const rustls_client_hello,
certified_keys: *const *const rustls_certified_key,
certified_keys_len: size_t,
out_key: *mut *const rustls_certified_key,
) -> rustls_result {
ffi_panic_boundary! {
let hello = try_ref_from_ptr!(hello);
let schemes: Vec<SignatureScheme> = sigschemes(try_slice!(hello.signature_schemes.data, hello.signature_schemes.len));
if out_key.is_null() {
return NullParameter
}
let keys_ptrs: &[*const rustls_certified_key] = try_slice!(certified_keys, certified_keys_len);
for &key_ptr in keys_ptrs {
let key_ref: &CertifiedKey = try_ref_from_ptr!(key_ptr);
if key_ref.key.choose_scheme(&schemes).is_some() {
unsafe {
*out_key = key_ptr;
}
return rustls_result::Ok;
}
}
rustls_result::NotFound
}
}
impl rustls_server_config_builder {
#[no_mangle]
pub extern "C" fn rustls_server_config_builder_set_persistence(
builder: *mut rustls_server_config_builder,
get_cb: rustls_session_store_get_callback,
put_cb: rustls_session_store_put_callback,
) -> rustls_result {
ffi_panic_boundary! {
let get_cb: SessionStoreGetCallback = match get_cb {
Some(cb) => cb,
None => return rustls_result::NullParameter,
};
let put_cb: SessionStorePutCallback = match put_cb {
Some(cb) => cb,
None => return rustls_result::NullParameter,
};
let builder: &mut ServerConfigBuilder = try_mut_from_ptr!(builder);
builder.session_storage = Some(Arc::new(SessionStoreBroker::new(
get_cb, put_cb
)));
rustls_result::Ok
}
}
}
#[cfg(test)]
mod tests {
use std::ptr::null_mut;
use super::*;
#[test]
fn test_config_builder() {
let builder: *mut rustls_server_config_builder =
rustls_server_config_builder::rustls_server_config_builder_new();
let h1 = "http/1.1".as_bytes();
let h2 = "h2".as_bytes();
let alpn: Vec<rustls_slice_bytes> = vec![h1.into(), h2.into()];
rustls_server_config_builder::rustls_server_config_builder_set_alpn_protocols(
builder,
alpn.as_ptr(),
alpn.len(),
);
let config = rustls_server_config_builder::rustls_server_config_builder_build(builder);
{
let config2 = try_ref_from_ptr!(config);
assert_eq!(config2.alpn_protocols, vec![h1, h2]);
}
rustls_server_config::rustls_server_config_free(config);
}
#[test]
fn test_server_config_builder_new_empty() {
let builder: *mut rustls_server_config_builder =
rustls_server_config_builder::rustls_server_config_builder_new();
assert_eq!(
rustls_server_config_builder::rustls_server_config_builder_build(builder),
null()
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn test_server_connection_new() {
let builder: *mut rustls_server_config_builder =
rustls_server_config_builder::rustls_server_config_builder_new();
let cert_pem = include_str!("../localhost/cert.pem").as_bytes();
let key_pem = include_str!("../localhost/key.pem").as_bytes();
let mut certified_key: *const rustls_certified_key = null();
let result = rustls_certified_key::rustls_certified_key_build(
cert_pem.as_ptr(),
cert_pem.len(),
key_pem.as_ptr(),
key_pem.len(),
&mut certified_key,
);
if !matches!(result, rustls_result::Ok) {
panic!(
"expected RUSTLS_RESULT_OK from rustls_certified_key_build, got {:?}",
result
);
}
rustls_server_config_builder::rustls_server_config_builder_set_certified_keys(
builder,
&certified_key,
1,
);
let config = rustls_server_config_builder::rustls_server_config_builder_build(builder);
assert_ne!(config, null());
let mut conn: *mut rustls_connection = null_mut();
let result = rustls_server_config::rustls_server_connection_new(config, &mut conn);
if !matches!(result, rustls_result::Ok) {
panic!("expected RUSTLS_RESULT_OK, got {:?}", result);
}
assert_eq!(rustls_connection::rustls_connection_wants_read(conn), true);
assert_eq!(
rustls_connection::rustls_connection_wants_write(conn),
false
);
assert_eq!(
rustls_connection::rustls_connection_is_handshaking(conn),
true
);
let some_byte = 42u8;
let mut alpn_protocol: *const u8 = &some_byte;
let mut alpn_protocol_len: usize = 1;
rustls_connection::rustls_connection_get_alpn_protocol(
conn,
&mut alpn_protocol,
&mut alpn_protocol_len,
);
assert_eq!(alpn_protocol, null());
assert_eq!(alpn_protocol_len, 0);
assert_eq!(
rustls_connection::rustls_connection_get_negotiated_ciphersuite(conn),
null()
);
assert_eq!(
rustls_connection::rustls_connection_get_peer_certificate(conn, 0),
null()
);
assert_eq!(
rustls_connection::rustls_connection_get_protocol_version(conn),
0
);
rustls_connection::rustls_connection_free(conn);
}
}