1use async_trait::async_trait;
11use bytes::{Buf, BufMut, BytesMut};
12use bytesize::ByteSize;
13use futures::{SinkExt, TryStreamExt, sink};
14use mz_ore::cast::CastFrom;
15use mz_ore::future::OreSinkExt;
16use mz_ore::netio::AsyncReady;
17use mz_pgwire_common::{
18 Conn, Cursor, DecodeState, ErrorResponse, FrontendMessage, MAX_PREAUTH_FRAME_SIZE,
19 MAX_REQUEST_SIZE, Pgbuf, parse_frame_len,
20};
21use tokio::io::{self, AsyncRead, AsyncWrite, Interest, Ready};
22use tokio_util::codec::{Decoder, Encoder, Framed};
23
24#[derive(Debug)]
28pub enum BackendMessage {
29 AuthenticationCleartextPassword,
30 ErrorResponse(ErrorResponse),
31}
32
33impl From<ErrorResponse> for BackendMessage {
34 fn from(err: ErrorResponse) -> BackendMessage {
35 BackendMessage::ErrorResponse(err)
36 }
37}
38
39pub struct FramedConn<A> {
46 inner: sink::Buffer<Framed<Conn<A>, Codec>, BackendMessage>,
47}
48
49impl<A> FramedConn<A>
50where
51 A: AsyncRead + AsyncWrite + Unpin,
52{
53 pub fn new(inner: Conn<A>) -> FramedConn<A> {
59 FramedConn {
60 inner: Framed::new(inner, Codec::new()).buffer(32),
61 }
62 }
63
64 pub async fn recv(&mut self) -> Result<Option<FrontendMessage>, io::Error> {
78 let message = self.inner.try_next().await?;
79 Ok(message)
80 }
81
82 pub async fn send<M>(&mut self, message: M) -> Result<(), io::Error>
91 where
92 M: Into<BackendMessage>,
93 {
94 let message = message.into();
95 self.inner.enqueue(message).await
96 }
97
98 pub async fn flush(&mut self) -> Result<(), io::Error> {
100 self.inner.flush().await
101 }
102}
103
104impl<A> FramedConn<A>
105where
106 A: AsyncRead + AsyncWrite + Unpin,
107{
108 pub fn inner(&self) -> &Conn<A> {
109 self.inner.get_ref().get_ref()
110 }
111 pub fn inner_mut(&mut self) -> &mut Conn<A> {
112 self.inner.get_mut().get_mut()
113 }
114}
115
116#[async_trait]
117impl<A> AsyncReady for FramedConn<A>
118where
119 A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
120{
121 async fn ready(&self, interest: Interest) -> io::Result<Ready> {
122 self.inner.get_ref().get_ref().ready(interest).await
123 }
124}
125
126struct Codec {
127 decode_state: DecodeState,
128}
129
130impl Codec {
131 pub fn new() -> Codec {
133 Codec {
134 decode_state: DecodeState::Head,
135 }
136 }
137}
138
139impl Default for Codec {
140 fn default() -> Codec {
141 Codec::new()
142 }
143}
144
145impl Encoder<BackendMessage> for Codec {
146 type Error = io::Error;
147
148 fn encode(&mut self, msg: BackendMessage, dst: &mut BytesMut) -> Result<(), io::Error> {
151 let start = dst.len();
155 match self.encode_inner(msg, dst) {
156 Ok(()) => Ok(()),
157 Err(e) => {
158 dst.truncate(start);
159 Err(e)
160 }
161 }
162 }
163}
164
165impl Codec {
166 fn encode_inner(&self, msg: BackendMessage, dst: &mut BytesMut) -> Result<(), io::Error> {
169 let byte = match &msg {
171 BackendMessage::AuthenticationCleartextPassword => b'R',
172 BackendMessage::ErrorResponse(r) => {
173 if r.severity.is_error() {
174 b'E'
175 } else {
176 b'N'
177 }
178 }
179 };
180 dst.put_u8(byte);
181
182 let base = dst.len();
184 dst.put_u32(0);
185
186 match msg {
188 BackendMessage::AuthenticationCleartextPassword => {
189 dst.put_u32(3);
190 }
191 BackendMessage::ErrorResponse(ErrorResponse {
192 severity,
193 code,
194 message,
195 detail,
196 hint,
197 position,
198 }) => {
199 dst.put_u8(b'S');
200 dst.put_string(severity.as_str());
201 dst.put_u8(b'C');
202 dst.put_string(code.code());
203 dst.put_u8(b'M');
204 dst.put_string(&message);
205 if let Some(detail) = &detail {
206 dst.put_u8(b'D');
207 dst.put_string(detail);
208 }
209 if let Some(hint) = &hint {
210 dst.put_u8(b'H');
211 dst.put_string(hint);
212 }
213 if let Some(position) = &position {
214 dst.put_u8(b'P');
215 dst.put_string(&position.to_string());
216 }
217 dst.put_u8(b'\0');
218 }
219 }
220
221 let len = dst.len() - base;
222
223 let len = i32::try_from(len).map_err(|_| {
225 io::Error::new(
226 io::ErrorKind::InvalidData,
227 "length of encoded message does not fit into an i32",
228 )
229 })?;
230 dst[base..base + 4].copy_from_slice(&len.to_be_bytes());
231
232 Ok(())
233 }
234}
235
236impl Decoder for Codec {
237 type Item = FrontendMessage;
238 type Error = io::Error;
239
240 fn decode(&mut self, src: &mut BytesMut) -> Result<Option<FrontendMessage>, io::Error> {
241 if src.len() > MAX_REQUEST_SIZE {
242 return Err(io::Error::new(
243 io::ErrorKind::InvalidData,
244 format!(
245 "request larger than {}",
246 ByteSize::b(u64::cast_from(MAX_REQUEST_SIZE))
247 ),
248 ));
249 }
250 loop {
251 match self.decode_state {
252 DecodeState::Head => {
253 if src.len() < 5 {
254 return Ok(None);
255 }
256 let msg_type = src[0];
257 let frame_len = parse_frame_len(&src[1..], MAX_PREAUTH_FRAME_SIZE)?;
258 src.advance(5);
259 src.reserve(frame_len);
260 self.decode_state = DecodeState::Data(msg_type, frame_len);
261 }
262
263 DecodeState::Data(msg_type, frame_len) => {
264 if src.len() < frame_len {
265 return Ok(None);
266 }
267 let buf = src.split_to(frame_len).freeze();
268 let buf = Cursor::new(&buf);
269 let msg = match msg_type {
270 b'X' => decode_terminate(buf)?,
272
273 b'p' => decode_password(buf)?,
275
276 _ => {
278 return Err(io::Error::new(
279 io::ErrorKind::InvalidData,
280 format!("unknown message type {}", msg_type),
281 ));
282 }
283 };
284 src.reserve(5);
285 self.decode_state = DecodeState::Head;
286 return Ok(Some(msg));
287 }
288 }
289 }
290 }
291}
292
293fn decode_terminate(mut _buf: Cursor) -> Result<FrontendMessage, io::Error> {
294 Ok(FrontendMessage::Terminate)
296}
297
298fn decode_password(mut buf: Cursor) -> Result<FrontendMessage, io::Error> {
299 Ok(FrontendMessage::Password {
300 password: buf.read_cstr()?.to_owned(),
301 })
302}