fix(dwh): align adapter sampling contract

This commit is contained in:
2026-07-11 20:14:23 +02:00
parent 717e5ecced
commit 216984aac8
11 changed files with 239 additions and 76 deletions
@@ -26,7 +26,7 @@
- Test: `harness/tests/test_dwh_port_contract.py`
**Interfaces:**
- Produces: `DwhCapabilities`, `DwhAdapter`, `DwhHealth`, and `UnsupportedCapability`.
- Produces: `DwhCapabilities`, `DwhAdapter`, `DwhHealth`, `DistinctValues`, and `UnsupportedCapability`.
- Consumes: existing catalog models from `tht.db.introspect` and execution result types from `tht.db.execute`.
- [ ] **Step 1: Write the failing protocol-shape test**
@@ -54,16 +54,21 @@ class DwhCapabilities:
sampling: bool = True
distinct_values: bool = True
@dataclass(frozen=True)
class DistinctValues:
values: list[object]
truncated: bool
@runtime_checkable
class DwhAdapter(Protocol):
@property
def capabilities(self) -> DwhCapabilities: ...
def health(self) -> DwhHealth: ...
def introspect(self) -> DatabaseCatalog: ...
def run_query(self, sql: str, *, limit: int | None = None) -> QueryResult: ...
def introspect(self) -> PhysicalSchema: ...
def run_query(self, sql: str, *, limit: int) -> ExecResult: ...
def explain(self, sql: str) -> PlanSummary: ...
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: ...
def distinct_values(self, table: str, column: str) -> list[object]: ...
def distinct_values(self, table: str, column: str) -> DistinctValues: ...
```
- [ ] **Step 4: Run contract test and type-oriented import smoke test**
@@ -97,8 +102,8 @@ git commit -m "refactor(dwh): define adapter contract"
```python
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
def test_adapter_rejects_write_sql(factory):
with pytest.raises(ReadOnlyViolation):
factory().run_query("delete from fact_sales")
with pytest.raises(ExecutionError):
factory().run_query("delete from fact_sales", limit=10)
```
- [ ] **Step 2: Verify failure**
@@ -111,12 +116,17 @@ Expected: FAIL because the adapter classes are absent.
```python
class PostgresDwhAdapter:
capabilities = DwhCapabilities()
def __init__(self, config: DatabaseConfig): self._config = config
def run_query(self, sql: str, *, limit: int | None = None) -> QueryResult:
return run_query(self._config, sql, limit=limit)
def __init__(self, config: DatabaseConfig):
self._config = config
self._engine = make_engine(config)
def run_query(self, sql: str, *, limit: int) -> ExecResult:
return run_query(self._engine, sql, limit=limit)
```
Implement the analogous REST wrapper by delegating to `tht.rest.*`; translate transport-specific errors only at the adapter boundary.
Implement the analogous REST wrapper by delegating to `tht.rest.*`; translate transport-specific
errors only at the adapter boundary. Both wrappers delegate frequency-ranked, distinct sampling to
the paired implementations in `tht.db.sampling`. A non-positive query limit is rejected, and
`distinct_values` reports any cap through `DistinctValues.truncated`.
- [ ] **Step 4: Run adapter, read-only, sampling, and REST tests**
+14 -1
View File
@@ -6,7 +6,7 @@ import pytest
from tht.config import LshConfig
from tht.db.introspect import introspect
from tht.db.sampling import is_text_type, unique_values_for_lsh
from tht.db.sampling import distinct_values, is_text_type, sample_column, unique_values_for_lsh
pytestmark = [pytest.mark.l0]
@@ -52,3 +52,16 @@ def test_unique_values_for_lsh_truncation_reported(admin_engine):
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
+86 -2
View File
@@ -1,8 +1,10 @@
import pytest
from sqlalchemy.exc import OperationalError
from tht.config import DatabaseConfig, RestConfig
from tht.execute import ExecutionError
from tht.ports import DwhAdapter
from tht.ports import DistinctValues, DwhAdapter
from tht.rest.client import RestError
def postgres_factory():
@@ -32,7 +34,7 @@ def rest_factory():
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
def test_adapter_rejects_write_sql_without_using_transport(factory):
with pytest.raises(ExecutionError, match="read-only enforcement"):
factory().run_query("delete from fact_sales")
factory().run_query("delete from fact_sales", limit=10)
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
@@ -40,3 +42,85 @@ def test_adapter_satisfies_dwh_protocol(factory):
adapter = factory()
assert isinstance(adapter, DwhAdapter)
assert adapter.capabilities.introspection is True
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
def test_run_query_requires_explicit_positive_limit(factory):
adapter = factory()
with pytest.raises(TypeError):
adapter.run_query("select 1")
with pytest.raises(ValueError, match="positive"):
adapter.run_query("select 1", limit=0)
def test_postgres_sampling_delegates_to_paired_sampling_functions(monkeypatch):
adapter = postgres_factory()
calls = []
expected = DistinctValues(values=["A"], truncated=True)
monkeypatch.setattr(
"tht.adapters.dwh.postgres.sampling.sample_column",
lambda engine, schema, table, column, *, limit: calls.append(
(engine, schema, table, column, limit)
)
or ["A", "B"],
)
monkeypatch.setattr(
"tht.adapters.dwh.postgres.sampling.distinct_values",
lambda engine, schema, table, column: expected,
)
assert adapter.sample_column("sales", "region", limit=2) == ["A", "B"]
assert calls == [(adapter._engine, "dw", "sales", "region", 2)]
assert adapter.distinct_values("sales", "region") is expected
def test_rest_sampling_delegates_and_translates_transport_errors(monkeypatch):
adapter = rest_factory()
expected = DistinctValues(values=["A", "B"], truncated=False)
monkeypatch.setattr(
"tht.adapters.dwh.thoth_rest.sampling.sample_column_rest",
lambda client, schema, table, column, *, limit: ["A", "B"],
)
monkeypatch.setattr(
"tht.adapters.dwh.thoth_rest.sampling.distinct_values_rest",
lambda client, schema, table, column: expected,
)
assert adapter.sample_column("sales", "region", limit=2) == ["A", "B"]
assert adapter.distinct_values("sales", "region") is expected
def fail(*args, **kwargs):
raise RestError("transport failed")
monkeypatch.setattr("tht.adapters.dwh.thoth_rest.sampling.sample_column_rest", fail)
monkeypatch.setattr("tht.adapters.dwh.thoth_rest.sampling.distinct_values_rest", fail)
with pytest.raises(ExecutionError, match="transport failed"):
adapter.sample_column("sales", "region", limit=2)
with pytest.raises(ExecutionError, match="transport failed"):
adapter.distinct_values("sales", "region")
def test_rest_distinct_values_reports_transport_truncation():
from tht.db.sampling import distinct_values_rest
class Client:
def top_values(self, schema, table, column, limit):
assert (schema, table, column, limit) == ("dw", "sales", "region", 3)
return [{"value": "A"}, {"value": "B"}, {"value": "C"}]
result = distinct_values_rest(Client(), "dw", "sales", "region", max_values=2)
assert result == DistinctValues(values=["A", "B"], truncated=True)
def test_postgres_health_only_normalizes_database_errors(monkeypatch):
adapter = postgres_factory()
database_error = OperationalError("select 1", {}, Exception("offline"))
monkeypatch.setattr("tht.adapters.dwh.postgres.ping", lambda engine: (_ for _ in ()).throw(database_error))
assert adapter.health().ok is False
monkeypatch.setattr(
"tht.adapters.dwh.postgres.ping",
lambda engine: (_ for _ in ()).throw(ValueError("programming bug")),
)
with pytest.raises(ValueError, match="programming bug"):
adapter.health()
+17 -2
View File
@@ -1,9 +1,14 @@
from dataclasses import FrozenInstanceError
import pytest
from tht.execute import ExecResult, PlanSummary
from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import (
DwhAdapter,
DwhCapabilities,
DwhHealth,
DistinctValues,
UnsupportedCapability,
)
@@ -17,7 +22,7 @@ class FakeDwhAdapter:
def introspect(self) -> PhysicalSchema:
raise NotImplementedError
def run_query(self, sql: str, *, limit: int | None = None) -> ExecResult:
def run_query(self, sql: str, *, limit: int) -> ExecResult:
raise NotImplementedError
def explain(self, sql: str) -> PlanSummary:
@@ -26,7 +31,7 @@ class FakeDwhAdapter:
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
raise NotImplementedError
def distinct_values(self, table: str, column: str) -> list[object]:
def distinct_values(self, table: str, column: str) -> DistinctValues:
raise NotImplementedError
@@ -45,3 +50,13 @@ def test_contract_types_are_public_and_capabilities_are_immutable():
assert capabilities.sampling is True
assert capabilities.distinct_values is True
assert issubclass(UnsupportedCapability, Exception)
with pytest.raises(FrozenInstanceError):
capabilities.explain = False
def test_all_contract_types_are_exported_from_public_package():
from tht.ports import DistinctValues as PublicDistinctValues
result = PublicDistinctValues(values=["a"], truncated=True)
assert result.values == ["a"]
assert result.truncated is True
+11 -7
View File
@@ -1,12 +1,14 @@
"""Direct PostgreSQL implementation of the DWH port."""
from tht.config import DatabaseConfig
from tht.db import execute
from sqlalchemy.exc import SQLAlchemyError
from tht.db import execute, sampling
from tht.db.connection import make_engine, ping
from tht.db.introspect import introspect
from tht.execute import ExecResult, PlanSummary
from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import DwhCapabilities, DwhHealth
from tht.ports.dwh import DistinctValues, DwhCapabilities, DwhHealth
class PostgresDwhAdapter:
@@ -19,23 +21,25 @@ class PostgresDwhAdapter:
def health(self) -> DwhHealth:
try:
ping(self._engine)
except Exception as exc:
except SQLAlchemyError as exc:
return DwhHealth(ok=False, detail=str(exc))
return DwhHealth(ok=True)
def introspect(self) -> PhysicalSchema:
return introspect(self._engine, self._config.database, self._config.db_schema)
def run_query(self, sql: str, *, limit: int | None = None) -> ExecResult:
def run_query(self, sql: str, *, limit: int) -> ExecResult:
return execute.run_query(self._engine, sql, limit=limit)
def explain(self, sql: str) -> PlanSummary:
return execute.explain(self._engine, sql)
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
return execute.sample_column(
return sampling.sample_column(
self._engine, self._config.db_schema, table, column, limit=limit
)
def distinct_values(self, table: str, column: str) -> list[object]:
return execute.distinct_values(self._engine, self._config.db_schema, table, column)
def distinct_values(self, table: str, column: str) -> DistinctValues:
return sampling.distinct_values(
self._engine, self._config.db_schema, table, column
)
+19 -17
View File
@@ -2,16 +2,12 @@
from tht.config import DatabaseConfig, RestConfig
from tht.db.introspect import introspect_rest
from tht.execute import ExecResult, PlanSummary
from tht.db import sampling
from tht.execute import ExecResult, ExecutionError, PlanSummary
from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import DwhCapabilities, DwhHealth
from tht.ports.dwh import DistinctValues, DwhCapabilities, DwhHealth
from tht.rest.client import RestClient, RestError
from tht.rest.execute import (
distinct_values_rest,
explain_rest,
run_controlled_rest,
sample_column_rest,
)
from tht.rest.execute import explain_rest, run_controlled_rest
class ThothRestDwhAdapter:
@@ -34,18 +30,24 @@ class ThothRestDwhAdapter:
self._client, self._database.database, self._database.db_schema
)
def run_query(self, sql: str, *, limit: int | None = None) -> ExecResult:
return run_controlled_rest(self._client, sql, limit=10 if limit is None else limit)
def run_query(self, sql: str, *, limit: int) -> ExecResult:
return run_controlled_rest(self._client, sql, limit=limit)
def explain(self, sql: str) -> PlanSummary:
return explain_rest(self._client, sql)
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
return sample_column_rest(
self._client, self._database.db_schema, table, column, limit=limit
)
try:
return sampling.sample_column_rest(
self._client, self._database.db_schema, table, column, limit=limit
)
except RestError as exc:
raise ExecutionError(str(exc)) from exc
def distinct_values(self, table: str, column: str) -> list[object]:
return distinct_values_rest(
self._client, self._database.db_schema, table, column
)
def distinct_values(self, table: str, column: str) -> DistinctValues:
try:
return sampling.distinct_values_rest(
self._client, self._database.db_schema, table, column
)
except RestError as exc:
raise ExecutionError(str(exc)) from exc
+4 -24
View File
@@ -4,39 +4,19 @@ from sqlalchemy import Engine
from tht.execute import ExecResult, PlanSummary, explain as _explain, run_controlled
DEFAULT_LIMIT = 10
DEFAULT_TIMEOUT_MS = 30_000
def run_query(engine: Engine, sql: str, *, limit: int | None = None) -> ExecResult:
def run_query(engine: Engine, sql: str, *, limit: int) -> ExecResult:
if limit <= 0:
raise ValueError("limit must be a positive integer")
return run_controlled(
engine,
sql,
limit=DEFAULT_LIMIT if limit is None else limit,
limit=limit,
timeout_ms=DEFAULT_TIMEOUT_MS,
)
def explain(engine: Engine, sql: str) -> PlanSummary:
return _explain(engine, sql, timeout_ms=DEFAULT_TIMEOUT_MS)
def sample_column(
engine: Engine, schema: str, table: str, column: str, *, limit: int
) -> list[object]:
quote = engine.dialect.identifier_preparer.quote
sql = (
f"SELECT {quote(column)} FROM {quote(schema)}.{quote(table)} "
f"WHERE {quote(column)} IS NOT NULL"
)
return [row[0] for row in run_query(engine, sql, limit=limit).rows]
def distinct_values(engine: Engine, schema: str, table: str, column: str) -> list[object]:
quote = engine.dialect.identifier_preparer.quote
sql = (
f"SELECT DISTINCT {quote(column)} FROM {quote(schema)}.{quote(table)} "
f"WHERE {quote(column)} IS NOT NULL ORDER BY {quote(column)}"
)
result = run_controlled(engine, sql, limit=1000, timeout_ms=DEFAULT_TIMEOUT_MS)
return [row[0] for row in result.rows]
+56
View File
@@ -5,10 +5,66 @@ from sqlalchemy import Engine, text
from tht.config import ExamplesConfig, LshConfig
from tht.mschema.models import Annotations, PhysicalSchema
from tht.ports.dwh import DistinctValues
logger = logging.getLogger(__name__)
TEXT_TYPE_PREFIXES = ("text", "varchar", "character", "char")
DEFAULT_DISTINCT_VALUES_LIMIT = 1000
def _quoted_top_values_query(engine: Engine, schema: str, table: str, column: str):
quote = engine.dialect.identifier_preparer.quote
identifier = quote(column)
return text(
f"SELECT {identifier} AS value FROM {quote(schema)}.{quote(table)} "
f"WHERE {identifier} IS NOT NULL GROUP BY {identifier} "
f"ORDER BY count(*) DESC, {identifier} LIMIT :lim"
)
def sample_column(
engine: Engine, schema: str, table: str, column: str, *, limit: int
) -> list[object]:
if limit <= 0:
raise ValueError("limit must be a positive integer")
query = _quoted_top_values_query(engine, schema, table, column)
with engine.connect() as conn:
rows = conn.execute(query, {"lim": limit}).fetchall()
return [row[0] for row in rows]
def sample_column_rest(
client, schema: str, table: str, column: str, *, limit: int
) -> list[object]:
if limit <= 0:
raise ValueError("limit must be a positive integer")
rows = client.top_values(schema, table, column, limit)
return [row["value"] for row in rows if row.get("value") is not None]
def distinct_values(
engine: Engine,
schema: str,
table: str,
column: str,
*,
max_values: int = DEFAULT_DISTINCT_VALUES_LIMIT,
) -> DistinctValues:
values = sample_column(engine, schema, table, column, limit=max_values + 1)
return DistinctValues(values=values[:max_values], truncated=len(values) > max_values)
def distinct_values_rest(
client,
schema: str,
table: str,
column: str,
*,
max_values: int = DEFAULT_DISTINCT_VALUES_LIMIT,
) -> DistinctValues:
values = sample_column_rest(client, schema, table, column, limit=max_values + 1)
return DistinctValues(values=values[:max_values], truncated=len(values) > max_values)
def is_text_type(pg_type: str) -> bool:
+2
View File
@@ -4,6 +4,7 @@ from tht.ports.dwh import (
DwhAdapter,
DwhCapabilities,
DwhHealth,
DistinctValues,
UnsupportedCapability,
)
@@ -11,5 +12,6 @@ __all__ = [
"DwhAdapter",
"DwhCapabilities",
"DwhHealth",
"DistinctValues",
"UnsupportedCapability",
]
+8 -2
View File
@@ -21,6 +21,12 @@ class DwhHealth:
detail: str | None = None
@dataclass(frozen=True)
class DistinctValues:
values: list[object]
truncated: bool
class UnsupportedCapability(Exception):
"""Raised when an adapter cannot provide an optional DWH operation."""
@@ -34,10 +40,10 @@ class DwhAdapter(Protocol):
def introspect(self) -> PhysicalSchema: ...
def run_query(self, sql: str, *, limit: int | None = None) -> ExecResult: ...
def run_query(self, sql: str, *, limit: int) -> ExecResult: ...
def explain(self, sql: str) -> PlanSummary: ...
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: ...
def distinct_values(self, table: str, column: str) -> list[object]: ...
def distinct_values(self, table: str, column: str) -> DistinctValues: ...
+2 -11
View File
@@ -17,6 +17,8 @@ from tht.rest.explain import parse_text_plan
def run_controlled_rest(client, sql: str, *, limit: int) -> ExecResult:
# Guard read-only client-side anche sul path REST (D7): non delegare l'unica verifica
# al server. Stesso check strutturale del path diretto.
if limit <= 0:
raise ValueError("limit must be a positive integer")
assert_read_only(sql)
final_sql, injected = _inject_limit(sql, limit)
start = time.monotonic()
@@ -46,14 +48,3 @@ def explain_rest(client, sql: str) -> PlanSummary:
except RestError as e:
raise ExecutionError(str(e)) from e
return parse_text_plan(lines)
def sample_column_rest(
client, schema: str, table: str, column: str, *, limit: int
) -> list[object]:
rows = client.top_values(schema, table, column, limit)
return [row["value"] for row in rows if row.get("value") is not None]
def distinct_values_rest(client, schema: str, table: str, column: str) -> list[object]:
return sample_column_rest(client, schema, table, column, limit=1000)