From 255c758d41a0c84133fb3e114bf2a92895ceec7e Mon Sep 17 00:00:00 2001 From: Steven Fackler Date: Sun, 14 Oct 2018 17:44:46 -0700 Subject: [PATCH] Add tokio-postgres-native-tls --- Cargo.toml | 1 + tokio-postgres-native-tls/Cargo.toml | 15 ++++ tokio-postgres-native-tls/src/lib.rs | 106 ++++++++++++++++++++++++++ tokio-postgres-native-tls/src/test.rs | 69 +++++++++++++++++ 4 files changed, 191 insertions(+) create mode 100644 tokio-postgres-native-tls/Cargo.toml create mode 100644 tokio-postgres-native-tls/src/lib.rs create mode 100644 tokio-postgres-native-tls/src/test.rs diff --git a/Cargo.toml b/Cargo.toml index d7a9186a..54c0fbbe 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -7,5 +7,6 @@ members = [ "postgres-openssl", "postgres-native-tls", "tokio-postgres", + "tokio-postgres-native-tls", "tokio-postgres-openssl", ] diff --git a/tokio-postgres-native-tls/Cargo.toml b/tokio-postgres-native-tls/Cargo.toml new file mode 100644 index 00000000..5e956202 --- /dev/null +++ b/tokio-postgres-native-tls/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "tokio-postgres-native-tls" +version = "0.1.0" +authors = ["Steven Fackler "] + +[dependencies] +bytes = "0.4" +futures = "0.1" +native-tls = "0.2" +tokio-io = "0.1" +tokio-tls = "0.2" +tokio-postgres = { version = "0.3", path = "../tokio-postgres" } + +[dev-dependencies] +tokio = "0.1.7" diff --git a/tokio-postgres-native-tls/src/lib.rs b/tokio-postgres-native-tls/src/lib.rs new file mode 100644 index 00000000..dee06f97 --- /dev/null +++ b/tokio-postgres-native-tls/src/lib.rs @@ -0,0 +1,106 @@ +extern crate bytes; +extern crate futures; +extern crate native_tls; +extern crate tokio_io; +extern crate tokio_postgres; +extern crate tokio_tls; + +#[cfg(test)] +extern crate tokio; + +use bytes::{Buf, BufMut}; +use futures::{Future, Poll}; +use std::error::Error; +use std::io::{self, Read, Write}; +use tokio_io::{AsyncRead, AsyncWrite}; +use tokio_postgres::tls::{Socket, TlsConnect, TlsStream}; + +#[cfg(test)] +mod test; + +pub struct TlsConnector { + connector: tokio_tls::TlsConnector, +} + +impl TlsConnector { + pub fn new() -> Result { + let connector = native_tls::TlsConnector::new()?; + Ok(TlsConnector::with_connector(connector)) + } + + pub fn with_connector(connector: native_tls::TlsConnector) -> TlsConnector { + TlsConnector { + connector: tokio_tls::TlsConnector::from(connector), + } + } +} + +impl TlsConnect for TlsConnector { + fn connect( + &self, + domain: &str, + socket: Socket, + ) -> Box, Error = Box> + Sync + Send> { + let f = self + .connector + .connect(domain, socket) + .map(|s| { + let s: Box = Box::new(SslStream(s)); + s + }).map_err(|e| { + let e: Box = Box::new(e); + e + }); + Box::new(f) + } +} + +struct SslStream(tokio_tls::TlsStream); + +impl Read for SslStream { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + self.0.read(buf) + } +} + +impl AsyncRead for SslStream { + unsafe fn prepare_uninitialized_buffer(&self, buf: &mut [u8]) -> bool { + self.0.prepare_uninitialized_buffer(buf) + } + + fn read_buf(&mut self, buf: &mut B) -> Poll + where + B: BufMut, + { + self.0.read_buf(buf) + } +} + +impl Write for SslStream { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.0.write(buf) + } + + fn flush(&mut self) -> io::Result<()> { + self.0.flush() + } +} + +impl AsyncWrite for SslStream { + fn shutdown(&mut self) -> Poll<(), io::Error> { + self.0.shutdown() + } + + fn write_buf(&mut self, buf: &mut B) -> Poll + where + B: Buf, + { + self.0.write_buf(buf) + } +} + +impl TlsStream for SslStream { + fn tls_server_end_point(&self) -> Option> { + self.0.get_ref().tls_server_end_point().unwrap_or(None) + } +} diff --git a/tokio-postgres-native-tls/src/test.rs b/tokio-postgres-native-tls/src/test.rs new file mode 100644 index 00000000..a860729a --- /dev/null +++ b/tokio-postgres-native-tls/src/test.rs @@ -0,0 +1,69 @@ +use futures::{Future, Stream}; +use native_tls::{self, Certificate}; +use tokio::runtime::current_thread::Runtime; +use tokio_postgres::{self, TlsMode}; + +use TlsConnector; + +fn smoke_test(url: &str, tls: TlsMode) { + let mut runtime = Runtime::new().unwrap(); + + let handshake = tokio_postgres::connect(url.parse().unwrap(), tls); + let (mut client, connection) = runtime.block_on(handshake).unwrap(); + let connection = connection.map_err(|e| panic!("{}", e)); + runtime.handle().spawn(connection).unwrap(); + + let prepare = client.prepare("SELECT 1::INT4"); + let statement = runtime.block_on(prepare).unwrap(); + let select = client.query(&statement, &[]).collect().map(|rows| { + assert_eq!(rows.len(), 1); + assert_eq!(rows[0].get::<_, i32>(0), 1); + }); + runtime.block_on(select).unwrap(); + + drop(statement); + drop(client); + runtime.run().unwrap(); +} + +#[test] +fn require() { + let connector = native_tls::TlsConnector::builder() + .add_root_certificate( + Certificate::from_pem(include_bytes!("../../test/server.crt")).unwrap(), + ).build() + .unwrap(); + let connector = TlsConnector::with_connector(connector); + smoke_test( + "postgres://ssl_user@localhost:5433/postgres", + TlsMode::Require(Box::new(connector)), + ); +} + +#[test] +fn prefer() { + let connector = native_tls::TlsConnector::builder() + .add_root_certificate( + Certificate::from_pem(include_bytes!("../../test/server.crt")).unwrap(), + ).build() + .unwrap(); + let connector = TlsConnector::with_connector(connector); + smoke_test( + "postgres://ssl_user@localhost:5433/postgres", + TlsMode::Prefer(Box::new(connector)), + ); +} + +#[test] +fn scram_user() { + let connector = native_tls::TlsConnector::builder() + .add_root_certificate( + Certificate::from_pem(include_bytes!("../../test/server.crt")).unwrap(), + ).build() + .unwrap(); + let connector = TlsConnector::with_connector(connector); + smoke_test( + "postgres://scram_user:password@localhost:5433/postgres", + TlsMode::Require(Box::new(connector)), + ); +}