Skip to main content

headgate_sql/
lib.rs

1use std::borrow::Cow;
2
3use sha2::{Digest, Sha256};
4
5// Every durable Postgres relation or type the driver or migrator may name. Index names
6// are handled separately: CREATE INDEX requires an unqualified name and places it beside
7// its explicitly qualified table, while DROP INDEX must identify the schema explicitly.
8const OBJECTS: &[&str] = &[
9    "headgate_state",
10    "headgate_job",
11    "headgate_rate_bucket",
12    "headgate_quarantine",
13    "headgate_partition_deficit",
14    "headgate_active_partition",
15    "headgate_inflight",
16    "headgate_concurrency_limit",
17    "headgate_queue_counter",
18    "headgate_partition_counter",
19    "headgate_queue_state",
20    "headgate_duty",
21    "headgate_schedule",
22    "headgate_schedule_event",
23    "headgate_worker",
24    "headgate_effect",
25    "headgate_operation",
26    "headgate_enqueue_policy",
27    "headgate_enqueue_counter",
28    "headgate_job_tag",
29    "headgate_queue_sample",
30    "headgate_archive_policy",
31    "headgate_job_archive",
32    "headgate_job_archive_before_2025",
33    "headgate_job_archive_202501",
34    "headgate_job_archive_202502",
35    "headgate_job_archive_202503",
36    "headgate_job_archive_202504",
37    "headgate_job_archive_202505",
38    "headgate_job_archive_202506",
39    "headgate_job_archive_202507",
40    "headgate_job_archive_202508",
41    "headgate_job_archive_202509",
42    "headgate_job_archive_202510",
43    "headgate_job_archive_202511",
44    "headgate_job_archive_202512",
45    "headgate_job_archive_202601",
46    "headgate_job_archive_202602",
47    "headgate_job_archive_202603",
48    "headgate_job_archive_202604",
49    "headgate_job_archive_202605",
50    "headgate_job_archive_202606",
51    "headgate_job_archive_202607",
52    "headgate_job_archive_202608",
53    "headgate_job_archive_202609",
54    "headgate_job_archive_202610",
55    "headgate_job_archive_202611",
56    "headgate_job_archive_202612",
57    "headgate_job_archive_202701",
58    "headgate_job_archive_202702",
59    "headgate_job_archive_202703",
60    "headgate_job_archive_202704",
61    "headgate_job_archive_202705",
62    "headgate_job_archive_202706",
63    "headgate_job_archive_202707",
64    "headgate_job_archive_202708",
65    "headgate_job_archive_202709",
66    "headgate_job_archive_202710",
67    "headgate_job_archive_202711",
68    "headgate_job_archive_202712",
69    "headgate_job_archive_202801",
70    "headgate_job_archive_202802",
71    "headgate_job_archive_202803",
72    "headgate_job_archive_202804",
73    "headgate_job_archive_202805",
74    "headgate_job_archive_202806",
75    "headgate_job_archive_202807",
76    "headgate_job_archive_202808",
77    "headgate_job_archive_202809",
78    "headgate_job_archive_202810",
79    "headgate_job_archive_202811",
80    "headgate_job_archive_202812",
81    "headgate_job_archive_202901",
82    "headgate_job_archive_202902",
83    "headgate_job_archive_202903",
84    "headgate_job_archive_202904",
85    "headgate_job_archive_202905",
86    "headgate_job_archive_202906",
87    "headgate_job_archive_202907",
88    "headgate_job_archive_202908",
89    "headgate_job_archive_202909",
90    "headgate_job_archive_202910",
91    "headgate_job_archive_202911",
92    "headgate_job_archive_202912",
93    "headgate_job_archive_203001",
94    "headgate_job_archive_203002",
95    "headgate_job_archive_203003",
96    "headgate_job_archive_203004",
97    "headgate_job_archive_203005",
98    "headgate_job_archive_203006",
99    "headgate_job_archive_203007",
100    "headgate_job_archive_203008",
101    "headgate_job_archive_203009",
102    "headgate_job_archive_203010",
103    "headgate_job_archive_203011",
104    "headgate_job_archive_203012",
105    "headgate_job_archive_203101",
106    "headgate_job_archive_203102",
107    "headgate_job_archive_203103",
108    "headgate_job_archive_203104",
109    "headgate_job_archive_203105",
110    "headgate_job_archive_203106",
111    "headgate_job_archive_203107",
112    "headgate_job_archive_203108",
113    "headgate_job_archive_203109",
114    "headgate_job_archive_203110",
115    "headgate_job_archive_203111",
116    "headgate_job_archive_203112",
117    "headgate_job_archive_after_2031",
118    "headgate_track_enqueue_depth",
119    "headgate_schema_migration",
120];
121
122const INDEXES: &[&str] = &[
123    "headgate_job_unique",
124    "headgate_job_unique_throttle",
125    "headgate_job_sticky_available",
126    "headgate_job_avail_sticky",
127    "headgate_job_archive_queue_time",
128];
129
130#[derive(Clone, Debug, Default)]
131pub struct PostgresNamespace {
132    name: Option<String>,
133    quoted: String,
134    wakeup_channel: String,
135}
136
137impl PostgresNamespace {
138    pub fn explicit(name: &str) -> Result<Self, String> {
139        if name.is_empty() {
140            return Err("Postgres schema must not be empty".into());
141        }
142        if name.as_bytes().contains(&0) {
143            return Err("Postgres schema must not contain NUL".into());
144        }
145        // NAMEDATALEN is 64, including the terminator. Reject instead of allowing the
146        // server to truncate two configured instances onto the same identifier.
147        if name.len() > 63 {
148            return Err("Postgres schema must be at most 63 UTF-8 bytes".into());
149        }
150        let digest = Sha256::digest(name.as_bytes());
151        let digest_hex = format!("{digest:x}");
152        let channel_hash = &digest_hex[..16];
153        Ok(Self {
154            name: Some(name.to_owned()),
155            quoted: quote_identifier(name),
156            wakeup_channel: format!("headgate_wakeup_{channel_hash}"),
157        })
158    }
159
160    pub fn name(&self) -> Option<&str> {
161        self.name.as_deref()
162    }
163
164    pub fn wakeup_channel(&self) -> &str {
165        if self.name.is_some() {
166            &self.wakeup_channel
167        } else {
168            "headgate_wakeup"
169        }
170    }
171
172    pub fn render<'a>(&self, sql: &'a str) -> Cow<'a, str> {
173        if self.name.is_none() {
174            return Cow::Borrowed(sql);
175        }
176        Cow::Owned(qualify_sql(sql, &self.quoted, self.wakeup_channel()))
177    }
178}
179
180pub fn quote_identifier(value: &str) -> String {
181    format!("\"{}\"", value.replace('"', "\"\""))
182}
183
184fn is_ident_start(byte: u8) -> bool {
185    byte == b'_' || byte.is_ascii_alphabetic()
186}
187
188fn is_ident_continue(byte: u8) -> bool {
189    is_ident_start(byte) || byte.is_ascii_digit() || byte == b'$'
190}
191
192fn dollar_delimiter(sql: &str, start: usize) -> Option<&str> {
193    let bytes = sql.as_bytes();
194    if bytes.get(start) != Some(&b'$') {
195        return None;
196    }
197    let mut end = start + 1;
198    while end < bytes.len() && (bytes[end].is_ascii_alphanumeric() || bytes[end] == b'_') {
199        end += 1;
200    }
201    (bytes.get(end) == Some(&b'$')).then(|| &sql[start..=end])
202}
203
204fn qualify_sql(sql: &str, quoted_schema: &str, wakeup_channel: &str) -> String {
205    let bytes = sql.as_bytes();
206    let mut out = String::with_capacity(sql.len() + 64);
207    let mut i = 0;
208    let mut previous_ident = "";
209    let mut last_ident = "";
210    while i < bytes.len() {
211        // Line comments.
212        if bytes[i..].starts_with(b"--") {
213            let end = sql[i..].find('\n').map_or(bytes.len(), |n| i + n + 1);
214            out.push_str(&sql[i..end]);
215            i = end;
216            continue;
217        }
218        // Nested block comments (Postgres permits nesting).
219        if bytes[i..].starts_with(b"/*") {
220            let start = i;
221            i += 2;
222            let mut depth = 1_u32;
223            while i < bytes.len() && depth != 0 {
224                if bytes[i..].starts_with(b"/*") {
225                    depth += 1;
226                    i += 2;
227                } else if bytes[i..].starts_with(b"*/") {
228                    depth -= 1;
229                    i += 2;
230                } else {
231                    i += 1;
232                }
233            }
234            out.push_str(&sql[start..i]);
235            continue;
236        }
237        // SQL strings. Handle doubled quotes and E'...' backslash escapes conservatively.
238        if bytes[i] == b'\'' {
239            let start = i;
240            i += 1;
241            while i < bytes.len() {
242                if bytes[i] == b'\\' {
243                    i = (i + 2).min(bytes.len());
244                } else if bytes[i] == b'\'' {
245                    i += 1;
246                    if i < bytes.len() && bytes[i] == b'\'' {
247                        i += 1;
248                    } else {
249                        break;
250                    }
251                } else {
252                    i += 1;
253                }
254            }
255            let literal = &sql[start..i];
256            if literal == "'headgate_wakeup'" {
257                out.push('\'');
258                out.push_str(wakeup_channel);
259                out.push('\'');
260            } else {
261                out.push_str(literal);
262            }
263            continue;
264        }
265        // Existing quoted identifiers are intentional and never rewritten.
266        if bytes[i] == b'"' {
267            let start = i;
268            i += 1;
269            while i < bytes.len() {
270                if bytes[i] == b'"' {
271                    i += 1;
272                    if i < bytes.len() && bytes[i] == b'"' {
273                        i += 1;
274                    } else {
275                        break;
276                    }
277                } else {
278                    i += 1;
279                }
280            }
281            out.push_str(&sql[start..i]);
282            continue;
283        }
284        // Dollar-quoted function bodies and strings.
285        if let Some(delimiter) = dollar_delimiter(sql, i) {
286            let start = i;
287            let body = i + delimiter.len();
288            i = sql[body..]
289                .find(delimiter)
290                .map_or(bytes.len(), |n| body + n + delimiter.len());
291            out.push_str(&sql[start..i]);
292            continue;
293        }
294        if is_ident_start(bytes[i]) {
295            let start = i;
296            i += 1;
297            while i < bytes.len() && is_ident_continue(bytes[i]) {
298                i += 1;
299            }
300            let token = &sql[start..i];
301            let dropping_index = INDEXES.contains(&token)
302                && previous_ident.eq_ignore_ascii_case("DROP")
303                && last_ident.eq_ignore_ascii_case("INDEX");
304            if OBJECTS.contains(&token) || dropping_index {
305                out.push_str(quoted_schema);
306                out.push('.');
307                out.push_str(token);
308            } else {
309                out.push_str(token);
310            }
311            previous_ident = last_ident;
312            last_ident = token;
313            continue;
314        }
315        let ch = sql[i..].chars().next().expect("valid UTF-8");
316        out.push(ch);
317        i += ch.len_utf8();
318    }
319    out
320}
321
322#[cfg(test)]
323mod tests {
324    use super::*;
325
326    #[test]
327    fn explicit_schema_quotes_objects_but_not_literals_comments_or_aliases() {
328        let namespace = PostgresNamespace::explicit("tenant-\"blue").unwrap();
329        let sql = "SELECT headgate_job.id, headgate_inflight_stale FROM headgate_job \
330                   JOIN headgate_rate_bucket b ON true \
331                   JOIN headgate_schedule_event e ON true \
332                   WHERE note = 'headgate_job' /* headgate_duty */ -- headgate_worker\n\
333                   AND state = 'available'::headgate_state";
334        let rendered = namespace.render(sql);
335        assert!(rendered.contains("\"tenant-\"\"blue\".headgate_job.id"));
336        assert!(rendered.contains("\"tenant-\"\"blue\".headgate_rate_bucket"));
337        assert!(rendered.contains("\"tenant-\"\"blue\".headgate_schedule_event"));
338        assert!(rendered.contains("::\"tenant-\"\"blue\".headgate_state"));
339        assert!(rendered.contains("headgate_inflight_stale FROM"));
340        assert!(rendered.contains("'headgate_job' /* headgate_duty */"));
341        assert!(rendered.contains("-- headgate_worker"));
342    }
343
344    #[test]
345    fn explicit_schema_namespaces_notifications_and_default_is_byte_identity() {
346        let sql = "SELECT pg_notify('headgate_wakeup', queue) FROM headgate_job";
347        assert_eq!(PostgresNamespace::default().render(sql), sql);
348        let namespace = PostgresNamespace::explicit("tenant").unwrap();
349        let rendered = namespace.render(sql);
350        assert!(rendered.contains(namespace.wakeup_channel()));
351        assert!(rendered.contains("\"tenant\".headgate_job"));
352        assert_eq!(
353            quote_identifier(namespace.wakeup_channel()),
354            format!("\"{}\"", namespace.wakeup_channel())
355        );
356    }
357
358    #[test]
359    fn invalid_schema_names_fail_instead_of_truncating_or_sharing() {
360        assert!(PostgresNamespace::explicit("").is_err());
361        assert!(PostgresNamespace::explicit("bad\0schema").is_err());
362        assert!(PostgresNamespace::explicit(&"x".repeat(64)).is_err());
363        assert!(PostgresNamespace::explicit(&"x".repeat(63)).is_ok());
364    }
365
366    #[test]
367    fn enqueue_backpressure_objects_are_namespaced_outside_the_trigger_body() {
368        let namespace = PostgresNamespace::explicit("tenant").unwrap();
369        let sql = "CREATE TABLE headgate_enqueue_policy (queue text);\n\
370                   CREATE TABLE headgate_enqueue_counter (queue text);\n\
371                   CREATE OR REPLACE FUNCTION headgate_track_enqueue_depth()\n\
372                   RETURNS trigger LANGUAGE plpgsql AS $$\n\
373                   BEGIN\n\
374                     EXECUTE format('INSERT INTO %I.headgate_enqueue_counter VALUES ($1)',\n\
375                                    TG_TABLE_SCHEMA) USING NEW.queue;\n\
376                     RETURN NEW;\n\
377                   END;\n\
378                   $$;\n\
379                   CREATE TRIGGER track AFTER INSERT ON headgate_job\n\
380                   FOR EACH ROW EXECUTE FUNCTION headgate_track_enqueue_depth();";
381        let rendered = namespace.render(sql);
382
383        assert!(rendered.contains("CREATE TABLE \"tenant\".headgate_enqueue_policy"));
384        assert!(rendered.contains("CREATE TABLE \"tenant\".headgate_enqueue_counter"));
385        assert!(
386            rendered
387                .contains("CREATE OR REPLACE FUNCTION \"tenant\".headgate_track_enqueue_depth()")
388        );
389        assert!(rendered.contains("ON \"tenant\".headgate_job"));
390        assert!(rendered.contains("EXECUTE FUNCTION \"tenant\".headgate_track_enqueue_depth()"));
391        assert!(rendered.contains("%I.headgate_enqueue_counter"));
392        assert!(rendered.contains("TG_TABLE_SCHEMA"));
393    }
394
395    #[test]
396    fn migration_indexes_and_new_metric_tables_stay_inside_the_explicit_schema() {
397        let namespace = PostgresNamespace::explicit("tenant").unwrap();
398        let rendered = namespace.render(
399            "DROP INDEX headgate_job_unique;\n\
400             CREATE UNIQUE INDEX headgate_job_unique ON headgate_job (unique_key);\n\
401             CREATE TABLE headgate_job_tag (job_id bigint);\n\
402             CREATE TABLE headgate_queue_sample (queue text);",
403        );
404        assert!(rendered.contains("DROP INDEX \"tenant\".headgate_job_unique"));
405        assert!(
406            rendered.contains("CREATE UNIQUE INDEX headgate_job_unique ON \"tenant\".headgate_job")
407        );
408        assert!(rendered.contains("CREATE TABLE \"tenant\".headgate_job_tag"));
409        assert!(rendered.contains("CREATE TABLE \"tenant\".headgate_queue_sample"));
410    }
411}