Coverage for snekql/query.py: 86%
515 statements
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 21:13 +0300
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 21:13 +0300
1"""Query Builder objects and factory functions."""
3from __future__ import annotations
5from collections.abc import Sequence
6from dataclasses import dataclass, replace
7from typing import Any, Protocol, Self, TypeVar, TypeVarTuple, cast, overload
9from snekql._model_materialization import (
10 decode_column_value,
11 encode_column_value,
12)
13from snekql._query_dialect import QueryDialect
14from snekql.errors import (
15 ModelDeclarationError,
16 QueryCompilationError,
17 QueryConstructionError,
18)
19from snekql.expressions import Assignment, OrderBy, Predicate
20from snekql.model import (
21 Model,
22 Table,
23 decode_model_row,
24 require_model_columns,
25 require_model_table_name,
26)
27from snekql.sqlite.identifiers import quote_identifier as quote_sqlite_identifier
28from snekql.storage import MISSING, Attr
29from snekql.validation import NonNegativeInt, validate_boundary
31ModelT = TypeVar("ModelT", bound=Table[Any])
32ReadModelT = TypeVar("ReadModelT", bound=Table[Any])
33SelectOwnerT = TypeVar("SelectOwnerT", bound=Table[Any])
34OwnerT = TypeVar("OwnerT", bound=Table[Any])
35SelectableOwnerT_co = TypeVar("SelectableOwnerT_co", bound=Table[Any], covariant=True)
36SelectableReadT_co = TypeVar("SelectableReadT_co", bound=Table[Any], covariant=True)
37T = TypeVar("T")
38T1 = TypeVar("T1")
39T2 = TypeVar("T2")
40T3 = TypeVar("T3")
41Ts = TypeVarTuple("Ts")
43_BINARY_PREDICATE_CHILD_COUNT = 2
44_UNARY_PREDICATE_CHILD_COUNT = 1
47def _sqlite_empty_insert_sql(quoted_table: str) -> str:
48 return "INSERT INTO " + quoted_table + " DEFAULT VALUES"
51def _encode_sqlite_column_value(
52 column: Attr[Any, Any, Any, Any, Any],
53 value: object,
54) -> object:
55 return encode_column_value(column, value, backend="sqlite")
58_SQLITE_QUERY_DIALECT = QueryDialect(
59 empty_insert_sql=_sqlite_empty_insert_sql,
60 encode_column_value=_encode_sqlite_column_value,
61 placeholder="?",
62 quote_identifier=quote_sqlite_identifier,
63)
66class _SelectableModelClass(Protocol[SelectableOwnerT_co, SelectableReadT_co]):
67 """Structural type for model classes accepted by `select(Model)`.
69 The protocol lets pyright connect the writable owner model type with the
70 fetched read model type exposed by table model classes.
71 """
73 @classmethod
74 def __owner_type__(cls) -> type[SelectableOwnerT_co]: ...
76 @classmethod
77 def __read_type__(cls) -> type[SelectableReadT_co]: ...
80@dataclass(frozen=True)
81class _SelectState:
82 model: type[Table[Any]]
83 fields: tuple[Attr[Any, Any, Any, Any, Any], ...]
84 returns_model: bool = False
85 explicit_all: bool = False
86 predicates: tuple[Predicate[Any], ...] = ()
87 orderings: tuple[OrderBy[Any], ...] = ()
88 limit_value: int | None = None
89 offset_value: int | None = None
92@dataclass(frozen=True)
93class _UpdateState:
94 model: type[Table[Any]]
95 assignments: tuple[Assignment[Any], ...] = ()
96 explicit_all: bool = False
97 predicates: tuple[Predicate[Any], ...] = ()
100@dataclass(frozen=True)
101class _DeleteState:
102 model: type[Table[Any]]
103 explicit_all: bool = False
104 predicates: tuple[Predicate[Any], ...] = ()
107class SelectModelQuery[SelectOwnerT: Table[Any], ReadModelT: Table[Any]]:
108 """Immutable select query that returns fetched table model instances."""
110 state: _SelectState
112 def __init__(self, state: _SelectState | None = None) -> None:
113 if state is None:
114 state = _empty_select_state()
115 self.state = state
117 def all(self) -> Self:
118 state = _select_all(self.state)
119 if state is self.state:
120 return self
121 return cast("Self", SelectModelQuery[SelectOwnerT, ReadModelT](state))
123 def where(self, *predicates: Predicate[SelectOwnerT]) -> Self:
124 state = _select_where(self.state, predicates)
125 return cast("Self", SelectModelQuery[SelectOwnerT, ReadModelT](state))
127 def order_by(self, *ordering: OrderBy[SelectOwnerT]) -> Self:
128 state = _select_order_by(self.state, ordering)
129 return cast("Self", SelectModelQuery[SelectOwnerT, ReadModelT](state))
131 @validate_boundary(error_type=QueryConstructionError)
132 def limit(self, value: NonNegativeInt) -> Self:
133 state = _select_limit(self.state, value)
134 return cast("Self", SelectModelQuery[SelectOwnerT, ReadModelT](state))
136 @validate_boundary(error_type=QueryConstructionError)
137 def offset(self, value: NonNegativeInt) -> Self:
138 state = _select_offset(self.state, value)
139 return cast("Self", SelectModelQuery[SelectOwnerT, ReadModelT](state))
142class SelectValueQuery[OwnerT: Table[Any], T]:
143 """Immutable select query that returns one scalar column value per row."""
145 state: _SelectState
147 def __init__(self, state: _SelectState | None = None) -> None:
148 if state is None:
149 state = _empty_select_state()
150 self.state = state
152 def all(self) -> Self:
153 state = _select_all(self.state)
154 if state is self.state:
155 return self
156 return cast("Self", SelectValueQuery[OwnerT, T](state))
158 def where(self, *predicates: Predicate[OwnerT]) -> Self:
159 state = _select_where(self.state, predicates)
160 return cast("Self", SelectValueQuery[OwnerT, T](state))
162 def order_by(self, *ordering: OrderBy[OwnerT]) -> Self:
163 state = _select_order_by(self.state, ordering)
164 return cast("Self", SelectValueQuery[OwnerT, T](state))
166 @validate_boundary(error_type=QueryConstructionError)
167 def limit(self, value: NonNegativeInt) -> Self:
168 state = _select_limit(self.state, value)
169 return cast("Self", SelectValueQuery[OwnerT, T](state))
171 @validate_boundary(error_type=QueryConstructionError)
172 def offset(self, value: NonNegativeInt) -> Self:
173 state = _select_offset(self.state, value)
174 return cast("Self", SelectValueQuery[OwnerT, T](state))
177class SelectTupleQuery[OwnerT: Table[Any], *Ts]:
178 """Immutable select query that returns selected column tuples per row."""
180 state: _SelectState
182 def __init__(self, state: _SelectState | None = None) -> None:
183 if state is None:
184 state = _empty_select_state()
185 self.state = state
187 def all(self) -> Self:
188 state = _select_all(self.state)
189 if state is self.state:
190 return self
191 return cast("Self", SelectTupleQuery[OwnerT, *Ts](state))
193 def where(self, *predicates: Predicate[OwnerT]) -> Self:
194 state = _select_where(self.state, predicates)
195 return cast("Self", SelectTupleQuery[OwnerT, *Ts](state))
197 def order_by(self, *ordering: OrderBy[OwnerT]) -> Self:
198 state = _select_order_by(self.state, ordering)
199 return cast("Self", SelectTupleQuery[OwnerT, *Ts](state))
201 @validate_boundary(error_type=QueryConstructionError)
202 def limit(self, value: NonNegativeInt) -> Self:
203 state = _select_limit(self.state, value)
204 return cast("Self", SelectTupleQuery[OwnerT, *Ts](state))
206 @validate_boundary(error_type=QueryConstructionError)
207 def offset(self, value: NonNegativeInt) -> Self:
208 state = _select_offset(self.state, value)
209 return cast("Self", SelectTupleQuery[OwnerT, *Ts](state))
212class InsertQuery[ModelT: Table[Any]]:
213 """Immutable insert statement for one pending table model instance."""
215 row: ModelT
217 def __init__(self, row: ModelT) -> None:
218 self.row: ModelT = row
221class UpdateQuery[ModelT: Table[Any]]:
222 """Immutable update statement for one table model."""
224 state: _UpdateState
226 def __init__(self, state: _UpdateState | None = None) -> None:
227 if state is None:
228 state = _UpdateState(model=Table[Any])
229 self.state: _UpdateState = state
231 def all(self) -> Self:
232 state = _update_all(self.state)
233 if state is self.state:
234 return self
235 return cast("Self", UpdateQuery[ModelT](state))
237 def set(self, *assignments: Assignment[ModelT]) -> Self:
238 state = _update_set(self.state, assignments)
239 return cast("Self", UpdateQuery[ModelT](state))
241 def where(self, *predicates: Predicate[ModelT]) -> Self:
242 state = _update_where(self.state, predicates)
243 return cast("Self", UpdateQuery[ModelT](state))
246class DeleteQuery[ModelT: Table[Any]]:
247 """Immutable delete statement for one table model."""
249 state: _DeleteState
251 def __init__(self, state: _DeleteState | None = None) -> None:
252 if state is None:
253 state = _DeleteState(model=Table[Any])
254 self.state: _DeleteState = state
256 def all(self) -> Self:
257 state = _delete_all(self.state)
258 if state is self.state:
259 return self
260 return cast("Self", DeleteQuery[ModelT](state))
262 def where(self, *predicates: Predicate[ModelT]) -> Self:
263 state = _delete_where(self.state, predicates)
264 return cast("Self", DeleteQuery[ModelT](state))
267type AnySelectQuery = (
268 SelectModelQuery[Any, Any]
269 | SelectValueQuery[Any, Any]
270 | SelectTupleQuery[Any, *tuple[Any, ...]]
271)
274def _empty_select_state() -> _SelectState:
275 return _SelectState(model=Table[Any], fields=())
278def _select_all(state: _SelectState) -> _SelectState:
279 if state.predicates:
280 msg = "all() cannot be combined with where()"
281 raise QueryConstructionError(msg)
282 if state.explicit_all:
283 return state
284 return replace(state, explicit_all=True)
287def _select_where(
288 state: _SelectState,
289 predicates: tuple[Predicate[Any], ...],
290) -> _SelectState:
291 if not predicates:
292 msg = "where() requires at least one predicate"
293 raise QueryConstructionError(msg)
294 if state.explicit_all:
295 msg = "where() cannot be combined with all()"
296 raise QueryConstructionError(msg)
297 for predicate in predicates:
298 _ensure_predicate_targets_model(predicate, state.model)
299 return replace(state, predicates=(*state.predicates, *predicates))
302def _select_order_by(
303 state: _SelectState,
304 orderings: tuple[OrderBy[Any], ...],
305) -> _SelectState:
306 if not orderings:
307 msg = "order_by() requires at least one ordering"
308 raise QueryConstructionError(msg)
309 for ordering in orderings:
310 _ensure_ordering_targets_model(ordering, state.model)
311 return replace(state, orderings=(*state.orderings, *orderings))
314def _select_limit(state: _SelectState, value: NonNegativeInt) -> _SelectState:
315 return replace(state, limit_value=value)
318def _select_offset(state: _SelectState, value: NonNegativeInt) -> _SelectState:
319 return replace(state, offset_value=value)
322def _update_all(state: _UpdateState) -> _UpdateState:
323 if state.predicates:
324 msg = "all() cannot be combined with where()"
325 raise QueryConstructionError(msg)
326 if state.explicit_all:
327 return state
328 return replace(state, explicit_all=True)
331def _update_set(
332 state: _UpdateState,
333 assignments: tuple[Assignment[Any], ...],
334) -> _UpdateState:
335 if not assignments:
336 msg = "set() requires at least one assignment"
337 raise QueryConstructionError(msg)
338 for assignment in assignments:
339 _ensure_assignment_targets_model(assignment, state.model)
340 return replace(state, assignments=(*state.assignments, *assignments))
343def _update_where(
344 state: _UpdateState,
345 predicates: tuple[Predicate[Any], ...],
346) -> _UpdateState:
347 if not predicates:
348 msg = "where() requires at least one predicate"
349 raise QueryConstructionError(msg)
350 if state.explicit_all:
351 msg = "where() cannot be combined with all()"
352 raise QueryConstructionError(msg)
353 for predicate in predicates:
354 _ensure_predicate_targets_model(predicate, state.model)
355 return replace(state, predicates=(*state.predicates, *predicates))
358def _delete_all(state: _DeleteState) -> _DeleteState:
359 if state.predicates:
360 msg = "all() cannot be combined with where()"
361 raise QueryConstructionError(msg)
362 if state.explicit_all:
363 return state
364 return replace(state, explicit_all=True)
367def _delete_where(
368 state: _DeleteState,
369 predicates: tuple[Predicate[Any], ...],
370) -> _DeleteState:
371 if not predicates:
372 msg = "where() requires at least one predicate"
373 raise QueryConstructionError(msg)
374 if state.explicit_all:
375 msg = "where() cannot be combined with all()"
376 raise QueryConstructionError(msg)
377 for predicate in predicates:
378 _ensure_predicate_targets_model(predicate, state.model)
379 return replace(state, predicates=(*state.predicates, *predicates))
382def _require_field(value: object) -> Attr[Any, Any, Any, Any, Any]:
383 if not isinstance(value, Attr):
384 msg = "select requires a model or field"
385 raise QueryConstructionError(msg)
386 return cast("Attr[Any, Any, Any, Any, Any]", value)
389def _require_column_name(column: Attr[Any, Any, Any, Any, Any]) -> str:
390 if column.name is None:
391 msg = "field is not bound to a model"
392 raise QueryConstructionError(msg)
393 return column.name
396def _require_column_model(column: Attr[Any, Any, Any, Any, Any]) -> type[Table[Any]]:
397 owner = column.owner
398 if owner is None:
399 msg = "field is not bound to a model"
400 raise QueryConstructionError(msg)
401 model = cast("type[Table[Any]]", owner)
402 try:
403 _ = require_model_columns(model)
404 except ModelDeclarationError as error:
405 msg = "field is not bound to a table model"
406 raise QueryConstructionError(msg) from error
407 return model
410def _ensure_predicate_targets_model(
411 predicate: Predicate[Any],
412 model: type[Table[Any]],
413) -> None:
414 if predicate.kind == "":
415 msg = "where predicates must be built from columns"
416 raise QueryConstructionError(msg)
417 if predicate.column is not None:
418 column = _require_field(predicate.column)
419 if _require_column_model(column) is not model:
420 msg = "joins are not supported in v1"
421 raise QueryConstructionError(msg)
422 for child in predicate.children:
423 _ensure_predicate_targets_model(child, model)
426def _ensure_ordering_targets_model(
427 ordering: OrderBy[Any],
428 model: type[Table[Any]],
429) -> None:
430 if ordering.column is None or ordering.direction not in {"ASC", "DESC"}:
431 msg = "orderings must be built from columns"
432 raise QueryConstructionError(msg)
433 column = _require_field(ordering.column)
434 if _require_column_model(column) is not model:
435 msg = "joins are not supported in v1"
436 raise QueryConstructionError(msg)
439def _ensure_assignment_targets_model(
440 assignment: Assignment[Any],
441 model: type[Table[Any]],
442) -> None:
443 if assignment.column is None:
444 msg = "assignments must be built from columns"
445 raise QueryConstructionError(msg)
446 column = _require_field(assignment.column)
447 if _require_column_model(column) is not model:
448 msg = "joins are not supported in v1"
449 raise QueryConstructionError(msg)
450 if column.is_generated or column.primary_key:
451 msg = "generated and primary key columns cannot update"
452 raise QueryConstructionError(msg)
455def _require_insert_model(row: object) -> type[Table[Any]]:
456 if not isinstance(row, Model):
457 msg = "insert requires a snekql model instance"
458 raise QueryConstructionError(msg)
459 model_row = cast("Model[Any, Any]", row)
460 return cast("type[Table[Any]]", model_row.__class__)
463def _encode_insert_row(
464 query: InsertQuery[Any],
465 dialect: QueryDialect,
466) -> tuple[type[Table[Any]], dict[str, object]]:
467 model_class = _require_insert_model(query.row)
468 row_values: dict[str, object] = {}
469 for name, column in require_model_columns(model_class).items():
470 value = getattr(query.row, name)
471 if value is MISSING:
472 continue
473 row_values[name] = dialect.encode_column_value(column, value)
474 return model_class, row_values
477def _compile_insert_sql(
478 query: InsertQuery[Any],
479 dialect: QueryDialect,
480) -> tuple[str, tuple[object, ...]]:
481 model_class, row_values = _encode_insert_row(query, dialect)
482 table_name = require_model_table_name(model_class)
483 quoted_table = dialect.quote_identifier(table_name)
484 if not row_values:
485 return dialect.empty_insert_sql(quoted_table), ()
486 names = tuple(row_values)
487 quoted_columns = ", ".join(dialect.quote_identifier(name) for name in names)
488 placeholders = ", ".join(dialect.placeholder for _ in names)
489 sql = "INSERT INTO " + quoted_table + f" ({quoted_columns}) VALUES ({placeholders})" # noqa: S608
490 params = tuple(row_values[name] for name in names)
491 return sql, params
494def _compile_compound_predicate_sql(
495 predicate: Predicate[Any],
496 model: type[Table[Any]],
497 dialect: QueryDialect,
498) -> tuple[str, tuple[object, ...]]:
499 if len(predicate.children) != _BINARY_PREDICATE_CHILD_COUNT:
500 msg = "compound predicate is malformed"
501 raise QueryCompilationError(msg)
502 left_sql, left_params = _compile_predicate_sql(
503 predicate.children[0],
504 model,
505 dialect,
506 )
507 right_sql, right_params = _compile_predicate_sql(
508 predicate.children[1],
509 model,
510 dialect,
511 )
512 operator = "AND" if predicate.kind == "and" else "OR"
513 return f"({left_sql}) {operator} ({right_sql})", (*left_params, *right_params)
516def _compile_negated_predicate_sql(
517 predicate: Predicate[Any],
518 model: type[Table[Any]],
519 dialect: QueryDialect,
520) -> tuple[str, tuple[object, ...]]:
521 if len(predicate.children) != _UNARY_PREDICATE_CHILD_COUNT:
522 msg = "negated predicate is malformed"
523 raise QueryCompilationError(msg)
524 child_sql, child_params = _compile_predicate_sql(
525 predicate.children[0],
526 model,
527 dialect,
528 )
529 return f"NOT ({child_sql})", child_params
532def _compile_equality_predicate_sql(
533 predicate: Predicate[Any],
534 column: Attr[Any, Any, Any, Any, Any],
535 column_name: str,
536 dialect: QueryDialect,
537) -> tuple[str, tuple[object, ...]]:
538 if predicate.value is None:
539 msg = f"{predicate.kind}(None) is invalid; use is_not_null()"
540 if predicate.kind == "eq":
541 msg = "eq(None) is invalid; use is_null()"
542 raise QueryCompilationError(msg)
543 operator = "=" if predicate.kind == "eq" else "!="
544 return (
545 f"{column_name} {operator} {dialect.placeholder}",
546 (dialect.encode_column_value(column, predicate.value),),
547 )
550def _compile_membership_predicate_sql(
551 predicate: Predicate[Any],
552 column: Attr[Any, Any, Any, Any, Any],
553 column_name: str,
554 dialect: QueryDialect,
555) -> tuple[str, tuple[object, ...]]:
556 if not predicate.values:
557 msg = "IN predicates require at least one value"
558 raise QueryCompilationError(msg)
559 if any(value is None for value in predicate.values):
560 msg = "IN predicate values cannot be None"
561 raise QueryCompilationError(msg)
562 placeholders = ", ".join(dialect.placeholder for _ in predicate.values)
563 operator = "IN" if predicate.kind == "in" else "NOT IN"
564 params = tuple(
565 dialect.encode_column_value(column, value) for value in predicate.values
566 )
567 return f"{column_name} {operator} ({placeholders})", params
570def _compile_like_predicate_sql(
571 predicate: Predicate[Any],
572 column: Attr[Any, Any, Any, Any, Any],
573 column_name: str,
574 dialect: QueryDialect,
575) -> tuple[str, tuple[object, ...]]:
576 if column.storage_type_name != "Text":
577 msg = f"{predicate.kind}() is only valid for text columns"
578 raise QueryCompilationError(msg)
579 operator = "LIKE" if predicate.kind == "like" else "NOT LIKE"
580 return (
581 f"{column_name} {operator} {dialect.placeholder}",
582 (dialect.encode_column_value(column, predicate.value),),
583 )
586def _compile_column_predicate_sql(
587 predicate: Predicate[Any],
588 dialect: QueryDialect,
589) -> tuple[str, tuple[object, ...]]:
590 column = _require_field(predicate.column)
591 column_name = dialect.quote_identifier(_require_column_name(column))
592 if predicate.kind in {"eq", "ne"}:
593 return _compile_equality_predicate_sql(
594 predicate,
595 column,
596 column_name,
597 dialect,
598 )
599 if predicate.kind == "is_null":
600 return f"{column_name} IS NULL", ()
601 if predicate.kind == "is_not_null":
602 return f"{column_name} IS NOT NULL", ()
603 if predicate.kind in {"in", "not_in"}:
604 return _compile_membership_predicate_sql(
605 predicate,
606 column,
607 column_name,
608 dialect,
609 )
610 if predicate.kind in {"like", "not_like"}:
611 return _compile_like_predicate_sql(predicate, column, column_name, dialect)
612 msg = "unknown predicate kind"
613 raise QueryCompilationError(msg)
616def _compile_predicate_sql(
617 predicate: Predicate[Any],
618 model: type[Table[Any]],
619 dialect: QueryDialect,
620) -> tuple[str, tuple[object, ...]]:
621 _ensure_predicate_targets_model(predicate, model)
622 if predicate.kind in {"and", "or"}:
623 return _compile_compound_predicate_sql(predicate, model, dialect)
624 if predicate.kind == "not":
625 return _compile_negated_predicate_sql(predicate, model, dialect)
626 return _compile_column_predicate_sql(predicate, dialect)
629def _compile_ordering_sql(
630 ordering: OrderBy[Any],
631 model: type[Table[Any]],
632 dialect: QueryDialect,
633) -> str:
634 _ensure_ordering_targets_model(ordering, model)
635 column = _require_field(ordering.column)
636 column_name = dialect.quote_identifier(_require_column_name(column))
637 return f"{column_name} {ordering.direction}"
640def _compile_predicates_sql(
641 predicates: tuple[Predicate[Any], ...],
642 model: type[Table[Any]],
643 dialect: QueryDialect,
644) -> tuple[str, tuple[object, ...]]:
645 predicate_sql_parts: list[str] = []
646 predicate_params: list[object] = []
647 for predicate in predicates:
648 predicate_sql, compiled_params = _compile_predicate_sql(
649 predicate,
650 model,
651 dialect,
652 )
653 predicate_sql_parts.append(f"({predicate_sql})")
654 predicate_params.extend(compiled_params)
655 return " AND ".join(predicate_sql_parts), tuple(predicate_params)
658def _compile_update_sql(
659 query: UpdateQuery[Any],
660 dialect: QueryDialect,
661) -> tuple[str, tuple[object, ...]]:
662 state = query.state
663 if not state.assignments:
664 msg = "update requires set() before execution"
665 raise QueryCompilationError(msg)
666 if not state.explicit_all and not state.predicates:
667 msg = "update requires all() or where() before execution"
668 raise QueryCompilationError(msg)
669 table_name = require_model_table_name(state.model)
670 set_sql_parts: list[str] = []
671 params: tuple[object, ...] = ()
672 for assignment in state.assignments:
673 _ensure_assignment_targets_model(assignment, state.model)
674 column = _require_field(assignment.column)
675 column_name = dialect.quote_identifier(_require_column_name(column))
676 set_sql_parts.append(f"{column_name} = {dialect.placeholder}")
677 params = (*params, dialect.encode_column_value(column, assignment.value))
678 sql_parts = [
679 "UPDATE " + dialect.quote_identifier(table_name) + " SET ", # noqa: S608
680 ", ".join(set_sql_parts),
681 ]
682 if state.predicates:
683 predicate_sql, predicate_params = _compile_predicates_sql(
684 state.predicates,
685 state.model,
686 dialect,
687 )
688 sql_parts.append(f" WHERE {predicate_sql}")
689 params = (*params, *predicate_params)
690 return "".join(sql_parts), params
693def _compile_delete_sql(
694 query: DeleteQuery[Any],
695 dialect: QueryDialect,
696) -> tuple[str, tuple[object, ...]]:
697 state = query.state
698 if not state.explicit_all and not state.predicates:
699 msg = "delete requires all() or where() before execution"
700 raise QueryCompilationError(msg)
701 table_name = require_model_table_name(state.model)
702 sql = "DELETE FROM " + dialect.quote_identifier(table_name) # noqa: S608
703 params: tuple[object, ...] = ()
704 if state.predicates:
705 predicate_sql, params = _compile_predicates_sql(
706 state.predicates,
707 state.model,
708 dialect,
709 )
710 sql = f"{sql} WHERE {predicate_sql}"
711 return sql, params
714def _compile_select_state(
715 state: _SelectState,
716 dialect: QueryDialect,
717) -> tuple[str, tuple[object, ...]]:
718 if not state.explicit_all and not state.predicates:
719 msg = "select requires all() or where() before execution"
720 raise QueryCompilationError(msg)
721 table_name = require_model_table_name(state.model)
722 quoted_columns = ", ".join(
723 dialect.quote_identifier(_require_column_name(column))
724 for column in state.fields
725 )
726 sql_parts = [
727 "SELECT " + quoted_columns + " FROM " + dialect.quote_identifier(table_name), # noqa: S608
728 ]
729 params: tuple[object, ...] = ()
730 if state.predicates:
731 predicate_sql, predicate_params = _compile_predicates_sql(
732 state.predicates,
733 state.model,
734 dialect,
735 )
736 sql_parts.append(f"WHERE {predicate_sql}")
737 params = (*params, *predicate_params)
738 if state.orderings:
739 order_by = ", ".join(
740 _compile_ordering_sql(ordering, state.model, dialect)
741 for ordering in state.orderings
742 )
743 sql_parts.append(f"ORDER BY {order_by}")
744 if state.limit_value is not None:
745 sql_parts.append(f"LIMIT {dialect.placeholder}")
746 params = (*params, state.limit_value)
747 if state.offset_value is not None:
748 if state.limit_value is None:
749 sql_parts.append("LIMIT -1")
750 sql_parts.append(f"OFFSET {dialect.placeholder}")
751 params = (*params, state.offset_value)
752 return " ".join(sql_parts), params
755def compile_select_sql_for_dialect(
756 query: AnySelectQuery,
757 dialect: QueryDialect,
758) -> tuple[str, tuple[object, ...]]:
759 """Compile a select query into backend Dialect SQL."""
761 return _compile_select_state(query.state, dialect)
764def compile_write_sql_for_dialect(
765 query: object,
766 dialect: QueryDialect,
767) -> tuple[str, tuple[object, ...]]:
768 """Compile a write query into backend Dialect SQL."""
770 if isinstance(query, InsertQuery):
771 return _compile_insert_sql(cast("InsertQuery[Any]", query), dialect)
772 if isinstance(query, UpdateQuery):
773 return _compile_update_sql(cast("UpdateQuery[Any]", query), dialect)
774 if isinstance(query, DeleteQuery):
775 return _compile_delete_sql(cast("DeleteQuery[Any]", query), dialect)
776 msg = "execute requires a write query"
777 raise QueryCompilationError(msg)
780def compile_select_sql(query: AnySelectQuery) -> tuple[str, tuple[object, ...]]:
781 """Compile a select query into parameterized SQLite SQL."""
783 return compile_select_sql_for_dialect(query, _SQLITE_QUERY_DIALECT)
786def compile_write_sql(query: object) -> tuple[str, tuple[object, ...]]:
787 """Compile a write query into parameterized SQLite SQL."""
789 return compile_write_sql_for_dialect(query, _SQLITE_QUERY_DIALECT)
792def materialize_select_row(
793 query: AnySelectQuery,
794 row: Sequence[object],
795) -> object:
796 """Decode one SQLite result row according to a select query."""
798 state = query.state
799 if len(row) != len(state.fields):
800 msg = "database row shape did not match select query"
801 raise QueryCompilationError(msg)
802 if state.returns_model:
803 values = {
804 _require_column_name(column): row[index]
805 for index, column in enumerate(state.fields)
806 }
807 return decode_model_row(state.model, values)
808 decoded_values = tuple(
809 decode_column_value(column, row[index], backend="sqlite")
810 for index, column in enumerate(state.fields)
811 )
812 if len(decoded_values) == 1:
813 return decoded_values[0]
814 return decoded_values
817@overload
818def select[SelectOwnerT: Table[Any], ReadModelT: Table[Any]](
819 model: _SelectableModelClass[SelectOwnerT, ReadModelT],
820 /,
821) -> SelectModelQuery[SelectOwnerT, ReadModelT]: ...
824@overload
825def select[OwnerT: Table[Any], T1](
826 field1: Attr[Any, Any, OwnerT, Any, T1],
827 /,
828) -> SelectValueQuery[OwnerT, T1]: ...
831@overload
832def select[OwnerT: Table[Any], T1, T2](
833 field1: Attr[Any, Any, OwnerT, Any, T1],
834 field2: Attr[Any, Any, OwnerT, Any, T2],
835 /,
836) -> SelectTupleQuery[OwnerT, T1, T2]: ...
839@overload
840def select[OwnerT: Table[Any], T1, T2, T3](
841 field1: Attr[Any, Any, OwnerT, Any, T1],
842 field2: Attr[Any, Any, OwnerT, Any, T2],
843 field3: Attr[Any, Any, OwnerT, Any, T3],
844 /,
845) -> SelectTupleQuery[OwnerT, T1, T2, T3]: ...
848def select(*args: object) -> object:
849 if len(args) == 0:
850 msg = "select requires a model or field"
851 raise QueryConstructionError(msg)
852 if any(isinstance(argument, type) for argument in args):
853 if len(args) != 1 or not isinstance(args[0], type):
854 msg = "mixed model and field selection is invalid"
855 raise QueryConstructionError(msg)
856 model = cast("type[Table[Any]]", args[0])
857 try:
858 columns = require_model_columns(model)
859 except ModelDeclarationError as error:
860 msg = "select requires a table model"
861 raise QueryConstructionError(msg) from error
862 state = _SelectState(
863 model=model,
864 fields=tuple(columns.values()),
865 returns_model=True,
866 )
867 return SelectModelQuery[Any, Any](state)
868 fields = tuple(_require_field(argument) for argument in args)
869 model = _require_column_model(fields[0])
870 for field in fields[1:]:
871 if _require_column_model(field) is not model:
872 msg = "joins are not supported in v1"
873 raise QueryConstructionError(msg)
874 state = _SelectState(model=model, fields=fields)
875 if len(fields) == 1:
876 return SelectValueQuery[Any, Any](state)
877 return SelectTupleQuery[Any, *tuple[Any, ...]](state)
880def insert[ModelT: Table[Any]](row: ModelT, /) -> InsertQuery[ModelT]:
881 return InsertQuery(row)
884def update[ModelT: Table[Any]](model: type[ModelT], /) -> UpdateQuery[ModelT]:
885 try:
886 _ = require_model_columns(model)
887 except ModelDeclarationError as error:
888 msg = "update requires a table model"
889 raise QueryConstructionError(msg) from error
890 return UpdateQuery(_UpdateState(model=cast("type[Table[Any]]", model)))
893def delete[ModelT: Table[Any]](model: type[ModelT], /) -> DeleteQuery[ModelT]:
894 try:
895 _ = require_model_columns(model)
896 except ModelDeclarationError as error:
897 msg = "delete requires a table model"
898 raise QueryConstructionError(msg) from error
899 return DeleteQuery(_DeleteState(model=cast("type[Table[Any]]", model)))