1use std::convert::TryInto;
2use std::{mem, str};
3
4use bytes::{self, BufMut, BytesMut};
5use futures::{sink::SinkExt, stream::StreamExt};
6use serde_bytes::Bytes;
7use tokio::net::{TcpStream, ToSocketAddrs};
8use tokio_util::codec::{Decoder, Encoder, Framed};
9
10use super::{amqp, de, sasl, ser, Error};
11
12pub struct Client {
13 transport: tokio_util::codec::Framed<TcpStream, Codec>,
14}
15
16impl Client {
17 pub async fn connect<A: ToSocketAddrs>(addr: A) -> Result<Self, ()> {
18 let stream = TcpStream::connect(addr).await.map_err(|_| ())?;
19 Ok(Self {
20 transport: Framed::new(stream, Codec),
21 })
22 }
23
24 pub async fn login(&mut self, user: &str, password: &str) -> Result<(), ()> {
28 self.transport
29 .send(&Frame::Header(Protocol::Sasl))
30 .await
31 .map_err(|_| ())?;
32 let _header = self.transport.next().await.ok_or(()).map_err(|_| ())?;
33 let _mechanisms = self.transport.next().await.ok_or(()).map_err(|_| ())?;
34
35 let mut response = vec![0u8];
36 response.extend_from_slice(user.as_bytes());
37 response.push(0);
38 response.extend_from_slice(password.as_bytes());
39
40 let init = Frame::Sasl(sasl::Frame::Init(sasl::Init {
41 mechanism: sasl::Mechanism::Plain,
42 initial_response: Some(Bytes::new(&response)),
43 hostname: None,
44 }));
45
46 self.transport.send(&init).await.map_err(|_| ())?;
47 let _outcome = self.transport.next().await.ok_or(()).map_err(|_| ())?;
48 let _header = self.transport.next().await.ok_or(()).map_err(|_| ())?;
49 self.transport
50 .send(&Frame::Header(Protocol::Amqp))
51 .await
52 .map_err(|_| ())
53 }
54
55 pub async fn open(&mut self, container_id: &str) -> Result<(), ()> {
56 let open = Frame::Amqp(amqp::Frame {
57 channel: 0,
58 extended_header: None,
59 performative: amqp::Performative::Open(amqp::Open {
60 container_id,
61 ..Default::default()
62 }),
63 message: None,
64 });
65
66 self.transport.send(&open).await.map_err(|_| ())?;
67 let _opened = self.transport.next().await.ok_or(()).map_err(|_| ())?;
68 Ok(())
69 }
70
71 pub async fn begin(&mut self) -> Result<(), ()> {
72 let begin = Frame::Amqp(amqp::Frame {
73 channel: 0,
74 extended_header: None,
75 performative: amqp::Performative::Begin(amqp::Begin {
76 remote_channel: None,
77 next_outgoing_id: 1,
78 incoming_window: 8,
79 outgoing_window: 8,
80 ..Default::default()
81 }),
82 message: None,
83 });
84
85 self.transport.send(&begin).await.map_err(|_| ())?;
86 let _begun = self.transport.next().await.ok_or(()).map_err(|_| ())?;
87 Ok(())
88 }
89
90 pub async fn attach(&mut self, attach: amqp::Attach<'_>) -> Result<(), ()> {
91 let is_sender = matches!(attach.role, amqp::Role::Sender);
92 let attach = Frame::Amqp(amqp::Frame {
93 channel: 0,
94 extended_header: None,
95 performative: amqp::Performative::Attach(attach),
96 message: None,
97 });
98
99 self.transport.send(&attach).await.map_err(|_| ())?;
100 let _attached = self.transport.next().await.ok_or(()).map_err(|_| ())?;
101 if is_sender {
102 let _flow = self.transport.next().await.ok_or(()).map_err(|_| ())?;
103 }
104
105 Ok(())
106 }
107
108 pub async fn flow(&mut self, flow: amqp::Flow<'_>) -> Result<(), ()> {
109 let flow = Frame::Amqp(amqp::Frame {
110 channel: 0,
111 extended_header: None,
112 performative: amqp::Performative::Flow(flow),
113 message: None,
114 });
115
116 self.transport.send(&flow).await.map_err(|_| ())?;
117 Ok(())
118 }
119
120 pub async fn transfer(
121 &mut self,
122 transfer: amqp::Transfer,
123 message: amqp::Message<'_>,
124 ) -> Result<(), ()> {
125 let transfer = Frame::Amqp(amqp::Frame {
126 channel: 0,
127 extended_header: None,
128 performative: amqp::Performative::Transfer(transfer),
129 message: Some(message),
130 });
131
132 self.transport.send(&transfer).await.map_err(|_| ())?;
133 let _transferred = self.transport.next().await.ok_or(()).map_err(|_| ())?;
134 Ok(())
135 }
136
137 #[allow(clippy::should_implement_trait)]
138 pub async fn next(&mut self) -> Option<Result<BytesFrame, Error>> {
139 self.transport.next().await
140 }
141}
142
143pub struct Codec;
144
145impl Decoder for Codec {
146 type Item = BytesFrame;
147 type Error = Error;
148
149 fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
150 if src.len() < 4 {
151 return Ok(None);
152 }
153
154 let length_or_proto_tag = &src[..4];
155 let bytes = if length_or_proto_tag == b"AMQP" && src.len() >= PROTO_HEADER_LENGTH {
156 src.split_to(PROTO_HEADER_LENGTH).freeze()
157 } else {
158 let len = u32::from_be_bytes((length_or_proto_tag).try_into().unwrap()) as usize;
159 if src.len() >= len {
160 src.split_to(len).freeze().split_off(4)
161 } else {
162 return Ok(None);
163 }
164 };
165
166 let frame = unsafe { mem::transmute(Frame::decode(&bytes)?) };
167 Ok(Some(BytesFrame { bytes, frame }))
168 }
169}
170
171impl Encoder<&Frame<'_>> for Codec {
172 type Error = Error;
173
174 fn encode(&mut self, item: &Frame<'_>, dst: &mut BytesMut) -> Result<(), Self::Error> {
175 let buf = item.to_vec().unwrap();
176 dst.put(&*buf);
177 Ok(())
178 }
179}
180
181pub struct BytesFrame {
182 #[allow(dead_code)]
183 bytes: bytes::Bytes,
184 frame: Frame<'static>,
185}
186
187impl BytesFrame {
188 #[allow(clippy::needless_lifetimes)]
189 pub fn frame<'a>(&'a self) -> &'a Frame<'a> {
190 &self.frame
191 }
192
193 #[allow(clippy::needless_lifetimes)]
194 pub fn body<'a>(&'a self) -> Option<&'a [u8]> {
195 let message = match self.frame() {
196 Frame::Amqp(amqp::Frame {
197 message: Some(msg), ..
198 }) => msg,
199 _ => return None,
200 };
201
202 match message.body {
203 Some(amqp::Body::Data(amqp::Data(data))) => Some(data),
204 Some(amqp::Body::Value(amqp::Value(amqp::Any::Bytes(data)))) => Some(data),
205 _ => None,
206 }
207 }
208}
209
210impl std::fmt::Debug for BytesFrame {
211 fn fmt(&self, fmt: &mut std::fmt::Formatter) -> std::fmt::Result {
212 self.frame.fmt(fmt)
213 }
214}
215
216#[allow(clippy::large_enum_variant)]
217#[derive(Debug, PartialEq)]
218pub enum Frame<'a> {
219 Amqp(amqp::Frame<'a>),
220 Header(Protocol),
221 Sasl(sasl::Frame<'a>),
222}
223
224impl<'a> Frame<'a> {
225 pub fn decode(buf: &'a [u8]) -> Result<Self, Error> {
226 if &buf[..4] == b"AMQP" {
227 return Ok(Frame::Header(Protocol::from_bytes(buf)));
228 }
229
230 let doff = buf[0];
231 if doff < 2 {
232 return Err(Error::InvalidData);
233 }
234
235 let result = match buf[1] {
236 0x00 => Ok(Frame::Amqp(amqp::Frame::decode(doff, &buf[2..])?)),
237 0x01 => {
238 assert_eq!(&buf[2..4], &[0, 0]);
239 let (sasl, rest) = de::deserialize(&buf[4..])?;
240 if !rest.is_empty() {
241 return Err(Error::TrailingCharacters);
242 }
243 Ok(Frame::Sasl(sasl))
244 }
245 _ => Err(Error::InvalidData),
246 };
247
248 if result.is_err() {
249 println!("failed to decode: {:?}", buf);
250 }
251 result
252 }
253
254 pub fn to_vec(&self) -> Result<Vec<u8>, Error> {
255 let mut buf = vec![0; 8];
256
257 match self {
258 Frame::Amqp(f) => {
259 buf[5] = 0x00;
260 ser::into_bytes(&f.performative, &mut buf)?;
261 if let Some(msg) = &f.message {
262 if let Some(header) = &msg.header {
263 ser::into_bytes(header, &mut buf)?;
264 }
265 if let Some(da) = &msg.delivery_annotations {
266 ser::into_bytes(da, &mut buf)?;
267 }
268 if let Some(ma) = &msg.message_annotations {
269 ser::into_bytes(ma, &mut buf)?;
270 }
271 if let Some(props) = &msg.properties {
272 ser::into_bytes(props, &mut buf)?;
273 }
274 if let Some(ap) = &msg.application_properties {
275 ser::into_bytes(ap, &mut buf)?;
276 }
277 ser::into_bytes(&msg.body, &mut buf)?;
278 if let Some(footer) = &msg.footer {
279 ser::into_bytes(footer, &mut buf)?;
280 }
281 }
282 (&mut buf[6..8]).copy_from_slice(&f.channel.to_be_bytes()[..]);
283 }
284 Frame::Header(p) => {
285 buf.copy_from_slice(p.header());
286 return Ok(buf);
287 }
288 Frame::Sasl(f) => {
289 buf[5] = 0x01;
290 ser::into_bytes(f, &mut buf).unwrap();
291 }
292 }
293
294 buf[4] = 2; let len = buf.len() as u32;
296 (&mut buf[..4]).copy_from_slice(&len.to_be_bytes()[..]);
297 Ok(buf)
298 }
299}
300
301#[derive(Copy, Clone, Debug, Eq, PartialEq)]
302pub enum Protocol {
303 Sasl,
304 Amqp,
305}
306
307impl Protocol {
308 fn from_bytes(bytes: &[u8]) -> Self {
309 match bytes {
310 SASL_PROTO_HEADER => Protocol::Sasl,
311 AMQP_PROTO_HEADER => Protocol::Amqp,
312 p => panic!("invalid protocol header {:?}", p),
313 }
314 }
315
316 fn header(self) -> &'static [u8] {
317 match self {
318 Protocol::Sasl => SASL_PROTO_HEADER,
319 Protocol::Amqp => AMQP_PROTO_HEADER,
320 }
321 }
322}
323
324pub const AMQP_PROTO_HEADER: &[u8] = b"AMQP\x00\x01\x00\x00";
366pub const SASL_PROTO_HEADER: &[u8] = b"AMQP\x03\x01\x00\x00";
367pub const PROTO_HEADER_LENGTH: usize = 8;