Skip to main content

wasi_pg_client/protocol/
frontend.rs

1use bytes::BytesMut;
2use postgres_protocol::IsNull;
3
4use super::error::ProtocolError;
5use super::types::FormatCode;
6use crate::protocol::Oid;
7
8#[derive(Debug, Clone, PartialEq)]
9pub enum FrontendMessage {
10    Startup {
11        params: Vec<(String, String)>,
12    },
13
14    SslRequest,
15
16    CancelRequest {
17        process_id: i32,
18        secret_key: i32,
19    },
20
21    Query {
22        sql: String,
23    },
24
25    Parse {
26        name: String,
27        sql: String,
28        param_types: Vec<Oid>,
29    },
30
31    Bind {
32        portal: String,
33        statement: String,
34        param_formats: Vec<FormatCode>,
35        params: Vec<Option<Vec<u8>>>,
36        result_formats: Vec<FormatCode>,
37    },
38
39    Describe {
40        variant: u8,
41        name: String,
42    },
43
44    Execute {
45        portal: String,
46        max_rows: i32,
47    },
48
49    Close {
50        variant: u8,
51        name: String,
52    },
53
54    Sync,
55
56    Flush,
57
58    CopyData {
59        data: Vec<u8>,
60    },
61
62    CopyDone,
63
64    CopyFail {
65        message: String,
66    },
67
68    PasswordMessage {
69        password: Vec<u8>,
70    },
71
72    SaslInitialResponse {
73        mechanism: String,
74        data: Vec<u8>,
75    },
76
77    SaslResponse {
78        data: Vec<u8>,
79    },
80
81    Terminate,
82}
83
84#[derive(Debug)]
85pub struct MessageEncoder;
86
87impl MessageEncoder {
88    pub fn encode(msg: &FrontendMessage, buf: &mut BytesMut) -> Result<(), ProtocolError> {
89        match msg {
90            FrontendMessage::Startup { params } => {
91                let pairs: Vec<(&str, &str)> = params
92                    .iter()
93                    .map(|(k, v)| (k.as_str(), v.as_str()))
94                    .collect();
95                postgres_protocol::message::frontend::startup_message(pairs, buf)?;
96            }
97            FrontendMessage::SslRequest => {
98                postgres_protocol::message::frontend::ssl_request(buf);
99            }
100            FrontendMessage::CancelRequest {
101                process_id,
102                secret_key,
103            } => {
104                postgres_protocol::message::frontend::cancel_request(*process_id, *secret_key, buf);
105            }
106            FrontendMessage::Query { sql } => {
107                postgres_protocol::message::frontend::query(sql, buf)?;
108            }
109            FrontendMessage::Parse {
110                name,
111                sql,
112                param_types,
113            } => {
114                postgres_protocol::message::frontend::parse(
115                    name,
116                    sql,
117                    param_types.iter().copied(),
118                    buf,
119                )?;
120            }
121            FrontendMessage::Bind {
122                portal,
123                statement,
124                param_formats,
125                params,
126                result_formats,
127            } => {
128                let pf: Vec<i16> = param_formats.iter().map(|f| *f as i16).collect();
129                let rf: Vec<i16> = result_formats.iter().map(|f| *f as i16).collect();
130                postgres_protocol::message::frontend::bind(
131                    portal,
132                    statement,
133                    pf,
134                    params.iter().map(|p| p.as_ref().map(|v| v.as_slice())),
135                    |item, buf| match item {
136                        Some(data) => {
137                            buf.extend_from_slice(data);
138                            Ok::<_, Box<dyn std::error::Error + Sync + Send>>(IsNull::No)
139                        }
140                        None => Ok(IsNull::Yes),
141                    },
142                    rf,
143                    buf,
144                )
145                .map_err(|_e| {
146                    ProtocolError::Io(std::io::Error::new(
147                        std::io::ErrorKind::InvalidInput,
148                        "bind encoding failed",
149                    ))
150                })?;
151            }
152            FrontendMessage::Describe { variant, name } => {
153                postgres_protocol::message::frontend::describe(*variant, name, buf)?;
154            }
155            FrontendMessage::Execute { portal, max_rows } => {
156                postgres_protocol::message::frontend::execute(portal, *max_rows, buf)?;
157            }
158            FrontendMessage::Close { variant, name } => {
159                postgres_protocol::message::frontend::close(*variant, name, buf)?;
160            }
161            FrontendMessage::Sync => {
162                postgres_protocol::message::frontend::sync(buf);
163            }
164            FrontendMessage::Flush => {
165                postgres_protocol::message::frontend::flush(buf);
166            }
167            FrontendMessage::CopyData { data } => {
168                postgres_protocol::message::frontend::CopyData::new(data.as_slice())?.write(buf);
169            }
170            FrontendMessage::CopyDone => {
171                postgres_protocol::message::frontend::copy_done(buf);
172            }
173            FrontendMessage::CopyFail { message } => {
174                postgres_protocol::message::frontend::copy_fail(message, buf)?;
175            }
176            FrontendMessage::PasswordMessage { password } => {
177                postgres_protocol::message::frontend::password_message(password, buf)?;
178            }
179            FrontendMessage::SaslInitialResponse { mechanism, data } => {
180                postgres_protocol::message::frontend::sasl_initial_response(mechanism, data, buf)?;
181            }
182            FrontendMessage::SaslResponse { data } => {
183                postgres_protocol::message::frontend::sasl_response(data, buf)?;
184            }
185            FrontendMessage::Terminate => {
186                postgres_protocol::message::frontend::terminate(buf);
187            }
188        }
189        Ok(())
190    }
191}