import pytest
from decimal import Decimal
from datetime import date
from unittest.mock import AsyncMock, MagicMock, patch
from app.services.rekonsiliasi_service import _compute_shift_breakdown


@pytest.mark.asyncio
async def test_compute_shift_breakdown_sums_per_shift():
    """Per-shift breakdown sums penjualan and penerimaan per tangki per shift."""
    s1 = MagicMock(id=1, nama="Shift 1")
    s3 = MagicMock(id=3, nama="Shift 3")
    tangki1 = MagicMock(id=10, nama="Tangki 1", produk=MagicMock(nama="Pertalite"), produk_id=1)
    period_pairs = [(s3.id, date(2026, 5, 6)), (s1.id, date(2026, 5, 7))]
    shifts_map = {s3.id: s3, s1.id: s1}
    tangki_list = [tangki1]
    db = AsyncMock()

    async def fake_penjualan(db, spbu_id, pairs, tangki_id):
        total = Decimal("0")
        for shift_id, d in pairs:
            if shift_id == s3.id:
                total += Decimal("100")
            elif shift_id == s1.id:
                total += Decimal("80")
        return total

    async def fake_penerimaan(db, spbu_id, pairs, tangki_id):
        return Decimal("5000") if any(sid == s3.id for sid, _ in pairs) else Decimal("0")

    async def fake_expected(db, spbu_id, pairs, tangki_id):
        return Decimal("5100") if any(sid == s3.id for sid, _ in pairs) else Decimal("0")

    with (
        patch("app.services.rekonsiliasi_service._get_penjualan_volume_period", fake_penjualan),
        patch("app.services.rekonsiliasi_service._get_penerimaan_volume_period", fake_penerimaan),
        patch("app.services.rekonsiliasi_service._get_penerimaan_expected_volume_period", fake_expected),
    ):
        result = await _compute_shift_breakdown(
            db=db,
            spbu_id=1,
            period_pairs=period_pairs,
            shifts_map=shifts_map,
            tangki_list=tangki_list,
        )

    assert len(result) == 2
    s3_breakdown = next(b for b in result if b.shift_id == s3.id)
    assert s3_breakdown.total_penjualan == Decimal("100")
    assert s3_breakdown.tangki[0].penerimaan_losses == Decimal("100")  # 5100 - 5000


@pytest.mark.asyncio
async def test_compute_shift_breakdown_empty_period_pairs():
    """Empty period_pairs returns empty list."""
    db = AsyncMock()
    result = await _compute_shift_breakdown(
        db=db,
        spbu_id=1,
        period_pairs=[],
        shifts_map={},
        tangki_list=[],
    )
    assert result == []


@pytest.mark.asyncio
async def test_compute_shift_breakdown_skips_unknown_shift():
    """shift_id not in shifts_map is skipped silently."""
    tangki1 = MagicMock(id=10, nama="Tangki 1", produk=None, produk_id=None)
    period_pairs = [(999, date(2026, 5, 7))]  # shift 999 not in map
    shifts_map = {}  # empty — no matching shift

    async def fake_zero(*args, **kwargs):
        return Decimal("0")

    db = AsyncMock()
    with (
        patch("app.services.rekonsiliasi_service._get_penjualan_volume_period", fake_zero),
        patch("app.services.rekonsiliasi_service._get_penerimaan_volume_period", fake_zero),
        patch("app.services.rekonsiliasi_service._get_penerimaan_expected_volume_period", fake_zero),
    ):
        result = await _compute_shift_breakdown(
            db=db,
            spbu_id=1,
            period_pairs=period_pairs,
            shifts_map=shifts_map,
            tangki_list=[tangki1],
        )

    assert result == []


@pytest.mark.asyncio
async def test_compute_shift_breakdown_tank_without_produk():
    """Tank with produk=None produces produk_nama=None in breakdown."""
    s1 = MagicMock(id=1, nama="Shift 1")
    tangki_no_produk = MagicMock(id=10, nama="Tangki 1", produk=None, produk_id=None)
    period_pairs = [(s1.id, date(2026, 5, 7))]
    shifts_map = {s1.id: s1}

    async def fake_zero(*args, **kwargs):
        return Decimal("0")

    db = AsyncMock()
    with (
        patch("app.services.rekonsiliasi_service._get_penjualan_volume_period", fake_zero),
        patch("app.services.rekonsiliasi_service._get_penerimaan_volume_period", fake_zero),
        patch("app.services.rekonsiliasi_service._get_penerimaan_expected_volume_period", fake_zero),
    ):
        result = await _compute_shift_breakdown(
            db=db,
            spbu_id=1,
            period_pairs=period_pairs,
            shifts_map=shifts_map,
            tangki_list=[tangki_no_produk],
        )

    assert len(result) == 1
    assert result[0].tangki[0].produk_nama is None
