1use std::borrow::Cow;
2
3use sha2::{Digest, Sha256};
4
5const 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 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 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 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 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 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 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}