File size: 5,019 Bytes
fa79ccd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""Unit tests for the post-load measure data-quality gate.

Run with: pytest tests/test_quality_gate.py -v
"""

import pytest

from demoprep_app.integrations.snowflake.quality_gate import (
    ZeroMeasureError,
    is_measure_column,
    run_measure_quality_gate,
)


class FakeCursor:
    """Answers the INFORMATION_SCHEMA query, then per-table aggregate queries."""

    def __init__(self, columns_rows, table_aggregates):
        # columns_rows: list of (table_name, column_name, data_type)
        # table_aggregates: {table_name: (row_count, nonzero_col1, nonzero_col2, ...)}
        self._columns_rows = columns_rows
        self._table_aggregates = table_aggregates
        self._last_result = None
        self.executed = []

    def execute(self, sql, params=None):
        self.executed.append(sql)
        if "INFORMATION_SCHEMA.COLUMNS" in sql:
            self._last_result = self._columns_rows
        else:
            table = next(t for t in self._table_aggregates if f'"{t}"' in sql)
            self._last_result = self._table_aggregates[table]

    def fetchall(self):
        return self._last_result

    def fetchone(self):
        return self._last_result

    def close(self):
        pass


class FakeConnection:
    def __init__(self, cursor):
        self._cursor = cursor

    def cursor(self):
        return self._cursor


def test_is_measure_column_heuristic():
    # Numeric non-key columns are measures
    assert is_measure_column("TOTAL_LOAD_REVENUE_USD", "NUMBER(12,2)")
    assert is_measure_column("BROKER_MARGIN_PCT", "FLOAT")
    # Join keys and codes are not measures
    assert not is_measure_column("LOAD_ID", "NUMBER(38,0)")
    assert not is_measure_column("LANE_KEY", "NUMBER(38,0)")
    assert not is_measure_column("CARRIER_CODE", "NUMBER(38,0)")
    # Non-numeric columns are not measures
    assert not is_measure_column("ORIGIN_CITY", "VARCHAR(100)")
    assert not is_measure_column("PICKUP_DATE", "DATE")


def test_gate_fails_loudly_on_all_zero_measure():
    cursor = FakeCursor(
        columns_rows=[
            ("FACT_LOAD_TRANSACTION", "LOAD_ID", "NUMBER"),
            ("FACT_LOAD_TRANSACTION", "LOADS_ACCEPTED", "NUMBER"),
            ("FACT_LOAD_TRANSACTION", "TOTAL_LOAD_REVENUE_USD", "NUMBER"),
            ("FACT_LOAD_TRANSACTION", "BROKER_MARGIN_USD", "NUMBER"),
        ],
        # 500 rows: LOADS_ACCEPTED populated, both derived measures all zero
        table_aggregates={"FACT_LOAD_TRANSACTION": (500, 500, 0, 0)},
    )

    with pytest.raises(ZeroMeasureError) as exc_info:
        run_measure_quality_gate(FakeConnection(cursor), "TRI_TEST_sch")

    err = exc_info.value
    assert len(err.failures) == 2
    failed_columns = {f["column"] for f in err.failures}
    assert failed_columns == {"TOTAL_LOAD_REVENUE_USD", "BROKER_MARGIN_USD"}
    # Error message must name the offending columns (fail loudly, no guessing)
    assert "FACT_LOAD_TRANSACTION.TOTAL_LOAD_REVENUE_USD" in str(err)
    assert "FACT_LOAD_TRANSACTION.BROKER_MARGIN_USD" in str(err)
    # ID column must not be profiled as a measure
    assert not any(f["column"] == "LOAD_ID" for f in err.failures)


def test_gate_passes_with_populated_measures():
    cursor = FakeCursor(
        columns_rows=[
            ("FACT_SALES", "SALE_ID", "NUMBER"),
            ("FACT_SALES", "REVENUE_USD", "NUMBER"),
            ("DIM_PRODUCT", "PRODUCT_ID", "NUMBER"),
            ("DIM_PRODUCT", "UNIT_PRICE", "NUMBER"),
        ],
        table_aggregates={
            "FACT_SALES": (500, 498),
            "DIM_PRODUCT": (50, 50),
        },
    )

    profile = run_measure_quality_gate(FakeConnection(cursor), "SALES_sch")
    assert profile["failures"] == []
    assert profile["tables_checked"] == 2
    assert profile["measure_columns_checked"] == 2
    assert profile["columns"]["FACT_SALES.REVENUE_USD"] == {
        "row_count": 500,
        "non_zero": 498,
    }


def test_gate_warns_on_mostly_zero_measure():
    cursor = FakeCursor(
        columns_rows=[("FACT_SALES", "DISCOUNT_USD", "NUMBER")],
        # 1 non-zero value in 500 rows (0.2%) — suspicious but not a failure
        table_aggregates={"FACT_SALES": (500, 1)},
    )

    profile = run_measure_quality_gate(FakeConnection(cursor), "SALES_sch")
    assert profile["failures"] == []
    assert len(profile["warnings"]) == 1
    assert "FACT_SALES.DISCOUNT_USD" in profile["warnings"][0]


def test_gate_warns_on_empty_table():
    cursor = FakeCursor(
        columns_rows=[("FACT_SALES", "REVENUE_USD", "NUMBER")],
        table_aggregates={"FACT_SALES": (0, 0)},
    )

    profile = run_measure_quality_gate(FakeConnection(cursor), "SALES_sch")
    assert profile["failures"] == []
    assert any("0 rows" in w for w in profile["warnings"])


def test_gate_errors_when_schema_has_no_columns():
    cursor = FakeCursor(columns_rows=[], table_aggregates={})
    with pytest.raises(RuntimeError, match="could not find any columns"):
        run_measure_quality_gate(FakeConnection(cursor), "MISSING_sch")