FuckingChat/dbworker.py

191 lines
6.6 KiB
Python

import sqlalchemy
from sqlalchemy import create_engine, String, Integer, Text, DateTime
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, Session
from sqlalchemy.orm import sessionmaker
from datetime import datetime
import time
import os
from cryptography.fernet import Fernet
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
import base64
# Encryption setup
SALT = b'salt_123456789'
PASSWORD = b'chat_secret_key_123'
def get_cipher():
kdf = PBKDF2HMAC(
algorithm=hashes.SHA256(),
length=32,
salt=SALT,
iterations=100000,
)
key = base64.urlsafe_b64encode(kdf.derive(PASSWORD))
return Fernet(key)
class Base(DeclarativeBase):
pass
class Messages(Base):
__tablename__ = "messages"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
ip: Mapped[str] = mapped_column(String(64), nullable=False)
content: Mapped[str] = mapped_column(Text, nullable=False)
timestamp: Mapped[int] = mapped_column(Integer, nullable=False)
# MariaDB configuration
DB_USER = 'chat'
DB_PASSWORD = 'uqhyUb5eBGg3qad.'
DB_HOST = 'localhost'
DB_PORT = '3306'
DB_NAME = 'chat_db'
# Create engine for MariaDB
DATABASE_URL = f"mariadb+mariadbconnector://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
# Alternative with pymysql (uncomment if needed):
# DATABASE_URL = f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
engine = create_engine(
DATABASE_URL,
echo=False,
pool_size=10,
max_overflow=20,
pool_pre_ping=True,
pool_recycle=3600
)
# Create tables
Base.metadata.create_all(engine)
class DBfrontend():
def __init__(self):
self.Session = sessionmaker(bind=engine)
self.cipher = get_cipher()
def encrypt_message(self, message: str) -> str:
"""Encrypt message before storing"""
return self.cipher.encrypt(message.encode()).decode()
def decrypt_message(self, encrypted_message: str) -> str:
"""Decrypt message after retrieval"""
try:
return self.cipher.decrypt(encrypted_message.encode()).decode()
except:
return "[Decryption failed]"
def add_message(self, ip: str, content: str) -> int:
"""Add a new message to the database and return its ID"""
with self.Session() as session:
encrypted_content = self.encrypt_message(content)
new_message = Messages(
ip=ip,
content=encrypted_content,
timestamp=int(time.time() * 1000)
)
session.add(new_message)
session.commit()
return new_message.id
def get_all_messages(self) -> list:
"""Get all messages from the database"""
with self.Session() as session:
messages = session.query(Messages).all()
for msg in messages:
msg.content = self.decrypt_message(msg.content)
return messages
def get_message_by_id(self, message_id: int) -> Messages | None:
"""Get a specific message by its ID"""
with self.Session() as session:
message = session.query(Messages).get(message_id)
if message:
message.content = self.decrypt_message(message.content)
return message
def get_messages_by_ip(self, ip: str) -> list:
"""Get all messages from a specific IP address"""
with self.Session() as session:
messages = session.query(Messages).filter(Messages.ip == ip).all()
for msg in messages:
msg.content = self.decrypt_message(msg.content)
return messages
def update_message(self, message_id: int, new_content: str) -> bool:
"""Update the content of a message"""
with self.Session() as session:
message = session.query(Messages).get(message_id)
if message:
message.content = self.encrypt_message(new_content)
session.commit()
return True
return False
def delete_message(self, message_id: int) -> bool:
"""Delete a message by its ID"""
with self.Session() as session:
message = session.query(Messages).get(message_id)
if message:
session.delete(message)
session.commit()
return True
return False
def delete_messages_by_ip(self, ip: str) -> int:
"""Delete all messages from a specific IP address and return count"""
with self.Session() as session:
messages = session.query(Messages).filter(Messages.ip == ip).all()
count = len(messages)
for message in messages:
session.delete(message)
session.commit()
return count
def get_message_count(self) -> int:
"""Get total number of messages"""
with self.Session() as session:
return session.query(Messages).count()
def get_latest_messages(self, limit: int = 10) -> list:
"""Get the latest N messages"""
with self.Session() as session:
messages = session.query(Messages).order_by(Messages.id.desc()).limit(limit).all()
for msg in messages:
msg.content = self.decrypt_message(msg.content)
return messages
def clear_all_messages(self) -> int:
"""Delete all messages and return count"""
with self.Session() as session:
count = session.query(Messages).count()
session.query(Messages).delete()
session.commit()
return count
def get_messages_by_date_range(self, start_timestamp: int, end_timestamp: int) -> list:
"""Get messages within a timestamp range"""
with self.Session() as session:
messages = session.query(Messages).filter(
Messages.timestamp >= start_timestamp,
Messages.timestamp <= end_timestamp
).order_by(Messages.timestamp.desc()).all()
for msg in messages:
msg.content = self.decrypt_message(msg.content)
return messages
# Optional: Test connection function
def test_connection():
"""Test if MariaDB connection is working"""
try:
with engine.connect() as conn:
result = conn.execute(sqlalchemy.text("SELECT 1"))
print("✅ MariaDB connection successful!")
return True
except Exception as e:
print(f"❌ MariaDB connection failed: {e}")
return False
# Run test if executed directly
if __name__ == "__main__":
test_connection()