import pandas as pd
import pytest

from engine.normalize import apply_steps, apply_derive


def s(*values):
    return pd.Series(list(values))


def test_trim_strips_and_drops_excel_float_suffix():
    out = apply_steps(s("  abc ", 123456.0), [{"step": "trim"}], {})
    assert out.tolist() == ["abc", "123456"]


def test_lpad_pads_short_values_only():
    out = apply_steps(s("12345", "1234567"), [{"step": "lpad", "length": 6, "fill": "0"}], {})
    assert out.tolist() == ["012345", "1234567"]


def test_last4_keeps_vpa_untouched():
    out = apply_steps(s("411111XXXXXX1111", "user@upi"), [{"step": "last4"}], {})
    assert out.tolist() == ["1111", "user@upi"]


def test_strip_chars_removes_all_tokens():
    out = apply_steps(s("1,234.50", "'0555'"), [{"step": "strip_chars", "chars": [",", "'"]}], {})
    assert out.tolist() == ["1234.50", "0555"]


def test_to_number_coerces_garbage_to_nan():
    out = apply_steps(s("500", "abc"), [{"step": "to_number"}], {})
    assert out.iloc[0] == 500
    assert pd.isna(out.iloc[1])


def test_date_parse_tries_formats_in_order():
    out = apply_steps(
        s("15/06/2026", "2026-06-16"),
        [{"step": "date_parse", "formats": ["%d/%m/%Y", "%Y-%m-%d"]}],
        {},
    )
    assert out.iloc[0] == pd.Timestamp("2026-06-15")
    assert out.iloc[1] == pd.Timestamp("2026-06-16")


def test_value_map_uses_context_location_map():
    ctx = {"location_map": {"SAHYADRI HOSPITAL BIBWEWADI": "BBW"}}
    out = apply_steps(
        s("SAHYADRI HOSPITAL BIBWEWADI", "Unknown"),
        [{"step": "value_map", "map_ref": "location_map"}],
        ctx,
    )
    assert out.tolist() == ["BBW", "Unknown"]


def test_upper_and_lower():
    assert apply_steps(s("Card"), [{"step": "upper"}], {}).tolist() == ["CARD"]
    assert apply_steps(s("User@UPI"), [{"step": "lower"}], {}).tolist() == ["user@upi"]


def test_coalesce_prefers_first_nonzero():
    df = pd.DataFrame({"domestic": ["500", "0", None], "intl": ["0", "750", "900"]})
    out = apply_derive(df, {"op": "coalesce", "from": ["domestic", "intl"]}, {})
    assert out.tolist() == ["500", "750", "900"]


def test_sum_derive():
    df = pd.DataFrame({"a": [1, 2], "b": [10, None]})
    out = apply_derive(df, {"op": "sum", "from": ["a", "b"]}, {})
    assert out.tolist() == [11.0, 2.0]


def test_unknown_step_raises():
    with pytest.raises(ValueError, match="Unknown normalization step"):
        apply_steps(s("x"), [{"step": "explode"}], {})
