use std::{
future::Future,
io::{self, Read, Write},
pin::Pin,
ptr,
task::{Context, Poll},
};
use bytes::{Bytes, BytesMut};
use futures::{ready, Sink, SinkExt, Stream, StreamExt};
use msf_rtp::{PacketMux, RtcpPacketType};
use openssl::ssl::{HandshakeError, Ssl, SslStream};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use crate::{
session::{DecodingError, SrtpSession},
Error, InternalError,
};
pub struct Connector {
inner: Ssl,
}
impl Connector {
pub fn new(ssl: Ssl) -> Self {
Self { inner: ssl }
}
pub async fn connect<S>(self, mut stream: S) -> Result<SrtpStream<S>, Error>
where
S: Stream<Item = io::Result<Bytes>> + Sink<Bytes, Error = io::Error> + Unpin,
{
let mut ssl_stream = InnerSslStream::new(&mut stream);
let connect = futures::future::lazy(move |cx| {
ssl_stream.set_async_context(Some(cx));
let mut res = HandshakeState::from(self.inner.connect(ssl_stream));
res.set_async_context(None);
res
});
let handshake = Handshake::from(connect.await);
let session = SrtpSession::client(handshake.await?.ssl())?;
Ok(SrtpStream::new(session, stream))
}
pub async fn accept<S>(self, mut stream: S) -> Result<SrtpStream<S>, Error>
where
S: Stream<Item = io::Result<Bytes>> + Sink<Bytes, Error = io::Error> + Unpin,
{
let mut ssl_stream = InnerSslStream::new(&mut stream);
let accept = futures::future::lazy(move |cx| {
ssl_stream.set_async_context(Some(cx));
let mut res = HandshakeState::from(self.inner.accept(ssl_stream));
res.set_async_context(None);
res
});
let handshake = Handshake::from(accept.await);
let session = SrtpSession::server(handshake.await?.ssl())?;
Ok(SrtpStream::new(session, stream))
}
}
pub struct SrtpStream<S> {
session: SrtpSession,
inner: S,
}
impl<S> SrtpStream<S> {
fn new(session: SrtpSession, stream: S) -> Self {
Self {
session,
inner: stream,
}
}
}
impl<S> Stream for SrtpStream<S>
where
S: Stream<Item = io::Result<Bytes>> + Unpin,
{
type Item = Result<PacketMux, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
if let Some(packet) = self.session.next() {
return Poll::Ready(Some(Ok(packet)));
} else if let Poll::Ready(ready) = self.inner.poll_next_unpin(cx) {
if let Some(frame) = ready.transpose()? {
if let Err(DecodingError::Other(err)) = self.session.decode(frame) {
return Poll::Ready(Some(Err(err)));
}
} else {
return Poll::Ready(None);
}
} else {
return Poll::Pending;
}
}
}
}
impl<S> Sink<PacketMux> for SrtpStream<S>
where
S: Sink<Bytes, Error = io::Error> + Unpin,
{
type Error = Error;
#[inline]
fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
ready!(self.inner.poll_ready_unpin(cx))?;
Poll::Ready(Ok(()))
}
fn start_send(mut self: Pin<&mut Self>, packet: PacketMux) -> Result<(), Self::Error> {
let frame = match packet {
PacketMux::Rtp(packet) => self.session.encode_rtp_packet(packet)?,
PacketMux::Rtcp(packet) => {
if let Some(first) = packet.first() {
match first.packet_type() {
RtcpPacketType::SR | RtcpPacketType::RR => {
self.session.encode_rtcp_packet(packet)?
}
_ => return Err(Error::from(InternalError::InvalidPacketType)),
}
} else {
return Ok(());
}
}
};
self.inner.start_send_unpin(frame)?;
Ok(())
}
#[inline]
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
ready!(self.inner.poll_flush_unpin(cx))?;
Poll::Ready(Ok(()))
}
#[inline]
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
ready!(self.inner.poll_close_unpin(cx))?;
Poll::Ready(Ok(()))
}
}
struct Handshake<'a, S> {
inner: Option<HandshakeState<'a, S>>,
}
impl<'a, S> From<HandshakeState<'a, S>> for Handshake<'a, S> {
fn from(state: HandshakeState<'a, S>) -> Self {
Self { inner: Some(state) }
}
}
impl<'a, S> Future for Handshake<'a, S>
where
S: Stream<Item = io::Result<Bytes>> + Sink<Bytes, Error = io::Error> + Unpin,
{
type Output = Result<SslStream<InnerSslStream<'a, S>>, InternalError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
let mut state = self
.inner
.take()
.expect("the future has been already resolved");
state.set_async_context(Some(cx));
match state.inner {
Ok(stream) => Poll::Ready(Ok(stream)),
Err(HandshakeError::SetupFailure(err)) => Poll::Ready(Err(err.into())),
Err(HandshakeError::Failure(m)) => {
Poll::Ready(Err(InternalError::from(m.into_error())))
}
Err(HandshakeError::WouldBlock(m)) => match m.handshake() {
Ok(stream) => Poll::Ready(Ok(stream)),
Err(HandshakeError::SetupFailure(err)) => Poll::Ready(Err(err.into())),
Err(HandshakeError::Failure(m)) => {
Poll::Ready(Err(InternalError::from(m.into_error())))
}
Err(HandshakeError::WouldBlock(m)) => {
let mut state = HandshakeState::from(HandshakeError::WouldBlock(m));
state.set_async_context(None);
self.inner = Some(state);
Poll::Pending
}
},
}
}
}
type HandshakeResult<'a, S> =
Result<SslStream<InnerSslStream<'a, S>>, HandshakeError<InnerSslStream<'a, S>>>;
struct HandshakeState<'a, S> {
inner: HandshakeResult<'a, S>,
}
impl<'a, S> HandshakeState<'a, S> {
fn set_async_context(&mut self, cx: Option<&mut Context<'_>>) {
let ssl_stream = match &mut self.inner {
Ok(ssl_stream) => Some(ssl_stream.get_mut()),
Err(HandshakeError::Failure(m)) => Some(m.get_mut()),
Err(HandshakeError::WouldBlock(m)) => Some(m.get_mut()),
_ => None,
};
if let Some(s) = ssl_stream {
s.set_async_context(cx);
}
}
}
impl<'a, S> From<HandshakeResult<'a, S>> for HandshakeState<'a, S> {
fn from(res: HandshakeResult<'a, S>) -> Self {
Self { inner: res }
}
}
impl<'a, S> From<HandshakeError<InnerSslStream<'a, S>>> for HandshakeState<'a, S> {
fn from(err: HandshakeError<InnerSslStream<'a, S>>) -> Self {
Self::from(Err(err))
}
}
struct InnerSslStream<'a, S> {
inner: RWStreamRef<'a, S>,
context: *mut (),
}
impl<'a, S> InnerSslStream<'a, S> {
fn new(stream: &'a mut S) -> Self {
Self {
inner: RWStreamRef::new(stream),
context: ptr::null_mut(),
}
}
fn set_async_context(&mut self, cx: Option<&mut Context<'_>>) {
if let Some(cx) = cx {
self.context = cx as *mut _ as *mut ();
} else {
self.context = ptr::null_mut();
}
}
}
unsafe impl<'a, S> Send for InnerSslStream<'a, S> where S: Send {}
unsafe impl<'a, S> Sync for InnerSslStream<'a, S> where S: Sync {}
impl<'a, S> Read for InnerSslStream<'a, S>
where
S: Stream<Item = io::Result<Bytes>> + Unpin,
{
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
debug_assert!(!self.context.is_null());
let cx = unsafe { &mut *(self.context as *mut Context<'_>) };
let pinned = Pin::new(&mut self.inner);
let mut buf = ReadBuf::new(buf);
let data = match pinned.poll_read(cx, &mut buf) {
Poll::Ready(Ok(())) => buf.filled(),
Poll::Ready(Err(err)) => return Err(err),
Poll::Pending => return Err(io::Error::from(io::ErrorKind::WouldBlock)),
};
Ok(data.len())
}
}
impl<'a, S> Write for InnerSslStream<'a, S>
where
S: Sink<Bytes, Error = io::Error> + Unpin,
{
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
debug_assert!(!self.context.is_null());
let cx = unsafe { &mut *(self.context as *mut Context<'_>) };
let pinned = Pin::new(&mut self.inner);
if let Poll::Ready(res) = pinned.poll_write(cx, buf) {
res
} else {
Err(io::Error::from(io::ErrorKind::WouldBlock))
}
}
fn flush(&mut self) -> io::Result<()> {
debug_assert!(!self.context.is_null());
let cx = unsafe { &mut *(self.context as *mut Context<'_>) };
let pinned = Pin::new(&mut self.inner);
if let Poll::Ready(res) = AsyncWrite::poll_flush(pinned, cx) {
res
} else {
Err(io::Error::from(io::ErrorKind::WouldBlock))
}
}
}
struct RWStreamRef<'a, S> {
stream: &'a mut S,
input: Bytes,
output: BytesMut,
}
impl<'a, S> RWStreamRef<'a, S> {
fn new(stream: &'a mut S) -> Self {
Self {
stream,
input: Bytes::new(),
output: BytesMut::new(),
}
}
}
impl<'a, S> AsyncRead for RWStreamRef<'a, S>
where
S: Stream<Item = io::Result<Bytes>> + Unpin,
{
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
loop {
if !self.input.is_empty() {
let remaining = buf.remaining();
let take = remaining.min(self.input.len());
let data = self.input.split_to(take);
buf.put_slice(&data);
return Poll::Ready(Ok(()));
} else if let Poll::Ready(ready) = self.stream.poll_next_unpin(cx) {
if let Some(chunk) = ready.transpose()? {
self.input = chunk;
} else {
return Poll::Ready(Ok(()));
}
} else {
return Poll::Pending;
}
}
}
}
impl<'a, S> AsyncWrite for RWStreamRef<'a, S>
where
S: Sink<Bytes, Error = io::Error> + Unpin,
{
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let this = &mut *self;
ready!(this.stream.poll_ready_unpin(cx))?;
this.output.extend_from_slice(buf);
let data = this.output.split_to(this.output.len()).freeze();
this.stream.start_send_unpin(data)?;
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
self.stream.poll_flush_unpin(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
self.stream.poll_close_unpin(cx)
}
}