sqlengine-lite 2.1.1__py3-none-any.whl → 2.2.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- sqlengine/__init__.py +3 -3
- sqlengine/core/__init__.py +4 -0
- sqlengine/core/connection.py +255 -0
- sqlengine/{utils → core}/repr.py +1 -30
- sqlengine/{utils → core}/sqlgen.py +2 -4
- sqlengine/{utils → core}/statements.py +36 -36
- sqlengine/{utils → core}/types.py +47 -8
- sqlengine/schema.py +1 -1
- sqlengine/sqltable.py +162 -276
- sqlengine/utils/__init__.py +2 -3
- sqlengine/utils/connection.py +19 -13
- sqlengine/utils/convert.py +24 -0
- {sqlengine_lite-2.1.1.dist-info → sqlengine_lite-2.2.0.dist-info}/METADATA +53 -9
- sqlengine_lite-2.2.0.dist-info/RECORD +17 -0
- sqlengine_lite-2.1.1.dist-info/RECORD +0 -14
- {sqlengine_lite-2.1.1.dist-info → sqlengine_lite-2.2.0.dist-info}/WHEEL +0 -0
- {sqlengine_lite-2.1.1.dist-info → sqlengine_lite-2.2.0.dist-info}/licenses/LICENSE +0 -0
- {sqlengine_lite-2.1.1.dist-info → sqlengine_lite-2.2.0.dist-info}/top_level.txt +0 -0
sqlengine/__init__.py
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
|
-
from .
|
|
2
|
-
from .
|
|
1
|
+
from .core import sqlgen, types
|
|
2
|
+
from .core.types import Schema, Primary
|
|
3
3
|
from .sqltable import SqlTableMixin
|
|
4
4
|
|
|
5
5
|
__author__ = "suffermuffin"
|
|
6
6
|
|
|
7
|
-
__all__ = ["types", "sqlgen", "Schema", "SqlTableMixin"]
|
|
7
|
+
__all__ = ["types", "sqlgen", "Schema", "SqlTableMixin", "Primary"]
|
|
@@ -0,0 +1,255 @@
|
|
|
1
|
+
import sqlite3
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
|
|
5
|
+
from typing import overload, Literal, Sequence
|
|
6
|
+
from contextlib import contextmanager
|
|
7
|
+
|
|
8
|
+
from .types import SqlValue, SqlRow
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
logger = logging.getLogger("sqlengine")
|
|
12
|
+
logger.setLevel(os.getenv("SQL_ENGINE_LOG_LEVEL", "WARNING").upper())
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class ConnectionManager:
|
|
16
|
+
|
|
17
|
+
_trans : sqlite3.Connection
|
|
18
|
+
_trans_cursor : sqlite3.Cursor
|
|
19
|
+
|
|
20
|
+
def __init__(self, database : str, **connection_params) -> None:
|
|
21
|
+
|
|
22
|
+
self.database = database
|
|
23
|
+
self.connection_params = connection_params
|
|
24
|
+
|
|
25
|
+
self._is_managed_transaction = False
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@overload
|
|
29
|
+
def _fetch(self, query : str, args : SqlRow,
|
|
30
|
+
method : Literal["fetchone"]) -> SqlRow: ...
|
|
31
|
+
@overload
|
|
32
|
+
def _fetch(self, query : str, args : SqlRow,
|
|
33
|
+
method : Literal["fetchall"]) -> list[SqlRow]: ...
|
|
34
|
+
|
|
35
|
+
def _fetch(
|
|
36
|
+
self,
|
|
37
|
+
query : str,
|
|
38
|
+
args : SqlRow = (),
|
|
39
|
+
method : Literal["fetchone", "fetchall"] = "fetchall"
|
|
40
|
+
) -> SqlRow | list[SqlRow]:
|
|
41
|
+
|
|
42
|
+
logger.debug(f"{self.database}: {query} {args}")
|
|
43
|
+
|
|
44
|
+
if self.in_transaction():
|
|
45
|
+
self._trans_cursor.execute(query, args)
|
|
46
|
+
return getattr(self._trans_cursor, method)()
|
|
47
|
+
|
|
48
|
+
with self.connect() as conn:
|
|
49
|
+
cursor = conn.cursor()
|
|
50
|
+
cursor.execute(query, args)
|
|
51
|
+
return getattr(cursor, method)()
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@overload
|
|
55
|
+
def _execute(self, query : str, args : tuple[SqlValue, ...], method : Literal["execute"]) -> None: ...
|
|
56
|
+
@overload
|
|
57
|
+
def _execute(self, query : str, args : Sequence[SqlRow], method : Literal["executemany"]) -> None: ...
|
|
58
|
+
|
|
59
|
+
def _execute(
|
|
60
|
+
self,
|
|
61
|
+
query : str,
|
|
62
|
+
args : tuple[SqlValue, ...] | Sequence[SqlRow] = (),
|
|
63
|
+
method : Literal["execute", "executemany"] = "execute"
|
|
64
|
+
) -> None:
|
|
65
|
+
"""
|
|
66
|
+
Shortcut to connect() -> execute[<many>]() -> commit() for single operations.
|
|
67
|
+
Can be used in transaction using `transaction()` manager.
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
query (str): SQL query to execute on SQLite3 DB
|
|
71
|
+
*args (Any): Arguments to the execution
|
|
72
|
+
method (str): "execute" or "executemany"
|
|
73
|
+
"""
|
|
74
|
+
logger.debug(f"{self.database}: {query} {args}")
|
|
75
|
+
|
|
76
|
+
if self.in_transaction():
|
|
77
|
+
getattr(self._trans_cursor, method)(query, args)
|
|
78
|
+
return
|
|
79
|
+
|
|
80
|
+
with self.connect() as conn:
|
|
81
|
+
cursor = conn.cursor()
|
|
82
|
+
getattr(cursor, method)(query, args)
|
|
83
|
+
conn.commit()
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def execute(self, query : str, *args : SqlValue) -> None:
|
|
87
|
+
"""
|
|
88
|
+
Shortcut to connect() -> execute() -> commit() for single operations.
|
|
89
|
+
Can be used in transaction using `transaction()` manager.
|
|
90
|
+
|
|
91
|
+
Args:
|
|
92
|
+
query (str): SQL query to execute on SQLite3 DB
|
|
93
|
+
*args (tuple[SqlValue, ...]): Arguments to the execution
|
|
94
|
+
"""
|
|
95
|
+
return self._execute(query, args, method="execute")
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def executemany(self, query : str, args : Sequence[SqlRow]) -> None:
|
|
99
|
+
"""
|
|
100
|
+
Shortcut to connect() -> executemany() -> commit() for single operations.
|
|
101
|
+
Can be used in transaction using `transaction()` manager.
|
|
102
|
+
|
|
103
|
+
Args:
|
|
104
|
+
query (str): SQL query to execute on SQLite3 DB
|
|
105
|
+
args (list[tuple[SqlValue, ...]]): Arguments to the execution
|
|
106
|
+
"""
|
|
107
|
+
return self._execute(query, args, method="executemany")
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def fetchone(self, query : str, *args : SqlValue) -> SqlRow:
|
|
111
|
+
"""
|
|
112
|
+
Fetch first row based on `query`
|
|
113
|
+
|
|
114
|
+
Args:
|
|
115
|
+
query (str): SQL query
|
|
116
|
+
*args (tuple[SqlValue, ...]): Arguments to the execution
|
|
117
|
+
|
|
118
|
+
Returns:
|
|
119
|
+
row (SqlRow): Single row
|
|
120
|
+
"""
|
|
121
|
+
return self._fetch(query, args, method="fetchone")
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
def fetchmany(self, query : str, *args : SqlValue, size : int = 1) -> list[SqlRow]:
|
|
125
|
+
"""
|
|
126
|
+
Fetch first `size` rows based on `query`
|
|
127
|
+
|
|
128
|
+
Args:
|
|
129
|
+
query (str): SQL query
|
|
130
|
+
*args (tuple[SqlValue, ...]): Arguments to the execution
|
|
131
|
+
size (str): Number of rows to return
|
|
132
|
+
|
|
133
|
+
Returns:
|
|
134
|
+
rows (list[SqlRow]): list of `size` rows
|
|
135
|
+
"""
|
|
136
|
+
logger.debug(f"{self.database}: {query} {args}")
|
|
137
|
+
|
|
138
|
+
if self.in_transaction():
|
|
139
|
+
self._trans_cursor.execute(query, args)
|
|
140
|
+
return self._trans_cursor.fetchmany(size)
|
|
141
|
+
|
|
142
|
+
with self.connect() as conn:
|
|
143
|
+
cursor = conn.cursor()
|
|
144
|
+
cursor.execute(query, args)
|
|
145
|
+
return cursor.fetchmany(size)
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def fetchall(self, query : str, *args : SqlValue) -> list[SqlRow]:
|
|
149
|
+
"""
|
|
150
|
+
Fetch all rows based on `query`
|
|
151
|
+
|
|
152
|
+
Args:
|
|
153
|
+
query (str): SQL query
|
|
154
|
+
*args (tuple[SqlValue, ...]): Arguments to the execution
|
|
155
|
+
|
|
156
|
+
Returns:
|
|
157
|
+
rows (list[SqlRow]): list of rows
|
|
158
|
+
"""
|
|
159
|
+
return self._fetch(query, args, method="fetchall")
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def in_transaction(self) -> bool:
|
|
163
|
+
""" Returns True if instance is in transaction """
|
|
164
|
+
return hasattr(self, "_trans") and hasattr(self, "_trans_cursor")
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def connect(self) -> sqlite3.Connection:
|
|
168
|
+
""" Shortcut to sqlite3 connection context manager """
|
|
169
|
+
return sqlite3.connect(self.database, **self.connection_params)
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def open(self) -> None:
|
|
173
|
+
""" Opens unmanaged transaction """
|
|
174
|
+
if self.in_transaction():
|
|
175
|
+
raise RuntimeError("Can't re-open existing connection")
|
|
176
|
+
|
|
177
|
+
self._trans = self.connect()
|
|
178
|
+
self._trans_cursor = self._trans.cursor()
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def close(self) -> None:
|
|
182
|
+
""" Closes unmanaged transaction """
|
|
183
|
+
if not self.in_transaction():
|
|
184
|
+
return
|
|
185
|
+
|
|
186
|
+
if self._is_managed_transaction:
|
|
187
|
+
raise RuntimeError("Can't manually close managed transaction")
|
|
188
|
+
|
|
189
|
+
self._trans_cursor.close()
|
|
190
|
+
self._trans.close()
|
|
191
|
+
del(self._trans_cursor)
|
|
192
|
+
del(self._trans)
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
def commit(self) -> None:
|
|
196
|
+
if not self.in_transaction():
|
|
197
|
+
raise RuntimeError("Can't commit outside transaction mode")
|
|
198
|
+
|
|
199
|
+
self._trans.commit()
|
|
200
|
+
|
|
201
|
+
|
|
202
|
+
def rollback(self) -> None:
|
|
203
|
+
if not self.in_transaction():
|
|
204
|
+
raise RuntimeError("Can't rollback outside transaction mode")
|
|
205
|
+
|
|
206
|
+
self._trans.rollback()
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
@contextmanager
|
|
210
|
+
def transaction(self, autocommit : bool = True):
|
|
211
|
+
"""
|
|
212
|
+
Creates context manager to use class methods in transaction
|
|
213
|
+
|
|
214
|
+
Args:
|
|
215
|
+
autocommit (bool): If `True`, will commit changes at the end of transaction
|
|
216
|
+
"""
|
|
217
|
+
|
|
218
|
+
self.open()
|
|
219
|
+
self._is_managed_transaction = True
|
|
220
|
+
logger.debug(f"{self.database}: Transaction started")
|
|
221
|
+
|
|
222
|
+
try:
|
|
223
|
+
yield
|
|
224
|
+
|
|
225
|
+
except Exception as e:
|
|
226
|
+
logger.error(f"{self.database}: Error while in transaction: {e}")
|
|
227
|
+
logger.debug(e, exc_info=True)
|
|
228
|
+
self._trans.rollback()
|
|
229
|
+
raise e
|
|
230
|
+
|
|
231
|
+
else:
|
|
232
|
+
if autocommit:
|
|
233
|
+
self._trans.commit()
|
|
234
|
+
|
|
235
|
+
finally:
|
|
236
|
+
self._is_managed_transaction = False
|
|
237
|
+
self.close()
|
|
238
|
+
logger.debug(f"{self.database}: Transaction finished")
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
@property
|
|
242
|
+
def tx_conn(self) -> sqlite3.Connection:
|
|
243
|
+
""" Gives access to connection while in transaction """
|
|
244
|
+
if not self.in_transaction():
|
|
245
|
+
raise RuntimeError("`tx_conn` is not available outside the transaction mode")
|
|
246
|
+
return self._trans
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
@property
|
|
250
|
+
def tx_cursor(self) -> sqlite3.Cursor:
|
|
251
|
+
""" Gives access to connection cursor while in transaction """
|
|
252
|
+
if not self.in_transaction():
|
|
253
|
+
raise RuntimeError("`tx_cursor` is not available outside the transaction mode")
|
|
254
|
+
return self._trans_cursor
|
|
255
|
+
|
sqlengine/{utils → core}/repr.py
RENAMED
|
@@ -1,12 +1,5 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
1
|
from html import escape
|
|
3
|
-
from typing import Sequence
|
|
4
|
-
|
|
5
|
-
import csv
|
|
6
|
-
|
|
7
|
-
if TYPE_CHECKING:
|
|
8
|
-
from .statements import Select, Where
|
|
9
|
-
from ..sqltable import SqlTableMixin
|
|
2
|
+
from typing import Sequence
|
|
10
3
|
|
|
11
4
|
from .types import SqlRow
|
|
12
5
|
|
|
@@ -57,25 +50,3 @@ def to_html(tablename : str, columns : list[str], repr_rows : Sequence[SqlRow],
|
|
|
57
50
|
|
|
58
51
|
return "".join(html)
|
|
59
52
|
|
|
60
|
-
|
|
61
|
-
def to_csv(builder : Select | Where[Select] | SqlTableMixin, path : str) -> None:
|
|
62
|
-
|
|
63
|
-
from ..sqltable import SqlTableMixin
|
|
64
|
-
from .statements import Where
|
|
65
|
-
|
|
66
|
-
match builder:
|
|
67
|
-
case Where():
|
|
68
|
-
builder = builder.then
|
|
69
|
-
case SqlTableMixin():
|
|
70
|
-
builder = builder.select
|
|
71
|
-
|
|
72
|
-
if builder._aggregate:
|
|
73
|
-
raise AssertionError("Aggregated queries are not supported")
|
|
74
|
-
|
|
75
|
-
columns = builder._table.columns if len(builder._columns) == 0 or "*" in builder._columns else builder._columns
|
|
76
|
-
repr_rows = builder.fetchall()
|
|
77
|
-
|
|
78
|
-
with open(path, 'w', newline='') as file:
|
|
79
|
-
writer = csv.writer(file)
|
|
80
|
-
writer.writerow(columns)
|
|
81
|
-
writer.writerows(repr_rows)
|
|
@@ -4,9 +4,7 @@ from typing import Sequence
|
|
|
4
4
|
def format_list(items : Sequence | set, brackets : bool = True) -> str:
|
|
5
5
|
""" Formats list into `(item1, item2, ...)` format """
|
|
6
6
|
items_str = ', '.join([str(i) for i in items])
|
|
7
|
-
if
|
|
8
|
-
return items_str
|
|
9
|
-
return f'({items_str})'
|
|
7
|
+
return f'({items_str})' if brackets else items_str
|
|
10
8
|
|
|
11
9
|
|
|
12
10
|
def create_table(
|
|
@@ -40,7 +38,7 @@ def values_placeholder(n_values : int) -> str:
|
|
|
40
38
|
def bulk_placeholder(n_values : int, n_rows : int) -> str:
|
|
41
39
|
""" Creates placeholders `(?, ?, ..), (?, ?, ..), ...` for each row """
|
|
42
40
|
place_holder = values_placeholder(n_values)
|
|
43
|
-
return
|
|
41
|
+
return format_list([place_holder]*n_rows, False)
|
|
44
42
|
|
|
45
43
|
|
|
46
44
|
def insert(tablename : str, columns : list[str], values : str) -> str:
|
|
@@ -1,20 +1,16 @@
|
|
|
1
|
-
from
|
|
2
|
-
from
|
|
3
|
-
from abc import ABC, abstractmethod
|
|
4
|
-
|
|
5
|
-
if TYPE_CHECKING:
|
|
6
|
-
from .statements import Statement
|
|
7
|
-
from ..sqltable import SqlTableMixin
|
|
1
|
+
from typing import Sequence, Literal, Generator, Self
|
|
2
|
+
from abc import ABC, abstractmethod
|
|
8
3
|
|
|
9
4
|
from . import sqlgen as sql
|
|
10
|
-
from .
|
|
5
|
+
from .connection import ConnectionManager
|
|
6
|
+
|
|
7
|
+
from .types import SqlValue, SqlRow, Schema
|
|
11
8
|
from .repr import to_html
|
|
12
9
|
|
|
13
10
|
|
|
14
|
-
class Where[T : Statement]:
|
|
11
|
+
class Where[T : "Statement"]:
|
|
15
12
|
""" Where clause build helper """
|
|
16
13
|
|
|
17
|
-
|
|
18
14
|
def __init__(self, statement : T):
|
|
19
15
|
|
|
20
16
|
self._statement = statement
|
|
@@ -146,11 +142,10 @@ class Statement(ABC):
|
|
|
146
142
|
Statement object that helps you build queries and execute them
|
|
147
143
|
"""
|
|
148
144
|
|
|
149
|
-
|
|
145
|
+
def __init__(self, connection : ConnectionManager, tableschema : Schema) -> None:
|
|
150
146
|
|
|
151
|
-
|
|
152
|
-
|
|
153
|
-
self._table = table
|
|
147
|
+
self._tableschema = tableschema
|
|
148
|
+
self._connection = connection
|
|
154
149
|
self._where: Where[Self] = Where(self)
|
|
155
150
|
|
|
156
151
|
self._custom_query : str | None = None
|
|
@@ -175,7 +170,7 @@ class Statement(ABC):
|
|
|
175
170
|
|
|
176
171
|
|
|
177
172
|
def reset(self) -> None:
|
|
178
|
-
"""
|
|
173
|
+
""" Reset statement to reuse the object """
|
|
179
174
|
self._where.reset()
|
|
180
175
|
self._custom_query = None
|
|
181
176
|
self._custom_args = ()
|
|
@@ -209,18 +204,16 @@ class Statement(ABC):
|
|
|
209
204
|
|
|
210
205
|
|
|
211
206
|
class MutationalStatement(Statement, ABC):
|
|
212
|
-
|
|
213
207
|
|
|
214
208
|
def execute(self) -> None:
|
|
215
209
|
query, args = self.build()
|
|
216
|
-
self.
|
|
210
|
+
self._connection.execute(query, *args)
|
|
217
211
|
|
|
218
212
|
|
|
219
213
|
class Select(Statement):
|
|
220
214
|
|
|
221
|
-
|
|
222
|
-
|
|
223
|
-
super().__init__(table)
|
|
215
|
+
def __init__(self, connection : ConnectionManager, tableschema : Schema) -> None:
|
|
216
|
+
super().__init__(connection, tableschema)
|
|
224
217
|
|
|
225
218
|
self._columns : list[str] = []
|
|
226
219
|
self._order_by : list[str] = []
|
|
@@ -260,17 +253,17 @@ class Select(Statement):
|
|
|
260
253
|
|
|
261
254
|
def fetchone(self) -> SqlRow:
|
|
262
255
|
query, args = self.build()
|
|
263
|
-
return self.
|
|
256
|
+
return self._connection.fetchone(query, *args)
|
|
264
257
|
|
|
265
258
|
|
|
266
259
|
def fetchmany(self, size : int = 1) -> list[SqlRow]:
|
|
267
260
|
query, args = self.build()
|
|
268
|
-
return self.
|
|
261
|
+
return self._connection.fetchmany(query, *args, size=size)
|
|
269
262
|
|
|
270
263
|
|
|
271
264
|
def fetchall(self) -> list[SqlRow]:
|
|
272
265
|
query, args = self.build()
|
|
273
|
-
return self.
|
|
266
|
+
return self._connection.fetchall(query, *args)
|
|
274
267
|
|
|
275
268
|
|
|
276
269
|
def fetchmany_iterator(self, batch_size: int) -> Generator[list[SqlRow], None, None]:
|
|
@@ -286,13 +279,13 @@ class Select(Statement):
|
|
|
286
279
|
>>> for batch in table.select.where.gt("Age", 30).then.fetchmany_iterator(1000):
|
|
287
280
|
>>> process_batch(batch)
|
|
288
281
|
"""
|
|
289
|
-
if not self.
|
|
282
|
+
if not self._connection.in_transaction():
|
|
290
283
|
raise RuntimeError("To use the `fetchall_iterator()` method you have \
|
|
291
284
|
to keep open the transaction of the table with `transaction()` manager")
|
|
292
285
|
|
|
293
286
|
query, exec_args = self.build()
|
|
294
287
|
|
|
295
|
-
iter_cursor = self.
|
|
288
|
+
iter_cursor = self._connection.tx_conn.cursor()
|
|
296
289
|
iter_cursor.execute(query, exec_args)
|
|
297
290
|
|
|
298
291
|
while batch := iter_cursor.fetchmany(batch_size):
|
|
@@ -302,13 +295,13 @@ class Select(Statement):
|
|
|
302
295
|
def __iter__(self) -> Generator[SqlRow, None, None]:
|
|
303
296
|
""" Select statement rows iterator """
|
|
304
297
|
|
|
305
|
-
if not self.
|
|
298
|
+
if not self._connection.in_transaction():
|
|
306
299
|
raise RuntimeError("To use the __iter__ method you have \
|
|
307
300
|
to keep open the transaction of the table with `transaction()` manager")
|
|
308
301
|
|
|
309
302
|
query, exec_args = self.build()
|
|
310
303
|
|
|
311
|
-
iter_cursor = self.
|
|
304
|
+
iter_cursor = self._connection.tx_conn.cursor()
|
|
312
305
|
iter_cursor.execute(query, exec_args)
|
|
313
306
|
|
|
314
307
|
while row := iter_cursor.fetchone():
|
|
@@ -330,7 +323,7 @@ class Select(Statement):
|
|
|
330
323
|
else:
|
|
331
324
|
limit = None
|
|
332
325
|
|
|
333
|
-
query = sql.select(self.
|
|
326
|
+
query = sql.select(self._tableschema["tablename"], columns, where_clause, order, limit)
|
|
334
327
|
|
|
335
328
|
return query, args
|
|
336
329
|
|
|
@@ -342,38 +335,45 @@ class Select(Statement):
|
|
|
342
335
|
self._limit = None
|
|
343
336
|
|
|
344
337
|
|
|
338
|
+
def _resolve_columns(self) -> list[str]:
|
|
339
|
+
return (
|
|
340
|
+
self._tableschema["columns"]
|
|
341
|
+
if len(self._columns) == 0 or "*" in self._columns
|
|
342
|
+
else self._columns
|
|
343
|
+
)
|
|
344
|
+
|
|
345
|
+
|
|
345
346
|
def _repr_html_(self) -> str | None:
|
|
346
347
|
|
|
347
348
|
if self._aggregate:
|
|
348
349
|
return None
|
|
349
350
|
|
|
350
351
|
limit = 26
|
|
351
|
-
columns = self.
|
|
352
|
+
columns = self._resolve_columns()
|
|
352
353
|
repr_rows = self.fetchmany(limit)
|
|
353
354
|
|
|
354
|
-
return to_html(self.
|
|
355
|
+
return to_html(self._tableschema["tablename"], columns, repr_rows, limit=limit-1)
|
|
355
356
|
|
|
356
357
|
|
|
357
358
|
class Delete(MutationalStatement):
|
|
358
359
|
|
|
359
|
-
|
|
360
360
|
def _build(self, where_clause : str, *args : SqlValue) -> tuple[str, tuple[SqlValue, ...]]:
|
|
361
361
|
|
|
362
362
|
if not where_clause:
|
|
363
363
|
raise ValueError("Delete statement must have a where clause")
|
|
364
364
|
|
|
365
|
-
query = sql.delete_rows(self.
|
|
365
|
+
query = sql.delete_rows(self._tableschema["tablename"], where_clause)
|
|
366
366
|
return query, args
|
|
367
367
|
|
|
368
|
+
|
|
368
369
|
def _reset(self) -> None:
|
|
369
370
|
pass
|
|
370
371
|
|
|
371
372
|
|
|
372
373
|
class Update(MutationalStatement):
|
|
373
374
|
|
|
374
|
-
|
|
375
|
-
|
|
376
|
-
super().__init__(table)
|
|
375
|
+
def __init__(self, connection : ConnectionManager, tableschema : Schema) -> None:
|
|
376
|
+
super().__init__(connection, tableschema)
|
|
377
377
|
self._set_clauses : list[str] = []
|
|
378
378
|
self._set_args : list[SqlValue] = []
|
|
379
379
|
|
|
@@ -391,7 +391,7 @@ class Update(MutationalStatement):
|
|
|
391
391
|
|
|
392
392
|
def _build(self, where_clause : str, *args : SqlValue) -> tuple[str, tuple[SqlValue, ...]]:
|
|
393
393
|
set_clause = sql.format_list(self._set_clauses, brackets=False)
|
|
394
|
-
query = f"UPDATE {self.
|
|
394
|
+
query = f"UPDATE {self._tableschema["tablename"]} SET {set_clause} WHERE {where_clause};"
|
|
395
395
|
return query, (*self._set_args, *args)
|
|
396
396
|
|
|
397
397
|
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
import sqlite3
|
|
2
|
-
from typing import Protocol, Self, TypeGuard, TypedDict
|
|
3
|
-
|
|
2
|
+
from typing import Protocol, Self, TypeGuard, TypedDict, Any
|
|
3
|
+
from types import UnionType
|
|
4
4
|
|
|
5
5
|
class CustomType(Protocol):
|
|
6
6
|
@classmethod
|
|
@@ -24,12 +24,20 @@ class Schema(TypedDict):
|
|
|
24
24
|
primary : list[str]
|
|
25
25
|
|
|
26
26
|
|
|
27
|
+
class Primary[T]:
|
|
28
|
+
__slots__ = ()
|
|
29
|
+
|
|
30
|
+
|
|
27
31
|
# https://docs.python.org/3/library/sqlite3.html#sqlite-and-python-types
|
|
28
|
-
_TYPES_MAP : dict[type, str] = {
|
|
29
|
-
int : "INTEGER",
|
|
30
|
-
float : "REAL",
|
|
31
|
-
str : "TEXT",
|
|
32
|
-
bytes : "BLOB",
|
|
32
|
+
_TYPES_MAP : dict[type | UnionType, str] = {
|
|
33
|
+
int : "INTEGER NOT NULL",
|
|
34
|
+
float : "REAL NOT NULL",
|
|
35
|
+
str : "TEXT NOT NULL",
|
|
36
|
+
bytes : "BLOB NOT NULL",
|
|
37
|
+
None | int : "INTEGER",
|
|
38
|
+
None | float : "REAL",
|
|
39
|
+
None | str : "TEXT",
|
|
40
|
+
None | bytes : "BLOB",
|
|
33
41
|
}
|
|
34
42
|
|
|
35
43
|
|
|
@@ -62,4 +70,35 @@ def pytype_to_sqltype(type_ : type) -> str:
|
|
|
62
70
|
if type_ not in _TYPES_MAP:
|
|
63
71
|
raise TypeError(f"{type_} is not natively supported by sqlite3")
|
|
64
72
|
|
|
65
|
-
return _TYPES_MAP[type_]
|
|
73
|
+
return _TYPES_MAP[type_]
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def register_resolve_types(types : list[SqlType | str], **connection_params) -> tuple[list[str], dict[str, Any]]:
|
|
77
|
+
""" Converts py types to sql types, registers custom types, resolves type names, updates connection params """
|
|
78
|
+
|
|
79
|
+
resolved : list[str] = []
|
|
80
|
+
assert_register_types = False
|
|
81
|
+
|
|
82
|
+
for type_ in types:
|
|
83
|
+
|
|
84
|
+
if isinstance(type_, str):
|
|
85
|
+
resolved.append(type_)
|
|
86
|
+
continue
|
|
87
|
+
|
|
88
|
+
if is_custom_type(type_):
|
|
89
|
+
ctname = type_.__name__.upper()
|
|
90
|
+
|
|
91
|
+
register_type(type_, ctname)
|
|
92
|
+
resolved.append(ctname)
|
|
93
|
+
|
|
94
|
+
if not assert_register_types:
|
|
95
|
+
assert_register_types = True
|
|
96
|
+
continue
|
|
97
|
+
|
|
98
|
+
sql_type = pytype_to_sqltype(type_)
|
|
99
|
+
resolved.append(sql_type)
|
|
100
|
+
|
|
101
|
+
if assert_register_types and ("detect_types" not in connection_params):
|
|
102
|
+
connection_params.update(dict(detect_types=sqlite3.PARSE_DECLTYPES))
|
|
103
|
+
|
|
104
|
+
return resolved, connection_params
|
sqlengine/schema.py
CHANGED