import os
from pathlib import Path

import psycopg
from dotenv import load_dotenv
from psycopg import sql
from psycopg.errors import SyntaxError as PgSyntaxError

load_dotenv(Path(__file__).resolve().parent / ".env")

PG_CONFIG = {
    "host": os.getenv("DB_HOST", "127.0.0.1"),
    "port": int(os.getenv("DB_PORT", "5432")),
    "dbname": os.getenv("DB_NAME", "paython_db"),
    "user": os.getenv("DB_USER", "postgres"),
    "password": os.getenv("DB_PASSWORD", ""),
}


def main():
    conn = psycopg.connect(**PG_CONFIG)
    conn.autocommit = True
    try:
        with conn.cursor() as cur:
            cur.execute(
                """
                SELECT
                    c.table_name,
                    c.data_type
                FROM information_schema.columns c
                WHERE c.table_schema = 'public'
                  AND c.column_name = 'id'
                  AND c.data_type IN ('integer', 'bigint')
                  AND c.is_nullable = 'NO'
                  AND c.column_default IS NULL
                ORDER BY c.table_name
                """
            )
            rows = cur.fetchall()

            for table_name, _ in rows:
                seq_name = f"{table_name}_id_seq"
                cur.execute(
                    sql.SQL("CREATE SEQUENCE IF NOT EXISTS {}").format(
                        sql.Identifier(seq_name)
                    )
                )
                cur.execute(
                    sql.SQL(
                        "SELECT COALESCE(MAX(id), 0) FROM {}"
                    ).format(sql.Identifier(table_name))
                )
                max_id = cur.fetchone()[0] or 0
                next_id = max_id + 1

                cur.execute(
                    sql.SQL("SELECT setval({}, {}, false)").format(
                        sql.Literal(seq_name),
                        sql.Literal(next_id),
                    )
                )
                try:
                    cur.execute(
                        sql.SQL(
                            "ALTER TABLE {} ALTER COLUMN id SET DEFAULT nextval({})"
                        ).format(
                            sql.Identifier(table_name),
                            sql.Literal(seq_name),
                        )
                    )
                    print(f"Fixed sequence for: {table_name} (next id: {next_id})")
                except PgSyntaxError:
                    # Identity columns already manage their own sequence.
                    print(f"Skipped identity id column: {table_name}")

    finally:
        conn.close()


if __name__ == "__main__":
    main()
