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` - Test: `harness/tests/test_dwh_port_contract.py`
**Interfaces:** **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`. - Consumes: existing catalog models from `tht.db.introspect` and execution result types from `tht.db.execute`.
- [ ] **Step 1: Write the failing protocol-shape test** - [ ] **Step 1: Write the failing protocol-shape test**
@@ -54,16 +54,21 @@ class DwhCapabilities:
sampling: bool = True sampling: bool = True
distinct_values: bool = True distinct_values: bool = True
@dataclass(frozen=True)
class DistinctValues:
values: list[object]
truncated: bool
@runtime_checkable @runtime_checkable
class DwhAdapter(Protocol): class DwhAdapter(Protocol):
@property @property
def capabilities(self) -> DwhCapabilities: ... def capabilities(self) -> DwhCapabilities: ...
def health(self) -> DwhHealth: ... def health(self) -> DwhHealth: ...
def introspect(self) -> DatabaseCatalog: ... def introspect(self) -> PhysicalSchema: ...
def run_query(self, sql: str, *, limit: int | None = None) -> QueryResult: ... def run_query(self, sql: str, *, limit: int) -> ExecResult: ...
def explain(self, sql: str) -> PlanSummary: ... def explain(self, sql: str) -> PlanSummary: ...
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: ... 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** - [ ] **Step 4: Run contract test and type-oriented import smoke test**
@@ -97,8 +102,8 @@ git commit -m "refactor(dwh): define adapter contract"
```python ```python
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory]) @pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
def test_adapter_rejects_write_sql(factory): def test_adapter_rejects_write_sql(factory):
with pytest.raises(ReadOnlyViolation): with pytest.raises(ExecutionError):
factory().run_query("delete from fact_sales") factory().run_query("delete from fact_sales", limit=10)
``` ```
- [ ] **Step 2: Verify failure** - [ ] **Step 2: Verify failure**
@@ -111,12 +116,17 @@ Expected: FAIL because the adapter classes are absent.
```python ```python
class PostgresDwhAdapter: class PostgresDwhAdapter:
capabilities = DwhCapabilities() capabilities = DwhCapabilities()
def __init__(self, config: DatabaseConfig): self._config = config def __init__(self, config: DatabaseConfig):
def run_query(self, sql: str, *, limit: int | None = None) -> QueryResult: self._config = config
return run_query(self._config, sql, limit=limit) 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** - [ ] **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.config import LshConfig
from tht.db.introspect import introspect 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] 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} truncated_cols = {(t.table, t.column) for t in truncated}
# fct_ricoveri has several eligible text columns with distinct values # fct_ricoveri has several eligible text columns with distinct values
assert any(t[0] == "fct_ricoveri" for t in truncated_cols) 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 import pytest
from sqlalchemy.exc import OperationalError
from tht.config import DatabaseConfig, RestConfig from tht.config import DatabaseConfig, RestConfig
from tht.execute import ExecutionError 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(): def postgres_factory():
@@ -32,7 +34,7 @@ def rest_factory():
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory]) @pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
def test_adapter_rejects_write_sql_without_using_transport(factory): def test_adapter_rejects_write_sql_without_using_transport(factory):
with pytest.raises(ExecutionError, match="read-only enforcement"): 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]) @pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
@@ -40,3 +42,85 @@ def test_adapter_satisfies_dwh_protocol(factory):
adapter = factory() adapter = factory()
assert isinstance(adapter, DwhAdapter) assert isinstance(adapter, DwhAdapter)
assert adapter.capabilities.introspection is True 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.execute import ExecResult, PlanSummary
from tht.mschema.models import PhysicalSchema from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import ( from tht.ports.dwh import (
DwhAdapter, DwhAdapter,
DwhCapabilities, DwhCapabilities,
DwhHealth, DwhHealth,
DistinctValues,
UnsupportedCapability, UnsupportedCapability,
) )
@@ -17,7 +22,7 @@ class FakeDwhAdapter:
def introspect(self) -> PhysicalSchema: def introspect(self) -> PhysicalSchema:
raise NotImplementedError raise NotImplementedError
def run_query(self, sql: str, *, limit: int | None = None) -> ExecResult: def run_query(self, sql: str, *, limit: int) -> ExecResult:
raise NotImplementedError raise NotImplementedError
def explain(self, sql: str) -> PlanSummary: def explain(self, sql: str) -> PlanSummary:
@@ -26,7 +31,7 @@ class FakeDwhAdapter:
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
raise NotImplementedError raise NotImplementedError
def distinct_values(self, table: str, column: str) -> list[object]: def distinct_values(self, table: str, column: str) -> DistinctValues:
raise NotImplementedError raise NotImplementedError
@@ -45,3 +50,13 @@ def test_contract_types_are_public_and_capabilities_are_immutable():
assert capabilities.sampling is True assert capabilities.sampling is True
assert capabilities.distinct_values is True assert capabilities.distinct_values is True
assert issubclass(UnsupportedCapability, Exception) 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.""" """Direct PostgreSQL implementation of the DWH port."""
from tht.config import DatabaseConfig 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.connection import make_engine, ping
from tht.db.introspect import introspect from tht.db.introspect import introspect
from tht.execute import ExecResult, PlanSummary from tht.execute import ExecResult, PlanSummary
from tht.mschema.models import PhysicalSchema from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import DwhCapabilities, DwhHealth from tht.ports.dwh import DistinctValues, DwhCapabilities, DwhHealth
class PostgresDwhAdapter: class PostgresDwhAdapter:
@@ -19,23 +21,25 @@ class PostgresDwhAdapter:
def health(self) -> DwhHealth: def health(self) -> DwhHealth:
try: try:
ping(self._engine) ping(self._engine)
except Exception as exc: except SQLAlchemyError as exc:
return DwhHealth(ok=False, detail=str(exc)) return DwhHealth(ok=False, detail=str(exc))
return DwhHealth(ok=True) return DwhHealth(ok=True)
def introspect(self) -> PhysicalSchema: def introspect(self) -> PhysicalSchema:
return introspect(self._engine, self._config.database, self._config.db_schema) 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) return execute.run_query(self._engine, sql, limit=limit)
def explain(self, sql: str) -> PlanSummary: def explain(self, sql: str) -> PlanSummary:
return execute.explain(self._engine, sql) return execute.explain(self._engine, sql)
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: 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 self._engine, self._config.db_schema, table, column, limit=limit
) )
def distinct_values(self, table: str, column: str) -> list[object]: def distinct_values(self, table: str, column: str) -> DistinctValues:
return execute.distinct_values(self._engine, self._config.db_schema, table, column) return sampling.distinct_values(
self._engine, self._config.db_schema, table, column
)
+15 -13
View File
@@ -2,16 +2,12 @@
from tht.config import DatabaseConfig, RestConfig from tht.config import DatabaseConfig, RestConfig
from tht.db.introspect import introspect_rest 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.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.client import RestClient, RestError
from tht.rest.execute import ( from tht.rest.execute import explain_rest, run_controlled_rest
distinct_values_rest,
explain_rest,
run_controlled_rest,
sample_column_rest,
)
class ThothRestDwhAdapter: class ThothRestDwhAdapter:
@@ -34,18 +30,24 @@ class ThothRestDwhAdapter:
self._client, self._database.database, self._database.db_schema self._client, self._database.database, self._database.db_schema
) )
def run_query(self, sql: str, *, limit: int | None = None) -> ExecResult: def run_query(self, sql: str, *, limit: int) -> ExecResult:
return run_controlled_rest(self._client, sql, limit=10 if limit is None else limit) return run_controlled_rest(self._client, sql, limit=limit)
def explain(self, sql: str) -> PlanSummary: def explain(self, sql: str) -> PlanSummary:
return explain_rest(self._client, sql) return explain_rest(self._client, sql)
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
return sample_column_rest( try:
return sampling.sample_column_rest(
self._client, self._database.db_schema, table, column, limit=limit 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]: def distinct_values(self, table: str, column: str) -> DistinctValues:
return distinct_values_rest( try:
return sampling.distinct_values_rest(
self._client, self._database.db_schema, table, column 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 from tht.execute import ExecResult, PlanSummary, explain as _explain, run_controlled
DEFAULT_LIMIT = 10
DEFAULT_TIMEOUT_MS = 30_000 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( return run_controlled(
engine, engine,
sql, sql,
limit=DEFAULT_LIMIT if limit is None else limit, limit=limit,
timeout_ms=DEFAULT_TIMEOUT_MS, timeout_ms=DEFAULT_TIMEOUT_MS,
) )
def explain(engine: Engine, sql: str) -> PlanSummary: def explain(engine: Engine, sql: str) -> PlanSummary:
return _explain(engine, sql, timeout_ms=DEFAULT_TIMEOUT_MS) 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.config import ExamplesConfig, LshConfig
from tht.mschema.models import Annotations, PhysicalSchema from tht.mschema.models import Annotations, PhysicalSchema
from tht.ports.dwh import DistinctValues
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
TEXT_TYPE_PREFIXES = ("text", "varchar", "character", "char") 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: def is_text_type(pg_type: str) -> bool:
+2
View File
@@ -4,6 +4,7 @@ from tht.ports.dwh import (
DwhAdapter, DwhAdapter,
DwhCapabilities, DwhCapabilities,
DwhHealth, DwhHealth,
DistinctValues,
UnsupportedCapability, UnsupportedCapability,
) )
@@ -11,5 +12,6 @@ __all__ = [
"DwhAdapter", "DwhAdapter",
"DwhCapabilities", "DwhCapabilities",
"DwhHealth", "DwhHealth",
"DistinctValues",
"UnsupportedCapability", "UnsupportedCapability",
] ]
+8 -2
View File
@@ -21,6 +21,12 @@ class DwhHealth:
detail: str | None = None detail: str | None = None
@dataclass(frozen=True)
class DistinctValues:
values: list[object]
truncated: bool
class UnsupportedCapability(Exception): class UnsupportedCapability(Exception):
"""Raised when an adapter cannot provide an optional DWH operation.""" """Raised when an adapter cannot provide an optional DWH operation."""
@@ -34,10 +40,10 @@ class DwhAdapter(Protocol):
def introspect(self) -> PhysicalSchema: ... 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 explain(self, sql: str) -> PlanSummary: ...
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]: ... 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: 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 # Guard read-only client-side anche sul path REST (D7): non delegare l'unica verifica
# al server. Stesso check strutturale del path diretto. # al server. Stesso check strutturale del path diretto.
if limit <= 0:
raise ValueError("limit must be a positive integer")
assert_read_only(sql) assert_read_only(sql)
final_sql, injected = _inject_limit(sql, limit) final_sql, injected = _inject_limit(sql, limit)
start = time.monotonic() start = time.monotonic()
@@ -46,14 +48,3 @@ def explain_rest(client, sql: str) -> PlanSummary:
except RestError as e: except RestError as e:
raise ExecutionError(str(e)) from e raise ExecutionError(str(e)) from e
return parse_text_plan(lines) 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)