"""Add seed data from seed.sql with separate prepared statements Revision ID: 0001_initial_migration Revises: Create Date: 2025-08-22 23:40:52.885449 """ from pathlib import Path from typing import Sequence from typing import Union from alembic import op from sqlalchemy import Column from sqlalchemy import DateTime from sqlalchemy import ForeignKey from sqlalchemy import Integer from sqlalchemy import String from sqlalchemy import UniqueConstraint # revision identifiers, used by Alembic. revision: str = "0001" down_revision: Union[str, Sequence[str], None] = None branch_labels: Union[str, Sequence[str], None] = None depends_on: Union[str, Sequence[str], None] = None def upgrade() -> None: """Create the database schema and seed initial data from seed.sql.""" # Create tables for users, rooms, bookings, invitees op.create_table( "users", Column("email", String, primary_key=True), Column("name", String, nullable=False), ) op.create_table( "rooms", Column("id", Integer, primary_key=True, autoincrement=True), Column("name", String, nullable=False, unique=True), Column("location", String, nullable=False), Column("equipment", String, nullable=False), Column("capacity", Integer, nullable=False), ) op.create_table( "bookings", Column("id", Integer, primary_key=True, autoincrement=True), Column("room_id", Integer, ForeignKey("rooms.id")), Column("start_time", DateTime(timezone=True), nullable=False), Column("end_time", DateTime(timezone=True), nullable=False), Column("title", String, nullable=True), ) op.create_table( "invitees", Column("id", Integer, primary_key=True, autoincrement=True), Column( "booking_id", Integer, ForeignKey("bookings.id", onupdate="CASCADE", ondelete="CASCADE"), ), Column( "user_email", String, ForeignKey("users.email", onupdate="CASCADE", ondelete="CASCADE"), ), UniqueConstraint("booking_id", "user_email", name="uq_invitee_booking_user"), ) # Seed data from seed.sql migration_dir = Path(__file__).parent.parent.resolve() seed_file_path = migration_dir / "seed.sql" with seed_file_path.open("r") as file: sql_content = file.read() sql_statements = [stmt.strip() for stmt in sql_content.split(";") if stmt.strip()] for stmt in sql_statements: op.execute(stmt) def downgrade() -> None: """Downgrade schema by removing data from users and rooms tables.""" op.execute("DELETE FROM users") op.execute("DELETE FROM rooms")