db.py (1789B)
1 from sqlalchemy import Column, String, BigInteger 2 from sqlalchemy.engine.base import Engine 3 from sqlalchemy.ext.declarative import declarative_base 4 from sqlalchemy.orm import Session, sessionmaker 5 6 from mautrix.types import UserID 7 8 from typing import NamedTuple, Optional 9 10 UserInfo = NamedTuple('UserInfo', text_warnings=int, kick_warnings=int) 11 Base = declarative_base() 12 13 14 class Warnings(Base): 15 __tablename__ = "warnings" 16 17 user_id: UserID = Column(String(255), primary_key=True, nullable=False) 18 text_warnings = Column(BigInteger, primary_key=False, nullable=False) 19 kick_warnings = Column(BigInteger, primary_key=False, nullable=False) 20 21 22 class Database: 23 db: Engine 24 25 def __init__(self, db: Engine) -> None: 26 self.db = db 27 Base.metadata.create_all(db) 28 self.Session = sessionmaker(bind=self.db) 29 30 def get_user(self, mxid: UserID) -> Optional[UserInfo]: 31 s: Session = self.Session() 32 try: 33 row = s.query(Warnings).filter(Warnings.user_id == mxid).one() 34 return UserInfo(text_warnings=row.text_warnings, kick_warnings=row.kick_warnings) 35 finally: 36 return None 37 38 def add_user(self, mxid: UserID) -> None: 39 token_row = Warnings(user_id=mxid, text_warnings=1, kick_warnings=0) 40 s: Session = self.Session() 41 s.add(token_row) 42 s.commit() 43 44 def increment_text_warnings(self, mxid: UserID, current_warnings: int) -> None: 45 s: Session = self.Session() 46 s.merge(Warnings(user_id=mxid, text_warnings=current_warnings+1)) 47 s.commit() 48 49 def increment_kick_warnings(self, mxid: UserID, current_warnings: int) -> None: 50 s: Session = self.Session() 51 s.merge(Warnings(user_id=mxid, kick_warnings=current_warnings+1)) 52 s.commit()