"""R7 feature tests: many-to-one aggregation matching + date-mismatch warnings."""

import pandas as pd

from engine import pipeline
from engine.matcher import run_stage
from helpers import csv_b64


def agg_stage(tolerance=None):
    tier = {"name": "batch", "keys": [["settlement_ref", "settlement_ref"]]}
    if tolerance:
        tier["amount_tolerance"] = tolerance
    return {
        "stage": "batch_settlement",
        "match_mode": "many_to_one",
        "side_a": {"source": "template:a"},
        "side_b": {"source": "template:b"},
        "tiers": [tier],
        "output": {"matched_sheet": "M", "unmatched_sheet": "U"},
    }


def test_many_to_one_sums_group_to_single_settlement():
    # 3 transactions under REF1 sum to 600 = one settlement row;
    # REF2's transactions sum to 999 ≠ settlement 1000 → unmatched.
    df_a = pd.DataFrame({
        "settlement_ref": ["REF1", "REF1", "REF1", "REF2", "REF2"],
        "amount": [100.0, 200.0, 300.0, 499.0, 500.0],
    })
    df_b = pd.DataFrame({
        "settlement_ref": ["REF1", "REF2"],
        "amount": [600.0, 1000.0],
    })

    result = run_stage(df_a, df_b, agg_stage(), {})

    assert len(result["matched"]) == 3          # all REF1 transactions
    assert set(result["matched"]["settlement_ref"]) == {"REF1"}
    assert result["matched"]["match_tier"].unique().tolist() == ["batch"]
    # matched_b repeats the settlement row per member transaction
    assert len(result["matched_b"]) == 3
    assert result["matched_b"]["amount"].unique().tolist() == [600.0]
    assert len(result["unmatched_a"]) == 2      # REF2 transactions
    assert result["unmatched_b"]["settlement_ref"].tolist() == ["REF2"]


def test_many_to_one_with_tolerance():
    df_a = pd.DataFrame({"settlement_ref": ["R1", "R1"], "amount": [100.0, 200.0]})
    df_b = pd.DataFrame({"settlement_ref": ["R1"], "amount": [300.4]})

    exact = run_stage(df_a, df_b, agg_stage(), {})
    assert len(exact["matched"]) == 0

    tolerant = run_stage(
        df_a, df_b,
        agg_stage(tolerance={"type": "abs", "value": 1.0, "fields": ["amount", "amount"]}),
        {},
    )
    assert len(tolerant["matched"]) == 2


def test_one_settlement_row_consumed_once():
    # Two groups with identical totals must not share one settlement row.
    df_a = pd.DataFrame({
        "settlement_ref": ["R1", "R2"],
        "amount": [500.0, 500.0],
    })
    df_b = pd.DataFrame({"settlement_ref": ["R1"], "amount": [500.0]})
    result = run_stage(df_a, df_b, agg_stage(), {})
    assert len(result["matched"]) == 1
    assert result["matched"]["settlement_ref"].tolist() == ["R1"]


TEMPLATE_DATED = {
    "name": "dated_src",
    "display_name": "Dated source",
    "source_role": "settlement_bank",
    "file_format": "csv",
    "fields": [
        {"canonical": "ref", "aliases": ["ref"]},
        {"canonical": "amount", "aliases": ["amount"], "normalize": [{"step": "to_number"}]},
        {"canonical": "txn_date", "aliases": ["date"], "dtype": "date",
         "normalize": [{"step": "date_parse", "formats": ["%Y-%m-%d"]}]},
    ],
}

SIMPLE_RULESET = {
    "name": "warn_test",
    "stages": [{
        "stage": "s1",
        "side_a": {"source": "template:dated_src"},
        "side_b": {"source": "template:dated_src"},
        "tiers": [{"name": "t", "keys": [["ref", "ref"]]}],
        "output": {"matched_sheet": "M", "unmatched_sheet": "U"},
    }],
}


def test_date_mismatch_warning_when_recon_date_differs():
    df = pd.DataFrame({
        "ref": ["A", "B", "C"],
        "amount": [1, 2, 3],
        "date": ["2026-06-09", "2026-06-09", "2026-06-09"],
    })
    files = {"dated_src": csv_b64(df)}
    templates = {"dated_src": TEMPLATE_DATED}

    result = pipeline.run(SIMPLE_RULESET, templates, files, {"recon_date": "2026-06-10"})
    assert len(result["warnings"]) == 1
    assert "2026-06-09" in result["warnings"][0]
    assert "2026-06-10" in result["warnings"][0]

    ok = pipeline.run(SIMPLE_RULESET, templates, files, {"recon_date": "2026-06-09"})
    assert ok["warnings"] == []


def test_multi_date_warning_without_recon_date():
    df = pd.DataFrame({
        "ref": ["A", "B", "C", "D"],
        "amount": [1, 2, 3, 4],
        "date": ["2026-06-09", "2026-06-10", "2026-06-11", "2026-06-12"],
    })
    result = pipeline.run(
        SIMPLE_RULESET, {"dated_src": TEMPLATE_DATED}, {"dated_src": csv_b64(df)}, {}
    )
    assert len(result["warnings"]) == 1
    assert "span multiple dates" in result["warnings"][0]
