"""Tests for database connection utilities."""

import os
from pathlib import Path
from unittest.mock import MagicMock, patch

import pytest
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session

from quber.db.connection import (
    get_database_url,
    get_engine,
    get_session,
    get_session_factory,
    init_db,
)
from quber.settings import get_settings


@pytest.fixture(autouse=True)
def isolate_settings(monkeypatch: pytest.MonkeyPatch, tmp_path: Path):
    # get_database_url now reads the cached get_settings().db; clear the cache
    # and chdir to an empty dir so each test sees only its patched os.environ
    # (no leftover cache, no real .env).
    monkeypatch.chdir(tmp_path)
    get_settings.cache_clear()
    yield
    get_settings.cache_clear()


class TestGetDatabaseUrl:
    """Tests for get_database_url function."""

    def test_default_values(self):
        """Test database URL with default environment variables."""
        with patch.dict(os.environ, {}, clear=True):
            url = get_database_url()

            assert url == "postgresql+psycopg://quber:quber_dev@localhost:5432/quber_rag"

    def test_custom_values(self):
        """Test database URL with custom environment variables."""
        env_vars = {
            "POSTGRES_USER": "custom_user",
            "POSTGRES_PASSWORD": "custom_pass",
            "POSTGRES_HOST": "db.example.com",
            "POSTGRES_PORT": "5433",
            "POSTGRES_DB": "custom_db",
        }

        with patch.dict(os.environ, env_vars, clear=True):
            url = get_database_url()

            assert url == "postgresql+psycopg://custom_user:custom_pass@db.example.com:5433/custom_db"

    def test_partial_custom_values(self):
        """Test database URL with some custom environment variables."""
        env_vars = {
            "POSTGRES_USER": "testuser",
            "POSTGRES_HOST": "testhost",
        }

        with patch.dict(os.environ, env_vars, clear=True):
            url = get_database_url()

            # Should use custom user and host, defaults for rest
            assert url == "postgresql+psycopg://testuser:quber_dev@testhost:5432/quber_rag"


class TestGetEngine:
    """Tests for get_engine function."""

    @patch("quber.db.connection.create_engine")
    @patch("quber.db.connection.get_database_url")
    def test_get_engine_default(self, mock_get_url, mock_create_engine):
        """Test creating engine with default settings."""
        mock_get_url.return_value = "postgresql+psycopg://test"
        mock_engine = MagicMock()
        mock_engine.url.database = "test_db"
        mock_engine.url.host = "localhost"
        mock_create_engine.return_value = mock_engine

        engine = get_engine()

        mock_get_url.assert_called_once()
        mock_create_engine.assert_called_once_with(
            "postgresql+psycopg://test", echo=False, pool_pre_ping=True
        )
        assert engine == mock_engine

    @patch("quber.db.connection.create_engine")
    @patch("quber.db.connection.get_database_url")
    def test_get_engine_with_echo(self, mock_get_url, mock_create_engine):
        """Test creating engine with echo=True."""
        mock_get_url.return_value = "postgresql+psycopg://test"
        mock_engine = MagicMock()
        mock_engine.url.database = "test_db"
        mock_engine.url.host = "localhost"
        mock_create_engine.return_value = mock_engine

        get_engine(echo=True)

        mock_create_engine.assert_called_once_with("postgresql+psycopg://test", echo=True, pool_pre_ping=True)


class TestInitDb:
    """Tests for init_db function."""

    @patch("quber.db.connection.Base")
    @patch("quber.db.connection.get_engine")
    def test_init_db_without_engine(self, mock_get_engine, mock_base):
        """Test init_db creates a new engine when none provided."""
        mock_engine = MagicMock(spec=Engine)
        mock_get_engine.return_value = mock_engine
        mock_metadata = MagicMock()
        mock_base.metadata = mock_metadata

        init_db()

        mock_get_engine.assert_called_once()
        mock_metadata.create_all.assert_called_once_with(mock_engine)
        mock_metadata.drop_all.assert_not_called()

    @patch("quber.db.connection.Base")
    def test_init_db_with_engine(self, mock_base):
        """Test init_db uses provided engine."""
        mock_engine = MagicMock(spec=Engine)
        mock_metadata = MagicMock()
        mock_base.metadata = mock_metadata

        init_db(engine=mock_engine)

        mock_metadata.create_all.assert_called_once_with(mock_engine)

    @patch("quber.db.connection.Base")
    def test_init_db_with_drop_all(self, mock_base):
        """Test init_db drops tables when drop_all=True."""
        mock_engine = MagicMock(spec=Engine)
        mock_metadata = MagicMock()
        mock_base.metadata = mock_metadata

        init_db(engine=mock_engine, drop_all=True)

        mock_metadata.drop_all.assert_called_once_with(mock_engine)
        mock_metadata.create_all.assert_called_once_with(mock_engine)


class TestGetSessionFactory:
    """Tests for get_session_factory function."""

    @patch("quber.db.connection.sessionmaker")
    @patch("quber.db.connection.get_engine")
    def test_get_session_factory_without_engine(self, mock_get_engine, mock_sessionmaker):
        """Test creating session factory without engine."""
        mock_engine = MagicMock(spec=Engine)
        mock_get_engine.return_value = mock_engine
        mock_factory = MagicMock()
        mock_sessionmaker.return_value = mock_factory

        factory = get_session_factory()

        mock_get_engine.assert_called_once()
        mock_sessionmaker.assert_called_once_with(bind=mock_engine, expire_on_commit=False)
        assert factory == mock_factory

    @patch("quber.db.connection.sessionmaker")
    def test_get_session_factory_with_engine(self, mock_sessionmaker):
        """Test creating session factory with provided engine."""
        mock_engine = MagicMock(spec=Engine)
        mock_factory = MagicMock()
        mock_sessionmaker.return_value = mock_factory

        factory = get_session_factory(engine=mock_engine)

        mock_sessionmaker.assert_called_once_with(bind=mock_engine, expire_on_commit=False)
        assert factory == mock_factory


class TestGetSession:
    """Tests for get_session context manager."""

    @patch("quber.db.connection.sessionmaker")
    @patch("quber.db.connection.get_engine")
    def test_get_session_success(self, mock_get_engine, mock_sessionmaker):
        """Test get_session context manager successful execution."""
        mock_engine = MagicMock(spec=Engine)
        mock_get_engine.return_value = mock_engine

        mock_session = MagicMock(spec=Session)
        mock_factory = MagicMock(return_value=mock_session)
        mock_sessionmaker.return_value = mock_factory

        with get_session() as session:
            assert session == mock_session

        mock_session.commit.assert_called_once()
        mock_session.close.assert_called_once()
        mock_session.rollback.assert_not_called()

    @patch("quber.db.connection.sessionmaker")
    @patch("quber.db.connection.get_engine")
    def test_get_session_with_exception(self, mock_get_engine, mock_sessionmaker):
        """Test get_session context manager with exception."""
        mock_engine = MagicMock(spec=Engine)
        mock_get_engine.return_value = mock_engine

        mock_session = MagicMock(spec=Session)
        mock_factory = MagicMock(return_value=mock_session)
        mock_sessionmaker.return_value = mock_factory

        with pytest.raises(ValueError):
            with get_session():
                raise ValueError("Test error")

        mock_session.rollback.assert_called_once()
        mock_session.close.assert_called_once()
        mock_session.commit.assert_not_called()

    @patch("quber.db.connection.sessionmaker")
    def test_get_session_with_provided_engine(self, mock_sessionmaker):
        """Test get_session with provided engine."""
        mock_engine = MagicMock(spec=Engine)
        mock_session = MagicMock(spec=Session)
        mock_factory = MagicMock(return_value=mock_session)
        mock_sessionmaker.return_value = mock_factory

        with get_session(engine=mock_engine) as session:
            assert session == mock_session

        mock_sessionmaker.assert_called_once_with(bind=mock_engine, expire_on_commit=False)
