Skip to main content

oasis_amqp/
proto.rs

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    /// Login with the given username and password
25    ///
26    /// Currently this only supports SASL PLAIN login.
27    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; // doff
295        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
324/*
325
326#[derive(Debug)]
327enum ConnectionState {
328    Start,
329    HdrRcvd,
330    HdrSent,
331    HdrExch,
332    OpenPipe,
333    OcPipe,
334    OpenRcvd,
335    OpenSent,
336    ClosePipe,
337    Opened,
338    CloseRcvd,
339    CloseSent,
340    Discarding,
341    End,
342}
343
344struct Session {
345    pub next_incoming_id: u32,
346    pub incoming_window: u32,
347    pub next_outgoing_id: u32,
348    pub outgoing_window: u32,
349    pub remote_incoming_window: u32,
350    pub remote_outgoing_window: u32,
351}
352
353enum SessionState {
354    Unmapped,
355    BeginSent,
356    BeginRcvd,
357    Mapped,
358    EndSent,
359    EndRcvd,
360    Discarding,
361}
362
363*/
364
365pub 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;