maubot-audio-preventer

git clone git://archive.git.mtrnord.blog/MTRNord/maubot-audio-preventer.git
Log | Files | Refs | README | LICENSE

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()