fix(core): restore adapter command behavior

This commit is contained in:
2026-07-11 21:04:53 +02:00
parent a35efa16de
commit 45d57756aa
7 changed files with 141 additions and 29 deletions
@@ -0,0 +1,84 @@
from types import SimpleNamespace
import pytest
import typer
from tht.cli import db_cmd
from tht.cli.lsh_cmd import _extract_lsh_values
from tht.mschema.models import Annotations, ColumnPhysical, PhysicalSchema, TablePhysical
from tht.ports.dwh import DistinctValues, DwhHealth
from tht.cli import memory_cmd
from tht.memory import MemoryRecord
from datetime import datetime
def _ping(monkeypatch, health, capsys):
monkeypatch.setattr(db_cmd, "load_config", lambda path: SimpleNamespace(database=SimpleNamespace(user="u")))
monkeypatch.setattr(db_cmd, "build_dwh", lambda cfg: SimpleNamespace(health=lambda: health))
try:
db_cmd.ping_cmd()
except typer.Exit as exc:
code = exc.exit_code
else:
code = 0
return code, capsys.readouterr()
def test_db_ping_public_health_success(monkeypatch, capsys):
code, output = _ping(monkeypatch, DwhHealth(ok=True, database="d", schema="s", read_only=True), capsys)
assert code == 0
assert "OK: connesso a d (schema s)" in output.out
def test_db_ping_rest_inaccessible_historical_wording(monkeypatch, capsys):
code, output = _ping(monkeypatch, DwhHealth(ok=False, detail="{'db_connected': False}", error_kind="inaccessible"), capsys)
assert code == 1
assert "ERRORE: DWH non accessibile via REST (risposta: {'db_connected': False})." in output.err
def test_db_ping_direct_connection_historical_wording(monkeypatch, capsys):
code, output = _ping(monkeypatch, DwhHealth(ok=False, detail="connection refused", error_kind="connection"), capsys)
assert code == 1
assert "ERRORE di connessione: connection refused" in output.err
@pytest.mark.parametrize(("limit", "truncated"), [(7, False), (1201, True)])
def test_lsh_extraction_honors_configured_limit(limit, truncated):
physical = PhysicalSchema(database="d", schema="s", introspected_at=datetime(2026, 1, 1), tables={
"t": TablePhysical(columns={"c": ColumnPhysical(type="text", eligible=True)})
})
calls = []
class Dwh:
def distinct_values(self, table, column, *, limit):
calls.append(limit)
return DistinctValues(values=list(range(limit)), truncated=truncated)
values, _, reports = _extract_lsh_values(Dwh(), physical, Annotations(), limit)
assert calls == [limit]
assert len(values["t"]["c"]) == limit
assert [report.indexed for report in reports] == ([limit] if truncated else [])
def test_memory_command_writes_through_factory_vector_store(monkeypatch):
store = SimpleNamespace(existing_hashes=lambda *args: {}, upsert=lambda table, rows: 1)
captured = []
original_upsert = store.upsert
store.upsert = lambda table, rows: captured.extend(rows) or original_upsert(table, rows)
cfg = SimpleNamespace(embeddings=object(), vector_write_rest=object())
manifest = SimpleNamespace(id="s1")
record = MemoryRecord(id="m1", ts=datetime(2026, 1, 1), session_id="s1",
decision_seq=7, type="table_promoted", subject="t",
question_context="q")
monkeypatch.setattr(memory_cmd, "_load_config_or_exit", lambda path: cfg)
monkeypatch.setattr(memory_cmd, "load_session_or_exit", lambda cfg, session: manifest)
monkeypatch.setattr(memory_cmd, "require_vector_write_allowed", lambda *args: None)
monkeypatch.setattr(memory_cmd, "has_vector_write_rest", lambda cfg: True)
monkeypatch.setattr(memory_cmd, "session_dir", lambda *args: None)
monkeypatch.setattr(memory_cmd, "registry_path", lambda cfg: None)
monkeypatch.setattr("tht.adapters.factory.build_vector_store", lambda cfg, require_write: store)
monkeypatch.setattr("tht.cli.vector_cmd.make_embedder",
lambda cfg: SimpleNamespace(embed_documents=lambda texts: [[0.1]]))
monkeypatch.setattr("tht.memory.promote", lambda *args, **kwargs: None)
monkeypatch.setattr("tht.memory.load_registry", lambda path: [record])
memory_cmd.save_one_cmd(session="s1", decision=7, json_out=True)
from tht.ports.vector import VectorWriteRecord
assert len(captured) == 1 and isinstance(captured[0], VectorWriteRecord)
+16
View File
@@ -6,6 +6,7 @@ from tht.db.sampling import distinct_values_rest, sample_column_rest
from tht.execute import ExecutionError from tht.execute import ExecutionError
from tht.ports import DistinctValues, DwhAdapter from tht.ports import DistinctValues, DwhAdapter
from tht.rest.client import RestError from tht.rest.client import RestError
from tht.adapters.dwh import PostgresDwhAdapter
def postgres_factory(): def postgres_factory():
@@ -145,3 +146,18 @@ def test_postgres_health_only_normalizes_database_errors(monkeypatch):
) )
with pytest.raises(ValueError, match="programming bug"): with pytest.raises(ValueError, match="programming bug"):
adapter.health() 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)]
+4 -2
View File
@@ -1,7 +1,7 @@
"""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 sqlalchemy.exc import SQLAlchemyError from sqlalchemy.exc import OperationalError, SQLAlchemyError
from tht.db import execute, sampling from tht.db import execute, sampling
from tht.db.connection import can_create_in_schema, make_engine, ping, writable_tables from tht.db.connection import can_create_in_schema, make_engine, ping, writable_tables
@@ -22,8 +22,10 @@ class PostgresDwhAdapter:
def health(self) -> DwhHealth: def health(self) -> DwhHealth:
try: try:
ping(self._engine) ping(self._engine)
except OperationalError as exc:
return DwhHealth(ok=False, detail=str(exc.orig), error_kind="connection")
except SQLAlchemyError as exc: except SQLAlchemyError as exc:
return DwhHealth(ok=False, detail=str(exc)) return DwhHealth(ok=False, detail=str(exc), error_kind="connection")
writable = tuple(writable_tables(self._engine, self._config.db_schema)) writable = tuple(writable_tables(self._engine, self._config.db_schema))
can_create = can_create_in_schema(self._engine, self._config.db_schema) can_create = can_create_in_schema(self._engine, self._config.db_schema)
return DwhHealth(ok=True, database=self._config.database, schema=self._config.db_schema, return DwhHealth(ok=True, database=self._config.database, schema=self._config.db_schema,
+3 -2
View File
@@ -21,11 +21,12 @@ class ThothRestDwhAdapter:
try: try:
result = self._client.ping() result = self._client.ping()
except RestError as exc: except RestError as exc:
return DwhHealth(ok=False, detail=str(exc)) return DwhHealth(ok=False, detail=str(exc), error_kind="connection")
ok = bool(result.get("db_connected") and result.get("schema_accessible")) ok = bool(result.get("db_connected") and result.get("schema_accessible"))
return DwhHealth(ok=ok, detail=None if ok else str(result), return DwhHealth(ok=ok, detail=None if ok else str(result),
database=self._database.database, schema=self._database.db_schema, database=self._database.database, schema=self._database.db_schema,
endpoint=self._client.cfg.base_url, read_only=True) endpoint=self._client.cfg.base_url, read_only=True,
error_kind=None if ok else "inaccessible")
def introspect(self) -> PhysicalSchema: def introspect(self) -> PhysicalSchema:
return introspect_rest( return introspect_rest(
+6 -1
View File
@@ -19,7 +19,12 @@ def ping_cmd(config: Path = CONFIG_OPT) -> None:
raise typer.Exit(code=1) raise typer.Exit(code=1)
health = build_dwh(cfg).health() health = build_dwh(cfg).health()
if not health.ok: if not health.ok:
typer.secho(f"ERRORE di connessione: {health.detail}", fg=typer.colors.RED, err=True) message = (
f"ERRORE: DWH non accessibile via REST (risposta: {health.detail})."
if health.error_kind == "inaccessible"
else f"ERRORE di connessione: {health.detail}"
)
typer.secho(message, fg=typer.colors.RED, err=True)
raise typer.Exit(code=1) raise typer.Exit(code=1)
if health.endpoint: if health.endpoint:
typer.secho(f"OK: connesso via REST a {health.endpoint} (schema {health.schema})", typer.secho(f"OK: connesso via REST a {health.endpoint} (schema {health.schema})",
+27 -24
View File
@@ -12,6 +12,30 @@ def _lsh_dir(cfg) -> Path:
return cfg.paths.indexes / "lsh" return cfg.paths.indexes / "lsh"
def _extract_lsh_values(dwh, physical, annotations, limit):
from tht.db.sampling import SkippedColumn, TruncatedColumn, is_text_type
from tht.mschema.eligibility import effective_eligibility
values, skipped, truncated = {}, [], []
for table_name, table in physical.tables.items():
table_ann = annotations.tables.get(table_name)
for column_name, column in table.columns.items():
ann_col = table_ann.columns.get(column_name) if table_ann else None
if not is_text_type(column.type) or not effective_eligibility(column, ann_col)[0]:
continue
try:
distinct = dwh.distinct_values(table_name, column_name, limit=limit)
except Exception as exc:
skipped.append(SkippedColumn(table_name, column_name, f"errore: {exc}"))
continue
vals = [str(value) for value in distinct.values if value not in (None, "")]
if vals:
values.setdefault(table_name, {})[column_name] = vals
if distinct.truncated:
truncated.append(TruncatedColumn(table_name, column_name, len(vals)))
return values, skipped, truncated
@lsh_app.command("build") @lsh_app.command("build")
def build_cmd(config: Path = CONFIG_OPT) -> None: def build_cmd(config: Path = CONFIG_OPT) -> None:
"""Costruisce l'indice LSH dai valori del database e lo salva su pickle.""" """Costruisce l'indice LSH dai valori del database e lo salva su pickle."""
@@ -33,31 +57,10 @@ def build_cmd(config: Path = CONFIG_OPT) -> None:
typer.echo("Estrazione valori (i più frequenti) dalle colonne testuali eligible...") typer.echo("Estrazione valori (i più frequenti) dalle colonne testuali eligible...")
from tht.adapters.factory import build_dwh from tht.adapters.factory import build_dwh
from tht.db.sampling import SkippedColumn, TruncatedColumn, is_text_type
from tht.mschema.eligibility import effective_eligibility
dwh = build_dwh(cfg) dwh = build_dwh(cfg)
values: dict[str, dict[str, list[str]]] = {} values, skipped, truncated = _extract_lsh_values(
skipped: list[SkippedColumn] = [] dwh, physical, annotations, cfg.lsh.max_values_per_column
truncated: list[TruncatedColumn] = [] )
for table_name, table in physical.tables.items():
table_ann = annotations.tables.get(table_name)
for column_name, column in table.columns.items():
ann_col = table_ann.columns.get(column_name) if table_ann else None
if not is_text_type(column.type) or not effective_eligibility(column, ann_col)[0]:
continue
try:
distinct = dwh.distinct_values(
table_name, column_name, limit=cfg.lsh.max_values_per_column
)
except Exception as e:
skipped.append(SkippedColumn(table_name, column_name, f"errore: {e}"))
continue
vals = [str(value) for value in distinct.values if value not in (None, "")]
if vals:
values.setdefault(table_name, {})[column_name] = vals
if distinct.truncated or len(vals) >= cfg.lsh.max_values_per_column:
truncated.append(TruncatedColumn(table_name, column_name, len(vals)))
n_values = sum(len(v) for t in values.values() for v in t.values()) n_values = sum(len(v) for t in values.values() for v in t.values())
typer.echo(f" {n_values} valori da {sum(len(t) for t in values.values())} colonne") typer.echo(f" {n_values} valori da {sum(len(t) for t in values.values())} colonne")
for s in skipped: for s in skipped:
+1
View File
@@ -25,6 +25,7 @@ class DwhHealth:
read_only: bool | None = None read_only: bool | None = None
writable_tables: tuple[str, ...] = () writable_tables: tuple[str, ...] = ()
can_create: bool = False can_create: bool = False
error_kind: str | None = None
@dataclass(frozen=True) @dataclass(frozen=True)