164 lines
6.3 KiB
Python
164 lines
6.3 KiB
Python
import pytest
|
|
from sqlalchemy.exc import OperationalError
|
|
|
|
from tht.config import DatabaseConfig, RestConfig
|
|
from tht.db.sampling import distinct_values_rest, sample_column_rest
|
|
from tht.execute import ExecutionError
|
|
from tht.ports import DistinctValues, DwhAdapter
|
|
from tht.rest.client import RestError
|
|
from tht.adapters.dwh import PostgresDwhAdapter
|
|
|
|
|
|
def postgres_factory():
|
|
from tht.adapters.dwh import PostgresDwhAdapter
|
|
|
|
return PostgresDwhAdapter(
|
|
DatabaseConfig(database="analytics", schema="dw", user="reader", password="secret")
|
|
)
|
|
|
|
|
|
def rest_factory():
|
|
from tht.adapters.dwh import ThothRestDwhAdapter
|
|
|
|
database = DatabaseConfig(
|
|
database="analytics",
|
|
schema="dw",
|
|
user="unused",
|
|
password="unused",
|
|
transport="rest",
|
|
)
|
|
return ThothRestDwhAdapter(
|
|
database,
|
|
RestConfig(base_url="https://dwh.example.test", api_key="secret"),
|
|
)
|
|
|
|
|
|
@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", limit=10)
|
|
|
|
|
|
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
|
|
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])
|
|
@pytest.mark.parametrize("invalid_limit", [True, 1.5, 0, -1])
|
|
def test_run_query_rejects_non_positive_integer_limit(factory, invalid_limit):
|
|
adapter = factory()
|
|
with pytest.raises(ValueError, match="positive integer"):
|
|
adapter.run_query("select 1", limit=invalid_limit)
|
|
|
|
|
|
@pytest.mark.parametrize("factory", [postgres_factory, rest_factory])
|
|
def test_run_query_requires_explicit_limit(factory):
|
|
with pytest.raises(TypeError):
|
|
factory().run_query("select 1")
|
|
|
|
|
|
@pytest.mark.parametrize("invalid_limit", [True, 1.5, 0, -1])
|
|
def test_rest_sampling_rejects_non_positive_integer_limit(invalid_limit):
|
|
class Client:
|
|
def top_values(self, *args):
|
|
raise AssertionError("transport must not be used")
|
|
|
|
with pytest.raises(ValueError, match="positive integer"):
|
|
sample_column_rest(Client(), "dw", "sales", "region", limit=invalid_limit)
|
|
with pytest.raises(ValueError, match="positive integer"):
|
|
distinct_values_rest(
|
|
Client(), "dw", "sales", "region", max_values=invalid_limit
|
|
)
|
|
|
|
|
|
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"],
|
|
)
|
|
distinct_calls = []
|
|
monkeypatch.setattr(
|
|
"tht.adapters.dwh.postgres.sampling.distinct_values",
|
|
lambda engine, schema, table, column, *, max_values: distinct_calls.append(max_values)
|
|
or 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", limit=17) is expected
|
|
assert distinct_calls == [17]
|
|
|
|
|
|
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, *, max_values: expected,
|
|
)
|
|
assert adapter.sample_column("sales", "region", limit=2) == ["A", "B"]
|
|
assert adapter.distinct_values("sales", "region", limit=17) 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", limit=17)
|
|
|
|
|
|
def test_rest_distinct_values_reports_transport_truncation():
|
|
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()
|
|
|
|
|
|
def test_non_default_timeout_reaches_query_and_explain(monkeypatch):
|
|
adapter = PostgresDwhAdapter(
|
|
DatabaseConfig(database="analytics", schema="dw", user="reader", password="secret"),
|
|
statement_timeout_ms=12_345,
|
|
)
|
|
calls = []
|
|
monkeypatch.setattr("tht.adapters.dwh.postgres.execute.run_query",
|
|
lambda engine, sql, *, limit, timeout_ms: calls.append(("run", timeout_ms)))
|
|
monkeypatch.setattr("tht.adapters.dwh.postgres.execute.explain",
|
|
lambda engine, sql, *, timeout_ms: calls.append(("explain", timeout_ms)))
|
|
adapter.run_query("select 1", limit=2)
|
|
adapter.explain("select 1")
|
|
assert calls == [("run", 12_345), ("explain", 12_345)]
|