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 CHANGED
@@ -1,7 +1,7 @@
1
- from .utils import sqlgen, types
2
- from .utils.types import Schema
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,4 @@
1
+ from .statements import Select, Delete, Update
2
+ from .connection import ConnectionManager
3
+
4
+ __all__ = ["Select", "Delete", "Update", "ConnectionManager"]
@@ -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
+
@@ -1,12 +1,5 @@
1
- from __future__ import annotations
2
1
  from html import escape
3
- from typing import Sequence, TYPE_CHECKING
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 not brackets:
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 f"{format_list([place_holder]*n_rows, False)}"
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 __future__ import annotations
2
- from typing import Sequence, Literal, Generator, Self, TYPE_CHECKING
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 .types import SqlValue, SqlRow
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
- __command__ : Literal["SELECT", "INSERT", "UPDATE", "DELETE"]
145
+ def __init__(self, connection : ConnectionManager, tableschema : Schema) -> None:
150
146
 
151
- def __init__(self, table : SqlTableMixin) -> None:
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
- """ Resets statement to reuse object """
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._table.execute(query, *args)
210
+ self._connection.execute(query, *args)
217
211
 
218
212
 
219
213
  class Select(Statement):
220
214
 
221
-
222
- def __init__(self, table : SqlTableMixin) -> None:
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._table.fetchone(query, *args)
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._table.fetchmany(query, *args, size=size)
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._table.fetchall(query, *args)
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._table.in_transaction():
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._table.tx_conn.cursor()
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._table.in_transaction():
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._table.tx_conn.cursor()
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._table.tablename, columns, where_clause, order, limit)
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._table.columns if len(self._columns) == 0 or "*" in self._columns else self._columns
352
+ columns = self._resolve_columns()
352
353
  repr_rows = self.fetchmany(limit)
353
354
 
354
- return to_html(self._table.tablename, columns, repr_rows, limit=limit-1)
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._table.tablename, where_clause)
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
- def __init__(self, table : SqlTableMixin) -> None:
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._table.tablename} SET {set_clause} WHERE {where_clause};"
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
@@ -3,7 +3,7 @@ import sqlite3
3
3
  from typing import overload
4
4
 
5
5
  from .sqltable import SqlTableMixin
6
- from .utils.types import Schema
6
+ from .core.types import Schema
7
7
 
8
8
 
9
9
  def get_database_tablenames(database : str, cursor : sqlite3.Cursor | None = None) -> list[str]: