use std::net::ToSocketAddrs;
use futures_util::stream::StreamExt;
use hyper::{Body, Response, Request};
use hyper::service::{make_service_fn, service_fn};
use hyper::server::conn::Http;
use hyper::server::Builder;
use std::sync::Arc;
use tokio::net::TcpListener;
use thruster_app::app::App;
use thruster_core::context::Context;
use thruster_context::basic_hyper_context::HyperRequest;
use native_tls;
use native_tls::Identity;
use tokio;
use crate::thruster_server::ThrusterServer;
pub struct SSLHyperServer<T: 'static + Context + Send> {
app: App<HyperRequest, T>,
cert: Option<Vec<u8>>,
cert_pass: &'static str,
}
impl<T: 'static + Context + Send> SSLHyperServer<T> {
pub fn cert(&mut self, cert: Vec<u8>) {
self.cert = Some(cert);
}
pub fn cert_pass(&mut self, cert_pass: &'static str) {
self.cert_pass = cert_pass;
}
}
impl<T: Context<Response = Response<Body>> + Send> ThrusterServer for SSLHyperServer<T> {
type Context = T;
type Response = Response<Body>;
type Request = HyperRequest;
fn new(app: App<Self::Request, T>) -> Self {
SSLHyperServer {
app,
cert: None,
cert_pass: "",
}
}
fn start(self, host: &str, port: u16) {
let addr = (host, port).to_socket_addrs().unwrap().next().unwrap();
let arc_app = Arc::new(self.app);
let cert = self.cert.unwrap().clone();
let cert_pass = self.cert_pass;
let cert = Identity::from_pkcs12(&cert, cert_pass)
.expect("Could not decrypt p12 file");
let tls_acceptor =
tokio_tls::TlsAcceptor::from(
native_tls::TlsAcceptor::builder(cert)
.build()
.expect("Could not create TLS acceptor.")
);
let _arc_acceptor = Arc::new(tls_acceptor);
let mut rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let service = make_service_fn(|_| {
let app = arc_app.clone();
async {
Ok::<_, hyper::Error>(service_fn(move |req: Request<Body>| {
let matched = app.resolve_from_method_and_path(
&req.method().to_string(),
&req.uri().to_string()
);
let req = HyperRequest::new(req);
app.resolve(req, matched)
}))
}
});
let listener = TcpListener::bind(&addr).await.unwrap();
let incoming = listener.incoming();
let server = Builder
::new(hyper::server::accept::from_stream(incoming.filter_map(|socket| {
async {
match socket {
Ok(stream) => {
match _arc_acceptor.clone().accept(stream).await {
Ok(val) => Some(Ok::<_, hyper::Error>(val)),
Err(e) => {
println!("TLS error: {}", e);
None
}
}
},
Err(e) => {
println!("TCP socket error: {}", e);
None
}
}
}
})), Http::new())
.serve(service);
server.await?;
Ok::<_, hyper::Error>(())
});
}
}