use crate::internal_alloc::Vec;
use crate::platform_bridge::noxtls_hmac_drbg_from_entropy;
use crate::protocol::{
CipherSuite, Connection, RecordContentType, TlsRecordDeframer, TlsVersion,
TLS_RECORD_HEADER_LEN,
};
type TransportResult<T> = core::result::Result<T, TransportError>;
use noxtls_core::{Error, Result as TlsResult};
use noxtls_io::transport::drive_async::{noxtls_read_record_async, noxtls_write_all_async};
use noxtls_io::transport::stream_async::AsyncByteStream;
use noxtls_io::transport::TransportError;
use noxtls_platform::EntropySource;
const TLS_RECORD_HANDSHAKE: u8 = 0x16;
const TLS_RECORD_APPLICATION_DATA: u8 = 0x17;
const TLS_RECORD_CHANGE_CIPHER_SPEC: u8 = 0x14;
fn tls_err<T>(result: TlsResult<T>) -> TransportResult<T> {
result.map_err(|_| TransportError::IoFailed("tls protocol error"))
}
pub struct AsyncTlsClientConfig {
pub cipher_suites: Vec<CipherSuite>,
}
impl Default for AsyncTlsClientConfig {
fn default() -> Self {
Self {
cipher_suites: vec![
CipherSuite::TlsAes256GcmSha384,
CipherSuite::TlsAes128GcmSha256,
],
}
}
}
pub struct AsyncTls13Client<S> {
stream: S,
pub conn: Connection,
deframer: TlsRecordDeframer,
}
impl<S: AsyncByteStream> AsyncTls13Client<S> {
pub async fn write_application(&mut self, plaintext: &[u8]) -> TransportResult<()> {
let packet = self
.seal_client_application(plaintext)
.map_err(|_| TransportError::IoFailed("seal failed"))?;
noxtls_write_all_async(&mut self.stream, &packet).await
}
pub async fn read_application(&mut self) -> TransportResult<Vec<u8>> {
let packet = noxtls_read_record_async(&mut self.stream, &mut self.deframer).await?;
self.open_server_application(&packet)
.map_err(|_| TransportError::IoFailed("open application failed"))
}
fn seal_client_application(&mut self, plaintext: &[u8]) -> TlsResult<Vec<u8>> {
let inner_len = plaintext.len().saturating_add(1);
let payload_len = inner_len.saturating_add(16);
let mut aad = [0_u8; TLS_RECORD_HEADER_LEN];
aad[0] = TLS_RECORD_APPLICATION_DATA;
aad[1] = 0x03;
aad[2] = 0x03;
aad[3] = ((payload_len >> 8) & 0xff) as u8;
aad[4] = (payload_len & 0xff) as u8;
self.conn
.noxtls_seal_record(plaintext, &aad)
.map(|protected| {
let mut packet =
Vec::with_capacity(TLS_RECORD_HEADER_LEN + protected.ciphertext.len());
packet.extend_from_slice(&aad);
packet.extend_from_slice(&protected.ciphertext);
packet
})
}
fn open_server_application(&mut self, packet: &[u8]) -> TlsResult<Vec<u8>> {
let aad = Connection::noxtls_tls13_packet_header_aad(packet)?;
let (plaintext, content_type) = self.conn.noxtls_open_tls13_record_packet(packet, &aad)?;
if content_type != RecordContentType::ApplicationData.to_u8() {
return Err(Error::ParseFailure("expected application_data"));
}
Ok(plaintext)
}
}
pub async fn noxtls_connect_tls13_async<S: AsyncByteStream>(
mut stream: S,
_config: &AsyncTlsClientConfig,
entropy: &mut dyn EntropySource,
) -> TransportResult<AsyncTls13Client<S>> {
let mut conn = Connection::noxtls_new(TlsVersion::Tls13);
let mut drbg = tls_err(noxtls_hmac_drbg_from_entropy(
entropy,
b"noxtls tls13 client",
))?;
let client_hello = tls_err(conn.noxtls_send_client_hello_auto(&mut drbg))?;
noxtls_write_all_async(&mut stream, &encode_handshake_record(&client_hello)?).await?;
let mut deframer = TlsRecordDeframer::noxtls_new();
let server_hello_record = noxtls_read_record_async(&mut stream, &mut deframer).await?;
let server_hello = extract_payload_message(&server_hello_record, 0x02)?;
tls_err(conn.noxtls_recv_server_hello(&server_hello))?;
if let Ok(ccs) = noxtls_read_record_async(&mut stream, &mut deframer).await {
if ccs.first() != Some(&TLS_RECORD_CHANGE_CIPHER_SPEC) {
deframer.push(&ccs);
}
}
tls_err(conn.noxtls_derive_handshake_secret())?;
let encrypted = noxtls_read_record_async(&mut stream, &mut deframer).await?;
let aad = tls_err(Connection::noxtls_tls13_packet_header_aad(&encrypted))?;
let packets = [encrypted];
tls_err(conn.noxtls_process_tls13_server_encrypted_handshake_flight(&packets, &aad))?;
let client_finished = tls_err(conn.noxtls_prepare_tls13_client_finished_message())?;
let finished_record = tls_err(seal_handshake_record(&mut conn, &client_finished))?;
noxtls_write_all_async(&mut stream, &finished_record).await?;
tls_err(conn.noxtls_activate_tls13_application_traffic_keys())?;
Ok(AsyncTls13Client {
stream,
conn,
deframer,
})
}
fn encode_handshake_record(body: &[u8]) -> TransportResult<Vec<u8>> {
let mut packet = Vec::with_capacity(TLS_RECORD_HEADER_LEN + body.len());
packet.push(TLS_RECORD_HANDSHAKE);
packet.extend_from_slice(&[0x03, 0x03]);
packet.extend_from_slice(&(body.len() as u16).to_be_bytes());
packet.extend_from_slice(body);
Ok(packet)
}
fn extract_payload_message(record: &[u8], msg_type: u8) -> TransportResult<Vec<u8>> {
if record.len() < TLS_RECORD_HEADER_LEN {
return Err(TransportError::IoFailed("short record"));
}
let payload = &record[TLS_RECORD_HEADER_LEN..];
if payload.first() != Some(&msg_type) {
return Err(TransportError::IoFailed("wrong handshake type"));
}
Ok(payload.to_vec())
}
fn seal_handshake_record(conn: &mut Connection, handshake: &[u8]) -> TlsResult<Vec<u8>> {
let inner_len = handshake.len().saturating_add(1);
let payload_len = inner_len.saturating_add(16);
let mut aad = [0_u8; TLS_RECORD_HEADER_LEN];
aad[0] = TLS_RECORD_APPLICATION_DATA;
aad[1] = 0x03;
aad[2] = 0x03;
aad[3] = ((payload_len >> 8) & 0xff) as u8;
aad[4] = (payload_len & 0xff) as u8;
conn.noxtls_seal_record(handshake, &aad).map(|p| {
let mut out = Vec::with_capacity(TLS_RECORD_HEADER_LEN + p.ciphertext.len());
out.extend_from_slice(&aad);
out.extend_from_slice(&p.ciphertext);
out
})
}