qail_pg/protocol/wire/
frontend.rs1use 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 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(¶m_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(¶m_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}