264 lines
9.4 KiB
Python
264 lines
9.4 KiB
Python
"""
|
|
models/db.py — DBfrontend: encrypted message store.
|
|
"""
|
|
|
|
import time
|
|
from typing import Optional
|
|
from cryptography.fernet import Fernet
|
|
from cryptography.hazmat.primitives import hashes
|
|
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
|
import base64
|
|
from sqlalchemy.orm import sessionmaker
|
|
from sqlalchemy.exc import IntegrityError
|
|
|
|
from .base import make_engine, Base, Messages, Poll, PollOption, PollVote
|
|
|
|
|
|
class DBfrontend:
|
|
def __init__(self, db_cfg, secret: bytes, salt: bytes):
|
|
self.engine = make_engine(db_cfg)
|
|
self.Session = sessionmaker(bind=self.engine)
|
|
self.cipher = self._get_cipher(secret, salt)
|
|
self._migrate()
|
|
|
|
def _migrate(self):
|
|
with self.Session() as session:
|
|
from sqlalchemy import inspect, text
|
|
inspector = inspect(session.connection())
|
|
|
|
columns = [c["name"] for c in inspector.get_columns("messages")]
|
|
if "poll_id" not in columns:
|
|
session.execute(text("ALTER TABLE messages ADD COLUMN poll_id INTEGER"))
|
|
|
|
constraints = [uq["name"] for uq in inspector.get_unique_constraints("poll_votes")]
|
|
if "uq_poll_voter" in constraints:
|
|
old = "poll_votes_old"
|
|
session.execute(text(f"ALTER TABLE poll_votes RENAME TO {old}"))
|
|
PollVote.__table__.create(session.connection())
|
|
session.execute(text(f"""
|
|
INSERT INTO poll_votes (id, poll_id, option_id, voter_ip, voted_at)
|
|
SELECT id, poll_id, option_id, voter_ip, voted_at FROM {old}
|
|
"""))
|
|
session.execute(text(f"DROP TABLE {old}"))
|
|
|
|
session.commit()
|
|
|
|
@staticmethod
|
|
def _get_cipher(secret: bytes, salt: bytes) -> Fernet:
|
|
kdf = PBKDF2HMAC(algorithm=hashes.SHA256(), length=32, salt=salt, iterations=100_000)
|
|
key = base64.urlsafe_b64encode(kdf.derive(secret))
|
|
return Fernet(key)
|
|
|
|
def _encrypt(self, message: str) -> str:
|
|
return self.cipher.encrypt(message.encode()).decode()
|
|
|
|
def _decrypt(self, encrypted: str) -> str:
|
|
try:
|
|
return self.cipher.decrypt(encrypted.encode()).decode()
|
|
except Exception:
|
|
return "[Decryption failed]"
|
|
|
|
def add_message(self, channel: str, ip: str, content: str, gif_url: Optional[str] = None, poll_id: Optional[int] = None) -> int:
|
|
with self.Session() as session:
|
|
msg = Messages(
|
|
channel=channel,
|
|
ip=ip,
|
|
content=self._encrypt(content),
|
|
gif_url=gif_url,
|
|
poll_id=poll_id,
|
|
timestamp=int(time.time() * 1000),
|
|
)
|
|
session.add(msg)
|
|
session.commit()
|
|
return msg.id
|
|
|
|
def get_messages(self, channel: str = "general", limit: int = 50) -> list:
|
|
with self.Session() as session:
|
|
q = (
|
|
session.query(Messages)
|
|
.filter(Messages.channel == channel)
|
|
.order_by(Messages.id.desc())
|
|
.limit(limit)
|
|
)
|
|
result = []
|
|
for msg in reversed(q.all()):
|
|
poll_data = None
|
|
if msg.poll_id:
|
|
poll = session.get(Poll, msg.poll_id)
|
|
if poll and poll.status != "deleted":
|
|
poll_data = self._serialize_poll(session, poll)
|
|
result.append(
|
|
{
|
|
"id": msg.id,
|
|
"channel": msg.channel,
|
|
"ip": msg.ip,
|
|
"content": self._decrypt(msg.content),
|
|
"gif_url": msg.gif_url,
|
|
"poll_id": msg.poll_id,
|
|
"poll": poll_data,
|
|
"timestamp": msg.timestamp,
|
|
}
|
|
)
|
|
return result
|
|
|
|
def delete_message(self, msg_id: int) -> bool:
|
|
with self.Session() as session:
|
|
msg = session.get(Messages, msg_id)
|
|
if msg:
|
|
session.delete(msg)
|
|
session.commit()
|
|
return True
|
|
return False
|
|
|
|
def add_poll(
|
|
self,
|
|
channel: str,
|
|
creator_ip: str,
|
|
question: str,
|
|
options: list[str],
|
|
expires_at: Optional[int] = None,
|
|
is_multiple: bool = False,
|
|
is_anonymous: bool = False,
|
|
) -> int:
|
|
with self.Session() as session:
|
|
poll = Poll(
|
|
channel=channel,
|
|
creator_ip=creator_ip,
|
|
question=self._encrypt(question),
|
|
created_at=int(time.time() * 1000),
|
|
expires_at=expires_at,
|
|
is_multiple=is_multiple,
|
|
is_anonymous=is_anonymous,
|
|
status="active",
|
|
)
|
|
session.add(poll)
|
|
session.flush()
|
|
|
|
for i, label in enumerate(options):
|
|
opt = PollOption(poll_id=poll.id, label=self._encrypt(label), position=i)
|
|
session.add(opt)
|
|
|
|
msg = Messages(
|
|
channel=channel,
|
|
ip=creator_ip,
|
|
content=self._encrypt(question),
|
|
poll_id=poll.id,
|
|
timestamp=int(time.time() * 1000),
|
|
)
|
|
session.add(msg)
|
|
|
|
session.commit()
|
|
return poll.id
|
|
|
|
def get_polls(self, channel: str = "general") -> list:
|
|
with self.Session() as session:
|
|
polls = (
|
|
session.query(Poll)
|
|
.filter(Poll.channel == channel, Poll.status != "deleted")
|
|
.order_by(Poll.created_at.desc())
|
|
.limit(20)
|
|
)
|
|
result = []
|
|
for poll in polls.all():
|
|
result.append(self._serialize_poll(session, poll))
|
|
return result
|
|
|
|
def get_poll(self, poll_id: int) -> Optional[dict]:
|
|
with self.Session() as session:
|
|
poll = session.get(Poll, poll_id)
|
|
if not poll or poll.status == "deleted":
|
|
return None
|
|
return self._serialize_poll(session, poll, include_voters=True)
|
|
|
|
def _serialize_poll(self, session, poll, include_voters=False) -> dict:
|
|
options_data = []
|
|
total_votes = 0
|
|
for opt in poll.options:
|
|
vote_count = len(opt.votes)
|
|
total_votes += vote_count
|
|
voters = [v.voter_ip for v in opt.votes] if include_voters else None
|
|
options_data.append({
|
|
"id": opt.id,
|
|
"label": self._decrypt(opt.label),
|
|
"position": opt.position,
|
|
"votes": vote_count,
|
|
"voters": voters,
|
|
})
|
|
|
|
return {
|
|
"id": poll.id,
|
|
"channel": poll.channel,
|
|
"creator_ip": poll.creator_ip,
|
|
"question": self._decrypt(poll.question),
|
|
"created_at": poll.created_at,
|
|
"expires_at": poll.expires_at,
|
|
"is_multiple": poll.is_multiple,
|
|
"is_anonymous": poll.is_anonymous,
|
|
"status": poll.status,
|
|
"options": options_data,
|
|
"total_votes": total_votes,
|
|
}
|
|
|
|
def vote(self, poll_id: int, option_id: int, voter_ip: str) -> tuple[bool, str]:
|
|
with self.Session() as session:
|
|
poll = session.get(Poll, poll_id)
|
|
if not poll or poll.status != "active":
|
|
return False, "Poll is not active"
|
|
|
|
if poll.expires_at and int(time.time() * 1000) > poll.expires_at:
|
|
poll.status = "closed"
|
|
session.commit()
|
|
return False, "Poll has expired"
|
|
|
|
option_valid = any(opt.id == option_id for opt in poll.options)
|
|
if not option_valid:
|
|
return False, "Invalid option"
|
|
|
|
if not poll.is_multiple:
|
|
existing = (
|
|
session.query(PollVote)
|
|
.filter(PollVote.poll_id == poll_id, PollVote.voter_ip == voter_ip)
|
|
.first()
|
|
)
|
|
if existing:
|
|
return False, "You have already voted in this poll"
|
|
|
|
existing_opt = (
|
|
session.query(PollVote)
|
|
.filter(
|
|
PollVote.poll_id == poll_id,
|
|
PollVote.option_id == option_id,
|
|
PollVote.voter_ip == voter_ip,
|
|
)
|
|
.first()
|
|
)
|
|
if existing_opt:
|
|
return False, "You have already voted for this option"
|
|
|
|
vote = PollVote(
|
|
poll_id=poll_id,
|
|
option_id=option_id,
|
|
voter_ip=voter_ip,
|
|
voted_at=int(time.time() * 1000),
|
|
)
|
|
session.add(vote)
|
|
try:
|
|
session.commit()
|
|
except IntegrityError:
|
|
session.rollback()
|
|
return False, "You have already voted in this poll"
|
|
return True, "Vote recorded"
|
|
|
|
def close_poll(self, poll_id: int, requester_ip: str) -> tuple[bool, str]:
|
|
with self.Session() as session:
|
|
poll = session.get(Poll, poll_id)
|
|
if not poll:
|
|
return False, "Poll not found"
|
|
if poll.status != "active":
|
|
return False, "Poll is not active"
|
|
if poll.creator_ip != requester_ip:
|
|
return False, "Only the creator can close this poll"
|
|
poll.status = "closed"
|
|
session.commit()
|
|
return True, "Poll closed"
|