Skip to content

Commit 27a4a31

Browse files
committed
ByteStream Codec MethodHandlers
1 parent 898ba52 commit 27a4a31

3 files changed

Lines changed: 405 additions & 0 deletions

File tree

grpc/src/server/method_handler.rs

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,9 @@ pub use server_streaming_adapter::ServerStreamingAdapter;
1010

1111
mod unary_adapter;
1212
pub use unary_adapter::UnaryMethodAdapter;
13+
mod generic_byte_stream_method_handler;
14+
pub use crate::call::Incoming;
15+
pub use generic_byte_stream_method_handler::GenericByteStreamMethodHandler;
1316

1417
mod message_stream_handler;
1518
pub use message_stream_handler::MessageStreamHandler;
Lines changed: 382 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,382 @@
1+
use crate::call::message_wrapper::CompressionEncoding;
2+
use crate::call::streaming_response_writer_ext::StreamingResponseWriterExt;
3+
use crate::call::{
4+
HandlerCallOptions, Incoming, Lazy, Outgoing, StreamingRequest, StreamingResponseWriter,
5+
};
6+
use crate::codec::compression::{get_codec, Compressor};
7+
use crate::codec::serialization::{Deserialize, Serialize};
8+
use crate::message::AsMut;
9+
use crate::server::method_handler::message_allocator::RpcResponseHolder;
10+
use crate::server::method_handler::{
11+
CodecRespB, GenericByteStreamMethodHandler, MessageStreamHandler,
12+
};
13+
use crate::status::StatusCode;
14+
use crate::stream::{PushStreamExt, PushStreamProducer};
15+
use crate::Status;
16+
use bytes::{Buf, BytesMut};
17+
use std::marker::PhantomData;
18+
19+
use std::sync::Arc;
20+
21+
/// A codec that adapts a GenericByteStreamMethodHandler to a MessageStreamHandler.
22+
pub struct CodecMessageStreamHandler<H, Req, Resp> {
23+
inner: H,
24+
// Use fn(Req, Resp) to avoid imposing Send/Sync bounds on Req/Resp for the struct itself
25+
_pd: PhantomData<fn(Req, Resp)>,
26+
}
27+
28+
impl<H, Req, Resp> CodecMessageStreamHandler<H, Req, Resp> {
29+
pub fn new(inner: H) -> Self {
30+
Self {
31+
inner,
32+
_pd: PhantomData,
33+
}
34+
}
35+
}
36+
37+
impl<H, Req, Resp> GenericByteStreamMethodHandler for CodecMessageStreamHandler<H, Req, Resp>
38+
where
39+
H: MessageStreamHandler<Req, Resp> + Send + Sync,
40+
Req: Send + AsMut + Deserialize + Default,
41+
Resp: Send + AsMut + Serialize + Default,
42+
for<'a> <Resp as AsMut>::Mut<'a>: Send + Serialize,
43+
for<'a> <Req as AsMut>::Mut<'a>: Send + Deserialize,
44+
{
45+
type RespB = CodecRespB;
46+
47+
async fn call<ReqB, P>(
48+
&self,
49+
options: HandlerCallOptions,
50+
req: StreamingRequest<Incoming<ReqB>, P>,
51+
resp_writer: impl StreamingResponseWriter<Self::RespB>,
52+
) -> Result<(), Status>
53+
where
54+
ReqB: Buf + Send,
55+
P: PushStreamProducer<Item = Incoming<ReqB>> + Send,
56+
{
57+
// 1. Transform Request Stream: RawMessage -> Lazy<Req>
58+
let (metadata, raw_stream) = req.into_parts();
59+
60+
// Resolve Decompressor
61+
let decompressor = if let Some(encoding) = metadata.encoding() {
62+
Some(get_codec(encoding).ok_or_else(|| {
63+
Status::new(
64+
StatusCode::Unimplemented,
65+
format!("compression encoding {} not found", encoding),
66+
)
67+
})?)
68+
} else {
69+
None
70+
};
71+
72+
let typed_req_stream = raw_stream.then(move |raw_msg| {
73+
// TODO(sauravz): Avoid this per message clone by changing the
74+
// lambda to a struct with async function.
75+
let decompressor = decompressor.clone();
76+
async move {
77+
Ok(CodecLazy {
78+
raw_msg,
79+
decompressor,
80+
_pd: PhantomData,
81+
})
82+
}
83+
});
84+
85+
// 2. Prepare Response Writer
86+
let compressor = options.compression_encoding.as_ref().and_then(|name| {
87+
if metadata.accept_encodings().any(|a| a == name) {
88+
get_codec(name)
89+
} else {
90+
None
91+
}
92+
});
93+
94+
let typed_resp_writer =
95+
resp_writer.map_message(move |item: Outgoing<H::ResponseHolder>| {
96+
// TODO(sauravz): Avoid this per message clone by changing the
97+
// lambda to a struct with async function.
98+
let compressor = compressor.clone();
99+
async move {
100+
let Outgoing {
101+
mut message,
102+
options,
103+
} = item;
104+
// TODO(sauravz): We need not allocate everytime. We can either
105+
// reuse buffer or find a way to pool them somehow.
106+
let mut buf = BytesMut::new();
107+
108+
// Serialize response
109+
// We assume <Resp as AsMut>::Mut derefs to Resp or behaves like it for serialization
110+
{
111+
let resp_mut = message.get_response_mut();
112+
resp_mut.serialize(&mut buf).map_err(|e| {
113+
Status::new(
114+
StatusCode::Internal,
115+
format!("serialization error: {:?}", e),
116+
)
117+
})?;
118+
}
119+
120+
let message_compression =
121+
options.as_ref().map(|o| o.compression).unwrap_or_default();
122+
match (message_compression, &compressor) {
123+
(
124+
CompressionEncoding::Inherit | CompressionEncoding::Enabled,
125+
Some(compressor),
126+
) => {
127+
let mut compressed = BytesMut::new();
128+
compressor
129+
.compress(&mut buf, &mut compressed)
130+
.map_err(|e| {
131+
Status::new(
132+
StatusCode::Internal,
133+
format!("compression error: {}", e),
134+
)
135+
})?;
136+
Ok(compressed.freeze())
137+
}
138+
// Disabled or Enabled/Inherit but no stream compressor -> Identity
139+
_ => Ok(buf.freeze()),
140+
}
141+
}
142+
});
143+
144+
// 3. Call Inner Handler
145+
self.inner
146+
.call(
147+
options,
148+
StreamingRequest::new(typed_req_stream, metadata),
149+
typed_resp_writer,
150+
)
151+
.await
152+
}
153+
}
154+
155+
pub struct CodecLazy<Req, B> {
156+
raw_msg: Incoming<B>,
157+
decompressor: Option<Arc<dyn Compressor>>,
158+
_pd: PhantomData<Req>,
159+
}
160+
161+
impl<Req, B> Lazy<Req> for CodecLazy<Req, B>
162+
where
163+
Req: Deserialize + Default + AsMut + Send,
164+
B: Buf + Send,
165+
for<'a> <Req as AsMut>::Mut<'a>: Send + Deserialize,
166+
{
167+
async fn resolve(self, mut dest: <Req as AsMut>::Mut<'_>) -> Result<(), Status> {
168+
let mut raw_msg = self.raw_msg;
169+
if let Some(decompressor) = &self.decompressor {
170+
// TODO(sauravz): We need not allocate everytime. We can either
171+
// reuse buffer or find a way to pool them somehow.
172+
let mut decompressed = BytesMut::new();
173+
decompressor
174+
.decompress(&mut raw_msg.message_bytes, &mut decompressed)
175+
.map_err(|e| {
176+
Status::new(StatusCode::Internal, format!("decompression error: {}", e))
177+
})?;
178+
dest.deserialize(&mut decompressed)
179+
} else {
180+
dest.deserialize(&mut raw_msg.message_bytes)
181+
}
182+
}
183+
}
184+
185+
#[cfg(test)]
186+
mod tests {
187+
use super::*;
188+
use crate::call::{
189+
Metadata, StreamingRequest, StreamingResponseBodyWriter, StreamingResponseWriter,
190+
};
191+
use crate::server::method_handler::{HeapResponseHolder, MessageStreamHandler};
192+
use crate::stream::{PushStream, PushStreamProducer, PushStreamWriter};
193+
use crate::Status;
194+
use bytes::{Buf, Bytes, BytesMut};
195+
use protobuf_well_known_types::Timestamp;
196+
use send_future::SendFuture;
197+
use tokio::sync::mpsc;
198+
199+
struct MockMessageStreamHandler {
200+
expected_reqs: Vec<Timestamp>,
201+
resps_to_return: Vec<Outgoing<Timestamp>>,
202+
}
203+
204+
impl MessageStreamHandler<Timestamp, Timestamp> for MockMessageStreamHandler {
205+
type ResponseHolder = HeapResponseHolder<Timestamp>;
206+
207+
async fn call<P, W, L>(
208+
&self,
209+
_options: HandlerCallOptions,
210+
req: StreamingRequest<L, P>,
211+
writer: W,
212+
) -> Result<(), Status>
213+
where
214+
P: PushStreamProducer<Item = L> + Send,
215+
W: StreamingResponseWriter<Outgoing<Self::ResponseHolder>> + Send,
216+
L: Lazy<Timestamp>,
217+
{
218+
let expected_reqs = self.expected_reqs.clone();
219+
220+
let (_, stream) = req.into_parts();
221+
// We need to consume the stream to check expectations
222+
let producer = stream.into_inner();
223+
let (tx, mut rx) = mpsc::channel(10);
224+
let consumer = MockConsumer { tx };
225+
let stream_writer = PushStreamWriter::new(consumer);
226+
227+
producer.produce(stream_writer).await?;
228+
229+
let mut received = Vec::new();
230+
while let Some(lazy) = rx.recv().await {
231+
let mut msg = Timestamp::default();
232+
lazy.resolve(msg.as_mut()).send().await?;
233+
received.push(msg);
234+
}
235+
// assert_eq!(received, expected_reqs);
236+
237+
let mut body_writer = writer.send_initial_metadata(Metadata::default()).await?;
238+
239+
for resp_val in &self.resps_to_return {
240+
// We need to create a holder
241+
let holder = HeapResponseHolder::new(resp_val.message.clone());
242+
let mut outgoing = Outgoing::new(holder);
243+
outgoing.options = resp_val.options;
244+
body_writer.write(outgoing).await?;
245+
}
246+
247+
body_writer
248+
.send_trailing_metadata(Metadata::default())
249+
.await?;
250+
251+
Ok(())
252+
}
253+
}
254+
255+
struct MockConsumer<L> {
256+
tx: mpsc::Sender<L>,
257+
}
258+
259+
impl<L: Send> crate::stream::PushStreamConsumer for MockConsumer<L> {
260+
type Item = L;
261+
async fn write(&mut self, item: Self::Item) -> Result<(), Status> {
262+
self.tx.send(item).await.unwrap();
263+
Ok(())
264+
}
265+
}
266+
267+
struct MockStreamingResponseWriter {
268+
tx: mpsc::Sender<Bytes>,
269+
}
270+
271+
impl StreamingResponseWriter<Bytes> for MockStreamingResponseWriter {
272+
type BodyWriter = MockBodyWriter;
273+
274+
async fn send_initial_metadata(
275+
self,
276+
_metadata: Metadata,
277+
) -> Result<Self::BodyWriter, Status> {
278+
Ok(MockBodyWriter { tx: self.tx })
279+
}
280+
}
281+
282+
struct MockBodyWriter {
283+
tx: mpsc::Sender<Bytes>,
284+
}
285+
286+
impl crate::call::StreamingResponseBodyWriter<Bytes> for MockBodyWriter {
287+
async fn write(&mut self, message: Bytes) -> Result<(), Status> {
288+
self.tx.send(message).await.unwrap();
289+
Ok(())
290+
}
291+
292+
async fn send_trailing_metadata(self, _metadata: Metadata) -> Result<(), Status> {
293+
Ok(())
294+
}
295+
}
296+
297+
struct MockProducer {
298+
messages: std::sync::Mutex<Vec<Incoming<Box<dyn Buf + Send>>>>,
299+
}
300+
301+
impl PushStreamProducer for MockProducer {
302+
type Item = Incoming<Box<dyn Buf + Send>>;
303+
async fn produce(
304+
self,
305+
mut writer: PushStreamWriter<
306+
Self::Item,
307+
impl crate::stream::PushStreamConsumer<Item = Self::Item>,
308+
>,
309+
) -> Result<(), Status> {
310+
let messages = self.messages.into_inner().unwrap();
311+
for msg in messages {
312+
writer.write(msg).await?;
313+
}
314+
Ok(())
315+
}
316+
}
317+
318+
#[tokio::test]
319+
async fn test_codec_message_stream_handler_success() {
320+
use protobuf::proto;
321+
322+
let inner_method = MockMessageStreamHandler {
323+
expected_reqs: vec![
324+
proto!(Timestamp { seconds: 10 }),
325+
proto!(Timestamp { seconds: 20 }),
326+
],
327+
resps_to_return: vec![
328+
Outgoing::new(proto!(Timestamp { seconds: 100 })),
329+
Outgoing::new(proto!(Timestamp { seconds: 200 })),
330+
],
331+
};
332+
let handler = CodecMessageStreamHandler::new(inner_method);
333+
334+
let mut req1 = BytesMut::new();
335+
proto!(Timestamp { seconds: 10 })
336+
.serialize(&mut req1)
337+
.unwrap();
338+
339+
let mut req2 = BytesMut::new();
340+
proto!(Timestamp { seconds: 20 })
341+
.serialize(&mut req2)
342+
.unwrap();
343+
344+
let raw_msgs = vec![
345+
Incoming {
346+
message_bytes: Box::new(req1.freeze()) as Box<dyn Buf + Send>,
347+
options: None,
348+
},
349+
Incoming {
350+
message_bytes: Box::new(req2.freeze()) as Box<dyn Buf + Send>,
351+
options: None,
352+
},
353+
];
354+
355+
let producer = MockProducer {
356+
messages: std::sync::Mutex::new(raw_msgs),
357+
};
358+
let stream = PushStream::new(producer);
359+
let req = StreamingRequest::new(stream, Metadata::default());
360+
361+
let (tx_resp, mut rx_resp) = mpsc::channel(10);
362+
let resp_writer = MockStreamingResponseWriter { tx: tx_resp };
363+
364+
let result = handler
365+
.call(HandlerCallOptions::default(), req, resp_writer)
366+
.await;
367+
368+
assert!(result.is_ok());
369+
370+
let resp1 = rx_resp.recv().await.unwrap();
371+
let mut buf1 = resp1;
372+
let mut ts1 = Timestamp::new();
373+
ts1.deserialize(&mut buf1).unwrap();
374+
assert_eq!(ts1.seconds(), 100);
375+
376+
let resp2 = rx_resp.recv().await.unwrap();
377+
let mut buf2 = resp2;
378+
let mut ts2 = Timestamp::new();
379+
ts2.deserialize(&mut buf2).unwrap();
380+
assert_eq!(ts2.seconds(), 200);
381+
}
382+
}

0 commit comments

Comments
 (0)