Skip to main content

safe_migrate/analysis/
settings.rs

1/// PostgreSQL stores `lock_timeout` and `statement_timeout` as signed 32-bit
2/// millisecond GUCs. Values above this limit are rejected by PostgreSQL.
3pub const MAX_TIMEOUT_MS: u64 = i32::MAX as u64;
4
5#[derive(Clone, Debug, PartialEq, Eq)]
6pub struct ScopedSetting<T> {
7    pub default: T,
8    pub session: T,
9    pub effective: T,
10}
11
12impl<T: Clone> ScopedSetting<T> {
13    pub fn new(default: T) -> Self {
14        Self {
15            session: default.clone(),
16            effective: default.clone(),
17            default,
18        }
19    }
20
21    pub fn reset_effective_to_session(&mut self) {
22        self.effective = self.session.clone();
23    }
24}
25
26/// Parse PostgreSQL's documented timeout syntax and normalize it to the
27/// integer millisecond representation used by its timeout GUCs.
28pub fn parse_timeout_ms(raw: &str) -> Result<u64, String> {
29    let value = raw.trim();
30    if value.is_empty() {
31        return Err("timeout value is empty".to_string());
32    }
33
34    let bytes = value.as_bytes();
35    let (mut number_end, sign) = match bytes.first() {
36        Some(b'+') => (1, 1.0),
37        Some(b'-') => (1, -1.0),
38        _ => (0, 1.0),
39    };
40    let has_explicit_sign = number_end == 1;
41    let unsigned_start = number_end;
42    let hexadecimal = bytes
43        .get(number_end..number_end + 2)
44        .is_some_and(|prefix| prefix == b"0x" || prefix == b"0X");
45    if hexadecimal {
46        number_end += 2;
47        let digits_start = number_end;
48        while bytes.get(number_end).is_some_and(u8::is_ascii_hexdigit) {
49            number_end += 1;
50        }
51        if number_end == digits_start {
52            return Err(format!("invalid timeout value '{raw}'"));
53        }
54    } else {
55        let mut integer_digits = 0;
56        while bytes.get(number_end).is_some_and(u8::is_ascii_digit) {
57            number_end += 1;
58            integer_digits += 1;
59        }
60        let mut mantissa_digits = integer_digits;
61        if bytes.get(number_end) == Some(&b'.') {
62            number_end += 1;
63            while bytes.get(number_end).is_some_and(u8::is_ascii_digit) {
64                number_end += 1;
65                mantissa_digits += 1;
66            }
67        }
68        if mantissa_digits == 0 {
69            return Err(format!("invalid timeout value '{raw}'"));
70        }
71        // PostgreSQL accepts `.5s`, but its signed-number path requires a
72        // digit before the decimal point.
73        if has_explicit_sign && integer_digits == 0 {
74            return Err(format!("invalid timeout value '{raw}'"));
75        }
76        if matches!(bytes.get(number_end), Some(b'e' | b'E')) {
77            number_end += 1;
78            if matches!(bytes.get(number_end), Some(b'+' | b'-')) {
79                number_end += 1;
80            }
81            let exponent_start = number_end;
82            while bytes.get(number_end).is_some_and(u8::is_ascii_digit) {
83                number_end += 1;
84            }
85            if number_end == exponent_start {
86                return Err(format!("invalid timeout value '{raw}'"));
87            }
88        }
89    }
90    let (number, unit) = value.split_at(number_end);
91    let unsigned_number = &number[unsigned_start..];
92    let integer_part_end = unsigned_number
93        .find(['.', 'e', 'E'])
94        .unwrap_or(unsigned_number.len());
95    if !hexadecimal
96        && unsigned_number.starts_with('0')
97        && unsigned_number[..integer_part_end]
98            .bytes()
99            .any(|digit| matches!(digit, b'8' | b'9'))
100    {
101        // PostgreSQL first calls strtol with base 0. An 8 or 9 terminates a
102        // leading-octal integer before it can fall back to decimal parsing.
103        return Err(format!("invalid timeout value '{raw}'"));
104    }
105    let numeric = if hexadecimal {
106        u64::from_str_radix(&unsigned_number[2..], 16)
107            .map(|value| value as f64)
108            .map_err(|_| format!("invalid timeout value '{raw}'"))?
109    } else if !unsigned_number.contains(['.', 'e', 'E']) && unsigned_number.starts_with('0') {
110        u64::from_str_radix(unsigned_number, 8)
111            .map(|value| value as f64)
112            .map_err(|_| format!("invalid timeout value '{raw}'"))?
113    } else {
114        unsigned_number
115            .parse::<f64>()
116            .map_err(|_| format!("invalid timeout value '{raw}'"))?
117    };
118    let numeric = numeric * sign;
119    if !numeric.is_finite() {
120        return Err(format!("invalid timeout value '{raw}'"));
121    }
122
123    let multiplier = match unit.trim() {
124        "" | "ms" => 1.0,
125        "us" => 0.001,
126        "s" => 1_000.0,
127        "min" => 60_000.0,
128        "h" => 3_600_000.0,
129        "d" => 86_400_000.0,
130        _ => return Err(format!("invalid timeout unit in '{raw}'")),
131    };
132    let milliseconds = numeric * multiplier;
133    if !milliseconds.is_finite() || milliseconds > MAX_TIMEOUT_MS as f64 + 0.5 {
134        return Err(format!(
135            "timeout value '{raw}' exceeds PostgreSQL's maximum"
136        ));
137    }
138    // PostgreSQL's integer GUC parser uses C `rint`, which rounds halfway
139    // values to the nearest even integer under its default rounding mode.
140    let rounded = milliseconds.round_ties_even();
141    if rounded < 0.0 || rounded > MAX_TIMEOUT_MS as f64 {
142        return Err(format!(
143            "timeout value '{raw}' is outside PostgreSQL's valid range"
144        ));
145    }
146    Ok(rounded as u64)
147}
148
149#[cfg(test)]
150mod tests {
151    use super::*;
152
153    #[test]
154    fn parses_documented_time_units_into_milliseconds() {
155        assert_eq!(parse_timeout_ms("0").unwrap(), 0);
156        assert_eq!(parse_timeout_ms("500").unwrap(), 500);
157        assert_eq!(parse_timeout_ms("1500 us").unwrap(), 2);
158        assert_eq!(parse_timeout_ms("1.5ms").unwrap(), 2);
159        assert_eq!(parse_timeout_ms("2.5ms").unwrap(), 2);
160        assert_eq!(parse_timeout_ms("3.5ms").unwrap(), 4);
161        assert_eq!(parse_timeout_ms("1e-3s").unwrap(), 1);
162        assert_eq!(parse_timeout_ms("0x10ms").unwrap(), 16);
163        assert_eq!(parse_timeout_ms("+0x10ms").unwrap(), 16);
164        assert_eq!(parse_timeout_ms("010ms").unwrap(), 8);
165        assert_eq!(parse_timeout_ms(".5s").unwrap(), 500);
166        assert_eq!(parse_timeout_ms("-0.5ms").unwrap(), 0);
167        assert_eq!(parse_timeout_ms("-500us").unwrap(), 0);
168        assert_eq!(parse_timeout_ms("-0x1us").unwrap(), 0);
169        assert_eq!(parse_timeout_ms("1.5s").unwrap(), 1_500);
170        assert_eq!(parse_timeout_ms("2min").unwrap(), 120_000);
171        assert_eq!(parse_timeout_ms("1h").unwrap(), 3_600_000);
172        assert_eq!(parse_timeout_ms("1d").unwrap(), 86_400_000);
173    }
174
175    #[test]
176    fn rejects_negative_unknown_and_out_of_range_values() {
177        assert!(parse_timeout_ms("-1").is_err());
178        assert!(parse_timeout_ms("-501us").is_err());
179        assert!(parse_timeout_ms("+.5s").is_err());
180        assert!(parse_timeout_ms("1sec").is_err());
181        assert!(parse_timeout_ms("NaN").is_err());
182        assert!(parse_timeout_ms("1e+s").is_err());
183        assert!(parse_timeout_ms("09ms").is_err());
184        assert!(parse_timeout_ms("08.0ms").is_err());
185        assert!(parse_timeout_ms("2147483648ms").is_err());
186    }
187}