diff --git a/harness/tests/test_set_schema_linking.py b/harness/tests/test_set_schema_linking.py new file mode 100644 index 00000000..ba8f0091 --- /dev/null +++ b/harness/tests/test_set_schema_linking.py @@ -0,0 +1,35 @@ +"""L1: store.set_schema_linking — validate against SchemaLinking, then write.""" +import json + +import pytest +from pydantic import ValidationError + +from tht.config import DatabaseConfig +from tht.session.models import SchemaLinking +from tht.session.store import create_session, set_schema_linking + + +def _db(): + return DatabaseConfig(database="testdb", user="u", password="p", **{"schema": "public"}) # noqa: S106 + + +def test_writes_and_revalidates(tmp_path): + m = create_session("q", _db(), tmp_path) + data = { + "question": "q riscritta", + "candidates": [{"kind": "table", "name": "fact_x", "decision": "promoted"}], + "joins": [{"from": "a.k", "to": "b.k"}], + } + path = set_schema_linking(m.id, data, tmp_path) + assert path.exists() + reloaded = json.loads(path.read_text()) + assert reloaded["candidates"][0]["name"] == "fact_x" + assert reloaded["joins"][0]["from"] == "a.k" # 'from' alias round-trips + SchemaLinking.model_validate(reloaded) # re-validates clean + + +def test_rejects_invalid_and_writes_nothing(tmp_path): + m = create_session("q", _db(), tmp_path) + with pytest.raises(ValidationError): + set_schema_linking(m.id, {"question": "q", "bogus": 1}, tmp_path) # extra=forbid + assert not (tmp_path / m.id / "schema_linking.json").exists() diff --git a/harness/tht/session/store.py b/harness/tht/session/store.py index c94958dd..859f1de3 100644 --- a/harness/tht/session/store.py +++ b/harness/tht/session/store.py @@ -1,3 +1,4 @@ +import json import os import shutil from datetime import UTC, datetime @@ -158,6 +159,28 @@ def set_question( return path +def set_schema_linking( + session_id: str, + data: dict, + sessions_root: Path, +) -> Path: + """Valida e scrive schema_linking.json (Fase 4) in modo deterministico. + + Valida `data` contro il modello SchemaLinking (ValidationError se invalido) PRIMA + di scrivere, così un artefatto malformato non tocca mai il disco. Ritorna il path. + """ + from tht.session.models import SchemaLinking + + load_session(session_id, sessions_root) + model = SchemaLinking.model_validate(data) + path = sessions_root / session_id / "schema_linking.json" + path.write_text( + json.dumps(model.model_dump(by_alias=True), indent=2, ensure_ascii=False) + ) + touch_manifest(session_id, sessions_root) + return path + + def load_session(session_id: str, sessions_root: Path) -> SessionManifest: path = sessions_root / session_id / MANIFEST if not path.exists():