1use uqa_core::{
10 json::{decode_json_string_with_control, JsonReadError},
11 memory::{Produced, ProductionControl, ProductionString, ProductionVec},
12 Value,
13};
14
15use crate::error::{Result, SQLError};
16
17use super::validate_named_argument_order_with_control;
18
19const PARAMETER_NAMES: [&str; 2] = ["target", "strip_in_arrays"];
20
21pub fn argument_positions(
23 name: &str,
24 argument_names: &[Option<&str>],
25) -> Result<Option<Vec<usize>>> {
26 argument_positions_with_control(name, argument_names, &ProductionControl::uncontrolled()).map(
27 |positions| {
28 positions.map(|positions| {
29 positions
30 .into_uncontrolled()
31 .expect("ordinary JSON null-stripping argument positions")
32 })
33 },
34 )
35}
36
37pub fn argument_positions_with_control(
38 name: &str,
39 argument_names: &[Option<&str>],
40 control: &ProductionControl<'_>,
41) -> Result<Option<Produced<Vec<usize>>>> {
42 control.check()?;
43 validate_named_argument_order_with_control(argument_names.iter().copied(), control)?;
44 let function = name
45 .get(..11)
46 .filter(|prefix| prefix.eq_ignore_ascii_case("pg_catalog."))
47 .map_or(name, |_| &name[11..]);
48 if !(function.eq_ignore_ascii_case("json_strip_nulls")
49 || function.eq_ignore_ascii_case("jsonb_strip_nulls"))
50 || !(1..=2).contains(&argument_names.len())
51 {
52 return Ok(None);
53 }
54 let mut occupied = [false; PARAMETER_NAMES.len()];
55 let mut positions = ProductionVec::new(*control);
56 positions.reserve(argument_names.len())?;
57 let mut positional = 0usize;
58 for argument_name in argument_names {
59 let position = if let Some(argument_name) = argument_name {
60 PARAMETER_NAMES
61 .iter()
62 .position(|candidate| candidate == argument_name)
63 } else {
64 let position = positional;
65 positional += 1;
66 Some(position)
67 };
68 let Some(position) = position.filter(|position| *position < occupied.len()) else {
69 return Ok(None);
70 };
71 if occupied[position] {
72 return Ok(None);
73 }
74 occupied[position] = true;
75 positions.push_copy(position)?;
76 }
77 Ok(occupied[0].then(|| positions.finish()).transpose()?)
78}
79
80pub(super) fn reorder_named_values_with_control(
81 function: &str,
82 call_args: &[(Option<String>, Value)],
83 control: &ProductionControl<'_>,
84) -> Result<Option<Produced<Vec<Value>>>> {
85 let names = super::call_arguments::evaluated_argument_names_with_control(call_args, control)?;
86 let positions = match argument_positions_with_control(function, &names, control) {
87 Ok(Some(positions)) => positions,
88 Err(error) if matches!(error.sqlstate(), Some("53200" | "57014")) => return Err(error),
89 Ok(None) | Err(_) => return Ok(None),
90 };
91 let mut values = [None; PARAMETER_NAMES.len()];
92 for ((_, value), position) in call_args.iter().zip(positions.iter().copied()) {
93 values[position] = Some(value);
94 }
95 let default = Value::Bool(false);
96 values[1].get_or_insert(&default);
97 let mut output = ProductionVec::new(*control);
98 output.reserve(values.len())?;
99 for value in values {
100 let Some(value) = value else { return Ok(None) };
101 output.push_produced(control.copy_value(value)?)?;
102 }
103 Ok(Some(output.finish()?))
104}
105
106pub(super) fn invalid_json_input(input: &str) -> SQLError {
107 SQLError::Routine {
108 sqlstate: "22P02".into(),
109 message: format!("invalid input syntax for type json: \"{input}\""),
110 }
111}
112
113#[cfg(test)]
115pub(super) fn strip_json_nulls_text(input: &str, strip_in_arrays: bool) -> Result<String> {
116 Ok(strip_json_nulls_text_with_control(
117 input,
118 strip_in_arrays,
119 &ProductionControl::uncontrolled(),
120 )?
121 .into_uncontrolled()
122 .expect("ordinary JSON stripping has no lease"))
123}
124
125pub(super) fn strip_json_nulls_text_with_control(
126 input: &str,
127 strip_in_arrays: bool,
128 control: &ProductionControl<'_>,
129) -> Result<Produced<String>> {
130 let mut parser = JsonStripParser {
131 input,
132 position: 0,
133 strip_in_arrays,
134 control: *control,
135 };
136 let rendered = parser.parse_value(0)?;
137 parser.skip_whitespace()?;
138 if parser.position != input.len() {
139 return Err(invalid_json_input(input));
140 }
141 Ok(rendered.text)
142}
143
144struct RenderedJson {
145 text: Produced<String>,
146 is_null: bool,
147}
148
149struct JsonStripParser<'a, 'c> {
150 input: &'a str,
151 position: usize,
152 strip_in_arrays: bool,
153 control: ProductionControl<'c>,
154}
155
156impl JsonStripParser<'_, '_> {
157 const MAX_DEPTH: usize = 128;
158
159 fn parse_value(&mut self, depth: usize) -> Result<RenderedJson> {
160 self.control.check()?;
161 if depth > Self::MAX_DEPTH {
162 return Err(invalid_json_input(self.input));
163 }
164 self.skip_whitespace()?;
165 match self.peek() {
166 Some(b'{') => self.parse_object(depth),
167 Some(b'[') => self.parse_array(depth),
168 Some(b'"') => self.parse_string().map(|text| RenderedJson {
169 text,
170 is_null: false,
171 }),
172 Some(b't') => self.parse_literal("true", false),
173 Some(b'f') => self.parse_literal("false", false),
174 Some(b'n') => self.parse_literal("null", true),
175 Some(b'-' | b'0'..=b'9') => self.parse_number(),
176 _ => Err(invalid_json_input(self.input)),
177 }
178 }
179
180 fn parse_object(&mut self, depth: usize) -> Result<RenderedJson> {
181 self.position += 1;
182 self.skip_whitespace()?;
183 let mut fields = ProductionString::new(self.control);
184 fields.push('{')?;
185 let mut emitted = false;
186 if self.consume(b'}') {
187 fields.push('}')?;
188 return Ok(RenderedJson {
189 text: fields.finish()?,
190 is_null: false,
191 });
192 }
193 loop {
194 self.skip_whitespace()?;
195 if self.peek() != Some(b'"') {
196 return Err(invalid_json_input(self.input));
197 }
198 let key = self.parse_string()?;
199 self.skip_whitespace()?;
200 if !self.consume(b':') {
201 return Err(invalid_json_input(self.input));
202 }
203 let value = self.parse_value(depth + 1)?;
204 if !value.is_null {
205 if emitted {
206 fields.push(',')?;
207 }
208 fields.push_str(&key)?;
209 fields.push(':')?;
210 fields.push_str(&value.text)?;
211 emitted = true;
212 }
213 self.skip_whitespace()?;
214 if self.consume(b'}') {
215 break;
216 }
217 if !self.consume(b',') {
218 return Err(invalid_json_input(self.input));
219 }
220 }
221 fields.push('}')?;
222 Ok(RenderedJson {
223 text: fields.finish()?,
224 is_null: false,
225 })
226 }
227
228 fn parse_array(&mut self, depth: usize) -> Result<RenderedJson> {
229 self.position += 1;
230 self.skip_whitespace()?;
231 let mut elements = ProductionString::new(self.control);
232 elements.push('[')?;
233 let mut emitted = false;
234 if self.consume(b']') {
235 elements.push(']')?;
236 return Ok(RenderedJson {
237 text: elements.finish()?,
238 is_null: false,
239 });
240 }
241 loop {
242 let value = self.parse_value(depth + 1)?;
243 if !self.strip_in_arrays || !value.is_null {
244 if emitted {
245 elements.push(',')?;
246 }
247 elements.push_str(&value.text)?;
248 emitted = true;
249 }
250 self.skip_whitespace()?;
251 if self.consume(b']') {
252 break;
253 }
254 if !self.consume(b',') {
255 return Err(invalid_json_input(self.input));
256 }
257 }
258 elements.push(']')?;
259 Ok(RenderedJson {
260 text: elements.finish()?,
261 is_null: false,
262 })
263 }
264
265 fn parse_string(&mut self) -> Result<Produced<String>> {
266 let start = self.position;
267 self.position += 1;
268 while let Some(byte) = self.peek() {
269 self.control.check()?;
270 match byte {
271 b'"' => {
272 self.position += 1;
273 let source = &self.input[start..self.position];
274 let decoded = decode_json_string_with_control(source.as_bytes(), &self.control)
275 .map_err(|error| match error {
276 JsonReadError::InvalidJson | JsonReadError::Malformed { .. } => {
277 invalid_json_input(self.input)
278 }
279 JsonReadError::Memory(error) => error.into(),
280 JsonReadError::Cancelled(error) => error.into(),
281 })?;
282 return super::json::quote_with_control(&decoded, &self.control);
283 }
284 b'\\' => {
285 self.position += 1;
286 if self.peek().is_none() {
287 return Err(invalid_json_input(self.input));
288 }
289 self.position += 1;
290 }
291 _ => self.position += 1,
292 }
293 }
294 Err(invalid_json_input(self.input))
295 }
296
297 fn parse_literal(&mut self, literal: &str, is_null: bool) -> Result<RenderedJson> {
298 if !self.input[self.position..].starts_with(literal) {
299 return Err(invalid_json_input(self.input));
300 }
301 self.position += literal.len();
302 Ok(RenderedJson {
303 text: self.control.copy_text(literal)?,
304 is_null,
305 })
306 }
307
308 fn parse_number(&mut self) -> Result<RenderedJson> {
309 let start = self.position;
310 self.consume(b'-');
311 match self.peek() {
312 Some(b'0') => self.position += 1,
313 Some(b'1'..=b'9') => {
314 self.position += 1;
315 self.consume_digits()?;
316 }
317 _ => return Err(invalid_json_input(self.input)),
318 }
319 if self.consume(b'.') {
320 let digits = self.position;
321 self.consume_digits()?;
322 if digits == self.position {
323 return Err(invalid_json_input(self.input));
324 }
325 }
326 if matches!(self.peek(), Some(b'e' | b'E')) {
327 self.position += 1;
328 if matches!(self.peek(), Some(b'+' | b'-')) {
329 self.position += 1;
330 }
331 let digits = self.position;
332 self.consume_digits()?;
333 if digits == self.position {
334 return Err(invalid_json_input(self.input));
335 }
336 }
337 Ok(RenderedJson {
338 text: self.control.copy_text(&self.input[start..self.position])?,
339 is_null: false,
340 })
341 }
342
343 fn consume_digits(&mut self) -> Result<()> {
344 while matches!(self.peek(), Some(b'0'..=b'9')) {
345 self.control.check()?;
346 self.position += 1;
347 }
348 Ok(())
349 }
350
351 fn skip_whitespace(&mut self) -> Result<()> {
352 while matches!(self.peek(), Some(b' ' | b'\n' | b'\r' | b'\t')) {
353 self.control.check()?;
354 self.position += 1;
355 }
356 Ok(())
357 }
358
359 fn consume(&mut self, expected: u8) -> bool {
360 if self.peek() == Some(expected) {
361 self.position += 1;
362 true
363 } else {
364 false
365 }
366 }
367
368 fn peek(&self) -> Option<u8> {
369 self.input.as_bytes().get(self.position).copied()
370 }
371}
372
373#[cfg(test)]
374mod tests {
375 use super::{argument_positions, strip_json_nulls_text};
376
377 #[test]
378 fn json_strip_positions_accept_the_default_and_declaration_order_names() {
379 assert_eq!(
380 argument_positions("json_strip_nulls", &[None]).unwrap(),
381 Some(vec![0])
382 );
383 assert_eq!(
384 argument_positions(
385 "jsonb_strip_nulls",
386 &[Some("strip_in_arrays"), Some("target")]
387 )
388 .unwrap(),
389 Some(vec![1, 0])
390 );
391 assert_eq!(
392 argument_positions("json_strip_nulls", &[Some("strip_in_arrays")]).unwrap(),
393 None
394 );
395 assert_eq!(
396 argument_positions("json_strip_nulls", &[Some("unknown"), Some("target")]).unwrap(),
397 None
398 );
399 }
400
401 #[test]
402 fn textual_json_null_stripping_preserves_order_duplicates_and_number_lexemes() {
403 let input = r#" { "z" : 1.2300e+02, "a" : null, "z" : 2, "s" : "\u0061", "nested" : [null,{"drop":null,"keep":3}] } "#;
404 assert_eq!(
405 strip_json_nulls_text(input, false).unwrap(),
406 r#"{"z":1.2300e+02,"z":2,"s":"a","nested":[null,{"keep":3}]}"#
407 );
408 assert_eq!(
409 strip_json_nulls_text(input, true).unwrap(),
410 r#"{"z":1.2300e+02,"z":2,"s":"a","nested":[{"keep":3}]}"#
411 );
412 assert_eq!(strip_json_nulls_text("null", true).unwrap(), "null");
413 }
414
415 #[test]
416 fn textual_json_null_stripping_rejects_malformed_input_with_json_sqlstate() {
417 for input in [r#"{"a":}"#, r#"{"a":01}"#, r"[1,]", r#""\uD800""#] {
418 assert_eq!(
419 strip_json_nulls_text(input, false).unwrap_err().sqlstate(),
420 Some("22P02")
421 );
422 }
423 }
424}