84 lines
3.2 KiB
Python
84 lines
3.2 KiB
Python
"""L0: db/sampling against known data (testcontainers).
|
|
Verifies unique_values_for_lsh returns the expected most-frequent values for text
|
|
columns, and that wide_text / non-text columns are excluded.
|
|
"""
|
|
import pytest
|
|
|
|
from tht.config import LshConfig
|
|
from tht.db.introspect import introspect
|
|
from tht.db.sampling import distinct_values, is_text_type, sample_column, unique_values_for_lsh
|
|
|
|
pytestmark = [pytest.mark.l0]
|
|
|
|
|
|
@pytest.mark.parametrize("invalid_limit", [True, 1.5, 0, -1])
|
|
def test_direct_sampling_rejects_non_positive_integer_limits(admin_engine, invalid_limit):
|
|
with pytest.raises(ValueError, match="positive integer"):
|
|
sample_column(
|
|
admin_engine, "dw", "fct_ricoveri", "reparto", limit=invalid_limit
|
|
)
|
|
with pytest.raises(ValueError, match="positive integer"):
|
|
distinct_values(
|
|
admin_engine,
|
|
"dw",
|
|
"fct_ricoveri",
|
|
"reparto",
|
|
max_values=invalid_limit,
|
|
)
|
|
|
|
|
|
def test_is_text_type():
|
|
assert is_text_type("text")
|
|
assert is_text_type("varchar(100)")
|
|
assert is_text_type("character varying")
|
|
assert not is_text_type("integer")
|
|
assert not is_text_type("bigint")
|
|
assert not is_text_type("timestamp without time zone")
|
|
|
|
|
|
def test_unique_values_for_lsh_returns_most_frequent(admin_engine):
|
|
schema = introspect(admin_engine, "testdb", "dw")
|
|
# Before classify_all, all text columns are eligible=True by default. Sampling
|
|
# only touches text types regardless.
|
|
values, _skipped, _truncated = unique_values_for_lsh(
|
|
admin_engine, schema, LshConfig(max_values_per_column=100)
|
|
)
|
|
# dim_pazienti.citta: Milano, Bergamo, Brescia (3 distinct, all eligible text)
|
|
citta = values.get("dim_pazienti", {}).get("citta")
|
|
assert citta is not None
|
|
assert set(citta) == {"Milano", "Bergamo", "Brescia"}
|
|
|
|
|
|
def test_unique_values_for_lsh_excludes_non_text(admin_engine):
|
|
schema = introspect(admin_engine, "testdb", "dw")
|
|
values, _, _ = unique_values_for_lsh(
|
|
admin_engine, schema, LshConfig(max_values_per_column=100)
|
|
)
|
|
# id_paziente is bigint — must never appear in the LSH values.
|
|
assert "id_paziente" not in values.get("dim_pazienti", {})
|
|
|
|
|
|
def test_unique_values_for_lsh_truncation_reported(admin_engine):
|
|
schema = introspect(admin_engine, "testdb", "dw")
|
|
# Force a tiny cap so procedura/diagnosi columns (which have >2 distinct values)
|
|
# are reported as truncated rather than silently cut.
|
|
_, _, truncated = unique_values_for_lsh(
|
|
admin_engine, schema, LshConfig(max_values_per_column=1)
|
|
)
|
|
truncated_cols = {(t.table, t.column) for t in truncated}
|
|
# fct_ricoveri has several eligible text columns with distinct values
|
|
assert any(t[0] == "fct_ricoveri" for t in truncated_cols)
|
|
|
|
|
|
def test_adapter_sampling_is_distinct_and_frequency_ranked(admin_engine):
|
|
values = sample_column(admin_engine, "dw", "fct_ricoveri", "reparto", limit=2)
|
|
assert values == ["cardiologia", "pronto soccorso"]
|
|
|
|
|
|
def test_adapter_distinct_values_reports_truncation(admin_engine):
|
|
result = distinct_values(
|
|
admin_engine, "dw", "fct_ricoveri", "reparto", max_values=1
|
|
)
|
|
assert result.values == ["cardiologia"]
|
|
assert result.truncated is True
|