import pytest

from app.services.infra.database import bootstrap_tenant, get_db_connection
from app.services.enrichment.job import enrich_tenant, select_products


def _insert(schema, key, name, description, content_hash):
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(
                f'INSERT INTO "{schema}".strategist_products '
                "(product_key, name, description, product_url, price_cents, "
                " content_hash) VALUES (%s,%s,%s,%s,%s,%s)",
                (key, name, description, f"https://x.example/{key}", 1000,
                 content_hash))
        conn.commit()
    finally:
        conn.close()


@pytest.fixture
def catalog(temp_tenant):
    bootstrap_tenant(temp_tenant)
    _insert(temp_tenant, "s:1", "Oxford Shirt", "A white cotton shirt.", "hash-a")
    _insert(temp_tenant, "s:2", "Belt", None, "hash-b")
    return temp_tenant


def test_unextracted_products_are_selected(catalog):
    keys = {p["product_key"] for p in select_products(catalog, force=False)}
    assert keys == {"s:1", "s:2"}


def test_an_already_extracted_product_is_skipped(catalog):
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(f'UPDATE "{catalog}".strategist_products '
                        "SET enriched_hash = content_hash WHERE product_key = 's:1'")
        conn.commit()
    finally:
        conn.close()

    keys = {p["product_key"] for p in select_products(catalog, force=False)}
    assert keys == {"s:2"}


def test_force_reselects_everything(catalog):
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(f'UPDATE "{catalog}".strategist_products '
                        "SET enriched_hash = content_hash")
        conn.commit()
    finally:
        conn.close()

    assert len(select_products(catalog, force=True)) == 2


def test_a_crawled_product_with_no_content_hash_is_extracted_only_once(temp_tenant,
                                                                      monkeypatch):
    # Crawled rows carry no content_hash. Before this was handled, they were
    # re-selected on every run, so a repeat enrich re-billed the whole crawled
    # set for nothing.
    import app.services.enrichment.job as mod
    bootstrap_tenant(temp_tenant)
    _insert(temp_tenant, "c:1", "Hand Woven Cotton Rug", "A rug.", None)

    monkeypatch.setattr(mod, "extract_batch",
                        lambda batch: {p["product_key"]: {"color": "white"}
                                       for p in batch})

    assert enrich_tenant(temp_tenant)["extracted"] == 1
    assert select_products(temp_tenant, force=False) == []
    assert enrich_tenant(temp_tenant)["extracted"] == 0


def test_an_unextractable_product_is_flagged(catalog, monkeypatch):
    import app.services.enrichment.job as mod
    monkeypatch.setattr(mod, "extract_batch",
                        lambda batch: {p["product_key"]: {"color": "white"}
                                       for p in batch})

    report = enrich_tenant(catalog)
    assert report["skipped_unextractable"] == 1

    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(f'SELECT missing_fields FROM "{catalog}".strategist_products '
                        "WHERE product_key = 's:2'")
            assert "attributes" in cur.fetchone()[0]
    finally:
        conn.close()


def test_extraction_writes_attributes_and_marks_the_hash(catalog, monkeypatch):
    import app.services.enrichment.job as mod
    monkeypatch.setattr(mod, "extract_batch",
                        lambda batch: {p["product_key"]:
                                       {"color": "white", "is_accessory": False}
                                       for p in batch})

    enrich_tenant(catalog)

    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(f'SELECT attributes, is_accessory, enriched_hash '
                        f'FROM "{catalog}".strategist_products '
                        "WHERE product_key = 's:1'")
            attributes, is_accessory, enriched_hash = cur.fetchone()
    finally:
        conn.close()

    assert any(a["key"] == "color" for a in attributes)
    assert is_accessory is False
    assert enriched_hash == "hash-a"
