-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsqlite.py
More file actions
74 lines (57 loc) · 2.35 KB
/
Copy pathsqlite.py
File metadata and controls
74 lines (57 loc) · 2.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
from abc import abstractmethod
from enum import Enum
from typing import Type, TypeVar
from sqlalchemy import create_engine
from sqlalchemy.exc import OperationalError, IntegrityError
from sqlalchemy.future import select
from sqlalchemy.orm import declarative_base, Session, Query, sessionmaker, scoped_session
Base = declarative_base()
class GenericTypeVar(TypeVar, _root=True):
def __getitem__(self, item): pass
# add more items to count as a "query" here
GenericQuery = GenericTypeVar("GenericQuery", Query, select)
M = GenericTypeVar("M")
# https://stackoverflow.com/questions/49581907/when-inheriting-sqlalchemy-class-from-abstract-class-exception-thrown-metaclass
class BaseWithMigrations(Base):
__abstract__ = True
@classmethod
@abstractmethod
def migrations(cls) -> list[str]:
pass
class SqliteStore:
class ParallelizationMode(Enum):
main = "main"
threaded = "threaded"
def __init__(self, db_filename: str, models: list[Type[BaseWithMigrations]], ):
self.engine = create_engine(f"sqlite:///{db_filename}.sqlite", echo=False, future=True)
Base.metadata.create_all(self.engine)
# self.session_factory = sessionmaker(bind=self.engine)
self.session: Session = Session(self.engine)
with self.session.begin() as tx:
for model in models:
for migration in model.migrations():
self.ddl_statement(self.session, migration)
print(f"migrations done")
tx.commit()
@staticmethod
def ddl_statement(session: Session, statement: str):
try:
session.execute(statement)
except OperationalError as oe:
msg = str(oe)
if "duplicate" not in msg:
raise oe
except IntegrityError as ie:
msg = str(ie)
if "UNIQUE constraint failed" not in msg:
raise ie
def store_row(self, row: Base):
return self.store_rows([row])
def store_rows(self, rows: list[Base]):
# session = scoped_session(self.session_factory) if new_session else self.session
session: Session = self.session
with session.begin():
session.add_all(rows)
def fetch_entities(self, stmt: GenericQuery[M], ) -> list[M]:
res: list[M] = self.session.execute(statement=stmt).scalars().all()
return res