Coverage for snekql/query.py: 86%

515 statements  

« prev     ^ index     » next       coverage.py v7.14.1, created at 2026-06-07 21:13 +0300

1"""Query Builder objects and factory functions.""" 

2 

3from __future__ import annotations 

4 

5from collections.abc import Sequence 

6from dataclasses import dataclass, replace 

7from typing import Any, Protocol, Self, TypeVar, TypeVarTuple, cast, overload 

8 

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 

30 

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") 

42 

43_BINARY_PREDICATE_CHILD_COUNT = 2 

44_UNARY_PREDICATE_CHILD_COUNT = 1 

45 

46 

47def _sqlite_empty_insert_sql(quoted_table: str) -> str: 

48 return "INSERT INTO " + quoted_table + " DEFAULT VALUES" 

49 

50 

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") 

56 

57 

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) 

64 

65 

66class _SelectableModelClass(Protocol[SelectableOwnerT_co, SelectableReadT_co]): 

67 """Structural type for model classes accepted by `select(Model)`. 

68 

69 The protocol lets pyright connect the writable owner model type with the 

70 fetched read model type exposed by table model classes. 

71 """ 

72 

73 @classmethod 

74 def __owner_type__(cls) -> type[SelectableOwnerT_co]: ... 

75 

76 @classmethod 

77 def __read_type__(cls) -> type[SelectableReadT_co]: ... 

78 

79 

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 

90 

91 

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], ...] = () 

98 

99 

100@dataclass(frozen=True) 

101class _DeleteState: 

102 model: type[Table[Any]] 

103 explicit_all: bool = False 

104 predicates: tuple[Predicate[Any], ...] = () 

105 

106 

107class SelectModelQuery[SelectOwnerT: Table[Any], ReadModelT: Table[Any]]: 

108 """Immutable select query that returns fetched table model instances.""" 

109 

110 state: _SelectState 

111 

112 def __init__(self, state: _SelectState | None = None) -> None: 

113 if state is None: 

114 state = _empty_select_state() 

115 self.state = state 

116 

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)) 

122 

123 def where(self, *predicates: Predicate[SelectOwnerT]) -> Self: 

124 state = _select_where(self.state, predicates) 

125 return cast("Self", SelectModelQuery[SelectOwnerT, ReadModelT](state)) 

126 

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)) 

130 

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)) 

135 

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)) 

140 

141 

142class SelectValueQuery[OwnerT: Table[Any], T]: 

143 """Immutable select query that returns one scalar column value per row.""" 

144 

145 state: _SelectState 

146 

147 def __init__(self, state: _SelectState | None = None) -> None: 

148 if state is None: 

149 state = _empty_select_state() 

150 self.state = state 

151 

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)) 

157 

158 def where(self, *predicates: Predicate[OwnerT]) -> Self: 

159 state = _select_where(self.state, predicates) 

160 return cast("Self", SelectValueQuery[OwnerT, T](state)) 

161 

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)) 

165 

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)) 

170 

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)) 

175 

176 

177class SelectTupleQuery[OwnerT: Table[Any], *Ts]: 

178 """Immutable select query that returns selected column tuples per row.""" 

179 

180 state: _SelectState 

181 

182 def __init__(self, state: _SelectState | None = None) -> None: 

183 if state is None: 

184 state = _empty_select_state() 

185 self.state = state 

186 

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)) 

192 

193 def where(self, *predicates: Predicate[OwnerT]) -> Self: 

194 state = _select_where(self.state, predicates) 

195 return cast("Self", SelectTupleQuery[OwnerT, *Ts](state)) 

196 

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)) 

200 

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)) 

205 

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)) 

210 

211 

212class InsertQuery[ModelT: Table[Any]]: 

213 """Immutable insert statement for one pending table model instance.""" 

214 

215 row: ModelT 

216 

217 def __init__(self, row: ModelT) -> None: 

218 self.row: ModelT = row 

219 

220 

221class UpdateQuery[ModelT: Table[Any]]: 

222 """Immutable update statement for one table model.""" 

223 

224 state: _UpdateState 

225 

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 

230 

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)) 

236 

237 def set(self, *assignments: Assignment[ModelT]) -> Self: 

238 state = _update_set(self.state, assignments) 

239 return cast("Self", UpdateQuery[ModelT](state)) 

240 

241 def where(self, *predicates: Predicate[ModelT]) -> Self: 

242 state = _update_where(self.state, predicates) 

243 return cast("Self", UpdateQuery[ModelT](state)) 

244 

245 

246class DeleteQuery[ModelT: Table[Any]]: 

247 """Immutable delete statement for one table model.""" 

248 

249 state: _DeleteState 

250 

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 

255 

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)) 

261 

262 def where(self, *predicates: Predicate[ModelT]) -> Self: 

263 state = _delete_where(self.state, predicates) 

264 return cast("Self", DeleteQuery[ModelT](state)) 

265 

266 

267type AnySelectQuery = ( 

268 SelectModelQuery[Any, Any] 

269 | SelectValueQuery[Any, Any] 

270 | SelectTupleQuery[Any, *tuple[Any, ...]] 

271) 

272 

273 

274def _empty_select_state() -> _SelectState: 

275 return _SelectState(model=Table[Any], fields=()) 

276 

277 

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) 

285 

286 

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)) 

300 

301 

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)) 

312 

313 

314def _select_limit(state: _SelectState, value: NonNegativeInt) -> _SelectState: 

315 return replace(state, limit_value=value) 

316 

317 

318def _select_offset(state: _SelectState, value: NonNegativeInt) -> _SelectState: 

319 return replace(state, offset_value=value) 

320 

321 

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) 

329 

330 

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)) 

341 

342 

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)) 

356 

357 

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) 

365 

366 

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)) 

380 

381 

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) 

387 

388 

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 

394 

395 

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 

408 

409 

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) 

424 

425 

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) 

437 

438 

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) 

453 

454 

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__) 

461 

462 

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 

475 

476 

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 

492 

493 

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) 

514 

515 

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 

530 

531 

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 ) 

548 

549 

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 

568 

569 

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 ) 

584 

585 

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) 

614 

615 

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) 

627 

628 

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}" 

638 

639 

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) 

656 

657 

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 

691 

692 

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 

712 

713 

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 

753 

754 

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.""" 

760 

761 return _compile_select_state(query.state, dialect) 

762 

763 

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.""" 

769 

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) 

778 

779 

780def compile_select_sql(query: AnySelectQuery) -> tuple[str, tuple[object, ...]]: 

781 """Compile a select query into parameterized SQLite SQL.""" 

782 

783 return compile_select_sql_for_dialect(query, _SQLITE_QUERY_DIALECT) 

784 

785 

786def compile_write_sql(query: object) -> tuple[str, tuple[object, ...]]: 

787 """Compile a write query into parameterized SQLite SQL.""" 

788 

789 return compile_write_sql_for_dialect(query, _SQLITE_QUERY_DIALECT) 

790 

791 

792def materialize_select_row( 

793 query: AnySelectQuery, 

794 row: Sequence[object], 

795) -> object: 

796 """Decode one SQLite result row according to a select query.""" 

797 

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 

815 

816 

817@overload 

818def select[SelectOwnerT: Table[Any], ReadModelT: Table[Any]]( 

819 model: _SelectableModelClass[SelectOwnerT, ReadModelT], 

820 /, 

821) -> SelectModelQuery[SelectOwnerT, ReadModelT]: ... 

822 

823 

824@overload 

825def select[OwnerT: Table[Any], T1]( 

826 field1: Attr[Any, Any, OwnerT, Any, T1], 

827 /, 

828) -> SelectValueQuery[OwnerT, T1]: ... 

829 

830 

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]: ... 

837 

838 

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]: ... 

846 

847 

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) 

878 

879 

880def insert[ModelT: Table[Any]](row: ModelT, /) -> InsertQuery[ModelT]: 

881 return InsertQuery(row) 

882 

883 

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))) 

891 

892 

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)))