Skip to main content

qail_pg/protocol/wire/
frontend.rs

1//! FrontendMessage encoder — client-to-server wire format.
2
3use super::types::*;
4
5impl FrontendMessage {
6    #[inline]
7    fn has_nul(s: &str) -> bool {
8        s.as_bytes().contains(&0)
9    }
10
11    #[inline]
12    fn content_len_to_wire_len(content_len: usize) -> Result<i32, FrontendEncodeError> {
13        let total = content_len
14            .checked_add(4)
15            .ok_or(FrontendEncodeError::MessageTooLarge(usize::MAX))?;
16        i32::try_from(total).map_err(|_| FrontendEncodeError::MessageTooLarge(total))
17    }
18
19    /// Fallible encoder that returns explicit reason on invalid input.
20    pub fn encode_checked(&self) -> Result<Vec<u8>, FrontendEncodeError> {
21        match self {
22            FrontendMessage::Startup {
23                user,
24                database,
25                protocol_version,
26                startup_params,
27            } => {
28                if Self::has_nul(user) {
29                    return Err(FrontendEncodeError::InteriorNul("user"));
30                }
31                if Self::has_nul(database) {
32                    return Err(FrontendEncodeError::InteriorNul("database"));
33                }
34                let mut seen_startup_keys = std::collections::HashSet::new();
35                let mut buf = Vec::new();
36                buf.extend_from_slice(&protocol_version.to_be_bytes());
37                buf.extend_from_slice(b"user\0");
38                buf.extend_from_slice(user.as_bytes());
39                buf.push(0);
40                buf.extend_from_slice(b"database\0");
41                buf.extend_from_slice(database.as_bytes());
42                buf.push(0);
43                for (key, value) in startup_params {
44                    let key_trimmed = key.trim();
45                    if key_trimmed.is_empty() {
46                        return Err(FrontendEncodeError::InvalidStartupParam(
47                            "key must not be empty".to_string(),
48                        ));
49                    }
50                    if key_trimmed != key {
51                        return Err(FrontendEncodeError::InvalidStartupParam(format!(
52                            "key contains leading/trailing whitespace: '{}'",
53                            key
54                        )));
55                    }
56                    let key_lc = key_trimmed.to_ascii_lowercase();
57                    if key_lc == "user" || key_lc == "database" {
58                        return Err(FrontendEncodeError::InvalidStartupParam(format!(
59                            "reserved key '{}'",
60                            key_trimmed
61                        )));
62                    }
63                    if !seen_startup_keys.insert(key_lc) {
64                        return Err(FrontendEncodeError::InvalidStartupParam(format!(
65                            "duplicate key '{}'",
66                            key_trimmed
67                        )));
68                    }
69                    if Self::has_nul(key) {
70                        return Err(FrontendEncodeError::InteriorNul("startup_param_key"));
71                    }
72                    if Self::has_nul(value) {
73                        return Err(FrontendEncodeError::InteriorNul("startup_param_value"));
74                    }
75                    buf.extend_from_slice(key_trimmed.as_bytes());
76                    buf.push(0);
77                    buf.extend_from_slice(value.as_bytes());
78                    buf.push(0);
79                }
80                buf.push(0);
81
82                let len = Self::content_len_to_wire_len(buf.len())?;
83                let mut result = len.to_be_bytes().to_vec();
84                result.extend(buf);
85                Ok(result)
86            }
87            FrontendMessage::Query(sql) => {
88                if Self::has_nul(sql) {
89                    return Err(FrontendEncodeError::InteriorNul("sql"));
90                }
91                let mut buf = Vec::new();
92                buf.push(b'Q');
93                let mut content = Vec::with_capacity(sql.len() + 1);
94                content.extend_from_slice(sql.as_bytes());
95                content.push(0);
96                let len = Self::content_len_to_wire_len(content.len())?;
97                buf.extend_from_slice(&len.to_be_bytes());
98                buf.extend_from_slice(&content);
99                Ok(buf)
100            }
101            FrontendMessage::Terminate => Ok(vec![b'X', 0, 0, 0, 4]),
102            FrontendMessage::SASLInitialResponse { mechanism, data } => {
103                if Self::has_nul(mechanism) {
104                    return Err(FrontendEncodeError::InteriorNul("mechanism"));
105                }
106                if data.len() > i32::MAX as usize {
107                    return Err(FrontendEncodeError::MessageTooLarge(data.len()));
108                }
109                let mut buf = Vec::new();
110                buf.push(b'p');
111
112                let mut content = Vec::new();
113                content.extend_from_slice(mechanism.as_bytes());
114                content.push(0);
115                let data_len = i32::try_from(data.len())
116                    .map_err(|_| FrontendEncodeError::MessageTooLarge(data.len()))?;
117                content.extend_from_slice(&data_len.to_be_bytes());
118                content.extend_from_slice(data);
119
120                let len = Self::content_len_to_wire_len(content.len())?;
121                buf.extend_from_slice(&len.to_be_bytes());
122                buf.extend_from_slice(&content);
123                Ok(buf)
124            }
125            FrontendMessage::SASLResponse(data) | FrontendMessage::GSSResponse(data) => {
126                if data.len() > i32::MAX as usize {
127                    return Err(FrontendEncodeError::MessageTooLarge(data.len()));
128                }
129                let mut buf = Vec::new();
130                buf.push(b'p');
131                let len = Self::content_len_to_wire_len(data.len())?;
132                buf.extend_from_slice(&len.to_be_bytes());
133                buf.extend_from_slice(data);
134                Ok(buf)
135            }
136            FrontendMessage::PasswordMessage(password) => {
137                if Self::has_nul(password) {
138                    return Err(FrontendEncodeError::InteriorNul("password"));
139                }
140                let mut buf = Vec::new();
141                buf.push(b'p');
142                let mut content = Vec::with_capacity(password.len() + 1);
143                content.extend_from_slice(password.as_bytes());
144                content.push(0);
145                let len = Self::content_len_to_wire_len(content.len())?;
146                buf.extend_from_slice(&len.to_be_bytes());
147                buf.extend_from_slice(&content);
148                Ok(buf)
149            }
150            FrontendMessage::Parse {
151                name,
152                query,
153                param_types,
154            } => {
155                if Self::has_nul(name) {
156                    return Err(FrontendEncodeError::InteriorNul("name"));
157                }
158                if Self::has_nul(query) {
159                    return Err(FrontendEncodeError::InteriorNul("query"));
160                }
161                if param_types.len() > i16::MAX as usize {
162                    return Err(FrontendEncodeError::TooManyParams(param_types.len()));
163                }
164                let mut buf = Vec::new();
165                buf.push(b'P');
166
167                let mut content = Vec::new();
168                content.extend_from_slice(name.as_bytes());
169                content.push(0);
170                content.extend_from_slice(query.as_bytes());
171                content.push(0);
172                let param_count = i16::try_from(param_types.len())
173                    .map_err(|_| FrontendEncodeError::TooManyParams(param_types.len()))?;
174                content.extend_from_slice(&param_count.to_be_bytes());
175                for oid in param_types {
176                    content.extend_from_slice(&oid.to_be_bytes());
177                }
178
179                let len = Self::content_len_to_wire_len(content.len())?;
180                buf.extend_from_slice(&len.to_be_bytes());
181                buf.extend_from_slice(&content);
182                Ok(buf)
183            }
184            FrontendMessage::Bind {
185                portal,
186                statement,
187                params,
188            } => {
189                if Self::has_nul(portal) {
190                    return Err(FrontendEncodeError::InteriorNul("portal"));
191                }
192                if Self::has_nul(statement) {
193                    return Err(FrontendEncodeError::InteriorNul("statement"));
194                }
195                if params.len() > i16::MAX as usize {
196                    return Err(FrontendEncodeError::TooManyParams(params.len()));
197                }
198                if let Some(too_large) = params
199                    .iter()
200                    .flatten()
201                    .find(|p| p.len() > i32::MAX as usize)
202                {
203                    return Err(FrontendEncodeError::MessageTooLarge(too_large.len()));
204                }
205
206                let mut buf = Vec::new();
207                buf.push(b'B');
208
209                let mut content = Vec::new();
210                content.extend_from_slice(portal.as_bytes());
211                content.push(0);
212                content.extend_from_slice(statement.as_bytes());
213                content.push(0);
214                content.extend_from_slice(&0i16.to_be_bytes());
215                let param_count = i16::try_from(params.len())
216                    .map_err(|_| FrontendEncodeError::TooManyParams(params.len()))?;
217                content.extend_from_slice(&param_count.to_be_bytes());
218                for param in params {
219                    match param {
220                        Some(data) => {
221                            let data_len = i32::try_from(data.len())
222                                .map_err(|_| FrontendEncodeError::MessageTooLarge(data.len()))?;
223                            content.extend_from_slice(&data_len.to_be_bytes());
224                            content.extend_from_slice(data);
225                        }
226                        None => content.extend_from_slice(&(-1i32).to_be_bytes()),
227                    }
228                }
229                content.extend_from_slice(&0i16.to_be_bytes());
230
231                let len = Self::content_len_to_wire_len(content.len())?;
232                buf.extend_from_slice(&len.to_be_bytes());
233                buf.extend_from_slice(&content);
234                Ok(buf)
235            }
236            FrontendMessage::Execute { portal, max_rows } => {
237                if Self::has_nul(portal) {
238                    return Err(FrontendEncodeError::InteriorNul("portal"));
239                }
240                if *max_rows < 0 {
241                    return Err(FrontendEncodeError::InvalidMaxRows(*max_rows));
242                }
243                let mut buf = Vec::new();
244                buf.push(b'E');
245                let mut content = Vec::new();
246                content.extend_from_slice(portal.as_bytes());
247                content.push(0);
248                content.extend_from_slice(&max_rows.to_be_bytes());
249                let len = Self::content_len_to_wire_len(content.len())?;
250                buf.extend_from_slice(&len.to_be_bytes());
251                buf.extend_from_slice(&content);
252                Ok(buf)
253            }
254            FrontendMessage::Sync => Ok(vec![b'S', 0, 0, 0, 4]),
255            FrontendMessage::CopyFail(msg) => {
256                if Self::has_nul(msg) {
257                    return Err(FrontendEncodeError::InteriorNul("copy_fail"));
258                }
259                let mut buf = Vec::new();
260                buf.push(b'f');
261                let mut content = Vec::with_capacity(msg.len() + 1);
262                content.extend_from_slice(msg.as_bytes());
263                content.push(0);
264                let len = Self::content_len_to_wire_len(content.len())?;
265                buf.extend_from_slice(&len.to_be_bytes());
266                buf.extend_from_slice(&content);
267                Ok(buf)
268            }
269            FrontendMessage::Close { is_portal, name } => {
270                if Self::has_nul(name) {
271                    return Err(FrontendEncodeError::InteriorNul("name"));
272                }
273                let mut buf = Vec::new();
274                buf.push(b'C');
275                let type_byte = if *is_portal { b'P' } else { b'S' };
276                let mut content = vec![type_byte];
277                content.extend_from_slice(name.as_bytes());
278                content.push(0);
279                let len = Self::content_len_to_wire_len(content.len())?;
280                buf.extend_from_slice(&len.to_be_bytes());
281                buf.extend_from_slice(&content);
282                Ok(buf)
283            }
284        }
285    }
286}