import pytest

import app.services.enrichment.job as enrichment
import app.services.pairing.job as pairing


class Recorder:
    def __init__(self):
        self.calls = []

    def __call__(self, percent, step):
        self.calls.append((percent, step))

    @property
    def percents(self):
        return [p for p, _ in self.calls]


def _products(n):
    return [{"product_key": f"k{i}", "name": f"Product Number {i}",
             "description": "A description.", "attributes": [],
             "price_cents": 100, "taxonomy_path": [], "brand": None,
             "content_hash": f"h{i}"} for i in range(n)]


@pytest.fixture
def stub_enrichment(monkeypatch):
    monkeypatch.setattr(enrichment, "migrate_products_table", lambda t: None)
    monkeypatch.setattr(enrichment, "compute_bands", lambda t: (100, 200))
    monkeypatch.setattr(enrichment, "persist_attributes", lambda *a, **kw: None)
    monkeypatch.setattr(enrichment, "flag_unextractable", lambda *a, **kw: None)
    monkeypatch.setattr(enrichment, "catalog_stats", lambda t: (45, 45))
    monkeypatch.setattr(enrichment, "extract_batch",
                        lambda batch: {p["product_key"]: {"color": "white"}
                                       for p in batch})
    monkeypatch.setattr(enrichment, "select_products",
                        lambda t, f: _products(45))


def test_enrichment_reports_progress(stub_enrichment):
    recorder = Recorder()
    enrichment.enrich_tenant("org_x", progress=recorder)
    assert recorder.calls


def test_percent_never_goes_backwards(stub_enrichment):
    # A bar that reverses teaches a user to distrust it.
    recorder = Recorder()
    enrichment.enrich_tenant("org_x", progress=recorder)
    assert recorder.percents == sorted(recorder.percents)


def test_percent_ends_at_one_hundred(stub_enrichment):
    recorder = Recorder()
    enrichment.enrich_tenant("org_x", progress=recorder)
    assert recorder.percents[-1] == 100


def test_percent_stays_in_range(stub_enrichment):
    recorder = Recorder()
    enrichment.enrich_tenant("org_x", progress=recorder)
    assert all(0 <= p <= 100 for p in recorder.percents)


def test_every_step_is_a_readable_string(stub_enrichment):
    # The step is shown to a person; "phase 2" tells them nothing.
    recorder = Recorder()
    enrichment.enrich_tenant("org_x", progress=recorder)
    assert all(isinstance(s, str) and s.strip() for _, s in recorder.calls)


def test_enrichment_works_with_no_callback(stub_enrichment):
    # Every existing test calls it this way.
    report = enrichment.enrich_tenant("org_x")
    assert report["extracted"] == 45


def test_a_failing_callback_does_not_kill_the_job(stub_enrichment):
    # Reporting must never abort work that is spending money.
    def boom(percent, step):
        raise RuntimeError("redis is down")

    report = enrichment.enrich_tenant("org_x", progress=boom)
    assert report["extracted"] == 45


def test_pairing_reports_progress_through_its_phases(monkeypatch):
    monkeypatch.setattr(pairing, "migrate_pairing_tables", lambda t: None)
    monkeypatch.setattr(pairing, "load_eligible", lambda t: [])
    recorder = Recorder()
    pairing.pair_tenant("org_x", progress=recorder)
    assert recorder.percents == sorted(recorder.percents)
    assert recorder.percents[-1] == 100


def test_pairing_does_not_report_the_same_percentage_twice(monkeypatch):
    # The live catalogue produces 9,187 candidate pairs across 45 distinct
    # percentages. Reporting per pair would be 9,187 Redis round trips to say
    # the same 45 things.
    products = [{"product_key": f"k{i}", "name": f"Product {i}",
                 "category": "shirts", "is_accessory": False,
                 "price_cents": 1000 + i, "rating": 4.0, "attributes": [],
                 "tenant_relations": {}} for i in range(40)]

    monkeypatch.setattr(pairing, "migrate_pairing_tables", lambda t: None)
    monkeypatch.setattr(pairing, "load_eligible", lambda t: products)
    monkeypatch.setattr(pairing, "vectors_for", lambda t, p: {})
    monkeypatch.setattr(pairing, "replace_neighbors", lambda t, pairs: None)
    monkeypatch.setattr(pairing, "load_decisions", lambda t: {})

    recorder = Recorder()
    pairing.pair_tenant("org_x", progress=recorder)

    scoring = [p for p in recorder.percents if 45 <= p < 90]
    assert len(scoring) == len(set(scoring))
    assert len(recorder.calls) < 120
