Coverage for snekql/mariadb/query.py: 90%

30 statements  

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

1"""MariaDB SQL compilation and row materialization.""" 

2 

3from __future__ import annotations 

4 

5from collections.abc import Sequence 

6from typing import Any 

7 

8from snekql._model_materialization import ( 

9 decode_column_value, 

10 decode_model_row, 

11 encode_column_value, 

12) 

13from snekql._query_dialect import QueryDialect 

14from snekql.errors import QueryCompilationError 

15from snekql.mariadb.identifiers import quote_identifier as quote_mariadb_identifier 

16from snekql.query import ( 

17 AnySelectQuery, 

18 compile_select_sql_for_dialect, 

19 compile_write_sql_for_dialect, 

20) 

21from snekql.storage import Attr 

22 

23 

24def _mariadb_empty_insert_sql(quoted_table: str) -> str: 

25 return f"INSERT INTO {quoted_table} () VALUES ()" # noqa: S608 

26 

27 

28def _encode_mariadb_column_value( 

29 column: Attr[Any, Any, Any, Any, Any], 

30 value: object, 

31) -> object: 

32 return encode_column_value(column, value, backend="mariadb") 

33 

34 

35_MARIADB_QUERY_DIALECT = QueryDialect( 

36 empty_insert_sql=_mariadb_empty_insert_sql, 

37 encode_column_value=_encode_mariadb_column_value, 

38 placeholder="%s", 

39 quote_identifier=quote_mariadb_identifier, 

40) 

41 

42 

43def compile_mariadb_select_sql( 

44 query: AnySelectQuery, 

45) -> tuple[str, tuple[object, ...]]: 

46 """Compile a select query into parameterized MariaDB SQL.""" 

47 

48 return compile_select_sql_for_dialect(query, _MARIADB_QUERY_DIALECT) 

49 

50 

51def compile_mariadb_write_sql(query: object) -> tuple[str, tuple[object, ...]]: 

52 """Compile a write query into parameterized MariaDB SQL.""" 

53 

54 return compile_write_sql_for_dialect(query, _MARIADB_QUERY_DIALECT) 

55 

56 

57def materialize_mariadb_select_row( 

58 query: AnySelectQuery, 

59 row: Sequence[object], 

60) -> object: 

61 """Decode one MariaDB result row according to a select query.""" 

62 

63 state = query.state 

64 if len(row) != len(state.fields): 

65 msg = "database row shape did not match select query" 

66 raise QueryCompilationError(msg) 

67 if state.returns_model: 

68 values = { 

69 column.name or "": row[index] for index, column in enumerate(state.fields) 

70 } 

71 return decode_model_row(state.model, values, backend="mariadb") 

72 decoded_values = tuple( 

73 decode_column_value(column, row[index], backend="mariadb") 

74 for index, column in enumerate(state.fields) 

75 ) 

76 if len(decoded_values) == 1: 

77 return decoded_values[0] 

78 return decoded_values