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>;