diff --git a/tonic/src/body.rs b/tonic/src/body.rs index 19f01c3..64b3808 100644 --- a/tonic/src/body.rs +++ b/tonic/src/body.rs @@ -1,6 +1,6 @@ use crate::{Code, Status}; use bytes::{Bytes, IntoBuf}; -use futures_core::TryStream; +use futures_core::Stream; use futures_util::{ready, TryStreamExt}; use http::HeaderMap; use http_body::Body; @@ -13,9 +13,18 @@ pub struct AsyncBody { error: Option, } +impl AsyncBody +where + S: Stream> + Unpin, +{ + pub fn new(inner: S) -> Self { + Self { inner, error: None } + } +} + impl Body for AsyncBody where - S: TryStream + Unpin, + S: Stream> + Unpin, { type Data = BytesBuf; type Error = Status; diff --git a/tonic/src/lib.rs b/tonic/src/lib.rs index 90c8a0d..b59cebf 100644 --- a/tonic/src/lib.rs +++ b/tonic/src/lib.rs @@ -18,6 +18,8 @@ pub use response::Response; pub use status::{Code, Status}; pub use tonic_macros::server; +pub(crate) use error::Error; + use std::future::Future; use std::sync::Arc; diff --git a/tonic/src/server/mod.rs b/tonic/src/server/mod.rs index 141c468..59336e0 100644 --- a/tonic/src/server/mod.rs +++ b/tonic/src/server/mod.rs @@ -1,14 +1,15 @@ -use crate::{Request, Response, Status}; +#![allow(dead_code)] + +use crate::{Code, Request, Response, Status}; use async_stream::stream; -use bytes::{Bytes, BytesMut, IntoBuf}; -use futures_core::TryStream; -use futures_util::{stream, StreamExt, TryStreamExt}; +use bytes::{Buf, BufMut, Bytes, BytesMut, IntoBuf}; +use futures_core::{Stream, TryStream}; +use futures_util::{future, stream, StreamExt, TryStreamExt}; +use http_body::Body; use std::future::Future; use tokio_codec::{Decoder, Encoder}; use tower_service::Service; - -#[allow(dead_code)] -type Result = std::result::Result, Status>; +use tracing::{debug, trace}; pub trait Codec { type Encode; @@ -16,7 +17,7 @@ pub trait Codec { } pub struct Encode { - inner: T, + encoder: T, source: U, } @@ -25,19 +26,19 @@ where T: Encoder, U: TryStream + Unpin, { - pub fn new(inner: T, source: U) -> Self { - Encode { inner, source } + pub fn new(encoder: T, source: U) -> Self { + Encode { encoder, source } } pub fn encode<'a>( &'a mut self, buf: &'a mut BytesMut, - ) -> impl TryStream + 'a { + ) -> impl Stream> + 'a { stream! { loop { match self.source.try_next().await { Ok(Some(item)) => { - self.inner.encode(item, buf).map_err(drop).unwrap(); + self.encoder.encode(item, buf).map_err(drop).unwrap(); let len = buf.len(); yield Ok(buf.split_to(len).freeze().into_buf()); }, @@ -49,17 +50,135 @@ where } } +pub struct Streaming { + decoder: T, + buf: BytesMut, + state: State, +} + +#[derive(Debug)] +enum State { + ReadHeader, + ReadBody { compression: bool, len: usize }, + Done, +} + +impl Streaming +where + T: Decoder, + T::Item: Unpin + 'static, +{ + pub fn decode<'a, B>( + &'a mut self, + source: &'a mut B, + ) -> impl Stream> + 'a + where + B: Body, + B::Error: Into, + { + stream! { + loop { + // TODO: use try_stream! and ? + if let Some(item) = self.decode_chunk().unwrap() { + yield Ok(item); + } + + let chunk = match future::poll_fn(|cx| source.poll_data(cx)).await { + Some(Ok(d)) => Some(d), + Some(Err(e)) => { + let err = e.into(); + debug!("decoder inner stream error: {:?}", err); + let status = Status::from_error(&*err); + yield Err(status); + break; + }, + None => None, + }; + + if let Some(data)= chunk { + self.buf.put(data); + } else { + if self.buf.has_remaining_mut() { + trace!("unexpected EOF decoding stream"); + yield Err(Status::new( + Code::Internal, + "Unexpected EOF decoding stream.".to_string(), + )); + } else { + break; + } + } + } + } + } + + fn decode_chunk(&mut self) -> Result, Status> { + let buf = (&self.buf).into_buf(); + + if let State::ReadHeader = self.state { + if buf.remaining() < 5 { + return Ok(None); + } + + let is_compressed = match buf.get_u8() { + 0 => false, + 1 => { + trace!("message compressed, compression not supported yet"); + return Err(crate::Status::new( + crate::Code::Unimplemented, + "Message compressed, compression not supported yet.".to_string(), + )); + } + f => { + trace!("unexpected compression flag"); + return Err(crate::Status::new( + crate::Code::Internal, + format!("Unexpected compression flag: {}", f), + )); + } + }; + let len = (&self.buf[..]).into_buf().get_u32_be() as usize; + + self.state = State::ReadBody { + compression: is_compressed, + len, + } + } + + if let State::ReadBody { len, .. } = self.state { + if buf.remaining() < len { + return Ok(None); + } + + match self.decoder.decode(&mut self.buf) { + Ok(Some(msg)) => { + self.state = State::ReadHeader; + return Ok(Some(msg)); + } + Err(e) => { + return Err(e); + } + } + } + + Ok(None) + } +} + #[cfg(test)] mod tests { use crate::body::AsyncBody; use crate::server::Encode; - use bytes::Bytes; + use bytes::{Bytes, BytesMut}; use tokio_codec::BytesCodec; #[test] fn body() { let stream = futures_util::stream::iter(vec![Ok(Bytes::new())]); - let encode = Encode::new(BytesCodec::new(), stream); + let mut encode = Encode::new(BytesCodec::new(), stream); + + let mut buf = BytesMut::with_capacity(1024); + AsyncBody::new(Box::pin(encode.encode(&mut buf))); } }