ab_engine 0.1.1__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.
@@ -0,0 +1,145 @@
1
+ import sys, os
2
+ sys.path.append(os.path.dirname(__file__))
3
+ from driver import Driver as BaseDriver, RowFactory
4
+ from sqlite3 import connect as db_connect
5
+ from collections import namedtuple, OrderedDict
6
+
7
+
8
+ def dict_factory(cursor, row):
9
+ fields = [column[0] for column in cursor.description]
10
+ return {key: value for key, value in zip(fields, row)}
11
+
12
+
13
+ def namedtuple_factory(cursor, row):
14
+ fields = [column[0] for column in cursor.description]
15
+ cls = namedtuple("Row", fields)
16
+ return cls._make(row)
17
+
18
+ _Info = namedtuple("Info", ["column_name", "data_type", "is_nullable", "character_maximum_length",
19
+ "numeric_precision", "numeric_scale", "column_default", "autoincrement", "pk"])
20
+
21
+ def _field_ifo_tuples(fields, str_fields):
22
+ str_fields = str_fields["sql"].split("\n", 1)[1][:-1].strip().split("\n")
23
+ str_def = {}
24
+ for x in str_fields:
25
+ n, x = x.strip().split(" ", 1)
26
+ str_def[n] = x
27
+ for n, row in enumerate(fields):
28
+ if "(" in row["type"]:
29
+ sz = row["type"].split("(")[1][:-1]
30
+ sz = int(sz)
31
+ else:
32
+ sz = None
33
+ ai = str_def[row["name"]]
34
+ fields[n] = _Info(column_name=row["name"], data_type=row["type"].lower(), character_maximum_length=sz,
35
+ column_default=row["dflt_value"], is_nullable=row["notnull"] == 0 and row["pk"] == 0,
36
+ autoincrement=" autoincrement" in ai, pk=row["pk"],
37
+ numeric_precision=None, numeric_scale=None)
38
+ return fields
39
+
40
+ _FACTORY_ = {
41
+ RowFactory.ANY.value: None,
42
+ RowFactory.TUPLE.value: None,
43
+ RowFactory.DICT.value: dict_factory,
44
+ RowFactory.NAMED_TUPLE.value: namedtuple_factory,
45
+ }
46
+
47
+ class Driver(BaseDriver):
48
+
49
+ def __init__(self, connection_string, on_open_close=None):
50
+ """
51
+ test_db.sqlite
52
+ """
53
+ super().__init__(connection_string, on_open_close)
54
+
55
+ async def begin(self):
56
+ await self._before_open()
57
+ self._conn = db_connect(self.connection_string)
58
+
59
+ async def sql(self, query, one_row=False, row_factory=RowFactory.DICT):
60
+ if self._conn is None:
61
+ await self.begin()
62
+ acur = self._conn.cursor()
63
+ if x:=_FACTORY_[row_factory.value]:
64
+ acur.row_factory = x
65
+ acur.execute(query)
66
+ descr = acur.description
67
+ if not descr:
68
+ return acur.rowcount
69
+ if one_row:
70
+ ret = acur.fetchone()
71
+ else:
72
+ ret = acur.fetchall()
73
+ if not ret:
74
+ ret = []
75
+ return ret
76
+
77
+ async def commit(self):
78
+ if not self._conn:
79
+ raise RuntimeError("Transaction is not open")
80
+ self._conn.commit()
81
+ await self.rollback()
82
+
83
+ async def rollback(self):
84
+ if not self._conn:
85
+ return
86
+ self._conn.close()
87
+ await super().rollback()
88
+
89
+ async def table_struct(self, table_name) -> dict:
90
+ defs = await self.sql(f"select sql from sqlite_master where type='table' and name='{table_name}'", one_row=True)
91
+ if not defs:
92
+ return None
93
+ fields = await self.sql(f"PRAGMA table_info('{table_name}')")
94
+ fields, pk, defs = [], [], _field_ifo_tuples(fields, defs)
95
+ for x in defs:
96
+ t = self._specify_type(x)
97
+ field = {
98
+ "name": x.column_name,
99
+ "type": t.type_name,
100
+ "not_null": x.is_nullable is not None and f'{x.is_nullable} '[0].upper() in "FNН",
101
+ "python_type": t.python_type,
102
+ }
103
+ if t.autoincrement or x.autoincrement:
104
+ field["autoincrement"] = True
105
+ n = self.ident_name(field["name"])
106
+ if n != field["name"]:
107
+ field["field"] = n
108
+ if x.character_maximum_length:
109
+ field["size"] = x.character_maximum_length
110
+ elif field["python_type"] != int and x.numeric_precision:
111
+ n = f'{x.numeric_precision}.{x.numeric_scale}'
112
+ field["size"] = float(n)
113
+ if x.pk:
114
+ while len(pk) < x.pk:
115
+ pk.append(None)
116
+ pk[x.pk-1] = x.column_name
117
+ fields.append(field)
118
+ constraints = []
119
+ if pk:
120
+ constraints.append({
121
+ "type": "primary key",
122
+ "fields": pk
123
+ })
124
+ defs = await self.sql(f"PRAGMA foreign_key_list('{table_name}')")
125
+ fk = {}
126
+ for x in defs:
127
+ k = fk.get(x["id"], {"type": "foreign key", "table":x["table"], "fields":OrderedDict()})
128
+ k["fields"][x["from"]] = x["to"]
129
+ fk[x["id"]] = k
130
+ for x in fk:
131
+ constraints.append(fk[x])
132
+
133
+ defs = await self.sql(f"select sql from sqlite_master where type = 'index' and sql like '%UNIQUE INDEX % on {table_name} %'",
134
+ row_factory=RowFactory.TUPLE)
135
+ for x in defs:
136
+ x = x[0].split("(",1)[1][:-1]
137
+ constraints.append({
138
+ "type": "unique",
139
+ "fields": [f.strip() for f in x.split(",")],
140
+ })
141
+ return {
142
+ "table": table_name,
143
+ "fields": fields,
144
+ "constraints": constraints
145
+ }
ab_engine/db/option.py ADDED
@@ -0,0 +1,391 @@
1
+ from abc import ABC
2
+ from ..class_tools import classproperty
3
+ from .driver import RowFactory, Driver, _set_is_option
4
+ from pathlib import Path
5
+ from importlib.machinery import SourceFileLoader
6
+ from json5 import loads
7
+ from asyncio import BoundedSemaphore, wait_for, TimeoutError, sleep
8
+ from ..error import raise_error
9
+ from inspect import iscoroutinefunction
10
+ from typing import Callable
11
+ from gc import collect as garbage_collect
12
+
13
+ DRIVER_CLASSES = {}
14
+ _DRIVERS_ = {}
15
+ _LIMITS_ = {}
16
+ _CFG_ = None
17
+
18
+ def _check_cfg(key=""):
19
+ global _CFG_
20
+ if _CFG_ is None:
21
+ from ..env import Config
22
+ _CFG_ = Config()
23
+ if key=="":
24
+ return
25
+ if _CFG_.hasattr("defaults"):
26
+ return _CFG_.defaults.get(key)
27
+
28
+
29
+ class Option(ABC):
30
+
31
+ @staticmethod
32
+ def is_option(x):
33
+ if isinstance(x, Option):
34
+ return True
35
+ while hasattr(x, "__base__"):
36
+ if x.__base__ == Option:
37
+ return True
38
+ x = x.__base__
39
+ return False
40
+
41
+ @classproperty
42
+ def row_factory(cls)->RowFactory:
43
+ return RowFactory.ANY
44
+
45
+ @classproperty
46
+ def can_process(cls):
47
+ return False
48
+
49
+ @classproperty
50
+ def one_row(cls):
51
+ return None
52
+
53
+ @staticmethod
54
+ async def process(ret_data, connection, row_factory):
55
+ """
56
+ Возвращает преобразованный набор данных и формат
57
+ если вместо формата вернулся None - дальнейшее преобразование невозможно
58
+ """
59
+ return ret_data, row_factory
60
+
61
+ _set_is_option(Option.is_option)
62
+
63
+ class ALL(Option):
64
+ ...
65
+
66
+ class DB(Option):
67
+ """
68
+ Соединение с БД
69
+ """
70
+ _TRASH_ = set()
71
+
72
+ def __init__(self, connection_string: str=""):
73
+ connection_string = connection_string.strip()
74
+ if connection_string.startswith("jdbc:"):
75
+ connection_string = connection_string[5:]
76
+ if connection_string.endswith("}"):
77
+ connection_string, params = connection_string.split("{",1)
78
+ self._params = loads(f"{{{params}")
79
+ for x in tuple(self._params.keys()):
80
+ if "$" in x:
81
+ v = self._params[x]
82
+ del self._params[x]
83
+ self._params[x.replace("$",".")] = v
84
+ else:
85
+ self._params = {}
86
+ if "://" not in connection_string:
87
+ _check_cfg()
88
+ connection_string = _CFG_.db_connection(connection_string)
89
+ driver_name, connection_string = connection_string.split("://", 1)
90
+
91
+ connection_string, conn_params = f"{connection_string}?".split("?",1)
92
+
93
+ driver = DRIVER_CLASSES.get(driver_name)
94
+ driver_path = None
95
+ if conn_params:
96
+ conn_params = conn_params[:-1].split("&")
97
+ if "driver_path" in conn_params:
98
+ driver_path = conn_params["driver_path"]
99
+ del conn_params["driver_path"]
100
+ conn_params = "?" + "&".join(conn_params)
101
+ connection_string += conn_params
102
+ if driver is None:
103
+ if driver_path is None:
104
+ driver_path = _check_cfg("db_driver_path")
105
+ if driver_path is None:
106
+ driver_path = Path(__file__).parent / f"driver_{driver_name}.py"
107
+ elif isinstance(driver_path, str):
108
+ driver_path = Path(driver_path) / f"driver_{driver_name}.py"
109
+ m = driver_path.name
110
+ driver_path = str(driver_path)
111
+ if driver_path in _DRIVERS_:
112
+ driver = _DRIVERS_[driver_path]
113
+ else:
114
+ m = m.split(".",1)[0] + f"_{len(_DRIVERS_)}"
115
+ m = SourceFileLoader(m, driver_path ).load_module()
116
+ driver = m.Driver
117
+ _DRIVERS_[driver_path] = driver
118
+ else:
119
+ driver_path = ""
120
+
121
+ self._hash = hash(driver_path + connection_string)
122
+ if "LIMIT" in self._params:
123
+ m = self._params["LIMIT"]
124
+ del self._params["LIMIT"]
125
+ else:
126
+ m = 0
127
+ self._conn_limit = None
128
+ self.connection_limit = m
129
+ self._connection = driver(connection_string, self._on_open_close)
130
+
131
+ @property
132
+ def connection_limit(self)->int:
133
+ lmt = _LIMITS_.get(self._hash)
134
+ return lmt._value if lmt else 0
135
+
136
+
137
+ @connection_limit.setter
138
+ def connection_limit(self, value:int):
139
+ """
140
+ Позволяет задать ограничение на количество таких соединений
141
+ """
142
+ if value<=0:
143
+ if self._hash in _LIMITS_:
144
+ del _LIMITS_[self._hash]
145
+ return
146
+ elif lmt:=_LIMITS_.get(self._hash):
147
+ if lmt._waiters:
148
+ lmt._value = value - len(lmt._waiters)
149
+ else:
150
+ lmt._value = value
151
+ else:
152
+ lmt = BoundedSemaphore(value)
153
+ _LIMITS_[self._hash] = lmt
154
+ self._conn_limit = lmt
155
+
156
+
157
+ async def _on_open_close(self, close=False)->dict:
158
+ if self._conn_limit:
159
+ if close:
160
+ await self._conn_limit.release()
161
+ if self._hash not in _LIMITS_:
162
+ self._conn_limit = None
163
+ elif self._hash not in _LIMITS_:
164
+ self._conn_limit = None
165
+ else:
166
+ if self._conn_limit.locked():
167
+ await self.garbage_collect()
168
+ await self._conn_limit.acquire()
169
+ return self._params
170
+
171
+ @property
172
+ def connection(self)->Driver:
173
+ return self._connection
174
+
175
+ def __del__(self):
176
+ if self.connection is not None and self.connection.in_transaction:
177
+ DB._TRASH_.add(self.connection)
178
+
179
+ @classmethod
180
+ async def garbage_collect(cls, clear_mem=True):
181
+ if clear_mem:
182
+ garbage_collect()
183
+ while len(DB._TRASH_)>0:
184
+ x = DB._TRASH_.pop()
185
+ await x.rollback()
186
+
187
+ class TIMEOUT(Option):
188
+
189
+ def __init__(self, time_sec:float, raise_error:bool=True):
190
+ self._time_sec = time_sec
191
+ self._raise_error = raise_error
192
+
193
+ async def __call__(self, func):
194
+ try:
195
+ return await wait_for(func, self._time_sec)
196
+ except TimeoutError as e:
197
+ if self._raise_error:
198
+ raise e
199
+ return None
200
+
201
+ class DICT(Option):
202
+ """
203
+ Указывает, что результат следует вернуть как список dict (по умолчанию)
204
+ """
205
+ @classproperty
206
+ def row_factory(cls) -> RowFactory:
207
+ return RowFactory.DICT
208
+
209
+ class TUPLE(Option):
210
+ """
211
+ Указывает, что результат следует вернуть как список tuple
212
+ """
213
+ @classproperty
214
+ def row_factory(cls) -> RowFactory:
215
+ return RowFactory.TUPLE
216
+
217
+ class OBJECT(Option):
218
+ """
219
+ Указывает, что результат следует вернуть как список namedtuple
220
+ """
221
+ @classproperty
222
+ def row_factory(cls) -> RowFactory:
223
+ return RowFactory.NAMED_TUPLE
224
+
225
+ class ROW(Option):
226
+ """
227
+ Указывает, что нужно вернуть только первую строку из выборки
228
+ """
229
+ @classproperty
230
+ def one_row(cls):
231
+ return True
232
+
233
+ class ONE(ROW):
234
+ """
235
+ Указывает, что нужно вернуть только значение первого поля из первой строки выборки
236
+ """
237
+ @classproperty
238
+ def can_process(cls):
239
+ return True
240
+
241
+ @staticmethod
242
+ async def process(ret_data, connection, row_factory):
243
+ if isinstance(ret_data, (int, None.__class__)):
244
+ return ret_data, None
245
+ match row_factory:
246
+ case RowFactory.TUPLE:
247
+ return ret_data[0], None
248
+ case RowFactory.DICT:
249
+ return ret_data[tuple(ret_data.keys())[0]], None
250
+ case RowFactory.NAMED_TUPLE:
251
+ x = ret_data._fields
252
+ if len(x) < 1:
253
+ return None, None
254
+ x = x[0]
255
+ return getattr(ret_data, x), None
256
+ case _:
257
+ return ret_data, None
258
+
259
+ class JSON(ONE):
260
+ """
261
+ Указывает, что значение первого поля в первой строке следует привести к dict или list
262
+ """
263
+ @staticmethod
264
+ async def process(ret_data, connection, row_factory):
265
+ ret_data, _ = await ONE.process(ret_data, connection, row_factory)
266
+ if isinstance(ret_data, (dict, list)):
267
+ return ret_data, None
268
+ elif isinstance(ret_data, str):
269
+ return loads(ret_data), None
270
+ else:
271
+ raise_error("NOT_CONV_DICT_LIST", data=ret_data)
272
+
273
+ class ROLLBACK(Option):
274
+ """
275
+ Указывает, что после выполнения запроса соединение должно быть закрыто с откатом транзакции
276
+ """
277
+ @classproperty
278
+ def can_process(cls):
279
+ return True
280
+
281
+ @staticmethod
282
+ async def process(ret_data, connection, row_factory):
283
+ if connection.in_transaction:
284
+ await connection.rollback()
285
+ return ret_data, row_factory
286
+
287
+ class COMMIT(Option):
288
+ """
289
+ Указывает, что после выполнения запроса соединение должно быть закрыто с подтверждением транзакции
290
+ """
291
+ @classproperty
292
+ def can_process(cls):
293
+ return True
294
+
295
+ @staticmethod
296
+ async def process(ret_data, connection, row_factory):
297
+ if connection.in_transaction:
298
+ await connection.commit()
299
+ return ret_data, row_factory
300
+
301
+ class PAGE(Option):
302
+
303
+ def __init__(self, limit:int, offset:int=0):
304
+ """
305
+ Ограничивает количество записей на странице и задает смещение первой записи страницы от начала выборки
306
+ :param limit: максимальное количество строк
307
+ :param offset: смещение от начала
308
+ """
309
+ if not(isinstance(limit, int) and isinstance(offset, int)) or limit < 0 or offset < 0:
310
+ raise_error("BAD_LIMIT_OFFSET")
311
+ self._limit = limit
312
+ self._offset = offset
313
+
314
+ async def __call__(self, connection, query):
315
+ query = await connection.page(query, self._limit, self._offset)
316
+ return query
317
+
318
+ def __str__(self):
319
+ return f"PAGE(limit={self._limit}, offset={self._offset})"
320
+
321
+
322
+ class CALLBACK(Option):
323
+
324
+ def __init__(self, callback_function:Callable, attribute_collection=None, *args, **kwargs):
325
+ """
326
+ Позволяет вернуть запрос с подставленными параметрами из функции sql и доработать его или вывести в лог
327
+ если callback_function возвращает строку, то именно эта строка станет запросом
328
+ :param callback_function: ссылка на функцию, которой передается текст запроса
329
+ """
330
+ self._callback = callback_function
331
+ self._args = args
332
+ self._kwargs = kwargs
333
+ self._attrs = attribute_collection
334
+
335
+ async def __call__(self, query):
336
+ if not self._callback:
337
+ return query
338
+ if iscoroutinefunction(self._callback):
339
+ x = await self._callback(query, *self._args, **self._kwargs)
340
+ else:
341
+ x = self._callback(query, *self._args, **self._kwargs)
342
+ return x or query
343
+
344
+ def __getitem__(self, item):
345
+ if self._attrs is None:
346
+ raise_error("ATTR_NOT_FOUND", name=item)
347
+ return self._attrs[item]
348
+
349
+
350
+ class ITERATOR(Option):
351
+
352
+ def __init__(self, page_size:int=200, async_delay=0.00000001):
353
+ self._page_size = page_size
354
+ self._delay = async_delay
355
+ self._ofs = 0
356
+ self._pos = page_size-1
357
+ self._buf = None
358
+ self._query = None
359
+ self._db = None
360
+ self._row_factory = None
361
+ self._process = None
362
+
363
+ def __call__(self, query, db=None, row_factory=RowFactory.DICT, process=None):
364
+ if self._query:
365
+ return None
366
+ self._query = query
367
+ self._db = db
368
+ self._row_factory = row_factory
369
+ self._process = process if isinstance(process, list) else []
370
+ return self
371
+
372
+ def __aiter__(self):
373
+ self._ofs = -self._page_size
374
+ self._pos = self._page_size - 1
375
+ self._buf = None
376
+ return self
377
+
378
+ async def __anext__(self):
379
+ await sleep(self._delay)
380
+ self._pos += 1
381
+ if self._buf and self._pos >= len(self._buf) and len(self._buf) < self._page_size:
382
+ raise StopAsyncIteration
383
+ elif self._pos >= self._page_size:
384
+ self._ofs += self._page_size
385
+ self._pos = 0
386
+ q = await self._db.connection.page(self._query, self._page_size, self._ofs)
387
+ self._buf = await self._db.connection.sql(q, one_row=False, row_factory=self._row_factory)
388
+ if not self._buf:
389
+ raise StopAsyncIteration
390
+ return self._buf[self._pos]
391
+
@@ -0,0 +1,90 @@
1
+ from .option import Option, DB, TIMEOUT, PAGE, CALLBACK, ITERATOR
2
+ from .driver import RowFactory
3
+ from ..error import raise_error
4
+
5
+ CONNECTION:str = ""
6
+
7
+ def set_connection(connection_string:str):
8
+ global CONNECTION
9
+ CONNECTION = connection_string
10
+
11
+ async def sql(query:str, *args, **kwargs):
12
+ process = []
13
+ db = callback = tm = row_factory = one_row = page = itr = None
14
+ for n, arg in enumerate(args):
15
+ if Option.is_option(arg):
16
+ if not process:
17
+ process.append(n)
18
+ if one_row is None and arg.one_row is not None:
19
+ one_row = arg
20
+ elif one_row and arg.one_row is not None and one_row.one_row!=arg.one_row:
21
+ raise_error("NEQ_ROWS", src1=one_row.__class__.__name__, src2=arg.__class__.__name__)
22
+ if row_factory is None or row_factory.row_factory == RowFactory.ANY:
23
+ row_factory = arg
24
+ elif row_factory and arg.row_factory != RowFactory.ANY:
25
+ raise_error("DEF_STR_FABRIC", src = row_factory.__class__.__name__)
26
+ if arg.can_process:
27
+ process.append(arg)
28
+ continue
29
+ if arg == DB:
30
+ arg = DB()
31
+ if isinstance(arg, DB) and db is None:
32
+ db = arg
33
+ elif isinstance(arg, CALLBACK) and callback is None:
34
+ callback = arg
35
+ elif isinstance(arg, DB):
36
+ raise_error("DB_ALRDY_DEF")
37
+ elif isinstance(arg, PAGE) and page is None:
38
+ page = arg
39
+ elif isinstance(arg, PAGE):
40
+ raise_error("PAGE_ALRDY_DEF")
41
+ elif isinstance(arg, TIMEOUT) and tm is None:
42
+ tm = arg
43
+ elif isinstance(arg, TIMEOUT):
44
+ raise_error("TIMEOUT_ALRDY_DEF")
45
+ elif arg == ITERATOR:
46
+ itr = ITERATOR()
47
+ elif isinstance(arg, ITERATOR):
48
+ itr = arg
49
+ elif process:
50
+ raise_error("PMT_BEF_OPT")
51
+ if process:
52
+ args = args[:process.pop(0)]
53
+ if db is None:
54
+ db = DB(CONNECTION)
55
+ in_self = True
56
+ else:
57
+ in_self = False
58
+ await db.garbage_collect(False)
59
+ one_row = one_row.one_row if one_row else False
60
+ row_factory = row_factory.row_factory if row_factory and row_factory.row_factory != RowFactory.ANY else RowFactory.DICT
61
+ if callback:
62
+ kwargs["__PARAM_CALLBACK_GETTER"] = callback
63
+ query = await db.connection.parse_query(query, *args, **kwargs)
64
+ if page:
65
+ query = await page(db.connection, query)
66
+ if callback is not None:
67
+ query = await callback(query)
68
+ if itr:
69
+ if page:
70
+ raise_error("PAGE_ITERATOR")
71
+ return itr(query, db, row_factory, process)
72
+ try:
73
+ if tm is None:
74
+ ret = await db.connection.sql(query, one_row=one_row, row_factory=row_factory)
75
+ else:
76
+ ret = await tm(db.connection.sql(query, one_row=one_row, row_factory=row_factory))
77
+ for f in process:
78
+ ret, row_factory = await f.process(ret, db.connection, row_factory)
79
+ if row_factory is not None:
80
+ break
81
+ except Exception as e:
82
+ try:
83
+ if db.connection.in_transaction:
84
+ await db.connection.rollback()
85
+ except Exception as e2:
86
+ ...
87
+ raise e
88
+ if in_self and db.connection.in_transaction:
89
+ await db.connection.commit()
90
+ return ret