use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Instant;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use tracing::{debug, trace};
use super::cred::{CredHandle, CredKind};
use super::handshake::{Handshake, StepOutcome};
use super::record_layer::{Decrypted, RecordLayer};
enum Mode {
Streaming {
record: RecordLayer,
enc_in: Vec<u8>,
plain_out: Vec<u8>,
pending_out: Vec<u8>,
pending_out_written: usize,
pending_plain_len: usize,
},
}
pub(crate) struct SchannelTlsStream<S> {
socket: S,
mode: Mode,
channel_binding: Option<Vec<u8>>,
}
impl<S> SchannelTlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
pub(crate) async fn connect(
mut socket: S,
cred: Arc<CredHandle>,
kind: CredKind,
server_name: &str,
alpn_blob: Option<Vec<u8>>,
) -> io::Result<Self> {
let cred_for_record = Arc::clone(&cred);
let alpn_requested = alpn_blob.is_some();
let mut handshake = Handshake::new(cred, kind, server_name, alpn_blob);
let mut enc_in: Vec<u8> = Vec::new();
let mut read_buf = vec![0u8; 8192];
let connect_started = Instant::now();
let mut read_count: u32 = 0;
let mut write_count: u32 = 0;
let mut bytes_read: usize = 0;
let mut bytes_written: usize = 0;
loop {
let mut consumed = 0;
let outcome = handshake.step(enc_in.as_mut_slice(), &mut consumed)?;
if consumed > 0 {
enc_in.drain(..consumed);
}
match outcome {
StepOutcome::Done { sizes } => {
log_negotiated_alpn(handshake.ctx_handle(), alpn_requested);
let ctx = handshake.into_ctx();
let channel_binding = extract_channel_binding(&ctx);
debug!(
elapsed_ms = connect_started.elapsed().as_millis() as u64,
socket_reads = read_count,
socket_writes = write_count,
bytes_read,
bytes_written,
extra_buffered = enc_in.len(),
"win_tls: stream entering Streaming mode (no flush needed)"
);
return Ok(SchannelTlsStream {
socket,
mode: Mode::Streaming {
record: RecordLayer::new(
ctx,
sizes,
cred_for_record,
kind,
server_name,
),
enc_in,
plain_out: Vec::new(),
pending_out: Vec::new(),
pending_out_written: 0,
pending_plain_len: 0,
},
channel_binding,
});
}
StepOutcome::DoneWithFlush { out, sizes } => {
socket.write_all(&out).await?;
socket.flush().await?;
write_count += 1;
bytes_written += out.len();
log_negotiated_alpn(handshake.ctx_handle(), alpn_requested);
let ctx = handshake.into_ctx();
let channel_binding = extract_channel_binding(&ctx);
debug!(
elapsed_ms = connect_started.elapsed().as_millis() as u64,
socket_reads = read_count,
socket_writes = write_count,
bytes_read,
bytes_written,
extra_buffered = enc_in.len(),
"win_tls: stream entering Streaming mode (after final flush)"
);
return Ok(SchannelTlsStream {
socket,
mode: Mode::Streaming {
record: RecordLayer::new(
ctx,
sizes,
cred_for_record,
kind,
server_name,
),
enc_in,
plain_out: Vec::new(),
pending_out: Vec::new(),
pending_out_written: 0,
pending_plain_len: 0,
},
channel_binding,
});
}
StepOutcome::WantWriteThenRead(out) => {
socket.write_all(&out).await?;
socket.flush().await?;
write_count += 1;
bytes_written += out.len();
}
StepOutcome::NeedMoreInput => {
}
}
let n = socket.read(&mut read_buf).await?;
if n == 0 {
debug!(
elapsed_ms = connect_started.elapsed().as_millis() as u64,
socket_reads = read_count,
socket_writes = write_count,
bytes_read,
bytes_written,
"win_tls: peer closed during handshake (EOF)"
);
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"peer closed connection during TLS handshake",
));
}
read_count += 1;
bytes_read += n;
trace!(
read_n = n,
bytes_read_total = bytes_read,
"win_tls: socket read"
);
enc_in.extend_from_slice(&read_buf[..n]);
}
}
pub(crate) fn get_ref(&self) -> &S {
&self.socket
}
pub(crate) fn get_mut(&mut self) -> &mut S {
&mut self.socket
}
pub(crate) fn ctx(&self) -> &super::handshake::SecCtx {
match &self.mode {
Mode::Streaming { record, .. } => record.ctx(),
}
}
pub(crate) fn channel_binding_token(&self) -> Option<Vec<u8>> {
self.channel_binding.clone()
}
}
fn extract_channel_binding(ctx: &super::handshake::SecCtx) -> Option<Vec<u8>> {
match super::bindings::query_unique_bindings(ctx) {
Ok(token) => {
debug!(
token_len = token.len(),
"win_tls: extracted tls-unique channel binding token"
);
Some(token)
}
Err(e) => {
debug!(error = %e, "win_tls: failed to extract channel binding token");
None
}
}
}
fn log_negotiated_alpn(
ctx: &windows_sys::Win32::Security::Credentials::SecHandle,
requested: bool,
) {
if !requested {
return;
}
match super::alpn::query_negotiated_alpn(ctx) {
Ok(Some(proto)) => {
debug!(
proto = %super::alpn::debug_proto(&proto),
"win_tls: server negotiated ALPN protocol"
);
}
Ok(None) => {
debug!("win_tls: server did not negotiate an ALPN protocol");
}
Err(e) => {
debug!(error = %e, "win_tls: failed to query negotiated ALPN");
}
}
}
fn read_eof_outcome(enc_in_empty: bool) -> Poll<io::Result<()>> {
if enc_in_empty {
Poll::Ready(Ok(()))
} else {
Poll::Ready(Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"win_tls: connection closed in the middle of a TLS record",
)))
}
}
impl<S> AsyncRead for SchannelTlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
let Mode::Streaming {
record,
enc_in,
plain_out,
..
} = &mut this.mode;
loop {
if !plain_out.is_empty() {
let n = std::cmp::min(plain_out.len(), buf.remaining());
buf.put_slice(&plain_out[..n]);
plain_out.drain(..n);
return Poll::Ready(Ok(()));
}
match record.decrypt(enc_in, plain_out) {
Ok(Decrypted::Ok) => continue, Ok(Decrypted::PeerClosed) => return Poll::Ready(Ok(())),
Ok(Decrypted::NeedMoreInput) => {
}
Err(e) => return Poll::Ready(Err(e)),
}
let mut tmp = [0u8; 8192];
let mut tmp_buf = ReadBuf::new(&mut tmp);
match Pin::new(&mut this.socket).poll_read(cx, &mut tmp_buf) {
Poll::Ready(Ok(())) => {
let filled = tmp_buf.filled().len();
if filled == 0 {
return read_eof_outcome(enc_in.is_empty());
}
enc_in.extend_from_slice(&tmp[..filled]);
}
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
}
}
impl<S> AsyncWrite for SchannelTlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
if buf.is_empty() {
return Poll::Ready(Ok(0));
}
let this = self.get_mut();
let Mode::Streaming {
record,
pending_out,
pending_out_written,
pending_plain_len,
..
} = &mut this.mode;
if !pending_out.is_empty() {
while *pending_out_written < pending_out.len() {
match Pin::new(&mut this.socket)
.poll_write(cx, &pending_out[*pending_out_written..])
{
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"socket accepted zero bytes",
)));
}
Poll::Ready(Ok(n)) => *pending_out_written += n,
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
let plain_len = *pending_plain_len;
debug_assert!(
plain_len <= buf.len(),
"win_tls: poll_write drained a pending record of {plain_len} \
plaintext bytes but the retry buffer is only {} bytes; the \
caller must retry with the same buffer after Poll::Pending",
buf.len()
);
trace!(
drained = pending_out.len(),
plain_len, "win_tls: poll_write drained pending record"
);
pending_out.clear();
*pending_out_written = 0;
*pending_plain_len = 0;
return Poll::Ready(Ok(plain_len));
}
let chunk = std::cmp::min(buf.len(), record.max_message());
let mut out = Vec::with_capacity(record.header_len() + chunk + record.trailer_len());
if let Err(e) = record.encrypt(&buf[..chunk], &mut out) {
return Poll::Ready(Err(e));
}
let mut written = 0;
while written < out.len() {
match Pin::new(&mut this.socket).poll_write(cx, &out[written..]) {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"socket accepted zero bytes",
)));
}
Poll::Ready(Ok(n)) => {
written += n;
}
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => {
debug!(
encrypted_total = out.len(),
written_so_far = written,
plain_chunk = chunk,
"win_tls: poll_write partial socket write, stashing pending record"
);
*pending_out = out;
*pending_out_written = written;
*pending_plain_len = chunk;
return Poll::Pending;
}
}
}
Poll::Ready(Ok(chunk))
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
let Mode::Streaming {
pending_out,
pending_out_written,
pending_plain_len,
..
} = &mut this.mode;
while *pending_out_written < pending_out.len() {
match Pin::new(&mut this.socket).poll_write(cx, &pending_out[*pending_out_written..]) {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"socket accepted zero bytes",
)));
}
Poll::Ready(Ok(n)) => *pending_out_written += n,
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
if !pending_out.is_empty() {
pending_out.clear();
*pending_out_written = 0;
*pending_plain_len = 0;
}
Pin::new(&mut this.socket).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
Pin::new(&mut this.socket).poll_shutdown(cx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::net::TcpListener;
#[tokio::test]
async fn handshake_returns_eof_when_peer_closes() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (sock, _) = listener.accept().await.unwrap();
let sock = sock;
let mut tmp = vec![0u8; 8192];
let _ = sock.readable().await;
let _ = sock.try_read(&mut tmp);
drop(sock);
});
let client_sock = tokio::net::TcpStream::connect(addr).await.unwrap();
let cred = super::super::cred::get_or_acquire(CredKind::NoValidate).unwrap();
let result =
SchannelTlsStream::connect(client_sock, cred, CredKind::NoValidate, "127.0.0.1", None)
.await;
assert!(result.is_err(), "expected handshake to fail on EOF");
let err = match result {
Ok(_) => unreachable!("checked above"),
Err(e) => e,
};
assert!(
err.kind() == io::ErrorKind::UnexpectedEof
|| err.to_string().contains("InitializeSecurityContextW")
|| err.kind() == io::ErrorKind::ConnectionReset,
"unexpected error: {err:?}"
);
server.await.unwrap();
}
use super::super::handshake::SecCtx;
use windows_sys::Win32::Security::Authentication::Identity;
#[test]
fn read_eof_outcome_clean_close_on_empty_buffer() {
match read_eof_outcome(true) {
Poll::Ready(Ok(())) => {}
other => panic!("expected graceful EOF, got {other:?}"),
}
}
#[test]
fn read_eof_outcome_truncated_record_is_error() {
match read_eof_outcome(false) {
Poll::Ready(Err(e)) => assert_eq!(e.kind(), io::ErrorKind::UnexpectedEof),
other => panic!("expected UnexpectedEof, got {other:?}"),
}
}
#[test]
fn extract_channel_binding_returns_none_when_query_fails() {
assert!(extract_channel_binding(&SecCtx::for_test_only()).is_none());
}
struct MockSocket {
written: Vec<u8>,
}
impl AsyncRead for MockSocket {
fn poll_read(
self: Pin<&mut Self>,
_: &mut Context<'_>,
_: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
impl AsyncWrite for MockSocket {
fn poll_write(
mut self: Pin<&mut Self>,
_: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
self.written.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
fn streaming_stream(
socket: MockSocket,
pending_out: Vec<u8>,
pending_plain_len: usize,
) -> SchannelTlsStream<MockSocket> {
let cred = super::super::cred::get_or_acquire(CredKind::NoValidate).unwrap();
let sizes: Identity::SecPkgContext_StreamSizes = unsafe { std::mem::zeroed() };
let record = RecordLayer::new(
SecCtx::for_test_only(),
sizes,
cred,
CredKind::NoValidate,
"test",
);
SchannelTlsStream {
socket,
mode: Mode::Streaming {
record,
enc_in: Vec::new(),
plain_out: Vec::new(),
pending_out,
pending_out_written: 0,
pending_plain_len,
},
channel_binding: None,
}
}
#[test]
fn poll_write_drains_pending_record_and_returns_plain_len() {
let waker = std::task::Waker::noop();
let mut cx = Context::from_waker(waker);
let mut s = streaming_stream(
MockSocket {
written: Vec::new(),
},
vec![10, 20, 30, 40],
7,
);
let buf = [0u8; 16];
match Pin::new(&mut s).poll_write(&mut cx, &buf) {
Poll::Ready(Ok(n)) => assert_eq!(n, 7, "must report the stashed plaintext length"),
other => panic!("expected Ready(Ok(7)), got {other:?}"),
}
let Mode::Streaming {
pending_out,
pending_out_written,
pending_plain_len,
..
} = &s.mode;
assert!(pending_out.is_empty(), "pending record must be cleared");
assert_eq!(*pending_out_written, 0);
assert_eq!(*pending_plain_len, 0);
assert_eq!(
s.socket.written,
vec![10, 20, 30, 40],
"record forwarded to socket"
);
}
#[cfg(debug_assertions)]
#[test]
#[should_panic(expected = "retry buffer is only")]
fn poll_write_pending_drain_asserts_buffer_invariant() {
let waker = std::task::Waker::noop();
let mut cx = Context::from_waker(waker);
let mut s = streaming_stream(
MockSocket {
written: Vec::new(),
},
vec![1, 2, 3, 4],
100,
);
let buf = [0u8; 4];
let _ = Pin::new(&mut s).poll_write(&mut cx, &buf);
}
#[test]
fn poll_read_drains_buffered_plaintext_first() {
let waker = std::task::Waker::noop();
let mut cx = Context::from_waker(waker);
let mut s = streaming_stream(
MockSocket {
written: Vec::new(),
},
Vec::new(),
0,
);
let Mode::Streaming { plain_out, .. } = &mut s.mode;
plain_out.extend_from_slice(&[1, 2, 3, 4, 5]);
let mut backing = [0u8; 3];
let mut rb = ReadBuf::new(&mut backing);
match Pin::new(&mut s).poll_read(&mut cx, &mut rb) {
Poll::Ready(Ok(())) => assert_eq!(rb.filled(), &[1, 2, 3]),
other => panic!("expected Ready(Ok), got {other:?}"),
}
let Mode::Streaming { plain_out, .. } = &s.mode;
assert_eq!(plain_out, &vec![4, 5], "leftover plaintext stays buffered");
}
#[test]
fn poll_read_returns_graceful_eof_when_socket_closes() {
let waker = std::task::Waker::noop();
let mut cx = Context::from_waker(waker);
let mut s = streaming_stream(
MockSocket {
written: Vec::new(),
},
Vec::new(),
0,
);
let mut backing = [0u8; 16];
let mut rb = ReadBuf::new(&mut backing);
match Pin::new(&mut s).poll_read(&mut cx, &mut rb) {
Poll::Ready(Ok(())) => assert_eq!(rb.filled().len(), 0, "clean EOF yields zero bytes"),
other => panic!("expected graceful EOF, got {other:?}"),
}
}
}