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