tako_rs_core/grpc/
streaming.rs1use std::convert::Infallible;
6use std::pin::Pin;
7use std::task::Context;
8use std::task::Poll;
9
10use bytes::Bytes;
11use bytes::BytesMut;
12use futures_util::Stream;
13use futures_util::StreamExt;
14use http::HeaderMap;
15use http::StatusCode;
16use http_body::Frame;
17use http_body_util::StreamBody;
18use prost::Message;
19
20use super::GrpcError;
21use super::framing::MAX_GRPC_MESSAGE_SIZE;
22use super::framing::grpc_encode;
23use super::status::GrpcStatus;
24use crate::body::TakoBody;
25use crate::extractors::FromRequest;
26use crate::responder::Responder;
27use crate::types::Request;
28use crate::types::Response;
29
30pub struct GrpcServerStream<S, T>
36where
37 S: Stream<Item = Result<T, GrpcStatus>> + Send + 'static,
38 T: Message + Send + 'static,
39{
40 pub stream: S,
41 pub initial_metadata: HeaderMap,
43}
44
45impl<S, T> GrpcServerStream<S, T>
46where
47 S: Stream<Item = Result<T, GrpcStatus>> + Send + 'static,
48 T: Message + Send + 'static,
49{
50 pub fn new(stream: S) -> Self {
51 Self {
52 stream,
53 initial_metadata: HeaderMap::new(),
54 }
55 }
56
57 pub fn with_metadata(mut self, headers: HeaderMap) -> Self {
58 self.initial_metadata = headers;
59 self
60 }
61}
62
63impl<S, T> Responder for GrpcServerStream<S, T>
64where
65 S: Stream<Item = Result<T, GrpcStatus>> + Send + 'static,
66 T: Message + Send + 'static,
67{
68 fn into_response(self) -> Response {
69 use std::sync::Arc;
70 use std::sync::atomic::AtomicBool;
71 use std::sync::atomic::Ordering;
72
73 let error_emitted = Arc::new(AtomicBool::new(false));
78 let mark_err = error_emitted.clone();
79 let stream = self.stream.map(move |item| match item {
80 Ok(msg) => {
81 let bytes = grpc_encode(&msg);
82 Ok::<_, Infallible>(Frame::data(Bytes::from(bytes)))
83 }
84 Err(status) => {
85 mark_err.store(true, Ordering::Release);
86 Ok(Frame::trailers(status.write_trailers()))
87 }
88 });
89
90 let check_err = error_emitted.clone();
93 let mut once = false;
94 let trailer = futures_util::stream::iter(std::iter::from_fn(move || {
95 if once {
96 None
97 } else {
98 once = true;
99 if check_err.load(Ordering::Acquire) {
100 None
101 } else {
102 Some(Ok::<_, Infallible>(Frame::trailers(
103 GrpcStatus::ok().write_trailers(),
104 )))
105 }
106 }
107 }));
108 let combined = stream.chain(trailer);
109
110 let mut resp = http::Response::builder()
122 .status(StatusCode::OK)
123 .header(
124 http::header::CONTENT_TYPE,
125 http::HeaderValue::from_static("application/grpc"),
126 )
127 .body(TakoBody::new(StreamBody::new(combined)))
128 .expect("static headers + body construction is infallible");
129 let headers = resp.headers_mut();
130 for (k, v) in &self.initial_metadata {
131 headers.insert(k.clone(), v.clone());
132 }
133 resp
134 }
135}
136
137pub struct GrpcClientStream<T: Message + Default + Send + 'static> {
142 pub stream: Pin<Box<dyn Stream<Item = Result<T, GrpcError>> + Send>>,
143}
144
145impl<'a, T> FromRequest<'a> for GrpcClientStream<T>
146where
147 T: Message + Default + Send + 'static,
148{
149 type Error = GrpcError;
150
151 fn from_request(
152 req: &'a mut Request,
153 ) -> impl core::future::Future<Output = core::result::Result<Self, Self::Error>> + Send + 'a {
154 async move {
155 let ct = req
156 .headers()
157 .get(http::header::CONTENT_TYPE)
158 .and_then(|v| v.to_str().ok())
159 .unwrap_or("");
160 if !ct.starts_with("application/grpc") {
161 return Err(GrpcError::InvalidContentType);
162 }
163
164 let body = std::mem::take(req.body_mut());
168 let stream = GrpcFrameStream::new(body);
169 Ok(GrpcClientStream {
170 stream: Box::pin(stream),
171 })
172 }
173 }
174}
175
176struct GrpcFrameStream<T> {
177 body: TakoBody,
178 buffer: BytesMut,
179 finished: bool,
180 _marker: std::marker::PhantomData<fn() -> T>,
181}
182
183impl<T> GrpcFrameStream<T> {
184 fn new(body: TakoBody) -> Self {
185 Self {
186 body,
187 buffer: BytesMut::new(),
188 finished: false,
189 _marker: std::marker::PhantomData,
190 }
191 }
192}
193
194impl<T> Stream for GrpcFrameStream<T>
195where
196 T: Message + Default,
197{
198 type Item = Result<T, GrpcError>;
199
200 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
201 let this = self.get_mut();
202 loop {
203 if this.buffer.len() >= 5 {
205 let msg_len = u32::from_be_bytes([
206 this.buffer[1],
207 this.buffer[2],
208 this.buffer[3],
209 this.buffer[4],
210 ]) as usize;
211 if msg_len > MAX_GRPC_MESSAGE_SIZE {
212 return Poll::Ready(Some(Err(GrpcError::MessageTooLarge)));
213 }
214 if this.buffer.len() >= 5 + msg_len {
215 if this.buffer[0] != 0 {
216 return Poll::Ready(Some(Err(GrpcError::CompressionUnsupported)));
217 }
218 let payload = this.buffer.split_to(5 + msg_len);
219 let msg_bytes = &payload[5..5 + msg_len];
220 return match T::decode(msg_bytes) {
221 Ok(m) => Poll::Ready(Some(Ok(m))),
222 Err(e) => Poll::Ready(Some(Err(GrpcError::DecodeError(e.to_string())))),
223 };
224 }
225 }
226
227 if this.finished {
228 return Poll::Ready(None);
229 }
230
231 let mut body = Pin::new(&mut this.body);
233 match http_body::Body::poll_frame(body.as_mut(), cx) {
234 Poll::Ready(Some(Ok(frame))) => {
235 if let Some(data) = frame.data_ref() {
236 this.buffer.extend_from_slice(data);
237 }
238 }
239 Poll::Ready(Some(Err(e))) => {
240 return Poll::Ready(Some(Err(GrpcError::BodyReadError(e.to_string()))));
241 }
242 Poll::Ready(None) => {
243 this.finished = true;
244 }
245 Poll::Pending => return Poll::Pending,
246 }
247 }
248 }
249}
250
251pub struct GrpcBidi<Req, Resp>
257where
258 Req: Message + Default + Send + 'static,
259 Resp: Message + Send + 'static,
260{
261 pub inbound: GrpcClientStream<Req>,
262 pub _phantom: std::marker::PhantomData<Resp>,
263}
264
265impl<'a, Req, Resp> FromRequest<'a> for GrpcBidi<Req, Resp>
266where
267 Req: Message + Default + Send + 'static,
268 Resp: Message + Send + 'static,
269{
270 type Error = GrpcError;
271
272 fn from_request(
273 req: &'a mut Request,
274 ) -> impl core::future::Future<Output = core::result::Result<Self, Self::Error>> + Send + 'a {
275 async move {
276 Ok(GrpcBidi {
277 inbound: GrpcClientStream::<Req>::from_request(req).await?,
278 _phantom: std::marker::PhantomData,
279 })
280 }
281 }
282}