Skip to main content

tonic_rpc/
codec.rs

1use std::{
2    io::{Read, Write},
3    marker::PhantomData,
4};
5
6use bytes::{buf::BufMut, Buf};
7use serde::{Deserialize, Serialize};
8use tonic::{codec, Status};
9
10pub trait SerdeCodec {
11    fn write<T, W>(item: T, w: W) -> Result<(), Status>
12    where
13        T: Serialize,
14        W: Write;
15
16    fn read<T, R>(r: R) -> Result<T, Status>
17    where
18        T: for<'de> Deserialize<'de>,
19        R: Read;
20}
21
22#[derive(Clone, Copy)]
23pub struct Encoder<C, T> {
24    _pd: PhantomData<(C, T)>,
25}
26
27impl<C, T> codec::Encoder for Encoder<C, T>
28where
29    T: Serialize,
30    C: SerdeCodec,
31{
32    type Item = T;
33    type Error = Status;
34    fn encode(
35        &mut self,
36        item: Self::Item,
37        dst: &mut codec::EncodeBuf<'_>,
38    ) -> Result<(), Self::Error> {
39        C::write(item, dst.writer())
40    }
41}
42
43#[derive(Clone, Copy)]
44pub struct Decoder<C, T> {
45    _pd: PhantomData<(C, T)>,
46}
47
48impl<C, T> codec::Decoder for Decoder<C, T>
49where
50    T: for<'de> Deserialize<'de>,
51    C: SerdeCodec,
52{
53    type Item = T;
54    type Error = Status;
55    fn decode(
56        &mut self,
57        src: &mut codec::DecodeBuf<'_>,
58    ) -> Result<Option<Self::Item>, Self::Error> {
59        Ok(Some(C::read::<T, _>(src.reader())?))
60    }
61}
62
63pub struct Codec<C, T, U> {
64    _pd: PhantomData<(C, T, U)>,
65}
66
67impl<C, T, U> Default for Codec<C, T, U> {
68    fn default() -> Self {
69        Codec { _pd: PhantomData }
70    }
71}
72
73impl<C, T, U> codec::Codec for Codec<C, T, U>
74where
75    C: SerdeCodec + Send + Sync + 'static,
76    T: Serialize + Send + Sync + 'static,
77    U: for<'de> Deserialize<'de> + Send + Sync + 'static,
78{
79    type Encode = T;
80    type Decode = U;
81    type Encoder = Encoder<C, T>;
82    type Decoder = Decoder<C, U>;
83
84    fn encoder(&mut self) -> Self::Encoder {
85        Encoder { _pd: PhantomData }
86    }
87
88    fn decoder(&mut self) -> Self::Decoder {
89        Decoder { _pd: PhantomData }
90    }
91}
92
93#[cfg(feature = "bincode")]
94#[cfg_attr(docsrs, doc(cfg(feature = "bincode")))]
95pub struct BincodeSerdeCodec;
96#[cfg(feature = "cbor")]
97#[cfg_attr(docsrs, doc(cfg(feature = "cbor")))]
98pub struct CborSerdeCodec;
99#[cfg(feature = "json")]
100#[cfg_attr(docsrs, doc(cfg(feature = "json")))]
101pub struct JsonSerdeCodec;
102#[cfg(feature = "messagepack")]
103#[cfg_attr(docsrs, doc(cfg(feature = "messagepack")))]
104pub struct MessagePackSerdeCodec;
105
106#[cfg(feature = "bincode")]
107#[cfg_attr(docsrs, doc(cfg(feature = "bincode")))]
108impl SerdeCodec for BincodeSerdeCodec {
109    fn write<T, W>(item: T, w: W) -> Result<(), Status>
110    where
111        T: Serialize,
112        W: Write,
113    {
114        bincode::serialize_into(w, &item)
115            .map_err(|bincode_err| Status::internal(format!("Error serializing {}", bincode_err)))
116    }
117
118    fn read<T, R>(r: R) -> Result<T, Status>
119    where
120        T: for<'de> Deserialize<'de>,
121        R: Read,
122    {
123        bincode::deserialize_from(r)
124            .map_err(|bincode_err| Status::internal(format!("Error deserializing {}", bincode_err)))
125    }
126}
127
128#[cfg(feature = "cbor")]
129#[cfg_attr(docsrs, doc(cfg(feature = "cbor")))]
130impl SerdeCodec for CborSerdeCodec {
131    fn write<T, W>(item: T, w: W) -> Result<(), Status>
132    where
133        T: Serialize,
134        W: Write,
135    {
136        serde_cbor::to_writer(w, &item)
137            .map_err(|serde_err| Status::internal(format!("Error serializing {}", serde_err)))
138    }
139
140    fn read<T, R>(r: R) -> Result<T, Status>
141    where
142        T: for<'de> Deserialize<'de>,
143        R: Read,
144    {
145        serde_cbor::from_reader(r)
146            .map_err(|serde_err| Status::internal(format!("Error deserializing {}", serde_err)))
147    }
148}
149
150#[cfg(feature = "json")]
151#[cfg_attr(docsrs, doc(cfg(feature = "json")))]
152impl SerdeCodec for JsonSerdeCodec {
153    fn write<T, W>(item: T, w: W) -> Result<(), Status>
154    where
155        T: Serialize,
156        W: Write,
157    {
158        serde_json::to_writer(w, &item)
159            .map_err(|serde_err| Status::internal(format!("Error serializing {}", serde_err)))
160    }
161
162    fn read<T, R>(r: R) -> Result<T, Status>
163    where
164        T: for<'de> Deserialize<'de>,
165        R: Read,
166    {
167        serde_json::from_reader(r)
168            .map_err(|serde_err| Status::internal(format!("Error deserializing {}", serde_err)))
169    }
170}
171
172#[cfg(feature = "messagepack")]
173#[cfg_attr(docsrs, doc(cfg(feature = "messagepack")))]
174impl SerdeCodec for MessagePackSerdeCodec {
175    fn write<T, W>(item: T, mut w: W) -> Result<(), Status>
176    where
177        T: Serialize,
178        W: Write,
179    {
180        rmp_serde::encode::write(&mut w, &item).map_err(|message_pack_err| {
181            Status::internal(format!("Error serializing {}", message_pack_err))
182        })
183    }
184
185    fn read<T, R>(r: R) -> Result<T, Status>
186    where
187        T: for<'de> Deserialize<'de>,
188        R: Read,
189    {
190        rmp_serde::from_read(r).map_err(|message_pack_err| {
191            Status::internal(format!("Error deserializing {}", message_pack_err))
192        })
193    }
194}
195
196#[cfg(feature = "bincode")]
197#[cfg_attr(docsrs, doc(cfg(feature = "bincode")))]
198pub type BincodeCodec<T, U> = Codec<BincodeSerdeCodec, T, U>;
199#[cfg(feature = "cbor")]
200#[cfg_attr(docsrs, doc(cfg(feature = "cbor")))]
201pub type CborCodec<T, U> = Codec<CborSerdeCodec, T, U>;
202#[cfg(feature = "json")]
203#[cfg_attr(docsrs, doc(cfg(feature = "json")))]
204pub type JsonCodec<T, U> = Codec<JsonSerdeCodec, T, U>;
205#[cfg(feature = "messagepack")]
206#[cfg_attr(docsrs, doc(cfg(feature = "messagepack")))]
207pub type MessagePackCodec<T, U> = Codec<MessagePackSerdeCodec, T, U>;