From f85c33f04251d6d3ea8ebb9442af9da11c9c2764 Mon Sep 17 00:00:00 2001 From: int 80h Date: Sun, 19 Apr 2020 22:10:44 -0400 Subject: [PATCH] Added tls.rs and resolver for sni --- Cargo.toml | 1 + src/main.rs | 20 +++----------- src/tls.rs | 77 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 82 insertions(+), 16 deletions(-) create mode 100644 src/tls.rs diff --git a/Cargo.toml b/Cargo.toml index d2fcbde..7f819ae 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,6 +9,7 @@ edition = "2018" [dependencies] tokio = { version = "0.2", features = [ "net", "io-util", "rt-threaded" ] } tokio-rustls = "*" +rustls = "*" futures-util = "*" toml = "*" serde = "*" diff --git a/src/main.rs b/src/main.rs index c6eeb0f..e67d8fd 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,3 +1,4 @@ +#![allow(unused_imports)] #[macro_use] extern crate serde_derive; @@ -20,20 +21,7 @@ use std::error::Error; mod config; - -fn load_certs(path: &String) -> io::Result> { - certs(&mut BufReader::new(File::open(path)?)) - .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "invalid cert")) -} - -fn load_keys(path: &String) -> PrivateKey { - let keyfile = File::open(path) - .expect("cannot open private key file"); - let mut reader = BufReader::new(keyfile); - let key = pkcs8_private_keys(&mut reader) - .expect("file contains invalid rsa private key"); - return key[0].clone(); -} +mod tls; fn get_content(request: String) -> String { @@ -75,8 +63,8 @@ fn main() -> io::Result<()> { .next() .ok_or_else(|| io::Error::from(io::ErrorKind::AddrNotAvailable))?; - let certs = load_certs(&cfg.server[0].cert)?; - let keys = load_keys(&cfg.server[0].key); + let certs = tls::load_certs(&cfg.server[0].cert)?; + let keys = tls::load_key(&cfg.server[0].key); let mut runtime = runtime::Builder::new() .threaded_scheduler() diff --git a/src/tls.rs b/src/tls.rs new file mode 100644 index 0000000..a26fb15 --- /dev/null +++ b/src/tls.rs @@ -0,0 +1,77 @@ +use std::fs::File; +use std::io::{ self, BufReader }; +use std::error::Error; +use std::collections::HashMap; +use std::sync::Arc; + +use rustls::{ResolvesServerCert, SignatureScheme}; +use rustls::ClientHello; +use rustls::sign::CertifiedKey; +use rustls::sign::{Signer, SigningKey, RSASigningKey}; +use tokio_rustls::rustls::{ Certificate, NoClientAuth, PrivateKey, ServerConfig }; +use tokio_rustls::rustls::internal::pemfile::{ certs, pkcs8_private_keys }; + +use crate::config; + +pub fn load_certs(path: &String) -> io::Result> { + certs(&mut BufReader::new(File::open(path)?)) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "invalid cert")) +} + +pub fn load_key(path: &String) -> PrivateKey { + let keyfile = File::open(path) + .expect("cannot open private key file"); + let mut reader = BufReader::new(keyfile); + let key = pkcs8_private_keys(&mut reader) + .expect("file contains invalid rsa private key"); + return key[0].clone(); +} + +pub struct CertResolver { + map: HashMap> +} + +impl CertResolver { + pub fn from_config(cfg: config::Config) -> Result { + let mut map = HashMap::new(); + + for server in cfg.server.iter() { + let key = load_key(&server.key); + // .chain_err(|| format!("Failed to load private key from {}", https.key_file))?; + let certs = load_certs(&server.cert).unwrap(); + // .chain_err(|| format!("Failed to load certificate from {}", server.cert)); + /* + let signer: Arc> = Arc::new(Box::new( + CertifiedKey::new(&key).map_err(|_| format!("Failed to create signer for {}", server.url)) + )); + */ + let signing_key = RSASigningKey::new( + &key).unwrap(); + + let signing_key_boxed: Arc> = Arc::new( + Box::new(signing_key)); + map.insert(server.url.clone(), Box::new(rustls::sign::CertifiedKey::new( + certs, signing_key_boxed))); + } + + println!("Successfully loaded {} HTTPS configurations", map.len()); + + Ok(CertResolver{ map }) + } +} + +impl ResolvesServerCert for CertResolver { + fn resolve( + &self, + client_hello: ClientHello, + ) -> Option { + if let Some(server_name) = client_hello.server_name() { + if let Some(cert) = self.map.get(server_name.into()) { + return Some(*cert.clone()) + } + } + + None + } +} +