Skip to main content

mago_analyzer/plugin/libraries/stdlib/string/
sprintf.rs

1//! `sprintf()` / format-string return type provider.
2//!
3//! When all arguments are known literals, resolves the exact result string.
4//! When the format string is known but arguments are not all literals,
5//! infers `non-empty-string` or `truthy-string` when the output is guaranteed
6//! to be non-empty or longer than one character.
7//!
8//! The core logic is exposed via [`resolve_sprintf`] so it can be reused by
9//! other providers (e.g. `Psl\Str\format`).
10
11use std::fmt::Write;
12
13use mago_codex::ttype::atomic::TAtomic;
14use mago_codex::ttype::atomic::scalar::TScalar;
15use mago_codex::ttype::get_literal_string;
16use mago_codex::ttype::get_non_empty_string;
17use mago_codex::ttype::get_truthy_string;
18use mago_codex::ttype::union::TUnion;
19use mago_word::word;
20
21use crate::plugin::context::InvocationInfo;
22use crate::plugin::context::ProviderContext;
23use crate::plugin::provider::Provider;
24use crate::plugin::provider::ProviderMeta;
25use crate::plugin::provider::function::FunctionReturnTypeProvider;
26use crate::plugin::provider::function::FunctionTarget;
27
28static META: ProviderMeta =
29    ProviderMeta::new("php::string::sprintf", "sprintf", "Resolves literal string for sprintf with literal args");
30
31#[derive(Default)]
32pub struct SprintfProvider;
33
34impl Provider for SprintfProvider {
35    fn meta() -> &'static ProviderMeta {
36        &META
37    }
38}
39
40impl FunctionReturnTypeProvider for SprintfProvider {
41    fn targets() -> FunctionTarget {
42        FunctionTarget::Exact(b"sprintf")
43    }
44
45    fn get_return_type(
46        &self,
47        context: &ProviderContext<'_, '_, '_>,
48        invocation: &InvocationInfo<'_, '_, '_>,
49    ) -> Option<TUnion> {
50        resolve_sprintf(context, invocation)
51    }
52}
53
54/// Resolve the return type of a sprintf-like call.
55///
56/// Expects the first argument to be the format string and subsequent arguments
57/// to be the format values (standard `sprintf` / `Psl\Str\format` signature).
58pub fn resolve_sprintf(
59    context: &ProviderContext<'_, '_, '_>,
60    invocation: &InvocationInfo<'_, '_, '_>,
61) -> Option<TUnion> {
62    let format_argument = invocation.get_argument(0, &[b"format"])?;
63    let format_type = context.get_expression_type(format_argument)?;
64    let format_str = format_type.get_single_literal_string_value()?;
65
66    if let Some(result) = resolve_literal(format_str, context, invocation) {
67        return Some(get_literal_string(word(&result)));
68    }
69
70    let min_len = analyze_min_length(format_str, context, invocation);
71    if min_len >= 2 {
72        Some(get_truthy_string())
73    } else if min_len >= 1 {
74        Some(get_non_empty_string())
75    } else {
76        None
77    }
78}
79
80fn argument_string_min_length(
81    context: &ProviderContext<'_, '_, '_>,
82    invocation: &InvocationInfo<'_, '_, '_>,
83    arg_index: usize,
84) -> usize {
85    let Some(arg) = invocation.get_argument(arg_index, &[]) else {
86        return 0;
87    };
88
89    let Some(arg_type) = context.get_expression_type(arg) else {
90        return 0;
91    };
92
93    if let Some(literal) = arg_type.get_single_literal_string_value() {
94        return literal.len();
95    }
96
97    let mut min_len = usize::MAX;
98    for atomic in arg_type.types.as_ref() {
99        let atomic_min = match atomic {
100            TAtomic::Scalar(TScalar::String(string)) if string.is_non_empty || string.is_numeric => 1,
101            _ => 0,
102        };
103
104        if atomic_min < min_len {
105            min_len = atomic_min;
106        }
107
108        if min_len == 0 {
109            return 0;
110        }
111    }
112
113    if min_len == usize::MAX { 0 } else { min_len }
114}
115
116/// Parse flags from a format specifier. Advances `i` past all flag characters.
117/// Returns `None` if the format string is malformed (unexpected end).
118fn parse_flags(bytes: &[u8], i: &mut usize) -> Option<(char, bool, bool)> {
119    let len = bytes.len();
120    let mut pad_char = ' ';
121    let mut left_align = false;
122    let mut show_sign = false;
123
124    loop {
125        if *i >= len {
126            return None;
127        }
128        match bytes[*i] {
129            b'-' => {
130                left_align = true;
131                *i += 1;
132            }
133            b'+' => {
134                show_sign = true;
135                *i += 1;
136            }
137            b' ' => *i += 1,
138            b'0' => {
139                pad_char = '0';
140                *i += 1;
141            }
142            b'\'' => {
143                *i += 1;
144                if *i >= len {
145                    return None;
146                }
147
148                pad_char = bytes[*i] as char;
149                *i += 1;
150            }
151            _ => break,
152        }
153    }
154
155    Some((pad_char, left_align, show_sign))
156}
157
158/// Parse a decimal number from `bytes` starting at `i`. Advances `i` past all digits.
159fn parse_number(bytes: &[u8], i: &mut usize) -> usize {
160    let len = bytes.len();
161    let mut n: usize = 0;
162    while *i < len && bytes[*i].is_ascii_digit() {
163        n = n * 10 + (bytes[*i] - b'0') as usize;
164        *i += 1;
165    }
166
167    n
168}
169
170/// Parse optional precision (`.N`). Advances `i` past the precision if present.
171fn parse_precision(bytes: &[u8], i: &mut usize) -> Option<usize> {
172    if *i < bytes.len() && bytes[*i] == b'.' {
173        *i += 1;
174        Some(parse_number(bytes, i))
175    } else {
176        None
177    }
178}
179
180/// Try to fully resolve sprintf to a literal string when all arguments are known literals.
181fn resolve_literal(
182    format_str: &[u8],
183    context: &ProviderContext<'_, '_, '_>,
184    invocation: &InvocationInfo<'_, '_, '_>,
185) -> Option<String> {
186    let format_str_utf8 = std::str::from_utf8(format_str).ok()?;
187    let mut result = String::with_capacity(format_str.len());
188    let mut buf = String::new();
189    let bytes = format_str;
190    let len = bytes.len();
191    let mut i = 0;
192    let mut arg_index: usize = 1;
193
194    while i < len {
195        if bytes[i] != b'%' {
196            let start = i;
197            i += 1;
198            while i < len && bytes[i] != b'%' {
199                i += 1;
200            }
201
202            result.push_str(&format_str_utf8[start..i]);
203            continue;
204        }
205
206        i += 1;
207        if i >= len {
208            return None;
209        }
210
211        if bytes[i] == b'%' {
212            result.push('%');
213            i += 1;
214            continue;
215        }
216
217        let (pad_char, left_align, show_sign) = parse_flags(bytes, &mut i)?;
218        let width = parse_number(bytes, &mut i);
219        let precision = parse_precision(bytes, &mut i);
220
221        if i >= len {
222            return None;
223        }
224
225        let specifier = bytes[i];
226        let arg = invocation.get_argument(arg_index, &[])?;
227        let arg_type = context.get_expression_type(arg)?;
228
229        i += 1;
230        arg_index += 1;
231
232        let needs_buf = width > 0 || specifier == b'e' || specifier == b'E';
233        let target = if needs_buf {
234            buf.clear();
235            &mut buf
236        } else {
237            &mut result
238        };
239
240        match specifier {
241            b's' => {
242                let value = arg_type.get_single_literal_string_value()?;
243                let value_str = std::str::from_utf8(value).ok()?;
244                if let Some(prec) = precision {
245                    target.push_str(&value_str[..value_str.len().min(prec)]);
246                } else {
247                    target.push_str(value_str);
248                }
249            }
250            b'd' => {
251                let value = arg_type.get_single_literal_int_value()?;
252                if show_sign && value >= 0 {
253                    target.push('+');
254                }
255
256                let _ = write!(target, "{value}");
257            }
258            b'u' => {
259                let value = arg_type.get_single_literal_int_value()?;
260                let _ = write!(target, "{}", value as u64);
261            }
262            b'f' | b'F' => {
263                let value = get_float_value(arg_type)?;
264                let prec = precision.unwrap_or(6);
265                if show_sign && value >= 0.0 {
266                    target.push('+');
267                }
268
269                let _ = write!(target, "{value:.prec$}");
270            }
271            b'e' | b'E' => {
272                let value = get_float_value(arg_type)?;
273                let prec = precision.unwrap_or(6);
274                if show_sign && value >= 0.0 {
275                    target.push('+');
276                }
277
278                let mark = target.len();
279                if specifier == b'e' {
280                    let _ = write!(target, "{value:.prec$e}");
281                } else {
282                    let _ = write!(target, "{value:.prec$E}");
283                }
284
285                // Rust writes e.g. `1e0`, PHP writes `1e+0`. Insert `+` if needed.
286                normalize_scientific_in_place(target, mark);
287            }
288            b'x' => {
289                let value = arg_type.get_single_literal_int_value()?;
290                let _ = write!(target, "{:x}", value as u64);
291            }
292            b'X' => {
293                let value = arg_type.get_single_literal_int_value()?;
294                let _ = write!(target, "{:X}", value as u64);
295            }
296            b'o' => {
297                let value = arg_type.get_single_literal_int_value()?;
298                let _ = write!(target, "{:o}", value as u64);
299            }
300            b'b' => {
301                let value = arg_type.get_single_literal_int_value()?;
302                let _ = write!(target, "{:b}", value as u64);
303            }
304            b'c' => {
305                let value = arg_type.get_single_literal_int_value()?;
306                target.push(char::from_u32(value as u32)?);
307            }
308            _ => return None,
309        }
310
311        if needs_buf {
312            if width > 0 && buf.len() < width {
313                let padding = width - buf.len();
314                if left_align {
315                    result.push_str(&buf);
316                    for _ in 0..padding {
317                        result.push(' ');
318                    }
319                } else {
320                    for _ in 0..padding {
321                        result.push(pad_char);
322                    }
323                    result.push_str(&buf);
324                }
325            } else {
326                result.push_str(&buf);
327            }
328        }
329    }
330
331    Some(result)
332}
333
334/// Extract a float value from a type union, accepting either a literal float or literal int.
335fn get_float_value(t: &TUnion) -> Option<f64> {
336    if let Some(v) = t.get_single_literal_float_value() {
337        Some(v)
338    } else {
339        t.get_single_literal_int_value().map(|v| v as f64)
340    }
341}
342
343/// Insert a `+` sign after `e`/`E` in scientific notation if Rust omitted it.
344/// Only scans bytes from `start` onward.
345fn normalize_scientific_in_place(s: &mut String, start: usize) {
346    let bytes = s.as_bytes();
347    for j in start..bytes.len() {
348        if bytes[j] == b'e' || bytes[j] == b'E' {
349            if j + 1 < bytes.len() && bytes[j + 1] != b'+' && bytes[j + 1] != b'-' {
350                s.insert(j + 1, '+');
351            }
352            return;
353        }
354    }
355}
356
357fn analyze_min_length(
358    format_str: &[u8],
359    context: &ProviderContext<'_, '_, '_>,
360    invocation: &InvocationInfo<'_, '_, '_>,
361) -> usize {
362    let bytes = format_str;
363    let len = bytes.len();
364    let mut i = 0;
365    let mut min_len: usize = 0;
366    let mut arg_index: usize = 1;
367
368    while i < len {
369        if bytes[i] != b'%' {
370            let start = i;
371            i += 1;
372            while i < len && bytes[i] != b'%' {
373                i += 1;
374            }
375
376            min_len += i - start;
377            continue;
378        }
379
380        i += 1;
381        if i >= len {
382            return min_len;
383        }
384
385        if bytes[i] == b'%' {
386            min_len += 1;
387            i += 1;
388            continue;
389        }
390
391        // Skip flags.
392        loop {
393            if i >= len {
394                return min_len;
395            }
396
397            match bytes[i] {
398                b'-' | b'+' | b' ' | b'0' => i += 1,
399                b'\'' => {
400                    i += 2;
401                    if i > len {
402                        return min_len;
403                    }
404                }
405                _ => break,
406            }
407        }
408
409        let width = parse_number(bytes, &mut i);
410        let precision = parse_precision(bytes, &mut i);
411
412        if i >= len {
413            return min_len;
414        }
415
416        let specifier = bytes[i];
417        i += 1;
418
419        let specifier_min = match specifier {
420            b's' => {
421                let mut from_arg = argument_string_min_length(context, invocation, arg_index);
422                if let Some(prec) = precision {
423                    from_arg = from_arg.min(prec);
424                }
425
426                from_arg
427            }
428            b'd' | b'u' | b'f' | b'F' | b'e' | b'E' | b'x' | b'X' | b'o' | b'b' | b'c' => 1,
429            _ => 0,
430        };
431
432        arg_index += 1;
433        min_len += specifier_min.max(width);
434    }
435
436    min_len
437}