"""Offline, non-mutating unit tests. No database or site credentials needed."""
import copy
import unittest
from collections import Counter
from unittest.mock import patch

import heilind_sku_pdf as tool


def candidate(sku, maker, prod=4, conversion=12):
    return {"sku": sku, "manufacturer_key": maker, "manufacturer": maker,
            "prod_count": prod, "conversion_count": conversion,
            "added_attributes": conversion - prod}


class PureTests(unittest.TestCase):
    def test_url_example(self):
        self.assertEqual(tool.product_url("COC120X10149X", tool.DEFAULTS["capture"]),
                         "https://www.heilind.com/coc120x10149x.html")

    def test_url_escapes_reserved_characters(self):
        self.assertEqual(tool.product_url("ABC/A#B?C", tool.DEFAULTS["capture"]),
                         "https://www.heilind.com/abc%2Fa%23b%3Fc.html")

    def test_identifier_rejects_injection(self):
        for value in ("a;DROP TABLE x", "a.b", "a`b", "", None):
            with self.assertRaises(tool.ConfigError):
                tool.ident(value)

    def test_populated_values(self):
        empty = {"", "null", "[]", "{}"}
        for value in (None, "  ", "\t\r\n", " NuLl ", "[]", "{}"):
            self.assertFalse(tool.populated(value, empty))
        for value in (0, False, "0", "1,2,3", "N/A", "steel"):
            self.assertTrue(tool.populated(value, empty))

    def test_distinct_attribute_codes_not_rows_or_options(self):
        rows = [
            {"entity_id": 7, "attribute_id": 1, "raw_value": "1,2,3"},
            {"entity_id": 7, "attribute_id": 1, "raw_value": "1,2,3"},
            {"entity_id": 7, "attribute_id": 2, "raw_value": "4"},
            {"entity_id": 7, "attribute_id": 3, "raw_value": "   "},
            {"entity_id": 7, "attribute_id": 4, "raw_value": 0},
            {"entity_id": 7, "attribute_id": 5, "raw_value": "MOL"},
            {"entity_id": 7, "attribute_id": 999, "raw_value": "orphan"},
        ]
        codes = {1: "finish", 2: "finish", 3: "length", 4: "rated_current", 5: "brand"}
        actual = tool.attribute_sets(rows, codes, tool.DEFAULTS["selection"])
        self.assertEqual(actual[7], {"finish", "rated_current"})

    def test_thresholds(self):
        settings = tool.DEFAULTS["selection"]
        self.assertTrue(tool.qualifies(8, 18, settings))
        self.assertFalse(tool.qualifies(20, 25, settings))
        self.assertFalse(tool.qualifies(5, 8, settings))
        self.assertFalse(tool.qualifies(0, 7, settings))
        self.assertTrue(tool.qualifies(0, 8, settings))
        self.assertTrue(tool.qualifies(10, 15, settings))
        self.assertFalse(tool.qualifies(20, 10, settings))

    def test_25_unique_diverse_skus(self):
        rows = [candidate(f"M{maker:02d}_{i:02d}", f"M{maker:02d}", 3, 10+i)
                for maker in range(12) for i in range(6)]
        actual = tool.choose_diverse(rows + [rows[0]], 25, 3)
        self.assertEqual(len(actual), 25)
        self.assertEqual(len({r["sku"] for r in actual}), 25)
        counts = Counter(r["manufacturer_key"] for r in actual)
        self.assertEqual(len(counts), 12)
        self.assertLessEqual(max(counts.values()), 3)
        self.assertEqual(len({r["manufacturer_key"] for r in actual[:12]}), 12)

    def test_does_not_relax_diversity(self):
        rows = [candidate(f"MOL{i}", "MOL") for i in range(25)]
        self.assertEqual(len(tool.choose_diverse(rows, 25, 3)), 3)

    def test_more_than_25_manufacturers(self):
        rows = [candidate(f"M{i:02d}", f"M{i:02d}", 3, 10+i) for i in range(30)]
        chosen = tool.choose_diverse(rows, 25, 3)
        self.assertEqual(len({r["manufacturer_key"] for r in chosen}), 25)
        self.assertEqual(chosen[0]["sku"], "M29")

    def test_sampling_ranges_cover_once(self):
        for lo, hi, n in ((0, 103, 20), (1, 1, 20), (23, 29, 3)):
            flattened = [value for start, end in tool.ranges(lo, hi, n) for value in range(start, end+1)]
            self.assertEqual(flattened, list(range(lo, hi+1)))

    def test_filename_no_paths_and_no_collisions(self):
        a = tool.safe_filename(1, "../ABC/DEF")
        b = tool.safe_filename(1, ".._ABC_DEF")
        self.assertNotIn("/", a)
        self.assertNotEqual(a, b)

    def test_html_escapes_data(self):
        manifest = fake_manifest()
        manifest["selected"][0]["manufacturer"] = '<script>alert("x")</script>'
        rendered = tool.summary_html(manifest)
        self.assertNotIn('<script>alert("x")</script>', rendered)
        self.assertIn("&lt;script&gt;", rendered)


def fake_manifest():
    row = candidate("COC120X10149X", "Example maker")
    row.update(percent_increase=200, url=tool.product_url(row["sku"], tool.DEFAULTS["capture"]),
               capture_status="pending")
    return {"format_version": 1, "databases": tool.DEFAULTS["databases"],
            "manufacturer_source": "TEST FIXTURE ONLY", "selection_settings": tool.DEFAULTS["selection"],
            "capture_settings": tool.DEFAULTS["capture"], "selected": [row], "warnings": [],
            "scan": {"scanned_conversion_products": 1, "shared_products": 1,
                     "stop_reason": "synthetic unit-test fixture", "started_at": "test", "finished_at": "test"}}


class FakeConnection:
    def __init__(self):
        self.rollbacks = 0

    def rollback(self):
        self.rollbacks += 1


class FakeDatabase:
    def __init__(self, conn, name, schema):
        self.conn, self.name, self.s = conn, name, schema
        self.codes = {i: f"attribute_{i}" for i in range(1, 21)}

    def preflight(self):
        pass

    def columns(self, _):
        return [{"Field": "entity_id"}, {"Field": "sku"}, {"Field": "mfg_code"}]

    def bounds(self):
        return (0, 59)

    def product_batch(self, start, end, limit, mfg_column=None):
        return [{"entity_id": i, "sku": f"TEST{i:03d}", "mfg_raw": f"M{i % 12:02d}"}
                for i in range(start, min(end + 1, start + limit))]

    def by_skus(self, skus):
        # Deliberately use different entity IDs across the databases.
        return {sku.casefold(): {"sku": sku, "entity_id": int(sku[4:]) + 1000} for sku in skus}

    def values(self, ids):
        length = 12 if self.name == "ecom_conversion" else 4
        return [{"entity_id": entity, "attribute_id": aid, "raw_value": "1,2,3"}
                for entity in ids for aid in range(1, length + 1)]


class ScanTests(unittest.TestCase):
    def test_sql_values_are_bound_and_entity_batch_is_bounded(self):
        db = tool.Database(None, "ecom_prod", tool.SCHEMA)
        statements = []

        def collect(conn, sql, params=()):
            statements.append((sql, params))
            return []

        with patch.object(tool, "query", side_effect=collect):
            db.by_skus(["x' OR 1=1 --"])
            db.values([4, 8, 15])
            db.product_batch(100, 200, 25)
        self.assertNotIn("OR 1=1", statements[0][0])
        self.assertEqual(statements[0][1], ["x' OR 1=1 --"])
        self.assertIn("IN (%s,%s,%s)", statements[1][0])
        self.assertEqual(statements[1][1], [4, 8, 15])
        self.assertIn("LIMIT %s", statements[2][0])
        self.assertEqual(statements[2][1], (100, 200, 25))
        self.assertTrue(all(sql.startswith("SELECT ") for sql, _ in statements))

    def test_batch_scan_independent_entity_ids(self):
        conn = FakeConnection()
        config = copy.deepcopy(tool.DEFAULTS)
        config["selection"].update(batch_size=5, scan_limit=60, segments=6)
        statements = []
        with patch.object(tool, "Database", FakeDatabase), patch.object(tool, "query", side_effect=lambda conn, sql: statements.append(sql)):
            manifest = tool.scan(conn, config)
        self.assertEqual(len(manifest["selected"]), 25)
        self.assertGreaterEqual(len({r["manufacturer_key"] for r in manifest["selected"]}), 9)
        for row in manifest["selected"]:
            self.assertEqual(row["conversion_entity_id"] + 1000, row["prod_entity_id"])
            self.assertEqual((row["prod_count"], row["conversion_count"], row["added_attributes"]), (4, 12, 8))
        self.assertEqual(len(statements), conn.rollbacks)
        self.assertTrue(all(sql == "START TRANSACTION READ ONLY" for sql in statements))
        self.assertLessEqual(manifest["scan"]["scanned_conversion_products"], 60)

    def test_unknown_manufacturer_excluded(self):
        db = FakeDatabase(FakeConnection(), "ecom_conversion", tool.SCHEMA)
        source = tool.Manufacturer(db, tool.DEFAULTS["manufacturer"])
        result = source.for_products([
            {"entity_id": 1, "mfg_raw": None}, {"entity_id": 2, "mfg_raw": "unknown"},
            {"entity_id": 3, "mfg_raw": "MOL"}, {"entity_id": 4, "mfg_raw": "1,2"}], [])
        self.assertEqual(result, {3: ("mol", "MOL")})

    def test_mfg_eav_id_is_not_joined_globally(self):
        db = FakeDatabase(FakeConnection(), "ecom_conversion", tool.SCHEMA)
        db.codes[99] = "brand"
        cfg = copy.deepcopy(tool.DEFAULTS["manufacturer"])
        cfg.update(source="attribute", attribute_code="brand", label_lookup={
            "table": "eav_attribute_option_value", "id_column": "option_id", "label_column": "value"})
        with self.assertRaisesRegex(tool.ConfigError, "ownership"):
            tool.Manufacturer(db, cfg)


if __name__ == "__main__":
    unittest.main()
