use crate::protocol::async_io::{ElefantAsyncRead, ElefantAsyncReadWrite, ElefantAsyncWrite};
use der::Decode;
use rustls::ClientConnection;
use sha2::Digest;
use std::io::{self, Read, Write};
use x509_cert::Certificate;
enum HashAlgorithm {
Sha256,
Sha384,
Sha512,
}
pub struct TlsStream<S> {
inner: S,
tls: ClientConnection,
write_buf: Vec<u8>,
}
impl<S: ElefantAsyncReadWrite> TlsStream<S> {
pub fn new(inner: S, tls: ClientConnection) -> Self {
Self {
inner,
tls,
write_buf: Vec::new(),
}
}
async fn drain_tls_writes(&mut self) -> io::Result<()> {
while self.tls.wants_write() {
self.write_buf.clear();
self.tls.write_tls(&mut self.write_buf)?;
self.inner.write_all(&self.write_buf).await?;
}
self.inner.flush().await?;
Ok(())
}
pub async fn handshake(&mut self) -> io::Result<()> {
let mut read_buf = [0u8; 4096];
loop {
self.drain_tls_writes().await?;
if !self.tls.is_handshaking() {
break;
}
let n = self.inner.read(&mut read_buf).await?;
if n == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"EOF during TLS handshake",
));
}
self.tls
.read_tls(&mut &read_buf[..n])
.expect("read_tls from slice cannot fail");
self.tls
.process_new_packets()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
}
Ok(())
}
pub async fn shutdown(&mut self) -> io::Result<()> {
self.tls.send_close_notify();
self.drain_tls_writes().await
}
pub fn channel_binding_data(&self) -> Option<Vec<u8>> {
let certs = self.tls.peer_certificates()?;
let end_entity_der = certs.first()?;
let hash_bytes = match Self::signature_hash_algorithm(end_entity_der.as_ref()) {
HashAlgorithm::Sha256 => sha2::Sha256::digest(end_entity_der.as_ref()).to_vec(),
HashAlgorithm::Sha384 => sha2::Sha384::digest(end_entity_der.as_ref()).to_vec(),
HashAlgorithm::Sha512 => sha2::Sha512::digest(end_entity_der.as_ref()).to_vec(),
};
Some(hash_bytes)
}
fn signature_hash_algorithm(der: &[u8]) -> HashAlgorithm {
let cert = Certificate::from_der(der).ok();
match cert {
Some(cert) => oid_to_hash(&cert.signature_algorithm.oid),
None => HashAlgorithm::Sha256, }
}
}
impl<S: ElefantAsyncReadWrite> ElefantAsyncRead for TlsStream<S> {
async fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
loop {
let read_result: io::Result<usize> = self.tls.reader().read(buf);
match read_result {
Ok(n) if n > 0 => return Ok(n),
Err(ref e) if e.kind() == io::ErrorKind::UnexpectedEof => return Ok(0),
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {}
Ok(_) => {} Err(e) => return Err(e),
}
self.drain_tls_writes().await?;
let mut read_buf = [0u8; 4096];
let n = self.inner.read(&mut read_buf).await?;
if n == 0 {
self.tls
.read_tls(&mut &[][..])
.expect("read_tls from empty slice cannot fail");
self.tls
.process_new_packets()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
let eof_result: io::Result<usize> = self.tls.reader().read(buf);
return match eof_result {
Ok(n) => Ok(n),
Err(ref e) if e.kind() == io::ErrorKind::UnexpectedEof => Ok(0),
Err(e) => Err(e),
};
}
self.tls
.read_tls(&mut &read_buf[..n])
.expect("read_tls from slice cannot fail");
self.tls
.process_new_packets()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
}
}
}
impl<S: ElefantAsyncReadWrite> ElefantAsyncWrite for TlsStream<S> {
async fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
self.tls.writer().write_all(buf)?;
self.drain_tls_writes().await
}
async fn flush(&mut self) -> io::Result<()> {
self.tls.writer().flush()?;
self.drain_tls_writes().await?;
self.inner.flush().await
}
}
pub enum TlsNegotiationResult<S> {
Tls(Box<TlsStream<S>>, Option<Vec<u8>>),
Declined(S),
}
pub async fn negotiate_tls<S: ElefantAsyncReadWrite>(
mut stream: S,
host: &str,
tls_config: &std::sync::Arc<rustls::ClientConfig>,
) -> Result<TlsNegotiationResult<S>, crate::ElefantClientError> {
use rustls_pki_types::ServerName;
let ssl_request: [u8; 8] = [
0x00, 0x00, 0x00, 0x08, 0x04, 0xd2, 0x16, 0x2f, ];
stream.write_all(&ssl_request).await?;
stream.flush().await?;
let mut response = [0u8; 1];
let n = stream.read(&mut response).await?;
if n == 0 {
return Err(crate::ElefantClientError::TlsError(
"server closed connection during SSL negotiation".into(),
));
}
match response[0] {
b'S' => {
let server_name = ServerName::try_from(host)
.map_err(|e| {
crate::ElefantClientError::TlsError(format!("invalid server name: {e}"))
})?
.to_owned();
let tls_conn = ClientConnection::new(tls_config.clone(), server_name).map_err(|e| {
crate::ElefantClientError::TlsError(format!("failed to create TLS connection: {e}"))
})?;
let mut tls_stream = TlsStream::new(stream, tls_conn);
tls_stream.handshake().await.map_err(|e| {
crate::ElefantClientError::TlsError(format!("TLS handshake failed: {e}"))
})?;
let channel_binding = tls_stream.channel_binding_data();
Ok(TlsNegotiationResult::Tls(
Box::new(tls_stream),
channel_binding,
))
}
b'N' => Ok(TlsNegotiationResult::Declined(stream)),
other => Err(crate::ElefantClientError::TlsError(format!(
"unexpected SSL response byte: 0x{other:02x}"
))),
}
}
fn oid_to_hash(oid: &der::oid::ObjectIdentifier) -> HashAlgorithm {
use der::oid::ObjectIdentifier as Oid;
const SHA256_RSA: Oid = Oid::new_unwrap("1.2.840.113549.1.1.11");
const SHA384_RSA: Oid = Oid::new_unwrap("1.2.840.113549.1.1.12");
const SHA512_RSA: Oid = Oid::new_unwrap("1.2.840.113549.1.1.13");
const ECDSA_SHA256: Oid = Oid::new_unwrap("1.2.840.10045.4.3.2");
const ECDSA_SHA384: Oid = Oid::new_unwrap("1.2.840.10045.4.3.3");
const ECDSA_SHA512: Oid = Oid::new_unwrap("1.2.840.10045.4.3.4");
if *oid == SHA384_RSA || *oid == ECDSA_SHA384 {
HashAlgorithm::Sha384
} else if *oid == SHA512_RSA || *oid == ECDSA_SHA512 {
HashAlgorithm::Sha512
} else if *oid == SHA256_RSA || *oid == ECDSA_SHA256 {
HashAlgorithm::Sha256
} else {
HashAlgorithm::Sha256
}
}