1use uqa_core::Value;
10
11use crate::error::{Result, SQLError};
12
13use super::validate_named_argument_order;
14
15const PARAMETER_NAMES: [&str; 2] = ["target", "strip_in_arrays"];
16
17pub fn argument_positions(
19 name: &str,
20 argument_names: &[Option<&str>],
21) -> Result<Option<Vec<usize>>> {
22 validate_named_argument_order(argument_names.iter().copied())?;
23 let lower = name.to_ascii_lowercase();
24 let function = lower.strip_prefix("pg_catalog.").unwrap_or(&lower);
25 if !matches!(function, "json_strip_nulls" | "jsonb_strip_nulls")
26 || !(1..=2).contains(&argument_names.len())
27 {
28 return Ok(None);
29 }
30 let mut occupied = [false; PARAMETER_NAMES.len()];
31 let mut positions = Vec::with_capacity(argument_names.len());
32 let mut positional = 0usize;
33 for argument_name in argument_names {
34 let position = if let Some(argument_name) = argument_name {
35 PARAMETER_NAMES
36 .iter()
37 .position(|candidate| candidate == argument_name)
38 } else {
39 let position = positional;
40 positional += 1;
41 Some(position)
42 };
43 let Some(position) = position.filter(|position| *position < occupied.len()) else {
44 return Ok(None);
45 };
46 if occupied[position] {
47 return Ok(None);
48 }
49 occupied[position] = true;
50 positions.push(position);
51 }
52 Ok(occupied[0].then_some(positions))
53}
54
55pub(super) fn reorder_named_values(
56 function: &str,
57 call_args: &[(Option<String>, Value)],
58) -> Option<Vec<Value>> {
59 let argument_names = call_args
60 .iter()
61 .map(|(name, _)| name.as_deref())
62 .collect::<Vec<_>>();
63 let positions = argument_positions(function, &argument_names)
64 .ok()
65 .flatten()?;
66 let mut values = vec![None; PARAMETER_NAMES.len()];
67 for ((_, value), position) in call_args.iter().zip(positions) {
68 values[position] = Some(value.clone());
69 }
70 values[1].get_or_insert(Value::Bool(false));
71 values.into_iter().collect()
72}
73
74pub(super) fn invalid_json_input(input: &str) -> SQLError {
75 SQLError::Routine {
76 sqlstate: "22P02".into(),
77 message: format!("invalid input syntax for type json: \"{input}\""),
78 }
79}
80
81pub(super) fn strip_json_nulls_text(input: &str, strip_in_arrays: bool) -> Result<String> {
83 let mut parser = JsonStripParser {
84 input,
85 position: 0,
86 strip_in_arrays,
87 };
88 let rendered = parser.parse_value(0)?;
89 parser.skip_whitespace();
90 if parser.position != input.len() {
91 return Err(invalid_json_input(input));
92 }
93 Ok(rendered.text)
94}
95
96struct RenderedJson {
97 text: String,
98 is_null: bool,
99}
100
101struct JsonStripParser<'a> {
102 input: &'a str,
103 position: usize,
104 strip_in_arrays: bool,
105}
106
107impl JsonStripParser<'_> {
108 const MAX_DEPTH: usize = 128;
109
110 fn parse_value(&mut self, depth: usize) -> Result<RenderedJson> {
111 if depth > Self::MAX_DEPTH {
112 return Err(invalid_json_input(self.input));
113 }
114 self.skip_whitespace();
115 match self.peek() {
116 Some(b'{') => self.parse_object(depth),
117 Some(b'[') => self.parse_array(depth),
118 Some(b'"') => self.parse_string().map(|text| RenderedJson {
119 text,
120 is_null: false,
121 }),
122 Some(b't') => self.parse_literal("true", false),
123 Some(b'f') => self.parse_literal("false", false),
124 Some(b'n') => self.parse_literal("null", true),
125 Some(b'-' | b'0'..=b'9') => self.parse_number(),
126 _ => Err(invalid_json_input(self.input)),
127 }
128 }
129
130 fn parse_object(&mut self, depth: usize) -> Result<RenderedJson> {
131 self.position += 1;
132 self.skip_whitespace();
133 let mut fields = Vec::new();
134 if self.consume(b'}') {
135 return Ok(RenderedJson {
136 text: "{}".into(),
137 is_null: false,
138 });
139 }
140 loop {
141 self.skip_whitespace();
142 if self.peek() != Some(b'"') {
143 return Err(invalid_json_input(self.input));
144 }
145 let key = self.parse_string()?;
146 self.skip_whitespace();
147 if !self.consume(b':') {
148 return Err(invalid_json_input(self.input));
149 }
150 let value = self.parse_value(depth + 1)?;
151 if !value.is_null {
152 fields.push(format!("{key}:{}", value.text));
153 }
154 self.skip_whitespace();
155 if self.consume(b'}') {
156 break;
157 }
158 if !self.consume(b',') {
159 return Err(invalid_json_input(self.input));
160 }
161 }
162 Ok(RenderedJson {
163 text: format!("{{{}}}", fields.join(",")),
164 is_null: false,
165 })
166 }
167
168 fn parse_array(&mut self, depth: usize) -> Result<RenderedJson> {
169 self.position += 1;
170 self.skip_whitespace();
171 let mut elements = Vec::new();
172 if self.consume(b']') {
173 return Ok(RenderedJson {
174 text: "[]".into(),
175 is_null: false,
176 });
177 }
178 loop {
179 let value = self.parse_value(depth + 1)?;
180 if !self.strip_in_arrays || !value.is_null {
181 elements.push(value.text);
182 }
183 self.skip_whitespace();
184 if self.consume(b']') {
185 break;
186 }
187 if !self.consume(b',') {
188 return Err(invalid_json_input(self.input));
189 }
190 }
191 Ok(RenderedJson {
192 text: format!("[{}]", elements.join(",")),
193 is_null: false,
194 })
195 }
196
197 fn parse_string(&mut self) -> Result<String> {
198 let start = self.position;
199 self.position += 1;
200 while let Some(byte) = self.peek() {
201 match byte {
202 b'"' => {
203 self.position += 1;
204 let source = &self.input[start..self.position];
205 let decoded = serde_json::from_str::<String>(source)
206 .map_err(|_| invalid_json_input(self.input))?;
207 return serde_json::to_string(&decoded)
208 .map_err(|_| invalid_json_input(self.input));
209 }
210 b'\\' => {
211 self.position += 1;
212 if self.peek().is_none() {
213 return Err(invalid_json_input(self.input));
214 }
215 self.position += 1;
216 }
217 _ => self.position += 1,
218 }
219 }
220 Err(invalid_json_input(self.input))
221 }
222
223 fn parse_literal(&mut self, literal: &str, is_null: bool) -> Result<RenderedJson> {
224 if !self.input[self.position..].starts_with(literal) {
225 return Err(invalid_json_input(self.input));
226 }
227 self.position += literal.len();
228 Ok(RenderedJson {
229 text: literal.into(),
230 is_null,
231 })
232 }
233
234 fn parse_number(&mut self) -> Result<RenderedJson> {
235 let start = self.position;
236 self.consume(b'-');
237 match self.peek() {
238 Some(b'0') => self.position += 1,
239 Some(b'1'..=b'9') => {
240 self.position += 1;
241 self.consume_digits();
242 }
243 _ => return Err(invalid_json_input(self.input)),
244 }
245 if self.consume(b'.') {
246 let digits = self.position;
247 self.consume_digits();
248 if digits == self.position {
249 return Err(invalid_json_input(self.input));
250 }
251 }
252 if matches!(self.peek(), Some(b'e' | b'E')) {
253 self.position += 1;
254 if matches!(self.peek(), Some(b'+' | b'-')) {
255 self.position += 1;
256 }
257 let digits = self.position;
258 self.consume_digits();
259 if digits == self.position {
260 return Err(invalid_json_input(self.input));
261 }
262 }
263 Ok(RenderedJson {
264 text: self.input[start..self.position].into(),
265 is_null: false,
266 })
267 }
268
269 fn consume_digits(&mut self) {
270 while matches!(self.peek(), Some(b'0'..=b'9')) {
271 self.position += 1;
272 }
273 }
274
275 fn skip_whitespace(&mut self) {
276 while matches!(self.peek(), Some(b' ' | b'\n' | b'\r' | b'\t')) {
277 self.position += 1;
278 }
279 }
280
281 fn consume(&mut self, expected: u8) -> bool {
282 if self.peek() == Some(expected) {
283 self.position += 1;
284 true
285 } else {
286 false
287 }
288 }
289
290 fn peek(&self) -> Option<u8> {
291 self.input.as_bytes().get(self.position).copied()
292 }
293}
294
295#[cfg(test)]
296mod tests {
297 use super::{argument_positions, strip_json_nulls_text};
298
299 #[test]
300 fn json_strip_positions_accept_the_default_and_declaration_order_names() {
301 assert_eq!(
302 argument_positions("json_strip_nulls", &[None]).unwrap(),
303 Some(vec![0])
304 );
305 assert_eq!(
306 argument_positions(
307 "jsonb_strip_nulls",
308 &[Some("strip_in_arrays"), Some("target")]
309 )
310 .unwrap(),
311 Some(vec![1, 0])
312 );
313 assert_eq!(
314 argument_positions("json_strip_nulls", &[Some("strip_in_arrays")]).unwrap(),
315 None
316 );
317 assert_eq!(
318 argument_positions("json_strip_nulls", &[Some("unknown"), Some("target")]).unwrap(),
319 None
320 );
321 }
322
323 #[test]
324 fn textual_json_null_stripping_preserves_order_duplicates_and_number_lexemes() {
325 let input = r#" { "z" : 1.2300e+02, "a" : null, "z" : 2, "s" : "\u0061", "nested" : [null,{"drop":null,"keep":3}] } "#;
326 assert_eq!(
327 strip_json_nulls_text(input, false).unwrap(),
328 r#"{"z":1.2300e+02,"z":2,"s":"a","nested":[null,{"keep":3}]}"#
329 );
330 assert_eq!(
331 strip_json_nulls_text(input, true).unwrap(),
332 r#"{"z":1.2300e+02,"z":2,"s":"a","nested":[{"keep":3}]}"#
333 );
334 assert_eq!(strip_json_nulls_text("null", true).unwrap(), "null");
335 }
336
337 #[test]
338 fn textual_json_null_stripping_rejects_malformed_input_with_json_sqlstate() {
339 for input in [r#"{"a":}"#, r#"{"a":01}"#, r"[1,]", r#""\uD800""#] {
340 assert_eq!(
341 strip_json_nulls_text(input, false).unwrap_err().sqlstate(),
342 Some("22P02")
343 );
344 }
345 }
346}