Skip to content

Commit 6254764

Browse files
committed
models update
1 parent 7f97259 commit 6254764

2 files changed

Lines changed: 104 additions & 87 deletions

File tree

src/models/base.py

Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
from __future__ import annotations
2+
3+
import re
4+
5+
from sqlalchemy import not_
6+
from sqlalchemy.exc import NoResultFound
7+
from sqlalchemy.orm import Query, Session, as_declarative, declared_attr
8+
9+
from src.utils.exceptions import AlreadyExists, ObjectNotFound
10+
11+
12+
@as_declarative()
13+
class Base:
14+
"""Base class for all database entities"""
15+
16+
@declared_attr
17+
def __tablename__(cls) -> str: # pylint: disable=no-self-argument
18+
"""Generate database table name automatically.
19+
Convert CamelCase class name to snake_case db table name.
20+
"""
21+
return re.sub(r"(?<!^)(?=[A-Z])", "_", cls.__name__).lower()
22+
23+
def __repr__(self):
24+
attrs = []
25+
for c in self.__table__.columns:
26+
attrs.append(f"{c.name}={getattr(self, c.name)}")
27+
return "{}({})".format(c.__class__.__name__, ', '.join(attrs))
28+
29+
30+
class BaseDbModel(Base):
31+
__abstract__ = True
32+
33+
@classmethod
34+
def create(cls, *, session: Session, **kwargs) -> BaseDbModel:
35+
obj = cls(**kwargs)
36+
session.add(obj)
37+
session.flush()
38+
return obj
39+
40+
@classmethod
41+
def query(cls, *, with_deleted: bool = False, session: Session) -> Query:
42+
"""Get all objects with soft deletes"""
43+
objs = session.query(cls)
44+
if not with_deleted and hasattr(cls, "is_deleted"):
45+
objs = objs.filter(not_(cls.is_deleted))
46+
return objs
47+
48+
@classmethod
49+
def get(cls, id: int | str, *, with_deleted=False, session: Session) -> BaseDbModel:
50+
"""Get object with soft deletes"""
51+
objs = session.query(cls)
52+
if not with_deleted and hasattr(cls, "is_deleted"):
53+
objs = objs.filter(not_(cls.is_deleted))
54+
try:
55+
if hasattr(cls, "uuid"):
56+
return objs.filter(cls.uuid == id).one()
57+
return objs.filter(cls.id == id).one()
58+
except NoResultFound:
59+
raise ObjectNotFound(cls, id)
60+
61+
@classmethod
62+
def update(cls, id: int | str, *, session: Session, **kwargs) -> BaseDbModel:
63+
"""Update model with new values from kwargs.
64+
If no new values are given, raise HTTP 409 error.
65+
"""
66+
get_new_values = False
67+
obj = cls.get(id, session=session)
68+
for k, v in kwargs.items():
69+
cur_v = getattr(obj, k)
70+
if cur_v != v:
71+
setattr(obj, k, v)
72+
get_new_values = True
73+
if not get_new_values:
74+
raise AlreadyExists(cls, id)
75+
session.add(obj)
76+
session.flush()
77+
return obj
78+
79+
@classmethod
80+
def delete(cls, id: int | str, *, session: Session) -> None:
81+
"""Soft delete object if possible, else hard delete"""
82+
obj = cls.get(id, session=session)
83+
if hasattr(obj, "is_deleted"):
84+
obj.is_deleted = True
85+
else:
86+
session.delete(obj)
87+
session.flush()

src/models/db.py

Lines changed: 17 additions & 87 deletions
Original file line numberDiff line numberDiff line change
@@ -1,87 +1,17 @@
1-
from __future__ import annotations
2-
3-
import re
4-
5-
from sqlalchemy import not_
6-
from sqlalchemy.exc import NoResultFound
7-
from sqlalchemy.orm import Query, Session, as_declarative, declared_attr
8-
9-
from src.utils.exceptions import AlreadyExists, ObjectNotFound
10-
11-
12-
@as_declarative()
13-
class Base:
14-
"""Base class for all database entities"""
15-
16-
@declared_attr
17-
def __tablename__(cls) -> str: # pylint: disable=no-self-argument
18-
"""Generate database table name automatically.
19-
Convert CamelCase class name to snake_case db table name.
20-
"""
21-
return re.sub(r"(?<!^)(?=[A-Z])", "_", cls.__name__).lower()
22-
23-
def __repr__(self):
24-
attrs = []
25-
for c in self.__table__.columns:
26-
attrs.append(f"{c.name}={getattr(self, c.name)}")
27-
return "{}({})".format(c.__class__.__name__, ', '.join(attrs))
28-
29-
30-
class BaseDbModel(Base):
31-
__abstract__ = True
32-
33-
@classmethod
34-
def create(cls, *, session: Session, **kwargs) -> BaseDbModel:
35-
obj = cls(**kwargs)
36-
session.add(obj)
37-
session.flush()
38-
return obj
39-
40-
@classmethod
41-
def query(cls, *, with_deleted: bool = False, session: Session) -> Query:
42-
"""Get all objects with soft deletes"""
43-
objs = session.query(cls)
44-
if not with_deleted and hasattr(cls, "is_deleted"):
45-
objs = objs.filter(not_(cls.is_deleted))
46-
return objs
47-
48-
@classmethod
49-
def get(cls, id: int | str, *, with_deleted=False, session: Session) -> BaseDbModel:
50-
"""Get object with soft deletes"""
51-
objs = session.query(cls)
52-
if not with_deleted and hasattr(cls, "is_deleted"):
53-
objs = objs.filter(not_(cls.is_deleted))
54-
try:
55-
if hasattr(cls, "uuid"):
56-
return objs.filter(cls.uuid == id).one()
57-
return objs.filter(cls.id == id).one()
58-
except NoResultFound:
59-
raise ObjectNotFound(cls, id)
60-
61-
@classmethod
62-
def update(cls, id: int | str, *, session: Session, **kwargs) -> BaseDbModel:
63-
"""Update model with new values from kwargs.
64-
If no new values are given, raise HTTP 409 error.
65-
"""
66-
get_new_values = False
67-
obj = cls.get(id, session=session)
68-
for k, v in kwargs.items():
69-
cur_v = getattr(obj, k)
70-
if cur_v != v:
71-
setattr(obj, k, v)
72-
get_new_values = True
73-
if not get_new_values:
74-
raise AlreadyExists(cls, id)
75-
session.add(obj)
76-
session.flush()
77-
return obj
78-
79-
@classmethod
80-
def delete(cls, id: int | str, *, session: Session) -> None:
81-
"""Soft delete object if possible, else hard delete"""
82-
obj = cls.get(id, session=session)
83-
if hasattr(obj, "is_deleted"):
84-
obj.is_deleted = True
85-
else:
86-
session.delete(obj)
87-
session.flush()
1+
from sqlalchemy.orm import Mapped, mapped_column, relationship
2+
from .base import BaseDBModel
3+
from sqlalchemy import Integer, String, ForeignKey
4+
5+
class Users(BaseDBModel):
6+
__tablename__ = "user_info"
7+
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
8+
tg_id: Mapped[int] = mapped_column(Integer, nullable=False, unique=True)
9+
name: Mapped[str] = mapped_column(String, nullable=False)
10+
birthday: Mapped[str] = mapped_column(String, default=None)
11+
about: Mapped[str] = mapped_column(String, nullable=False)
12+
13+
class Holidays(BaseDBModel):
14+
__tablename__ = "holidays_status"
15+
id: Mapped[int] = mapped_column(Integer, ForeignKey("user_info.id"), primary_key=True)
16+
status: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
17+
till_date: Mapped[str] = mapped_column(String, nullable=False, default="null")

0 commit comments

Comments
 (0)