1"""Written queries — the statement you wrote, rendered and checked by Postgres.
2
3A query is either *built* — the code assembles it at runtime, which is what the
4QuerySet API does — or *written*: known in full when you type it, however
5complex. `Model.query.sql()` is the written side. The two meet at one seam: a
6built query drops into a written one as a subquery.
7
8The statement is a t-string (PEP 750), so Python does the interpolation and
9hands this module the literal text and each interpolated *object* separately.
10What the object is decides what it renders as:
11
12 {Model} the table, quoted
13 {Model:*} every column of the model; rows become model instances
14 {Model.field} the qualified column, "table"."column"
15 {Model.field:name} just the column, for an INSERT list or an UPDATE SET
16 {queryset} the built query, embedded as a subquery
17 {statement} another sql() statement, embedded the same way
18 {template} another t-string, rendered inline
19 {value} anything else, bound as a parameter
20
21A row is an instance only when `{Model:*}` is the whole select list; every
22other shape is a `result_type` dataclass, where a field annotated with a model
23class takes that model's expansion — the one in the outer select list, since
24an expansion inside a subquery is just columns.
25
26What that guarantees, exactly: the literal halves of a t-string are SQL the
27author wrote in the source, and every interpolated object is dispatched on
28its type, where a value always binds as a parameter. `sql()` takes a
29`Template`, so a `str` can't be passed at all — the type checker refuses a
30literal, an f-string and a runtime-built string alike.
31
32The one way to put runtime text into a statement is to build a `Template`
33from a string yourself (`Template(text)`, or concatenating one onto a
34t-string). That is a deliberate escape hatch, and it is exactly what must
35never be done with anything that came from outside the program.
36"""
37
38import ast
39import copy
40import dataclasses
41import datetime
42import decimal
43import json
44import types
45import typing
46import uuid
47import weakref
48import zoneinfo
49from collections.abc import Iterator, Sequence
50from contextlib import contextmanager
51from string.templatelib import Interpolation, Template
52from typing import TYPE_CHECKING, Any, Self
53
54import psycopg
55from plain.postgres import transaction
56from plain.postgres.base import Model
57from plain.postgres.db import get_connection
58from plain.postgres.dialect import adapt_json_value, quote_name
59from plain.postgres.fields import Field
60from plain.postgres.fields.encrypted import _ENCRYPTED_PREFIX
61from plain.postgres.fields.json import JSONField
62from plain.postgres.fields.related_descriptors import (
63 ForwardForeignKeyDescriptor,
64 ForwardManyToManyDescriptor,
65)
66from plain.postgres.fields.related_managers import BaseRelatedManager
67from plain.postgres.fields.reverse_descriptors import BaseReverseDescriptor
68from plain.postgres.fields.timezones import TimeZoneField
69from plain.postgres.otel import db_span, suppress_db_tracing
70from plain.postgres.query import QuerySet, prefetch_objects
71from plain.postgres.registry import models_registry
72from plain.postgres.sql.compiler import apply_converters, get_converters
73
74if TYPE_CHECKING:
75 from collections.abc import Generator
76
77 from plain.exceptions import ValidationError
78 from plain.postgres.connection import DatabaseConnection
79 from plain.postgres.query import Prefetch
80
81__all__ = ["Written"]
82
83
84# --------------------------------------------------------------------------
85# Rendering
86# --------------------------------------------------------------------------
87
88
89@dataclasses.dataclass(frozen=True)
90class _StarExpansion:
91 """What one `{Model:*}` put into the select list.
92
93 The fields are in declared order and the column names are theirs, which is
94 how the result columns are found again — the expansion decides which
95 positions are the instance, never the catalog.
96
97 `depth` is how many parentheses were open where it was written. Depth 0 is
98 the statement's own select list (each branch of a UNION included); deeper
99 is inside a subquery, a CTE or a function call, where the expansion is
100 just columns and the outer statement names the ones it wants. `None` means
101 the statement's text couldn't be scanned to the end, so nobody knows.
102 """
103
104 model: type[Model]
105 fields: tuple[Field, ...]
106 columns: tuple[str, ...]
107 depth: int | None
108
109
110@dataclasses.dataclass(frozen=True)
111class _Rendered:
112 sql: str
113 bind_sql: str
114 params: tuple[Any, ...]
115 stars: tuple[_StarExpansion, ...]
116
117
118class _ParenDepth:
119 """How many parentheses are open, scanning the statement as it is written.
120
121 Only the author's own text is scanned. Everything this module renders is
122 either parenthesis-free (an identifier, a `%s`) or balanced (an embedded
123 queryset or statement, wrapped in its own `(` `)`), so skipping it leaves
124 the depth exactly where the author's text put it — and a stray `(` inside
125 an embedded query's string literal can't shift it.
126
127 What is skipped, because a parenthesis inside it is text and not
128 structure: `'a (quoted) string'` (a doubled `''` is an escaped quote and
129 the string goes on), `E'it\\'s ('` — an escape string, where a backslash
130 also escapes the character after it — `"a (quoted) identifier"` like a
131 plain string, `$$ ( $$` and `$tag$ ( $tag$` dollar-quoted strings (only
132 their own tag ends them), `-- ( to end of line`, and `/* ( */` block
133 comments, which nest in Postgres.
134
135 The state carries across calls: one string or comment can span the text on
136 either side of an interpolation. A statement that ends with the scan still
137 inside a string or block comment is one this scan misread — Postgres
138 would have refused it — so `unterminated` says its depths can't be
139 trusted.
140 """
141
142 def __init__(self) -> None:
143 self.depth = 0
144 self._inside = "" # "", "'", '"', "--", "/*", or a dollar-quote tag
145 self._backslash_escapes = False # inside an E'...' string
146 self._comment_depth = 0
147
148 @property
149 def unterminated(self) -> bool:
150 # A trailing `-- comment` is a legitimate way for a statement to end.
151 return self._inside not in ("", "--")
152
153 def scan(self, text: str) -> None:
154 index = 0
155 while index < len(text):
156 if self._inside == "":
157 index = self._scan_sql(text, index)
158 elif self._inside in ("'", '"'):
159 index = self._scan_quoted(text, index)
160 elif self._inside == "--":
161 if text[index] == "\n":
162 self._inside = ""
163 index += 1
164 elif self._inside == "/*":
165 index = self._scan_comment(text, index)
166 else:
167 index = self._scan_dollar_quoted(text, index)
168
169 def _scan_sql(self, text: str, index: int) -> int:
170 char = text[index]
171 if char == "(":
172 self.depth += 1
173 elif char == ")":
174 # A statement can't have more `)` than `(`, but a fragment handed
175 # in on its own can -- never go below the outer select list.
176 self.depth = max(0, self.depth - 1)
177 elif char == "'":
178 self._inside = char
179 self._backslash_escapes = _opens_escape_string(text, index)
180 elif char == '"':
181 self._inside = char
182 self._backslash_escapes = False
183 elif text.startswith("--", index):
184 self._inside = "--"
185 return index + 2
186 elif text.startswith("/*", index):
187 self._inside = "/*"
188 self._comment_depth = 1
189 return index + 2
190 elif (tag := _dollar_quote_tag(text, index)) is not None:
191 self._inside = tag
192 return index + len(tag)
193 return index + 1
194
195 def _scan_quoted(self, text: str, index: int) -> int:
196 quote = self._inside
197 if self._backslash_escapes and text[index] == "\\":
198 return index + 2 # E'...': the next character is taken literally
199 if not text.startswith(quote, index):
200 return index + 1
201 if text.startswith(quote * 2, index):
202 return index + 2 # an escaped quote: the literal goes on
203 self._inside = ""
204 return index + 1
205
206 def _scan_comment(self, text: str, index: int) -> int:
207 if text.startswith("/*", index):
208 self._comment_depth += 1
209 return index + 2
210 if text.startswith("*/", index):
211 self._comment_depth -= 1
212 if self._comment_depth == 0:
213 self._inside = ""
214 return index + 2
215 return index + 1
216
217 def _scan_dollar_quoted(self, text: str, index: int) -> int:
218 if text.startswith(self._inside, index):
219 tag_length = len(self._inside)
220 self._inside = ""
221 return index + tag_length
222 return index + 1
223
224
225def _follows_identifier(text: str, index: int) -> bool:
226 """Whether the character before `index` is part of a word."""
227 return index > 0 and (text[index - 1].isalnum() or text[index - 1] == "_")
228
229
230def _opens_escape_string(text: str, index: int) -> bool:
231 """Whether the `'` at `index` opens an `E'...'` escape string.
232
233 The `E` has to stand alone: in `TYPE'x'` it ends a word, and the string
234 is an ordinary one.
235 """
236 return (
237 index > 0
238 and text[index - 1] in "Ee"
239 and not _follows_identifier(text, index - 1)
240 )
241
242
243def _dollar_quote_tag(text: str, index: int) -> str | None:
244 """The `$$` or `$tag$` that opens a dollar-quoted string at `index`.
245
246 A `$` inside a word is part of an identifier (`a$b$` is one), never a
247 quote.
248 """
249 if text[index] != "$" or _follows_identifier(text, index):
250 return None
251 end = text.find("$", index + 1)
252 if end == -1:
253 return None
254 tag = text[index + 1 : end]
255 if tag and not tag.isidentifier():
256 return None
257 return text[index : end + 1]
258
259
260class _Sql:
261 """The statement being rendered, in the two forms it needs.
262
263 `bind` is what psycopg is handed: it parses `%s` placeholders whether it
264 binds client- or server-side, so a literal `%` in author text has to arrive
265 doubled. `display` is the same statement as written — what `.sql`, the
266 query span and the logs show.
267 """
268
269 def __init__(self) -> None:
270 self.bind: list[str] = []
271 self.display: list[str] = []
272 self.paren_depth = _ParenDepth()
273
274 def author(self, text: str) -> None:
275 """Text the author wrote: template literals and fragments."""
276 self.bind.append(text.replace("%", "%%"))
277 self.display.append(text)
278 self.paren_depth.scan(text)
279
280 def rendered(self, sql: str, display: str | None = None) -> None:
281 """SQL this module produced: identifiers, placeholders, subqueries."""
282 self.bind.append(sql)
283 self.display.append(sql if display is None else display)
284
285
286def _qualified(table: str, column: str) -> str:
287 return f"{quote_name(table)}.{quote_name(column)}"
288
289
290def _render(template: Template) -> _Rendered:
291 """Turn a t-string into one SQL statement and its positional parameters.
292
293 Everything binds positionally: psycopg refuses a statement that mixes `%s`
294 with `%(name)s`, and an embedded queryset always brings `%s`. A value
295 interpolated twice binds twice.
296 """
297 sql = _Sql()
298 params: list[Any] = []
299 stars: list[_StarExpansion] = []
300
301 _render_template(template, sql, params, stars)
302
303 if sql.paren_depth.unterminated:
304 # The scan ended inside a string or comment Postgres would have seen
305 # closed, so it misread something and no depth it recorded can be
306 # trusted. An unknown depth is never the outer select list.
307 stars = [dataclasses.replace(star, depth=None) for star in stars]
308
309 binds_parameters = bool(params)
310 return _Rendered(
311 # The trailing `;` comes off both forms, so `.sql` and the span show
312 # the statement that actually ran.
313 sql=_one_statement("".join(sql.display), binds_parameters=binds_parameters),
314 bind_sql=_one_statement("".join(sql.bind), binds_parameters=binds_parameters),
315 params=tuple(params),
316 stars=tuple(stars),
317 )
318
319
320def _render_template(
321 template: Template,
322 sql: _Sql,
323 params: list[Any],
324 stars: list[_StarExpansion],
325) -> None:
326 """Render one t-string into the statement being built.
327
328 A `Template` alternates literal text and interpolations, starting and
329 ending with text — `strings` always has exactly one more entry than
330 `interpolations`.
331 """
332 for literal, interpolation in zip(template.strings, template.interpolations):
333 sql.author(literal)
334 _render_interpolation(interpolation, sql, params, stars)
335 sql.author(template.strings[-1])
336
337
338# The relation accessors a model class carries. A forward foreign key is not
339# one of them: it has a column of its own, and `_field_of` renders it.
340_RELATIONS = (
341 ForwardManyToManyDescriptor,
342 BaseReverseDescriptor,
343 BaseRelatedManager,
344)
345
346
347def _brace_lookalike_error(expression: str) -> str | None:
348 """The message for braces that were meant to stay braces, if these were.
349
350 `'^\\d{2}$'` is a regex whose braces weren't doubled, and Python read
351 `{2}` as an interpolation of the int 2. Nothing downstream can tell that
352 from a parameter, so the source text between the braces answers it — but
353 only for the two shapes a forgotten brace actually makes: a bare integer
354 (a quantifier) and a comma-separated run of literals (an array literal,
355 a regex range). Every other literal is left alone to bind, because
356 `{None}`, `{"active"}` and `{True}` are values someone meant.
357 """
358 try:
359 node = ast.parse(expression, mode="eval").body
360 except SyntaxError:
361 return None
362
363 doubled = f"{{{{{expression}}}}}"
364
365 # `type(...) is int` rather than isinstance: True is an int, and a bool
366 # is a value to bind.
367 if isinstance(node, ast.Constant) and type(node.value) is int:
368 return (
369 f"{{{expression}}} interpolates the number {node.value} as a bound "
370 "parameter. Write the number into the SQL, or — if this was meant "
371 "to be a literal brace, a regex quantifier like \\d{2} or an array "
372 f"literal — write {doubled}."
373 )
374
375 # `'{1,2,3}'` parses as a tuple. A set can only come from braces the
376 # author typed on purpose, but psycopg can't adapt one either way.
377 if isinstance(node, ast.Tuple | ast.Set) and all(
378 isinstance(element, ast.Constant) for element in node.elts
379 ):
380 kind = type(node).__name__.lower()
381 return (
382 f"{{{expression}}} interpolates a {kind} of literals as a bound "
383 "parameter. Write them into the SQL, or — if this was meant to be "
384 "a literal brace, an array literal or a regex range — write "
385 f"{doubled}."
386 )
387
388 return None
389
390
391def _render_interpolation(
392 interpolation: Interpolation,
393 sql: _Sql,
394 params: list[Any],
395 stars: list[_StarExpansion],
396) -> None:
397 """Render one `{...}` on what its value *is*. The value type decides.
398
399 Errors quote `interpolation.expression` — the source text between the
400 braces — so a message names what the author wrote.
401 """
402 written = interpolation.expression
403 value = interpolation.value
404 format_spec = interpolation.format_spec
405
406 if (lookalike := _brace_lookalike_error(written)) is not None:
407 raise ValueError(lookalike)
408
409 if interpolation.conversion:
410 raise ValueError(
411 f"{{{written}!{interpolation.conversion}}} uses a conversion. A "
412 "written query interpolates models and binds values — there is "
413 "nothing to convert."
414 )
415
416 if isinstance(value, type) and issubclass(value, Model):
417 _render_model(value, format_spec, written, sql, stars)
418 return
419
420 field = _field_of(value)
421 if field is not None:
422 _render_field(field, format_spec, written, sql)
423 return
424
425 if isinstance(value, Template):
426 if format_spec:
427 raise ValueError(
428 f"{{{written}:{format_spec}}} — a nested template takes no "
429 "format spec; it renders as the SQL it spells out."
430 )
431 _render_template(value, sql, params, stars)
432 return
433
434 if isinstance(value, Model):
435 # Caught here rather than at execute, where psycopg's "cannot adapt
436 # type" names the class and nothing else.
437 raise TypeError(
438 f"{{{written}}} is a {type(value).__name__} instance, not something "
439 f"a statement can hold. Interpolate a field of it "
440 f"({{{written}.id}}), or the value you meant."
441 )
442
443 if isinstance(value, _RELATIONS):
444 raise TypeError(
445 f"{{{written}}} is a relation, not a column. A written query has no "
446 "relations to follow — write the JOIN out and interpolate the "
447 "related model's own columns."
448 )
449
450 if format_spec:
451 raise ValueError(
452 f"{{{written}:{format_spec}}} — a value takes no format spec. It "
453 "binds as a parameter; formatting it would put it in the SQL. (If "
454 "you meant a literal brace in the SQL, double it: `{{` and `}}`.)"
455 )
456
457 _render_value(value, sql, params)
458
459
460def _field_of(value: Any) -> Field | None:
461 """The model field this interpolation names, if it names one.
462
463 `Widget.name` at class level *is* the `Field`. A foreign key is a
464 descriptor instead — that is what serves `Post.author.email` traversal —
465 and the field it wraps is the one that owns the `_id` column.
466 """
467 if isinstance(value, Field):
468 return value
469 if isinstance(value, ForwardForeignKeyDescriptor):
470 return value._field
471 return None
472
473
474def _render_model(
475 model: type[Model],
476 format_spec: str,
477 written: str,
478 sql: _Sql,
479 stars: list[_StarExpansion],
480) -> None:
481 """`{Model}` is the table; `{Model:*}` is every column of it."""
482 table = model.model_options.db_table
483
484 if not format_spec:
485 sql.rendered(quote_name(table))
486 return
487
488 if format_spec != "*":
489 raise ValueError(
490 f"{{{written}:{format_spec}}} — the only format spec a model takes "
491 "is `:*`, which expands to every column of it."
492 )
493
494 fields = tuple(model._model_meta.fields)
495 stars.append(
496 _StarExpansion(
497 model=model,
498 fields=fields,
499 columns=tuple(field.column for field in fields),
500 depth=sql.paren_depth.depth,
501 )
502 )
503 sql.rendered(", ".join(_qualified(table, field.column) for field in fields))
504
505
506def _render_field(field: Field, format_spec: str, written: str, sql: _Sql) -> None:
507 """`{Model.field}` is the qualified column; `{Model.field:name}` is bare."""
508 if field.is_lookup_reference:
509 # `WidgetTag.widget.name` is the traversal `where()` follows through a
510 # join it builds. A written statement builds nothing, and the field
511 # handed back doesn't even carry the table its column lives on.
512 raise ValueError(
513 f"{{{written}}} is a traversal, not a column of this statement. A "
514 "written query has no relations to follow — write the JOIN out and "
515 "interpolate the related model's own field."
516 )
517
518 if "model" not in field.__dict__:
519 raise ValueError(
520 f"{{{written}}} is a field that belongs to no model, so there is no "
521 "table to qualify its column with. Interpolate a model's own field."
522 )
523
524 if not format_spec:
525 sql.rendered(_qualified(field.model.model_options.db_table, field.column))
526 return
527
528 if format_spec != "name":
529 raise ValueError(
530 f"{{{written}:{format_spec}}} — the only format spec a column takes "
531 "is `:name`, which renders the column on its own for an INSERT list "
532 "or an UPDATE SET target."
533 )
534
535 # The bare column: an INSERT column list and an UPDATE SET target can't
536 # take a qualified name.
537 sql.rendered(quote_name(field.column))
538
539
540def _render_value(value: Any, sql: _Sql, params: list[Any]) -> None:
541 """Render a value that names nothing in the models."""
542 if isinstance(value, Written):
543 # The newlines matter, the same way they do in `_wrap`: the embedded
544 # statement can end in a `-- comment`, and the `)` would be inside it.
545 sql.rendered(f"(\n{value._bind_sql}\n)", f"(\n{value.sql}\n)")
546 params.extend(value.params)
547 elif isinstance(value, QuerySet):
548 # elide_empty=False so a queryset that can't match anything (an empty
549 # `is_in`, a `none()`) still compiles to SQL that returns no rows,
550 # instead of raising EmptyResultSet out of the middle of a render.
551 compiled, queryset_params = value.sql_query.get_compiler(
552 elide_empty=False
553 ).as_sql()
554 sql.rendered(f"({compiled})")
555 params.extend(queryset_params)
556 elif isinstance(value, dict):
557 sql.rendered("%s")
558 params.append(adapt_json_value(value, None))
559 else:
560 # Lists included: psycopg binds one as an array, which is what
561 # `= ANY({ids})` wants. Nothing is ever expanded into `IN (...)`.
562 sql.rendered("%s")
563 params.append(value)
564
565
566def _one_statement(sql: str, *, binds_parameters: bool) -> str:
567 """One written query is one statement.
568
569 A statement that binds a parameter goes to the server over the extended
570 query protocol, which carries exactly one command, so Postgres refuses
571 `SELECT 1; DROP TABLE x` itself. With no parameters to bind psycopg sends
572 the statement as a simple query instead, and the server will happily run
573 both halves — so that case is refused here.
574
575 The check is deliberately blunt: any `;` left after the trailing one comes
576 off is refused, quoted or not. A template is the statement the author
577 wrote, so the cost of being wrong is a rewrite, not a mystery.
578 """
579 sql = sql.rstrip().removesuffix(";")
580 if not binds_parameters and ";" in sql:
581 raise ValueError(
582 "A written query is a single statement, and this one contains a "
583 "';'. Split it into separate sql() calls. (A ';' inside a SQL "
584 "string literal counts too: interpolate the string instead, so it "
585 "binds as a parameter.)"
586 )
587 return sql
588
589
590def _wrap(sql: str, suffix: str = "", *, prefix: str = "SELECT * FROM ") -> str:
591 """Put a statement in a derived table.
592
593 The newlines matter: a statement can end in a `-- comment`, and anything
594 appended to that line would be inside it.
595 """
596 return f'{prefix}(\n{sql}\n) "written"{" " + suffix if suffix else ""}'
597
598
599def _leading_keyword(sql: str) -> str:
600 """The first word of a statement, past whitespace and leading comments."""
601 index = 0
602 while index < len(sql):
603 if sql[index].isspace():
604 index += 1
605 elif sql.startswith("--", index):
606 end = sql.find("\n", index)
607 index = len(sql) if end == -1 else end + 1
608 elif sql.startswith("/*", index):
609 end = sql.find("*/", index + 2)
610 index = len(sql) if end == -1 else end + 2
611 else:
612 break
613 rest = sql[index:]
614 if rest.startswith("("):
615 return "("
616 return rest.split(maxsplit=1)[0].upper() if rest.split() else ""
617
618
619# --------------------------------------------------------------------------
620# Executing
621# --------------------------------------------------------------------------
622
623
624@contextmanager
625def _written_cursor(connection: DatabaseConnection) -> Generator[Any]:
626 """A cursor that binds this statement's parameters server-side.
627
628 Plain's connections default to `ClientCursor` — client-side binding, no
629 server-side statement of any kind, which is what keeps them safe behind a
630 transaction-mode pooler like pgbouncer. A written statement uses psycopg's
631 ordinary `Cursor` instead, for one reason: the extended query protocol
632 carries exactly one command, so Postgres refuses a second statement
633 smuggled into a template whenever the statement binds a parameter. (With
634 no parameters psycopg sends a simple query, so `_one_statement` covers
635 that case.) It stays pooler-safe because the statement is UNNAMED and
636 one-shot — `prepare=True` is the named kind that isn't, and every execute
637 here passes `prepare=False` to say so.
638
639 Wrapped in Plain's own cursor wrapper so the statement is logged and
640 guarded like every other query.
641 """
642 connection.ensure_connection()
643 assert connection.connection is not None
644 with connection._prepare_cursor(psycopg.Cursor(connection.connection)) as cursor:
645 yield cursor
646
647
648# --------------------------------------------------------------------------
649# What the result columns are, and where they came from
650# --------------------------------------------------------------------------
651
652_CATALOG_SQL = """
653SELECT a.attrelid, a.attnum, c.relname, a.attname
654FROM pg_attribute a JOIN pg_class c ON c.oid = a.attrelid
655WHERE (a.attrelid, a.attnum) IN (SELECT unnest(%s::oid[]), unnest(%s::int2[]))
656"""
657
658# (table oid, column number) -> the model field that column belongs to, per
659# connection because the oids are a property of the database.
660_catalog_cache: weakref.WeakKeyDictionary[
661 DatabaseConnection, dict[tuple[int, int], Field | None]
662] = weakref.WeakKeyDictionary()
663
664# One plan per result-column signature — the names, types and sources the
665# statement actually came back with. Two renders of the same template that
666# return the same columns share a plan; anything that changes the columns
667# (a different embedded queryset, a schema change) builds a new one, and the
668# cache can only grow to the number of distinct result shapes.
669#
670# Per connection, like the catalog cache and for the same reason: a signature
671# carries table OIDs, and those belong to one database. Two databases with the
672# same schema -- a checkout and the fork it came from -- assign them
673# independently.
674_plans: weakref.WeakKeyDictionary[DatabaseConnection, dict[_Signature, _Plan]] = (
675 weakref.WeakKeyDictionary()
676)
677
678
679@dataclasses.dataclass(frozen=True)
680class _Column:
681 """One result column, as the statement that produced it described it."""
682
683 name: str
684 type_oid: int
685 table_oid: int
686 table_column: int
687
688
689type _Signature = tuple[tuple[str, int, int, int], ...]
690
691
692def _signature(columns: list[_Column]) -> _Signature:
693 return tuple(
694 (column.name, column.type_oid, column.table_oid, column.table_column)
695 for column in columns
696 )
697
698
699def _describe_columns(cursor: Any) -> list[_Column]:
700 """Read the result columns off the cursor.
701
702 `cursor.description` carries the name and the type OID; the source table
703 and column live only on the libpq result underneath it, so `pgresult` is
704 the route to them. Both are read while the result is still the cursor's.
705 """
706 result = cursor.pgresult
707 return [
708 _Column(
709 name=column.name,
710 type_oid=column.type_code,
711 table_oid=result.ftable(position),
712 table_column=result.ftablecol(position),
713 )
714 for position, column in enumerate(cursor.description)
715 ]
716
717
718def _strip_alias_marker(name: str) -> str:
719 """`n!` and `oldest?` map to `n` and `oldest`.
720
721 The markers are the sqlx convention for declaring in the statement what
722 the catalog can't know: `!` this is not null despite appearances, `?` this
723 can be null. Today they are documentation — nothing enforces them — but
724 they are stripped so the column still maps onto a field by name.
725 """
726 if name.endswith(("!", "?")):
727 return name[:-1]
728 return name
729
730
731def _fields_for_columns(
732 columns: list[_Column], connection: DatabaseConnection
733) -> list[Field | None]:
734 """The model field behind each result column, where there is one.
735
736 A column that comes from a table — through aliases, joins, derived tables
737 and CTEs — carries its source in `table_oid`/`table_column`. An aggregate,
738 an expression, or a branch of a UNION carries no source, and gets no field:
739 provenance is what attaches converters, so a column without it is returned
740 exactly as Postgres sent it.
741 """
742 pairs = {
743 (column.table_oid, column.table_column)
744 for column in columns
745 if column.table_oid and column.table_column
746 }
747 resolved = _catalog_fields(connection, pairs)
748 return [resolved.get((column.table_oid, column.table_column)) for column in columns]
749
750
751def _catalog_fields(
752 connection: DatabaseConnection, pairs: set[tuple[int, int]]
753) -> dict[tuple[int, int], Field | None]:
754 """Resolve `(table oid, column number)` pairs to model fields.
755
756 One batched catalog query for everything not already known on this
757 connection. Untraced: this is bookkeeping about the statement, not the
758 statement.
759 """
760 cache = _catalog_cache.setdefault(connection, {})
761 missing = sorted(pair for pair in pairs if pair not in cache)
762
763 if missing:
764 with suppress_db_tracing(), connection.cursor() as cursor:
765 cursor.execute(
766 _CATALOG_SQL,
767 [[pair[0] for pair in missing], [pair[1] for pair in missing]],
768 )
769 rows = cursor.fetchall()
770
771 models_by_table = {
772 model.model_options.db_table: model
773 for model in models_registry.get_models()
774 }
775 located = {
776 (relation_oid, column_number): (table_name, column_name)
777 for relation_oid, column_number, table_name, column_name in rows
778 }
779 for pair in missing:
780 cache[pair] = None
781 if pair not in located:
782 continue
783 table_name, column_name = located[pair]
784 model = models_by_table.get(table_name)
785 if model is None:
786 continue
787 cache[pair] = next(
788 (
789 field
790 for field in model._model_meta.fields
791 if field.column == column_name
792 ),
793 None,
794 )
795
796 return {pair: cache[pair] for pair in pairs}
797
798
799# --------------------------------------------------------------------------
800# The plan: how this statement's rows become results
801# --------------------------------------------------------------------------
802
803
804@dataclasses.dataclass(frozen=True)
805class _ModelField:
806 """A `result_type` field that a `{Model:*}` expansion fills.
807
808 `start` is where the expansion's run of columns begins in the result, and
809 `pk_position` is the primary key's place inside that run — the one column
810 a real row can never have NULL in, so an all-NULL run (the outer side of a
811 join that matched nothing) is recognised by it alone.
812 """
813
814 name: str
815 star: _StarExpansion
816 start: int
817 pk_position: int
818 allow_null: bool
819
820
821@dataclasses.dataclass(frozen=True)
822class _Plan:
823 """How to turn this statement's raw rows into results.
824
825 Built from the first execution's result columns and reused for as long as
826 a statement comes back with the same ones.
827 """
828
829 signature: _Signature
830 names: tuple[str, ...]
831 converters: dict[int, tuple[list[Any], Any]]
832 decryptable: tuple[int, ...]
833 star: _StarExpansion | None
834 star_start: int
835 model_fields: tuple[_ModelField, ...]
836 named_columns: tuple[tuple[int, str], ...]
837 stars: tuple[_StarExpansion, ...]
838 result_type: Any
839
840 def build(
841 self, rows: list[tuple[Any, ...]], connection: DatabaseConnection
842 ) -> list[Any]:
843 self._refuse_ciphertext(rows)
844
845 converted: Any = rows
846 if self.converters:
847 converted = apply_converters(iter(rows), self.converters, connection)
848
849 if self.star is not None:
850 return [self._hydrate(self.star, self.star_start, row) for row in converted]
851
852 assert self.result_type is not None
853 return [self._row(row) for row in converted]
854
855 def _refuse_ciphertext(self, rows: list[tuple[Any, ...]]) -> None:
856 """Refuse a column that lost its source and came back as ciphertext.
857
858 Provenance is what attaches the decrypting converter, and an
859 expression, an aggregate or a UNION drops it — so
860 `coalesce({Secret.api_key}, '')` would otherwise hand back the stored
861 token as a perfectly well-typed `str`.
862
863 Every value in a column that *could* hold one is checked, on every
864 execution: the first row of the first execution is no evidence (it can
865 be NULL, or there can be no rows at all), and a plan is reused by any
866 statement with the same column signature.
867
868 The cost is a false positive on a plain text column whose value
869 happens to start with the prefix, which is a refusal to hand back a
870 string that looks exactly like a leaked secret.
871 """
872 if not self.decryptable:
873 return
874 for row in rows:
875 for position in self.decryptable:
876 value = row[position]
877 if isinstance(value, str) and value.startswith(_ENCRYPTED_PREFIX):
878 raise TypeError(
879 f"Column {self.names[position]!r} holds an encrypted "
880 "value but lost track of the column it came from, so "
881 "nothing can decrypt it — an expression, an aggregate "
882 "or a UNION does that. Select it as {Model.field} on "
883 "its own."
884 )
885
886 def _row(self, row: Sequence[Any]) -> Any:
887 """One `result_type` row: columns by name, expansions by model class."""
888 values = {name: row[position] for position, name in self.named_columns}
889 for model_field in self.model_fields:
890 values[model_field.name] = self._model_value(model_field, row)
891 return self.result_type(**values)
892
893 def _model_value(
894 self, model_field: _ModelField, row: Sequence[Any]
895 ) -> Model | None:
896 model = model_field.star.model
897 if row[model_field.start + model_field.pk_position] is None:
898 # The expansion came back all NULL, which only an outer join that
899 # matched nothing does -- a real row always has a primary key.
900 if model_field.allow_null:
901 return None
902 raise TypeError(
903 f"{self.result_type.__name__}.{model_field.name} came back with "
904 f"every {model.__name__} column NULL -- the outer side of a "
905 f"join that matched nothing. Annotate it "
906 f"`{model.__name__} | None` to get None there."
907 )
908 return self._hydrate(model_field.star, model_field.start, row)
909
910 def _hydrate(self, star: _StarExpansion, start: int, row: Sequence[Any]) -> Model:
911 """The instance one `{Model:*}` expansion's run of columns describes."""
912 return star.model.from_db(
913 [field.name for field in star.fields],
914 list(row[start : start + len(star.fields)]),
915 )
916
917
918def _build_plan(
919 *,
920 model: type[Model],
921 stars: tuple[_StarExpansion, ...],
922 result_type: Any,
923 columns: list[_Column],
924 connection: DatabaseConnection,
925) -> _Plan:
926 names = tuple(_strip_alias_marker(column.name) for column in columns)
927 fields = _fields_for_columns(columns, connection)
928 converters = _converters_for(columns, fields, connection)
929 decryptable = _decryptable_positions(columns, fields)
930
931 if result_type is None:
932 star, star_start = _instance_star(
933 model=model, stars=stars, names=names, fields=fields
934 )
935 return _Plan(
936 signature=_signature(columns),
937 names=names,
938 converters=converters,
939 decryptable=decryptable,
940 star=star,
941 star_start=star_start,
942 model_fields=(),
943 named_columns=(),
944 stars=stars,
945 result_type=None,
946 )
947
948 # Only an expansion in the outer select list puts its columns in the
949 # result; a deeper one is just columns of a subquery, whatever the outer
950 # statement then names them.
951 outer = tuple(star for star in stars if star.depth == 0)
952 starts = _allocate_star_runs(names, fields, outer)
953
954 model_fields = _model_fields_for(
955 result_type=result_type,
956 stars=stars,
957 outer=outer,
958 starts=starts,
959 names=names,
960 )
961 _refuse_an_unmapped_expansion(
962 result_type=result_type,
963 outer=outer,
964 starts=starts,
965 model_fields=model_fields,
966 )
967 expanded = {
968 position
969 for model_field in model_fields
970 for position in range(
971 model_field.start, model_field.start + len(model_field.star.fields)
972 )
973 }
974 named_columns = tuple(
975 (position, name)
976 for position, name in enumerate(names)
977 if position not in expanded
978 )
979 _check_result_type(
980 result_type=result_type,
981 names=names,
982 model_fields=model_fields,
983 named_columns=named_columns,
984 columns=columns,
985 fields=fields,
986 connection=connection,
987 )
988 return _Plan(
989 signature=_signature(columns),
990 names=names,
991 converters=converters,
992 decryptable=decryptable,
993 star=None,
994 star_start=0,
995 model_fields=model_fields,
996 named_columns=named_columns,
997 stars=stars,
998 result_type=result_type,
999 )
1000
1001
1002def _instance_star(
1003 *,
1004 model: type[Model],
1005 stars: tuple[_StarExpansion, ...],
1006 names: tuple[str, ...],
1007 fields: list[Field | None],
1008) -> tuple[_StarExpansion, int]:
1009 """The expansion a row *is*, for a statement with no `result_type`.
1010
1011 An instance is always complete and carries only its own columns, so this
1012 is the one shape a statement can have without declaring it: `{Model:*}`
1013 and nothing else. One more column and the row has a shape of its own.
1014 """
1015 if not stars:
1016 raise TypeError(
1017 f"This statement returns columns ({', '.join(names)}) and nothing "
1018 f"says what a row is. Select {{{model.__name__}:*}} to get "
1019 "instances, or pass result_type= a dataclass."
1020 )
1021
1022 star = _instance_stars(stars)[0]
1023 start = _locate_star(names, fields, star, claimed=[])
1024 if start is None:
1025 raise TypeError(
1026 f"This statement returns ({', '.join(names)}), which doesn't "
1027 f"contain {star.model.__name__}'s columns "
1028 f"({', '.join(star.columns)}) in order. If the "
1029 f"{{{star.model.__name__}:*}} is inside a subquery, select the "
1030 "columns you want in the outer statement and pass result_type= a "
1031 "dataclass; otherwise something renamed them, and a "
1032 f"{{{star.model.__name__}:*}} row can't be aliased."
1033 )
1034
1035 extra = [
1036 name
1037 for position, name in enumerate(names)
1038 if not start <= position < start + len(star.fields)
1039 ]
1040 if extra:
1041 raise TypeError(
1042 f"This statement selects {{{star.model.__name__}:*}} and other "
1043 f"columns ({', '.join(extra)}), so a row is not a "
1044 f"{star.model.__name__} -- it has a shape of its own. Declare that "
1045 f"shape: a dataclass with a `{star.model.__name__}` field for the "
1046 "instance and a field per extra column, passed as result_type=."
1047 )
1048 return star, start
1049
1050
1051def _instance_stars(stars: tuple[_StarExpansion, ...]) -> tuple[_StarExpansion, ...]:
1052 """The expansions that could be the row of a statement with no `result_type`.
1053
1054 The outer select list first — that is where a row is declared. With no
1055 expansion there, a statement that passes a subquery's `{Model:*}` straight
1056 through (`SELECT * FROM (SELECT {Model:*} ...) sub`) still means those
1057 instances and nothing else, so the deeper ones are read as a fallback.
1058 """
1059 outer = tuple(star for star in stars if star.depth == 0)
1060 return outer or stars
1061
1062
1063def _allocate_star_runs(
1064 names: tuple[str, ...],
1065 fields: list[Field | None],
1066 stars: tuple[_StarExpansion, ...],
1067) -> list[int | None]:
1068 """Where each expansion's run of columns is, in the order they were written.
1069
1070 Two expansions never read the same columns, so each takes the first run
1071 left to it and the ones behind it have to look elsewhere. The exception is
1072 the same model expanded in each branch of a UNION: the branches collapse
1073 onto one set of result columns, so the second takes the very run the first
1074 did.
1075 """
1076 starts: list[int | None] = []
1077 claimed: list[tuple[int, int, type[Model]]] = []
1078 for star in stars:
1079 start = _locate_star(names, fields, star, claimed=claimed)
1080 starts.append(start)
1081 if start is not None:
1082 run = (start, len(star.columns), star.model)
1083 if run not in claimed:
1084 claimed.append(run)
1085 return starts
1086
1087
1088def _run_is_unclaimed(
1089 start: int,
1090 width: int,
1091 model: type[Model],
1092 claimed: list[tuple[int, int, type[Model]]],
1093) -> bool:
1094 """Whether this run of columns is still free for this expansion to take."""
1095 for claimed_start, claimed_width, claimed_model in claimed:
1096 if (start, width, model) == (claimed_start, claimed_width, claimed_model):
1097 continue # the same model's run in another UNION branch
1098 if start < claimed_start + claimed_width and claimed_start < start + width:
1099 return False
1100 return True
1101
1102
1103def _star_runs(names: tuple[str, ...], star: _StarExpansion) -> list[int]:
1104 """Every position where this expansion's run of columns could start.
1105
1106 The expansion emitted the model's columns, in declared order, as one run —
1107 so a run of result columns with those names is where it could be. Reading
1108 it back by name means the instance's fields are the ones the template asked
1109 for, not whichever columns the catalog happens to trace to this model: a
1110 self-join, a UNION branch, or another table with the same column names
1111 can't shift them.
1112 """
1113 width = len(star.columns)
1114 return [
1115 start
1116 for start in range(len(names) - width + 1)
1117 if tuple(names[start : start + width]) == star.columns
1118 ]
1119
1120
1121def _locate_star(
1122 names: tuple[str, ...],
1123 fields: list[Field | None],
1124 star: _StarExpansion,
1125 *,
1126 claimed: list[tuple[int, int, type[Model]]],
1127) -> int | None:
1128 """Where `{Model:*}`'s columns start in the result, or None if they aren't there.
1129
1130 Not there means the expansion's columns never reached the outer select
1131 list under their own names — it is inside a subquery, or another
1132 expansion already took the only run they match.
1133 """
1134 width = len(star.columns)
1135 candidates = [
1136 start
1137 for start in _star_runs(names, star)
1138 if _run_is_unclaimed(start, width, star.model, claimed)
1139 ]
1140 if not candidates:
1141 return None
1142
1143 if len(candidates) == 1:
1144 start = candidates[0]
1145 # Column names alone are weak evidence: every model's columns start
1146 # with `id`, so one model's list is often a prefix of another's. Every
1147 # column that kept its source has to be this expansion's own column in
1148 # that place; one that traces anywhere else means this run belongs to
1149 # something else. (A UNION keeps no sources, so names are all it has.)
1150 traced = [
1151 (offset, fields[start + offset])
1152 for offset in range(width)
1153 if fields[start + offset] is not None
1154 ]
1155 if not all(field is star.fields[offset] for offset, field in traced):
1156 return None
1157 return start
1158
1159 # Two runs of columns have those names. Provenance breaks the tie when it
1160 # survived; when it didn't, the statement is genuinely ambiguous.
1161 confirmed = [
1162 start
1163 for start in candidates
1164 if all(fields[start + offset] is star.fields[offset] for offset in range(width))
1165 ]
1166 if len(confirmed) != 1:
1167 raise TypeError(
1168 f"This statement returns {star.model.__name__}'s columns "
1169 f"({', '.join(star.columns)}) more than once, so which run is "
1170 f"the {star.model.__name__} is ambiguous. Alias the other one's "
1171 "columns."
1172 )
1173 return confirmed[0]
1174
1175
1176def _converters_for(
1177 columns: list[_Column],
1178 fields: list[Field | None],
1179 connection: DatabaseConnection,
1180) -> dict[int, tuple[list[Any], Any]]:
1181 """The converter each column needs before it becomes a value.
1182
1183 A column with a field gets that field's converters — decryption, JSON
1184 parsing, a timezone. A `json`/`jsonb` column with no field gets parsed
1185 anyway: Plain loads `jsonb` as text on purpose (a `JSONField`'s converter
1186 is normally what parses it), and handing back the raw text because the
1187 column came out of an expression would be a surprise, not a rule.
1188 """
1189 converters = get_converters(
1190 [
1191 None if field is None else field.get_col(field.model.model_options.db_table)
1192 for field in fields
1193 ],
1194 connection,
1195 )
1196 for position, (column, field) in enumerate(zip(columns, fields, strict=True)):
1197 if field is None and column.type_oid in _JSON_OIDS:
1198 converters[position] = ([_parse_json], None)
1199 return converters
1200
1201
1202_JSON_OIDS = frozenset({114, 3802}) # json, jsonb
1203
1204
1205def _parse_json(value: Any, expression: Any, connection: Any) -> Any:
1206 """Parse a `json`/`jsonb` column that no field converter claimed."""
1207 if isinstance(value, str | bytes):
1208 return json.loads(value)
1209 return value
1210
1211
1212# The column types an encrypted value could arrive in: it is stored as text,
1213# and an expression can put it through a json type on the way out.
1214_MAYBE_CIPHERTEXT_OIDS = frozenset({25, 1043, 114, 3802}) # text, varchar, json, jsonb
1215
1216
1217def _decryptable_positions(
1218 columns: list[_Column], fields: list[Field | None]
1219) -> tuple[int, ...]:
1220 """The positions where an undecrypted value could turn up.
1221
1222 A column with a field is decrypted by that field's converter. One without
1223 has nothing to decrypt it, so if it is text-shaped its values get checked.
1224 """
1225 return tuple(
1226 position
1227 for position, (column, field) in enumerate(zip(columns, fields, strict=True))
1228 if field is None and column.type_oid in _MAYBE_CIPHERTEXT_OIDS
1229 )
1230
1231
1232# --------------------------------------------------------------------------
1233# Result type verification
1234# --------------------------------------------------------------------------
1235
1236# The Python type psycopg hands back for each type OID, read against the
1237# adapters Plain installs.
1238_OID_TO_PYTHON: dict[int, type] = {
1239 16: bool,
1240 17: bytes,
1241 20: int,
1242 21: int,
1243 23: int,
1244 25: str,
1245 700: float,
1246 701: float,
1247 869: str, # inet, loaded as text by plain.postgres.adapters
1248 1043: str,
1249 1082: datetime.date,
1250 1083: datetime.time,
1251 1114: datetime.datetime,
1252 1184: datetime.datetime,
1253 1186: datetime.timedelta,
1254 1700: decimal.Decimal,
1255 2950: uuid.UUID,
1256}
1257
1258
1259class _AnyJson:
1260 """The stand-in for a parsed JSON value, which can be anything."""
1261
1262
1263def _python_type_for_column(
1264 column: _Column, field: Field | None, connection: DatabaseConnection
1265) -> Any:
1266 """The Python type a row will actually carry in this column.
1267
1268 The OID says what psycopg loads; a converter then says what it turns into —
1269 JSON is parsed, a TimeZoneField builds a ZoneInfo, an encrypted field
1270 decrypts to its own type.
1271 """
1272 if column.type_oid in _JSON_OIDS:
1273 return _AnyJson
1274
1275 if field is not None and field.get_db_converters(connection):
1276 if isinstance(field, JSONField):
1277 return _AnyJson
1278 if isinstance(field, TimeZoneField):
1279 return zoneinfo.ZoneInfo
1280
1281 assert connection.connection is not None
1282 info = connection.connection.adapters.types.get(column.type_oid)
1283 if info is not None and info.array_oid == column.type_oid:
1284 return list
1285
1286 return _OID_TO_PYTHON.get(column.type_oid)
1287
1288
1289def _model_annotation(annotation: Any) -> tuple[type[Model], bool] | None:
1290 """The model a `result_type` field is annotated with, and whether it allows None.
1291
1292 `widget: Widget` is `(Widget, False)` and `widget: Widget | None` is
1293 `(Widget, True)`. Anything else is an ordinary column field.
1294 """
1295 allow_null = False
1296 if typing.get_origin(annotation) in (typing.Union, types.UnionType):
1297 args = typing.get_args(annotation)
1298 named = [arg for arg in args if arg is not type(None)]
1299 allow_null = len(named) < len(args)
1300 if len(named) != 1:
1301 return None # a real union: nothing single to fill
1302 annotation = named[0]
1303
1304 if isinstance(annotation, type) and issubclass(annotation, Model):
1305 return annotation, allow_null
1306 return None
1307
1308
1309def _union_models(annotation: Any) -> list[type[Model]]:
1310 """The model classes a union annotation names, `None` aside."""
1311 if typing.get_origin(annotation) not in (typing.Union, types.UnionType):
1312 return []
1313 return [
1314 arg
1315 for arg in typing.get_args(annotation)
1316 if isinstance(arg, type) and issubclass(arg, Model)
1317 ]
1318
1319
1320def _model_fields_for(
1321 *,
1322 result_type: Any,
1323 stars: tuple[_StarExpansion, ...],
1324 outer: tuple[_StarExpansion, ...],
1325 starts: list[int | None],
1326 names: tuple[str, ...],
1327) -> tuple[_ModelField, ...]:
1328 """Pair each model-annotated field of `result_type` with its expansion.
1329
1330 The annotation names the model and an expansion is *of* a model, so the two
1331 find each other by model class. There is nothing else to match on, which is
1332 why one model can fill one field, from one run of columns.
1333 """
1334 hints = _type_hints(result_type)
1335 annotated: list[tuple[str, type[Model], bool]] = []
1336 for field in dataclasses.fields(result_type):
1337 if not field.init:
1338 continue
1339 model_annotation = _model_annotation(hints.get(field.name))
1340 if model_annotation is not None:
1341 model, allow_null = model_annotation
1342 annotated.append((field.name, model, allow_null))
1343 continue
1344 union = _union_models(hints.get(field.name))
1345 if len(union) > 1:
1346 raise TypeError(
1347 f"{result_type.__name__}.{field.name} is annotated "
1348 f"{' | '.join(model.__name__ for model in union)}. A row has "
1349 "one shape, so a model field names one model -- the one this "
1350 "statement expands -- or you select the columns you want "
1351 "instead."
1352 )
1353
1354 fields_by_model: dict[type[Model], list[str]] = {}
1355 for name, model, _ in annotated:
1356 fields_by_model.setdefault(model, []).append(name)
1357 for model, field_names in fields_by_model.items():
1358 if len(field_names) > 1:
1359 raise TypeError(
1360 f"{result_type.__name__} declares more than one "
1361 f"{model.__name__} field ({', '.join(field_names)}), and one "
1362 f"{{{model.__name__}:*}} expansion can only fill one of them. A "
1363 "self-join needs aliases an expansion can't give -- select the "
1364 "columns you want instead."
1365 )
1366
1367 if annotated:
1368 _refuse_models_sharing_a_run(outer=outer, starts=starts)
1369
1370 model_fields = []
1371 for name, model, allow_null in annotated:
1372 if not any(star.model is model for star in stars):
1373 raise TypeError(
1374 f"{result_type.__name__}.{name} is a {model.__name__}, and "
1375 f"nothing in this statement selects {{{model.__name__}:*}} to "
1376 f"fill it. Declare {{{model.__name__}:*}} in the select list, or "
1377 "drop the field."
1378 )
1379
1380 located = [
1381 (star, start)
1382 for star, start in zip(outer, starts, strict=True)
1383 if star.model is model and start is not None
1384 ]
1385 if not located and any(
1386 star.model is model and star.depth is None for star in stars
1387 ):
1388 raise TypeError(
1389 f"{result_type.__name__}.{name} is a {model.__name__}, but this "
1390 "statement's SQL couldn't be scanned to the end -- it reads as "
1391 "ending inside an unterminated string or comment -- so there is "
1392 f"no telling whether its {{{model.__name__}:*}} is in the outer "
1393 "select list. Check the quoting (an E'...' string, a $tag$ "
1394 "string, a /* comment */), or select the columns you want "
1395 "instead."
1396 )
1397 if not located:
1398 # Either every expansion of this model is inside a subquery, or
1399 # the outer one's columns never arrived under their own names.
1400 raise TypeError(
1401 f"{result_type.__name__}.{name} is a {model.__name__}, but this "
1402 f"statement returns ({', '.join(names)}), which doesn't contain "
1403 f"{model.__name__}'s columns in order. Select "
1404 f"{{{model.__name__}:*}} in the outer select list -- an "
1405 "expansion inside a subquery is just columns, and a "
1406 f"{{{model.__name__}:*}} row can't be aliased."
1407 )
1408
1409 runs = sorted({start for _, start in located})
1410 if len(runs) > 1:
1411 positions = ", ".join(str(start) for start in runs)
1412 raise TypeError(
1413 f"This statement expands {{{model.__name__}:*}} into two "
1414 f"different runs of columns (starting at {positions}), and "
1415 f"{result_type.__name__}.{name} can only be one of them. Select "
1416 "the columns you want instead."
1417 )
1418
1419 star, start = located[0]
1420 model_fields.append(
1421 _ModelField(
1422 name=name,
1423 star=star,
1424 start=start,
1425 pk_position=_pk_position(star),
1426 allow_null=allow_null,
1427 )
1428 )
1429
1430 return tuple(model_fields)
1431
1432
1433def _refuse_models_sharing_a_run(
1434 *,
1435 outer: tuple[_StarExpansion, ...],
1436 starts: list[int | None],
1437) -> None:
1438 """Two models expanded onto one run of columns can't fill a model field.
1439
1440 An outer expansion that found no run of its own is in a later branch of a
1441 UNION: its rows arrive in the same columns as the first branch's, by
1442 position. When it is a different model from the one that did find the
1443 run, a row can't say which model it is — hydrating it as the first would
1444 turn the other's rows into that model's instances.
1445 """
1446 placed = [
1447 star for star, start in zip(outer, starts, strict=True) if start is not None
1448 ]
1449 unplaced = [
1450 star for star, start in zip(outer, starts, strict=True) if start is None
1451 ]
1452 for other in unplaced:
1453 for star in placed:
1454 if star.model is other.model:
1455 continue
1456 raise TypeError(
1457 f"This statement expands {{{star.model.__name__}:*}} and "
1458 f"{{{other.model.__name__}:*}} onto the same run of result "
1459 "columns -- the branches of a UNION share one -- so a row "
1460 "can't say which model it is. Give each model its own "
1461 "statement, or select the columns you want and map them by "
1462 "name."
1463 )
1464
1465
1466def _pk_position(star: _StarExpansion) -> int:
1467 """Where the primary key sits inside the expansion's run of columns."""
1468 return next(
1469 position for position, field in enumerate(star.fields) if field.primary_key
1470 )
1471
1472
1473def _refuse_an_unmapped_expansion(
1474 *,
1475 result_type: Any,
1476 outer: tuple[_StarExpansion, ...],
1477 starts: list[int | None],
1478 model_fields: tuple[_ModelField, ...],
1479) -> None:
1480 """A `{Model:*}` whose columns the dataclass has no room for.
1481
1482 An expansion no model field claimed is columns like any others: they map
1483 onto same-named fields, which is what `SELECT {Model:*}` under a dataclass
1484 of its columns means and what a UNION of two star selects needs. When some
1485 of them have no field at all, the expansion is what the dataclass is
1486 missing, and saying so is more use than naming its columns one by one.
1487 """
1488 claimed = {model_field.star.model for model_field in model_fields}
1489 declared = {field.name for field in dataclasses.fields(result_type) if field.init}
1490 for star, start in zip(outer, starts, strict=True):
1491 if star.model in claimed or start is None:
1492 continue
1493 unmapped = [column for column in star.columns if column not in declared]
1494 if not unmapped:
1495 continue
1496 raise TypeError(
1497 f"This statement selects {{{star.model.__name__}:*}} and "
1498 f"{result_type.__name__} has no field for "
1499 f"{', '.join(unmapped)}. Declare "
1500 f"`{star.model.__name__.lower()}: {star.model.__name__}` to take "
1501 "the whole expansion, or select the columns you want instead."
1502 )
1503
1504
1505def _annotation_base(annotation: Any) -> Any:
1506 """`str | None` -> `str`, `list[int]` -> `list`, `dict[str, Any]` -> `dict`."""
1507 origin = typing.get_origin(annotation)
1508 if origin in (typing.Union, types.UnionType):
1509 args = [arg for arg in typing.get_args(annotation) if arg is not type(None)]
1510 if len(args) == 1:
1511 annotation = args[0]
1512 else:
1513 return None # a real union: nothing single to compare against
1514 return typing.get_origin(annotation) or annotation
1515
1516
1517def _type_hints(result_type: Any) -> dict[str, Any]:
1518 """`result_type`'s annotations, resolved.
1519
1520 They are resolved late — the dataclass may be defined anywhere — so a name
1521 that doesn't resolve surfaces here, where the statement can say so.
1522 """
1523 try:
1524 return typing.get_type_hints(result_type)
1525 except NameError as exc:
1526 raise TypeError(
1527 f"sql(result_type={result_type.__name__}) can't read its "
1528 f"annotations: {exc}. Every name they use has to be importable at "
1529 f"runtime — move it out of `if TYPE_CHECKING` for this dataclass."
1530 ) from exc
1531
1532
1533def _describe_result_type(result_type: Any) -> str:
1534 hints = _type_hints(result_type)
1535 lines = [
1536 f" {field.name}: {_format_annotation(hints.get(field.name, Any))}"
1537 for field in dataclasses.fields(result_type)
1538 if field.init
1539 ]
1540 return f"{result_type.__name__}:\n" + "\n".join(lines)
1541
1542
1543def _format_annotation(annotation: Any) -> str:
1544 return getattr(annotation, "__name__", str(annotation)).replace("typing.", "")
1545
1546
1547def _has_default(field: Any) -> bool:
1548 return (
1549 field.default is not dataclasses.MISSING
1550 or field.default_factory is not dataclasses.MISSING
1551 )
1552
1553
1554def _check_result_type(
1555 *,
1556 result_type: Any,
1557 names: tuple[str, ...],
1558 model_fields: tuple[_ModelField, ...],
1559 named_columns: tuple[tuple[int, str], ...],
1560 columns: list[_Column],
1561 fields: list[Field | None],
1562 connection: DatabaseConnection,
1563) -> None:
1564 """Check the result columns against the dataclass, by name then by type.
1565
1566 The columns an expansion filled are already spoken for — they went into a
1567 model-annotated field whole — so what is left is matched by name.
1568 """
1569 by_name = tuple(name for _, name in named_columns)
1570 if len(set(by_name)) != len(by_name):
1571 raise TypeError(
1572 f"This statement returns duplicate column names ({', '.join(by_name)}), "
1573 f"and {result_type.__name__} maps columns by name. Alias them apart."
1574 )
1575
1576 expanded = {model_field.name for model_field in model_fields}
1577 declared = [
1578 field
1579 for field in dataclasses.fields(result_type)
1580 if field.init and field.name not in expanded
1581 ]
1582 missing = [
1583 field.name
1584 for field in declared
1585 if field.name not in by_name and not _has_default(field)
1586 ]
1587 unexpected = [
1588 name for name in by_name if name not in {field.name for field in declared}
1589 ]
1590 if missing or unexpected:
1591 problems = []
1592 if missing:
1593 problems.append(f"no column for {', '.join(missing)}")
1594 if unexpected:
1595 problems.append(f"no field for column {', '.join(unexpected)}")
1596 # Every column the statement returned, and which of them a model field
1597 # already took whole — otherwise a statement whose expansion took
1598 # everything reads as having returned nothing.
1599 taken = "".join(
1600 f" ({model_field.name} took "
1601 f"{', '.join(names[model_field.start : model_field.start + len(model_field.star.columns)])})"
1602 for model_field in model_fields
1603 )
1604 raise TypeError(
1605 f"This statement doesn't match {result_type.__name__} — "
1606 f"{'; '.join(problems)}. It returned {', '.join(names)}{taken}, "
1607 f"and the dataclass declares\n{_describe_result_type(result_type)}"
1608 )
1609
1610 hints = _type_hints(result_type)
1611 for position, name in named_columns:
1612 expected = _python_type_for_column(
1613 columns[position], fields[position], connection
1614 )
1615 if expected is None or expected is _AnyJson:
1616 continue
1617 annotated = _annotation_base(hints.get(name, Any))
1618 if annotated in (Any, object, None):
1619 continue
1620 if annotated is not expected:
1621 raise TypeError(
1622 f"Column {name!r} comes back as {expected.__name__}, but "
1623 f"{result_type.__name__} declares it "
1624 f"{_format_annotation(hints[name])}. The statement returns\n"
1625 f"{_describe_result_type(result_type)}"
1626 )
1627
1628
1629# --------------------------------------------------------------------------
1630# Constraint violations
1631# --------------------------------------------------------------------------
1632
1633
1634def _integrity_error_to_validation_error(
1635 exc: psycopg.IntegrityError,
1636) -> ValidationError | None:
1637 """The ValidationError the violated constraint describes.
1638
1639 The same mapping `Model.create()`/`Model.update()` do, found the same way —
1640 by the constraint name Postgres reports — except that the name is looked up
1641 across every registered model, because a written statement can write any
1642 table, not just the one its queryset named.
1643 """
1644 constraint_name = exc.diag.constraint_name
1645 if not constraint_name:
1646 return None
1647
1648 for model in models_registry.get_models():
1649 meta = model._model_meta
1650 constraint = meta.constraints_by_name.get(
1651 constraint_name
1652 ) or meta.foreign_keys_by_constraint_name.get(constraint_name)
1653 if constraint is None:
1654 continue
1655 # No instance: a written statement has rows, not objects. The
1656 # constraint describes itself without one.
1657 error = constraint._db_violation_error(None, model)
1658 if error is None:
1659 return None
1660 from plain.exceptions import ValidationError
1661
1662 return ValidationError(error.update_error_dict({}))
1663
1664 return None
1665
1666
1667# --------------------------------------------------------------------------
1668# The statement
1669# --------------------------------------------------------------------------
1670
1671
1672class Written[R]:
1673 """One written statement, rendered and ready to run.
1674
1675 Immutable, and it runs at most once: iterate it and the rows are cached,
1676 the way a queryset caches its result. Call `sql()` again to run it again —
1677 which matters for writes, where a second iteration must not insert twice.
1678
1679 It renders when it is constructed, so a bad interpolation fails at the
1680 call site, and so it can be interpolated into another statement without
1681 running.
1682 """
1683
1684 def __init__(
1685 self,
1686 *,
1687 model: type[Model],
1688 template: Template,
1689 result_type: Any = None,
1690 ) -> None:
1691 if not isinstance(template, Template):
1692 # The type checker already says so; this is what an untyped call
1693 # site gets, and it names the one thing that would be an
1694 # injection if it were allowed through.
1695 raise TypeError(
1696 "sql() takes a t-string. A str -- a literal, an f-string, or "
1697 f"one built at runtime -- is not one; got "
1698 f"{type(template).__name__}."
1699 )
1700
1701 if result_type is not None and not (
1702 isinstance(result_type, type) and dataclasses.is_dataclass(result_type)
1703 ):
1704 raise TypeError("sql(result_type=...) requires a dataclass.")
1705
1706 rendered = _render(template)
1707
1708 # With a `result_type` each expansion fills the field annotated with
1709 # its model, so several models can be starred -- and a row is the
1710 # dataclass, never an instance. Without one, the row *is* the
1711 # instance: the same model's columns can be expanded more than once
1712 # (that is what each branch of a UNION needs), but two models can't
1713 # both be the row.
1714 instance_model = None
1715 if result_type is None and rendered.stars:
1716 star_models = dict.fromkeys(
1717 star.model for star in _instance_stars(rendered.stars)
1718 )
1719 if len(star_models) > 1:
1720 names = ", ".join(model.__name__ for model in star_models)
1721 raise TypeError(
1722 f"sql() expands {{Model:*}} for more than one model "
1723 f"({names}), so a row can't be one instance. Pass "
1724 "result_type= a dataclass with a field for each model."
1725 )
1726 instance_model = next(iter(star_models), None)
1727
1728 self._model = model
1729 self._result_type = result_type
1730 self._stars = rendered.stars
1731 self._instance_model = instance_model
1732 self._sql = rendered.sql
1733 self._bind_sql = rendered.bind_sql
1734 self._params = rendered.params
1735 self._prefetch_lookups: tuple[str | Prefetch, ...] = ()
1736 self._executed = False
1737 self._row_count = 0
1738 self._raw_rows: list[tuple[Any, ...]] = []
1739 self._columns: list[_Column] | None = None
1740 self._result_cache: list[R] | None = None
1741 self._count_cache: int | None = None
1742 self._exists_cache: bool | None = None
1743
1744 # -- what it renders to -------------------------------------------------
1745
1746 @property
1747 def sql(self) -> str:
1748 """The rendered statement, with `%s` where each parameter binds."""
1749 return self._sql
1750
1751 @property
1752 def params(self) -> tuple[Any, ...]:
1753 """The parameters, in the order they bind."""
1754 return self._params
1755
1756 def __repr__(self) -> str:
1757 return f"<Written: {self._sql}>"
1758
1759 def prefetch(self, *lookups: str | Prefetch) -> Self:
1760 """Load related objects for the instances this statement returns.
1761
1762 The same as `QuerySet.prefetch()`, and like it this returns a new
1763 statement — the one it was called on is untouched and still unrun.
1764 """
1765 if self._instance_model is None:
1766 raise TypeError(
1767 "prefetch() attaches related objects to model instances, and "
1768 "this statement returns rows. Select {Model:*}, or join the "
1769 "related table into the statement."
1770 )
1771 # The clone carries whatever this statement already ran -- a write
1772 # must not run a second time just because a prefetch was added -- and
1773 # only drops the hydrated rows, which is what the prefetch changes.
1774 clone = copy.copy(self)
1775 clone._prefetch_lookups = self._prefetch_lookups + lookups
1776 clone._result_cache = None
1777 return clone
1778
1779 # -- running it ---------------------------------------------------------
1780
1781 @contextmanager
1782 def _run_statement(self, sql: str, display: str) -> Generator[Any]:
1783 """Run one statement on the ORM's connection, yielding its cursor.
1784
1785 The span carries the statement as written — the `%` doubling psycopg
1786 needs is an artifact of binding, not something to read in a trace — so
1787 the cursor wrapper's own span is suppressed and this one replaces it.
1788 """
1789 connection = get_connection()
1790 params = list(self._params)
1791 with _written_cursor(connection) as cursor:
1792 try:
1793 with (
1794 db_span(
1795 connection,
1796 display,
1797 params=params,
1798 row_count_provider=lambda: cursor.rowcount,
1799 ),
1800 transaction.mark_for_rollback_on_error(),
1801 suppress_db_tracing(),
1802 ):
1803 # prepare=False: a named statement is the thing a
1804 # transaction-mode pooler can't follow, and
1805 # `prepare_threshold` is configurable, so say it outright
1806 # rather than relying on the connection's default.
1807 cursor.execute(sql, params, prepare=False)
1808 except psycopg.IntegrityError as exc:
1809 error = _integrity_error_to_validation_error(exc)
1810 if error is not None:
1811 raise error from exc
1812 raise
1813 yield cursor
1814
1815 def _run(self) -> None:
1816 """Execute the statement, once, and keep its raw rows."""
1817 if self._executed:
1818 return
1819
1820 with self._run_statement(self._bind_sql, self._sql) as cursor:
1821 self._row_count = cursor.rowcount
1822 # `description` is None for a statement with no result at all —
1823 # a write with no RETURNING clause.
1824 if cursor.description is not None:
1825 self._columns = _describe_columns(cursor)
1826 self._raw_rows = cursor.fetchall()
1827 self._executed = True
1828
1829 def _plan_for(self, columns: list[_Column]) -> _Plan:
1830 """The plan for these result columns, built once per column signature.
1831
1832 A cached plan is only reused when the statement came back with exactly
1833 the columns it was built from — same names, same types, same sources.
1834 A different embedded queryset, or a changed schema, gets a new plan
1835 instead of the wrong one.
1836 """
1837 connection = get_connection()
1838 plans = _plans.setdefault(connection, {})
1839 signature = _signature(columns)
1840 plan = plans.get(signature)
1841 if (
1842 plan is not None
1843 and plan.result_type is self._result_type
1844 and plan.stars == self._stars
1845 ):
1846 return plan
1847
1848 plan = _build_plan(
1849 model=self._model,
1850 stars=self._stars,
1851 result_type=self._result_type,
1852 columns=columns,
1853 connection=connection,
1854 )
1855 plans[signature] = plan
1856 return plan
1857
1858 def _fetch(self) -> list[R]:
1859 """The rows, hydrated — instances or `result_type` rows."""
1860 self._run()
1861 if self._result_cache is not None:
1862 return self._result_cache
1863
1864 if self._columns is None:
1865 self._result_cache = []
1866 return self._result_cache
1867
1868 plan = self._plan_for(self._columns)
1869 self._result_cache = plan.build(self._raw_rows, get_connection())
1870 if self._prefetch_lookups:
1871 prefetch_objects(self._result_cache, *self._prefetch_lookups)
1872 return self._result_cache
1873
1874 def _limited(self, limit: int) -> list[R]:
1875 """The first `limit` rows, without spending the statement's one run.
1876
1877 Only a read takes this path — `first()` and `get()` on a write go
1878 through the single execution, because running a write twice is not a
1879 cheaper way to look at fewer rows.
1880 """
1881 if self._executed:
1882 return self._fetch()[:limit]
1883
1884 # The statement's own ORDER BY is inside the subquery. Postgres
1885 # doesn't promise to keep a subquery's ordering, but it does keep it
1886 # in practice for a plain wrapper like this one -- and `first()` on an
1887 # unordered statement was already arbitrary.
1888 sql = _wrap(self._bind_sql, f"LIMIT {limit}")
1889 display = _wrap(self._sql, f"LIMIT {limit}")
1890 with self._run_statement(sql, display) as cursor:
1891 columns = _describe_columns(cursor)
1892 rows = cursor.fetchall()
1893
1894 results = self._plan_for(columns).build(rows, get_connection())
1895 if self._prefetch_lookups:
1896 prefetch_objects(results, *self._prefetch_lookups)
1897 return results
1898
1899 def _scalar(self, sql: str, display: str) -> Any:
1900 with self._run_statement(sql, display) as cursor:
1901 row = cursor.fetchone()
1902 assert row is not None
1903 return row[0]
1904
1905 def _reads_rows(self) -> bool:
1906 """Whether this statement is a read, and so safe to run twice.
1907
1908 A `SELECT`, a `VALUES`, a parenthesised one, or a `WITH` — a `WITH`
1909 that holds a write can't be wrapped at all, because Postgres requires
1910 a data-modifying CTE to be at the top level of its statement and says
1911 so rather than running it twice.
1912 """
1913 return _leading_keyword(self._sql) in ("SELECT", "WITH", "VALUES", "TABLE", "(")
1914
1915 def _require_a_result(self) -> None:
1916 """Refuse to answer a question about rows a statement doesn't have.
1917
1918 A write with no RETURNING clause produces no result at all. Counting
1919 it would answer 0, which reads like "the UPDATE matched nothing" —
1920 `execute()` is the question that has an answer.
1921 """
1922 self._run()
1923 if self._columns is None:
1924 raise TypeError(
1925 "This statement returns no rows; use execute() for the number "
1926 "of rows affected."
1927 )
1928
1929 def __iter__(self) -> Iterator[R]:
1930 return iter(self._fetch())
1931
1932 def __len__(self) -> int:
1933 self._require_a_result()
1934 return len(self._fetch())
1935
1936 def __bool__(self) -> bool:
1937 self._require_a_result()
1938 return bool(self._fetch())
1939
1940 def all(self) -> list[R]:
1941 """Every row, as a list."""
1942 return list(self._fetch())
1943
1944 def first(self) -> R | None:
1945 """The first row, or None if the statement returned none."""
1946 rows = self._limited(1) if self._reads_rows() else self._fetch()
1947 return rows[0] if rows else None
1948
1949 def get(self) -> R:
1950 """The one row the statement returned.
1951
1952 A statement that returns model instances raises that model's
1953 `DoesNotExist`/`MultipleObjectsReturned`, the same as `QuerySet.get()`.
1954 A `result_type` statement has no model to raise for, so it raises
1955 `ValueError`.
1956 """
1957 rows = self._limited(2) if self._reads_rows() else self._fetch()
1958 if len(rows) == 1:
1959 return rows[0]
1960 if self._instance_model is not None:
1961 if not rows:
1962 raise self._instance_model.DoesNotExist(
1963 f"{self._instance_model.model_options.object_name} matching "
1964 "query does not exist."
1965 )
1966 raise self._instance_model.MultipleObjectsReturned(
1967 f"get() returned more than one "
1968 f"{self._instance_model.model_options.object_name} -- it "
1969 f"returned {len(rows)}!"
1970 )
1971 if not rows:
1972 raise ValueError("get() found no rows.")
1973 raise ValueError(f"get() found {len(rows)} rows, not one.")
1974
1975 def count(self) -> int:
1976 """How many rows the statement returns."""
1977 if self._executed or not self._reads_rows():
1978 # A write counts the rows it returned, from its one execution —
1979 # wrapping it in a count() would run the write to throw the rows
1980 # away.
1981 self._require_a_result()
1982 return len(self._raw_rows)
1983 if self._count_cache is None:
1984 self._count_cache = self._scalar(
1985 _wrap(self._bind_sql, prefix="SELECT count(*) FROM "),
1986 _wrap(self._sql, prefix="SELECT count(*) FROM "),
1987 )
1988 return self._count_cache
1989
1990 def exists(self) -> bool:
1991 """Whether the statement returns any row at all."""
1992 if self._executed or not self._reads_rows():
1993 return self.count() > 0 # runs it once, and refuses a resultless one
1994 if self._count_cache is not None:
1995 return self._count_cache > 0
1996 if self._exists_cache is None:
1997 self._exists_cache = self._scalar(
1998 _wrap(
1999 self._bind_sql, prefix="SELECT EXISTS(SELECT 1 FROM ", suffix=")"
2000 ),
2001 _wrap(self._sql, prefix="SELECT EXISTS(SELECT 1 FROM ", suffix=")"),
2002 )
2003 return self._exists_cache
2004
2005 def execute(self) -> int:
2006 """Run the statement and return how many rows it affected.
2007
2008 This is the write without a RETURNING clause — the count is what an
2009 `UPDATE` or `DELETE` has to say. It runs the statement without shaping
2010 any rows, so a statement with no `result_type` and no `{Model:*}` is
2011 still runnable this way.
2012 """
2013 self._run()
2014 return self._row_count