"""A2: сторож знака денег возврата на границе ``fact_returns``.

Тесты считают реальные выражения проекции в маленьком SQL-движке, поэтому
проверяют результат, а не наличие слова ``ABS`` в исходнике. Отрицательная
сумма не должна попасть в факт: либо она превращается в положительную величину,
либо источник блокируется до DML.
"""
from __future__ import annotations

from contextlib import contextmanager
from datetime import date
from decimal import Decimal
import re
import sqlite3
from typing import Any

import pytest

from kernel.etl import (
    build_fact_rating,
    build_fact_returns,
    build_fact_turnover,
    validate_invariants,
)

NEGATIVE_SOURCE_AMOUNT = -123.45


def _amount_expression(sql: str, marker: str) -> str:
    """Взять выражение, из которого конкретный план строит ``amount``."""
    start = sql.index(marker)
    end = sql.index(" AS amount", start)
    return sql[start:end].strip()


def _wb_amount_expression() -> str:
    """В WB-подготовке amount не имеет алиаса, потому что это INSERT в temp-таблицу."""
    sql = build_fact_returns.WB_PREPARE_SQLS[2]
    start = sql.index("       latest.order_id,") + len("       latest.order_id,")
    return sql[start:].lstrip().splitlines()[0].rstrip(",").strip()


def _ozon_invalid_price_expression() -> str | None:
    """Получить фактический предикат, который не допускает отрицательную цену Ozon."""
    sql = build_fact_returns.OZON_ATTEST_SQL
    marker = "current_return.unit_price IS NULL"
    if marker not in sql:
        return None
    start = sql.index(marker)
    end = sql.index("), 0) AS invalid_prices", start)
    return sql[start:end].strip()


def _sql_value(
    expression: str,
    rows: tuple[tuple[str, dict[str, Any]], ...],
) -> Any:
    """Вычислить производственное выражение на одной синтетической строке."""
    connection = sqlite3.connect(":memory:")
    try:
        from_parts: list[str] = []
        for index, (alias, row) in enumerate(rows):
            table = f"source_{index}"
            columns = ", ".join(f'"{name}" REAL' for name in row)
            connection.execute(f'CREATE TABLE "{table}" ({columns})')
            names = tuple(row)
            placeholders = ", ".join("?" for _ in names)
            connection.execute(
                f'INSERT INTO "{table}" ({", ".join(f"{name!r}" for name in names)}) '
                f"VALUES ({placeholders})",
                tuple(row[name] for name in names),
            )
            from_parts.append(f'"{table}" AS {alias}')

        from_clause = " CROSS JOIN ".join(from_parts)
        query = f"SELECT {expression} AS amount FROM {from_clause}"
        return connection.execute(query).fetchone()[0]
    finally:
        connection.close()


def _fact_amounts(
    amount_expression: str,
    rows: tuple[tuple[str, dict[str, Any]], ...],
    gate_expression: str | None = None,
) -> list[Any]:
    """Смоделировать только границу «гейт → INSERT amount», без подмены формулы."""
    if gate_expression is not None and bool(_sql_value(gate_expression, rows)):
        # Отрицательный источник признан непригодным и до INSERT не доходит.
        return []
    return [_sql_value(amount_expression, rows)]


@pytest.mark.parametrize(
    ("marketplace", "amount_expression", "rows", "gate_expression"),
    (
        pytest.param(
            "WB",
            _wb_amount_expression(),
            (("latest", {"for_pay": NEGATIVE_SOURCE_AMOUNT}),),
            None,
            id="wb",
        ),
        pytest.param(
            "Ozon",
            _amount_expression(
                build_fact_returns.OZON_SQL,
                "current_return.unit_price * current_return.quantity",
            ),
            (
                (
                    "current_return",
                    {"unit_price": NEGATIVE_SOURCE_AMOUNT, "quantity": 1},
                ),
            ),
            _ozon_invalid_price_expression(),
            id="ozon",
        ),
        pytest.param(
            "Lamoda",
            _amount_expression(
                build_fact_returns.LAMODA_PREPARE_SQLS[10],
                "ABS(CASE\n         WHEN COALESCE(fact_order.order_matches",
            ),
            (
                (
                    "fact_order",
                    {
                        "order_matches": 1,
                        "revenue": NEGATIVE_SOURCE_AMOUNT,
                        "quantity": 1,
                    },
                ),
                ("projected_return", {"quantity": 1}),
            ),
            None,
            id="lamoda",
        ),
        pytest.param(
            "YM",
            _amount_expression(
                build_fact_returns.YM_PREPARE_SQLS[11],
                "ABS(SUM(\n           (current_return.order_revenue / current_return.order_quantity)",
            ),
            (
                (
                    "current_return",
                    {
                        "order_revenue": NEGATIVE_SOURCE_AMOUNT,
                        "order_quantity": 1,
                        "item_count": 1,
                    },
                ),
            ),
            None,
            id="ym",
        ),
    ),
)
def test_negative_return_source_cannot_publish_negative_fact_amount(
    marketplace: str,
    amount_expression: str,
    rows: tuple[tuple[str, dict[str, Any]], ...],
    gate_expression: str | None,
) -> None:
    """Отрицательный источник не должен становиться отрицательным возвратом в факте."""
    amounts = _fact_amounts(amount_expression, rows, gate_expression)

    # Для WB проверяем не только знак: смена ABS на clamp/ноль тоже была бы
    # денежным регрессом, потому что возврат должен сохранить величину выплаты.
    if marketplace == "WB":
        assert amounts == [pytest.approx(abs(NEGATIVE_SOURCE_AMOUNT))]
    assert all(amount is not None and amount >= 0 for amount in amounts), marketplace


_NEGATIVE_AMOUNT_SHIELD = re.compile(
    r"(?m)^[ \t]*AND[ \t]+amount[ \t]*>=[ \t]*0[ \t]*$"
)


class _InvariantCursor:
    def __init__(self, row: dict[str, object]) -> None:
        self.row = row
        self.statements: list[str] = []

    def execute(self, sql: str, _params: object = None) -> None:
        self.statements.append(sql)

    def fetchone(self) -> dict[str, object]:
        return self.row


def test_negative_return_amount_is_registered_as_a_hard_invariant() -> None:
    checks = dict(
        (name, (severity, sql))
        for name, severity, sql in validate_invariants.CHECKS
    )
    severity, sql = checks["fact_returns_amount_nonnegative"]

    assert severity == "hard"
    assert "FROM gwptd_kernel.fact_returns" in sql
    assert "WHERE amount < 0" in sql
    assert "SUM(amount)" in sql
    assert "GROUP_CONCAT(DISTINCT mp_id" in sql


@pytest.mark.parametrize(
    ("module", "run"),
    (
        pytest.param(
            build_fact_turnover,
            lambda: build_fact_turnover.build(calc_date=date(2026, 8, 4)),
            id="turnover",
        ),
        pytest.param(
            build_fact_rating,
            lambda: build_fact_rating.build(
                month="2026-08",
                rate=Decimal("96"),
                as_of_date=date(2026, 8, 4),
            ),
            id="rating",
        ),
    ),
)
def test_negative_return_amount_aborts_metric_build_before_dml(
    monkeypatch: pytest.MonkeyPatch,
    module: Any,
    run: Any,
) -> None:
    """Оба money-reader-а падают до первого DELETE/UPSERT с диагностикой суммы и MP."""
    cursor = _InvariantCursor(
        {
            "violations": 129,
            "sample": "negative_rows=129; amount_sum=-1831668.94; mp_ids=1",
        }
    )

    @contextmanager
    def cursor_context():
        yield cursor

    monkeypatch.setattr(module, "get_cursor", cursor_context)

    with pytest.raises(
        validate_invariants.NegativeReturnAmountInvariantError,
        match=r"negative_rows=129; amount_sum=-1831668\.94; mp_ids=1",
    ):
        run()

    assert len(cursor.statements) == 1
    assert "fact_returns" in cursor.statements[0]
    assert all(
        "DELETE" not in statement.upper() and "INSERT" not in statement.upper()
        for statement in cursor.statements
    )


@pytest.mark.parametrize(
    ("module", "step_name", "run"),
    (
        pytest.param(
            build_fact_turnover,
            "_build_for_date",
            lambda: build_fact_turnover.build(calc_date=date(2026, 8, 4)),
            id="turnover",
        ),
        pytest.param(
            build_fact_rating,
            "_build_for_month",
            lambda: build_fact_rating.build(
                month="2026-08",
                rate=Decimal("96"),
                as_of_date=date(2026, 8, 4),
            ),
            id="rating",
        ),
    ),
)
def test_nonnegative_return_amount_allows_the_existing_metric_calculation(
    monkeypatch: pytest.MonkeyPatch,
    module: Any,
    step_name: str,
    run: Any,
) -> None:
    """Когда инвариант зелёный, builder продолжает обычный расчёт без нового пути."""
    cursor = _InvariantCursor(
        {
            "violations": 0,
            "sample": "negative_rows=0; amount_sum=0; mp_ids=none",
        }
    )
    called: list[tuple[tuple[object, ...], dict[str, object]]] = []

    @contextmanager
    def cursor_context():
        yield cursor

    def next_step(*args: object, **kwargs: object) -> int:
        called.append((args, kwargs))
        return 0

    monkeypatch.setattr(module, "get_cursor", cursor_context)
    monkeypatch.setattr(module, step_name, next_step)

    run()

    assert len(cursor.statements) == 1
    assert len(called) == 1


@pytest.mark.parametrize(
    ("name", "sql"),
    (
        pytest.param("rating", build_fact_rating._upsert_sql(1, (1, 2), "AND is_realized = 1")),
        pytest.param("turnover", build_fact_turnover._upsert_sql()),
    ),
)
def test_nonnegative_return_amounts_keep_the_previous_metric_total(
    name: str,
    sql: str,
) -> None:
    """На области, которую пропускает gate, снятый предикат не меняет сумму возвратов."""
    returns = ((1, 0.0), (2, 1_200.0), (3, 1_000.0))
    previous_total = (
        sum(quantity for quantity, amount in returns if amount >= 0),
        sum(amount for _quantity, amount in returns if amount >= 0),
    )
    current_total = (
        sum(quantity for quantity, _amount in returns),
        sum(amount for _quantity, amount in returns),
    )

    assert current_total == previous_total == (6, 2_200.0), name
    assert _NEGATIVE_AMOUNT_SHIELD.search(sql) is None, name
    assert re.search(r"SUM\(amount\)\s+AS\s+returns_amount", sql), name
