Merge origin/codex/portable-deployment into feat/docker-local-deploy

Unisce gli internals di Codex (secret-bundle, provider-credentials, auth upstream,
security hardening, CI multiarch) mantenendo le fix portal-specific:
- backend: configPath da THT_CONFIG (fix sessioni) + dataRoot di Codex; authMode 'upstream'
- Docker/compose: TENUTO il mio (verificato live: omics_network+alias, env_file, pi npm-g)
  perche' il compose/Dockerfile/entrypoint di Codex sono accoppiati al suo modello
  secret-bundle (tht doctor inesistente, secret-policy.sh). Adottabile in futuro.
- config.test.ts: preso Codex (superset)
Verificato: tsc clean, 132/132 vitest.
This commit is contained in:
User
2026-07-12 21:13:20 +02:00
211 changed files with 23270 additions and 415 deletions
+4
View File
@@ -22,6 +22,7 @@ dependencies = [
tht = "tht.cli:app"
[project.optional-dependencies]
s3 = ["boto3>=1.34,<2"]
dev = [
"pytest>=8.0",
"testcontainers[postgres]>=4.0",
@@ -31,6 +32,9 @@ dev = [
[tool.setuptools.packages.find]
include = ["tht*"]
[tool.setuptools.package-data]
tht = ["migrations/vector/*.sql"]
[tool.ruff]
line-length = 100
+27 -6
View File
@@ -28,7 +28,9 @@ $$;
create or replace function public.search_similar(
table_name text,
query_embedding vector,
limit_count integer
limit_count integer,
kinds text[] default null,
metadata_filter jsonb default null
)
returns table(id bigint, similarity real, metadata jsonb)
language plpgsql
@@ -37,15 +39,33 @@ set search_path = public, vectors, extensions
as $$
begin
perform public._assert_vector_read_table(table_name);
if metadata_filter is not null and (
table_name <> 'evidence'
or not (metadata_filter ? 'vector_generation')
or not (metadata_filter ? 'document_ids')
or not (metadata_filter ? 'workspace_id')
or jsonb_object_length(metadata_filter) <> 3
or jsonb_typeof(metadata_filter->'document_ids') <> 'array'
) then
raise exception 'invalid Evidence metadata filter';
end if;
return query execute format(
'select t.id,
(1 - (t.embedding <=> $1))::real as similarity,
t.metadata
from vectors.%I t
where ($3 is null or t.kind = any ($3))
and ($4 is null or (
t.metadata->>''vector_generation'' = $4->>''vector_generation''
and t.metadata->>''workspace_id'' = $4->>''workspace_id''
and t.metadata->>''document_id'' in (
select jsonb_array_elements_text($4->''document_ids'')
)
))
order by t.embedding <=> $1
limit $2',
table_name
) using query_embedding, limit_count;
) using query_embedding, limit_count, kinds, metadata_filter;
end;
$$;
@@ -75,7 +95,7 @@ end;
$$;
revoke all on function public._assert_vector_read_table(text) from public;
revoke all on function public.search_similar(text, vector, integer) from public;
revoke all on function public.search_similar(text, vector, integer, text[], jsonb) from public;
revoke all on function public.list_tables() from public;
-- Revoca dai ruoli client generici, poi abilita solo il reader dedicato
@@ -83,15 +103,16 @@ revoke all on function public.list_tables() from public;
do $$
begin
if exists (select 1 from pg_roles where rolname = 'anon') then
revoke all on function public.search_similar(text, vector, integer) from anon;
revoke all on function public.search_similar(text, vector, integer, text[], jsonb) from anon;
revoke all on function public.list_tables() from anon;
end if;
if exists (select 1 from pg_roles where rolname = 'authenticated') then
revoke all on function public.search_similar(text, vector, integer) from authenticated;
revoke all on function public.search_similar(text, vector, integer, text[], jsonb) from authenticated;
revoke all on function public.list_tables() from authenticated;
end if;
if exists (select 1 from pg_roles where rolname = 'vector_reader') then
grant execute on function public.search_similar(text, vector, integer) to vector_reader;
grant execute on function public.search_similar(text, vector, integer, text[], jsonb)
to vector_reader;
grant execute on function public.list_tables() to vector_reader;
end if;
end $$;
@@ -40,6 +40,30 @@ begin
end;
$$;
drop function if exists public.list_evidence_generations(text, text);
create or replace function public.list_evidence_generations(
table_name text, kind text, workspace_id text
)
returns table(generation text)
language plpgsql
security definer
set search_path = public, vectors, extensions
as $$
begin
if table_name <> 'evidence' or kind <> 'evidence' or workspace_id !~ '^[a-z][a-z0-9_-]{0,63}$' then
raise exception 'only exact Evidence generations may be listed';
end if;
return query
select distinct e.metadata->>'vector_generation'
from vectors.evidence e
where e.kind = 'evidence'
and e.metadata->>'vector_generation' ~ '^gen:[0-9a-f]{32}$'
and e.metadata->>'workspace_id' = workspace_id
order by 1;
end;
$$;
create or replace function public.existing_vector_hashes(table_name text, kinds text[])
returns table(record_key text, content_hash text)
language plpgsql
@@ -101,9 +125,35 @@ begin
end;
$$;
drop function if exists public.delete_vector_generation(text, text, text);
create or replace function public.delete_vector_generation(
table_name text, kind text, generation text, workspace_id text
)
returns jsonb
language plpgsql
security definer
set search_path = public, vectors, extensions
as $$
declare affected integer;
begin
if table_name <> 'evidence' or kind <> 'evidence' or generation !~ '^gen:[0-9a-f]{32}$'
or workspace_id !~ '^[a-z][a-z0-9_-]{0,63}$' then
raise exception 'only an exact Evidence generation may be deleted';
end if;
delete from vectors.evidence e
where e.kind = 'evidence' and e.metadata->>'vector_generation' = generation
and e.metadata->>'workspace_id' = workspace_id;
get diagnostics affected = row_count;
return jsonb_build_object('deleted', affected);
end;
$$;
revoke all on function public._assert_vector_write_table(text, text[]) from public;
revoke all on function public.existing_vector_hashes(text, text[]) from public;
revoke all on function public.upsert_vector_records(text, jsonb) from public;
revoke all on function public.delete_vector_generation(text, text, text, text) from public;
revoke all on function public.list_evidence_generations(text, text, text) from public;
-- Su alcuni progetti Supabase le funzioni in `public` ricevono grant automatici: revoca
-- esplicitamente dai ruoli client generici, poi abilita solo il writer dedicato.
@@ -112,14 +162,20 @@ begin
if exists (select 1 from pg_roles where rolname = 'anon') then
revoke all on function public.existing_vector_hashes(text, text[]) from anon;
revoke all on function public.upsert_vector_records(text, jsonb) from anon;
revoke all on function public.delete_vector_generation(text, text, text, text) from anon;
revoke all on function public.list_evidence_generations(text, text, text) from anon;
end if;
if exists (select 1 from pg_roles where rolname = 'authenticated') then
revoke all on function public.existing_vector_hashes(text, text[]) from authenticated;
revoke all on function public.upsert_vector_records(text, jsonb) from authenticated;
revoke all on function public.delete_vector_generation(text, text, text, text) from authenticated;
revoke all on function public.list_evidence_generations(text, text, text) from authenticated;
end if;
if exists (select 1 from pg_roles where rolname = 'vector_writer') then
grant execute on function public.existing_vector_hashes(text, text[]) to vector_writer;
grant execute on function public.upsert_vector_records(text, jsonb) to vector_writer;
grant execute on function public.delete_vector_generation(text, text, text, text) to vector_writer;
grant execute on function public.list_evidence_generations(text, text, text) to vector_writer;
end if;
end $$;
+30 -1
View File
@@ -6,11 +6,27 @@ 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]
@pytest.mark.parametrize("invalid_limit", [True, 1.5, 0, -1])
def test_direct_sampling_rejects_non_positive_integer_limits(admin_engine, invalid_limit):
with pytest.raises(ValueError, match="positive integer"):
sample_column(
admin_engine, "dw", "fct_ricoveri", "reparto", limit=invalid_limit
)
with pytest.raises(ValueError, match="positive integer"):
distinct_values(
admin_engine,
"dw",
"fct_ricoveri",
"reparto",
max_values=invalid_limit,
)
def test_is_text_type():
assert is_text_type("text")
assert is_text_type("varchar(100)")
@@ -52,3 +68,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
@@ -0,0 +1,258 @@
"""L0 gate for the complete durable Evidence/pgvector lifecycle."""
import hashlib
import json
import pytest
from sqlalchemy import create_engine, text
from testcontainers.postgres import PostgresContainer
from tht.adapters.evidence import FilesystemEvidenceSource
from tht.adapters.vector.pgvector import PgVectorStore
from tht.cli.vector_migrate_cmd import migrate
from tht.config import DatabaseConfig
from tht.corpus.chunk import ChunkPolicy
from tht.corpus.pipeline import CorpusPipeline
from tht.corpus.store import CorpusStore
from tht.ports.vector import VectorRecord, VectorWriteRecord
from tht.search import combined_search
from tht.search.evidence import ActiveEvidenceSearcher, resolve_evidence_file
DIMENSIONS = 768
class DeterministicEmbedder:
def embed_documents(self, texts):
return [self.embed_query(text) for text in texts]
def embed_query(self, text):
vector = [0.0] * DIMENSIONS
vector[0] = 0.8
vector[1] = 0.6
return vector
class EvidenceDelegate:
"""Adapt the real multi-collection port to the runtime search protocol."""
def __init__(self, store):
self.store = store
def search(self, embedding, top_n=10, kinds=None, metadata_filter=None):
return self.store.search(
["evidence"], embedding, limit=top_n, kinds=kinds,
metadata_filter=metadata_filter,
)
class InterruptAfterRealPartialUpsert:
"""Crash after a committed real row, as a process death would."""
def __init__(self, store):
self.store = store
self.interrupt = True
def __getattr__(self, name):
return getattr(self.store, name)
def upsert(self, collection, records):
if self.interrupt and len(records) > 1:
self.interrupt = False
self.store.upsert(collection, records[:1])
raise KeyboardInterrupt("injected process death after committed vector row")
return self.store.upsert(collection, records)
@pytest.fixture(scope="module")
def persistent_pgvector():
with PostgresContainer("pgvector/pgvector:pg16") as postgres:
migrate(postgres.get_connection_url())
admin = create_engine(postgres.get_connection_url())
with admin.begin() as connection:
connection.exec_driver_sql(
"ALTER ROLE vector_reader LOGIN PASSWORD 'reader-lifecycle'"
)
connection.exec_driver_sql(
"ALTER ROLE vector_writer LOGIN PASSWORD 'writer-lifecycle'"
)
url = admin.url
common = dict(
host=url.host, port=url.port, database=url.database, schema="vectors"
)
reader = DatabaseConfig(
**common, user="vector_reader", password="reader-lifecycle"
)
writer = DatabaseConfig(
**common, user="vector_writer", password="writer-lifecycle"
)
yield postgres, admin, reader, writer
admin.dispose()
def _pipeline(root, source_root, vectors):
return CorpusPipeline(
store=CorpusStore(root / "corpus"),
sources=[FilesystemEvidenceSource(source_root)],
embedder=DeterministicEmbedder(),
vector_store=vectors,
embedding_model="deterministic-l0",
embedding_dimensions=DIMENSIONS,
chunk_policy=ChunkPolicy(version="lifecycle-v1", max_chars=48),
pipeline_version="evidence-v1",
retain_published_generations=2,
)
def _publish(pipeline, root, serial):
return pipeline.run_as_job(
workspace_id="pgvector-lifecycle",
workspace_root=root,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + f"{serial:x}" * 64,
)
@pytest.mark.l0
def test_real_pgvector_corpus_job_lifecycle(tmp_path, persistent_pgvector):
postgres, admin, reader_config, writer_config = persistent_pgvector
source_root = tmp_path / "sources"
source_root.mkdir()
kept = source_root / "kept.md"
stable = source_root / "stable.md"
stable.write_text("unchanged dependency evidence", encoding="utf-8")
removed = source_root / "removed.md"
removed.write_text("removed evidence generation zero", encoding="utf-8")
vectors = PgVectorStore(reader_config, writer_config, expected_dimension=DIMENSIONS)
generations = []
for serial in range(3):
kept.write_text(f"active evidence generation {serial}", encoding="utf-8")
result = _publish(_pipeline(tmp_path, source_root, vectors), tmp_path, serial + 1)
assert result.status == "succeeded"
generations.append(result.generation)
removed_document = next(
doc for doc in CorpusStore(tmp_path / "corpus").active_manifest().documents
if "removed.md" in doc.source_uri
)
removed_document_id = removed_document.document_id
removed_ref = removed_document.document_id
# A stale, closer row must not consume LIMIT before ACTIVE filtering.
stale_generation = generations[-2]
stale = VectorWriteRecord(
record=VectorRecord(
id="chunk:stale-perfect-match", kind="evidence", ref="doc:stale",
title="stale forbidden", content="stale forbidden",
metadata={
"document_id": "doc:stale",
"vector_generation": stale_generation,
},
),
embedding=[1.0] + [0.0] * (DIMENSIONS - 1),
content_hash="sha256:" + "a" * 64,
)
vectors.upsert("evidence", [stale])
runtime = ActiveEvidenceSearcher(CorpusStore(tmp_path / "corpus"), EvidenceDelegate(vectors))
query = DeterministicEmbedder().embed_query("active")
hits = runtime.search(query, top_n=1, kinds=["evidence"])
assert len(hits) == 1 and hits[0].title != "stale forbidden"
packed = combined_search(
"active", lsh_hits=None, store=runtime, embedder=DeterministicEmbedder(),
top=1, rrf_k=60, kinds=["evidence"], query_vec=query,
)
assert len(packed) == 1 and packed[0].label != "stale forbidden"
# Fourth publication removes a document and creates multiple chunks for crash recovery.
removed.unlink()
kept.write_text("active fourth generation " * 8, encoding="utf-8")
crashing = InterruptAfterRealPartialUpsert(vectors)
candidate = _pipeline(tmp_path, source_root, crashing)
with pytest.raises(KeyboardInterrupt, match="injected process death"):
_publish(candidate, tmp_path, 4)
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
crashed_run = max(runs.iterdir(), key=lambda path: path.stat().st_mtime_ns).name
before = vectors.existing_hashes("evidence", ["evidence"])
intent = json.loads(
(runs / crashed_run / "artifacts" / "vector-intent.json").read_text()
)["records"]
already_present = set(intent) & set(before)
assert len(already_present) == 1
resumed = candidate.run_as_job(
workspace_id="pgvector-lifecycle", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "4" * 64,
resume_run_id=crashed_run,
)
assert resumed.status == "succeeded" and resumed.resumed_from == crashed_run
generations.append(resumed.generation)
after = vectors.existing_hashes("evidence", ["evidence"])
assert {key: after[key] for key in already_present} == {
key: before[key] for key in already_present
}
assert set(intent).issubset(after)
with admin.connect() as connection:
duplicate_count = connection.execute(text(
"SELECT count(*) - count(DISTINCT record_key) FROM vectors.evidence"
)).scalar_one()
assert duplicate_count == 0
runtime = ActiveEvidenceSearcher(CorpusStore(tmp_path / "corpus"), EvidenceDelegate(vectors))
active_hits = runtime.search(query, top_n=20, kinds=["evidence"])
assert active_hits
assert any("active fourth generation" in hit.content for hit in active_hits)
assert all(hit.ref not in {"doc:stale", removed_ref} for hit in active_hits)
assert all(hit.metadata.get("document_id") != removed_document_id for hit in active_hits)
active_pack = combined_search(
"active", lsh_hits=None, store=runtime, embedder=DeterministicEmbedder(),
top=20, rrf_k=60, kinds=["evidence"], query_vec=query,
)
assert active_pack
assert any("active fourth generation" in result.content for result in active_pack)
assert all(removed_document.content not in result.content for result in active_pack)
assert any("unchanged dependency evidence" in result.content for result in active_pack)
manifest = CorpusStore(tmp_path / "corpus").active_manifest()
assert resolve_evidence_file(
CorpusStore(tmp_path / "corpus"), removed_document_id,
materialized_root=tmp_path / "session",
) == ""
active_document = manifest.documents[0]
owned = CorpusStore(tmp_path / "corpus").materialize_document(
active_document.document_id, tmp_path / "session" / "active-evidence.md"
)
assert owned.read_bytes() == active_document.content.encode()
assert hashlib.sha256(owned.read_bytes()).hexdigest() == active_document.content_hash[7:]
# Recreate engines and stores against the same persisted database.
vectors._reader.dispose()
vectors._writer.dispose()
recreated = PgVectorStore(reader_config, writer_config, expected_dimension=DIMENSIONS)
assert recreated.health().ok is True
recreated_hits = ActiveEvidenceSearcher(
CorpusStore(tmp_path / "corpus"), EvidenceDelegate(recreated)
).search(query, top_n=2, kinds=["evidence"])
active_dependencies = set(manifest.metadata["document_generations"].values())
assert recreated_hits
assert all(hit.metadata["vector_generation"] in active_dependencies for hit in recreated_hits)
orphan = "gen:" + "f" * 32
recreated.upsert("evidence", [VectorWriteRecord(
record=VectorRecord(
id="chunk:exact-vector-orphan", kind="evidence", ref="doc:orphan",
title="orphan", content="orphan",
metadata={"document_id": "doc:orphan", "vector_generation": orphan,
"workspace_id": "pgvector-lifecycle"},
),
embedding=query, content_hash="sha256:" + "f" * 64,
)])
final_pipeline = _pipeline(tmp_path, source_root, recreated)
report = final_pipeline.gc(workspace_root=tmp_path)
assert report["evicted"] == [orphan]
expected_fs = set(generations[-2:])
assert set(CorpusStore(tmp_path / "corpus").list_generations()) == expected_fs
expected_vectors = expected_fs | {generations[0]}
assert set(recreated.list_evidence_generations(
"evidence", "pgvector-lifecycle"
)) == expected_vectors
assert final_pipeline.gc(workspace_root=tmp_path)["evicted"] == []
+406
View File
@@ -0,0 +1,406 @@
import pytest
from psycopg2.errors import InsufficientPrivilege
from sqlalchemy import create_engine, text
from sqlalchemy.exc import ProgrammingError
from testcontainers.postgres import PostgresContainer
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
from tht.config import DatabaseConfig
from tht.ports.vector import (
VectorReadUnavailable,
VectorRecord,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
)
def _record(content_hash: str, embedding: list[float], *, kind: str = "memory"):
return VectorWriteRecord(
record=VectorRecord(
id=f"record:{content_hash}",
kind=kind,
ref="session:test",
title=content_hash,
content=f"content {content_hash}",
metadata={"content_hash": content_hash},
),
embedding=embedding,
content_hash=content_hash,
)
@pytest.fixture(scope="module")
def vector_configs():
with PostgresContainer("pgvector/pgvector:pg16") as pg:
host = pg.get_container_host_ip()
port = int(pg.get_exposed_port(5432))
admin_config = DatabaseConfig(
host=host,
port=port,
database=pg.dbname,
schema="vectors",
user=pg.username,
password=pg.password,
)
engine = create_engine(pg.get_connection_url())
with engine.begin() as connection:
connection.exec_driver_sql("CREATE SCHEMA vectors")
connection.exec_driver_sql("CREATE EXTENSION vector WITH SCHEMA vectors")
for table in ("schema_records", "evidence", "memory"):
connection.exec_driver_sql(f"""
CREATE TABLE vectors.{table} (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(2) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT now()
)
""")
connection.exec_driver_sql("CREATE ROLE vector_l0_reader LOGIN PASSWORD 'reader'")
connection.exec_driver_sql("CREATE ROLE vector_l0_writer LOGIN PASSWORD 'writer'")
connection.exec_driver_sql(
"CREATE ROLE vector_l0_no_sequence LOGIN PASSWORD 'no_sequence'"
)
connection.exec_driver_sql(
"GRANT USAGE ON SCHEMA vectors TO vector_l0_reader, vector_l0_writer, "
"vector_l0_no_sequence"
)
connection.exec_driver_sql(
"GRANT SELECT ON ALL TABLES IN SCHEMA vectors TO vector_l0_reader"
)
connection.exec_driver_sql(
"GRANT USAGE, SELECT ON ALL SEQUENCES IN SCHEMA vectors TO vector_l0_writer"
)
for table in ("schema_records", "evidence", "memory"):
connection.exec_driver_sql(
f"GRANT INSERT, UPDATE ON vectors.{table} "
"TO vector_l0_writer, vector_l0_no_sequence"
)
if table == "evidence":
connection.exec_driver_sql(
"GRANT DELETE ON vectors.evidence TO vector_l0_writer"
)
connection.exec_driver_sql(
"GRANT SELECT (kind, metadata) ON vectors.evidence TO vector_l0_writer"
)
connection.exec_driver_sql(
f"GRANT SELECT (record_key, kind, content_hash) "
f"ON vectors.{table} TO vector_l0_writer, vector_l0_no_sequence"
)
engine.dispose()
reader_config = admin_config.model_copy(
update={"user": "vector_l0_reader", "password": "reader"}
)
writer_config = admin_config.model_copy(
update={"user": "vector_l0_writer", "password": "writer"}
)
no_sequence_config = admin_config.model_copy(
update={"user": "vector_l0_no_sequence", "password": "no_sequence"}
)
yield admin_config, reader_config, writer_config, no_sequence_config
@pytest.fixture
def store(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
_, reader_config, writer_config, _ = vector_configs
store = PgVectorStore(reader_config, writer_config, expected_dimension=2)
store.upsert("memory", [_record("reset", [0.0, 1.0])])
yield store
def test_pgvector_round_trip_hash_and_upsert(store):
assert store.upsert("memory", [_record("a", [1.0, 0.0])]) == 1
assert store.existing_hashes("memory", ["memory"])["record:a"] == "a"
hits = store.search(["memory"], [1.0, 0.0], limit=5, kinds=["memory"])
assert hits[0].metadata["content_hash"] == "a"
assert hits[0].id == "record:a"
assert store.upsert("memory", [_record("a", [0.8, 0.2])]) == 1
assert store.search(["memory"], [0.8, 0.2], limit=1)[0].id == "record:a"
def test_pgvector_lists_and_deletes_exact_evidence_generation(store):
generation = "gen:" + "a" * 32
value = VectorWriteRecord(
record=VectorRecord(
id="evidence-generation-a", kind="evidence", ref="doc:a", title="a",
content="content", metadata={"vector_generation": generation, "workspace_id": "default"},
),
embedding=[1.0, 0.0], content_hash="sha256:" + "a" * 64,
)
store.upsert("evidence", [value])
assert generation in store.list_evidence_generations("evidence", "default")
assert store.delete_generation("evidence", generation, "default") == 1
assert generation not in store.list_evidence_generations("evidence", "default")
def test_pgvector_generation_cleanup_isolated_between_workspaces(store):
generation = "gen:" + "b" * 32
records = [VectorWriteRecord(
record=VectorRecord(
id=f"evidence-{workspace}", kind="evidence", ref=f"doc:{workspace}",
title=workspace, content=workspace,
metadata={"vector_generation": generation, "workspace_id": workspace},
), embedding=[1.0, 0.0], content_hash="sha256:" + key * 64,
) for workspace, key in (("workspace-a", "b"), ("workspace-b", "c"))]
store.upsert("evidence", records)
assert store.delete_generation("evidence", generation, "workspace-a") == 1
assert generation not in store.list_evidence_generations("evidence", "workspace-a")
assert generation in store.list_evidence_generations("evidence", "workspace-b")
def test_pgvector_search_filters_kinds_before_limit(store):
store.upsert("memory", [_record("solved", [1.0, 0.0], kind="solved_question")])
hits = store.search("memory".split(), [1.0, 0.0], limit=1, kinds=["memory"])
assert len(hits) == 1
assert hits[0].kind == "memory"
def test_pgvector_multi_collection_search_skips_collections_unrelated_to_kinds(store):
store.upsert("evidence", [_record("evidence", [1.0, 0.0], kind="evidence")])
hits = store.search(["evidence", "memory"], [1.0, 0.0], limit=3, kinds=["memory"])
assert hits
assert {hit.kind for hit in hits} == {"memory"}
def test_pgvector_multi_collection_kind_filter_matches_http_adapter(store):
class Reader:
def search_similar(self, collection, embedding, limit, kinds=None):
if collection != "memory" or "memory" not in (kinds or []):
return []
return [
{
"similarity": 1.0,
"metadata": {
"record_key": "record:a",
"kind": "memory",
"ref": "session:test",
"title": "a",
"content": "content a",
"content_hash": "a",
},
}
]
direct = store.search(["evidence", "memory"], [1.0, 0.0], limit=1, kinds=["memory"])
http = ThothHttpVectorStore(Reader(), None).search(
["evidence", "memory"], [1.0, 0.0], limit=1, kinds=["memory"]
)
assert [(hit.id, hit.kind) for hit in direct] == [(hit.id, hit.kind) for hit in http]
def test_pgvector_search_rejects_unknown_kind_globally(store):
with pytest.raises(VectorStoreError, match="Kind not allowed"):
store.search(["memory"], [1.0, 0.0], limit=1, kinds=["unknown"])
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
def test_pgvector_search_requires_strict_positive_limit(store, limit):
with pytest.raises(ValueError, match="positive integer"):
store.search(["memory"], [1.0, 0.0], limit=limit)
def test_pgvector_allowlists_collections(store):
with pytest.raises(VectorStoreError, match="Collection not allowed"):
store.search(["memory; DROP SCHEMA vectors"], [1.0, 0.0], limit=1)
with pytest.raises(VectorStoreError, match="Collection not allowed"):
store.upsert("unknown", [])
def test_pgvector_rejects_kinds_not_belonging_to_collection(store):
with pytest.raises(VectorStoreError, match="Kind not allowed"):
store.existing_hashes("evidence", ["memory"])
with pytest.raises(VectorStoreError, match="Kind not allowed"):
store.upsert("evidence", [_record("wrong", [1.0, 0.0])])
def test_pgvector_separates_read_and_write_credentials(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
_, reader_config, writer_config, _ = vector_configs
reader = PgVectorStore(reader_config, expected_dimension=2)
assert reader.capabilities.search is True
assert reader.capabilities.upsert is False
with pytest.raises(VectorWriteUnavailable):
reader.upsert("memory", [])
writer = PgVectorStore(None, writer_config, expected_dimension=2)
assert writer.capabilities.search is False
assert writer.capabilities.upsert is True
with pytest.raises(VectorReadUnavailable):
writer.search(["memory"], [1.0, 0.0], limit=1)
def test_pgvector_database_roles_are_least_privilege(vector_configs):
_, reader_config, writer_config, _ = vector_configs
reader_engine = create_engine(
f"postgresql+psycopg2://{reader_config.user}:{reader_config.password}"
f"@{reader_config.host}:{reader_config.port}/{reader_config.database}"
)
writer_engine = create_engine(
f"postgresql+psycopg2://{writer_config.user}:{writer_config.password}"
f"@{writer_config.host}:{writer_config.port}/{writer_config.database}"
)
with pytest.raises(ProgrammingError):
with reader_engine.begin() as connection:
connection.execute(
text(
"INSERT INTO vectors.memory "
"(record_key, kind, content_hash, metadata, embedding) "
"VALUES ('forbidden', 'memory', 'x', '{}', '[1,0]')"
)
)
with pytest.raises(ProgrammingError):
with writer_engine.connect() as connection:
connection.execute(
text(
"SELECT metadata, 1 - (embedding <=> '[1,0]'::vector) AS similarity "
"FROM vectors.memory ORDER BY embedding <=> '[1,0]'::vector LIMIT 1"
)
)
reader_engine.dispose()
writer_engine.dispose()
def test_pgvector_writer_health_requires_sequence_usage(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
admin_config, _, _, no_sequence_config = vector_configs
store = PgVectorStore(None, no_sequence_config, expected_dimension=2)
health = store.health()
assert health.ok is False
assert health.write_reachable is False
assert health.write_detail == (
"vector schema incomplete: missing sequence privileges evidence, memory, schema_records"
)
with pytest.raises(VectorWriteUnavailable) as error:
store.upsert("memory", [_record("needs-sequence", [1.0, 0.0])])
assert isinstance(error.value.__cause__, InsufficientPrivilege)
admin_engine = create_engine(
f"postgresql+psycopg2://{admin_config.user}:{admin_config.password}"
f"@{admin_config.host}:{admin_config.port}/{admin_config.database}"
)
with admin_engine.begin() as connection:
connection.exec_driver_sql(
"GRANT USAGE ON ALL SEQUENCES IN SCHEMA vectors TO vector_l0_no_sequence"
)
admin_engine.dispose()
assert store.health().ok is True
assert store.upsert("memory", [_record("has-sequence", [1.0, 0.0])]) == 1
def test_pgvector_health_requires_schema_usage_for_reader_and_writer(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
admin_config, reader_config, writer_config, _ = vector_configs
admin_engine = create_engine(
f"postgresql+psycopg2://{admin_config.user}:{admin_config.password}"
f"@{admin_config.host}:{admin_config.port}/{admin_config.database}"
)
store = PgVectorStore(reader_config, writer_config, expected_dimension=2)
with admin_engine.begin() as connection:
connection.exec_driver_sql(
f"REVOKE USAGE ON SCHEMA vectors FROM {reader_config.user}, {writer_config.user}"
)
health = store.health()
assert health.read_reachable is False and health.write_reachable is False
assert "missing schema usage" in health.read_detail
assert "missing schema usage" in health.write_detail
with pytest.raises(VectorReadUnavailable, match="Vector read operation unavailable"):
store.search(["memory"], [1.0, 0.0], limit=1)
with pytest.raises(VectorWriteUnavailable, match="Vector write operation unavailable"):
store.upsert("memory", [_record("blocked", [1.0, 0.0])])
with admin_engine.begin() as connection:
connection.exec_driver_sql(
f"GRANT USAGE ON SCHEMA vectors TO {reader_config.user}, {writer_config.user}"
)
admin_engine.dispose()
assert store.health().ok is True
def test_pgvector_maps_unavailable_connections_without_leaking_password(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
_, reader_config, writer_config, _ = vector_configs
password = "never-leak-this"
reader = reader_config.model_copy(update={"port": 1, "password": password})
writer = writer_config.model_copy(update={"port": 1, "password": password})
with pytest.raises(VectorReadUnavailable) as read_error:
PgVectorStore(reader, None).search(["memory"], [1.0, 0.0], limit=1)
with pytest.raises(VectorWriteUnavailable) as hash_error:
PgVectorStore(None, writer).existing_hashes("memory", ["memory"])
with pytest.raises(VectorWriteUnavailable) as write_error:
PgVectorStore(None, writer).upsert("memory", [_record("x", [1.0, 0.0])])
assert password not in str(read_error.value)
assert password not in str(hash_error.value)
assert password not in str(write_error.value)
def test_pgvector_health_reports_dimension_and_each_connection(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
_, reader_config, writer_config, _ = vector_configs
health = PgVectorStore(reader_config, writer_config, expected_dimension=2).health()
assert health.ok is True
assert health.read_reachable is True
assert health.write_reachable is True
assert health.observed_dimensions == (2,)
assert health.dimension_compatible is True
mismatch = PgVectorStore(reader_config, None, expected_dimension=3).health()
assert mismatch.ok is False
assert mismatch.read_reachable is False
assert mismatch.read_detail == (
"embedding dimension mismatch: evidence=2, memory=2, schema_records=2"
)
assert mismatch.dimension_compatible is False
def test_pgvector_health_rejects_clean_and_partial_schemas(vector_configs):
from tht.adapters.vector.pgvector import PgVectorStore
admin_config, _, _, _ = vector_configs
engine = create_engine(
f"postgresql+psycopg2://{admin_config.user}:{admin_config.password}"
f"@{admin_config.host}:{admin_config.port}/{admin_config.database}"
)
with engine.begin() as connection:
connection.exec_driver_sql("CREATE SCHEMA clean_vectors")
connection.exec_driver_sql("CREATE SCHEMA partial_vectors")
connection.exec_driver_sql(
"CREATE TABLE partial_vectors.memory "
"(record_key text, kind text, content_hash text, metadata jsonb)"
)
engine.dispose()
clean = PgVectorStore(
admin_config.model_copy(update={"db_schema": "clean_vectors"}),
expected_dimension=2,
).health()
assert clean.ok is False
assert clean.read_reachable is False
assert clean.read_detail == (
"vector schema incomplete: missing tables evidence, memory, schema_records"
)
partial = PgVectorStore(
admin_config.model_copy(update={"db_schema": "partial_vectors"}),
expected_dimension=2,
).health()
assert partial.ok is False
assert partial.read_reachable is False
assert partial.read_detail == (
"vector schema incomplete: missing tables evidence, schema_records; "
"missing embedding columns memory"
)
@@ -0,0 +1,288 @@
import math
import pytest
from sqlalchemy import create_engine
from testcontainers.postgres import PostgresContainer
from tht.adapters.vector.pgvector import PgVectorStore
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
from tht.config import DatabaseConfig, RestConfig
from tht.ports.vector import VectorRecord, VectorStoreError, VectorWriteRecord
from tht.vectorstore.rest_client import VectorRestClient, VectorRestError
def _write(record_id, kind, embedding, content_hash):
return VectorWriteRecord(
VectorRecord(
id=record_id,
kind=kind,
ref="fixture",
title=record_id,
content=f"content {record_id}",
metadata={"fixture": True},
),
embedding,
content_hash,
)
FIXTURE = [
_write("memory:a", "memory", [1.0, 0.0], "hash-a"),
_write("memory:b", "memory", [1.0, 0.0], "hash-b"),
_write("solved:a", "solved_question", [0.8, 0.2], "hash-solved"),
]
class Response:
def __init__(self, payload=None, status=200):
self.status_code = status
self.payload = payload
self.text = "" if payload is None else "json"
@property
def ok(self):
return self.status_code < 400
def json(self):
return self.payload
class FixtureHttpTransport:
def __init__(self):
self.rows = {}
self.calls = []
def post(self, url, json, headers, **kwargs):
assert headers == {"X-API-Key": "parity-key"}
self.calls.append((url.rsplit("/", 1)[-1], json))
function = self.calls[-1][0]
if function == "list_tables":
return Response([{"table_name": "memory", "vector_dimensions": 2}])
if function == "upsert_vector_records":
for row in json["rows"]:
self.rows[(json["table_name"], row["record_key"])] = row
return Response({"upserted": len(json["rows"])})
if function == "existing_vector_hashes":
return Response([
{"record_key": row["record_key"], "content_hash": row["content_hash"]}
for (table, _), row in self.rows.items()
if table == json["table_name"] and row["kind"] in json["kinds"]
])
assert function == "search_similar"
table_name = json["table_name"]
embedding = json["query_embedding"]
kinds = json.get("kinds")
def similarity(row):
left, right = row["embedding"], embedding
return sum(a * b for a, b in zip(left, right)) / (
math.sqrt(sum(a * a for a in left))
* math.sqrt(sum(b * b for b in right))
)
rows = [
{"metadata": row["metadata"], "similarity": similarity(row)}
for (table, _), row in self.rows.items()
if table == table_name and (not kinds or row["kind"] in kinds)
]
payload = sorted(
rows,
key=lambda row: (-row["similarity"], row["metadata"]["record_key"]),
)[: json["limit_count"]]
return Response(payload)
@pytest.fixture
def direct_store():
with PostgresContainer("pgvector/pgvector:pg16") as postgres:
config = DatabaseConfig(
host=postgres.get_container_host_ip(),
port=int(postgres.get_exposed_port(5432)),
database=postgres.dbname,
schema="vectors",
user=postgres.username,
password=postgres.password,
)
engine = create_engine(postgres.get_connection_url())
with engine.begin() as connection:
connection.exec_driver_sql("CREATE SCHEMA vectors")
connection.exec_driver_sql("CREATE EXTENSION vector WITH SCHEMA vectors")
connection.exec_driver_sql(
"CREATE TABLE vectors.memory ("
"id bigserial PRIMARY KEY, record_key text UNIQUE NOT NULL, "
"kind text NOT NULL, content_hash text NOT NULL, metadata jsonb NOT NULL, "
"embedding vectors.vector(2) NOT NULL, indexed_at timestamptz NOT NULL "
"DEFAULT now())"
)
engine.dispose()
reader, writer = config, config
store = PgVectorStore(reader, writer, expected_dimension=2)
store.upsert("memory", FIXTURE)
yield store
@pytest.fixture
def http_store(monkeypatch):
transport = FixtureHttpTransport()
monkeypatch.setattr("tht.vectorstore.rest_client.requests.post", transport.post)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
store = ThothHttpVectorStore(client, client, expected_dimension=2)
store.upsert("memory", FIXTURE)
store.transport = transport
return store
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
def test_kind_filtered_search_has_identical_order(request, store_fixture):
store = request.getfixturevalue(store_fixture)
hits = store.search(["memory"], [1.0, 0.0], limit=3, kinds=["memory"])
assert [(hit.id, hit.kind, round(hit.similarity, 6)) for hit in hits] == [
("memory:a", "memory", 1.0),
("memory:b", "memory", 1.0),
]
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
def test_hash_and_upsert_parity(request, store_fixture):
store = request.getfixturevalue(store_fixture)
assert store.existing_hashes("memory", ["memory"]) == {
"memory:a": "hash-a",
"memory:b": "hash-b",
}
replacement = _write("memory:a", "memory", [0.0, 1.0], "hash-a-2")
assert store.upsert("memory", [replacement]) == 1
assert store.existing_hashes("memory", ["memory"])["memory:a"] == "hash-a-2"
assert store.search(["memory"], [0.0, 1.0], limit=1, kinds=["memory"])[0].id == "memory:a"
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
def test_validation_error_parity(request, store_fixture):
store = request.getfixturevalue(store_fixture)
with pytest.raises(VectorStoreError, match="Collection not allowed"):
store.search(["not_allowed"], [1.0, 0.0], limit=1)
with pytest.raises(VectorStoreError, match="Kind not allowed"):
store.search(["memory"], [1.0, 0.0], limit=1, kinds=["not_allowed"])
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
def test_dimension_error_parity(request, store_fixture):
store = request.getfixturevalue(store_fixture)
with pytest.raises(VectorStoreError, match="Query embedding dimension"):
store.search(["memory"], [1.0], limit=1)
with pytest.raises(VectorStoreError, match="Embedding dimension"):
store.upsert("memory", [_write("bad", "memory", [1.0], "bad")])
def test_http_parity_exercises_rpc_kinds_payload(http_store):
http_store.search(["memory"], [1.0, 0.0], limit=2, kinds=["memory"])
search_calls = [payload for function, payload in http_store.transport.calls if function == "search_similar"]
assert search_calls[-1] == {
"query_embedding": [1.0, 0.0],
"limit_count": 2,
"table_name": "memory",
"kinds": ["memory"],
}
def test_http_adapter_maps_transport_error(monkeypatch):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: Response({"message": "server broke"}, status=500),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
store = ThothHttpVectorStore(client, client, expected_dimension=2)
with pytest.raises(VectorStoreError, match="HTTP 500"):
store.search(["memory"], [1.0, 0.0], limit=1, kinds=["memory"])
def test_http_adapter_tolerates_malformed_metadata(monkeypatch):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: Response([{"similarity": 0.5, "metadata": None}]),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
hit = ThothHttpVectorStore(client, None, expected_dimension=2).search(
["memory"], [1.0, 0.0], limit=1
)[0]
assert (hit.id, hit.kind, hit.metadata) == ("", "", {})
def test_http_adapter_legacy_fallback_preserves_kind_semantics(monkeypatch):
calls = []
def post(url, json, **kwargs):
calls.append(json)
if "kinds" in json:
return Response({"message": "function not found"}, status=404)
return Response([
{"similarity": 1.0, "metadata": {"record_key": "wrong", "kind": "solved_question"}},
{"similarity": 0.9, "metadata": {"record_key": "right", "kind": "memory"}},
])
monkeypatch.setattr("tht.vectorstore.rest_client.requests.post", post)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
hits = ThothHttpVectorStore(client, None, expected_dimension=2).search(
["memory"], [1.0, 0.0], limit=2, kinds=["memory"]
)
assert [hit.id for hit in hits] == ["right"]
assert "kinds" in calls[0] and "kinds" not in calls[1]
def test_http_delete_generation_uses_exact_allowlisted_rpc_payload(monkeypatch):
calls = []
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda url, json, **kwargs: calls.append((url, json)) or Response({"deleted": 2}),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
assert client.delete_generation("evidence", "gen:" + "a" * 32, "default") == 2
assert calls == [("https://vectors.test/rpc/delete_vector_generation", {
"table_name": "evidence", "kind": "evidence", "generation": "gen:" + "a" * 32,
"workspace_id": "default",
})]
def test_http_delete_generation_legacy_404_fails_closed_without_body_leak(monkeypatch):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: Response({"message": "secret legacy endpoint detail"}, status=404),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
with pytest.raises(VectorRestError, match="delete_vector_generation RPC is unavailable") as error:
client.delete_generation("evidence", "gen:" + "a" * 32, "default")
assert "secret" not in str(error.value)
def test_http_list_evidence_generations_exact_rpc_and_legacy_fail_closed(monkeypatch):
calls = []
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda url, json, **kwargs: calls.append((url, json)) or Response([
{"generation": "gen:" + "a" * 32}
]),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
assert client.list_evidence_generations("evidence", "default") == ["gen:" + "a" * 32]
assert calls[0][0].endswith("/rpc/list_evidence_generations")
assert calls[0][1] == {"table_name": "evidence", "kind": "evidence", "workspace_id": "default"}
@pytest.mark.parametrize("generation", ["gen:a", "gen:" + "A" * 32, "gen:" + "a" * 33])
def test_http_generation_operations_reject_noncanonical_values(monkeypatch, generation):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: pytest.fail("invalid generation reached transport"),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
with pytest.raises(ValueError, match="canonical"):
client.delete_generation("evidence", generation, "default")
def test_http_inventory_rejects_malformed_rpc_output(monkeypatch):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: Response([{"generation": "gen:../escape"}]),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
with pytest.raises(VectorRestError, match="malformed"):
client.list_evidence_generations("evidence", "default")
+303
View File
@@ -0,0 +1,303 @@
import json
from pathlib import Path
import pytest
from sqlalchemy import create_engine, text
from sqlalchemy.exc import ProgrammingError
from testcontainers.postgres import PostgresContainer
from typer.testing import CliRunner
from tht.cli import app
from tht.config import DatabaseConfig
from tht.ports.vector import VectorRecord, VectorWriteRecord
@pytest.fixture(scope="module")
def database_url():
with PostgresContainer("pgvector/pgvector:pg16") as postgres:
yield postgres.get_connection_url()
def test_migrations_are_clean_and_idempotent(database_url):
from tht.cli.vector_migrate_cmd import migrate, migration_status
before = migration_status(database_url)
assert [item.version for item in before.pending] == ["001", "002", "003", "004"]
migrate(database_url)
migrate(database_url)
status = migration_status(database_url)
assert status.pending == ()
assert status.drifted == ()
assert [item.version for item in status.applied] == ["001", "002", "003", "004"]
def test_schema_matches_direct_adapter_contract(database_url):
engine = create_engine(database_url)
with engine.connect() as connection:
rows = connection.execute(
text(
"SELECT table_name, column_name, data_type, udt_name "
"FROM information_schema.columns WHERE table_schema = 'vectors' "
"ORDER BY table_name, ordinal_position"
)
).all()
vector_types = connection.execute(
text(
"SELECT c.relname, format_type(a.atttypid, a.atttypmod) "
"FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace "
"JOIN pg_attribute a ON a.attrelid = c.oid AND a.attname = 'embedding' "
"WHERE n.nspname = 'vectors' ORDER BY c.relname"
)
).all()
engine.dispose()
tables = {row.table_name for row in rows}
assert tables == {"evidence", "memory", "schema_records"}
required = {"id", "record_key", "kind", "content_hash", "metadata", "embedding", "indexed_at"}
for table in tables:
assert {row.column_name for row in rows if row.table_name == table} == required
assert vector_types == [
("evidence", "vectors.vector(768)"),
("memory", "vectors.vector(768)"),
("schema_records", "vectors.vector(768)"),
]
def test_roles_have_runtime_privileges_only(database_url):
from tht.cli.vector_migrate_cmd import migrate
migrate(database_url)
admin = create_engine(database_url)
with admin.begin() as connection:
connection.exec_driver_sql("ALTER ROLE vector_reader LOGIN PASSWORD 'reader-test-only'")
connection.exec_driver_sql("ALTER ROLE vector_writer LOGIN PASSWORD 'writer-test-only'")
url = admin.url
reader = create_engine(url.set(username="vector_reader", password="reader-test-only"))
writer = create_engine(url.set(username="vector_writer", password="writer-test-only"))
from tht.adapters.vector.pgvector import PgVectorStore
common = {
"host": url.host,
"port": url.port,
"database": url.database,
"schema": "vectors",
}
reader_config = DatabaseConfig(
**common, user="vector_reader", password="reader-test-only"
)
writer_config = DatabaseConfig(
**common, user="vector_writer", password="writer-test-only"
)
store = PgVectorStore(reader_config, writer_config, expected_dimension=768)
assert store.health().ok is True
assert store.upsert(
"memory",
[
VectorWriteRecord(
record=VectorRecord(
id="adapter-write",
kind="memory",
ref="session:test",
title="test",
content="test",
),
embedding=[0.0] * 768,
content_hash="adapter-hash",
)
],
) == 1
with reader.connect() as connection:
connection.execute(text("SELECT metadata, embedding FROM vectors.memory")).all()
with pytest.raises(ProgrammingError):
with reader.begin() as connection:
connection.execute(
text(
"INSERT INTO vectors.memory "
"(record_key, kind, content_hash, metadata, embedding) "
"VALUES ('reader-write', 'memory', 'x', '{}', "
"array_fill(0, ARRAY[768])::vectors.vector)"
)
)
with writer.begin() as connection:
connection.execute(
text(
"INSERT INTO vectors.memory "
"(record_key, kind, content_hash, metadata, embedding) "
"VALUES ('writer-ok', 'memory', 'x', '{}', "
"array_fill(0, ARRAY[768])::vectors.vector)"
)
)
assert connection.execute(
text("SELECT content_hash FROM vectors.memory WHERE record_key = 'writer-ok'")
).scalar_one() == "x"
connection.execute(
text("UPDATE vectors.memory SET content_hash = 'y' WHERE record_key = 'writer-ok'")
)
with pytest.raises(ProgrammingError):
with writer.connect() as connection:
connection.execute(text("SELECT metadata FROM vectors.memory")).all()
with pytest.raises(ProgrammingError):
with writer.begin() as connection:
connection.execute(text("DELETE FROM vectors.memory WHERE record_key = 'writer-ok'"))
reader.dispose()
writer.dispose()
admin.dispose()
def test_status_json_is_pristine(database_url, monkeypatch):
monkeypatch.setenv("THT_VECTOR_ADMIN_URL", database_url)
result = CliRunner().invoke(app, ["vector", "migrate", "--status", "--json"])
assert result.exit_code == 0, result.output
assert json.loads(result.stdout) == {
"applied": ["001", "002", "003", "004"],
"drifted": [],
"pending": [],
}
assert result.stderr == ""
def test_checksum_drift_is_reported_and_refused(database_url, tmp_path):
from tht.cli.vector_migrate_cmd import MigrationError, migrate, migration_status
migrations = _copy_migrations(tmp_path)
migrate(database_url, migrations)
(migrations / "002_schema_tables.sql").write_text("SELECT 2;\n")
assert [item.version for item in migration_status(database_url, migrations).drifted] == [
"002"
]
with pytest.raises(MigrationError, match="checksum drift"):
migrate(database_url, migrations)
def test_unknown_applied_version_is_downgrade_drift(database_url):
from tht.cli.vector_migrate_cmd import MigrationError, migrate, migration_status
migrate(database_url)
engine = create_engine(database_url)
with engine.begin() as connection:
connection.execute(
text(
"INSERT INTO public.tht_vector_migrations (version, name, checksum) "
"VALUES ('999', 'future', 'future-checksum'), "
"('future_x', 'future_named', 'future-checksum')"
)
)
try:
with pytest.raises(
MigrationError, match="absent from local manifest: 999, future_x"
):
migration_status(database_url)
with pytest.raises(
MigrationError, match="absent from local manifest: 999, future_x"
):
migrate(database_url)
finally:
with engine.begin() as connection:
connection.execute(
text(
"DELETE FROM public.tht_vector_migrations "
"WHERE version IN ('999', 'future_x')"
)
)
engine.dispose()
def test_migration_versions_sort_numerically_and_reject_numeric_duplicates(tmp_path):
from tht.cli.vector_migrate_cmd import MigrationError, _discover
migrations = tmp_path / "ordered"
migrations.mkdir()
(migrations / "10_tenth.sql").write_text("SELECT 10;\n")
(migrations / "2_second.sql").write_text("SELECT 2;\n")
assert [item.version for item in _discover(migrations)] == ["2", "10"]
(migrations / "02_duplicate.sql").write_text("SELECT 2;\n")
with pytest.raises(MigrationError, match="Duplicate migration version: 2"):
_discover(migrations)
def test_hostile_admin_search_path_cannot_shadow_migration_objects(database_url):
from tht.cli.vector_migrate_cmd import migrate
admin = create_engine(database_url, isolation_level="AUTOCOMMIT")
with admin.connect() as connection:
connection.exec_driver_sql("DROP DATABASE IF EXISTS vector_hostile")
connection.exec_driver_sql("CREATE DATABASE vector_hostile")
hostile_url = admin.url.set(database="vector_hostile")
hostile = create_engine(hostile_url)
try:
with hostile.begin() as connection:
connection.exec_driver_sql("CREATE SCHEMA shadow")
connection.exec_driver_sql(
"CREATE TABLE shadow.tht_vector_migrations "
"(version text, checksum text, poisoned boolean DEFAULT true)"
)
connection.exec_driver_sql("ALTER ROLE test SET search_path = shadow, public")
hostile.dispose()
migrate(hostile_url.render_as_string(hide_password=False))
verification = create_engine(hostile_url)
with verification.connect() as connection:
assert connection.execute(
text("SELECT count(*) FROM public.tht_vector_migrations")
).scalar_one() == 4
assert connection.execute(
text("SELECT count(*) FROM shadow.tht_vector_migrations")
).scalar_one() == 0
assert connection.execute(
text(
"SELECT format_type(a.atttypid, a.atttypmod) "
"FROM pg_catalog.pg_attribute a "
"WHERE a.attrelid = 'vectors.memory'::pg_catalog.regclass "
"AND a.attname = 'embedding'"
)
).scalar_one() == "vectors.vector(768)"
verification.dispose()
finally:
cleanup = create_engine(database_url, isolation_level="AUTOCOMMIT")
with cleanup.connect() as connection:
connection.exec_driver_sql("ALTER ROLE test RESET search_path")
connection.exec_driver_sql(
"SELECT pg_catalog.pg_terminate_backend(pid) FROM pg_catalog.pg_stat_activity "
"WHERE datname = 'vector_hostile' AND pid <> pg_catalog.pg_backend_pid()"
)
connection.exec_driver_sql("DROP DATABASE IF EXISTS vector_hostile")
cleanup.dispose()
admin.dispose()
def test_failed_batch_rolls_back_schema_and_ledger(database_url, tmp_path):
from tht.cli.vector_migrate_cmd import MigrationError, migrate, migration_status
migrations = _copy_migrations(tmp_path)
(migrations / "005_first.sql").write_text("CREATE TABLE public.must_rollback (id int);\n")
(migrations / "006_broken.sql").write_text("THIS IS NOT SQL;\n")
with pytest.raises(MigrationError, match="006_broken.sql"):
migrate(database_url, migrations)
engine = create_engine(database_url)
with engine.connect() as connection:
assert connection.execute(text("SELECT to_regclass('public.must_rollback')")).scalar() is None
engine.dispose()
status = migration_status(database_url, migrations)
assert [item.version for item in status.applied] == ["001", "002", "003", "004"]
assert [item.version for item in status.pending] == ["005", "006"]
def _copy_migrations(tmp_path: Path) -> Path:
source = Path(__file__).parents[2] / "tht" / "migrations" / "vector"
target = tmp_path / "migrations"
target.mkdir()
for migration in source.glob("*.sql"):
(target / migration.name).write_bytes(migration.read_bytes())
return target
@@ -39,7 +39,10 @@ def test_save_one_upserts_to_real_pgvector(l2_env):
detail="ablazione", rationale="L2 self-test (idempotent)",
question_context="ablazione 2025", tables=["fct_ricoveri"], concepts=[],
)
upserted = save_one_memory([record], decision_seq=999, writer=writer, embedder=embedder)
from tht.adapters.vector import ThothHttpVectorStore
store = ThothHttpVectorStore(reader=writer, writer=writer)
upserted = save_one_memory([record], decision_seq=999, store=store, embedder=embedder)
assert upserted >= 0 # idempotent: 0 on unchanged, >=1 on new/updated
# read it back via the READER key (vector_rest, path /vector/v1/)
@@ -0,0 +1,116 @@
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)
def test_solved_index_writes_through_writer_only_factory_store(monkeypatch):
writer_only_store = SimpleNamespace(
capabilities=SimpleNamespace(search=False, upsert=True),
existing_hashes=lambda *args: {},
upsert=lambda table, rows: 1,
)
cfg = SimpleNamespace(embeddings=object(), vector_write_rest=object())
manifest = SimpleNamespace(id="s1")
solved_record = object()
calls = []
monkeypatch.setattr(memory_cmd, "has_vector_write_rest", lambda cfg: True)
monkeypatch.setattr(memory_cmd, "load_session_or_exit", lambda cfg, session: manifest)
monkeypatch.setattr(memory_cmd, "session_dir", lambda *args: None)
monkeypatch.setattr(
"tht.adapters.factory.build_vector_store",
lambda cfg, require_write: calls.append(require_write) or writer_only_store,
)
monkeypatch.setattr("tht.cli.sql_cmd.promoted_tables_for", lambda *args: [])
monkeypatch.setattr("tht.solved.build_solved_record", lambda *args: solved_record)
monkeypatch.setattr(
"tht.solved.save_solved_question",
lambda record, *, store, embedder: int(
record is solved_record and store is writer_only_store
),
)
monkeypatch.setattr("tht.cli.vector_cmd.make_embedder", lambda cfg: object())
assert memory_cmd.index_solved_session(cfg, "s1") == 1
assert calls == [True]
+131
View File
@@ -0,0 +1,131 @@
import pytest
from tht.adapters.dwh import PostgresDwhAdapter, ThothRestDwhAdapter
from tht.adapters.vector import PgVectorStore, ThothHttpVectorStore
from tht.adapters.factory import build_dwh, build_vector_store
from tht.config import Config, ConfigError
def _config(*, dwh_type="thoth_rest", vector_type="thoth_vector_http", reader=True, writer=True):
dwh = (
{
"type": "thoth_rest",
"database": {"database": "analytics", "schema": "mart"},
"endpoint": {"base_url": "https://dwh.test/", "api_key": "reader"},
}
if dwh_type == "thoth_rest"
else {
"type": "postgres_direct",
"connection": {
"host": "db",
"database": "analytics",
"schema": "mart",
"user": "reader",
"password": "secret",
},
}
)
vectors = (
{
"type": "thoth_vector_http",
**(
{"reader": {"base_url": "https://vectors.test/", "api_key": "reader"}}
if reader
else {}
),
**(
{"writer": {"base_url": "https://vectors.test/", "api_key": "writer"}}
if writer
else {}
),
}
if vector_type == "thoth_vector_http"
else {
"type": "pgvector_direct",
**(
{
"reader": {
"host": "vector-db",
"database": "postgres",
"schema": "vectors",
"user": "reader",
"password": "secret",
}
}
if reader
else {}
),
**(
{
"writer": {
"host": "vector-db",
"database": "postgres",
"schema": "vectors",
"user": "writer",
"password": "secret",
}
}
if writer
else {}
),
}
)
legacy_database = (
dwh["connection"]
if dwh_type == "postgres_direct"
else {
**dwh["database"],
"user": "rest",
"password": "",
"transport": "rest",
}
)
return Config.model_validate({"dwh": dwh, "vectors": vectors, "database": legacy_database})
@pytest.mark.parametrize(
("dwh_type", "adapter_type"),
[("postgres_direct", PostgresDwhAdapter), ("thoth_rest", ThothRestDwhAdapter)],
)
def test_factory_selects_dwh_adapter(dwh_type, adapter_type):
assert isinstance(build_dwh(_config(dwh_type=dwh_type)), adapter_type)
def test_factory_selects_http_vector_and_requires_writer():
config = _config(writer=False)
assert isinstance(build_vector_store(config), ThothHttpVectorStore)
with pytest.raises(ConfigError, match="writer"):
build_vector_store(config, require_write=True)
def test_factory_builds_writer_only_http_vector_when_write_is_required():
config = _config(reader=False, writer=True)
store = build_vector_store(config, require_write=True)
assert isinstance(store, ThothHttpVectorStore)
assert store.capabilities.search is False
assert store.capabilities.upsert is True
def test_factory_selects_direct_vector_store_and_requires_writer():
config = _config(vector_type="pgvector_direct", writer=False)
assert isinstance(build_vector_store(config), PgVectorStore)
with pytest.raises(ConfigError, match="writer"):
build_vector_store(config, require_write=True)
def test_factory_builds_writer_only_direct_vector_when_write_is_required():
store = build_vector_store(
_config(vector_type="pgvector_direct", reader=False), require_write=True
)
assert isinstance(store, PgVectorStore)
assert store.capabilities.search is False
assert store.capabilities.upsert is True
def test_factory_propagates_non_default_statement_timeout():
config = _config(dwh_type="postgres_direct")
config.execution.statement_timeout_ms = 12_345
assert build_dwh(config)._statement_timeout_ms == 12_345
+148
View File
@@ -0,0 +1,148 @@
import json
import os
import subprocess
from pathlib import Path
import pytest
from typer.testing import CliRunner
from tht.cli import app
from tht.config import load_config
from tht.adapters.factory import build_vector_store
def _write_old_workspace(tmp_path):
path = tmp_path / "old.yaml"
path.write_text(
"""
database:
host: ignored-for-rest
database: analytics
schema: mart
user: legacy-user
password: legacy-password
transport: rest
rest:
base_url: https://dwh.example.test/
api_key: dwh-reader
vector_db:
host: vector-db
database: postgres
schema: vectors
user: vector-user
password: vector-password
vector_rest:
base_url: https://vectors.example.test/
api_key: vector-reader
vector_write_rest:
base_url: https://vectors.example.test/
api_key: vector-writer
paths:
artifacts: build/artifacts
indexes: build/indexes
sessions: build/sessions
"""
)
return path
def _write_new_workspace(tmp_path):
path = tmp_path / "new.yaml"
path.write_text(
"""
dwh:
type: thoth_rest
database:
database: analytics
schema: mart
endpoint:
base_url: https://dwh.example.test/
api_key: dwh-reader
vectors:
type: thoth_vector_http
reader:
base_url: https://vectors.example.test/
api_key: vector-reader
writer:
base_url: https://vectors.example.test/
api_key: vector-writer
direct:
host: vector-db
database: postgres
schema: vectors
user: vector-user
password: vector-password
roots:
artifacts: build/artifacts
indexes: build/indexes
sessions: build/sessions
"""
)
return path
def test_legacy_rest_workspace_equals_new_resource_schema(tmp_path, capsys):
with pytest.warns(FutureWarning, match="DEPRECATION") as warnings:
old = load_config(_write_old_workspace(tmp_path))
captured = capsys.readouterr()
new = load_config(_write_new_workspace(tmp_path))
assert old.dwh.model_dump() == new.dwh.model_dump()
assert old.vectors.model_dump() == new.vectors.model_dump()
assert old.roots.model_dump() == new.roots.model_dump()
assert captured.out == ""
assert captured.err == ""
assert len(warnings) == 1
def test_legacy_warning_does_not_contaminate_cli_json(tmp_path):
with pytest.warns(FutureWarning, match="DEPRECATION") as warnings:
result = CliRunner().invoke(
app,
["session", "list", "--json", "-c", str(_write_old_workspace(tmp_path))],
)
assert result.exit_code == 0
json.loads(result.stdout)
assert "DEPRECATION" not in result.stdout
assert result.stderr == ""
assert len(warnings) == 1
def test_legacy_cli_subprocess_warns_once_on_stderr_and_keeps_json_stdout(tmp_path):
workspace = _write_old_workspace(tmp_path)
result = subprocess.run(
[
str(Path(__file__).parents[1] / ".venv" / "bin" / "tht"),
"session",
"list",
"--json",
"-c",
str(workspace),
],
cwd=tmp_path,
env={**os.environ, "PYTHONWARNINGS": "default"},
text=True,
capture_output=True,
check=False,
)
assert result.returncode == 0
json.loads(result.stdout)
assert "DEPRECATION" not in result.stdout
assert result.stderr.count("DEPRECATION") == 1
def test_legacy_writer_only_vector_config_builds_for_targeted_writes(tmp_path):
workspace = _write_old_workspace(tmp_path)
content = workspace.read_text().replace(
"vector_rest:\n base_url: https://vectors.example.test/\n api_key: vector-reader\n",
"",
)
workspace.write_text(content)
with pytest.warns(FutureWarning):
cfg = load_config(workspace)
store = build_vector_store(cfg, require_write=True)
assert store.capabilities.search is False
assert store.capabilities.upsert is True
+171
View File
@@ -0,0 +1,171 @@
import pytest
from tht.config import (
ConfigError,
PgvectorDirectConfig,
PostgresDwhConfig,
ThothRestDwhConfig,
ThothVectorHttpConfig,
load_config,
)
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
from tht.adapters.factory import build_evidence_sources
def test_direct_vector_passwords_load_from_file_references(monkeypatch, tmp_path):
reader = tmp_path / "reader"
writer = tmp_path / "writer"
reader.write_text("reader-secret")
writer.write_text("writer-secret")
monkeypatch.setenv("READER_FILE", str(reader))
monkeypatch.setenv("WRITER_FILE", str(writer))
workspace = tmp_path / "workspace.yaml"
workspace.write_text("""
dwh:
type: postgres_direct
connection: {database: d, schema: public, user: u, password: p}
vectors:
type: pgvector_direct
reader: {database: d, schema: vectors, user: r, password_file: '${READER_FILE}'}
writer: {database: d, schema: vectors, user: w, password_file: '${WRITER_FILE}'}
""")
config = load_config(workspace)
assert config.vectors.reader.password == "reader-secret"
assert config.vectors.writer.password == "writer-secret"
def test_direct_vector_secret_file_rejects_whitespace(tmp_path):
secret = tmp_path / "reader"
secret.write_text("bad secret")
workspace = tmp_path / "workspace.yaml"
workspace.write_text(f"""
dwh:
type: postgres_direct
connection: {{database: d, schema: public, user: u, password: p}}
vectors:
type: pgvector_direct
reader: {{database: d, schema: vectors, user: r, password_file: {secret}}}
""")
with pytest.raises(ConfigError, match="secret file"):
load_config(workspace)
def test_loads_discriminated_dwh_and_vector_resources(tmp_path):
workspace = tmp_path / "workspace.yaml"
workspace.write_text(
"""
dwh:
type: thoth_rest
database:
database: analytics
schema: mart
endpoint:
base_url: https://dwh.example.test/
api_key: dwh-reader
vectors:
type: thoth_vector_http
reader:
base_url: https://vectors.example.test/
api_key: vector-reader
writer:
base_url: https://vectors.example.test/
api_key: vector-writer
roots:
artifacts: build/artifacts
indexes: build/indexes
sessions: build/sessions
"""
)
cfg = load_config(workspace)
assert isinstance(cfg.dwh, ThothRestDwhConfig)
assert cfg.dwh.database.db_schema == "mart"
assert isinstance(cfg.vectors, ThothVectorHttpConfig)
assert cfg.vectors.writer.api_key == "vector-writer"
assert cfg.roots.sessions.as_posix() == "build/sessions"
def test_loads_direct_discriminated_resources(tmp_path):
workspace = tmp_path / "workspace.yaml"
workspace.write_text(
"""
dwh:
type: postgres_direct
connection: &database
host: db
database: analytics
schema: mart
user: reader
password: secret
vectors:
type: pgvector_direct
connection:
<<: *database
schema: vectors
"""
)
cfg = load_config(workspace)
assert isinstance(cfg.dwh, PostgresDwhConfig)
assert cfg.database.transport == "direct"
assert isinstance(cfg.vectors, PgvectorDirectConfig)
assert cfg.vector_db.db_schema == "vectors"
def test_loads_writer_only_http_vector_resource(tmp_path):
workspace = tmp_path / "workspace.yaml"
workspace.write_text(
"""
dwh:
type: thoth_rest
database: {database: analytics, schema: mart}
endpoint: {base_url: https://dwh.test/, api_key: reader}
vectors:
type: thoth_vector_http
writer: {base_url: https://vectors.test/, api_key: writer}
embeddings: {base_url: http://ollama:11434, dim: 768}
"""
)
cfg = load_config(workspace)
assert cfg.vectors.reader is None
assert cfg.vectors.writer.api_key == "writer"
def test_builds_typed_evidence_sources_and_keeps_legacy_compatible(tmp_path):
common = """
dwh:
type: postgres_direct
connection: {database: d, schema: public, user: u, password: p}
"""
modern = tmp_path / "modern.yaml"
modern.write_text(common + f"""
evidence:
sources:
- type: filesystem
root: {tmp_path}
max_bytes: 123
- type: http
urls: ['https://example.test/doc.md?token=transport-only']
""")
cfg = load_config(modern)
assert "transport-only" not in repr(cfg.evidence)
assert "transport-only" not in cfg.evidence.model_dump_json()
assert cfg.evidence.sources[1].allow_private_hosts is False
sources = build_evidence_sources(cfg)
assert isinstance(sources[0], FilesystemEvidenceSource)
assert isinstance(sources[1], HttpManifestEvidenceSource)
assert "transport-only" not in repr(sources[1])
legacy = tmp_path / "legacy.yaml"
(tmp_path / "curated").mkdir()
legacy.write_text(common + f"""
evidence:
source_root: {tmp_path}
evidence_dir: curated
""")
legacy_source = build_evidence_sources(load_config(legacy))[0]
assert isinstance(legacy_source, FilesystemEvidenceSource)
assert legacy_source.root == (tmp_path / "curated").resolve()
+118
View File
@@ -0,0 +1,118 @@
import hashlib
import pytest
from tht.corpus.chunk import ChunkPolicy, chunk
from tht.corpus.models import CanonicalDocument, CorpusManifest
def document(content: str) -> CanonicalDocument:
normalized = content.replace("\r\n", "\n").replace("\r", "\n")
digest = hashlib.sha256(normalized.encode()).hexdigest()
return CanonicalDocument(
document_id="doc:abc",
source_id="source:a",
source_uri="https://host/a.md",
source_fingerprint="etag:abc",
content_hash=f"sha256:{digest}",
title="A",
content=normalized,
media_type="text/markdown",
pipeline_version="pipe:v1",
metadata={"owner": "docs"},
)
def other_document(content: str) -> CanonicalDocument:
return document(content).model_copy(
update={
"document_id": "doc:def",
"source_id": "source:b",
"source_uri": "https://host/b.md",
}
)
def test_chunk_ids_are_stable_for_same_content_and_repeat_runs():
policy = ChunkPolicy(version="paragraph:v1", max_chars=8)
first = chunk(document("A\n\nB"), policy)
second = chunk(document("A\r\n\r\nB"), policy)
repeated = chunk(document("A\n\nB"), policy)
assert first == second == repeated
def test_policy_version_changes_ids_without_changing_boundaries():
doc = document("alpha\n\nbeta")
first = chunk(doc, ChunkPolicy(version="paragraph:v1", max_chars=6))
second = chunk(doc, ChunkPolicy(version="paragraph:v2", max_chars=6))
assert [item.content for item in first] == [item.content for item in second]
assert [item.chunk_id for item in first] != [item.chunk_id for item in second]
def test_same_policy_version_with_different_boundary_config_changes_ids():
doc = document("alpha beta")
first = chunk(doc, ChunkPolicy(version="paragraph:v1", max_chars=6))
second = chunk(doc, ChunkPolicy(version="paragraph:v1", max_chars=7))
assert first[0].chunk_id != second[0].chunk_id
def test_identical_content_in_different_documents_cannot_collide_in_manifest():
policy = ChunkPolicy(version="paragraph:v1", max_chars=20)
first = document("same")
second = other_document("same")
chunks = [*chunk(first, policy), *chunk(second, policy)]
manifest = CorpusManifest(
pipeline_version="pipe:v1", documents=[first, second], chunks=chunks
)
assert len({item.chunk_id for item in manifest.chunks}) == 2
def test_long_non_ascii_tokens_are_hard_split_by_unicode_characters():
chunks = chunk(document("ééééé世界"), ChunkPolicy(version="chars:v1", max_chars=3))
assert [item.content for item in chunks] == ["ééé", "éé世", "界"]
assert all(len(item.content) <= 3 for item in chunks)
@pytest.mark.parametrize(
"content",
[
"alpha beta\tgamma\n\ndelta",
"line with markdown hard break \nnext line\n```\na b\n```",
" \t\n\n \n",
"supercalifragilisticexpialidocious",
"é 世界\r\nnext",
],
)
def test_chunks_preserve_every_character_and_respect_max_chars(content):
doc = document(content)
chunks = chunk(doc, ChunkPolicy(version="exact:v1", max_chars=9))
assert "".join(item.content for item in chunks) == doc.content
assert all(0 < len(item.content) <= 9 for item in chunks)
def test_chunks_have_contiguous_ordinals_hashes_and_provenance_metadata():
doc = document("alpha beta gamma")
chunks = chunk(doc, ChunkPolicy(version="words:v1", max_chars=7))
assert [item.ordinal for item in chunks] == list(range(len(chunks)))
assert len({item.chunk_id for item in chunks}) == len(chunks)
for item in chunks:
assert item.source_uri == doc.source_uri
assert item.document_id == doc.document_id
assert item.pipeline_version == doc.pipeline_version
assert item.metadata["chunk_policy"]["max_chars"] == 7
assert item.metadata["chunk_policy"]["version"] == "words:v1"
assert item.metadata["chunk_policy"]["fingerprint"].startswith("sha256:")
assert item.metadata["document"] == {"owner": "docs"}
assert item.content_hash == "sha256:" + hashlib.sha256(item.content.encode()).hexdigest()
def test_duplicate_chunk_content_cannot_collide_across_ordinals():
chunks = chunk(document("samesame"), ChunkPolicy(version="paragraph:v1", max_chars=4))
assert [item.content for item in chunks] == ["same", "same"]
assert chunks[0].chunk_id != chunks[1].chunk_id
def test_empty_document_has_no_chunks_and_invalid_policy_is_rejected():
assert chunk(document(""), ChunkPolicy(version="v1", max_chars=4)) == []
with pytest.raises(ValueError):
ChunkPolicy(version="v1", max_chars=0)
+217
View File
@@ -0,0 +1,217 @@
import hashlib
from datetime import UTC, datetime, timedelta, timezone
import pytest
from pydantic import ValidationError
from tht.corpus.models import CanonicalChunk, CanonicalDocument, CorpusManifest
def document(source_uri: str = "https://host/a.md") -> CanonicalDocument:
content = "# A"
return CanonicalDocument(
document_id="doc:abc",
source_id="source:a",
source_uri=source_uri,
source_fingerprint="etag:abc",
content_hash=f"sha256:{hashlib.sha256(content.encode()).hexdigest()}",
title="A",
content=content,
media_type="text/markdown",
pipeline_version="evidence-v1",
)
def chunk() -> CanonicalChunk:
content = "# A"
return CanonicalChunk(
chunk_id="chunk:abc:0",
document_id="doc:abc",
ordinal=0,
content=content,
content_hash=f"sha256:{hashlib.sha256(content.encode()).hexdigest()}",
source_uri="https://host/a.md",
pipeline_version="evidence-v1",
)
def test_manifest_contains_provenance_without_credentials():
manifest = CorpusManifest(
manifest_id="manifest:abc",
created_at=datetime(2026, 7, 12, tzinfo=UTC),
pipeline_version="evidence-v1",
embedding_model="nomic-embed-text",
embedding_dimensions=768,
documents=[document()],
chunks=[chunk()],
)
payload = manifest.model_dump_json()
assert "https://host/a.md" in payload
assert "etag:abc" in payload
assert "evidence-v1" in payload
assert "nomic-embed-text" in payload
assert "api_key" not in payload
def test_manifest_can_be_assembled_before_publish_identifiers_are_assigned():
manifest = CorpusManifest(documents=[document()])
assert manifest.documents[0].source_uri == "https://host/a.md"
assert manifest.manifest_id is None
@pytest.mark.parametrize("model", [document(), chunk()])
def test_canonical_records_are_frozen(model):
with pytest.raises(ValidationError):
model.pipeline_version = "changed" # type: ignore[misc]
def test_manifest_collections_are_immutable_tuples_with_json_arrays():
first = CorpusManifest(
manifest_id="manifest:one", created_at=datetime.now(UTC), pipeline_version="v1"
)
second = CorpusManifest(
manifest_id="manifest:two", created_at=datetime.now(UTC), pipeline_version="v1"
)
with pytest.raises(AttributeError):
first.documents.append(document())
assert first.documents == ()
assert second.documents == ()
assert '"documents":[]' in first.model_dump_json()
def test_manifest_validates_embedding_compatibility_fields():
with pytest.raises(ValidationError):
CorpusManifest(
manifest_id="bad",
created_at=datetime.now(UTC),
pipeline_version="v1",
embedding_dimensions=0,
)
def test_canonical_metadata_rejects_secrets_and_non_json_values():
with pytest.raises(ValidationError, match="credential-like"):
CanonicalDocument.model_validate(
{**document().model_dump(), "metadata": {"password": "secret"}}
)
with pytest.raises(ValidationError):
CanonicalChunk(
chunk_id="c",
document_id="d",
ordinal=0,
content="x",
content_hash="sha256:x",
source_uri="file:///x",
pipeline_version="v1",
metadata={"bad": object()},
)
def test_manifest_rejects_duplicate_ids_and_source_ids():
first = document()
duplicate_source = first.model_copy(
update={"document_id": "doc:other", "source_uri": "https://host/b.md"}
)
with pytest.raises(ValidationError, match="source_id"):
CorpusManifest(pipeline_version="evidence-v1", documents=[first, duplicate_source])
with pytest.raises(ValidationError, match="chunk_id"):
CorpusManifest(
pipeline_version="evidence-v1", documents=[first], chunks=[chunk(), chunk()]
)
def test_manifest_rejects_orphan_noncontiguous_and_inconsistent_chunks():
with pytest.raises(ValidationError, match="unknown document"):
CorpusManifest(pipeline_version="evidence-v1", chunks=[chunk()])
second = chunk().model_copy(update={"chunk_id": "chunk:abc:2", "ordinal": 2})
with pytest.raises(ValidationError, match="contiguous"):
CorpusManifest(
pipeline_version="evidence-v1", documents=[document()], chunks=[chunk(), second]
)
wrong_uri = chunk().model_copy(update={"source_uri": "https://host/wrong.md"})
with pytest.raises(ValidationError, match="source_uri"):
CorpusManifest(
pipeline_version="evidence-v1", documents=[document()], chunks=[wrong_uri]
)
def test_manifest_rejects_inconsistent_pipeline_versions():
wrong = document().model_copy(update={"pipeline_version": "other-v1"})
with pytest.raises(ValidationError, match="pipeline_version"):
CorpusManifest(pipeline_version="evidence-v1", documents=[wrong])
def test_vector_generation_requires_embedding_compatibility():
with pytest.raises(ValidationError, match="vector_generation"):
CorpusManifest(pipeline_version="evidence-v1", vector_generation="generation:one")
@pytest.mark.parametrize(
("field", "value"),
[
("document_id", "not-namespaced"),
("content_hash", "sha256:not-hex"),
("source_uri", "https://user:pass@host/a"),
],
)
def test_canonical_document_rejects_malformed_or_sensitive_provenance(field, value):
with pytest.raises(ValidationError):
CanonicalDocument.model_validate({**document().model_dump(), field: value})
@pytest.mark.parametrize(
"source_uri",
[
"https://host/a?X-Amz-Credential=abc&X-Amz-Signature=secret#access_token=bad",
"https://host/a?sig=sas-secret&sp=r#section",
],
)
def test_canonical_provenance_strips_query_and_fragment(source_uri):
doc = document(source_uri=source_uri)
canonical_chunk = chunk().model_copy(update={"source_uri": source_uri})
manifest = CorpusManifest(
pipeline_version="evidence-v1", documents=[doc], chunks=[canonical_chunk]
)
assert doc.source_uri == "https://host/a"
assert canonical_chunk.source_uri == "https://host/a"
payload = manifest.model_dump_json()
assert "X-Amz" not in payload
assert "sas-secret" not in payload
assert "access_token" not in payload
@pytest.mark.parametrize("factory", [document, chunk])
def test_content_hash_must_match_exact_canonical_utf8(factory):
record = factory()
with pytest.raises(ValidationError, match="exact canonical UTF-8 content"):
type(record).model_validate({**record.model_dump(), "content": record.content + "\n"})
def test_model_copy_revalidates_records_and_manifests():
with pytest.raises(ValidationError, match="namespaced"):
document().model_copy(update={"document_id": "invalid"})
manifest = CorpusManifest(
pipeline_version="evidence-v1",
embedding_model="embed-v1",
embedding_dimensions=768,
)
with pytest.raises(ValidationError, match="set together"):
manifest.model_copy(update={"embedding_dimensions": None})
def test_manifest_datetimes_are_aware_and_normalized_to_utc():
with pytest.raises(ValidationError, match="timezone-aware"):
CorpusManifest(created_at=datetime(2026, 7, 12), pipeline_version="evidence-v1")
plus_two = datetime(2026, 7, 12, 12, tzinfo=timezone(timedelta(hours=2)))
manifest = CorpusManifest(created_at=plus_two, pipeline_version="evidence-v1")
assert manifest.created_at.tzinfo is UTC
assert manifest.created_at.hour == 10
+90
View File
@@ -0,0 +1,90 @@
import hashlib
from datetime import UTC, datetime
import pytest
from tht.corpus.normalize import MAX_DOCUMENT_BYTES, PermanentNormalizationError, normalize
from tht.ports.evidence import AcquiredDocument, SourceObject
def acquired(content: bytes, *, media_type: str = "text/markdown") -> AcquiredDocument:
return AcquiredDocument(
source=SourceObject(
source_id="source:guide",
uri="https://host/guide.md?signature=transport#part",
fingerprint="etag:abc",
modified_at=datetime(2026, 7, 12, 12, 0, tzinfo=UTC),
metadata={"owner": "docs"},
),
content=content,
media_type=media_type,
metadata={"transport": "http"},
)
def test_normalize_utf8_bom_newlines_unicode_and_frontmatter():
raw = (
"\ufeff---\r\ntitle: Café\r\ntags: [uno, due]\r\n---\r\n"
"Cafe\u0301\rBody\r\n"
).encode()
document = normalize(acquired(raw, media_type="text/markdown; charset=UTF-8"), "pipe:v1")
assert document.content == "Café\nBody\n"
assert document.title == "Café"
assert document.metadata["frontmatter"] == {"tags": ("uno", "due"), "title": "Café"}
assert document.metadata["source"] == {"owner": "docs"}
assert document.metadata["acquisition"] == {"transport": "http"}
assert document.source_uri == "https://host/guide.md"
assert document.modified_at == datetime(2026, 7, 12, 12, 0, tzinfo=UTC)
assert document.content_hash == "sha256:" + hashlib.sha256(document.content.encode()).hexdigest()
def test_plain_text_that_only_resembles_frontmatter_is_not_dropped():
document = normalize(acquired(b"---\nnot: closed\nbody"), "pipe:v1")
assert document.content == "---\nnot: closed\nbody"
assert "frontmatter" not in document.metadata
def test_frontmatter_can_end_at_eof_without_inventing_content():
document = normalize(acquired(b"---\ntitle: Empty\n---"), "pipe:v1")
assert document.title == "Empty"
assert document.content == ""
@pytest.mark.parametrize(
"frontmatter",
[
"title: first\ntitle: second",
"title: &shared value\ncopy: *shared",
"nested: " + "[" * 25 + "x" + "]" * 25,
"items: [" + ",".join("x" for _ in range(1100)) + "]",
"api_key: secret",
],
)
def test_rejects_unsafe_frontmatter_as_typed_permanent_error(frontmatter):
raw = f"---\n{frontmatter}\n---\nbody".encode()
with pytest.raises(PermanentNormalizationError) as caught:
normalize(acquired(raw), "pipe:v1")
assert caught.value.reason == "invalid_frontmatter"
def test_pipeline_policy_errors_are_not_misclassified_as_bad_frontmatter():
with pytest.raises(ValueError, match="pipeline_version") as caught:
normalize(acquired(b"---\ntitle: valid\n---\nbody"), "")
assert not isinstance(caught.value, PermanentNormalizationError)
@pytest.mark.parametrize(
("content", "media_type", "reason"),
[
(b"bad: \xff", "text/plain", "undecodable"),
(b"hello", "text/plain; charset=iso-8859-1", "unsupported_charset"),
(b"x" * (MAX_DOCUMENT_BYTES + 1), "text/plain", "oversized"),
],
)
def test_rejects_invalid_input_as_typed_permanent_error(content, media_type, reason):
with pytest.raises(PermanentNormalizationError) as caught:
normalize(acquired(content, media_type=media_type), "pipe:v1")
assert caught.value.permanent is True
assert caught.value.reason == reason
+902
View File
@@ -0,0 +1,902 @@
from datetime import UTC, datetime, timedelta
import pytest
from tht.corpus.chunk import ChunkPolicy
from tht.corpus.pipeline import CorpusPipeline, PipelineError, PipelineResult
from tht.corpus.store import CorpusStore
from tht.corpus.models import CanonicalChunk, CanonicalDocument, CorpusManifest
from tht.ports.evidence import AcquiredDocument, SourceObject
from tht.ports.vector import VectorCapabilities, VectorHealth
class Source:
def __init__(self, documents):
self.documents = documents
self.acquire_calls = []
def discover(self):
return [item[0] for item in self.documents]
def acquire(self, item):
self.acquire_calls.append(item.source_id)
payload = next(payload for source, payload in self.documents if source.source_id == item.source_id)
if isinstance(payload, Exception):
raise payload
return AcquiredDocument(
source=item, content=payload.encode(), media_type=item.metadata.get("media_type")
)
class Embedder:
def __init__(self, dim=3, fail=False):
self.dim = dim
self.fail = fail
self.calls = []
def embed_documents(self, texts):
self.calls.extend(texts)
if self.fail:
raise RuntimeError("embed failed")
return [[float(i) for i in range(self.dim)] for _ in texts]
class Vectors:
capabilities = VectorCapabilities(search=True, existing_hashes=True, upsert=True)
def __init__(self, fail=False):
self.fail = fail
self.records = []
self.dimension = 3
def upsert(self, collection, records):
self.records.extend(records[:1] if self.fail else records)
if self.fail:
raise RuntimeError("partial write")
return len(records)
def existing_hashes(self, collection, kinds):
return {
value.record.id: value.content_hash for value in self.records
}
def health(self):
return VectorHealth(
ok=True, expected_dimension=3, observed_dimensions=(self.dimension,),
dimension_compatible=self.dimension == 3,
)
def delete_generation(self, collection, generation, workspace_id):
self.records = [
value for value in self.records
if not (value.record.metadata["vector_generation"] == generation
and value.record.metadata.get("workspace_id") == workspace_id)
]
return 0
def list_evidence_generations(self, collection, workspace_id):
return sorted({
value.record.metadata["vector_generation"] for value in self.records
if value.record.kind == "evidence"
and value.record.metadata.get("workspace_id") == workspace_id
})
class InterruptingVectors(Vectors):
def __init__(self):
super().__init__()
self.batches = []
self.interrupt = True
def upsert(self, collection, records):
self.batches.append([value.record.id for value in records])
if self.interrupt:
self.interrupt = False
self.records.append(records[0])
raise KeyboardInterrupt("process interruption after partial write")
self.records.extend(records)
return len(records)
def item(name, fingerprint):
return SourceObject(
source_id=f"fs:{name}", uri=f"file:///safe/{name}.md", fingerprint=f"sha256:{fingerprint}"
)
def pipeline(tmp_path, source, *, embedder=None, vectors=None, model="model-a", policy=None,
retain=3):
return CorpusPipeline(
store=CorpusStore(tmp_path / "corpus"), sources=[source],
embedder=embedder or Embedder(), vector_store=vectors or Vectors(),
embedding_model=model, embedding_dimensions=3,
chunk_policy=policy or ChunkPolicy(version="chunk-v1", max_chars=100),
pipeline_version="evidence-v1",
retain_published_generations=retain,
)
def test_retention_bounds_generations_and_purges_vectors_after_publish(tmp_path):
vectors = Vectors()
generations = []
for index in range(4):
result = pipeline(
tmp_path, Source([(item("one", str(index)), f"version {index}")]),
vectors=vectors, retain=2,
).run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + str(index) * 64,
)
generations.append(result.generation)
store = CorpusStore(tmp_path / "corpus")
assert store.list_generations() == generations[-2:]
assert {r.record.metadata["vector_generation"] for r in vectors.records} == set(generations[-2:])
assert store.active_generation() == generations[-1]
def test_retention_keeps_filesystem_when_vector_purge_fails_then_retries(tmp_path):
class FailingDelete(Vectors):
def __init__(self):
super().__init__()
self.fail_delete = True
def delete_generation(self, collection, generation, workspace_id):
if self.fail_delete:
raise RuntimeError("credential secret")
return super().delete_generation(collection, generation, workspace_id)
vectors = FailingDelete()
for index in range(2):
pipeline(tmp_path, Source([(item("one", str(index)), str(index))]), vectors=vectors,
retain=1).run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + str(index) * 64,
)
assert len(CorpusStore(tmp_path / "corpus").list_generations()) == 2
vectors.fail_delete = False
report = pipeline(tmp_path, Source([(item("one", "1"), "1")]), vectors=vectors,
retain=1).gc(workspace_root=tmp_path)
assert report["status"] == "succeeded"
assert len(CorpusStore(tmp_path / "corpus").list_generations()) == 1
def test_gc_reconciles_vector_only_generation(tmp_path):
vectors = Vectors()
orphan = "gen:" + "f" * 32
from tht.ports.vector import VectorWriteRecord
from tht.vectorstore.records import VectorRecord
vectors.records.append(VectorWriteRecord(
record=VectorRecord(id="orphan", kind="evidence", ref="doc:x", title="", content="x",
metadata={"vector_generation": orphan, "workspace_id": "default"}),
embedding=[0.0, 0.0, 0.0], content_hash="sha256:" + "0" * 64,
))
candidate = pipeline(tmp_path, Source([]), vectors=vectors, retain=1)
report = candidate.gc(workspace_root=tmp_path)
assert report["evicted"] == [orphan]
assert vectors.list_evidence_generations("evidence", "default") == []
assert candidate.gc(workspace_root=tmp_path)["evicted"] == []
@pytest.mark.parametrize("status", ["running", "failed"])
def test_gc_protects_generations_referenced_by_resumable_checkpoints(tmp_path, status):
generation = "gen:" + "e" * 32
store = CorpusStore(tmp_path / "corpus")
store.stage(CorpusManifest(), {}, generation=generation)
run = tmp_path / ".tht-jobs" / "evidence" / "runs" / ("a" * 32)
(run / "artifacts").mkdir(parents=True)
(run / "checkpoint.json").write_text(__import__("json").dumps({"status": status}))
(run / "artifacts" / "plan.json").write_text(
__import__("json").dumps({"generation": generation})
)
candidate = pipeline(tmp_path, Source([]), vectors=Vectors(), retain=1)
report = candidate.gc(workspace_root=tmp_path)
assert generation in report["protected"]
assert store.generation_path(generation).exists()
def test_explicit_gc_blocks_while_job_holds_corpus_writer_lock(tmp_path):
import threading
candidate = pipeline(tmp_path, Source([(item("one", "a"), "one")]), vectors=Vectors())
entered = threading.Event()
release = threading.Event()
gc_finished = threading.Event()
def pause(_context, stage):
if stage == "discover":
entered.set()
assert release.wait(5)
job = threading.Thread(target=lambda: candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
after_stage_return=pause,
))
job.start()
assert entered.wait(5)
def collect():
candidate.gc(workspace_root=tmp_path)
gc_finished.set()
gc_thread = threading.Thread(target=collect)
gc_thread.start()
assert not gc_finished.wait(0.1)
release.set()
job.join(5)
gc_thread.join(5)
assert gc_finished.is_set()
assert candidate.store.active_generation() is not None
def test_gc_preserves_vector_dependencies_of_retained_manifests(tmp_path):
vectors = Vectors()
one = item("one", "a")
first = pipeline(tmp_path, Source([(one, "stable")]), vectors=vectors, retain=2).run().generation
second = pipeline(
tmp_path, Source([(one, "stable"), (item("two", "b"), "two")]),
vectors=vectors, retain=2,
).run().generation
third = pipeline(
tmp_path, Source([(one, "stable"), (item("two", "c"), "changed")]),
vectors=vectors, retain=2,
).run().generation
assert CorpusStore(tmp_path / "corpus").list_generations() == [second, third]
assert first in vectors.list_evidence_generations("evidence", "default")
def test_active_searcher_without_active_fails_closed_for_evidence(tmp_path):
from types import SimpleNamespace
from tht.search.evidence import active_searcher
class Delegate:
def search(self, embedding, top_n=10, kinds=None, metadata_filter=None):
return ["legacy"]
cfg = SimpleNamespace(paths=SimpleNamespace(artifacts=tmp_path / "artifacts"))
wrapped = active_searcher(cfg, Delegate())
assert wrapped.search([1.0], kinds=["evidence"]) == []
assert wrapped.search([1.0], kinds=["memory"]) == ["legacy"]
def test_active_searcher_splits_default_and_mixed_kinds_before_global_limit(tmp_path):
from types import SimpleNamespace
from tht.search.evidence import ActiveEvidenceSearcher
store = CorpusStore(tmp_path / "corpus")
generation = store.stage(
CorpusManifest(metadata={"workspace_id": "default"}), {},
generation="gen:" + "a" * 32,
)
store.publish(generation)
calls = []
class Delegate:
def search(self, embedding, top_n=10, kinds=None, metadata_filter=None):
calls.append((kinds, metadata_filter))
if kinds == ["evidence"]:
return [SimpleNamespace(id="active", similarity=0.8)]
return [SimpleNamespace(id="memory", similarity=0.9)]
searcher = ActiveEvidenceSearcher(store, Delegate())
hits = searcher.search([1.0], top_n=1, kinds=["evidence", "memory"])
assert [hit.id for hit in hits] == ["memory"]
assert calls[0] == (["memory"], None)
# Empty manifest means no Evidence query, but the split remains explicit and safe.
assert all(call[0] != ["evidence"] for call in calls)
calls.clear()
searcher.search([1.0], top_n=1)
assert calls[0][0] == ["memory", "schema_column", "schema_table", "solved_question"]
assert all(call[0] is not None for call in calls)
def test_active_evidence_query_holds_lock_against_publish(tmp_path):
import threading
from types import SimpleNamespace
from tht.search.evidence import ActiveEvidenceSearcher
first_pipeline = pipeline(tmp_path, Source([(item("one", "a"), "old")]), vectors=Vectors())
first_pipeline.run()
store = first_pipeline.store
entered = threading.Event()
release = threading.Event()
published = threading.Event()
class Delegate:
def search(self, embedding, top_n=10, kinds=None, metadata_filter=None):
entered.set()
assert release.wait(5)
return [SimpleNamespace(id="active", similarity=1.0)]
search = threading.Thread(
target=lambda: ActiveEvidenceSearcher(store, Delegate()).search(
[1.0], kinds=["evidence"]
)
)
search.start()
assert entered.wait(5)
next_generation = store.stage(CorpusManifest(), {})
def publish():
with store.writer_lock():
store.publish(next_generation)
published.set()
publisher = threading.Thread(target=publish)
publisher.start()
assert not published.wait(0.1)
release.set()
search.join(5)
publisher.join(5)
assert published.is_set()
def test_pipeline_result_dump_does_not_deepcopy_frozen_metadata():
manifest = CorpusManifest(metadata={"nested": {"value": ["safe"]}})
payload = PipelineResult(
"succeeded", None, False, (), (), (), manifest,
).model_dump(mode="json")
assert payload["manifest_id"] is None
assert "manifest" not in payload
def test_pipeline_result_public_dump_is_bounded_and_excludes_evidence_content(tmp_path):
import json
result = pipeline(
tmp_path, Source([(item("one", "a"), "SENSITIVE EVIDENCE CONTENT")])
).run()
payload = result.model_dump(mode="json")
encoded = json.dumps(payload)
assert "SENSITIVE EVIDENCE CONTENT" not in encoded
assert "documents" not in payload and "chunks" not in payload
large = PipelineResult(
"failed", None, False,
tuple(f"fs:item-{index}" for index in range(1000)), (), (), result.manifest,
).model_dump(mode="json")
assert len(large["changed"]) == 100
assert large["counts"]["changed"] == 1000
assert len(json.dumps(large)) < 25_000
def test_pipeline_result_repr_is_bounded_and_excludes_manifest_secrets():
secret = "TOP_SECRET_CONTENT"
manifest = CorpusManifest.model_construct(
manifest_id="gen:" + "a" * 64,
documents=tuple(CanonicalDocument.model_construct(content=secret) for _ in range(1000)),
chunks=tuple(CanonicalChunk.model_construct(content=secret) for _ in range(1000)),
metadata={"password": secret, "credential": "Bearer " + secret},
)
result = PipelineResult(
"succeeded", "gen:" + "a" * 64, True, (), (), (), manifest,
run_id="b" * 32,
)
rendered = repr(result)
assert str(result) == rendered
assert len(rendered) < 1000
assert secret not in rendered
assert "password" not in rendered
assert "credential" not in rendered
assert "manifest" not in rendered.lower()
assert "documents" not in rendered
assert "chunks" not in rendered
def test_reused_corpus_root_rejects_workspace_rename_before_any_mutation(tmp_path):
vectors = Vectors()
first = pipeline(tmp_path, Source([(item("one", "a"), "stable")]), vectors=vectors)
first.run_as_job(
workspace_id="workspace-a", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
active = first.store.active_generation()
records = list(vectors.records)
renamed_source = Source([(item("one", "a"), "stable")])
renamed = pipeline(tmp_path, renamed_source, vectors=vectors)
with pytest.raises(PipelineError, match="different workspace"):
renamed.run_as_job(
workspace_id="workspace-b", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
assert renamed_source.acquire_calls == []
assert renamed.store.active_generation() == active
assert vectors.records == records
def test_gc_rejects_workspace_mismatch_without_deleting(tmp_path):
vectors = Vectors()
owner = pipeline(tmp_path, Source([(item("one", "a"), "stable")]), vectors=vectors)
owner.run_as_job(
workspace_id="workspace-a", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
generations = owner.store.list_generations()
wrong = pipeline(tmp_path, Source([]), vectors=vectors)
wrong.workspace_id = "workspace-b"
with pytest.raises(PipelineError, match="different workspace"):
wrong.gc(workspace_root=tmp_path)
assert wrong.store.list_generations() == generations
@pytest.mark.parametrize("kinds", [None, ["evidence", "memory"], ["memory"]])
def test_active_search_rejects_workspace_mismatch_before_delegate(tmp_path, kinds):
from tht.search.evidence import ActiveEvidenceSearcher, CorpusWorkspaceMismatchError
vectors = Vectors()
owner = pipeline(tmp_path, Source([(item("one", "a"), "stable")]), vectors=vectors)
owner.run_as_job(
workspace_id="workspace-a", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
class Delegate:
def search(self, *args, **kwargs):
raise AssertionError("workspace mismatch reached vector delegate")
with pytest.raises(CorpusWorkspaceMismatchError, match="different workspace"):
ActiveEvidenceSearcher(
owner.store, Delegate(), expected_workspace_id="workspace-b",
).search([1.0], kinds=kinds)
def test_unscoped_active_manifest_is_never_adopted_by_direct_run_or_gc(tmp_path):
store = CorpusStore(tmp_path / "corpus")
generation = store.stage(CorpusManifest(), {})
store.publish(generation)
class ForbiddenSource:
def discover(self):
raise AssertionError("invalid corpus reached source discovery")
class ForbiddenVectors:
def __getattr__(self, name):
raise AssertionError(f"invalid corpus reached vector operation {name}")
candidate = CorpusPipeline(
store=store, sources=[ForbiddenSource()], embedder=Embedder(),
vector_store=ForbiddenVectors(), embedding_model="model", embedding_dimensions=3,
chunk_policy=ChunkPolicy(version="chunk-v1", max_chars=100),
pipeline_version="evidence-v1",
)
with pytest.raises(PipelineError, match="missing or invalid"):
candidate.run()
with pytest.raises(PipelineError, match="missing or invalid"):
candidate.gc(workspace_root=tmp_path)
assert store.active_generation() == generation
def test_unchanged_documents_skip_acquire_normalize_chunk_and_embed(tmp_path):
one = item("one", "a")
first_source = Source([(one, "hello")])
first = pipeline(tmp_path, first_source)
first.run()
second_source = Source([(one, "ignored")])
second_embedder = Embedder()
result = pipeline(tmp_path, second_source, embedder=second_embedder).run()
assert result.unchanged == ("fs:one",)
assert second_source.acquire_calls == []
assert second_embedder.calls == []
def test_unchanged_job_reuses_active_generation_without_new_directory(tmp_path):
source = Source([(
SourceObject(
source_id="fs:one", uri="file:///safe/one.md", fingerprint="sha256:a",
modified_at=datetime(2026, 1, 1, tzinfo=UTC),
metadata={"media_type": "text/markdown", "size": 5, "nested": {"b": 2, "a": 1}},
),
"hello",
)])
candidate = pipeline(tmp_path, source)
args = dict(workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64)
first = candidate.run_as_job(**args)
snapshot = first.manifest.metadata["source_snapshot"]["fs:one"]
assert snapshot == {
"source_id": "fs:one", "uri": "file:///safe/one.md", "fingerprint": "sha256:a",
"modified_at": "2026-01-01T00:00:00Z",
"metadata": {"media_type": "text/markdown", "size": 5,
"nested": {"a": 1, "b": 2}},
"media_type": "text/markdown", "size": 5,
}
count = len(candidate.store.list_generations())
second = candidate.run_as_job(**args)
assert second.generation == first.generation
assert second.published is False
assert len(candidate.store.list_generations()) == count
@pytest.mark.parametrize("field", ["uri", "modified_at", "metadata"])
def test_job_source_snapshot_change_forces_publish_with_same_fingerprint(tmp_path, field):
original = SourceObject(
source_id="fs:one", uri="file:///safe/one.md", fingerprint="sha256:a",
modified_at=datetime(2026, 1, 1, tzinfo=UTC),
metadata={"media_type": "text/markdown", "size": 5, "label": "original"},
)
vectors = Vectors()
args = dict(workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64)
first = pipeline(tmp_path, Source([(original, "hello")]), vectors=vectors).run_as_job(**args)
updates = {
"uri": "file:///safe/renamed.md",
"modified_at": original.modified_at + timedelta(seconds=1),
"metadata": {"media_type": "text/markdown", "size": 5, "label": "changed"},
}
changed = original.model_copy(update={field: updates[field]})
source = Source([(changed, "hello")])
result = pipeline(tmp_path, source, vectors=vectors).run_as_job(**args)
assert result.published is True
assert result.generation != first.generation
assert source.acquire_calls == ["fs:one"]
@pytest.mark.parametrize("fingerprint_name", ["config_fingerprint", "input_fingerprint"])
def test_job_binding_change_forces_publish(tmp_path, fingerprint_name):
vectors = Vectors()
source_object = item("one", "a")
args = dict(workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64)
first = pipeline(tmp_path, Source([(source_object, "hello")]), vectors=vectors).run_as_job(**args)
args[fingerprint_name] = "sha256:" + "3" * 64
source = Source([(source_object, "hello")])
result = pipeline(tmp_path, source, vectors=vectors).run_as_job(**args)
assert result.published is True
assert result.generation != first.generation
assert source.acquire_calls == []
@pytest.mark.parametrize(
"damage", ["legacy_metadata", "corrupt_document_sources", "document", "vector"]
)
def test_job_incomplete_active_contract_never_noops(tmp_path, damage):
import json
vectors = Vectors()
source_object = item("one", "a")
args = dict(workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64)
candidate = pipeline(tmp_path, Source([(source_object, "hello")]), vectors=vectors)
first = candidate.run_as_job(**args)
if damage == "legacy_metadata":
manifest_path = candidate.store.generation_path(first.generation) / "manifest.json"
payload = json.loads(manifest_path.read_text())
payload["metadata"].pop("source_snapshot")
manifest_path.write_text(json.dumps(payload))
elif damage == "corrupt_document_sources":
manifest_path = candidate.store.generation_path(first.generation) / "manifest.json"
payload = json.loads(manifest_path.read_text())
payload["metadata"]["document_sources"] = {}
manifest_path.write_text(json.dumps(payload))
elif damage == "document":
path = candidate.store.resolve_document(first.manifest.documents[0].document_id)
path.unlink()
else:
vectors.records.clear()
source = Source([(source_object, "hello")])
result = pipeline(tmp_path, source, vectors=vectors).run_as_job(**args)
assert result.published is True
assert result.generation != first.generation
assert source.acquire_calls == ["fs:one"]
@pytest.mark.parametrize(
"damage", ["modified_at", "source_metadata", "media_type", "missing_chunk",
"altered_chunk", "extra_chunk", "vector_dimension"]
)
def test_job_corrupt_canonical_document_or_chunk_never_noops(tmp_path, damage):
import hashlib
import json
vectors = Vectors()
source_object = SourceObject(
source_id="fs:one", uri="file:///safe/one.md", fingerprint="sha256:a",
modified_at=datetime(2026, 1, 1, tzinfo=UTC),
metadata={"media_type": "text/markdown", "size": 11, "owner": "docs"},
)
args = dict(workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64)
candidate = pipeline(
tmp_path, Source([(source_object, "hello world")]), vectors=vectors,
policy=ChunkPolicy(version="chunk-v1", max_chars=6),
)
first = candidate.run_as_job(**args)
manifest_path = candidate.store.generation_path(first.generation) / "manifest.json"
payload = json.loads(manifest_path.read_text())
document = payload["documents"][0]
chunks = payload["chunks"]
if damage == "modified_at":
document["modified_at"] = "2026-01-01T00:00:01Z"
elif damage == "source_metadata":
document["metadata"]["source"]["owner"] = "attacker"
elif damage == "media_type":
document["media_type"] = "text/plain"
elif damage == "missing_chunk":
payload["chunks"] = chunks[:-1]
elif damage == "altered_chunk":
chunks[0]["content"] = "HELLO "
chunks[0]["content_hash"] = "sha256:" + hashlib.sha256(b"HELLO ").hexdigest()
chunks[0]["chunk_id"] = "chunk:" + "a" * 64
elif damage == "extra_chunk":
extra = CanonicalChunk(
chunk_id="chunk:" + "b" * 64, document_id=document["document_id"],
ordinal=len(chunks), content="", content_hash="sha256:" + hashlib.sha256(b"").hexdigest(),
source_uri=document["source_uri"], pipeline_version=document["pipeline_version"],
)
chunks.append(extra.model_dump(mode="json"))
else:
vectors.dimension = 4
manifest_path.write_text(json.dumps(payload))
source = Source([(source_object, "hello world")])
result = pipeline(
tmp_path, source, vectors=vectors,
policy=ChunkPolicy(version="chunk-v1", max_chars=6),
).run_as_job(**args)
assert result.published is True
assert result.generation != first.generation
assert source.acquire_calls == ["fs:one"]
def test_removed_documents_are_marked_and_absent_from_new_manifest(tmp_path):
one, two = item("one", "a"), item("two", "b")
pipeline(tmp_path, Source([(one, "one"), (two, "two")])).run()
result = pipeline(tmp_path, Source([(one, "one")])).run()
assert result.removed == ("fs:two",)
assert {doc.source_id for doc in result.manifest.documents} == {"fs:one"}
def test_model_or_chunk_policy_change_forces_full_rebuild(tmp_path):
one = item("one", "a")
pipeline(tmp_path, Source([(one, "hello")])).run()
source = Source([(one, "hello")])
changed = pipeline(tmp_path, source, model="model-b").run()
assert changed.changed == ("fs:one",)
assert source.acquire_calls == ["fs:one"]
def test_partial_vector_failure_never_changes_active_or_exposes_generation(tmp_path):
one = item("one", "a")
good = pipeline(tmp_path, Source([(one, "old")]))
old = good.run().generation
changed = item("one", "b")
vectors = Vectors(fail=True)
broken = pipeline(tmp_path, Source([(changed, "new")]), vectors=vectors)
with pytest.raises(PipelineError):
broken.run()
assert broken.store.active_generation() == old
assert vectors.records[0].record.metadata["vector_generation"] != old
def test_dimension_mismatch_fails_before_vector_write_and_publish(tmp_path):
one = item("one", "a")
vectors = Vectors()
candidate = pipeline(tmp_path, Source([(one, "hello")]), embedder=Embedder(dim=2), vectors=vectors)
with pytest.raises(PipelineError, match="dimension"):
candidate.run()
assert vectors.records == []
assert candidate.store.active_generation() is None
def test_dry_run_and_failed_acquire_never_change_active(tmp_path):
one = item("one", "a")
active = pipeline(tmp_path, Source([(one, "old")])).run().generation
changed = item("one", "b")
dry = pipeline(tmp_path, Source([(changed, "new")])).run(dry_run=True)
assert dry.published is False
assert dry.generation is None
assert dry.manifest.documents[0].content == "old"
with pytest.raises(PipelineError):
pipeline(tmp_path, Source([(changed, RuntimeError("boom"))])).run()
assert CorpusStore(tmp_path / "corpus").active_generation() == active
def test_job_pipeline_uses_ordered_plan_and_returns_run_id(tmp_path):
one = item("one", "a")
candidate = pipeline(tmp_path, Source([(one, "hello")]))
result = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
assert result.status == "succeeded"
assert result.run_id and len(result.run_id) == 32
checkpoint = tmp_path / ".tht-jobs" / "evidence" / "runs" / result.run_id / "checkpoint.json"
payload = __import__("json").loads(checkpoint.read_text())
assert [stage["name"] for stage in payload["stages"]] == [
"discover", "acquire_normalize_chunk", "embed", "vector_upsert",
"stage_validate", "publish", "retention_cleanup",
]
def test_job_pipeline_dry_run_only_discovers_and_reports_changes(tmp_path):
one = item("one", "a")
source = Source([(one, "hello")])
embedder = Embedder()
vectors = Vectors()
result = pipeline(tmp_path, source, embedder=embedder, vectors=vectors).run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
dry_run=True,
)
assert result.changed == ("fs:one",)
assert source.acquire_calls == []
assert embedder.calls == []
assert vectors.records == []
assert result.generation is None and result.published is False
@pytest.mark.parametrize("crash_stage", [
"discover", "acquire_normalize_chunk", "embed", "vector_upsert",
"stage_validate", "publish", "retention_cleanup",
])
def test_job_pipeline_crash_after_each_stage_resumes_without_duplicate_effects(tmp_path, crash_stage):
one = item("one", "a")
source = Source([(one, "hello")])
embedder = Embedder()
vectors = Vectors()
candidate = pipeline(tmp_path, source, embedder=embedder, vectors=vectors)
class Crash(BaseException):
pass
def fault(_context, stage):
if stage == crash_stage:
raise Crash()
with pytest.raises(Crash):
candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
after_stage_return=fault,
)
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
crashed = next(runs.iterdir()).name
result = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
resume_run_id=crashed,
)
assert result.status == "succeeded"
assert source.acquire_calls == ["fs:one"]
assert len(embedder.calls) == 1
assert len(vectors.records) == 1
def test_job_pipeline_raw_upsert_failure_compensates_and_resumes_with_new_generation(tmp_path):
one = item("one", "a")
vectors = Vectors(fail=True)
candidate = pipeline(tmp_path, Source([(one, "hello")]), vectors=vectors)
first = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
assert first.status == "failed"
assert vectors.records == []
old_generation = first.generation
vectors.fail = False
resumed = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
resume_run_id=first.run_id,
)
assert resumed.status == "succeeded", resumed
assert resumed.generation != old_generation
assert candidate.store.active_generation() == resumed.generation
def test_job_pipeline_raw_stage_failure_compensates_vectors_and_resumes(tmp_path, monkeypatch):
one = item("one", "a")
vectors = Vectors()
candidate = pipeline(tmp_path, Source([(one, "hello")]), vectors=vectors)
real_stage = candidate.store.stage
calls = 0
def fail_once(*args, **kwargs):
nonlocal calls
calls += 1
if calls == 1:
raise OSError("raw stage failure")
return real_stage(*args, **kwargs)
monkeypatch.setattr(candidate.store, "stage", fail_once)
first = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
assert first.status == "failed" and vectors.records == []
resumed = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
resume_run_id=first.run_id,
)
assert resumed.status == "succeeded", resumed
@pytest.mark.parametrize("stage,filename", [
("discover", "plan.json"),
("acquire_normalize_chunk", "manifest.json"),
("embed", "embeddings.json"),
])
@pytest.mark.parametrize("mutation", ["missing", "tampered"])
def test_job_pipeline_rejects_corrupt_required_artifacts_before_resume(
tmp_path, stage, filename, mutation,
):
one = item("one", "a")
candidate = pipeline(tmp_path, Source([(one, "hello")]))
class Crash(BaseException):
pass
with pytest.raises(Crash):
candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
after_stage_return=lambda _context, name: (
(_ for _ in ()).throw(Crash()) if name == stage else None
),
)
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
crashed = next(runs.iterdir())
target = crashed / "artifacts" / filename
target.unlink() if mutation == "missing" else target.write_text("tampered")
from tht.jobs.runner import CorruptCheckpointError
with pytest.raises(CorruptCheckpointError, match="artifact"):
candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
resume_run_id=crashed.name,
)
def test_vector_intent_is_reconciled_after_process_interruption_without_duplicate_upsert(tmp_path):
one = item("one", "a")
vectors = InterruptingVectors()
candidate = pipeline(
tmp_path, Source([(one, "a" * 250)]), vectors=vectors,
policy=ChunkPolicy(version="chunk-v1", max_chars=100),
)
with pytest.raises(KeyboardInterrupt):
candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
interrupted = next(runs.iterdir())
checkpoint = __import__("json").loads((interrupted / "checkpoint.json").read_text())
vector_stage = checkpoint["stages"][3]
assert vector_stage["status"] == "running"
assert vector_stage["effect_state"] == "intent"
first_written = vectors.batches[0][0]
result = candidate.run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
resume_run_id=interrupted.name,
)
assert result.status == "succeeded" and result.published is True
assert first_written not in vectors.batches[1]
assert len(vectors.records) == 3
+166
View File
@@ -0,0 +1,166 @@
import pytest
import os
from tht.corpus.models import CorpusManifest
from tht.corpus.store import CorpusStore, UnsafeCorpusPath
def test_publish_switches_active_atomically_and_resolves_materialized_files(tmp_path):
store = CorpusStore(tmp_path / "corpus")
generation = store.stage(CorpusManifest(), {})
seen = []
store._replace = lambda source, target: (seen.append(source.read_text()), source.replace(target))
published = store.publish(generation)
assert published == generation
assert store.active_generation() == generation
assert seen == [generation + "\n"]
def test_active_manifest_is_a_consistent_reader_snapshot(tmp_path):
store = CorpusStore(tmp_path / "corpus")
first = store.stage(CorpusManifest(metadata={"name": "first"}), {})
second = store.stage(CorpusManifest(metadata={"name": "second"}), {})
store.publish(first)
snapshot = store.active_manifest()
store.publish(second)
assert snapshot.metadata["name"] == "first"
assert store.active_manifest().metadata["name"] == "second"
def test_store_rejects_symlinked_generation_root(tmp_path):
outside = tmp_path / "outside"
outside.mkdir()
root = tmp_path / "corpus"
root.symlink_to(outside, target_is_directory=True)
with pytest.raises(UnsafeCorpusPath):
CorpusStore(root)
def test_active_pointer_cannot_escape_generation_root(tmp_path):
store = CorpusStore(tmp_path / "corpus")
store.root.mkdir(parents=True, exist_ok=True)
store.active_path.write_text("../outside\n")
with pytest.raises(UnsafeCorpusPath):
store.active_manifest()
def test_publish_restores_previous_active_when_directory_fsync_fails_after_replace(tmp_path, monkeypatch):
store = CorpusStore(tmp_path / "corpus")
first = store.stage(CorpusManifest(), {})
second = store.stage(CorpusManifest(), {})
store.publish(first)
def fail_once():
store._fsync_directory = store._sync_root
raise OSError("post replace crash")
store._fsync_directory = fail_once
with pytest.raises(OSError, match="post replace"):
store.publish(second)
assert store.active_generation() == first
def test_read_document_rejects_symlink_hardlink_and_hash_mismatch(tmp_path):
from tht.corpus.models import CanonicalDocument
content = "trusted"
digest = "sha256:" + __import__("hashlib").sha256(content.encode()).hexdigest()
document = CanonicalDocument(
document_id="doc:" + "a" * 64, source_id="fs:one", source_uri="file:///one",
source_fingerprint="sha256:" + "b" * 64, content_hash=digest, content=content,
pipeline_version="evidence-v1",
)
store = CorpusStore(tmp_path / "corpus")
generation = store.stage(CorpusManifest(documents=(document,)), {document.document_id: content})
path = store.resolve_document(document.document_id, generation)
assert store.read_document(document.document_id, generation) == content
path.unlink()
path.symlink_to(tmp_path / "outside")
(tmp_path / "outside").write_text(content)
with pytest.raises(UnsafeCorpusPath):
store.read_document(document.document_id, generation)
path.unlink()
os.link(tmp_path / "outside", path)
with pytest.raises(UnsafeCorpusPath):
store.read_document(document.document_id, generation)
path.unlink()
path.write_text("tampered")
with pytest.raises(UnsafeCorpusPath):
store.read_document(document.document_id, generation)
def test_generation_inventory_is_validated_and_sorted(tmp_path):
store = CorpusStore(tmp_path / "corpus")
first = store.stage(CorpusManifest(), {}, generation="gen:" + "1" * 32)
second = store.stage(CorpusManifest(), {}, generation="gen:" + "2" * 32)
(store.root / "unrelated").mkdir()
assert store.list_generations() == [first, second]
def test_published_inventory_excludes_staged_and_invalid_newer_directories(tmp_path):
store = CorpusStore(tmp_path / "corpus")
first = store.stage(CorpusManifest(), {}, generation="gen:" + "1" * 32)
store.publish(first)
store.stage(CorpusManifest(), {}, generation="gen:" + "2" * 32)
invalid = store.generation_path("gen:" + "3" * 32)
invalid.mkdir()
(invalid / "PUBLISHED").write_text("2026-01-01T00:00:00Z\n")
assert store.published_generations() == [first]
def test_owned_copy_uses_validated_descriptor_bytes_when_source_is_replaced(tmp_path, monkeypatch):
from tht.corpus.models import CanonicalDocument
import hashlib
content = "active bytes"
document = CanonicalDocument(
document_id="doc:" + "c" * 64, source_id="fs:copy", source_uri="file:///copy",
source_fingerprint="sha256:" + "d" * 64,
content_hash="sha256:" + hashlib.sha256(content.encode()).hexdigest(),
content=content, pipeline_version="evidence-v1",
)
store = CorpusStore(tmp_path / "corpus")
generation = store.stage(CorpusManifest(documents=(document,)), {document.document_id: content})
store.publish(generation)
source = store.resolve_document(document.document_id)
real_read = os.read
def replace_after_read(fd, size):
payload = real_read(fd, size)
source.unlink()
source.write_text("replacement")
return payload
monkeypatch.setattr(os, "read", replace_after_read)
owned = store.materialize_document(document.document_id, tmp_path / "session" / "evidence.md")
assert owned.read_text() == content
assert hashlib.sha256(owned.read_bytes()).hexdigest() == document.content_hash.removeprefix("sha256:")
def test_materialized_snapshot_uses_identified_manifest_when_active_changes(tmp_path):
from tht.corpus.models import CanonicalDocument
import hashlib
def doc(content, fingerprint):
return CanonicalDocument(
document_id="doc:" + hashlib.sha256(content.encode()).hexdigest(),
source_id="fs:item", source_uri="file:///item",
source_fingerprint="sha256:" + fingerprint * 64,
content_hash="sha256:" + hashlib.sha256(content.encode()).hexdigest(),
content=content, pipeline_version="evidence-v1",
)
store = CorpusStore(tmp_path / "corpus")
old = doc("old", "a")
old_generation = store.stage(CorpusManifest(documents=(old,)), {old.document_id: old.content})
store.publish(old_generation)
snapshot = store.active_manifest()
new = doc("new", "b")
new_generation = store.stage(CorpusManifest(documents=(new,)), {new.document_id: new.content})
store.publish(new_generation)
path = store.materialize_document(
snapshot.documents[0].document_id, tmp_path / "owned.md", generation=snapshot.manifest_id,
)
assert path.read_text() == "old"
+156
View File
@@ -0,0 +1,156 @@
import json
from pathlib import Path
from typer.testing import CliRunner
from tht.cli import app
runner = CliRunner()
def _config(path: Path, *, absolute_sessions: Path | None = None) -> Path:
sessions = absolute_sessions or Path("sessions")
path.write_text(
"dwh:\n"
" type: postgres_direct\n"
" connection:\n"
" database: patient_db\n"
" schema: private_schema\n"
" user: pii_user\n"
" password: super-secret\n"
"roots:\n"
" artifacts: artifacts\n"
" indexes: indexes\n"
f" sessions: {sessions}\n"
)
return path
def test_doctor_json_reports_portable_path_statuses_without_secrets(monkeypatch, tmp_path):
cfg = _config(tmp_path / "demo.yaml")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
result = runner.invoke(app, ["doctor", "--json", "--config", str(cfg)])
assert result.exit_code == 0
payload = json.loads(result.stdout)
assert payload == {
"ok": True,
"components": {
"config": {"status": "ok"},
"data_root": {"status": "ok"},
"workspace_paths": {"status": "ok", "legacy_absolute": []},
},
}
assert "super-secret" not in result.stdout
assert "patient_db" not in result.stdout
assert "pii_user" not in result.stdout
assert result.stderr == ""
def test_doctor_json_flags_absolute_legacy_paths(monkeypatch, tmp_path):
cfg = _config(tmp_path / "demo.yaml", absolute_sessions=tmp_path / "old-sessions")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
result = runner.invoke(app, ["doctor", "--json", "--config", str(cfg)])
assert result.exit_code == 0
payload = json.loads(result.stdout)
assert payload["components"]["workspace_paths"] == {
"status": "warning",
"legacy_absolute": ["sessions"],
}
assert str(tmp_path) not in result.stdout
def test_doctor_json_returns_structured_config_error(monkeypatch, tmp_path):
cfg = _config(tmp_path / "demo.yaml")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
cfg.write_text(cfg.read_text().replace("sessions: sessions", "sessions: ../../private"))
result = runner.invoke(app, ["doctor", "--json", "--config", str(cfg)])
assert result.exit_code == 1
payload = json.loads(result.stdout)
assert payload["ok"] is False
assert payload["components"]["workspace_paths"]["status"] == "error"
assert "outside workspace root" in payload["components"]["workspace_paths"]["message"]
assert result.stderr == ""
def test_doctor_json_does_not_echo_invalid_config_values(monkeypatch, tmp_path):
cfg = _config(tmp_path / "demo.yaml")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
cfg.write_text(cfg.read_text().replace("database: patient_db", ""))
result = runner.invoke(app, ["doctor", "--json", "--config", str(cfg)])
assert result.exit_code == 1
payload = json.loads(result.stdout)
assert payload["components"]["config"] == {
"status": "error",
"message": "configuration is invalid or unreadable",
}
assert "super-secret" not in result.stdout
assert "pii_user" not in result.stdout
def test_doctor_json_normalizes_malformed_yaml(monkeypatch, tmp_path):
cfg = tmp_path / "demo.yaml"
cfg.write_text("password: super-secret\nroots: [unterminated")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
result = runner.invoke(app, ["doctor", "--json", "--config", str(cfg)])
assert result.exit_code == 1
assert json.loads(result.stdout)["components"]["config"] == {
"status": "error",
"message": "configuration is invalid or unreadable",
}
assert result.stderr == ""
assert "super-secret" not in result.stdout
assert "Traceback" not in result.stdout
def test_doctor_json_normalizes_unreadable_config(monkeypatch, tmp_path):
cfg = tmp_path / "demo.yaml"
cfg.mkdir()
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
result = runner.invoke(app, ["doctor", "--json", "--config", str(cfg)])
assert result.exit_code == 1
assert json.loads(result.stdout)["components"]["config"] == {
"status": "error",
"message": "configuration is invalid or unreadable",
}
assert result.stderr == ""
def test_doctor_human_output_is_actionable_and_redacted(monkeypatch, tmp_path):
cfg = _config(tmp_path / "demo.yaml", absolute_sessions=tmp_path / "patient-private")
monkeypatch.delenv("THT_DATA_ROOT", raising=False)
result = runner.invoke(app, ["doctor", "--config", str(cfg)])
assert result.exit_code == 0
assert "data_root: warning - set THT_DATA_ROOT to enable portable storage" in result.stdout
assert "workspace_paths: warning - absolute legacy roots: sessions" in result.stdout
assert str(tmp_path) not in result.stdout
assert "patient_db" not in result.stdout
assert "super-secret" not in result.stdout
def test_doctor_human_config_error_is_actionable_and_redacted(monkeypatch, tmp_path):
cfg = tmp_path / "patient-private.yaml"
cfg.write_text("password: super-secret\nroots: [unterminated")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
result = runner.invoke(app, ["doctor", "--config", str(cfg)])
assert result.exit_code == 1
assert "config: error - configuration is invalid or unreadable" in result.stdout
assert str(tmp_path) not in result.stdout
assert "super-secret" not in result.stdout
assert result.stderr == ""
+163
View File
@@ -0,0 +1,163 @@
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)]
+70
View File
@@ -0,0 +1,70 @@
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,
)
class FakeDwhAdapter:
capabilities = DwhCapabilities()
def health(self) -> DwhHealth:
return DwhHealth(ok=True)
def introspect(self) -> PhysicalSchema:
raise NotImplementedError
def run_query(self, sql: str, *, limit: int) -> ExecResult:
raise NotImplementedError
def explain(self, sql: str) -> PlanSummary:
raise NotImplementedError
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
raise NotImplementedError
def distinct_values(self, table: str, column: str, *, limit: int) -> DistinctValues:
raise NotImplementedError
def test_fake_adapter_satisfies_runtime_protocol():
adapter = FakeDwhAdapter()
assert isinstance(adapter, DwhAdapter)
assert adapter.capabilities.explain is True
assert adapter.health().ok is True
def test_contract_types_are_public_and_capabilities_are_immutable():
capabilities = DwhCapabilities()
assert capabilities.introspection is True
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 DwhAdapter as PublicDwhAdapter
from tht.ports import DwhCapabilities as PublicDwhCapabilities
from tht.ports import DwhHealth as PublicDwhHealth
from tht.ports import DistinctValues as PublicDistinctValues
from tht.ports import UnsupportedCapability as PublicUnsupportedCapability
result = PublicDistinctValues(values=["a"], truncated=True)
assert result.values == ["a"]
assert result.truncated is True
assert PublicDwhAdapter is DwhAdapter
assert PublicDwhCapabilities is DwhCapabilities
assert PublicDwhHealth is DwhHealth
assert PublicUnsupportedCapability is UnsupportedCapability
+928
View File
@@ -0,0 +1,928 @@
import hashlib
import json
from pathlib import Path
from typer.testing import CliRunner
from tht.cli import app
from tht.jobs.dwh_pipeline import DwhPreprocessPipeline
from tht.jobs.dwh_pipeline import active_generation_dir, config_dwh_binding
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
from tht.jobs.dwh_pipeline import lease_dwh_snapshot
from tht.jobs.locking import _lock_name
FP = "sha256:" + hashlib.sha256(b"test").hexdigest()
def snapshot_config(tmp_path, workspace_id="demo"):
from types import SimpleNamespace
cfg = SimpleNamespace(
paths=SimpleNamespace(artifacts=tmp_path / "artifacts", indexes=tmp_path / "indexes"),
_workspace_id=workspace_id,
_config_source="test",
)
cfg.model_dump_json = lambda: "test"
return cfg
def test_dwh_and_evidence_jobs_have_distinct_lock_names():
assert _lock_name("demo", "dwh") != _lock_name("demo", "evidence")
def test_unowned_reads_fail_closed_without_creating_any_files(tmp_path):
import pytest
cfg = snapshot_config(tmp_path)
with pytest.raises(Exception, match="not initialized"):
resolve_dwh_snapshot(cfg)
with pytest.raises(Exception, match="not initialized"):
with lease_dwh_snapshot(cfg):
pass
assert not (tmp_path / ".tht-dwh").exists()
def test_writer_claim_allows_only_lock_and_empty_generations(tmp_path):
import pytest
for name, make_entry in (
("unexpected", lambda root: (root / "unexpected").write_text("x")),
("stale-temp", lambda root: (root / ".OWNER.json.stale.tmp").write_text("x")),
("unexpected-dir", lambda root: (root / "other").mkdir()),
):
root = tmp_path / name / ".tht-dwh"
root.mkdir(parents=True, mode=0o700)
make_entry(root)
calls = []
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=root.parent,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: calls.append("called"),
build_lsh=lambda physical, output: None,
)
with pytest.raises(Exception, match="unbound"):
pipeline.run()
assert calls == []
assert not (root / "OWNER.json").exists()
allowed = tmp_path / "allowed"
(allowed / ".tht-dwh" / "generations").mkdir(parents=True, mode=0o700)
report = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=allowed,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
).run()
assert report.status == "succeeded"
def test_writer_rejects_unbound_legacy_artifacts_before_building(tmp_path):
import pytest
legacy = tmp_path / "artifacts" / "mschema" / "physical.yaml"
legacy.parent.mkdir(parents=True)
legacy.write_text("legacy")
calls = []
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: calls.append("called"),
build_lsh=lambda physical, output: None,
current_physical=legacy,
)
with pytest.raises(Exception, match="legacy artifacts are unbound"):
pipeline.run()
assert calls == []
assert not (tmp_path / ".tht-dwh" / "OWNER.json").exists()
def test_writer_rejects_dangling_legacy_symlinks_before_claim_or_callback(tmp_path):
import pytest
legacy = tmp_path / "artifacts" / "mschema" / "physical.yaml"
legacy.parent.mkdir(parents=True)
legacy.symlink_to(tmp_path / "missing-catalog")
calls = []
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: calls.append("called"),
build_lsh=lambda physical, output: None,
current_physical=legacy,
)
with pytest.raises(Exception, match="legacy artifacts are unbound"):
pipeline.run()
assert calls == []
assert not (tmp_path / ".tht-dwh" / "OWNER.json").exists()
def test_owner_publication_remains_on_locked_root_when_path_is_swapped(
monkeypatch, tmp_path,
):
import pytest
import tht.jobs.dwh_pipeline as module
real_replace = module.os.replace
moved = tmp_path / "locked-root"
replacement = tmp_path / ".tht-dwh"
swapped = False
def swapping_replace(source, destination, *args, **kwargs):
nonlocal swapped
if destination == "OWNER.json" and kwargs.get("dst_dir_fd") is not None:
swapped = True
replacement.rename(moved)
replacement.mkdir(mode=0o700)
(moved / "generation.lock").rename(replacement / "generation.lock")
return real_replace(source, destination, *args, **kwargs)
monkeypatch.setattr(module.os, "replace", swapping_replace)
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: (_ for _ in ()).throw(AssertionError("callback called")),
build_lsh=lambda physical, output: None,
)
with pytest.raises(Exception):
pipeline.run()
assert swapped
assert (moved / "OWNER.json").is_file()
assert not (replacement / "OWNER.json").exists()
assert (replacement / "generation.lock").is_file()
def test_owner_requires_exact_read_only_owner_mode_and_active_requires_binding(tmp_path):
import pytest
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
pipeline.run()
cfg = snapshot_config(tmp_path)
binding = config_dwh_binding(cfg)
assert active_generation_dir(tmp_path, binding) is not None
with pytest.raises(Exception, match="different workspace configuration"):
active_generation_dir(tmp_path, {**binding, "workspace_id": "other"})
marker = tmp_path / ".tht-dwh" / "OWNER.json"
marker.chmod(0o440)
with pytest.raises(Exception, match="ownership marker"):
resolve_dwh_snapshot(cfg)
def test_selected_dwh_stages_run_in_declared_order(tmp_path):
calls = []
pipeline = DwhPreprocessPipeline(
workspace_id="demo",
workspace_root=tmp_path,
config_fingerprint=FP,
input_fingerprint=FP,
introspect=lambda output: (calls.append("introspect"), output.write_text("catalog")),
build_lsh=lambda physical, output: _write_lsh(calls, physical, output),
)
report = pipeline.run(("introspect", "lsh"))
assert report.status == "succeeded"
assert calls == ["introspect", "lsh"]
assert [stage.name for stage in report.stages] == ["introspect", "lsh"]
active = (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip()
published = tmp_path / ".tht-dwh" / "generations" / active
assert (published / "physical.yaml").read_text() == "catalog"
assert sorted(path.name for path in published.iterdir()) == [
"demo_lsh.pkl", "demo_meta.json", "demo_minhashes.pkl",
"generation-manifest.json", "physical.yaml",
]
def test_shared_root_rejects_other_workspace_before_builder_or_read(tmp_path):
calls = []
owner = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
published = owner.run()
contender = DwhPreprocessPipeline(
workspace_id="other", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: calls.append("introspect"),
build_lsh=lambda physical, output: calls.append("lsh"),
)
import pytest
with pytest.raises(Exception, match="different workspace configuration"):
contender.run()
with pytest.raises(Exception, match="different workspace configuration"):
resolve_dwh_snapshot(snapshot_config(tmp_path, "other"))
assert calls == []
assert resolve_dwh_snapshot(snapshot_config(tmp_path)).generation == published.run_id
def test_shared_root_mismatch_fails_without_deadlock_while_owner_reader_is_active(tmp_path):
import threading
owner = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
owner.run()
contender = DwhPreprocessPipeline(
workspace_id="other", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: (_ for _ in ()).throw(AssertionError("builder called")),
build_lsh=lambda physical, output: None,
)
finished = threading.Event()
errors = []
with lease_dwh_snapshot(snapshot_config(tmp_path)):
thread = threading.Thread(
target=lambda: (errors.append(_capture_error(contender.run)), finished.set())
)
thread.start()
assert not finished.wait(0.1)
thread.join(2)
assert finished.is_set()
assert "different workspace configuration" in str(errors[0])
def test_concurrent_brand_new_shared_root_has_one_atomic_owner_and_loser_never_builds(tmp_path):
import threading
calls = {"alpha": 0, "beta": 0}
results = []
barrier = threading.Barrier(2)
def run(workspace):
def introspect(output):
calls[workspace] += 1
output.write_text("catalog")
def build(physical, output):
calls[workspace] += 1
_write_lsh([], physical, output)
candidate = DwhPreprocessPipeline(
workspace_id=workspace, workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=introspect, build_lsh=build,
)
barrier.wait()
results.append((workspace, _capture_error(candidate.run)))
threads = [threading.Thread(target=run, args=(name,)) for name in ("alpha", "beta")]
for thread in threads:
thread.start()
for thread in threads:
thread.join(5)
assert all(not thread.is_alive() for thread in threads)
winner = next(name for name, result in results if not isinstance(result, Exception))
loser = next(name for name, result in results if isinstance(result, Exception))
assert calls[winner] == 2
assert calls[loser] == 0
def test_missing_active_with_generations_and_symlink_owner_marker_fail_closed(tmp_path):
import pytest
calls = []
owner = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
owner.run()
(tmp_path / ".tht-dwh" / "ACTIVE").unlink()
owner.introspect = lambda output: calls.append("called")
with pytest.raises(Exception, match="without a consistent ACTIVE"):
owner.run()
assert calls == []
other_root = tmp_path / "other"
marker_root = other_root / ".tht-dwh"
marker_root.mkdir(parents=True)
external = tmp_path / "external-owner"
external.write_text("foreign")
(marker_root / "OWNER.json").symlink_to(external)
contender = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=other_root,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: calls.append("symlink-called"),
build_lsh=lambda physical, output: None,
)
with pytest.raises(Exception, match="ownership marker"):
contender.run()
assert calls == []
def _capture_error(operation):
try:
return operation()
except Exception as error:
return error
def _write_lsh(calls, physical: Path, output: Path):
calls.append("lsh")
assert physical.read_text() == "catalog"
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json"):
(output / name).write_text(name)
def test_preprocess_dwh_json_is_pristine(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
class Report:
status = "succeeded"
def model_dump(self, mode=None):
return {"status": "succeeded", "run_id": "a" * 32, "stages": []}
seen = {}
def run(config, *, steps, resume):
seen.update(config=config, steps=steps, resume=resume)
return Report()
monkeypatch.setattr(command, "run_dwh_from_config", run)
response = CliRunner().invoke(
app,
[
"preprocess", "dwh", "--steps", "introspect,lsh", "--json",
"-c", str(tmp_path / "workspace.yaml"),
],
)
assert response.exit_code == 0, response.output
assert json.loads(response.output)["run_id"] == "a" * 32
assert seen["steps"] == ("introspect", "lsh")
def test_preprocess_dwh_rejects_unknown_or_duplicate_steps(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
called = False
def forbidden(*args, **kwargs):
nonlocal called
called = True
monkeypatch.setattr(command, "run_dwh_from_config", forbidden)
runner = CliRunner()
for value in ("introspect,unknown", "lsh,lsh", ""):
response = runner.invoke(
app,
["preprocess", "dwh", "--steps", value, "--json", "-c", str(tmp_path / "w.yaml")],
)
assert response.exit_code == 2
assert json.loads(response.output)["status"] == "failed"
assert called is False
def test_failed_multi_file_build_never_replaces_active_generation(tmp_path):
def catalog(output):
output.write_text("old-catalog")
first = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=catalog,
build_lsh=lambda physical, output: [
(output / name).write_text(name)
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
).run(("introspect", "lsh"))
assert first.status == "succeeded"
old_active = (tmp_path / ".tht-dwh" / "ACTIVE").read_text()
def partial_lsh(physical, output):
(output / "demo_lsh.pkl").write_text("new-but-partial")
raise RuntimeError("crash between LSH files")
failed = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("new-catalog"),
build_lsh=partial_lsh,
).run(("introspect", "lsh"))
assert failed.status == "failed"
assert (tmp_path / ".tht-dwh" / "ACTIVE").read_text() == old_active
resumed = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: (_ for _ in ()).throw(
AssertionError("completed introspection must not repeat")
),
build_lsh=lambda physical, output: [
(output / name).write_text("recovered")
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
).run(("introspect", "lsh"), resume_run_id=failed.run_id)
assert resumed.status == "succeeded"
def test_unsafe_lsh_filename_is_rejected(tmp_path):
import pytest
with pytest.raises(ValueError, match="flat safe"):
DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: None, build_lsh=lambda physical, output: None,
lsh_filenames=("../escape.pkl", "ok.pkl", "meta.json"),
)
def test_active_fsync_failure_restores_previous_pointer(monkeypatch, tmp_path):
import os
import tht.jobs.dwh_pipeline as module
def build(physical, output):
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json"):
(output / name).write_text(name)
first_pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("old"), build_lsh=build,
)
first = first_pipeline.run()
root = tmp_path / ".tht-dwh"
root_identity = (root.stat().st_dev, root.stat().st_ino)
original_fsync = module.os.fsync
failed_once = False
def fail_active_once(fd):
nonlocal failed_once
info = os.fstat(fd)
if (
(info.st_dev, info.st_ino) == root_identity
and "ACTIVE" in os.listdir(fd)
and not failed_once
):
failed_once = True
raise OSError("injected directory fsync failure")
original_fsync(fd)
second = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("new"), build_lsh=build,
)
monkeypatch.setattr(module.os, "fsync", fail_active_once)
failed = second.run()
assert failed.status == "failed"
assert (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip() == first.run_id
def test_snapshot_root_swap_after_lease_never_reads_replacement(monkeypatch, tmp_path):
import tht.jobs.dwh_pipeline as module
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("trusted"),
build_lsh=lambda physical, output: [
(output / name).write_text("trusted")
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
)
first = pipeline.run()
assert first.status == "succeeded"
root = tmp_path / ".tht-dwh"
moved = tmp_path / "moved-read-root"
replacement = root
real_read = module._read_owned_at
swapped = False
def swapping_read(directory_fd, name, *, readonly):
nonlocal swapped
if name == "ACTIVE" and not swapped:
swapped = True
replacement.rename(moved)
replacement.mkdir(mode=0o700)
(replacement / "sentinel").write_text("replacement-secret")
return real_read(directory_fd, name, readonly=readonly)
monkeypatch.setattr(module, "_read_owned_at", swapping_read)
try:
with lease_dwh_snapshot(snapshot_config(tmp_path)) as snapshot:
assert snapshot.physical.read_text() == "trusted"
except Exception as error:
assert "ACTIVE" in str(error) or "root" in str(error)
assert swapped
assert (replacement / "sentinel").read_text() == "replacement-secret"
def test_snapshot_copies_each_validated_artifact_once_without_reopen(monkeypatch, tmp_path):
import tht.jobs.dwh_pipeline as module
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("trusted"),
build_lsh=lambda physical, output: [
(output / name).write_text("trusted")
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
)
assert pipeline.run().status == "succeeded"
real_read = module._read_owned_at
reads = {}
def mutate_on_reopen(directory_fd, name, *, readonly):
reads[name] = reads.get(name, 0) + 1
if name.endswith(".pkl") and reads[name] > 1:
return b"MALICIOUS_PICKLE"
return real_read(directory_fd, name, readonly=readonly)
monkeypatch.setattr(module, "_read_owned_at", mutate_on_reopen)
with lease_dwh_snapshot(snapshot_config(tmp_path)) as snapshot:
assert (snapshot.lsh_dir / "demo_lsh.pkl").read_text() == "trusted"
assert "MALICIOUS" not in (snapshot.lsh_dir / "demo_lsh.pkl").read_text()
assert all(count == 1 for count in reads.values())
def test_reconcile_mismatch_closes_active_generation_fd(monkeypatch, tmp_path):
import os
from types import SimpleNamespace
import pytest
import tht.jobs.dwh_pipeline as module
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
report = pipeline.run()
run_dir = tmp_path / ".tht-jobs" / "dwh" / "runs" / report.run_id
real_active = module._active_generation_fd
def mismatched_active(root_fd, binding):
generation, generation_fd = real_active(root_fd, binding)
return "f" * 32, generation_fd
monkeypatch.setattr(module, "_active_generation_fd", mismatched_active)
source = SimpleNamespace(
run_id=report.run_id,
stages=(SimpleNamespace(
status="running", effect_state="intent", name="lsh",
artifact_files=("physical.yaml", "demo_lsh.pkl", "demo_minhashes.pkl",
"demo_meta.json"),
),),
)
before = len(os.listdir("/dev/fd"))
with pytest.raises(Exception, match="not ACTIVE"):
pipeline._reconcile_effects(source, run_dir)
assert len(os.listdir("/dev/fd")) == before
def test_pipeline_releases_materialized_snapshot_after_every_run(tmp_path):
import tht.jobs.dwh_pipeline as module
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
baseline = set(module._SNAPSHOT_DIRS)
for _ in range(3):
assert pipeline.run().status == "succeeded"
assert set(module._SNAPSHOT_DIRS) == baseline
assert pipeline._snapshot_holder is None
pipeline.introspect = lambda output: (_ for _ in ()).throw(RuntimeError("injected"))
assert pipeline.run().status == "failed"
assert set(module._SNAPSHOT_DIRS) == baseline
assert pipeline._snapshot_holder is None
def test_corrupt_resume_checkpoint_releases_materialized_snapshot(tmp_path):
import tht.jobs.dwh_pipeline as module
import pytest
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
assert pipeline.run().status == "succeeded"
baseline = set(module._SNAPSHOT_DIRS)
run_id = "e" * 32
run_dir = tmp_path / ".tht-jobs" / "dwh" / "runs" / run_id
run_dir.mkdir(parents=True)
(run_dir / "checkpoint.json").write_text("not-json")
with pytest.raises(Exception, match="checkpoint is invalid"):
pipeline.run(resume_run_id=run_id)
assert pipeline._snapshot_holder is None
assert set(module._SNAPSHOT_DIRS) == baseline
def test_job_spec_construction_failure_releases_materialized_snapshot(monkeypatch, tmp_path):
import tht.jobs.dwh_pipeline as module
import pytest
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _write_lsh([], physical, output),
)
assert pipeline.run().status == "succeeded"
baseline = set(module._SNAPSHOT_DIRS)
captured = []
real_materialize = module._materialize_generation_fd
def capture(*args, **kwargs):
holder, root = real_materialize(*args, **kwargs)
captured.append(holder)
return holder, root
monkeypatch.setattr(module, "_materialize_generation_fd", capture)
monkeypatch.setattr(
module, "JobSpec",
lambda **kwargs: (_ for _ in ()).throw(RuntimeError("job spec injected")),
)
with pytest.raises(RuntimeError, match="job spec injected"):
pipeline.run()
assert pipeline._snapshot_holder is None
assert set(module._SNAPSHOT_DIRS) == baseline
assert captured and all(not path.exists() for path in captured)
def test_publish_root_swap_after_lease_never_writes_replacement(monkeypatch, tmp_path):
import tht.jobs.dwh_pipeline as module
def make(content):
return DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text(content),
build_lsh=lambda physical, output: [
(output / name).write_text(content)
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
)
first = make("old").run()
assert first.status == "succeeded", first
root = tmp_path / ".tht-dwh"
moved = tmp_path / "moved-publish-root"
real_replace = module.os.replace
swapped = False
def swapping_replace(source, destination, *args, **kwargs):
nonlocal swapped
if destination == "ACTIVE" and kwargs.get("dst_dir_fd") is not None and not swapped:
swapped = True
root.rename(moved)
root.mkdir(mode=0o700)
(root / "sentinel").write_text("replacement-safe")
return real_replace(source, destination, *args, **kwargs)
monkeypatch.setattr(module.os, "replace", swapping_replace)
result = make("new").run()
assert result.status in {"succeeded", "failed"}
assert swapped
assert (root / "sentinel").read_text() == "replacement-safe"
moved_active = (moved / "ACTIVE").read_text().strip()
assert len(moved_active) == 32
assert (moved / "generations" / moved_active).is_dir()
def test_cleanup_root_swap_after_lease_never_deletes_replacement(monkeypatch, tmp_path):
import tht.jobs.dwh_pipeline as module
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("trusted"),
build_lsh=lambda physical, output: [
(output / name).write_text("trusted")
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
retain_generations=1,
)
pipeline.run()
root = tmp_path / ".tht-dwh"
moved = tmp_path / "moved-cleanup-root"
real_open = module.os.open
swapped = False
def swapping_open(path, flags, *args, **kwargs):
nonlocal swapped
if path == "generations" and kwargs.get("dir_fd") is not None and not swapped:
swapped = True
root.rename(moved)
root.mkdir(mode=0o700)
(root / "sentinel").write_text("replacement-safe")
return real_open(path, flags, *args, **kwargs)
monkeypatch.setattr(module.os, "open", swapping_open)
pipeline._cleanup_generations()
assert swapped
assert (root / "sentinel").read_text() == "replacement-safe"
def test_snapshot_stays_on_one_generation_across_publish(tmp_path):
def pipeline(content):
return DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text(content),
build_lsh=lambda physical, output: [
(output / name).write_text(content)
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
)
first = pipeline("old").run()
cfg = snapshot_config(tmp_path)
snapshot = resolve_dwh_snapshot(cfg)
pipeline("new").run()
assert snapshot.generation == first.run_id
assert snapshot.physical.read_text() == "old"
assert (snapshot.lsh_dir / "demo_meta.json").read_text() == "old"
def test_generation_retention_keeps_active_and_one_rollback(tmp_path):
run_ids = []
for index in range(5):
report = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output, i=index: output.write_text(str(i)),
build_lsh=lambda physical, output, i=index: [
(output / name).write_text(str(i))
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
retain_generations=2,
).run()
run_ids.append(report.run_id)
remaining = {path.name for path in (tmp_path / ".tht-dwh" / "generations").iterdir()}
assert remaining == set(run_ids[-2:])
def test_corrupt_newer_directory_does_not_consume_rollback_slot(tmp_path):
run_ids = []
pipeline = None
for index in range(3):
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output, i=index: output.write_text(str(i)),
build_lsh=lambda physical, output, i=index: [
(output / name).write_text(str(i))
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
retain_generations=3,
)
run_ids.append(pipeline.run().run_id)
generations = tmp_path / ".tht-dwh" / "generations"
corrupt = generations / ("f" * 32)
corrupt.mkdir(mode=0o700)
(corrupt / "junk").write_text("not a published generation")
pipeline.retain_generations = 2
pipeline._cleanup_generations()
assert (generations / run_ids[-1]).is_dir()
assert (generations / run_ids[-2]).is_dir()
assert not (generations / run_ids[0]).exists()
assert corrupt.is_dir()
def test_retention_n_counts_active_plus_n_minus_one_rollbacks_even_if_active_is_old(tmp_path):
import os
run_ids = []
pipeline = None
for index in range(3):
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output, i=index: output.write_text(str(i)),
build_lsh=lambda physical, output, i=index: [
(output / name).write_text(str(i))
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
], retain_generations=3,
)
run_ids.append(pipeline.run().run_id)
generations = tmp_path / ".tht-dwh" / "generations"
os.utime(generations / run_ids[-1], ns=(1, 1))
pipeline.retain_generations = 2
pipeline._cleanup_generations()
remaining = {path.name for path in generations.iterdir() if path.is_dir()}
assert remaining == {run_ids[-1], run_ids[-2]}
def test_retention_candidate_swap_to_symlink_is_never_followed(monkeypatch, tmp_path):
import tht.jobs.dwh_pipeline as module
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("active"),
build_lsh=lambda physical, output: [
(output / name).write_text("active")
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
], retain_generations=1,
)
pipeline.run()
generations = tmp_path / ".tht-dwh" / "generations"
candidate_name = "e" * 32
candidate = generations / candidate_name
candidate.mkdir(mode=0o700)
external = tmp_path / "external-crafted"
external.mkdir()
sentinel = external / "sentinel"
sentinel.write_text("must-not-read-or-mutate")
real_open = module.os.open
swapped = False
def swapping_open(path, flags, *args, **kwargs):
nonlocal swapped
if path == candidate_name and kwargs.get("dir_fd") is not None and not swapped:
swapped = True
candidate.rmdir()
candidate.symlink_to(external, target_is_directory=True)
return real_open(path, flags, *args, **kwargs)
monkeypatch.setattr(module.os, "open", swapping_open)
pipeline._cleanup_generations()
assert swapped
assert sentinel.read_text() == "must-not-read-or-mutate"
assert candidate.is_symlink()
def test_reader_lease_blocks_retain_one_publisher_until_file_reads_finish(tmp_path):
import threading
import time
def make(content, retain=1):
return DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text(content),
build_lsh=lambda physical, output: [
(output / name).write_text(content)
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
], retain_generations=retain,
)
first = make("old").run()
cfg = snapshot_config(tmp_path)
completed = threading.Event()
with lease_dwh_snapshot(cfg) as snapshot:
thread = threading.Thread(target=lambda: (make("new").run(), completed.set()))
thread.start()
time.sleep(0.05)
assert not completed.is_set()
assert snapshot.physical.read_text() == "old"
assert snapshot.generation == first.run_id
thread.join(timeout=2)
assert completed.is_set()
assert not (tmp_path / ".tht-dwh" / "generations" / first.run_id).exists()
def test_cleanup_never_follows_top_level_or_child_symlinks(tmp_path):
external = tmp_path / "external"
external.mkdir()
victim = external / "victim"
victim.write_text("safe")
def make(content, retain=1):
return DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text(content),
build_lsh=lambda physical, output: [
(output / name).write_text(content)
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
], retain_generations=retain,
)
first = make("one", retain=2).run()
generations = tmp_path / ".tht-dwh" / "generations"
(generations / ("a" * 32)).symlink_to(external, target_is_directory=True)
make("two", retain=2).run()
old = generations / first.run_id
old.chmod(0o700)
(old / "hostile-link").symlink_to(victim)
make("three").run()
assert victim.read_text() == "safe"
assert victim.stat().st_mode & 0o200
@@ -0,0 +1,239 @@
from datetime import UTC, datetime, timedelta, timezone
import pytest
from pydantic import ValidationError
from tht.ports.evidence import (
AcquiredDocument,
EvidenceSource,
EvidenceSourceError,
EvidenceSourceErrorCategory,
SourceObject,
)
class StubSource:
def discover(self):
return iter(
[
SourceObject(
source_id="source:handbook",
uri="https://host/handbook.md",
fingerprint="sha256:abc",
)
]
)
def acquire(self, item: SourceObject) -> AcquiredDocument:
return AcquiredDocument(
source=item,
content=b"# Handbook",
acquired_at=datetime(2026, 7, 12, tzinfo=UTC),
media_type="text/markdown",
)
def test_runtime_checkable_source_protocol():
source = StubSource()
assert isinstance(source, EvidenceSource)
assert source.acquire(next(source.discover())).content == b"# Handbook"
def test_source_objects_are_frozen_and_metadata_defaults_are_independent():
first = SourceObject(source_id="source:a", uri="file:///a", fingerprint="sha256:a")
second = SourceObject(source_id="source:b", uri="file:///b", fingerprint="sha256:b")
with pytest.raises(ValidationError):
first.uri = "file:///changed" # type: ignore[misc]
with pytest.raises(TypeError):
first.metadata["owner"] = "team-a"
assert second.metadata == {}
@pytest.mark.parametrize(
"key",
[
"password",
"PassWd",
"api_key",
"x-api-key",
"accessToken",
"refresh.token",
"client secret",
"privateKey",
"session_cookie",
"Authorization",
],
)
def test_source_metadata_rejects_credential_specific_keys(key):
with pytest.raises(ValidationError, match="credential-like"):
SourceObject(
source_id="source:a",
uri="https://host/a",
fingerprint="etag:abc",
metadata={"nested": [{key: "secret"}]},
)
def test_source_metadata_allows_benign_generic_token_and_secret_labels():
source = SourceObject(
source_id="source:a",
uri="https://host/a",
fingerprint="etag:abc",
metadata={"token": "word count token", "secret": False},
)
assert source.metadata["token"] == "word count token"
def test_nested_metadata_is_recursively_immutable_and_serializes_as_json():
source = SourceObject(
source_id="source:a",
uri="https://host/a",
fingerprint="etag:abc",
metadata={"nested": {"items": [1, {"ok": True}]}},
)
with pytest.raises(TypeError):
source.metadata["nested"]["items"][1]["ok"] = False
assert '"items":[1,{"ok":true}]' in source.model_dump_json()
@pytest.mark.parametrize(
"uri",
[
"https://user:pass@host/a",
"https://host/a?api_key=secret",
"https://host/a?accessToken=secret",
],
)
def test_source_uri_rejects_embedded_credentials(uri):
with pytest.raises(ValidationError, match="credentials"):
SourceObject(source_id="source:a", uri=uri, fingerprint="etag:abc")
def test_source_metadata_must_be_json_safe():
with pytest.raises(ValidationError):
SourceObject(
source_id="source:a",
uri="file:///a",
fingerprint="sha256:a",
metadata={"path": object()},
)
def test_source_identity_and_fingerprint_must_be_namespaced():
with pytest.raises(ValidationError, match="namespaced"):
SourceObject(source_id="plain", uri="file:///a", fingerprint="sha256:a")
with pytest.raises(ValidationError, match="namespaced"):
SourceObject(source_id="source:a", uri="file:///a", fingerprint="plain")
def test_acquired_document_does_not_accept_credentials_as_extra_fields():
item = SourceObject(source_id="source:a", uri="https://host/a", fingerprint="etag:abc")
with pytest.raises(ValidationError):
AcquiredDocument(source=item, content=b"a", api_key="secret")
def test_acquired_binary_content_has_explicit_json_round_trip():
item = SourceObject(source_id="source:a", uri="https://host/a", fingerprint="etag:abc")
acquired = AcquiredDocument(source=item, content=b"\x00\xffbinary\x80")
payload = acquired.model_dump_json()
restored = AcquiredDocument.model_validate_json(payload)
assert restored.content == acquired.content
assert "binary" not in payload
def test_datetimes_must_be_aware_and_are_normalized_to_utc():
with pytest.raises(ValidationError, match="timezone-aware"):
SourceObject(
source_id="source:a",
uri="https://host/a",
fingerprint="etag:abc",
modified_at=datetime(2026, 7, 12),
)
source = SourceObject(
source_id="source:a",
uri="https://host/a",
fingerprint="etag:abc",
modified_at=datetime(2026, 7, 12, 4, tzinfo=timezone(timedelta(hours=2))),
)
assert source.modified_at.tzinfo is UTC
assert source.modified_at.hour == 2
acquired = AcquiredDocument(
source=source,
content=b"a",
acquired_at=datetime(2026, 7, 12, 2, tzinfo=UTC) + timedelta(hours=0),
)
assert acquired.acquired_at.utcoffset() == timedelta(0)
def test_source_errors_are_typed_retryable_and_safe():
transient = EvidenceSourceError(
"password=hunter2 at https://user:secret@host",
category=EvidenceSourceErrorCategory.TRANSIENT,
details={"status": 503},
)
permanent = EvidenceSourceError(
"unsupported media type",
category=EvidenceSourceErrorCategory.PERMANENT,
)
assert transient.retryable is True
assert permanent.retryable is False
assert transient.details["status"] == 503
assert str(transient) == "evidence source operation failed"
assert transient.args == ("evidence source operation failed",)
assert "hunter2" not in repr(transient)
with pytest.raises(AttributeError):
transient.category = EvidenceSourceErrorCategory.PERMANENT
with pytest.raises(AttributeError):
transient.args = ("leak",)
with pytest.raises(AttributeError):
transient.details = {"unsafe": True}
assert "hunter2" not in repr(transient.__dict__)
with pytest.raises(TypeError):
transient.details["status"] = 200
with pytest.raises(ValueError, match="credential-like"):
EvidenceSourceError(
"bad",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"apiKey": "must-not-leak"},
)
with pytest.raises(ValidationError):
EvidenceSourceError(
"bad",
category=EvidenceSourceErrorCategory.PERMANENT,
details={"not_json": object()},
)
def test_source_error_preserves_original_only_through_exception_chaining():
cause = RuntimeError("transport diagnostic with password=hunter2")
error = EvidenceSourceError(
"ignored unsafe diagnostic",
category=EvidenceSourceErrorCategory.TRANSIENT,
)
try:
raise error from cause
except EvidenceSourceError as caught:
assert caught.__cause__ is cause
assert "hunter2" not in str(caught)
assert "hunter2" not in caught.args
def test_model_copy_revalidates_source_and_acquired_records():
source = SourceObject(source_id="source:a", uri="file:///a", fingerprint="sha256:a")
acquired = AcquiredDocument(source=source, content=b"a")
with pytest.raises(ValidationError, match="namespaced"):
source.model_copy(update={"source_id": "invalid"})
with pytest.raises(ValidationError, match="timezone-aware"):
acquired.model_copy(update={"acquired_at": datetime(2026, 7, 12)})
@@ -0,0 +1,113 @@
import os
import pytest
from tht.adapters.evidence import FilesystemEvidenceSource
from tht.ports.evidence import EvidenceSourceError
def test_filesystem_discovery_is_stable_and_acquisition_is_bounded(tmp_path):
(tmp_path / "z.md").write_text("z")
(tmp_path / "nested").mkdir()
(tmp_path / "nested" / "a.md").write_text("alpha")
source = FilesystemEvidenceSource(tmp_path, max_bytes=5)
first = list(source.discover())
assert [item.uri for item in first] == sorted(item.uri for item in first)
assert all(item.source_id.startswith("filesystem:") for item in first)
assert all(item.fingerprint.startswith("sha256:") for item in first)
assert source.acquire(first[0]).content in {b"alpha", b"z"}
(tmp_path / "large.md").write_bytes(b"123456")
with pytest.raises(EvidenceSourceError) as caught:
list(source.discover())
assert not caught.value.retryable
assert "large.md" not in str(caught.value)
def test_filesystem_rejects_symlink_escape(tmp_path):
root = tmp_path / "root"
root.mkdir()
outside = tmp_path / "secret.md"
outside.write_text("secret")
(root / "escape.md").symlink_to(outside)
with pytest.raises(EvidenceSourceError) as caught:
list(FilesystemEvidenceSource(root).discover())
assert not caught.value.retryable
assert str(outside) not in str(caught.value)
def test_filesystem_acquire_rejects_object_from_another_source(tmp_path):
left = tmp_path / "left"
right = tmp_path / "right"
left.mkdir()
right.mkdir()
(left / "doc.md").write_text("left")
(right / "doc.md").write_text("right")
item = next(iter(FilesystemEvidenceSource(left).discover()))
with pytest.raises(EvidenceSourceError):
FilesystemEvidenceSource(right).acquire(item)
def test_filesystem_acquire_rejects_content_changed_since_discovery(tmp_path):
path = tmp_path / "doc.md"
path.write_text("first")
source = FilesystemEvidenceSource(tmp_path)
item = next(iter(source.discover()))
path.write_text("second")
with pytest.raises(EvidenceSourceError) as caught:
source.acquire(item)
assert not caught.value.retryable
def test_filesystem_open_is_safe_when_file_is_swapped_for_symlink(tmp_path, monkeypatch):
root = tmp_path / "root"
root.mkdir()
path = root / "doc.md"
path.write_text("safe")
outside = tmp_path / "outside.md"
outside.write_text("secret")
source = FilesystemEvidenceSource(root)
real_open = os.open
swapped = False
def racing_open(name, flags, *args, **kwargs):
nonlocal swapped
if name == "doc.md" and not swapped:
swapped = True
path.unlink()
path.symlink_to(outside)
return real_open(name, flags, *args, **kwargs)
monkeypatch.setattr(os, "open", racing_open)
with pytest.raises(EvidenceSourceError):
list(source.discover())
def test_filesystem_open_is_safe_when_ancestor_is_swapped_for_symlink(tmp_path, monkeypatch):
root = tmp_path / "root"
nested = root / "nested"
nested.mkdir(parents=True)
(nested / "doc.md").write_text("safe")
outside = tmp_path / "outside"
outside.mkdir()
(outside / "doc.md").write_text("secret")
source = FilesystemEvidenceSource(root)
real_open = os.open
swapped = False
def racing_open(name, flags, *args, **kwargs):
nonlocal swapped
if name == "nested" and not swapped and kwargs.get("dir_fd") is not None:
swapped = True
(nested / "doc.md").unlink()
nested.rmdir()
nested.symlink_to(outside, target_is_directory=True)
return real_open(name, flags, *args, **kwargs)
monkeypatch.setattr(os, "open", racing_open)
with pytest.raises(EvidenceSourceError):
list(source.discover())
+331
View File
@@ -0,0 +1,331 @@
import threading
import socket
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
import pytest
from tht.adapters.evidence import HttpManifestEvidenceSource
from tht.ports.evidence import EvidenceSourceError
class Handler(BaseHTTPRequestHandler):
etag_requests = 0
etag_body_responses = 0
redirect_target = "/redirected-v1"
redirect_request_validators = []
final_request_validators = []
def do_GET(self):
if self.path.startswith("/etag"):
type(self).etag_requests += 1
if self.headers.get("If-None-Match") == '"abc"':
self.send_response(304)
self.end_headers()
return
self.send_response(200)
self.send_header("ETag", '"abc"')
self.send_header("Content-Type", "text/markdown")
self.end_headers()
type(self).etag_body_responses += 1
self.wfile.write(b"hello")
elif self.path == "/large":
self.send_response(200)
self.send_header("Content-Length", "20")
self.end_headers()
self.wfile.write(b"x" * 20)
elif self.path == "/busy":
self.send_response(503)
self.end_headers()
elif self.path == "/missing":
self.send_response(404)
self.end_headers()
elif self.path == "/redirect-private":
self.send_response(302)
self.send_header("Location", f"http://127.0.0.1:{self.server.server_port}/etag")
self.end_headers()
elif self.path == "/redirect-userinfo":
self.send_response(302)
self.send_header(
"Location", f"http://user:password@127.0.0.1:{self.server.server_port}/etag"
)
self.end_headers()
elif self.path == "/stable-redirect":
type(self).redirect_request_validators.append(self.headers.get("If-None-Match"))
self.send_response(302)
self.send_header("Location", type(self).redirect_target)
self.end_headers()
elif self.path in {"/redirected-v1", "/redirected-v2"}:
type(self).final_request_validators.append(
(self.path, self.headers.get("If-None-Match"))
)
etag = '"v1"' if self.path.endswith("v1") else '"v2"'
if self.headers.get("If-None-Match") == etag:
self.send_response(304)
self.end_headers()
return
self.send_response(200)
self.send_header("ETag", etag)
self.end_headers()
self.wfile.write(self.path.encode())
else:
self.send_response(200)
self.send_header("Last-Modified", "Wed, 21 Oct 2015 07:28:00 GMT")
self.end_headers()
self.wfile.write(b"fallback")
def log_message(self, format, *args):
pass
@pytest.fixture
def server_url():
Handler.etag_requests = 0
Handler.etag_body_responses = 0
Handler.redirect_target = "/redirected-v1"
Handler.redirect_request_validators = []
Handler.final_request_validators = []
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_port}"
finally:
server.shutdown()
thread.join()
def test_http_uses_etag_and_strips_query_from_provenance(server_url):
source = HttpManifestEvidenceSource(
[f"{server_url}/etag?token=secret"], allow_private_hosts=True
)
item = next(iter(source.discover()))
assert item.fingerprint.startswith("etag:")
assert item.fingerprint != "etag:abc"
assert item.uri == f"{server_url}/etag"
assert "secret" not in item.model_dump_json()
assert source.acquire(item).content == b"hello"
def test_http_uses_last_modified_then_content_hash(server_url):
modified = next(iter(HttpManifestEvidenceSource(
[f"{server_url}/modified"], allow_private_hosts=True
).discover()))
assert modified.fingerprint.startswith("last-modified:")
class NoValidators(Handler):
def do_GET(self):
self.send_response(200)
self.end_headers()
self.wfile.write(b"content")
server = ThreadingHTTPServer(("127.0.0.1", 0), NoValidators)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
item = next(iter(HttpManifestEvidenceSource(
[f"http://127.0.0.1:{server.server_port}/doc"], allow_private_hosts=True
).discover()))
assert item.fingerprint.startswith("sha256:")
finally:
server.shutdown()
thread.join()
@pytest.mark.parametrize("path,retryable", [("/busy", True), ("/missing", False)])
def test_http_classifies_status_errors(server_url, path, retryable):
with pytest.raises(EvidenceSourceError) as caught:
list(HttpManifestEvidenceSource(
[server_url + path], allow_private_hosts=True
).discover())
assert caught.value.retryable is retryable
assert server_url not in str(caught.value)
def test_http_rejects_oversize_and_private_redirect(server_url):
with pytest.raises(EvidenceSourceError) as large:
list(HttpManifestEvidenceSource(
[server_url + "/large"], max_bytes=10, allow_private_hosts=True
).discover())
assert not large.value.retryable
with pytest.raises(EvidenceSourceError) as redirect:
list(HttpManifestEvidenceSource([server_url + "/redirect-private"]).discover())
assert not redirect.value.retryable
def test_http_rejects_unsupported_manifest_scheme():
with pytest.raises(ValueError, match="http"):
HttpManifestEvidenceSource(["file:///tmp/secret"])
def test_http_conditional_discovery_reuses_cached_verified_bytes(server_url):
source = HttpManifestEvidenceSource([server_url + "/etag"], allow_private_hosts=True)
first = next(iter(source.discover()))
second = next(iter(source.discover()))
assert second == first
assert source.acquire(second).content == b"hello"
assert Handler.etag_requests == 3
assert Handler.etag_body_responses == 1
def test_http_rejects_mixed_public_private_dns_answers(monkeypatch):
monkeypatch.setattr(socket, "getaddrinfo", lambda *args, **kwargs: [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 80)),
(socket.AF_INET6, socket.SOCK_STREAM, 6, "", ("::1", 80, 0, 0)),
])
with pytest.raises(EvidenceSourceError) as caught:
list(HttpManifestEvidenceSource(["http://example.test/doc"]).discover())
assert not caught.value.retryable
def test_http_rejects_userinfo_redirect(server_url):
with pytest.raises(EvidenceSourceError) as caught:
list(HttpManifestEvidenceSource(
[server_url + "/redirect-userinfo"], allow_private_hosts=True
).discover())
assert not caught.value.retryable
class FakeSocket:
def __init__(self, address):
self.address = address
def getpeername(self):
return (self.address, 443)
class FakeResponse:
status_code = 200
headers = {}
is_redirect = False
def __init__(self, *, peer="127.0.0.1", stream_error=None, location=None):
connection = type("Connection", (), {"sock": FakeSocket(peer)})()
self.raw = type("Raw", (), {"_connection": connection})()
self.stream_error = stream_error
self.closed = False
if location:
self.is_redirect = True
self.status_code = 302
self.headers = {"Location": location}
else:
self.is_redirect = False
self.status_code = 200
self.headers = {}
def iter_content(self, chunk_size):
if self.stream_error:
raise self.stream_error
yield b"ok"
def close(self):
self.closed = True
class FakeSession:
def __init__(self, response):
self.response = response
def get(self, *args, **kwargs):
return self.response
def test_http_rejects_public_to_private_rebind(monkeypatch):
monkeypatch.setattr(socket, "getaddrinfo", lambda *args, **kwargs: [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)),
])
source = HttpManifestEvidenceSource(["https://example.test/doc"])
response = FakeResponse(peer="127.0.0.1")
source._session = FakeSession(response)
with pytest.raises(EvidenceSourceError):
list(source.discover())
assert response.closed
def test_http_rejects_public_redirect_to_private_destination(monkeypatch):
monkeypatch.setattr(socket, "getaddrinfo", lambda host, *args, **kwargs: [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", (
"93.184.216.34" if host == "example.test" else "127.0.0.1", 443
)),
])
source = HttpManifestEvidenceSource(["https://example.test/doc"])
response = FakeResponse(
peer="93.184.216.34", location="https://private.test/secret"
)
source._session = FakeSession(response)
with pytest.raises(EvidenceSourceError) as caught:
list(source.discover())
assert not caught.value.retryable
assert response.closed
def test_http_closes_response_when_streaming_fails():
source = HttpManifestEvidenceSource(
["https://example.test/doc"], allow_private_hosts=True
)
response = FakeResponse(stream_error=socket.timeout("read timed out"))
source._session = FakeSession(response)
with pytest.raises(EvidenceSourceError):
list(source.discover())
assert response.closed
def test_http_binds_validators_to_exact_final_redirect_url(server_url):
source = HttpManifestEvidenceSource(
[server_url + "/stable-redirect"], allow_private_hosts=True
)
first = next(iter(source.discover()))
second = next(iter(source.discover()))
assert first == second
assert Handler.redirect_request_validators == [None, None]
assert Handler.final_request_validators == [
("/redirected-v1", None),
("/redirected-v1", '"v1"'),
]
def test_http_redirect_path_change_fetches_and_replaces_body(server_url):
source = HttpManifestEvidenceSource(
[server_url + "/stable-redirect"], allow_private_hosts=True
)
first = next(iter(source.discover()))
assert source.acquire(first).content == b"/redirected-v1"
Handler.redirect_target = "/redirected-v2"
with pytest.raises(EvidenceSourceError):
source.acquire(first)
current = next(iter(source.discover()))
assert source.acquire(current).content == b"/redirected-v2"
assert ("/redirected-v2", None) in Handler.final_request_validators
def test_http_rejects_unsolicited_304_without_bound_validator(monkeypatch):
monkeypatch.setattr(socket, "getaddrinfo", lambda *args, **kwargs: [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)),
])
source = HttpManifestEvidenceSource(["https://example.test/doc"])
response = FakeResponse(peer="93.184.216.34")
response.status_code = 304
source._session = FakeSession(response)
with pytest.raises(EvidenceSourceError) as caught:
list(source.discover())
assert not caught.value.retryable
assert response.closed
def test_http_rejects_cross_origin_304_for_cached_provenance(server_url):
source = HttpManifestEvidenceSource([server_url + "/etag"], allow_private_hosts=True)
item = next(iter(source.discover()))
response = FakeResponse()
response.status_code = 304
source._session = FakeSession(response)
with pytest.raises(EvidenceSourceError) as caught:
source._download("https://other.example/doc", item.uri)
assert not caught.value.retryable
assert response.closed
+92
View File
@@ -0,0 +1,92 @@
import multiprocessing
import os
from pathlib import Path
import pytest
from tht.jobs.locking import JobAlreadyRunningError, WorkspaceJobLock
def _hold_lock(root: str, ready, release):
with WorkspaceJobLock(Path(root), "demo", "evidence"):
ready.set()
release.wait(10)
def _crash_with_lock(root: str, ready):
lock = WorkspaceJobLock(Path(root), "demo", "evidence")
lock.acquire()
ready.set()
raise SystemExit(7)
def test_same_workspace_and_job_are_exclusive_across_processes(tmp_path):
context = multiprocessing.get_context("spawn")
ready = context.Event()
release = context.Event()
process = context.Process(target=_hold_lock, args=(str(tmp_path), ready, release))
process.start()
assert ready.wait(10)
try:
with pytest.raises(JobAlreadyRunningError):
WorkspaceJobLock(tmp_path, "demo", "evidence").acquire()
finally:
release.set()
process.join(10)
assert process.exitcode == 0
def test_evidence_and_dwh_jobs_have_distinct_locks(tmp_path):
with WorkspaceJobLock(tmp_path, "demo", "evidence"):
with WorkspaceJobLock(tmp_path, "demo", "dwh"):
pass
def test_lock_keys_cannot_escape_lock_directory(tmp_path):
with pytest.raises(ValueError, match="filesystem-safe"):
WorkspaceJobLock(tmp_path, "demo", "../evidence")
def test_preexisting_lock_symlink_is_rejected(tmp_path):
lock = WorkspaceJobLock(tmp_path, "demo", "evidence")
lock.path.parent.mkdir(parents=True)
target = tmp_path / "target"
target.write_text("do not modify")
lock.path.symlink_to(target)
with pytest.raises(OSError):
lock.acquire()
assert target.read_text() == "do not modify"
def test_preexisting_locks_directory_symlink_is_rejected(tmp_path):
jobs = tmp_path / ".tht-jobs"
jobs.mkdir()
outside = tmp_path / "outside"
outside.mkdir()
(jobs / ".locks").symlink_to(outside, target_is_directory=True)
with pytest.raises(OSError):
WorkspaceJobLock(tmp_path, "demo", "evidence").acquire()
assert list(outside.iterdir()) == []
def test_lock_file_is_owner_only_regular_single_link(tmp_path):
with WorkspaceJobLock(tmp_path, "demo", "evidence") as lock:
stat = os.stat(lock.path, follow_symlinks=False)
assert stat.st_uid == os.getuid()
assert stat.st_nlink == 1
assert stat.st_mode & 0o777 == 0o600
def test_lock_is_recoverable_after_process_crash_without_stale_deletion(tmp_path):
context = multiprocessing.get_context("spawn")
ready = context.Event()
process = context.Process(target=_crash_with_lock, args=(str(tmp_path), ready))
process.start()
assert ready.wait(10)
process.join(10)
assert process.exitcode == 7
lock_path = WorkspaceJobLock(tmp_path, "demo", "evidence").path
assert lock_path.exists()
with WorkspaceJobLock(tmp_path, "demo", "evidence"):
assert lock_path.exists()
+400
View File
@@ -0,0 +1,400 @@
import json
import os
import pytest
from pydantic import ValidationError
from tht.jobs.models import JobSpec
from tht.jobs.runner import CorruptCheckpointError, StageArtifacts, run_job
import tht.jobs.runner as runner_module
def _spec(tmp_path, **updates):
values = {
"workspace_id": "demo",
"job_type": "evidence",
"workspace_root": tmp_path,
"spec_version": "jobs-v1",
"pipeline_version": "evidence-v1",
"config_fingerprint": "sha256:" + "1" * 64,
"input_fingerprint": "sha256:" + "2" * 64,
"stage_ids": ("stage",),
}
values.update(updates)
return JobSpec(**values)
def test_job_models_are_immutable(tmp_path):
spec = _spec(tmp_path)
with pytest.raises(ValidationError, match="Instance is frozen"):
spec.job_type = "dwh"
resumed = spec.with_resume("a" * 32)
assert resumed.workspace_root == tmp_path
assert resumed.resume_run_id == "a" * 32
def test_failed_stage_is_resumable_and_skips_completed_stage(tmp_path):
calls = []
def discover(context):
calls.append(("discover", context.dry_run))
def acquire(_context):
calls.append(("acquire", False))
raise RuntimeError("source /customer/alice token=secret unavailable")
first = run_job(_spec(tmp_path, stage_ids=("discover", "acquire")), [discover, acquire])
assert first.status == "failed"
assert [stage.status for stage in first.stages] == ["succeeded", "failed"]
assert first.stages[1].error.model_dump() == {
"category": "internal",
"code": "stage_exception",
"message": "stage execution failed",
}
def acquire(_context):
calls.append(("recovered", False))
second = run_job(
_spec(tmp_path, resume_run_id=first.run_id, stage_ids=("discover", "acquire")),
[discover, acquire],
)
assert second.resumed_from == first.run_id
assert second.status == "succeeded"
assert calls == [("discover", False), ("acquire", False), ("recovered", False)]
def test_resume_carries_successful_stage_artifacts_into_new_run(tmp_path):
def discover(context):
artifacts = context.run_dir / "artifacts"
artifacts.mkdir()
(artifacts / "discovery.json").write_text('{"source":"one"}')
return StageArtifacts(("discovery.json",))
first = run_job(
_spec(tmp_path, stage_ids=("discover", "acquire")),
[discover, lambda _context: (_ for _ in ()).throw(RuntimeError("crash"))],
)
def acquire(context):
assert (context.run_dir / "artifacts" / "discovery.json").read_text() == '{"source":"one"}'
resumed = run_job(
_spec(tmp_path, resume_run_id=first.run_id, stage_ids=("discover", "acquire")),
[discover, acquire],
)
assert resumed.status == "succeeded"
def test_crash_after_stage_effect_resumes_without_repeating_stage(tmp_path):
calls = []
def stage(context):
calls.append("stage")
artifacts = context.run_dir / "artifacts"
artifacts.mkdir()
(artifacts / "effect.json").write_text("ok")
return StageArtifacts(("effect.json",))
class Crash(BaseException):
pass
with pytest.raises(Crash):
run_job(
_spec(tmp_path), [stage],
after_stage_return=lambda *_: (_ for _ in ()).throw(Crash()),
)
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
crashed_run = next(runs.iterdir()).name
resumed = run_job(_spec(tmp_path).with_resume(crashed_run), [stage])
assert resumed.status == "succeeded"
assert calls == ["stage"]
def test_resume_rejects_tampered_successful_stage_artifact(tmp_path):
def stage(context):
artifacts = context.run_dir / "artifacts"
artifacts.mkdir()
(artifacts / "effect.json").write_text("ok")
return StageArtifacts(("effect.json",))
report = run_job(_spec(tmp_path), [stage])
path = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id / "artifacts" / "effect.json"
path.write_text("tampered")
with pytest.raises(CorruptCheckpointError, match="artifact"):
run_job(_spec(tmp_path).with_resume(report.run_id), [stage])
def test_resume_rejects_extra_symlink_before_any_stage(tmp_path):
report = run_job(_spec(tmp_path), [lambda _context: StageArtifacts()])
artifacts = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id / "artifacts"
(artifacts / "unsafe").symlink_to(tmp_path)
called = False
def forbidden(_context):
nonlocal called
called = True
with pytest.raises(CorruptCheckpointError, match="artifact"):
run_job(_spec(tmp_path).with_resume(report.run_id), [forbidden])
assert called is False
def test_nonexistent_well_formed_resume_run_id_is_rejected(tmp_path):
with pytest.raises(CorruptCheckpointError, match="checkpoint"):
run_job(_spec(tmp_path).with_resume("a" * 32), [lambda _context: None])
@pytest.mark.parametrize("tamper", ["artifact_and_manifest", "spec", "producer"])
def test_resume_rejects_manifest_root_or_binding_tamper(tmp_path, tamper):
def stage(context):
artifacts = context.run_dir / "artifacts"
artifacts.mkdir()
(artifacts / "effect.json").write_text("ok")
return StageArtifacts(("effect.json",))
report = run_job(_spec(tmp_path), [stage])
artifacts = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id / "artifacts"
manifest_path = artifacts / "artifact-manifest.json"
manifest = json.loads(manifest_path.read_text())
if tamper == "artifact_and_manifest":
(artifacts / "effect.json").write_text("evil")
digest = __import__("hashlib").sha256(b"evil").hexdigest()
manifest["stages"]["stage"]["files"]["effect.json"] = {
"sha256": digest, "size": 4,
}
elif tamper == "spec":
manifest["spec_fingerprint"] = "sha256:" + "0" * 64
else:
manifest["stages"]["other"] = manifest["stages"].pop("stage")
manifest_path.write_text(json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n")
with pytest.raises(CorruptCheckpointError, match="artifact"):
run_job(_spec(tmp_path).with_resume(report.run_id), [stage])
def test_successful_job_is_idempotently_resumable(tmp_path):
calls = []
def normalize(_context):
calls.append("normalize")
first = run_job(_spec(tmp_path), [normalize])
second = run_job(_spec(tmp_path, resume_run_id=first.run_id), [normalize])
assert first.status == second.status == "succeeded"
assert second.resumed_from == first.run_id
assert calls == ["normalize"]
def test_checkpoints_and_report_are_json_safe_and_do_not_disclose_workspace_path(tmp_path):
def publish(_context):
return {"ignored": "/customer/alice", "password": "secret"}
report = run_job(_spec(tmp_path), [publish])
run_dir = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id
checkpoint = json.loads((run_dir / "checkpoint.json").read_text())
payload = (run_dir / "report.json").read_text()
parsed = json.loads(payload)
assert checkpoint["status"] == "succeeded"
assert parsed["schema_version"] == 1
assert parsed["run_id"] == report.run_id
assert str(tmp_path) not in payload
assert "alice" not in payload
assert "secret" not in payload
assert parsed["started_at"].endswith("Z")
assert parsed["finished_at"].endswith("Z")
def test_dry_run_is_exposed_to_stages_and_report(tmp_path):
observed = []
def plan(context):
observed.append(context.dry_run)
report = run_job(_spec(tmp_path, dry_run=True), [plan])
assert observed == [True]
assert report.dry_run is True
assert report.status == "succeeded"
def test_corrupt_checkpoint_is_rejected_without_running_stages(tmp_path):
first = run_job(_spec(tmp_path), [lambda _context: None])
checkpoint = (
tmp_path / ".tht-jobs" / "evidence" / "runs" / first.run_id / "checkpoint.json"
)
checkpoint.write_text("{not-json")
called = False
def stage(_context):
nonlocal called
called = True
with pytest.raises(CorruptCheckpointError, match="checkpoint is invalid"):
run_job(_spec(tmp_path, resume_run_id=first.run_id), [stage])
assert called is False
def test_stage_timestamps_are_aware_and_ordered(tmp_path):
report = run_job(_spec(tmp_path), [lambda _context: None])
stage = report.stages[0]
assert stage.started_at.tzinfo is not None
assert stage.finished_at.tzinfo is not None
assert stage.started_at <= stage.finished_at
assert report.started_at <= stage.started_at <= report.finished_at
@pytest.mark.parametrize(
("update", "replacement"),
[
("dry_run", True),
("spec_version", "jobs-v2"),
("pipeline_version", "evidence-v2"),
("config_fingerprint", "sha256:" + "3" * 64),
("input_fingerprint", "sha256:" + "4" * 64),
("workspace_id", "other"),
("job_type", "dwh"),
],
)
def test_resume_rejects_changed_identity_or_inputs_before_stage_execution(
tmp_path, update, replacement
):
first = run_job(_spec(tmp_path), [lambda _context: None])
called = False
def stage(_context):
nonlocal called
called = True
values = {update: replacement, "resume_run_id": first.run_id}
with pytest.raises(CorruptCheckpointError, match="incompatible"):
run_job(_spec(tmp_path, **values), [stage])
assert called is False
@pytest.mark.parametrize("stages", [[], [lambda _context: None, lambda _context: None]])
def test_resume_rejects_removed_or_inserted_stages(tmp_path, stages):
def first_stage(_context):
pass
first = run_job(_spec(tmp_path), [first_stage])
with pytest.raises(CorruptCheckpointError, match="incompatible"):
run_job(_spec(tmp_path, resume_run_id=first.run_id), stages)
def test_resume_rejects_reordered_stages(tmp_path):
def one(_context):
pass
def two(_context):
pass
first = run_job(_spec(tmp_path, stage_ids=("one", "two")), [one, two])
with pytest.raises(CorruptCheckpointError, match="incompatible"):
run_job(
_spec(tmp_path, resume_run_id=first.run_id, stage_ids=("two", "one")),
[two, one],
)
def test_hostile_exception_identity_never_enters_terminal_report(tmp_path):
Hostile = type("ApiKey_secret_/customer/alice", (Exception,), {})
def fail(_context):
raise Hostile("password=hunter2")
report = run_job(_spec(tmp_path), [fail])
payload = report.model_dump_json()
assert report.status == "failed"
assert report.stages[0].error.model_dump() == {
"category": "internal",
"code": "stage_exception",
"message": "stage execution failed",
}
assert "secret" not in payload
assert "alice" not in payload
assert "hunter2" not in payload
def test_new_run_without_resume_allows_intentional_spec_change(tmp_path):
first = run_job(_spec(tmp_path), [lambda _context: None])
second = run_job(_spec(tmp_path, input_fingerprint="sha256:" + "9" * 64), [lambda _context: None])
assert second.status == "succeeded"
assert second.run_id != first.run_id
assert second.resumed_from is None
def test_run_directories_are_private_and_fsynced_before_atomic_replace(tmp_path, monkeypatch):
events = []
real_replace = os.replace
monkeypatch.setattr(runner_module.os, "fsync", lambda _fd: events.append("fsync"))
def tracked_replace(source, destination):
events.append("replace")
real_replace(source, destination)
monkeypatch.setattr(runner_module.os, "replace", tracked_replace)
report = run_job(_spec(tmp_path), [lambda _context: None])
run_dir = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id
assert run_dir.stat().st_mode & 0o777 == 0o700
first_replace = events.index("replace")
assert "fsync" in events[:first_replace]
assert "fsync" in events[first_replace + 1 :]
def _tamper_checkpoint(tmp_path, report, transform):
path = tmp_path / ".tht-jobs" / report.job_type / "runs" / report.run_id / "checkpoint.json"
payload = json.loads(path.read_text())
transform(payload)
path.write_text(json.dumps(payload))
@pytest.mark.parametrize(
"transform",
[
lambda payload: payload["stages"].pop(),
lambda payload: payload["stages"].reverse(),
lambda payload: payload["stages"].append(dict(payload["stages"][0])),
lambda payload: payload["stages"][0].update(name="substitute"),
lambda payload: payload.update(input_fingerprint="sha256:" + "f" * 64),
lambda payload: payload["stages"][0].update(status="pending"),
],
)
def test_semantically_tampered_checkpoint_fails_without_orphan_run(tmp_path, transform):
def one(_context):
pass
def two(_context):
pass
spec = _spec(tmp_path, stage_ids=("one", "two"))
first = run_job(spec, [one, two])
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
before = {path.name for path in runs.iterdir()}
_tamper_checkpoint(tmp_path, first, transform)
called = False
def forbidden(_context):
nonlocal called
called = True
with pytest.raises(CorruptCheckpointError, match="invalid|incompatible"):
run_job(spec.with_resume(first.run_id), [forbidden, forbidden])
assert called is False
assert {path.name for path in runs.iterdir()} == before
def test_stored_compatibility_fingerprint_tamper_fails_without_orphan(tmp_path):
first = run_job(_spec(tmp_path), [lambda _context: None])
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
before = {path.name for path in runs.iterdir()}
_tamper_checkpoint(
tmp_path,
first,
lambda payload: payload.update(compatibility_fingerprint="sha256:" + "0" * 64),
)
with pytest.raises(CorruptCheckpointError, match="invalid"):
run_job(_spec(tmp_path, resume_run_id=first.run_id), [lambda _context: None])
assert {path.name for path in runs.iterdir()} == before
+151
View File
@@ -0,0 +1,151 @@
import hashlib
import pytest
from tht.jobs.dwh_pipeline import DwhPreprocessPipeline
FP = "sha256:" + hashlib.sha256(b"test").hexdigest()
def test_lsh_failure_resumes_exact_run_without_repeating_introspection(tmp_path):
calls = []
def introspect(output):
calls.append("introspect")
output.write_text("catalog")
def fail_lsh(physical, output):
calls.append("lsh-failed")
raise RuntimeError("database detail that must not leak")
failed = DwhPreprocessPipeline(
workspace_id="demo",
workspace_root=tmp_path,
config_fingerprint=FP,
input_fingerprint=FP,
introspect=introspect,
build_lsh=fail_lsh,
).run(("introspect", "lsh"))
assert failed.status == "failed"
resumed = DwhPreprocessPipeline(
workspace_id="demo",
workspace_root=tmp_path,
config_fingerprint=FP,
input_fingerprint=FP,
introspect=introspect,
build_lsh=lambda physical, output: _recover_lsh(calls, output),
).run(("introspect", "lsh"), resume_run_id=failed.run_id)
assert resumed.status == "succeeded"
assert resumed.resumed_from == failed.run_id
assert calls == ["introspect", "lsh-failed", "lsh-recovered"]
def _recover_lsh(calls, output):
calls.append("lsh-recovered")
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json"):
(output / name).write_text(name)
def test_resume_rejects_a_different_stage_selection(tmp_path):
failed = DwhPreprocessPipeline(
workspace_id="demo",
workspace_root=tmp_path,
config_fingerprint=FP,
input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: (_ for _ in ()).throw(RuntimeError()),
).run(("introspect", "lsh"))
pipeline = DwhPreprocessPipeline(
workspace_id="demo",
workspace_root=tmp_path,
config_fingerprint=FP,
input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: None,
)
try:
pipeline.run(("lsh",), resume_run_id=failed.run_id)
except Exception as error:
assert "incompatible" in str(error)
else:
raise AssertionError("resume with different stages must fail")
def test_resume_rejects_tampered_succeeded_stage_artifact(tmp_path):
failed = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: (_ for _ in ()).throw(RuntimeError()),
).run(("introspect", "lsh"))
artifact = (
tmp_path / ".tht-jobs" / "dwh" / "runs" / failed.run_id
/ "artifacts" / "physical.yaml"
)
artifact.write_text("tampered")
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _recover_lsh([], output),
)
with pytest.raises(Exception, match="artifact manifest is invalid"):
pipeline.run(("introspect", "lsh"), resume_run_id=failed.run_id)
def test_post_publish_crash_reconciles_same_generation_on_resume(tmp_path):
crashed = False
builder_calls = 0
def crash_once(_generation):
nonlocal crashed
if not crashed:
crashed = True
raise KeyboardInterrupt("simulated process death")
def build(physical, output):
nonlocal builder_calls
builder_calls += 1
_recover_lsh([], output)
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=build,
after_publish=crash_once,
)
with pytest.raises(KeyboardInterrupt):
pipeline.run(("introspect", "lsh"))
active = (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip()
checkpoint = next((tmp_path / ".tht-jobs" / "dwh" / "runs").glob("*/checkpoint.json"))
source_run_id = checkpoint.parent.name
assert active == source_run_id
resumed = pipeline.run(("introspect", "lsh"), resume_run_id=source_run_id)
assert resumed.status == "succeeded"
assert builder_calls == 1
assert (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip() == source_run_id
def test_resume_of_succeeded_run_detects_tampered_published_file(tmp_path):
pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("catalog"),
build_lsh=lambda physical, output: _recover_lsh([], output),
)
succeeded = pipeline.run(("introspect", "lsh"))
published = (
tmp_path / ".tht-dwh" / "generations" / succeeded.run_id / "demo_meta.json"
)
published.chmod(0o600)
published.write_text("tampered")
with pytest.raises(Exception, match="published DWH"):
pipeline.run(("introspect", "lsh"), resume_run_id=succeeded.run_id)
+4 -4
View File
@@ -56,12 +56,12 @@ def test_save_one_memory_preserves_subject_through_upsert_row():
from tht.memory import save_one_memory
writer = MagicMock()
writer.upsert_records.return_value = 1
writer.upsert.return_value = 1
embedder = MagicMock()
embedder.embed_documents.return_value = [[0.0] * 8]
save_one_memory([_record(decision_seq=1)], decision_seq=1, writer=writer, embedder=embedder)
row = writer.upsert_records.call_args[0][1][0]
md = row["metadata"]
save_one_memory([_record(decision_seq=1)], decision_seq=1, store=writer, embedder=embedder)
row = writer.upsert.call_args[0][1][0]
md = row.record.metadata
assert md["subject"] == "dim_pazienti"
assert md["detail"] == "promossa"
assert md["rationale"] == "perche' serve"
+11 -14
View File
@@ -8,11 +8,8 @@ pgvector as a one-row upsert. This test pins the pure core of that behavior:
- writer.sync is NEVER called (that is the full-resync path)
"""
from datetime import datetime
from pathlib import Path
from unittest.mock import MagicMock
import pytest
from tht.memory import MemoryRecord, memory_vector_record_for_decision, save_one_memory
@@ -47,31 +44,31 @@ def test_returns_none_for_unknown_decision_seq():
def test_save_one_calls_upsert_with_single_row_never_sync():
records = [_record(seq=7)]
writer = MagicMock()
writer.upsert_records.return_value = 1
writer.upsert.return_value = 1
embedder = MagicMock()
embedder.embed_documents.return_value = [[0.1] * 8]
upserted = save_one_memory(records, decision_seq=7, writer=writer, embedder=embedder)
upserted = save_one_memory(records, decision_seq=7, store=writer, embedder=embedder)
assert upserted == 1
writer.sync.assert_not_called() # the whole point of D11: no full resync
writer.upsert_records.assert_called_once()
args = writer.upsert_records.call_args
writer.upsert.assert_called_once()
args = writer.upsert.call_args
# table is memory, exactly one row
assert args[0][0] == "memory"
rows = args[0][1]
assert len(rows) == 1
assert rows[0]["record_key"] == "memory:mem-0007"
assert "embedding" in rows[0]
assert rows[0].record.id == "memory:mem-0007"
assert rows[0].embedding
def test_save_one_no_record_for_seq_is_noop():
records = [_record(seq=7)]
writer = MagicMock()
embedder = MagicMock()
upserted = save_one_memory(records, decision_seq=42, writer=writer, embedder=embedder)
upserted = save_one_memory(records, decision_seq=42, store=writer, embedder=embedder)
assert upserted == 0
writer.upsert_records.assert_not_called()
writer.upsert.assert_not_called()
writer.sync.assert_not_called()
embedder.embed_documents.assert_not_called()
@@ -81,9 +78,9 @@ def test_save_one_uses_writer_key_for_upsert():
Verified indirectly: save_one_memory takes the writer as its client argument."""
records = [_record(seq=7)]
writer = MagicMock()
writer.upsert_records.return_value = 1
writer.upsert.return_value = 1
embedder = MagicMock()
embedder.embed_documents.return_value = [[0.0] * 4]
save_one_memory(records, decision_seq=7, writer=writer, embedder=embedder)
save_one_memory(records, decision_seq=7, store=writer, embedder=embedder)
# one upsert call, single row, table=memory
assert writer.upsert_records.call_count == 1
assert writer.upsert.call_count == 1
+110
View File
@@ -0,0 +1,110 @@
from pathlib import Path
import pytest
from tht.config import ConfigError, load_config
from tht.paths import resolve_workspace_paths
def _write_config(path: Path, *, sessions: str = "sessions", absolute: Path | None = None) -> Path:
root = absolute or Path("artifacts")
path.write_text(
"dwh:\n"
" type: postgres_direct\n"
" connection:\n"
" database: db\n"
" schema: public\n"
" user: user\n"
" password: secret\n"
"roots:\n"
f" artifacts: {root}\n"
f" indexes: {root if absolute else 'indexes'}\n"
f" sessions: {absolute if absolute else sessions}\n"
)
return path
def test_relative_paths_resolve_under_workspace_root(tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml")
cfg = load_config(cfg_path)
resolved = resolve_workspace_paths(cfg_path, cfg, tmp_path / "data")
assert resolved.workspace == tmp_path / "data/workspaces/demo"
assert resolved.sessions == tmp_path / "data/workspaces/demo/sessions"
assert resolved.artifacts == tmp_path / "data/workspaces/demo/artifacts"
assert resolved.indexes == tmp_path / "data/workspaces/demo/indexes"
assert resolved.corpus == tmp_path / "data/workspaces/demo/corpus"
def test_absolute_legacy_paths_are_preserved(tmp_path):
legacy = tmp_path / "existing-workspace"
cfg_path = _write_config(tmp_path / "demo.yaml", absolute=legacy)
cfg = load_config(cfg_path)
resolved = resolve_workspace_paths(cfg_path, cfg, tmp_path / "data")
assert resolved.sessions == legacy
assert resolved.artifacts == legacy
assert resolved.indexes == legacy
def test_path_escape_is_rejected(tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml", sessions="../../private")
cfg = load_config(cfg_path)
with pytest.raises(ConfigError, match="outside workspace root"):
resolve_workspace_paths(cfg_path, cfg, tmp_path / "data")
def test_workspace_symlink_escape_is_rejected(tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml")
cfg = load_config(cfg_path)
data_root = tmp_path / "data"
(data_root / "workspaces").mkdir(parents=True)
(data_root / "workspaces" / "demo").symlink_to(tmp_path / "private", target_is_directory=True)
with pytest.raises(ConfigError, match="outside workspaces root"):
resolve_workspace_paths(cfg_path, cfg, data_root)
def test_nested_root_symlink_escape_is_rejected(tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml")
cfg = load_config(cfg_path)
workspace = tmp_path / "data/workspaces/demo"
workspace.mkdir(parents=True)
(workspace / "sessions").symlink_to(tmp_path / "private", target_is_directory=True)
with pytest.raises(ConfigError, match="outside workspace root"):
resolve_workspace_paths(cfg_path, cfg, tmp_path / "data")
def test_corpus_symlink_escape_is_rejected(tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml")
cfg = load_config(cfg_path)
workspace = tmp_path / "data/workspaces/demo"
workspace.mkdir(parents=True)
(workspace / "corpus").symlink_to(tmp_path / "private", target_is_directory=True)
with pytest.raises(ConfigError, match="outside workspace root"):
resolve_workspace_paths(cfg_path, cfg, tmp_path / "data")
def test_data_root_environment_activates_portable_paths(monkeypatch, tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml")
monkeypatch.setenv("THT_DATA_ROOT", str(tmp_path / "data"))
cfg = load_config(cfg_path)
assert cfg.paths.sessions == tmp_path / "data/workspaces/demo/sessions"
assert cfg.paths.artifacts == tmp_path / "data/workspaces/demo/artifacts"
def test_no_data_root_preserves_legacy_relative_paths(monkeypatch, tmp_path):
cfg_path = _write_config(tmp_path / "demo.yaml")
monkeypatch.delenv("THT_DATA_ROOT", raising=False)
cfg = load_config(cfg_path)
assert cfg.paths.sessions == Path("sessions")
assert cfg.paths.artifacts == Path("artifacts")
+131
View File
@@ -0,0 +1,131 @@
import json
from types import SimpleNamespace
from typer.testing import CliRunner
from tht.cli import app
def test_preprocess_evidence_json_is_pristine(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
result = SimpleNamespace(model_dump=lambda mode=None: {
"status": "succeeded", "generation": "gen:abc", "published": True
})
monkeypatch.setattr(command, "run_from_config", lambda *args, **kwargs: result)
response = CliRunner().invoke(
app, ["preprocess", "evidence", "--json", "-c", str(tmp_path / "workspace.yaml")]
)
assert response.exit_code == 0, response.output
assert json.loads(response.output)["generation"] == "gen:abc"
def test_preprocess_failure_is_structured_and_nonzero(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
monkeypatch.setattr(command, "run_from_config", lambda *a, **k: (_ for _ in ()).throw(RuntimeError("secret detail")))
response = CliRunner().invoke(
app, ["preprocess", "evidence", "--json", "-c", str(tmp_path / "workspace.yaml")]
)
assert response.exit_code != 0
assert json.loads(response.output) == {"status": "failed", "error": "preprocessing failed"}
assert "secret detail" not in response.output
def test_preprocess_failed_job_report_is_sanitized_json_and_nonzero(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
result = SimpleNamespace(model_dump=lambda mode=None: {
"status": "failed", "run_id": "a" * 32, "published": False,
"generation": "gen:" + "b" * 32, "changed": ["fs:one"],
})
monkeypatch.setattr(command, "run_from_config", lambda *args, **kwargs: result)
response = CliRunner().invoke(
app, ["preprocess", "evidence", "--json", "-c", str(tmp_path / "workspace.yaml")]
)
assert response.exit_code == 1
payload = json.loads(response.output)
assert payload["status"] == "failed"
assert payload["error"] == "preprocessing job failed"
assert "traceback" not in response.output.lower()
def test_preprocess_real_failed_stage_result_exits_nonzero(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
from test_corpus_pipeline import Source, item, pipeline
result = pipeline(
tmp_path, Source([(item("one", "a"), RuntimeError("SENSITIVE EVIDENCE secret"))])
).run_as_job(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint="sha256:" + "1" * 64,
input_fingerprint="sha256:" + "2" * 64,
)
assert result.status == "failed"
monkeypatch.setattr(command, "run_from_config", lambda *args, **kwargs: result)
response = CliRunner().invoke(
app, ["preprocess", "evidence", "--json", "-c", str(tmp_path / "workspace.yaml")]
)
assert response.exit_code == 1
assert json.loads(response.output)["status"] == "failed"
assert "SENSITIVE EVIDENCE" not in response.output
assert "secret" not in response.output
def test_preprocess_evidence_text_uses_uncapped_result_counts(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
result = SimpleNamespace(model_dump=lambda mode=None: {
"status": "succeeded", "run_id": "a" * 32,
"generation": "gen:" + "b" * 64, "published": True,
"changed": ["fs:item"] * 100,
"unchanged": ["fs:item"] * 100,
"removed": ["fs:item"] * 100,
"counts": {"changed": 1001, "unchanged": 902, "removed": 803},
})
monkeypatch.setattr(command, "run_from_config", lambda *args, **kwargs: result)
response = CliRunner().invoke(
app, ["preprocess", "evidence", "-c", str(tmp_path / "workspace.yaml")]
)
assert response.exit_code == 0, response.output
assert "changed=1001 unchanged=902 removed=803" in response.output
def test_preprocess_resume_rejects_generation_id_before_configuration(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
called = False
def forbidden(*args, **kwargs):
nonlocal called
called = True
monkeypatch.setattr(command, "run_from_config", forbidden)
response = CliRunner().invoke(
app,
[
"preprocess", "evidence", "--resume", "gen:" + "a" * 32,
"--json", "-c", str(tmp_path / "workspace.yaml"),
],
)
assert response.exit_code != 0
assert json.loads(response.output) == {
"status": "failed", "error": "resume requires a preprocessing run id"
}
assert called is False
def test_preprocess_evidence_gc_json_is_pristine(monkeypatch, tmp_path):
import tht.cli.preprocess_cmd as command
monkeypatch.setattr(command, "gc_from_config", lambda *a, **k: {
"status": "succeeded", "dry_run": True, "evicted": [], "failures": [],
})
response = CliRunner().invoke(
app, ["preprocess", "evidence", "gc", "--dry-run", "--json", "-c",
str(tmp_path / "workspace.yaml")]
)
assert response.exit_code == 0, response.output
assert json.loads(response.output)["dry_run"] is True
+199
View File
@@ -0,0 +1,199 @@
from datetime import UTC, datetime
import pytest
from tht.ports.evidence import EvidenceSourceError
class Body:
def __init__(self, data): self.data, self.closed = data, False
def read(self, amount): return self.data[:amount]
def close(self): self.closed = True
class Client:
def __init__(self): self.body, self.list_calls = Body(b"hello"), 0
def get_paginator(self, name): return self
def list_objects_v2(self, **kwargs):
self.list_calls += 1
return {"Contents": [{"Key": "clinical/a.md", "ETag": '"abc"',
"Size": 5, "LastModified": datetime(2026, 1, 1, tzinfo=UTC)}]}
def paginate(self, **kwargs):
yield {"Contents": [{"Key": "clinical/a.md", "ETag": '"abc"',
"Size": 5, "LastModified": datetime(2026, 1, 1, tzinfo=UTC)}]}
def get_object(self, **kwargs):
assert kwargs == {"Bucket": "evidence", "Key": "clinical/a.md"}
return {"Body": self.body, "ContentLength": 5, "ContentType": "text/markdown",
"ETag": '"abc"'}
def test_s3_canonical_uri_version_fingerprint_and_closed_body():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
source = S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client)
item = next(iter(source.discover()))
assert item.uri == "s3://evidence/clinical/a.md"
assert item.fingerprint.startswith("etag:")
assert source.acquire(item).content == b"hello"
assert client.body.closed
def test_s3_etag_fallback_and_bounds():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
with pytest.raises(ValueError):
S3EvidenceSource(bucket="evidence", client=client, max_objects=0)
def test_s3_rejects_private_or_insecure_endpoint_without_explicit_opt_in():
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="trusted"):
S3EvidenceSource(bucket="evidence", endpoint_url="https://127.0.0.1:9000", client=Client())
with pytest.raises(ValueError, match="HTTPS"):
S3EvidenceSource(bucket="evidence", endpoint_url="http://s3.example.test", client=Client())
source = S3EvidenceSource(bucket="evidence", endpoint_url="http://127.0.0.1:9000",
trusted_endpoint=True, allow_private_endpoint=True,
allow_insecure_endpoint=True, client=Client())
assert source is not None
@pytest.mark.parametrize("bucket", ["UPPER", "bad_bucket", "-start", "end-", "a..b"])
def test_s3_rejects_invalid_bucket_names(bucket):
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="bucket"):
S3EvidenceSource(bucket=bucket, client=Client())
@pytest.mark.parametrize("bucket", ["127.0.0.1", "192.168.1.1"])
def test_s3_rejects_ip_shaped_bucket(bucket):
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="bucket"):
S3EvidenceSource(bucket=bucket, client=Client())
def test_s3_rejects_endpoint_query_path_fragment_and_untrusted_custom_host():
from tht.adapters.evidence.s3 import S3EvidenceSource
for endpoint in ("https://s3.example.test/path", "https://s3.example.test/?x=1",
"https://s3.example.test/#x"):
with pytest.raises(ValueError, match="root"):
S3EvidenceSource(bucket="evidence", endpoint_url=endpoint,
trusted_endpoint=True, client=Client())
with pytest.raises(ValueError, match="trusted"):
S3EvidenceSource(bucket="evidence", endpoint_url="https://s3.example.test", client=Client())
def test_s3_rejects_out_of_prefix_key_and_missing_validator():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
client.list_objects_v2 = lambda **kwargs: {"Contents": [{"Key": "other/a.md", "ETag": '"x"'}]}
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
client.list_objects_v2 = lambda **kwargs: {"Contents": [{"Key": "clinical/a.md"}]}
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
def test_s3_rejects_leading_slash_prefix_empty_and_control_keys():
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="prefix"):
S3EvidenceSource(bucket="evidence", prefix="/clinical", client=Client())
for key in ("", "clinical/a\x00.md", "clinical/a\x7f.md"):
client = Client()
client.list_objects_v2 = lambda **kwargs: {"Contents": [{"Key": key, "ETag": '"x"'}]}
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", client=client).discover())
@pytest.mark.parametrize("prefix", ["/bad", "x" * 1025, "bad\x00prefix", "bad\x7fprefix"])
def test_s3_rejects_invalid_prefix_before_client_request(prefix):
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
with pytest.raises(ValueError, match="prefix"):
S3EvidenceSource(bucket="evidence", prefix=prefix, client=client)
assert client.list_calls == 0
def test_s3_hard_page_limit_never_requests_page_max_plus_one():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
def listing(**kwargs):
client.list_calls += 1
return {"Contents": [{"Key": f"clinical/{client.list_calls}.md", "ETag": '"x"'}],
"IsTruncated": True, "NextContinuationToken": str(client.list_calls)}
client.list_objects_v2 = listing
with pytest.raises(EvidenceSourceError):
list(S3EvidenceSource(bucket="evidence", prefix="clinical/", max_pages=2,
client=client).discover())
assert client.list_calls == 2
def test_s3_acquire_rejects_exact_etag_drift_and_closes_body():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
source = S3EvidenceSource(bucket="evidence", client=client)
item = next(iter(source.discover()))
client.get_object = lambda **kwargs: {"Body": client.body, "ContentLength": 5,
"ETag": '"changed"'}
with pytest.raises(EvidenceSourceError):
source.acquire(item)
assert client.body.closed
def test_s3_acquire_rejects_forged_reconstructed_item_before_get():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
source = S3EvidenceSource(bucket="evidence", client=client)
item = next(iter(source.discover()))
forged = item.model_copy(update={"fingerprint": "etag:" + "0" * 64})
client.get_object = lambda **kwargs: (_ for _ in ()).throw(AssertionError("called"))
with pytest.raises(EvidenceSourceError):
source.acquire(forged)
@pytest.mark.parametrize("host", ["127.0.0.1", "10.0.0.1", "169.254.1.1", "0.0.0.0",
"[::1]", "[fe80::1]", "[::]"])
def test_s3_literal_non_global_endpoint_requires_private_opt_in(host):
from tht.adapters.evidence.s3 import S3EvidenceSource
with pytest.raises(ValueError, match="private"):
S3EvidenceSource(bucket="evidence", endpoint_url=f"https://{host}:9000",
trusted_endpoint=True, client=Client())
def test_s3_size_limit_closes_body():
from tht.adapters.evidence.s3 import S3EvidenceSource
client = Client()
source = S3EvidenceSource(bucket="evidence", client=client, max_bytes=4)
item = next(iter(source.discover()))
with pytest.raises(EvidenceSourceError):
source.acquire(item)
assert client.body.closed
def test_s3_config_serialization_masks_credentials():
from tht.config import S3EvidenceSourceConfig
config = S3EvidenceSourceConfig(type="s3", bucket="evidence",
access_key="access-secret", secret_key="write-secret")
assert "access-secret" not in repr(config)
assert "write-secret" not in repr(config)
def test_s3_config_loads_credentials_from_secret_files(tmp_path):
from tht.config import load_config
access, secret = tmp_path / "access", tmp_path / "secret"
access.write_text("access-value")
secret.write_text("secret-value")
workspace = tmp_path / "workspace.yaml"
workspace.write_text(f"""
dwh:
type: postgres_direct
connection: {{database: d, schema: public, user: u, password: p}}
evidence:
sources:
- type: s3
bucket: evidence
access_key_file: {access}
secret_key_file: {secret}
""")
source = load_config(workspace).evidence.sources[0]
assert source.access_key.get_secret_value() == "access-value"
assert source.secret_key.get_secret_value() == "secret-value"
+73 -6
View File
@@ -4,6 +4,8 @@ from typer.testing import CliRunner
from tht.cli import app
from tht.mschema.models import ColumnPhysical, PhysicalSchema, TablePhysical
from tht.config import ExamplesConfig
from tht.cli.schema_cmd import _add_examples
def _write_catalog(tmp_path):
@@ -32,17 +34,62 @@ def _write_config(tmp_path):
return cfg
def test_introspect_cache_hit_skips_dwh(tmp_path):
# Le credenziali sono fasulle: se la guardia non scattasse PRIMA del branch
# transport, il comando tenterebbe la connessione e fallirebbe.
def test_introspect_rejects_unbound_legacy_cache(tmp_path):
catalog = _write_catalog(tmp_path)
before = catalog.read_bytes()
cfg = _write_config(tmp_path)
res = CliRunner().invoke(app, ["schema", "introspect", "-c", str(cfg)])
assert res.exit_code == 1
assert "legacy artifacts are unbound" in res.output
assert catalog.exists()
def test_introspect_fresh_root_initializes_through_writer_job(tmp_path, monkeypatch):
import tht.cli.schema_cmd as module
cfg = _write_config(tmp_path)
physical = PhysicalSchema(
database="d", schema="s", introspected_at=datetime(2026, 1, 1),
tables={"dim_patient": TablePhysical(columns={"id": ColumnPhysical(type="bigint")})},
)
def refresh(_cfg, *, output_path=None, **_kwargs):
physical.to_yaml(output_path)
return physical
monkeypatch.setattr(module, "refresh_catalog", refresh)
res = CliRunner().invoke(app, ["schema", "introspect", "-c", str(cfg)])
assert res.exit_code == 0, res.output
assert "OK (cache)" in res.output
assert (tmp_path / ".tht-dwh" / "OWNER.json").is_file()
assert "1 tabelle" in res.output
assert catalog.read_bytes() == before
def test_lsh_build_fresh_root_initializes_introspection_and_lsh(tmp_path, monkeypatch):
import tht.cli.lsh_cmd as lsh_module
import tht.cli.schema_cmd as schema_module
import tht.lshindex as lshindex_module
cfg = _write_config(tmp_path)
physical = PhysicalSchema(
database="d", schema="s", introspected_at=datetime(2026, 1, 1),
tables={"dim_patient": TablePhysical(columns={"id": ColumnPhysical(type="bigint")})},
)
def refresh(_cfg, *, output_path=None, **_kwargs):
physical.to_yaml(output_path)
return physical
def build(_cfg, *, physical_file, output_dir, **_kwargs):
assert physical_file.is_file()
for name in ("s_lsh.pkl", "s_minhashes.pkl", "s_meta.json"):
(output_dir / name).write_text("index")
return {}, [], [], {}
monkeypatch.setattr(schema_module, "refresh_catalog", refresh)
monkeypatch.setattr(lsh_module, "build_lsh_artifacts", build)
monkeypatch.setattr(lshindex_module, "load_index", lambda *_args, **_kwargs: (None, {}, None))
res = CliRunner().invoke(app, ["lsh", "build", "-c", str(cfg)])
assert res.exit_code == 0, res.output
assert (tmp_path / ".tht-dwh" / "OWNER.json").is_file()
def test_introspect_refresh_bypasses_cache(tmp_path):
@@ -68,3 +115,23 @@ def test_render_without_catalog_guides_fallback(tmp_path):
res = CliRunner().invoke(app, ["schema", "render", "-c", str(cfg)])
assert res.exit_code == 1
assert "Esegui prima" in res.output
def test_examples_skip_one_unreadable_column_and_continue(caplog):
physical = PhysicalSchema(
database="d", schema="s", introspected_at=datetime(2026, 1, 1),
tables={"t": TablePhysical(columns={
"bad": ColumnPhysical(type="text"), "good": ColumnPhysical(type="text")
})},
)
class Dwh:
def sample_column(self, table, column, *, limit):
if column == "bad":
raise RuntimeError("denied")
return ["kept"]
_add_examples(Dwh(), physical, ExamplesConfig(max_per_column=3))
assert physical.tables["t"].columns["bad"].examples == []
assert physical.tables["t"].columns["good"].examples == ["kept"]
assert "Campionamento saltato" in caplog.text
+19 -3
View File
@@ -5,6 +5,8 @@ from types import SimpleNamespace
from typer.testing import CliRunner
from tht.cli import app
from tht.config import load_config
from tht.jobs.dwh_pipeline import DwhPreprocessPipeline, config_dwh_binding
from tht.mschema.models import ColumnPhysical, PhysicalSchema, TablePhysical
from tht.vectorstore.embeddings import EmbeddingsError
@@ -43,11 +45,11 @@ class _FakeSearcher:
def _workspace(tmp_path, with_session=None):
PhysicalSchema(
physical = PhysicalSchema(
database="d", schema="s", introspected_at=datetime(2026, 1, 1),
tables={"fact_ablazione": TablePhysical(
comment="Ablazioni", columns={"cod_paz": ColumnPhysical(type="bigint")})},
).to_yaml(tmp_path / "artifacts" / "mschema" / "physical.yaml")
)
cfg = tmp_path / "workspace.yaml"
cfg.write_text(
"database: {database: d, schema: s, user: u, password: p, transport: direct}\n"
@@ -56,6 +58,18 @@ def _workspace(tmp_path, with_session=None):
f"paths: {{artifacts: {tmp_path/'artifacts'}, indexes: {tmp_path/'i'}, "
f"sessions: {tmp_path/'sessions'}}}\n"
)
binding = config_dwh_binding(load_config(cfg))
DwhPreprocessPipeline(
workspace_id=binding["workspace_id"], workspace_root=tmp_path,
config_fingerprint=binding["config_fingerprint"],
input_fingerprint=binding["input_fingerprint"],
introspect=lambda output: physical.to_yaml(output),
build_lsh=lambda _physical, output: [
(output / name).write_text("index")
for name in ("s_lsh.pkl", "s_minhashes.pkl", "s_meta.json")
],
lsh_filenames=("s_lsh.pkl", "s_minhashes.pkl", "s_meta.json"),
).run()
if with_session:
sdir = tmp_path / "sessions" / with_session
sdir.mkdir(parents=True)
@@ -81,7 +95,9 @@ def test_pack_single_embed_and_sections(tmp_path, monkeypatch):
assert res.exit_code == 0, res.output
assert emb.calls == 1 # UN solo embedding per le tre ricerche
assert "fact_ablazione" in res.output and "Ablazioni" in res.output
assert "Dominio ablazione" in res.output
# Evidence is fail-closed until an ACTIVE corpus exists; legacy vector rows
# must not leak into a new search pack.
assert "Dominio ablazione" not in res.output
assert "SELECT 1" in res.output
@@ -90,6 +90,25 @@ def test_search_similar_reraises_non_404_with_kinds(monkeypatch):
_client().search_similar("memory", [0.1] * 4, 5, kinds=["memory"])
def test_generation_filter_is_sent_exactly_and_legacy_404_fails_closed(monkeypatch):
calls = []
def fake_call(self, function, payload):
calls.append(payload)
raise VectorRestError("HTTP 404 missing filtered RPC")
monkeypatch.setattr(VectorRestClient, "_call", fake_call)
metadata_filter = {"vector_generation": "gen:abc", "document_ids": ["doc:1"]}
with pytest.raises(VectorRestError, match="404"):
_client().search_similar(
"evidence", [0.1] * 4, 5, kinds=["evidence"], metadata_filter=metadata_filter
)
assert calls == [{
"query_embedding": [0.1] * 4, "limit_count": 5, "table_name": "evidence",
"kinds": ["evidence"], "metadata_filter": metadata_filter,
}]
def test_rest_searcher_forwards_kinds_to_client():
calls = []
+10 -10
View File
@@ -43,18 +43,18 @@ def test_solved_kind_maps_to_memory_table():
def test_save_upserts_single_row_into_memory_table():
writer = MagicMock()
writer.existing_hashes.return_value = {}
writer.upsert_records.return_value = 1
writer.upsert.return_value = 1
embedder = MagicMock()
embedder.embed_documents.return_value = [[0.1] * 8]
assert save_solved_question(_rec(), writer=writer, embedder=embedder) == 1
assert save_solved_question(_rec(), store=writer, embedder=embedder) == 1
writer.sync.assert_not_called()
table, rows = writer.upsert_records.call_args[0]
table, rows = writer.upsert.call_args[0]
assert table == "memory"
assert len(rows) == 1
assert rows[0]["record_key"] == "solved:s1"
assert rows[0]["metadata"]["kind"] == SOLVED_KIND
assert rows[0]["metadata"]["sql"].startswith("SELECT")
assert rows[0].record.id == "solved:s1"
assert rows[0].record.kind == SOLVED_KIND
assert rows[0].record.metadata["sql"].startswith("SELECT")
def test_save_skips_when_question_and_sql_unchanged():
@@ -62,9 +62,9 @@ def test_save_skips_when_question_and_sql_unchanged():
writer = MagicMock()
writer.existing_hashes.return_value = {r.id: _solved_hash(r)}
embedder = MagicMock()
assert save_solved_question(r, writer=writer, embedder=embedder) == 0
assert save_solved_question(r, store=writer, embedder=embedder) == 0
embedder.embed_documents.assert_not_called()
writer.upsert_records.assert_not_called()
writer.upsert.assert_not_called()
def test_sql_change_alone_triggers_reupsert():
@@ -72,7 +72,7 @@ def test_sql_change_alone_triggers_reupsert():
new = _rec(sql="SELECT 1") # stessa domanda, SQL diverso
writer = MagicMock()
writer.existing_hashes.return_value = {old.id: _solved_hash(old)}
writer.upsert_records.return_value = 1
writer.upsert.return_value = 1
embedder = MagicMock()
embedder.embed_documents.return_value = [[0.0] * 4]
assert save_solved_question(new, writer=writer, embedder=embedder) == 1
assert save_solved_question(new, store=writer, embedder=embedder) == 1
@@ -0,0 +1,59 @@
import os
import shutil
import subprocess
import sys
import zipfile
from pathlib import Path
def test_built_wheel_installs_vector_migrations_and_discovers_cli(tmp_path):
harness = Path(__file__).parents[1]
wheelhouse = tmp_path / "wheelhouse"
target = tmp_path / "site"
wheelhouse.mkdir()
uv = shutil.which("uv")
assert uv is not None, "uv is required to verify the production wheel"
build_env = {**os.environ, "UV_CACHE_DIR": str(tmp_path / "uv-cache")}
subprocess.run(
[
uv,
"build",
"--wheel",
"--out-dir",
str(wheelhouse),
str(harness),
],
check=True,
capture_output=True,
text=True,
env=build_env,
)
wheel = next(wheelhouse.glob("tht-*.whl"))
with zipfile.ZipFile(wheel) as archive:
names = set(archive.namelist())
assert "tht/migrations/vector/001_extensions.sql" in names
assert "tht/migrations/vector/003_roles.sql" in names
subprocess.run(
[sys.executable, "-m", "pip", "install", "--no-deps", "--target", str(target), wheel],
check=True,
capture_output=True,
text=True,
)
env = {**os.environ, "PYTHONPATH": str(target)}
probe = subprocess.run(
[
sys.executable,
"-c",
"from typer.testing import CliRunner; from tht.cli import app; "
"r=CliRunner().invoke(app, ['vector','migrate','--help']); "
"print(r.output); raise SystemExit(r.exit_code)",
],
env=env,
check=False,
capture_output=True,
text=True,
cwd=tmp_path,
)
assert probe.returncode == 0, probe.stderr + probe.stdout
assert "--status" in probe.stdout
+250
View File
@@ -0,0 +1,250 @@
from dataclasses import FrozenInstanceError
from unittest.mock import MagicMock
import pytest
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
from tht.evidence.model import EvidenceDoc
from tht.ports.vector import (
VectorHit,
VectorRecord,
VectorStore,
VectorReadUnavailable,
VectorWriteRecord,
VectorWriteUnavailable,
)
from tht.vectorstore.records import evidence_records
def test_http_store_reports_reader_without_writer():
reader = MagicMock()
store = ThothHttpVectorStore(reader=reader, writer=None)
assert store.capabilities.search is True
assert store.capabilities.upsert is False
with pytest.raises(VectorWriteUnavailable):
store.upsert("memory", [])
def test_http_store_supports_writer_without_reader():
writer = MagicMock()
store = ThothHttpVectorStore(reader=None, writer=writer, expected_dimension=768)
assert store.capabilities.search is False
assert store.capabilities.existing_hashes is True
assert store.capabilities.upsert is True
with pytest.raises(VectorReadUnavailable):
store.search(["memory"], [0.1], limit=1)
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
def test_http_search_requires_a_strict_positive_integer_limit(limit):
store = ThothHttpVectorStore(reader=MagicMock(), writer=None)
with pytest.raises(ValueError, match="positive integer"):
store.search(["memory"], [0.1], limit=limit)
def test_http_store_keeps_reader_and_writer_operations_separate():
reader = MagicMock()
reader.search_similar.return_value = [
{
"similarity": 0.75,
"metadata": {
"record_key": "m1",
"kind": "memory",
"ref": "session:s1",
"title": "Choice",
"content": "Use the curated table",
},
}
]
writer = MagicMock()
writer.existing_hashes.return_value = {"m1": "abc"}
writer.upsert_records.return_value = 1
store = ThothHttpVectorStore(reader=reader, writer=writer)
hits = store.search(["memory"], [0.1, 0.2], limit=3, kinds=["memory"])
assert hits == [
VectorHit(
id="m1",
kind="memory",
ref="session:s1",
title="Choice",
content="Use the curated table",
metadata={
"record_key": "m1",
"kind": "memory",
"ref": "session:s1",
"title": "Choice",
"content": "Use the curated table",
},
similarity=0.75,
)
]
reader.search_similar.assert_called_once_with(
"memory", [0.1, 0.2], 3, kinds=["memory"]
)
writer.search_similar.assert_not_called()
assert store.existing_hashes("memory", ["memory"]) == {"m1": "abc"}
writer.existing_hashes.assert_called_once_with("memory", ["memory"])
records = [
VectorWriteRecord(
record=VectorRecord(
id="m1",
kind="memory",
ref="session:s1",
title="Choice",
content="Use the curated table",
),
embedding=[0.1, 0.2],
content_hash="abc",
)
]
assert store.upsert("memory", records) == 1
writer.upsert_records.assert_called_once()
reader.upsert_records.assert_not_called()
def test_http_upsert_serializes_a_canonical_builder_record():
record = evidence_records(
[EvidenceDoc(id="joins", title="Join guidance", body="Use the curated join")],
max_chunk_chars=1000,
)[0]
writer = MagicMock()
writer.upsert_records.return_value = 1
store = ThothHttpVectorStore(reader=MagicMock(), writer=writer)
assert store.upsert(
"evidence",
[VectorWriteRecord(record=record, embedding=[0.2, 0.3], content_hash="digest")],
) == 1
row = writer.upsert_records.call_args.args[1][0]
assert row["record_key"] == "evidence:joins:0"
assert row["metadata"]["status"] == "reviewed"
assert row["embedding"] == [0.2, 0.3]
assert row["content_hash"] == "digest"
def test_http_upsert_preserves_metadata_named_like_transport_fields():
record = VectorRecord(
id="collision",
kind="memory",
ref="session:s1",
title="Collision",
content="Semantic metadata must survive",
metadata={"embedding": "semantic embedding", "content_hash": "semantic hash"},
)
writer = MagicMock()
store = ThothHttpVectorStore(reader=MagicMock(), writer=writer)
store.upsert(
"memory",
[VectorWriteRecord(record=record, embedding=[0.4], content_hash="transport hash")],
)
row = writer.upsert_records.call_args.args[1][0]
assert row["embedding"] == [0.4]
assert row["content_hash"] == "transport hash"
assert row["metadata"]["embedding"] == "semantic embedding"
assert row["metadata"]["content_hash"] == "semantic hash"
def test_http_store_is_runtime_vector_store():
store = ThothHttpVectorStore(reader=MagicMock(), writer=None)
assert isinstance(store, VectorStore)
def test_vector_contract_is_exported_from_public_packages():
from tht.adapters.vector import ThothHttpVectorStore as PublicHttpStore
from tht.ports import VectorStore as PublicVectorStore
from tht.ports import VectorWriteRecord as PublicVectorWriteRecord
from tht.ports import VectorReadUnavailable as PublicVectorReadUnavailable
assert PublicHttpStore is ThothHttpVectorStore
assert PublicVectorStore is VectorStore
assert PublicVectorWriteRecord is VectorWriteRecord
assert PublicVectorReadUnavailable is VectorReadUnavailable
capabilities = store_capabilities = ThothHttpVectorStore(
reader=MagicMock(), writer=None
).capabilities
assert capabilities.search is True
with pytest.raises(FrozenInstanceError):
store_capabilities.search = False
def test_http_health_uses_reader_list_tables_and_reports_failure():
reader = MagicMock()
store = ThothHttpVectorStore(reader=reader, writer=None)
assert store.health().ok is True
reader.list_tables.side_effect = RuntimeError("offline")
health = store.health()
assert health.ok is False
assert health.detail == "offline"
def test_http_health_reports_read_write_and_dimension_status_independently():
reader = MagicMock()
reader.list_tables.return_value = [
{"table_name": "memory", "vector_dimensions": 768}
]
writer = MagicMock()
writer.list_tables.return_value = [
{"table_name": "memory", "vector_dimensions": 768}
]
store = ThothHttpVectorStore(reader, writer, expected_dimension=768)
health = store.health()
assert health.ok is True
assert health.read_configured is True
assert health.read_reachable is True
assert health.write_configured is True
assert health.write_reachable is True
assert health.expected_dimension == 768
assert health.observed_dimensions == (768,)
assert health.dimension_compatible is True
def test_http_health_does_not_hide_writer_failure_behind_reader_success():
reader = MagicMock()
reader.list_tables.return_value = []
writer = MagicMock()
writer.list_tables.side_effect = RuntimeError("writer offline")
store = ThothHttpVectorStore(reader, writer, expected_dimension=768)
health = store.health()
assert health.ok is False
assert health.read_reachable is True
assert health.write_reachable is False
assert health.write_detail == "writer offline"
assert health.dimension_compatible is None
def test_http_health_covers_read_only_and_write_only_configuration():
reader = MagicMock()
reader.list_tables.return_value = [{"vector_dimensions": 384}]
read_health = ThothHttpVectorStore(reader, None, expected_dimension=768).health()
assert read_health.ok is False
assert read_health.write_configured is False
assert read_health.write_reachable is None
assert read_health.dimension_compatible is False
writer = MagicMock()
writer.list_tables.return_value = [{"vector_dimensions": 768}]
write_health = ThothHttpVectorStore(None, writer, expected_dimension=768).health()
assert write_health.ok is True
assert write_health.read_configured is False
assert write_health.read_reachable is None
assert write_health.dimension_compatible is True
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
def test_legacy_direct_search_requires_a_strict_positive_integer_limit(limit):
store = LegacyDirectVectorStore(engine=MagicMock())
with pytest.raises(ValueError, match="positive integer"):
store.search(["memory"], [0.1], limit=limit)
+6
View File
@@ -0,0 +1,6 @@
"""Data-warehouse adapter implementations."""
from tht.adapters.dwh.postgres import PostgresDwhAdapter
from tht.adapters.dwh.thoth_rest import ThothRestDwhAdapter
__all__ = ["PostgresDwhAdapter", "ThothRestDwhAdapter"]
+54
View File
@@ -0,0 +1,54 @@
"""Direct PostgreSQL implementation of the DWH port."""
from tht.config import DatabaseConfig
from sqlalchemy.exc import OperationalError, SQLAlchemyError
from tht.db import execute, sampling
from tht.db.connection import can_create_in_schema, make_engine, ping, writable_tables
from tht.db.introspect import introspect
from tht.execute import ExecResult, PlanSummary
from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import DistinctValues, DwhCapabilities, DwhHealth
class PostgresDwhAdapter:
capabilities = DwhCapabilities()
def __init__(self, config: DatabaseConfig, *, statement_timeout_ms: int = 30_000):
self._config = config
self._engine = make_engine(config)
self._statement_timeout_ms = statement_timeout_ms
def health(self) -> DwhHealth:
try:
ping(self._engine)
except OperationalError as exc:
return DwhHealth(ok=False, detail=str(exc.orig), error_kind="connection")
except SQLAlchemyError as exc:
return DwhHealth(ok=False, detail=str(exc), error_kind="connection")
writable = tuple(writable_tables(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,
read_only=not writable and not can_create,
writable_tables=writable, can_create=can_create)
def introspect(self) -> PhysicalSchema:
return introspect(self._engine, self._config.database, self._config.db_schema)
def run_query(self, sql: str, *, limit: int) -> ExecResult:
return execute.run_query(
self._engine, sql, limit=limit, timeout_ms=self._statement_timeout_ms
)
def explain(self, sql: str) -> PlanSummary:
return execute.explain(self._engine, sql, timeout_ms=self._statement_timeout_ms)
def sample_column(self, table: str, column: str, *, limit: int) -> list[object]:
return sampling.sample_column(
self._engine, self._config.db_schema, table, column, limit=limit
)
def distinct_values(self, table: str, column: str, *, limit: int) -> DistinctValues:
return sampling.distinct_values(
self._engine, self._config.db_schema, table, column, max_values=limit
)
+56
View File
@@ -0,0 +1,56 @@
"""Thoth/PostgREST implementation of the DWH port."""
from tht.config import DatabaseIdentityConfig, RestConfig
from tht.db.introspect import introspect_rest
from tht.db import sampling
from tht.execute import ExecResult, ExecutionError, PlanSummary
from tht.mschema.models import PhysicalSchema
from tht.ports.dwh import DistinctValues, DwhCapabilities, DwhHealth
from tht.rest.client import RestClient, RestError
from tht.rest.execute import explain_rest, run_controlled_rest
class ThothRestDwhAdapter:
capabilities = DwhCapabilities()
def __init__(self, database: DatabaseIdentityConfig, rest: RestConfig):
self._database = database
self._client = RestClient(rest)
def health(self) -> DwhHealth:
try:
result = self._client.ping()
except RestError as exc:
return DwhHealth(ok=False, detail=str(exc), error_kind="connection")
ok = bool(result.get("db_connected") and result.get("schema_accessible"))
return DwhHealth(ok=ok, detail=None if ok else str(result),
database=self._database.database, schema=self._database.db_schema,
endpoint=self._client.cfg.base_url, read_only=True,
error_kind=None if ok else "inaccessible")
def introspect(self) -> PhysicalSchema:
return introspect_rest(
self._client, self._database.database, self._database.db_schema
)
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]:
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, *, limit: int) -> DistinctValues:
try:
return sampling.distinct_values_rest(
self._client, self._database.db_schema, table, column, max_values=limit
)
except RestError as exc:
raise ExecutionError(str(exc)) from exc
@@ -0,0 +1,7 @@
"""Evidence source adapter implementations."""
from tht.adapters.evidence.filesystem import FilesystemEvidenceSource
from tht.adapters.evidence.http import HttpManifestEvidenceSource
from tht.adapters.evidence.s3 import S3EvidenceSource
__all__ = ["FilesystemEvidenceSource", "HttpManifestEvidenceSource", "S3EvidenceSource"]
+148
View File
@@ -0,0 +1,148 @@
"""Contained, race-safe filesystem Evidence source."""
import hashlib
import os
import stat
from datetime import UTC, datetime
from pathlib import Path, PurePosixPath
from urllib.parse import unquote, urlsplit
from tht.ports.evidence import (
AcquiredDocument,
EvidenceSourceError,
EvidenceSourceErrorCategory,
SourceObject,
)
class FilesystemEvidenceSource:
def __init__(
self,
root: Path | str,
*,
patterns: tuple[str, ...] | list[str] = ("**/*.md",),
max_bytes: int = 10 * 1024 * 1024,
) -> None:
if max_bytes < 1:
raise ValueError("max_bytes must be positive")
if not patterns or any(
not pattern or Path(pattern).is_absolute() or ".." in Path(pattern).parts
for pattern in patterns
):
raise ValueError("at least one non-empty discovery pattern is required")
try:
self.root = Path(root).expanduser().resolve(strict=True)
self._root_fd = os.open(
self.root,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW | os.O_CLOEXEC,
)
except OSError as error:
raise ValueError("filesystem evidence root is unavailable") from error
self.patterns = tuple(patterns)
self.max_bytes = max_bytes
def __del__(self):
root_fd = getattr(self, "_root_fd", None)
if root_fd is not None:
try:
os.close(root_fd)
except OSError:
pass
@staticmethod
def _safe_error(operation: str, *, transient: bool = False, **details):
return EvidenceSourceError(
"filesystem source operation failed",
category=(
EvidenceSourceErrorCategory.TRANSIENT
if transient
else EvidenceSourceErrorCategory.PERMANENT
),
details={"operation": operation, **details},
)
def _open_read(self, relative: PurePosixPath) -> tuple[bytes, os.stat_result]:
parts = relative.parts
if not parts or any(part in {"", ".", ".."} for part in parts):
raise self._safe_error("path_validation")
directory_fd = os.dup(self._root_fd)
file_fd = None
try:
for component in parts[:-1]:
next_fd = os.open(
component,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW | os.O_CLOEXEC,
dir_fd=directory_fd,
)
os.close(directory_fd)
directory_fd = next_fd
file_fd = os.open(
parts[-1],
os.O_RDONLY | os.O_NOFOLLOW | os.O_CLOEXEC,
dir_fd=directory_fd,
)
file_stat = os.fstat(file_fd)
if not stat.S_ISREG(file_stat.st_mode):
raise self._safe_error("path_validation")
if file_stat.st_size > self.max_bytes:
raise self._safe_error("read", limit_bytes=self.max_bytes)
content = bytearray()
while len(content) <= self.max_bytes:
chunk = os.read(file_fd, min(64 * 1024, self.max_bytes + 1 - len(content)))
if not chunk:
break
content.extend(chunk)
if len(content) > self.max_bytes:
raise self._safe_error("read", limit_bytes=self.max_bytes)
return bytes(content), file_stat
except EvidenceSourceError:
raise
except OSError as error:
raise self._safe_error("open") from error
finally:
if file_fd is not None:
os.close(file_fd)
os.close(directory_fd)
def _item(
self, relative: PurePosixPath, content: bytes, file_stat: os.stat_result
) -> SourceObject:
relative_text = relative.as_posix()
return SourceObject(
source_id=f"filesystem:{hashlib.sha256(relative_text.encode()).hexdigest()}",
uri=(self.root / relative_text).as_uri(),
fingerprint=f"sha256:{hashlib.sha256(content).hexdigest()}",
modified_at=datetime.fromtimestamp(file_stat.st_mtime, tz=UTC),
metadata={"relative_path": relative_text},
)
def discover(self):
candidates = {
path.relative_to(self.root).as_posix()
for pattern in self.patterns
for path in self.root.glob(pattern)
}
for relative_text in sorted(candidates):
relative = PurePosixPath(relative_text)
content, file_stat = self._open_read(relative)
yield self._item(relative, content, file_stat)
def acquire(self, item: SourceObject) -> AcquiredDocument:
parsed = urlsplit(item.uri)
if parsed.scheme != "file" or parsed.netloc or parsed.query or parsed.fragment:
raise self._safe_error("acquire")
try:
relative = Path(unquote(parsed.path)).relative_to(self.root)
except ValueError as error:
raise self._safe_error("acquire") from error
pure_relative = PurePosixPath(relative.as_posix())
content, file_stat = self._open_read(pure_relative)
expected = self._item(pure_relative, content, file_stat)
if item.source_id != expected.source_id or item.fingerprint != expected.fingerprint:
raise self._safe_error("acquire")
return AcquiredDocument(
source=expected,
content=content,
media_type="text/markdown" if relative.suffix.lower() == ".md" else None,
acquired_at=datetime.now(UTC),
)
+287
View File
@@ -0,0 +1,287 @@
"""Explicit-manifest HTTP Evidence source with SSRF-safe bounded acquisition."""
import hashlib
import ipaddress
import socket
from collections import OrderedDict
from datetime import UTC, datetime
from email.utils import parsedate_to_datetime
from urllib.parse import urljoin, urlsplit
import requests
from tht.ports.evidence import (
AcquiredDocument,
EvidenceSourceError,
EvidenceSourceErrorCategory,
SourceObject,
canonical_provenance_uri,
)
class HttpManifestEvidenceSource:
def __init__(
self,
urls: list[str] | tuple[str, ...],
*,
connect_timeout: float = 5,
read_timeout: float = 30,
max_bytes: int = 10 * 1024 * 1024,
max_redirects: int = 5,
allow_private_hosts: bool = False,
max_cache_bytes: int = 64 * 1024 * 1024,
) -> None:
if not urls:
raise ValueError("HTTP evidence manifest must contain at least one URL")
if (
connect_timeout <= 0
or read_timeout <= 0
or max_bytes < 1
or max_redirects < 0
or max_cache_bytes < 1
):
raise ValueError("HTTP evidence limits must be positive")
self._transport_by_uri: dict[str, str] = {}
for url in urls:
self._validate_url_shape(url)
provenance = canonical_provenance_uri(url)
if provenance in self._transport_by_uri:
raise ValueError("HTTP evidence manifest contains duplicate canonical provenance")
self._transport_by_uri[provenance] = url
self.connect_timeout = connect_timeout
self.read_timeout = read_timeout
self.max_bytes = max_bytes
self.max_redirects = max_redirects
self.allow_private_hosts = allow_private_hosts
self.max_cache_bytes = max_cache_bytes
self._session = requests.Session()
self._session.trust_env = False
self._cache: OrderedDict[str, AcquiredDocument] = OrderedDict()
# provenance -> (exact final effective URL, ETag, Last-Modified)
self._validators: dict[str, tuple[str, str | None, str | None]] = {}
self._cache_bytes = 0
def __repr__(self) -> str:
return f"HttpManifestEvidenceSource(objects={len(self._transport_by_uri)})"
@staticmethod
def _safe_error(operation: str, *, transient: bool = False, **details):
return EvidenceSourceError(
"HTTP source operation failed",
category=(
EvidenceSourceErrorCategory.TRANSIENT
if transient
else EvidenceSourceErrorCategory.PERMANENT
),
details={"operation": operation, **details},
)
@staticmethod
def _validate_url_shape(url: str) -> None:
parsed = urlsplit(url)
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
raise ValueError("HTTP evidence URLs must use http or https")
if parsed.username is not None or parsed.password is not None:
raise ValueError("HTTP evidence URLs must not contain userinfo credentials")
@staticmethod
def _source_id(uri: str) -> str:
return f"http:{hashlib.sha256(uri.encode()).hexdigest()}"
@staticmethod
def _normalized_ip(value: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address:
address = ipaddress.ip_address(value.split("%", 1)[0])
if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped:
return address.ipv4_mapped
return address
def _resolve_allowed(self, url: str) -> set[ipaddress.IPv4Address | ipaddress.IPv6Address]:
try:
self._validate_url_shape(url)
except ValueError as error:
raise self._safe_error("url_validation") from error
if self.allow_private_hosts:
return set()
parsed = urlsplit(url)
port = parsed.port or (443 if parsed.scheme == "https" else 80)
try:
rows = socket.getaddrinfo(parsed.hostname, port, type=socket.SOCK_STREAM)
addresses = {self._normalized_ip(row[4][0]) for row in rows}
except (OSError, ValueError) as error:
raise self._safe_error("resolution", transient=True) from error
if not addresses:
raise self._safe_error("resolution", transient=True)
# Reject the entire answer set if any address is private/reserved. Choosing only a public
# member would leave DNS ordering as a policy bypass.
if any(not address.is_global for address in addresses):
raise self._safe_error("network_policy")
return addresses
def _verify_peer(
self,
response,
allowed: set[ipaddress.IPv4Address | ipaddress.IPv6Address],
) -> None:
if self.allow_private_hosts:
return
try:
connection = response.raw._connection
peer = self._normalized_ip(connection.sock.getpeername()[0])
except (AttributeError, OSError, TypeError, ValueError) as error:
raise self._safe_error("peer_validation", transient=True) from error
if not peer.is_global or peer not in allowed:
raise self._safe_error("network_policy")
@staticmethod
def _status_category(status: int) -> EvidenceSourceErrorCategory:
if status in {408, 425, 429} or 500 <= status <= 599:
return EvidenceSourceErrorCategory.TRANSIENT
return EvidenceSourceErrorCategory.PERMANENT
def _conditional_headers(self, provenance: str, request_url: str) -> dict[str, str]:
cached = self._cache.get(self._source_id(provenance))
validators = self._validators.get(provenance)
if cached is None or validators is None:
return {}
final_url, etag, last_modified = validators
if request_url != final_url:
return {}
headers = {}
if etag:
headers["If-None-Match"] = etag
if last_modified:
headers["If-Modified-Since"] = last_modified
return headers
def _remember(
self,
provenance: str,
final_url: str,
document: AcquiredDocument,
validators: tuple[str | None, str | None],
) -> None:
source_id = document.source.source_id
old = self._cache.pop(source_id, None)
if old is not None:
self._cache_bytes -= len(old.content)
self._cache[source_id] = document
self._cache_bytes += len(document.content)
self._validators[provenance] = (final_url, *validators)
while self._cache and self._cache_bytes > self.max_cache_bytes:
evicted_id, evicted = self._cache.popitem(last=False)
self._cache_bytes -= len(evicted.content)
for uri in tuple(self._validators):
if self._source_id(uri) == evicted_id:
del self._validators[uri]
def _download(self, transport_url: str, provenance: str) -> AcquiredDocument:
current = transport_url
try:
for redirect_count in range(self.max_redirects + 1):
headers = self._conditional_headers(provenance, current)
allowed = self._resolve_allowed(current)
response = None
try:
response = self._session.get(
current,
headers=headers,
stream=True,
allow_redirects=False,
timeout=(self.connect_timeout, self.read_timeout),
)
self._verify_peer(response, allowed)
if response.is_redirect:
if redirect_count == self.max_redirects:
raise self._safe_error("redirect")
destination = urljoin(current, response.headers.get("Location", ""))
try:
self._validate_url_shape(destination)
except ValueError as error:
raise self._safe_error("redirect") from error
# The next iteration binds validators to the exact destination URL.
current = destination
continue
if response.status_code == 304:
cached = self._cache.get(self._source_id(provenance))
binding = self._validators.get(provenance)
if (
cached is None
or not headers
or binding is None
or binding[0] != current
):
raise self._safe_error("conditional_response")
self._cache.move_to_end(cached.source.source_id)
return cached
if not 200 <= response.status_code <= 299:
raise EvidenceSourceError(
"HTTP status failure",
category=self._status_category(response.status_code),
details={"operation": "download", "status": response.status_code},
)
length = response.headers.get("Content-Length")
if length is not None and int(length) > self.max_bytes:
raise self._safe_error("download", limit_bytes=self.max_bytes)
content = bytearray()
for chunk in response.iter_content(
chunk_size=min(64 * 1024, self.max_bytes + 1)
):
content.extend(chunk)
if len(content) > self.max_bytes:
raise self._safe_error("download", limit_bytes=self.max_bytes)
etag = response.headers.get("ETag")
last_modified = response.headers.get("Last-Modified")
media_type = (
response.headers.get("Content-Type", "").split(";", 1)[0] or None
)
break
finally:
if response is not None:
response.close()
except EvidenceSourceError:
raise
except (requests.Timeout, requests.ConnectionError, TimeoutError) as error:
raise self._safe_error("download", transient=True) from error
except requests.RequestException as error:
raise self._safe_error("download", transient=True) from error
except (OSError, ValueError) as error:
raise self._safe_error("download") from error
modified_at = None
if etag:
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
elif last_modified:
try:
modified_at = parsedate_to_datetime(last_modified).astimezone(UTC)
fingerprint = f"last-modified:{int(modified_at.timestamp())}"
except (TypeError, ValueError, OverflowError):
fingerprint = f"sha256:{hashlib.sha256(content).hexdigest()}"
else:
fingerprint = f"sha256:{hashlib.sha256(content).hexdigest()}"
item = SourceObject(
source_id=self._source_id(provenance),
uri=provenance,
fingerprint=fingerprint,
modified_at=modified_at,
)
document = AcquiredDocument(
source=item,
content=bytes(content),
media_type=media_type,
acquired_at=datetime.now(UTC),
)
self._remember(provenance, current, document, (etag, last_modified))
return document
def discover(self):
for provenance in sorted(self._transport_by_uri):
yield self._download(self._transport_by_uri[provenance], provenance).source
def acquire(self, item: SourceObject) -> AcquiredDocument:
transport = self._transport_by_uri.get(item.uri)
if transport is None or item.source_id != self._source_id(item.uri):
raise self._safe_error("acquire")
document = self._download(transport, item.uri)
if document.source.fingerprint != item.fingerprint:
raise self._safe_error("acquire")
return document
+152
View File
@@ -0,0 +1,152 @@
"""Bounded S3-compatible Evidence source using the supported boto3 client."""
import hashlib
import ipaddress
import re
from datetime import UTC, datetime
from urllib.parse import quote, urlsplit
from tht.ports.evidence import (
AcquiredDocument, EvidenceSourceError, EvidenceSourceErrorCategory, SourceObject,
)
class S3EvidenceSource:
def __init__(self, *, bucket: str, prefix: str = "", endpoint_url: str | None = None,
region: str | None = None, access_key: str | None = None,
secret_key: str | None = None, session_token: str | None = None,
trusted_endpoint: bool = False,
allow_private_endpoint: bool = False, allow_insecure_endpoint: bool = False,
max_bytes: int = 10 * 1024 * 1024, max_objects: int = 10_000,
max_pages: int = 100, page_size: int = 1000, client=None) -> None:
bucket_valid = re.fullmatch(r"(?=.{3,63}$)(?!-)(?!.*\.\.)(?!.*\.-)(?!.*-\.)"
r"[a-z0-9](?:[a-z0-9.-]*[a-z0-9])?", bucket)
try:
ipaddress.ip_address(bucket)
bucket_is_ip = True
except ValueError:
bucket_is_ip = False
if (not bucket_valid or bucket_is_ip
or any(value < 1 for value in (max_bytes, max_objects, max_pages, page_size))):
raise ValueError("S3 evidence limits and bucket must be non-empty and positive")
if (prefix.startswith("/") or len(prefix.encode()) > 1024
or any(ord(char) < 32 or ord(char) == 127 for char in prefix)):
raise ValueError("S3 prefix is invalid")
if endpoint_url:
parsed = urlsplit(endpoint_url)
if parsed.username or parsed.password:
raise ValueError("S3 endpoint must not contain credentials")
if parsed.scheme not in {"http", "https"}:
raise ValueError("S3 endpoint scheme must be exactly https or explicitly allowed http")
if parsed.scheme == "http" and not allow_insecure_endpoint:
raise ValueError("S3 endpoint must use HTTPS unless explicitly allowed")
if not parsed.hostname:
raise ValueError("S3 endpoint must include a hostname")
if parsed.path not in {"", "/"} or parsed.query or parsed.fragment:
raise ValueError("S3 custom endpoint must be an origin root without query/fragment")
if not trusted_endpoint:
raise ValueError("S3 custom endpoint requires explicit trusted_endpoint opt-in")
try:
literal = ipaddress.ip_address(parsed.hostname)
except ValueError:
literal = None
if literal is not None and not literal.is_global and not allow_private_endpoint:
raise ValueError("S3 private endpoint requires explicit opt-in")
self.bucket, self.prefix = bucket, prefix
self.max_bytes, self.max_objects = max_bytes, max_objects
self.max_pages, self.page_size = max_pages, min(page_size, 1000)
if client is None:
try:
import boto3
from botocore.config import Config as BotoConfig
except ImportError as exc: # pragma: no cover - deployment optional dependency
raise RuntimeError("Install tht[s3] to use S3 Evidence") from exc
client = boto3.client("s3", endpoint_url=endpoint_url, region_name=region,
aws_access_key_id=access_key,
aws_secret_access_key=secret_key,
aws_session_token=session_token, verify=True,
config=BotoConfig(s3={"addressing_style": "path"}))
self._client = client
self._items: dict[str, tuple[SourceObject, str]] = {}
@staticmethod
def _error(operation: str, transient: bool = False):
return EvidenceSourceError("S3 source operation failed",
category=(EvidenceSourceErrorCategory.TRANSIENT if transient
else EvidenceSourceErrorCategory.PERMANENT),
details={"operation": operation})
def discover(self):
count = pages = 0
try:
token = None
for _ in range(self.max_pages):
params = {"Bucket": self.bucket, "Prefix": self.prefix,
"MaxKeys": self.page_size}
if token is not None:
params["ContinuationToken"] = token
page = self._client.list_objects_v2(**params)
pages += 1
for row in page.get("Contents", []):
count += 1
if count > self.max_objects:
raise self._error("object_limit")
key, etag = row.get("Key"), row.get("ETag")
if (not isinstance(key, str) or not key or not key.startswith(self.prefix)
or len(key.encode()) > 1024
or any(ord(char) < 32 or ord(char) == 127 for char in key)):
raise self._error("invalid_key")
if not isinstance(etag, str) or not etag or len(etag) > 1024:
raise self._error("missing_validator")
uri = f"s3://{self.bucket}/{quote(key, safe='/')}"
fingerprint = f"etag:{hashlib.sha256(etag.encode()).hexdigest()}"
source_id = "s3:" + hashlib.sha256(uri.encode()).hexdigest()
modified = row.get("LastModified")
if modified is not None:
modified = modified.astimezone(UTC)
item = SourceObject(source_id=source_id, uri=uri, fingerprint=fingerprint,
modified_at=modified,
metadata={"size": int(row.get("Size", 0))})
self._items[source_id] = (item, etag)
yield item
if not page.get("IsTruncated"):
return
token = page.get("NextContinuationToken")
if not isinstance(token, str) or not token:
raise self._error("list_continuation")
raise self._error("list_limit")
except EvidenceSourceError:
raise
except Exception as exc:
raise self._error("list", transient=True) from exc
def acquire(self, item: SourceObject) -> AcquiredDocument:
binding = self._items.get(item.source_id)
if binding is None or item != binding[0]:
raise self._error("acquire")
discovered, etag = binding
key = discovered.uri.split(f"s3://{self.bucket}/", 1)[1]
from urllib.parse import unquote
key = unquote(key)
kwargs = {"Bucket": self.bucket, "Key": key}
body = None
try:
response = self._client.get_object(**kwargs)
body = response["Body"]
if response.get("ETag") != etag:
raise self._error("etag_changed")
if int(response.get("ContentLength", 0)) > self.max_bytes:
raise self._error("download_limit")
content = body.read(self.max_bytes + 1)
if len(content) > self.max_bytes:
raise self._error("download_limit")
return AcquiredDocument(source=item, content=content,
media_type=response.get("ContentType"),
acquired_at=datetime.now(UTC))
except EvidenceSourceError:
raise
except Exception as exc:
raise self._error("download", transient=True) from exc
finally:
if body is not None:
body.close()
+141
View File
@@ -0,0 +1,141 @@
"""Central construction of deployment-specific adapters."""
from tht.adapters.dwh import PostgresDwhAdapter, ThothRestDwhAdapter
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
from tht.adapters.evidence.s3 import S3EvidenceSource
from tht.adapters.vector import PgVectorStore, ThothHttpVectorStore
from tht.config import Config, ConfigError
from tht.db.connection import make_engine
from tht.ports.dwh import DwhAdapter
from tht.ports.vector import VectorStore
from tht.vectorstore.rest_client import VectorRestClient
def build_dwh(cfg: Config) -> DwhAdapter:
"""Build the DWH adapter selected by the validated workspace resource."""
resource = cfg.dwh
match resource.type:
case "postgres_direct":
return PostgresDwhAdapter(
resource.connection,
statement_timeout_ms=cfg.execution.statement_timeout_ms,
)
case "thoth_rest":
return ThothRestDwhAdapter(resource.database, resource.endpoint)
case other: # pragma: no cover - Pydantic's discriminator rejects this first.
raise ConfigError(f"Adapter DWH non supportato: {other}")
def build_vector_store(cfg: Config, *, require_write: bool = False) -> VectorStore:
"""Build the vector adapter, optionally requiring an HTTP writer credential."""
resource = cfg.vectors
if resource is None:
raise ConfigError("Risorsa vectors non configurata")
match resource.type:
case "pgvector_direct":
reader = resource.reader or resource.connection
if require_write and resource.writer is None:
raise ConfigError("Vector writer non configurato per pgvector_direct")
return PgVectorStore(
reader,
resource.writer,
expected_dimension=cfg.embeddings.dim if cfg.embeddings is not None else None,
)
case "thoth_vector_http":
if require_write and resource.writer is None:
raise ConfigError("Vector writer non configurato")
return ThothHttpVectorStore(
VectorRestClient(resource.reader) if resource.reader is not None else None,
VectorRestClient(resource.writer) if resource.writer is not None else None,
expected_dimension=cfg.embeddings.dim if cfg.embeddings is not None else None,
)
case other: # pragma: no cover - Pydantic's discriminator rejects this first.
raise ConfigError(f"Adapter vector non supportato: {other}")
def build_vector_loader(cfg: Config, collection: str):
"""Compatibility construction for legacy collection sync commands."""
resource = cfg.vectors
if resource is None:
raise ConfigError("Risorsa vectors non configurata")
if cfg.embeddings is None:
raise ConfigError("Embeddings non configurati")
if (
resource.type == "thoth_vector_http"
and resource.writer is not None
and (cfg.profile == "workstation" or resource.direct is None)
):
from tht.vectorstore.rest_writer import RestVectorWriter
return RestVectorWriter(VectorRestClient(resource.writer), table=collection)
connection = (
resource.writer or resource.connection
if resource.type == "pgvector_direct"
else resource.direct
)
if connection is None:
raise ConfigError("Vector writer non configurato")
from tht.vectorstore.store import VectorStore as TableVectorStore
return TableVectorStore(
make_engine(connection),
schema=connection.db_schema,
table=collection,
dim=cfg.embeddings.dim,
)
def build_evidence_sources(cfg: Config):
"""Build configured Evidence sources, including the legacy curated filesystem tree."""
evidence = cfg.evidence
if evidence is None:
return []
sources = []
if evidence.source_root is not None:
sources.append(FilesystemEvidenceSource(evidence.source_root / evidence.evidence_dir))
for resource in evidence.sources:
match resource.type:
case "filesystem":
sources.append(
FilesystemEvidenceSource(
resource.root,
patterns=resource.patterns,
max_bytes=resource.max_bytes,
)
)
case "http":
sources.append(
HttpManifestEvidenceSource(
[url.get_secret_value() for url in resource.urls],
connect_timeout=resource.connect_timeout,
read_timeout=resource.read_timeout,
max_bytes=resource.max_bytes,
max_redirects=resource.max_redirects,
allow_private_hosts=resource.allow_private_hosts,
max_cache_bytes=resource.max_cache_bytes,
)
)
case "s3":
def secret(value):
return value.get_secret_value() if value is not None else None
sources.append(S3EvidenceSource(
bucket=resource.bucket, prefix=resource.prefix,
endpoint_url=resource.endpoint_url, region=resource.region,
access_key=secret(resource.access_key), secret_key=secret(resource.secret_key),
session_token=secret(resource.session_token),
trusted_endpoint=resource.trusted_endpoint,
allow_private_endpoint=resource.allow_private_endpoint,
allow_insecure_endpoint=resource.allow_insecure_endpoint,
max_bytes=resource.max_bytes, max_objects=resource.max_objects,
max_pages=resource.max_pages, page_size=resource.page_size,
))
case other: # pragma: no cover - Pydantic rejects unsupported discriminators.
raise ConfigError(f"Adapter evidence non supportato: {other}")
return sources
__all__ = ["build_dwh", "build_evidence_sources", "build_vector_loader", "build_vector_store"]
+7
View File
@@ -0,0 +1,7 @@
"""Vector-store adapter implementations."""
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
from tht.adapters.vector.pgvector import PgVectorStore
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
__all__ = ["LegacyDirectVectorStore", "PgVectorStore", "ThothHttpVectorStore"]
@@ -0,0 +1,70 @@
"""Compatibility adapter for the existing direct PostgreSQL vector reader."""
from sqlalchemy import Engine
from tht.ports.vector import (
VectorCapabilities,
VectorHealth,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
require_positive_limit,
)
from tht.vectorstore.store import VectorHit, VectorStore as TableVectorStore
class LegacyDirectVectorStore:
"""Read-only port wrapper around the legacy table-scoped pgvector store."""
capabilities = VectorCapabilities(search=True, existing_hashes=False, upsert=False)
def __init__(self, engine: Engine, schema: str = "vectors", dim: int = 768):
self._engine = engine
self._schema = schema
self._dim = dim
def health(self) -> VectorHealth:
try:
with self._engine.connect() as connection:
connection.exec_driver_sql("SELECT 1")
except Exception as exc:
return VectorHealth(
ok=False,
detail=str(exc),
read_configured=True,
read_reachable=False,
read_detail=str(exc),
expected_dimension=self._dim,
)
return VectorHealth(
ok=True,
read_configured=True,
read_reachable=True,
expected_dimension=self._dim,
)
def search(
self,
collections: list[str],
embedding: list[float],
*,
limit: int,
kinds: list[str] | None = None,
metadata_filter: dict[str, object] | None = None,
) -> list[VectorHit]:
require_positive_limit(limit)
if metadata_filter is not None:
raise VectorStoreError("Legacy vector store cannot enforce metadata filtering")
hits: list[VectorHit] = []
for collection in collections:
table = TableVectorStore(
self._engine, schema=self._schema, table=collection, dim=self._dim
)
hits.extend(table.search(embedding, top_n=limit, kinds=kinds))
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[:limit]
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
raise VectorWriteUnavailable("Legacy direct reader has no writer interface")
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
raise VectorWriteUnavailable("Legacy direct reader has no writer interface")
+462
View File
@@ -0,0 +1,462 @@
"""Direct PostgreSQL/pgvector implementation of the vector port."""
import json
import re
from psycopg2 import sql
from sqlalchemy import Engine
from tht.config import DatabaseConfig
from tht.db.connection import make_engine
from tht.ports.vector import (
VectorCapabilities,
VectorHealth,
VectorReadUnavailable,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
require_positive_limit,
)
from tht.vectorstore.store import VectorHit, hit_from_metadata
COLLECTION_KINDS = {
"schema_records": {"schema_table", "schema_column"},
"evidence": {"evidence"},
"memory": {"memory", "solved_question"},
}
ALLOWED_COLLECTIONS = frozenset(COLLECTION_KINDS)
ALLOWED_KINDS = frozenset().union(*COLLECTION_KINDS.values())
_VECTOR_DIMENSION = re.compile(r"^(?:[a-z_][a-z0-9_]*\.)?vector\((\d+)\)$")
def _collection(schema: str, name: str) -> sql.Identifier:
if name not in ALLOWED_COLLECTIONS:
raise VectorStoreError(f"Collection not allowed: {name}")
return sql.Identifier(schema, name)
def _vector_literal(values: list[float]) -> str:
return "[" + ",".join(str(float(value)) for value in values) + "]"
def _vector_type(schema: str) -> sql.Identifier:
return sql.Identifier(schema, "vector")
def _cosine_operator(schema: str) -> sql.Composed:
return sql.SQL("OPERATOR({}.<=>)").format(sql.Identifier(schema))
def _validate_collection_kinds(collection: str, kinds: list[str]) -> None:
invalid = set(kinds) - COLLECTION_KINDS[collection]
if invalid:
raise VectorStoreError(f"Kind not allowed for {collection}: {', '.join(sorted(invalid))}")
def _validate_known_kinds(kinds: list[str]) -> None:
invalid = set(kinds) - ALLOWED_KINDS
if invalid:
raise VectorStoreError(f"Kind not allowed: {', '.join(sorted(invalid))}")
class PgVectorStore:
"""Direct store with independent reader and writer database credentials."""
def __init__(
self,
read_config: DatabaseConfig | None,
write_config: DatabaseConfig | None = None,
*,
expected_dimension: int | None = None,
):
self._reader = make_engine(read_config) if read_config is not None else None
self._writer = make_engine(write_config) if write_config is not None else None
config = read_config or write_config
self._schema = config.db_schema if config is not None else "vectors"
if read_config and write_config and read_config.db_schema != write_config.db_schema:
raise VectorStoreError("Reader and writer vector schemas must match")
self._expected_dimension = expected_dimension
@property
def capabilities(self) -> VectorCapabilities:
writable = self._writer is not None
return VectorCapabilities(
search=self._reader is not None,
existing_hashes=writable,
upsert=writable,
metadata_filter=self._reader is not None,
delete_generation=writable,
list_evidence_generations=writable,
)
def _probe(
self, engine: Engine | None, *, writable: bool
) -> tuple[bool | None, str | None, set[int]]:
if engine is None:
return None, None, set()
try:
raw = engine.raw_connection()
try:
with raw.cursor() as cursor:
cursor.execute("SELECT 1")
cursor.execute(
"SELECT has_schema_privilege(current_user, %s, 'USAGE')",
(self._schema,),
)
schema_usage = bool(cursor.fetchone()[0])
if not schema_usage:
return False, "vector schema incomplete: missing schema usage", set()
cursor.execute(
"""SELECT c.relname, format_type(a.atttypid, a.atttypmod),
has_table_privilege(current_user, c.oid, 'SELECT'),
has_table_privilege(current_user, c.oid, 'INSERT'),
has_table_privilege(current_user, c.oid, 'UPDATE'),
has_column_privilege(current_user, c.oid, 'record_key', 'SELECT')
AND has_column_privilege(
current_user, c.oid, 'content_hash', 'SELECT'
)
AND has_column_privilege(current_user, c.oid, 'kind', 'SELECT'),
CASE WHEN id_attr.attname IS NOT NULL THEN
pg_get_serial_sequence(
format('%%I.%%I', n.nspname, c.relname), 'id'
)
END AS id_sequence,
CASE WHEN id_attr.attname IS NOT NULL THEN
has_sequence_privilege(
current_user,
pg_get_serial_sequence(
format('%%I.%%I', n.nspname, c.relname), 'id'
),
'USAGE'
)
END AS sequence_usage
FROM pg_class c
JOIN pg_namespace n ON n.oid = c.relnamespace
LEFT JOIN pg_attribute a ON a.attrelid = c.oid
AND a.attname = 'embedding' AND NOT a.attisdropped
LEFT JOIN pg_attribute id_attr ON id_attr.attrelid = c.oid
AND id_attr.attname = 'id' AND NOT id_attr.attisdropped
WHERE n.nspname = %s AND c.relname = ANY(%s)
AND c.relkind IN ('r', 'p')""",
(self._schema, list(ALLOWED_COLLECTIONS)),
)
rows = cursor.fetchall()
present = {row[0] for row in rows}
missing_tables = sorted(ALLOWED_COLLECTIONS - present)
missing_embeddings = sorted(row[0] for row in rows if row[1] is None)
privilege_missing = sorted(
row[0]
for row in rows
if (writable and not (row[3] and row[4] and row[5]))
or (not writable and not row[2])
)
missing_sequences = sorted(
row[0] for row in rows if writable and row[6] is None
)
sequence_privilege_missing = sorted(
row[0] for row in rows if writable and row[6] is not None and not row[7]
)
problems = []
if missing_tables:
problems.append("missing tables " + ", ".join(missing_tables))
if missing_embeddings:
problems.append(
"missing embedding columns " + ", ".join(missing_embeddings)
)
if privilege_missing:
authority = "write" if writable else "read"
problems.append(
f"missing {authority} privileges " + ", ".join(privilege_missing)
)
if missing_sequences:
problems.append("missing id sequences " + ", ".join(missing_sequences))
if sequence_privilege_missing:
problems.append(
"missing sequence privileges " + ", ".join(sequence_privilege_missing)
)
if problems:
return False, "vector schema incomplete: " + "; ".join(problems), set()
dimensions = {
int(match.group(1))
for _, type_name, *_ in rows
if (match := _VECTOR_DIMENSION.match(type_name))
}
invalid_types = sorted(
row[0]
for row in rows
if row[1] is not None and not _VECTOR_DIMENSION.match(row[1])
)
if invalid_types:
return (
False,
"vector schema incomplete: invalid embedding types "
+ ", ".join(invalid_types),
set(),
)
if self._expected_dimension is not None:
mismatches = sorted(
f"{name}={int(match.group(1))}"
for name, type_name, *_ in rows
if (match := _VECTOR_DIMENSION.match(type_name))
and int(match.group(1)) != self._expected_dimension
)
if mismatches:
return (
False,
"embedding dimension mismatch: " + ", ".join(mismatches),
dimensions,
)
return True, None, dimensions
finally:
raw.close()
except Exception as exc:
return False, f"vector database probe failed: {type(exc).__name__}", set()
def health(self) -> VectorHealth:
read_ok, read_detail, read_dimensions = self._probe(self._reader, writable=False)
write_ok, write_detail, write_dimensions = self._probe(self._writer, writable=True)
dimensions = tuple(sorted(read_dimensions | write_dimensions))
compatible = (
None
if self._expected_dimension is None or not dimensions
else dimensions == (self._expected_dimension,)
)
reachable = [value for value in (read_ok, write_ok) if value is not None]
details = [value for value in (read_detail, write_detail) if value]
return VectorHealth(
ok=bool(reachable) and all(reachable) and compatible is not False,
detail="; ".join(details) or None,
read_configured=self._reader is not None,
read_reachable=read_ok,
read_detail=read_detail,
write_configured=self._writer is not None,
write_reachable=write_ok,
write_detail=write_detail,
expected_dimension=self._expected_dimension,
observed_dimensions=dimensions,
dimension_compatible=compatible,
)
def search(
self,
collections: list[str],
embedding: list[float],
*,
limit: int,
kinds: list[str] | None = None,
metadata_filter: dict[str, object] | None = None,
) -> list[VectorHit]:
require_positive_limit(limit)
if self._reader is None:
raise VectorReadUnavailable("Vector reader credential is not configured")
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
raise VectorStoreError("Query embedding dimension does not match configured dimension")
if kinds:
_validate_known_kinds(kinds)
hits: list[VectorHit] = []
raw = None
try:
raw = self._reader.raw_connection()
with raw.cursor() as cursor:
for collection in collections:
table = _collection(self._schema, collection)
collection_kinds = (
sorted(set(kinds) & COLLECTION_KINDS[collection]) if kinds else None
)
if kinds and not collection_kinds:
continue
clauses = []
filter_params = []
if collection_kinds:
clauses.append(sql.SQL("kind = ANY(%s)"))
filter_params.append(collection_kinds)
if metadata_filter is not None:
if collection != "evidence" or set(metadata_filter) != {
"vector_generation", "document_ids", "workspace_id"
}:
raise VectorStoreError("Unsupported vector metadata filter")
generation = metadata_filter["vector_generation"]
document_ids = metadata_filter["document_ids"]
workspace_id = metadata_filter["workspace_id"]
if not isinstance(generation, str) or not isinstance(document_ids, list) or not isinstance(workspace_id, str):
raise VectorStoreError("Invalid vector metadata filter")
clauses.append(sql.SQL("metadata->>'vector_generation' = %s"))
clauses.append(sql.SQL("metadata->>'document_id' = ANY(%s)"))
clauses.append(sql.SQL("metadata->>'workspace_id' = %s"))
filter_params.extend((generation, document_ids, workspace_id))
where = (
sql.SQL(" WHERE ") + sql.SQL(" AND ").join(clauses)
if clauses else sql.SQL("")
)
query = sql.SQL(
"SELECT metadata, 1 - (embedding {} %s::{}) AS similarity "
"FROM {}{} ORDER BY embedding {} %s::{}, record_key LIMIT %s"
).format(
_cosine_operator(self._schema),
_vector_type(self._schema),
table,
where,
_cosine_operator(self._schema),
_vector_type(self._schema),
)
params = [_vector_literal(embedding)]
params.extend(filter_params)
params.extend((_vector_literal(embedding), limit))
cursor.execute(query, params)
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
except VectorStoreError:
raise
except Exception as exc:
raise VectorReadUnavailable("Vector read operation unavailable") from exc
finally:
if raw is not None:
raw.close()
return sorted(hits, key=lambda hit: (-hit.similarity, hit.id))[:limit]
def _require_writer(self) -> Engine:
if self._writer is None:
raise VectorWriteUnavailable("Vector writer credential is not configured")
return self._writer
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
engine = self._require_writer()
table = _collection(self._schema, collection)
_validate_collection_kinds(collection, kinds)
raw = None
try:
raw = engine.raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL("SELECT record_key, content_hash FROM {} WHERE kind = ANY(%s)").format(
table
),
(kinds,),
)
return dict(cursor.fetchall())
except VectorStoreError:
raise
except Exception as exc:
raise VectorWriteUnavailable("Vector write operation unavailable") from exc
finally:
if raw is not None:
raw.close()
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
engine = self._require_writer()
table = _collection(self._schema, collection)
for write_record in records:
_validate_collection_kinds(collection, [write_record.record.kind])
if (
self._expected_dimension is not None
and len(write_record.embedding) != self._expected_dimension
):
raise VectorStoreError("Embedding dimension does not match configured dimension")
insert = sql.SQL(
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
"VALUES (%s, %s, %s, %s::jsonb, %s::{}) "
"ON CONFLICT (record_key) DO NOTHING"
).format(table, _vector_type(self._schema))
update = sql.SQL(
"UPDATE {} SET kind = %s, content_hash = %s, metadata = %s::jsonb, "
"embedding = %s::{}, indexed_at = pg_catalog.now() WHERE record_key = %s"
).format(table, _vector_type(self._schema))
raw = None
try:
raw = engine.raw_connection()
with raw.cursor() as cursor:
for write_record in records:
record = write_record.record
metadata = {
"kind": record.kind,
"ref": record.ref,
"record_key": record.id,
"title": record.title,
"content": record.content,
**record.metadata,
}
metadata_json = json.dumps(metadata)
vector = _vector_literal(write_record.embedding)
cursor.execute(
insert,
(record.id, record.kind, write_record.content_hash, metadata_json, vector),
)
if cursor.rowcount == 0:
cursor.execute(
update,
(
record.kind,
write_record.content_hash,
metadata_json,
vector,
record.id,
),
)
raw.commit()
except VectorStoreError:
if raw is not None:
raw.rollback()
raise
except Exception as exc:
if raw is not None:
raw.rollback()
raise VectorWriteUnavailable("Vector write operation unavailable") from exc
finally:
if raw is not None:
raw.close()
return len(records)
def delete_generation(self, collection: str, generation: str, workspace_id: str) -> int:
if collection != "evidence" or re.fullmatch(r"gen:[0-9a-f]{32}", generation) is None:
raise VectorStoreError("Only exact Evidence generations may be deleted")
if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None:
raise VectorStoreError("Invalid Evidence workspace namespace")
raw = None
try:
raw = self._require_writer().raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL(
"DELETE FROM {} WHERE kind = 'evidence' "
"AND metadata->>'vector_generation' = %s "
"AND metadata->>'workspace_id' = %s"
).format(_collection(self._schema, collection)),
(generation, workspace_id),
)
count = cursor.rowcount
raw.commit()
return count
except Exception as exc:
if raw is not None:
raw.rollback()
raise VectorWriteUnavailable("Vector generation cleanup unavailable") from exc
finally:
if raw is not None:
raw.close()
def list_evidence_generations(self, collection: str, workspace_id: str) -> list[str]:
if collection != "evidence":
raise VectorStoreError("Only exact Evidence generations may be listed")
if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None:
raise VectorStoreError("Invalid Evidence workspace namespace")
raw = None
try:
raw = self._require_writer().raw_connection()
with raw.cursor() as cursor:
cursor.execute(
sql.SQL(
"SELECT DISTINCT metadata->>'vector_generation' FROM {} "
"WHERE kind = 'evidence' AND metadata->>'vector_generation' "
"~ '^gen:[0-9a-f]{{32}}$' AND metadata->>'workspace_id' = %s ORDER BY 1"
).format(_collection(self._schema, collection)),
(workspace_id,),
)
return [row[0] for row in cursor.fetchall()]
except Exception as exc:
raise VectorWriteUnavailable("Vector generation inventory unavailable") from exc
finally:
if raw is not None:
raw.close()
__all__ = ["ALLOWED_COLLECTIONS", "PgVectorStore"]
+191
View File
@@ -0,0 +1,191 @@
"""Thoth vector HTTP adapter using distinct read and write clients."""
import re
from tht.ports.vector import (
VectorCapabilities,
VectorHealth,
VectorHit,
VectorReadUnavailable,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
require_positive_limit,
)
from tht.vectorstore.rest_client import VectorRestClient, VectorRestError
from tht.vectorstore.store import hit_from_metadata
from tht.adapters.vector.pgvector import (
_collection,
_validate_collection_kinds,
_validate_known_kinds,
)
def _merge(hits: list[VectorHit], limit: int) -> list[VectorHit]:
return sorted(hits, key=lambda hit: (-hit.similarity, hit.id))[:limit]
class ThothHttpVectorStore:
"""Vector port backed by the existing allowlisted REST RPCs."""
def __init__(
self,
reader: VectorRestClient | None,
writer: VectorRestClient | None,
expected_dimension: int | None = None,
):
self._reader = reader
self._writer = writer
self._expected_dimension = expected_dimension
@property
def capabilities(self) -> VectorCapabilities:
writable = self._writer is not None
return VectorCapabilities(
search=self._reader is not None, existing_hashes=writable, upsert=writable,
metadata_filter=self._reader is not None, delete_generation=writable,
list_evidence_generations=writable,
)
def health(self) -> VectorHealth:
read_reachable, read_detail, read_tables = self._probe(self._reader)
write_reachable, write_detail, write_tables = self._probe(self._writer)
dimensions = tuple(sorted({
dimension
for row in [*read_tables, *write_tables]
if type(dimension := row.get("vector_dimensions")) is int
}))
compatible = (
None
if self._expected_dimension is None or not dimensions
else dimensions == (self._expected_dimension,)
)
reachable = [
status for status in (read_reachable, write_reachable) if status is not None
]
ok = bool(reachable) and all(reachable) and compatible is not False
details = [detail for detail in (read_detail, write_detail) if detail]
return VectorHealth(
ok=ok,
detail="; ".join(details) or None,
read_configured=self._reader is not None,
read_reachable=read_reachable,
read_detail=read_detail,
write_configured=self._writer is not None,
write_reachable=write_reachable,
write_detail=write_detail,
expected_dimension=self._expected_dimension,
observed_dimensions=dimensions,
dimension_compatible=compatible,
)
@staticmethod
def _probe(client: VectorRestClient | None) -> tuple[bool | None, str | None, list[dict]]:
if client is None:
return None, None, []
try:
return True, None, client.list_tables()
except Exception as exc:
return False, str(exc), []
def search(
self,
collections: list[str],
embedding: list[float],
*,
limit: int,
kinds: list[str] | None = None,
metadata_filter: dict[str, object] | None = None,
) -> list[VectorHit]:
require_positive_limit(limit)
if self._reader is None:
raise VectorReadUnavailable("Vector reader credential is not configured")
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
raise VectorStoreError("Query embedding dimension does not match configured dimension")
if kinds:
_validate_known_kinds(kinds)
hits: list[VectorHit] = []
for collection in collections:
_collection("vectors", collection)
try:
if metadata_filter is None:
rows = self._reader.search_similar(collection, embedding, limit, kinds=kinds)
else:
rows = self._reader.search_similar(
collection, embedding, limit, kinds=kinds,
metadata_filter=metadata_filter,
)
except VectorRestError as exc:
raise VectorStoreError(str(exc)) from exc
hits.extend(
hit_from_metadata(row.get("similarity", 0.0), row.get("metadata"))
for row in rows
)
if kinds:
allowed = set(kinds)
hits = [hit for hit in hits if hit.kind in allowed]
return _merge(hits, limit)
def _require_writer(self) -> VectorRestClient:
if self._writer is None:
raise VectorWriteUnavailable("Vector writer credential is not configured")
return self._writer
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
_collection("vectors", collection)
_validate_collection_kinds(collection, kinds)
try:
return self._require_writer().existing_hashes(collection, kinds)
except VectorRestError as exc:
raise VectorStoreError(str(exc)) from exc
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
writer = self._require_writer()
_collection("vectors", collection)
for record in records:
_validate_collection_kinds(collection, [record.record.kind])
if (
self._expected_dimension is not None
and len(record.embedding) != self._expected_dimension
):
raise VectorStoreError("Embedding dimension does not match configured dimension")
rows = [self._row(record) for record in records]
try:
return writer.upsert_records(collection, rows)
except VectorRestError as exc:
raise VectorStoreError(str(exc)) from exc
def delete_generation(self, collection: str, generation: str, workspace_id: str) -> int:
if collection != "evidence" or re.fullmatch(r"gen:[0-9a-f]{32}", generation) is None:
raise VectorStoreError("Only exact Evidence generations may be deleted")
try:
return self._require_writer().delete_generation(collection, generation, workspace_id)
except VectorRestError as exc:
raise VectorStoreError(str(exc)) from exc
def list_evidence_generations(self, collection: str, workspace_id: str) -> list[str]:
if collection != "evidence":
raise VectorStoreError("Only exact Evidence generations may be listed")
try:
return self._require_writer().list_evidence_generations(collection, workspace_id)
except VectorRestError as exc:
raise VectorWriteUnavailable("Vector generation inventory unavailable") from exc
@staticmethod
def _row(write_record: VectorWriteRecord) -> dict:
record = write_record.record
metadata = {
"kind": record.kind,
"ref": record.ref,
"record_key": record.id,
"title": record.title,
"content": record.content,
**record.metadata,
}
return {
"record_key": record.id,
"kind": record.kind,
"content_hash": write_record.content_hash,
"metadata": metadata,
"embedding": write_record.embedding,
}
+5
View File
@@ -41,19 +41,23 @@ from tht.cli.cte_cmd import cte_app # noqa: E402
from tht.cli.datamart_cmd import datamart_app # noqa: E402
from tht.cli.db_cmd import db_app # noqa: E402
from tht.cli.decision_cmd import decision_app # noqa: E402
from tht.cli.doctor_cmd import doctor # noqa: E402
from tht.cli.evidence_cmd import evidence_app # noqa: E402
from tht.cli.formula_cmd import formula_app # noqa: E402
from tht.cli.lsh_cmd import lsh_app # noqa: E402
from tht.cli.memory_cmd import memory_app # noqa: E402
from tht.cli.ollama_cmd import ollama_app # noqa: E402
from tht.cli.phase_cmd import phase_app # noqa: E402
from tht.cli.preprocess_cmd import preprocess_app # noqa: E402
from tht.cli.schema_cmd import schema_app # noqa: E402
from tht.cli.search_cmd import search_app # noqa: E402
from tht.cli.session_cmd import session_app # noqa: E402
from tht.cli.sql_cmd import sql_app # noqa: E402
from tht.cli.vector_cmd import vector_app # noqa: E402
import tht.cli.vector_migrate_cmd # noqa: E402, F401
app.add_typer(phase_app, name="phase")
app.add_typer(preprocess_app, name="preprocess")
app.add_typer(config_app, name="config")
app.add_typer(schema_app, name="schema")
app.add_typer(session_app, name="session")
@@ -69,3 +73,4 @@ app.add_typer(cte_app, name="cte")
app.add_typer(datamart_app, name="datamart")
app.add_typer(lsh_app, name="lsh")
app.add_typer(ollama_app, name="ollama")
app.command("doctor")(doctor)
+22 -45
View File
@@ -1,40 +1,14 @@
from pathlib import Path
import typer
from sqlalchemy.exc import OperationalError
from tht.adapters.factory import build_dwh
from tht.cli.config_cmd import CONFIG_OPT
from tht.config import ConfigError, load_config
from tht.db.connection import can_create_in_schema, make_engine, ping, writable_tables
from tht.db.fetch_ca import CaFetchError, describe_pem, fetch_chain_pem, parse_host_port
db_app = typer.Typer(help="Operazioni sul database target")
def _ping_rest(cfg, schema: str) -> None:
"""Health check via REST. Il read-only è garantito strutturalmente dall'API
(ammette solo SELECT/WITH): non serve il controllo dei privilegi di scrittura."""
from tht.rest.client import RestClient, RestError
try:
info = RestClient(cfg.rest).ping()
except RestError as e:
typer.secho(f"ERRORE di connessione: {e}", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
if not info.get("db_connected") or not info.get("schema_accessible"):
typer.secho(
f"ERRORE: DWH non accessibile via REST (risposta: {info}).",
fg=typer.colors.RED, err=True,
)
raise typer.Exit(code=1)
typer.secho(
f"OK: connesso via REST a {cfg.rest.base_url} (schema {schema})", fg=typer.colors.GREEN
)
typer.secho(
"OK: accesso read-only garantito dall'API (solo SELECT/WITH).", fg=typer.colors.GREEN
)
@db_app.command("ping")
def ping_cmd(config: Path = CONFIG_OPT) -> None:
"""Testa la connessione e verifica che l'utente sia effettivamente read-only."""
@@ -43,28 +17,31 @@ def ping_cmd(config: Path = CONFIG_OPT) -> None:
except ConfigError as e:
typer.secho(f"ERRORE: {e}", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
schema = cfg.database.db_schema
if cfg.database.transport == "rest":
_ping_rest(cfg, schema)
return
engine = make_engine(cfg.database)
try:
ping(engine)
except OperationalError as e:
typer.secho(f"ERRORE di connessione: {e.orig}", fg=typer.colors.RED, err=True)
health = build_dwh(cfg).health()
if not health.ok:
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)
typer.secho(f"OK: connesso a {cfg.database.database} (schema {schema})", fg=typer.colors.GREEN)
writable = writable_tables(engine, schema)
can_create = can_create_in_schema(engine, schema)
if writable or can_create:
if health.endpoint:
typer.secho(f"OK: connesso via REST a {health.endpoint} (schema {health.schema})",
fg=typer.colors.GREEN)
typer.secho("OK: accesso read-only garantito dall'API (solo SELECT/WITH).",
fg=typer.colors.GREEN)
return
typer.secho(f"OK: connesso a {health.database} (schema {health.schema})",
fg=typer.colors.GREEN)
if not health.read_only:
typer.secho(
f"ERRORE: l'utente '{cfg.database.user}' NON e' read-only.", fg=typer.colors.RED, err=True
)
if writable:
typer.echo(f" Tabelle scrivibili: {', '.join(writable[:10])}", err=True)
if can_create:
typer.echo(f" L'utente puo' creare oggetti nello schema {schema}.", err=True)
if health.writable_tables:
typer.echo(f" Tabelle scrivibili: {', '.join(health.writable_tables[:10])}", err=True)
if health.can_create:
typer.echo(f" L'utente puo' creare oggetti nello schema {health.schema}.", err=True)
typer.echo(" Crea un ruolo read-only con scripts/create_readonly_role.sql.", err=True)
raise typer.Exit(code=2)
typer.secho("OK: l'utente e' read-only sullo schema target.", fg=typer.colors.GREEN)
+62
View File
@@ -0,0 +1,62 @@
from __future__ import annotations
import json
import os
from pathlib import Path
from typing import Any
import typer
import yaml
from tht.cli.config_cmd import CONFIG_OPT
from tht.config import ConfigError, load_config
def _emit(payload: dict[str, Any], as_json: bool) -> None:
if as_json:
# A single serializer call keeps stdout valid for machine consumers.
typer.echo(json.dumps(payload, sort_keys=True))
return
for component, result in payload["components"].items():
detail = ""
if result["status"] == "error":
detail = f" - {result['message']}"
elif component == "data_root" and result["status"] == "warning":
detail = " - set THT_DATA_ROOT to enable portable storage"
elif component == "workspace_paths" and result["status"] == "warning":
detail = " - absolute legacy roots: " + ", ".join(result["legacy_absolute"])
typer.echo(f"{component}: {result['status']}{detail}")
def doctor(
config: Path = CONFIG_OPT,
as_json: bool = typer.Option(False, "--json", help="Emette diagnostica JSON."),
) -> None:
"""Validate portable storage configuration without contacting external services."""
data_root = os.environ.get("THT_DATA_ROOT")
components: dict[str, dict[str, Any]] = {
"config": {"status": "ok"},
"data_root": {"status": "ok" if data_root else "warning"},
}
try:
cfg = load_config(config)
except (ConfigError, yaml.YAMLError, OSError) as exc:
path_error = isinstance(exc, ConfigError) and "outside workspace" in str(exc)
target = "workspace_paths" if path_error else "config"
# Validation errors can contain Pydantic input excerpts, including credentials.
message = str(exc) if path_error else "configuration is invalid or unreadable"
components[target] = {"status": "error", "message": message}
payload = {"ok": False, "components": components}
_emit(payload, as_json)
raise typer.Exit(code=1)
legacy_absolute = [
name
for name in ("sessions", "artifacts", "indexes")
if getattr(cfg.roots, name).is_absolute()
]
components["workspace_paths"] = {
"status": "warning" if legacy_absolute else "ok",
"legacy_absolute": legacy_absolute,
}
_emit({"ok": True, "components": components}, as_json)
+77 -33
View File
@@ -9,45 +9,91 @@ lsh_app = typer.Typer(help="Indice LSH su valori dei campi (derivato, rigenerabi
def _lsh_dir(cfg) -> Path:
return cfg.paths.indexes / "lsh"
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
return resolve_dwh_snapshot(cfg).lsh_dir
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
def build_lsh_artifacts(
cfg, *, dwh=None, verbose: bool = False, physical_file: Path | None = None,
output_dir: Path | None = None,
):
"""Run the existing LSH extraction/build algorithm and persist its outputs."""
from tht.adapters.factory import build_dwh
from tht.cli.schema_cmd import annotations_path
from tht.lshindex import build_index, save_index
from tht.mschema.models import Annotations, PhysicalSchema
phys_file = physical_file or physical_path(cfg)
if not phys_file.exists():
raise FileNotFoundError("physical catalog is missing; run schema introspect first")
physical = PhysicalSchema.from_yaml(phys_file)
annotations = Annotations.from_yaml(annotations_path(cfg))
target = dwh if dwh is not None else build_dwh(cfg)
values, skipped, truncated = _extract_lsh_values(
target, physical, annotations, cfg.lsh.max_values_per_column
)
lsh, minhashes = build_index(values, cfg.lsh, verbose=verbose)
save_index(
lsh, minhashes, cfg.lsh, output_dir or (cfg.paths.indexes / "lsh"),
name=cfg.database.db_schema,
)
return minhashes, skipped, truncated, values
@lsh_app.command("build")
def build_cmd(config: Path = CONFIG_OPT) -> None:
"""Costruisce l'indice LSH dai valori del database e lo salva su pickle."""
from tht.lshindex import build_index, save_index
from tht.mschema.models import Annotations, PhysicalSchema
cfg = _load_config_or_exit(config)
phys_file = physical_path(cfg)
if not phys_file.exists():
typer.secho(
f"ERRORE: {phys_file} non trovato. Esegui prima `tht schema introspect`.",
fg=typer.colors.RED, err=True,
)
raise typer.Exit(code=1)
physical = PhysicalSchema.from_yaml(phys_file)
from tht.cli.schema_cmd import annotations_path
annotations = Annotations.from_yaml(annotations_path(cfg))
dwh_root = cfg.paths.artifacts.parent / ".tht-dwh"
initialized = dwh_root.exists() or dwh_root.is_symlink()
if initialized:
phys_file = physical_path(cfg)
if not phys_file.exists():
typer.secho(
f"ERRORE: {phys_file} non trovato. Esegui prima `tht schema introspect`.",
fg=typer.colors.RED, err=True,
)
raise typer.Exit(code=1)
typer.echo("Estrazione valori (i più frequenti) dalle colonne testuali eligible...")
if cfg.database.transport == "rest":
from tht.db.sampling import unique_values_for_lsh_rest
from tht.rest.client import RestClient
from tht.cli.preprocess_cmd import run_dwh_from_config
from tht.lshindex import load_index
values, skipped, truncated = unique_values_for_lsh_rest(
RestClient(cfg.rest), physical, cfg.lsh, annotations
)
else:
from tht.db.connection import make_engine
from tht.db.sampling import unique_values_for_lsh
values, skipped, truncated = unique_values_for_lsh(
make_engine(cfg.database), physical, cfg.lsh, annotations
)
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")
report = run_dwh_from_config(
config, steps=("lsh",) if initialized else ("introspect", "lsh")
)
if report.status != "succeeded":
typer.secho("ERRORE: DWH preprocessing failed", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
_, minhashes, _ = load_index(_lsh_dir(cfg), name=cfg.database.db_schema)
skipped, truncated = [], []
n_values = len(minhashes)
n_columns = len({(entry[1], entry[2]) for entry in minhashes.values()})
typer.echo(f" {n_values} valori da {n_columns} colonne")
for s in skipped:
typer.secho(f" saltata {s.table}.{s.column}: {s.reason}", fg=typer.colors.YELLOW)
for t in truncated:
@@ -57,8 +103,6 @@ def build_cmd(config: Path = CONFIG_OPT) -> None:
fg=typer.colors.YELLOW,
)
lsh, minhashes = build_index(values, cfg.lsh, verbose=True)
save_index(lsh, minhashes, cfg.lsh, _lsh_dir(cfg), name=cfg.database.db_schema)
typer.secho(
f"OK: indice LSH ({len(minhashes)} entry) -> {_lsh_dir(cfg)}", fg=typer.colors.GREEN
)
+6 -5
View File
@@ -153,9 +153,9 @@ def save_one_cmd(
"""
import json as _json
from tht.adapters.factory import build_vector_store
from tht.cli.vector_cmd import make_embedder
from tht.memory import load_registry, promote, save_one_memory
from tht.vectorstore.rest_client import VectorRestClient
cfg = _load_config_or_exit(config)
manifest = load_session_or_exit(cfg, session)
@@ -167,6 +167,7 @@ def save_one_cmd(
fg=typer.colors.RED, err=True,
)
raise typer.Exit(code=4)
store = build_vector_store(cfg, require_write=True)
sdir = session_dir(cfg, session)
# Promuove la decisione scelta nel registro locale (idempotente: salta se gia' presente
@@ -174,9 +175,8 @@ def save_one_cmd(
promote(sdir, manifest, seqs=[decision], registry_path=registry_path(cfg))
records = [r for r in load_registry(registry_path(cfg)) if r.session_id == manifest.id]
writer = VectorRestClient(cfg.vector_write_rest)
embedder = make_embedder(cfg.embeddings)
count = save_one_memory(records, decision, writer=writer, embedder=embedder)
count = save_one_memory(records, decision, store=store, embedder=embedder)
msg = (
f"{count} memoria salvata su pgvector (decision_seq {decision})."
@@ -449,23 +449,24 @@ def index_solved_session(cfg, session_id: str) -> int:
Solleva RuntimeError se manca la writer key e SolvedIndexError se mancano gli
artefatti: il finalize li degrada a warning, il comando CLI li converte in
errori espliciti."""
from tht.adapters.factory import build_vector_store
from tht.cli.sql_cmd import promoted_tables_for
from tht.cli.vector_cmd import make_embedder
from tht.solved import build_solved_record, save_solved_question
from tht.vectorstore.rest_client import VectorRestClient
if not has_vector_write_rest(cfg):
raise RuntimeError(
"vector_write_rest assente: la coppia domanda->SQL si indicizza solo con la "
"writer key configurata nel workspace yaml"
)
store = build_vector_store(cfg, require_write=True)
manifest = load_session_or_exit(cfg, session_id)
record = build_solved_record(
session_dir(cfg, session_id), manifest, promoted_tables_for(cfg, session_id)
)
return save_solved_question(
record,
writer=VectorRestClient(cfg.vector_write_rest),
store=store,
embedder=make_embedder(cfg.embeddings),
)
+222
View File
@@ -0,0 +1,222 @@
"""One-shot preprocessing commands."""
from __future__ import annotations
import json
import re
import hashlib
from pathlib import Path
import typer
from tht.cli.config_cmd import CONFIG_OPT
preprocess_app = typer.Typer(help="Materialize versioned preprocessing artifacts")
def run_dwh_from_config(
config: Path, *, steps: tuple[str, ...], resume: str | None = None,
):
from tht.cli.lsh_cmd import build_lsh_artifacts
from tht.cli.schema_cmd import _load_config_or_exit, refresh_catalog
from tht.jobs.dwh_pipeline import (
DwhPreprocessPipeline, config_dwh_binding,
)
cfg = _load_config_or_exit(config)
binding = config_dwh_binding(cfg)
workspace_root = cfg.paths.artifacts.parent
lsh_names = (
f"{cfg.database.db_schema}_lsh.pkl",
f"{cfg.database.db_schema}_minhashes.pkl",
f"{cfg.database.db_schema}_meta.json",
)
pipeline = DwhPreprocessPipeline(
workspace_id=binding["workspace_id"],
workspace_root=workspace_root,
config_fingerprint=binding["config_fingerprint"],
input_fingerprint=binding["input_fingerprint"],
introspect=lambda output: refresh_catalog(cfg, output_path=output),
build_lsh=lambda physical, output: build_lsh_artifacts(
cfg, physical_file=physical, output_dir=output
),
lsh_filenames=lsh_names,
current_physical=cfg.paths.artifacts / "mschema" / "physical.yaml",
current_lsh_dir=cfg.paths.indexes / "lsh",
)
return pipeline.run(steps, resume_run_id=resume)
def _parse_dwh_steps(value: str) -> tuple[str, ...]:
allowed = ("introspect", "lsh")
steps = tuple(part.strip() for part in value.split(",") if part.strip())
if (
not steps
or len(steps) != len(set(steps))
or any(step not in allowed for step in steps)
or tuple(sorted(steps, key=allowed.index)) != steps
):
raise ValueError("steps must be a unique ordered subset of introspect,lsh")
return steps
def run_from_config(config: Path, *, dry_run: bool = False, resume: str | None = None):
from tht.adapters.factory import build_evidence_sources, build_vector_store
from tht.cli.schema_cmd import _load_config_or_exit
from tht.cli.vector_cmd import make_embedder
from tht.corpus.chunk import ChunkPolicy
from tht.corpus.pipeline import CorpusPipeline
from tht.corpus.store import CorpusStore
cfg = _load_config_or_exit(config)
if cfg.embeddings is None:
raise RuntimeError("embeddings are not configured")
corpus_root = cfg.paths.artifacts.parent / "corpus"
pipeline = CorpusPipeline(
store=CorpusStore(corpus_root), sources=build_evidence_sources(cfg),
embedder=make_embedder(cfg.embeddings),
vector_store=build_vector_store(cfg, require_write=True),
embedding_model=cfg.embeddings.model, embedding_dimensions=cfg.embeddings.dim,
chunk_policy=ChunkPolicy(version="chunk-v1", max_chars=cfg.vector.max_chunk_chars),
pipeline_version="evidence-v1",
retain_published_generations=cfg.vector.retain_published_generations,
)
def fingerprint(value: str) -> str:
return "sha256:" + hashlib.sha256(value.encode()).hexdigest()
return pipeline.run_as_job(
workspace_id=config.stem.lower().replace(".", "-").replace("_", "-"),
workspace_root=corpus_root.parent,
config_fingerprint=fingerprint(cfg.model_dump_json()),
input_fingerprint=fingerprint(config.resolve().as_posix()),
dry_run=dry_run,
resume_run_id=resume,
)
def gc_from_config(config: Path, *, dry_run: bool = False):
from tht.adapters.factory import build_evidence_sources, build_vector_store
from tht.cli.schema_cmd import _load_config_or_exit
from tht.cli.vector_cmd import make_embedder
from tht.corpus.chunk import ChunkPolicy
from tht.corpus.pipeline import CorpusPipeline
from tht.corpus.store import CorpusStore
cfg = _load_config_or_exit(config)
if cfg.embeddings is None:
raise RuntimeError("embeddings are not configured")
corpus_root = cfg.paths.artifacts.parent / "corpus"
pipeline = CorpusPipeline(
store=CorpusStore(corpus_root), sources=build_evidence_sources(cfg),
embedder=make_embedder(cfg.embeddings), vector_store=build_vector_store(cfg, require_write=True),
embedding_model=cfg.embeddings.model, embedding_dimensions=cfg.embeddings.dim,
chunk_policy=ChunkPolicy(version="chunk-v1", max_chars=cfg.vector.max_chunk_chars),
pipeline_version="evidence-v1",
retain_published_generations=cfg.vector.retain_published_generations,
)
pipeline.workspace_id = config.stem.lower().replace(".", "-").replace("_", "-")
return pipeline.gc(workspace_root=corpus_root.parent, dry_run=dry_run)
@preprocess_app.command("evidence")
def evidence_cmd(
action: str | None = typer.Argument(None),
config: Path = CONFIG_OPT,
dry_run: bool = typer.Option(False, "--dry-run"),
resume: str | None = typer.Option(None, "--resume"),
json_output: bool = typer.Option(False, "--json"),
) -> None:
if action is not None and action != "gc":
raise typer.BadParameter("only the optional 'gc' action is supported")
if action == "gc":
try:
payload = gc_from_config(config, dry_run=dry_run)
except Exception:
payload = {"status": "failed", "error": "evidence cleanup failed"}
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.secho("ERRORE: evidence cleanup failed", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1) from None
if json_output:
typer.echo(json.dumps(payload, ensure_ascii=False, sort_keys=True))
else:
typer.echo(f"OK: evicted={len(payload['evicted'])} failures={len(payload['failures'])}")
return
if resume is not None and re.fullmatch(r"[0-9a-f]{32}", resume) is None:
payload = {"status": "failed", "error": "resume requires a preprocessing run id"}
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.secho("ERRORE: resume requires a preprocessing run id", fg=typer.colors.RED, err=True)
raise typer.Exit(code=2)
try:
result = run_from_config(config, dry_run=dry_run, resume=resume)
except Exception:
payload = {"status": "failed", "error": "preprocessing failed"}
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.secho("ERRORE: preprocessing failed", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1) from None
payload = result.model_dump(mode="json")
if payload.get("status") != "succeeded":
payload["error"] = "preprocessing job failed"
if json_output:
typer.echo(json.dumps(payload, ensure_ascii=False, sort_keys=True))
else:
typer.secho("ERRORE: preprocessing job failed", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
if json_output:
typer.echo(json.dumps(payload, ensure_ascii=False, sort_keys=True))
else:
counts = payload["counts"]
typer.echo(
f"OK: run={payload['run_id']} generation={payload['generation']} "
f"changed={counts['changed']} unchanged={counts['unchanged']} "
f"removed={counts['removed']}"
)
@preprocess_app.command("dwh")
def dwh_cmd(
config: Path = CONFIG_OPT,
steps: str = typer.Option("introspect,lsh", "--steps"),
resume: str | None = typer.Option(None, "--resume"),
json_output: bool = typer.Option(False, "--json"),
) -> None:
try:
selected = _parse_dwh_steps(steps)
except ValueError:
payload = {"status": "failed", "error": "invalid DWH preprocessing steps"}
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.secho("ERRORE: invalid DWH preprocessing steps", fg=typer.colors.RED, err=True)
raise typer.Exit(code=2) from None
if resume is not None and re.fullmatch(r"[0-9a-f]{32}", resume) is None:
payload = {"status": "failed", "error": "resume requires a preprocessing run id"}
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.secho("ERRORE: resume requires a preprocessing run id", fg=typer.colors.RED, err=True)
raise typer.Exit(code=2)
try:
result = run_dwh_from_config(config, steps=selected, resume=resume)
except Exception:
payload = {"status": "failed", "error": "DWH preprocessing failed"}
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.secho("ERRORE: DWH preprocessing failed", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1) from None
payload = result.model_dump(mode="json")
if json_output:
typer.echo(json.dumps(payload, ensure_ascii=False, sort_keys=True))
elif result.status == "succeeded":
typer.echo(f"OK: run={result.run_id} stages={','.join(selected)}")
else:
typer.secho(f"ERRORE: run={result.run_id} DWH preprocessing failed", fg=typer.colors.RED, err=True)
if result.status != "succeeded":
raise typer.Exit(code=1)
+51 -34
View File
@@ -1,16 +1,31 @@
from pathlib import Path
import logging
import typer
from sqlalchemy.exc import OperationalError
from tht.adapters.factory import build_dwh
from tht.cli.config_cmd import CONFIG_OPT
from tht.config import ConfigError, load_config
from tht.db.connection import make_engine
from tht.db.introspect import IntrospectionError, introspect
from tht.db.sampling import add_examples
from tht.db.sampling import is_text_type
from tht.mschema.eligibility import classify_all
schema_app = typer.Typer(help="Gestione mschema (rappresentazione canonica dello schema)")
logger = logging.getLogger(__name__)
def _add_examples(dwh, phys, examples) -> None:
for table_name, table in phys.tables.items():
for column_name, column in table.columns.items():
if not is_text_type(column.type):
continue
try:
sampled = dwh.sample_column(
table_name, column_name, limit=examples.max_per_column
)
except Exception as exc:
logger.warning("Campionamento saltato per %s.%s: %s",
table_name, column_name, exc)
continue
column.examples = [str(value) for value in sampled if value not in (None, "")]
def _load_config_or_exit(config: Path):
@@ -22,13 +37,27 @@ def _load_config_or_exit(config: Path):
def physical_path(cfg) -> Path:
return cfg.paths.artifacts / "mschema" / "physical.yaml"
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
if not (cfg.paths.artifacts.parent / ".tht-dwh").exists():
return cfg.paths.artifacts / "mschema" / "physical.yaml"
return resolve_dwh_snapshot(cfg).physical
def annotations_path(cfg) -> Path:
return cfg.paths.artifacts / "mschema" / "annotations.yaml"
def refresh_catalog(cfg, *, dwh=None, output_path: Path | None = None):
"""Run the existing catalog algorithm and persist its canonical output."""
target = dwh if dwh is not None else build_dwh(cfg)
physical = target.introspect()
_add_examples(target, physical, cfg.examples)
classify_all(physical, cfg.eligibility)
physical.to_yaml(output_path or (cfg.paths.artifacts / "mschema" / "physical.yaml"))
return physical
@schema_app.command("introspect")
def introspect_cmd(
config: Path = CONFIG_OPT,
@@ -43,8 +72,11 @@ def introspect_cmd(
Se physical.yaml esiste già, esce subito (cache); usa --refresh per rigenerarlo.
"""
cfg = _load_config_or_exit(config)
out = physical_path(cfg)
if out.exists() and not refresh:
dwh_root = cfg.paths.artifacts.parent / ".tht-dwh"
out = cfg.paths.artifacts / "mschema" / "physical.yaml"
if dwh_root.exists() or dwh_root.is_symlink():
out = physical_path(cfg)
if (dwh_root.exists() or dwh_root.is_symlink()) and out.exists() and not refresh:
from datetime import UTC, datetime
from tht.mschema.models import PhysicalSchema
@@ -64,33 +96,18 @@ def introspect_cmd(
fg=typer.colors.GREEN,
)
return
if cfg.database.transport == "rest":
from tht.db.introspect import introspect_rest
from tht.db.sampling import add_examples_rest
from tht.rest.client import RestClient, RestError
try:
from tht.cli.preprocess_cmd import run_dwh_from_config
from tht.mschema.models import PhysicalSchema
client = RestClient(cfg.rest)
try:
phys = introspect_rest(
client, database=cfg.database.database, schema=cfg.database.db_schema
)
add_examples_rest(client, phys, cfg.examples)
classify_all(phys, cfg.eligibility)
except RestError as e:
typer.secho(f"ERRORE: {e}", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
else:
engine = make_engine(cfg.database)
try:
phys = introspect(
engine, database=cfg.database.database, schema=cfg.database.db_schema
)
add_examples(engine, phys, cfg.examples)
classify_all(phys, cfg.eligibility)
except (OperationalError, IntrospectionError) as e:
typer.secho(f"ERRORE: {e}", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
phys.to_yaml(out)
report = run_dwh_from_config(config, steps=("introspect",))
if report.status != "succeeded":
raise RuntimeError("DWH preprocessing failed")
out = physical_path(cfg)
phys = PhysicalSchema.from_yaml(out)
except Exception as e:
typer.secho(f"ERRORE: {e}", fg=typer.colors.RED, err=True)
raise typer.Exit(code=1)
n_cols = sum(len(t.columns) for t in phys.tables.values())
n_ignored = sum(
1 for t in phys.tables.values() for c in t.columns.values() if not c.eligible
+39 -9
View File
@@ -21,8 +21,18 @@ DEFAULT_TOP_FALLBACK = 10
search_app = typer.Typer(help="Ricerca semantica (evidence/schema/values) nel vectorstore")
def _leased_dwh_snapshot(cfg, context: typer.Context):
from tht.jobs.dwh_pipeline import lease_dwh_snapshot
lease = lease_dwh_snapshot(cfg)
snapshot = lease.__enter__()
context.call_on_close(lambda: lease.__exit__(None, None, None))
return snapshot
@search_app.command("find")
def search_cmd(
ctx: typer.Context,
keyword: str = typer.Argument(..., help="Termine da cercare, es. 'ablazione'."),
config: Path = CONFIG_OPT,
top: int | None = typer.Option(
@@ -47,7 +57,18 @@ def search_cmd(
from tht.search import combined_search
cfg = _load_config_or_exit(config)
from tht.search.evidence import validate_corpus_workspace
workspace_id = config.stem.lower().replace(".", "-").replace("_", "-")
validate_corpus_workspace(cfg, workspace_id)
dwh_snapshot = _leased_dwh_snapshot(cfg, ctx)
require_vector_cfg(cfg)
from tht.search.evidence import active_searcher
runtime_searcher = active_searcher(
cfg, open_searcher(cfg),
workspace_id=workspace_id,
)
if kind is not None and kind not in KIND_MAP:
typer.secho(
f"ERRORE: --kind sconosciuto: {kind} (validi: {', '.join(KIND_MAP)})",
@@ -85,7 +106,7 @@ def search_cmd(
lsh_hits = None
try:
lsh, minhashes, meta = load_index(
cfg.paths.indexes / "lsh", name=cfg.database.db_schema
dwh_snapshot.lsh_dir, name=cfg.database.db_schema
)
hits = query_index(lsh, minhashes, keyword, meta, top_n=top * 3)
lsh_hits = [(h.table, h.column, h.value, h.score) for h in hits]
@@ -97,12 +118,12 @@ def search_cmd(
)
if kind == "schema":
from tht.cli.schema_cmd import annotations_path, physical_path
from tht.cli.schema_cmd import annotations_path
from tht.mschema.models import Annotations, PhysicalSchema
from tht.mschema.render import to_mschema_text
from tht.search import schema_tables
phys_file = physical_path(cfg)
phys_file = dwh_snapshot.physical
if not phys_file.exists():
typer.secho(
f"ERRORE: {phys_file} non trovato. Esegui prima `tht schema introspect`.",
@@ -112,7 +133,7 @@ def search_cmd(
candidates = combined_search(
keyword=keyword, lsh_hits=lsh_hits,
store=open_searcher(cfg), embedder=make_embedder(cfg.embeddings),
store=runtime_searcher, embedder=make_embedder(cfg.embeddings),
top=cfg.search.schema_chunk_pool, rrf_k=cfg.search.rrf_k,
kinds=KIND_MAP["schema"],
)
@@ -156,7 +177,7 @@ def search_cmd(
kinds = KIND_MAP.get(kind) if kind else None
results = combined_search(
keyword=keyword, lsh_hits=lsh_hits if kind != "evidence" else None,
store=open_searcher(cfg), embedder=make_embedder(cfg.embeddings),
store=runtime_searcher, embedder=make_embedder(cfg.embeddings),
top=top, rrf_k=cfg.search.rrf_k, kinds=kinds,
)
@@ -216,6 +237,7 @@ PACK_EXCERPT_CHARS = 400
@search_app.command("pack")
def pack_cmd(
ctx: typer.Context,
question: str = typer.Argument(..., help="La domanda in linguaggio naturale."),
config: Path = CONFIG_OPT,
session: str = typer.Option(
@@ -238,6 +260,11 @@ def pack_cmd(
from tht.vectorstore.rest_client import VectorRestError
cfg = _load_config_or_exit(config)
from tht.search.evidence import validate_corpus_workspace
workspace_id = config.stem.lower().replace(".", "-").replace("_", "-")
validate_corpus_workspace(cfg, workspace_id)
dwh_snapshot = _leased_dwh_snapshot(cfg, ctx)
require_vector_cfg(cfg)
tables: list[dict] = []
@@ -249,17 +276,20 @@ def pack_cmd(
vec = None
searcher = embedder = None
try:
searcher = open_searcher(cfg)
from tht.search.evidence import active_searcher
searcher = active_searcher(
cfg, open_searcher(cfg),
workspace_id=workspace_id,
)
embedder = make_embedder(cfg.embeddings)
vec = embedder.embed_query(question)
except degrade as e:
warnings.append(f"retrieval non disponibile ({e}): prosegui con le ricerche live")
if vec is not None:
from tht.cli.schema_cmd import physical_path
descriptions: dict[str, str] = {}
phys_file = physical_path(cfg)
phys_file = dwh_snapshot.physical
if phys_file.exists():
from tht.mschema.models import PhysicalSchema
+6 -34
View File
@@ -90,46 +90,18 @@ def validate_or_exit(cfg, sql: str, session_id: str | None):
return result
def _ro_engine(cfg):
"""Engine sul target con search_path impostato allo schema (nomi non qualificati)."""
from sqlalchemy import create_engine
db = cfg.database
url = f"postgresql+psycopg2://{db.user}:{db.password}@{db.host}:{db.port}/{db.database}"
return create_engine(
url, echo=False,
connect_args={"options": f"-csearch_path={db.db_schema}"},
)
def _rest_client(cfg):
from tht.rest.client import RestClient
return RestClient(cfg.rest)
def do_explain(cfg, sql: str):
"""EXPLAIN secondo il transport configurato (direct|rest)."""
if cfg.database.transport == "rest":
from tht.rest.execute import explain_rest
"""EXPLAIN through the configured DWH adapter."""
from tht.adapters.factory import build_dwh
return explain_rest(_rest_client(cfg), sql)
from tht.execute import explain
return explain(_ro_engine(cfg), sql, timeout_ms=cfg.execution.statement_timeout_ms)
return build_dwh(cfg).explain(sql)
def _run_transport(cfg, sql: str, *, limit: int):
"""Dispatch all'esecutore controllato secondo il transport (direct|rest)."""
if cfg.database.transport == "rest":
from tht.rest.execute import run_controlled_rest
"""Dispatch through the configured DWH adapter."""
from tht.adapters.factory import build_dwh
return run_controlled_rest(_rest_client(cfg), sql, limit=limit)
from tht.execute import run_controlled
return run_controlled(
_ro_engine(cfg), sql, limit=limit, timeout_ms=cfg.execution.statement_timeout_ms
)
return build_dwh(cfg).run_query(sql, limit=limit)
def do_run(cfg, sql: str, *, limit: int, offset: int = 0):
+13 -21
View File
@@ -46,35 +46,27 @@ def open_store(cfg, table: str):
Sul server preferisce la connessione diretta. In profilo workstation usa `vector_write_rest`
se configurato, con upsert remoto non distruttivo.
"""
if has_vector_write_rest(cfg) and (cfg.profile == "workstation" or cfg.vector_db is None):
from tht.vectorstore.rest_client import VectorRestClient
from tht.vectorstore.rest_writer import RestVectorWriter
from tht.adapters.factory import build_vector_loader
return RestVectorWriter(VectorRestClient(cfg.vector_write_rest), table=table)
from tht.db.connection import make_engine
from tht.vectorstore.store import VectorStore
engine = make_engine(cfg.vector_db)
return VectorStore(
engine, schema=cfg.vector_db.db_schema, table=table, dim=cfg.embeddings.dim
)
return build_vector_loader(cfg, table)
def open_searcher(cfg):
"""Searcher per la LETTURA (similarity search): via REST se `vector_rest` è configurato,
altrimenti connessione diretta (dev/test)."""
if cfg.vector_rest is not None:
from tht.vectorstore.reader import RestSearcher
from tht.vectorstore.rest_client import VectorRestClient
from tht.adapters.factory import build_vector_store
from tht.vectorstore.reader import tables_for_kinds
return RestSearcher(VectorRestClient(cfg.vector_rest))
from tht.db.connection import make_engine
from tht.vectorstore.reader import DirectSearcher
store = build_vector_store(cfg)
return DirectSearcher(
make_engine(cfg.vector_db), schema=cfg.vector_db.db_schema, dim=cfg.embeddings.dim
)
class AdapterSearcher:
def search(self, query_vec, top_n=10, kinds=None, metadata_filter=None):
return store.search(
tables_for_kinds(kinds), query_vec, limit=top_n, kinds=kinds,
metadata_filter=metadata_filter,
)
return AdapterSearcher()
def _print_stats(stats) -> None:
+227
View File
@@ -0,0 +1,227 @@
"""Versioned, transactional migrations for the direct pgvector schema."""
from __future__ import annotations
import hashlib
import json
import re
from dataclasses import dataclass
from importlib.resources import files
from importlib.resources.abc import Traversable
from pathlib import Path
import typer
from sqlalchemy import create_engine, text
from sqlalchemy.exc import SQLAlchemyError
from tht.cli.vector_cmd import vector_app
MIGRATIONS_DIR = files("tht").joinpath("migrations", "vector")
_MIGRATION_NAME = re.compile(r"^(?P<version>\d+)_(?P<name>[a-z0-9_]+)\.sql$")
_LOCK_KEY = 7_304_708_654_221_909_028
class MigrationError(RuntimeError):
"""Raised when migration discovery or application is unsafe."""
@dataclass(frozen=True)
class Migration:
version: str
name: str
path: Traversable
checksum: str
@dataclass(frozen=True)
class MigrationStatus:
applied: tuple[Migration, ...]
pending: tuple[Migration, ...]
drifted: tuple[Migration, ...]
def _migration_source(directory: Traversable | Path | str) -> Traversable:
return Path(directory) if isinstance(directory, (str, Path)) else directory
def _discover(directory: Traversable | Path | str) -> tuple[Migration, ...]:
source = _migration_source(directory)
migrations = []
seen_versions: set[int] = set()
paths = [path for path in source.iterdir() if path.name.endswith(".sql")]
parsed = []
for path in paths:
match = _MIGRATION_NAME.fullmatch(path.name)
if match is None:
raise MigrationError(f"Invalid migration filename: {path.name}")
version = match.group("version")
numeric_version = int(version)
if numeric_version in seen_versions:
raise MigrationError(f"Duplicate migration version: {numeric_version}")
seen_versions.add(numeric_version)
parsed.append((numeric_version, version, match.group("name"), path))
for _, version, name, path in sorted(parsed, key=lambda item: item[0]):
migrations.append(
Migration(
version=version,
name=name,
path=path,
checksum=hashlib.sha256(path.read_bytes()).hexdigest(),
)
)
if not migrations:
raise MigrationError(f"No migrations found in {source}")
return tuple(migrations)
def _applied(connection) -> dict[str, str]:
exists = connection.execute(
text("SELECT pg_catalog.to_regclass('public.tht_vector_migrations')")
).scalar()
if exists is None:
return {}
return dict(
connection.execute(
text("SELECT version, checksum FROM public.tht_vector_migrations")
).all()
)
def _reject_unknown_versions(
migrations: tuple[Migration, ...], applied_checksums: dict[str, str]
) -> None:
local_versions = {migration.version for migration in migrations}
unknown = sorted(
set(applied_checksums) - local_versions,
key=lambda version: (0, int(version)) if version.isdigit() else (1, version),
)
if unknown:
raise MigrationError(
"Database migration versions absent from local manifest: " + ", ".join(unknown)
)
def migration_status(
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
) -> MigrationStatus:
migrations = _discover(migrations_dir)
engine = create_engine(database_url)
try:
with engine.connect() as connection:
connection.exec_driver_sql("SET LOCAL search_path = pg_catalog, pg_temp")
applied_checksums = _applied(connection)
finally:
engine.dispose()
_reject_unknown_versions(migrations, applied_checksums)
applied = tuple(
migration
for migration in migrations
if applied_checksums.get(migration.version) == migration.checksum
)
drifted = tuple(
migration
for migration in migrations
if migration.version in applied_checksums
and applied_checksums[migration.version] != migration.checksum
)
pending = tuple(
migration for migration in migrations if migration.version not in applied_checksums
)
return MigrationStatus(applied=applied, pending=pending, drifted=drifted)
def migrate(
database_url: str, migrations_dir: Traversable | Path | str = MIGRATIONS_DIR
) -> MigrationStatus:
migrations = _discover(migrations_dir)
engine = create_engine(database_url)
current: Migration | None = None
try:
with engine.begin() as connection:
connection.exec_driver_sql("SET LOCAL search_path = pg_catalog, pg_temp")
connection.execute(
text("SELECT pg_catalog.pg_advisory_xact_lock(:key)"), {"key": _LOCK_KEY}
)
connection.exec_driver_sql(
"""CREATE TABLE IF NOT EXISTS public.tht_vector_migrations (
version text PRIMARY KEY,
name text NOT NULL,
checksum text NOT NULL,
applied_at timestamptz NOT NULL DEFAULT pg_catalog.now()
)"""
)
connection.exec_driver_sql(
"REVOKE ALL ON public.tht_vector_migrations FROM PUBLIC"
)
applied_checksums = _applied(connection)
_reject_unknown_versions(migrations, applied_checksums)
drifted = [
item
for item in migrations
if item.version in applied_checksums
and applied_checksums[item.version] != item.checksum
]
if drifted:
versions = ", ".join(item.version for item in drifted)
raise MigrationError(f"Migration checksum drift: {versions}")
for current in migrations:
if current.version in applied_checksums:
continue
connection.exec_driver_sql(current.path.read_text())
connection.execute(
text(
"INSERT INTO public.tht_vector_migrations (version, name, checksum) "
"VALUES (:version, :name, :checksum)"
),
{
"version": current.version,
"name": current.name,
"checksum": current.checksum,
},
)
except MigrationError:
raise
except SQLAlchemyError as exc:
filename = current.path.name if current is not None else "migration setup"
raise MigrationError(f"Failed to apply {filename}: {type(exc).__name__}") from exc
finally:
engine.dispose()
return migration_status(database_url, migrations_dir)
def _payload(status: MigrationStatus) -> dict[str, list[str]]:
return {
"applied": [item.version for item in status.applied],
"drifted": [item.version for item in status.drifted],
"pending": [item.version for item in status.pending],
}
@vector_app.command("migrate")
def migrate_cmd(
database_url: str = typer.Option(
..., "--database-url", envvar="THT_VECTOR_ADMIN_URL", help="Admin PostgreSQL URL."
),
status_only: bool = typer.Option(False, "--status", help="Inspect without applying."),
json_output: bool = typer.Option(False, "--json", help="Emit pristine JSON."),
) -> None:
"""Apply or inspect the local pgvector schema migrations."""
try:
status = migration_status(database_url) if status_only else migrate(database_url)
except (MigrationError, SQLAlchemyError) as exc:
if json_output:
typer.echo(json.dumps({"error": str(exc)}, sort_keys=True))
else:
typer.echo(f"ERROR: {exc}", err=True)
raise typer.Exit(code=1) from None
payload = _payload(status)
if json_output:
typer.echo(json.dumps(payload, sort_keys=True))
else:
typer.echo(
f"Applied: {len(status.applied)}; pending: {len(status.pending)}; "
f"drifted: {len(status.drifted)}"
)
__all__ = ["MigrationError", "MigrationStatus", "migrate", "migration_status"]
+232 -11
View File
@@ -1,10 +1,13 @@
import os
import re
import warnings
from pathlib import Path
from typing import Any, Literal
from typing import Annotated, Any, Literal
import yaml
from pydantic import BaseModel, Field, ValidationError
from pydantic import BaseModel, Field, PrivateAttr, SecretStr, model_validator, ValidationError
from tht.config_compat import translate_legacy_config
_ENV_RE = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}")
@@ -15,6 +18,7 @@ class ConfigError(Exception):
def _expand_env(value: Any) -> Any:
if isinstance(value, str):
def repl(m: re.Match) -> str:
var = m.group(1)
if var not in os.environ:
@@ -32,6 +36,29 @@ def _expand_env(value: Any) -> Any:
return value
def _resolve_secret_files(value: Any) -> Any:
if isinstance(value, dict):
resolved = {key: _resolve_secret_files(item) for key, item in value.items()}
for secret_name in ("password", "access_key", "secret_key", "session_token"):
file_name = f"{secret_name}_file"
if file_name not in resolved:
continue
if secret_name in resolved:
raise ConfigError(f"{secret_name} and {file_name} are mutually exclusive")
path = Path(resolved.pop(file_name))
try:
secret = path.read_text()
except (OSError, UnicodeError) as exc:
raise ConfigError(f"Cannot read secret file: {path}") from exc
if not secret or any(char.isspace() for char in secret) or "\x00" in secret:
raise ConfigError(f"Invalid secret file: {path}")
resolved[secret_name] = secret
return resolved
if isinstance(value, list):
return [_resolve_secret_files(item) for item in value]
return value
class DatabaseConfig(BaseModel):
host: str = "localhost"
port: int = 5432
@@ -56,12 +83,68 @@ class RestConfig(BaseModel):
ssl_ca: str | None = None # path al certificato CA (per server con CA interna)
class DatabaseIdentityConfig(BaseModel):
database: str
db_schema: str = Field(alias="schema")
model_config = {"populate_by_name": True}
class PostgresDwhConfig(BaseModel):
type: Literal["postgres_direct"]
connection: DatabaseConfig
class ThothRestDwhConfig(BaseModel):
type: Literal["thoth_rest"]
database: DatabaseIdentityConfig
endpoint: RestConfig
DwhResourceConfig = Annotated[
PostgresDwhConfig | ThothRestDwhConfig,
Field(discriminator="type"),
]
class PgvectorDirectConfig(BaseModel):
type: Literal["pgvector_direct"]
reader: DatabaseConfig | None = None
writer: DatabaseConfig | None = None
# Deprecated compatibility: a single direct connection historically meant read-only.
connection: DatabaseConfig | None = None
@model_validator(mode="after")
def validate_connections(self):
if self.reader is None and self.writer is None and self.connection is None:
raise ValueError("pgvector_direct requires a reader or writer connection")
return self
class ThothVectorHttpConfig(BaseModel):
type: Literal["thoth_vector_http"]
reader: RestConfig | None = None
writer: RestConfig | None = None
# Transitional direct loading path used by the server profile.
direct: DatabaseConfig | None = None
VectorResourceConfig = Annotated[
PgvectorDirectConfig | ThothVectorHttpConfig,
Field(discriminator="type"),
]
class PathsConfig(BaseModel):
artifacts: Path = Path("artifacts")
indexes: Path = Path("indexes")
sessions: Path = Path("sessions")
class WorkspaceRoots(PathsConfig):
pass
class ExamplesConfig(BaseModel):
max_per_column: int = 10
@@ -84,21 +167,73 @@ class LshConfig(BaseModel):
class EligibilityConfig(BaseModel):
# Soglie del principio di column eligibility (testo ampio ignorato ovunque).
max_declared_len: int = 128 # char/varchar dichiarati <= soglia: eligible senza campionare
max_avg_length: int = 40 # fallback data-driven: lunghezza media valori campionati
max_sampled_len: int = 200 # fallback data-driven: lunghezza massima valore campionato
max_declared_len: int = 128 # char/varchar dichiarati <= soglia: eligible senza campionare
max_avg_length: int = 40 # fallback data-driven: lunghezza media valori campionati
max_sampled_len: int = 200 # fallback data-driven: lunghezza massima valore campionato
# Colonne di servizio sempre ignorate per nome (match case-insensitive), a prescindere
# dal tipo: metadati ETL/audit non analitici (es. timestamp di ultimo aggiornamento).
ignore_columns: list[str] = ["etl_last_update"]
class FilesystemEvidenceSourceConfig(BaseModel):
type: Literal["filesystem"]
root: Path
patterns: list[str] = ["**/*.md"]
max_bytes: int = Field(default=10 * 1024 * 1024, gt=0)
class HttpEvidenceSourceConfig(BaseModel):
type: Literal["http"]
# Manifest URLs may contain signed query parameters. Treat the complete transport URL as
# secret-bearing configuration; adapters derive a query-free provenance URI from it.
urls: list[SecretStr] = Field(min_length=1)
connect_timeout: float = Field(default=5, gt=0)
read_timeout: float = Field(default=30, gt=0)
max_bytes: int = Field(default=10 * 1024 * 1024, gt=0)
max_redirects: int = Field(default=5, ge=0)
allow_private_hosts: bool = False
max_cache_bytes: int = Field(default=64 * 1024 * 1024, gt=0)
class S3EvidenceSourceConfig(BaseModel):
type: Literal["s3"]
bucket: str = Field(min_length=1)
prefix: str = ""
endpoint_url: str | None = None
region: str | None = None
access_key: SecretStr | None = None
secret_key: SecretStr | None = None
session_token: SecretStr | None = None
trusted_endpoint: bool = False
allow_private_endpoint: bool = False
allow_insecure_endpoint: bool = False
max_bytes: int = Field(default=10 * 1024 * 1024, gt=0)
max_objects: int = Field(default=10_000, gt=0)
max_pages: int = Field(default=100, gt=0)
page_size: int = Field(default=1000, gt=0, le=1000)
EvidenceSourceConfig = Annotated[
FilesystemEvidenceSourceConfig | HttpEvidenceSourceConfig | S3EvidenceSourceConfig,
Field(discriminator="type"),
]
class EvidenceSourcesConfig(BaseModel):
source_root: Path
# Legacy curated-tree configuration remains accepted during migration.
source_root: Path | None = None
# cartella curata a mano nell'ETL (relativa a source_root): unica fonte delle
# evidence. Niente piu' estrazione automatica dalle schede tabella: i documenti
# qui dentro sono gia' evidence pronte (frontmatter + corpo), scelte e arricchite
# dall'autore ETL e organizzate in sottocartelle per dominio.
evidence_dir: str = "evidence"
sources: list[EvidenceSourceConfig] = []
@model_validator(mode="after")
def require_a_source(self):
if self.source_root is None and not self.sources:
raise ValueError("evidence requires source_root or sources")
return self
class EmbeddingsConfig(BaseModel):
@@ -114,11 +249,13 @@ class EmbeddingsConfig(BaseModel):
class VectorConfig(BaseModel):
max_chunk_chars: int = 4000
# ACTIVE plus the two most recent rollback generations by default.
retain_published_generations: int = Field(default=3, ge=1)
class SearchConfig(BaseModel):
rrf_k: int = 60
top_schema_tables: int = 12 # default `--top` per `tht search --kind schema` (n. tabelle)
top_schema_tables: int = 12 # default `--top` per `tht search --kind schema` (n. tabelle)
schema_chunk_pool: int = 150 # chunk tabella/colonna fusi prima dell'aggregazione a tabella
@@ -131,13 +268,27 @@ class ExecutionConfig(BaseModel):
max_aggregate_cells: int = 20
max_export_rows: int = 100000
forbidden_functions: list[str] = [
"setval", "nextval", "pg_advisory_lock", "pg_advisory_xact_lock",
"dblink", "dblink_exec", "pg_terminate_backend", "pg_cancel_backend",
"lo_import", "lo_export", "pg_reload_conf",
"setval",
"nextval",
"pg_advisory_lock",
"pg_advisory_xact_lock",
"dblink",
"dblink_exec",
"pg_terminate_backend",
"pg_cancel_backend",
"lo_import",
"lo_export",
"pg_reload_conf",
]
class Config(BaseModel):
_workspace_id: str = PrivateAttr(default="default")
_config_source: str = PrivateAttr(default="direct")
dwh: DwhResourceConfig
vectors: VectorResourceConfig | None = None
roots: WorkspaceRoots = WorkspaceRoots()
# Compatibility views retained until all call sites consume typed resources.
database: DatabaseConfig
# Profilo dell'installazione, letto da THT_PROFILE (.env), non dallo yaml versionato.
# server: ricostruisce i derivati (artefatti, LSH, vettori schema nel vectordb).
@@ -166,6 +317,15 @@ class Config(BaseModel):
# una API key separata dalla lettura; espone solo upsert/hash via RPC allowlist.
vector_write_rest: RestConfig | None = None
@model_validator(mode="before")
@classmethod
def accept_legacy_constructor_fields(cls, value: Any) -> Any:
if not isinstance(value, dict) or "dwh" in value:
return value
translated, _ = translate_legacy_config(value)
_populate_legacy_views(translated)
return translated
def load_config(path: Path) -> Config:
if not path.exists():
@@ -173,8 +333,11 @@ def load_config(path: Path) -> Config:
raw = yaml.safe_load(path.read_text())
if not isinstance(raw, dict):
raise ConfigError(f"Configurazione non valida (atteso un mapping YAML): {path}")
expanded = _resolve_secret_files(_expand_env(raw))
translated, used_legacy = translate_legacy_config(expanded)
_populate_legacy_views(translated)
try:
cfg = Config.model_validate(_expand_env(raw))
cfg = Config.model_validate(translated)
except ValidationError as e:
raise ConfigError(f"Configurazione non valida in {path}:\n{e}") from e
env_profile = os.environ.get("THT_PROFILE")
@@ -188,4 +351,62 @@ def load_config(path: Path) -> Config:
raise ConfigError(
f"transport: rest richiede la sezione `rest` (base_url, api_key) in {path}."
)
data_root = os.environ.get("THT_DATA_ROOT")
if data_root:
# Import locally: paths owns resolution, while ConfigError remains the public
# configuration exception callers already handle.
from tht.paths import resolve_workspace_paths
resolved = resolve_workspace_paths(path, cfg, Path(data_root))
cfg = cfg.model_copy(
update={
"paths": PathsConfig(
sessions=resolved.sessions,
artifacts=resolved.artifacts,
indexes=resolved.indexes,
)
}
)
elif not used_legacy:
# Modern `roots` replace `paths`; without a mounted data root retain the old
# working-directory-relative behavior used by local development.
cfg = cfg.model_copy(update={"paths": PathsConfig(**cfg.roots.model_dump())})
if used_legacy:
warnings.warn(
"DEPRECATION: legacy workspace resource keys are deprecated; "
"use dwh, vectors, and roots.",
FutureWarning,
stacklevel=2,
)
cfg._workspace_id = path.stem.lower().replace(".", "-").replace("_", "-")
cfg._config_source = path.resolve().as_posix()
return cfg
def _populate_legacy_views(raw: dict[str, Any]) -> None:
"""Populate old Config attributes for command compatibility during migration."""
dwh = raw.get("dwh")
if "database" not in raw and isinstance(dwh, dict):
if dwh.get("type") == "postgres_direct":
raw["database"] = {**dwh["connection"], "transport": "direct"}
elif dwh.get("type") == "thoth_rest":
raw["database"] = {
**dwh["database"],
"user": "rest",
"password": "",
"transport": "rest",
}
raw["rest"] = dwh["endpoint"]
vectors = raw.get("vectors")
if isinstance(vectors, dict):
if vectors.get("type") == "pgvector_direct":
raw.setdefault(
"vector_db",
vectors.get("writer") or vectors.get("reader") or vectors.get("connection"),
)
elif vectors.get("type") == "thoth_vector_http":
raw.setdefault("vector_rest", vectors.get("reader"))
raw.setdefault("vector_write_rest", vectors.get("writer"))
raw.setdefault("vector_db", vectors.get("direct"))
raw.setdefault("paths", raw.get("roots", {}))
+74
View File
@@ -0,0 +1,74 @@
from __future__ import annotations
from copy import deepcopy
from typing import Any
_LEGACY_RESOURCE_KEYS = {
"database",
"rest",
"vector_db",
"vector_rest",
"vector_write_rest",
"paths",
}
def _as_mapping(value: Any) -> dict[str, Any] | None:
if isinstance(value, dict):
return deepcopy(value)
model_dump = getattr(value, "model_dump", None)
if callable(model_dump):
return model_dump(by_alias=True)
return None
def translate_legacy_config(raw: dict[str, Any]) -> tuple[dict[str, Any], bool]:
"""Translate the legacy flat resource keys without validating their contents."""
translated = deepcopy(raw)
legacy = any(key in raw for key in _LEGACY_RESOURCE_KEYS)
if not legacy:
return translated, False
database = _as_mapping(raw.get("database"))
rest = raw.get("rest")
if "dwh" not in translated and database is not None:
if database.get("transport", "direct") == "rest":
identity = {
key: database[key]
for key in ("database", "schema")
if key in database
}
translated["dwh"] = {
"type": "thoth_rest",
"database": identity,
"endpoint": rest,
}
else:
connection = database
connection.pop("transport", None)
translated["dwh"] = {
"type": "postgres_direct",
"connection": connection,
}
if "vectors" not in translated:
vector_db = raw.get("vector_db")
reader = raw.get("vector_rest")
writer = raw.get("vector_write_rest")
if reader is not None or writer is not None:
translated["vectors"] = {
"type": "thoth_vector_http",
"reader": reader,
"writer": writer,
"direct": vector_db,
}
elif vector_db is not None:
translated["vectors"] = {
"type": "pgvector_direct",
"connection": vector_db,
}
if "roots" not in translated and "paths" in raw:
translated["roots"] = deepcopy(raw["paths"])
return translated, True
+1
View File
@@ -0,0 +1 @@
"""Canonical, transport-independent Evidence corpus."""
+84
View File
@@ -0,0 +1,84 @@
"""Versioned deterministic chunking for canonical corpus documents."""
import hashlib
import json
import re
from dataclasses import asdict, dataclass
from tht.corpus.models import CanonicalChunk, CanonicalDocument
@dataclass(frozen=True, slots=True)
class ChunkPolicy:
version: str
max_chars: int
def __post_init__(self) -> None:
if not self.version:
raise ValueError("chunk policy version must not be empty")
if self.max_chars <= 0:
raise ValueError("max_chars must be greater than zero")
def _hash(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def _contents(content: str, maximum: int) -> list[str]:
result: list[str] = []
start = 0
while start < len(content):
end = min(start + maximum, len(content))
if end < len(content):
boundaries = list(re.finditer(r"\s+", content[start:end]))
if boundaries:
end = start + boundaries[-1].end()
result.append(content[start:end])
start = end
return result
def _policy_fingerprint(policy: ChunkPolicy) -> str:
serialized = json.dumps(asdict(policy), ensure_ascii=False, sort_keys=True, separators=(",", ":"))
return f"sha256:{_hash(serialized)}"
def chunk(document: CanonicalDocument, policy: ChunkPolicy) -> list[CanonicalChunk]:
"""Split canonical text with stable character-count boundaries and identifiers."""
chunks: list[CanonicalChunk] = []
policy_fingerprint = _policy_fingerprint(policy)
for ordinal, content in enumerate(_contents(document.content, policy.max_chars)):
chunk_hash = f"sha256:{_hash(content)}"
identifier = _hash(
":".join(
(
document.document_id,
document.content_hash,
policy_fingerprint,
str(ordinal),
chunk_hash,
)
)
)
chunks.append(
CanonicalChunk(
chunk_id=f"chunk:{identifier}",
document_id=document.document_id,
ordinal=ordinal,
content=content,
content_hash=chunk_hash,
source_uri=document.source_uri,
pipeline_version=document.pipeline_version,
metadata={
"chunk_policy": {
"version": policy.version,
"max_chars": policy.max_chars,
"fingerprint": policy_fingerprint,
},
"document": document.model_dump(mode="json")["metadata"],
"source_fingerprint": document.source_fingerprint,
"title": document.title,
},
)
)
return chunks
+163
View File
@@ -0,0 +1,163 @@
"""Immutable records emitted by the Evidence preprocessing pipeline."""
import hashlib
import re
from collections.abc import Mapping
from datetime import UTC, datetime
from typing import Self
from pydantic import BaseModel, ConfigDict, Field, JsonValue, field_validator, model_validator
from tht.ports.evidence import (
canonical_provenance_uri,
normalize_aware_datetime,
validate_namespaced_value,
validate_safe_metadata,
)
_NAMESPACED_ID = re.compile(r"^[a-z][a-z0-9_-]*:[A-Za-z0-9._:-]+$")
_SHA256 = re.compile(r"^sha256:[0-9a-f]{64}$")
def _validate_namespaced_id(value: str) -> str:
if not _NAMESPACED_ID.fullmatch(value):
raise ValueError("identifier must be namespaced as '<kind>:<stable-value>'")
return value
def _validate_hash(value: str) -> str:
if not _SHA256.fullmatch(value):
raise ValueError("content hash must be 'sha256:' followed by 64 lowercase hex digits")
return value
def _require_content_hash(content: str, content_hash: str) -> None:
expected = f"sha256:{hashlib.sha256(content.encode('utf-8')).hexdigest()}"
if content_hash != expected:
raise ValueError("content_hash must match the exact canonical UTF-8 content")
class _CanonicalValue(BaseModel):
model_config = ConfigDict(
frozen=True, extra="forbid", validate_default=True, revalidate_instances="always"
)
def model_copy(self, *, update: Mapping[str, object] | None = None, deep: bool = False) -> Self:
"""Copy through full field and model validation, including manifest invariants."""
data = self.model_dump(round_trip=True)
if update:
data.update(update)
return type(self).model_validate(data)
class _WithMetadata(_CanonicalValue):
metadata: dict[str, JsonValue] = Field(default_factory=dict)
_frozen_metadata = field_validator("metadata")(validate_safe_metadata)
class CanonicalDocument(_WithMetadata):
"""Normalized text whose hash covers the exact stored UTF-8 content bytes."""
document_id: str
source_id: str
source_uri: str
source_fingerprint: str = Field(min_length=1)
content_hash: str
title: str = ""
content: str
media_type: str = "text/plain"
modified_at: datetime | None = None
pipeline_version: str = Field(min_length=1)
_document_id = field_validator("document_id")(_validate_namespaced_id)
_source_id = field_validator("source_id")(_validate_namespaced_id)
_source_uri = field_validator("source_uri")(canonical_provenance_uri)
_source_fingerprint = field_validator("source_fingerprint")(validate_namespaced_value)
_content_hash = field_validator("content_hash")(_validate_hash)
_modified_at = field_validator("modified_at")(normalize_aware_datetime)
@model_validator(mode="after")
def content_hash_matches(self) -> "CanonicalDocument":
_require_content_hash(self.content, self.content_hash)
return self
class CanonicalChunk(_WithMetadata):
"""Chunk text whose hash covers the exact stored UTF-8 content bytes."""
chunk_id: str
document_id: str
ordinal: int = Field(ge=0)
content: str
content_hash: str
source_uri: str
pipeline_version: str = Field(min_length=1)
_chunk_id = field_validator("chunk_id")(_validate_namespaced_id)
_document_id = field_validator("document_id")(_validate_namespaced_id)
_content_hash = field_validator("content_hash")(_validate_hash)
_source_uri = field_validator("source_uri")(canonical_provenance_uri)
@model_validator(mode="after")
def content_hash_matches(self) -> "CanonicalChunk":
_require_content_hash(self.content, self.content_hash)
return self
class CorpusManifest(_WithMetadata):
"""Description of one internally consistent publishable generation."""
schema_version: int = Field(default=1, ge=1)
manifest_id: str | None = None
created_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
pipeline_version: str = Field(default="evidence-v1", min_length=1)
embedding_model: str | None = None
embedding_dimensions: int | None = Field(default=None, gt=0)
vector_generation: str | None = None
documents: tuple[CanonicalDocument, ...] = Field(default_factory=tuple)
chunks: tuple[CanonicalChunk, ...] = Field(default_factory=tuple)
_manifest_id = field_validator("manifest_id")(
lambda value: _validate_namespaced_id(value) if value is not None else None
)
_vector_generation = field_validator("vector_generation")(
lambda value: _validate_namespaced_id(value) if value is not None else None
)
_created_at = field_validator("created_at")(normalize_aware_datetime)
@model_validator(mode="after")
def validate_generation(self) -> "CorpusManifest":
if (self.embedding_model is None) != (self.embedding_dimensions is None):
raise ValueError("embedding_model and embedding_dimensions must be set together")
if self.vector_generation is not None and self.embedding_model is None:
raise ValueError("vector_generation requires embedding model and dimension compatibility")
document_ids = [document.document_id for document in self.documents]
source_ids = [document.source_id for document in self.documents]
chunk_ids = [chunk.chunk_id for chunk in self.chunks]
self._require_unique("document_id", document_ids)
self._require_unique("source_id", source_ids)
self._require_unique("chunk_id", chunk_ids)
documents = {document.document_id: document for document in self.documents}
ordinals: dict[str, list[int]] = {}
for document in self.documents:
if document.pipeline_version != self.pipeline_version:
raise ValueError("document pipeline_version must match manifest pipeline_version")
for chunk in self.chunks:
document = documents.get(chunk.document_id)
if document is None:
raise ValueError(f"chunk references unknown document: {chunk.document_id}")
if chunk.pipeline_version != self.pipeline_version:
raise ValueError("chunk pipeline_version must match manifest pipeline_version")
if chunk.source_uri != document.source_uri:
raise ValueError("chunk source_uri must match its document provenance")
ordinals.setdefault(chunk.document_id, []).append(chunk.ordinal)
for document_id, values in ordinals.items():
if sorted(values) != list(range(len(values))):
raise ValueError(f"chunk ordinals must be unique and contiguous for {document_id}")
return self
@staticmethod
def _require_unique(field: str, values: list[str]) -> None:
if len(values) != len(set(values)):
raise ValueError(f"{field} values must be unique")
+147
View File
@@ -0,0 +1,147 @@
"""Pure, deterministic conversion of acquired bytes into canonical text."""
import hashlib
import re
import unicodedata
from collections.abc import Mapping
import yaml
from pydantic import JsonValue, TypeAdapter, ValidationError
from yaml.events import AliasEvent
from yaml.nodes import MappingNode
from tht.corpus.models import CanonicalDocument
from tht.ports.evidence import AcquiredDocument, canonical_provenance_uri
MAX_DOCUMENT_BYTES = 10 * 1024 * 1024
_CHARSET = re.compile(r"(?:^|;)\s*charset\s*=\s*[\"']?([^;\s\"']+)", re.IGNORECASE)
_FRONTMATTER = re.compile(r"\A---\n(.*?)\n---(?:\n|\Z)", re.DOTALL)
_JSON_OBJECT = TypeAdapter(dict[str, JsonValue])
_MAX_FRONTMATTER_DEPTH = 20
_MAX_FRONTMATTER_NODES = 1000
class _FrontmatterLoader(yaml.SafeLoader):
"""SafeLoader with bounded structure and no YAML graph features."""
def __init__(self, stream) -> None:
super().__init__(stream)
self._depth = 0
self._nodes = 0
def compose_node(self, parent, index):
event = self.peek_event()
if isinstance(event, AliasEvent) or getattr(event, "anchor", None) is not None:
raise yaml.constructor.ConstructorError(None, None, "aliases are not allowed")
self._depth += 1
self._nodes += 1
if self._depth > _MAX_FRONTMATTER_DEPTH or self._nodes > _MAX_FRONTMATTER_NODES:
raise yaml.constructor.ConstructorError(None, None, "frontmatter is too complex")
try:
return super().compose_node(parent, index)
finally:
self._depth -= 1
def construct_mapping(self, node, deep=False):
if not isinstance(node, MappingNode):
return super().construct_mapping(node, deep=deep)
seen: set[object] = set()
for key_node, _ in node.value:
key = self.construct_object(key_node, deep=deep)
try:
duplicate = key in seen
seen.add(key)
except TypeError as error:
raise yaml.constructor.ConstructorError(
None, None, "mapping keys must be scalar"
) from error
if duplicate:
raise yaml.constructor.ConstructorError(None, None, "duplicate mapping key")
return super().construct_mapping(node, deep=deep)
class PermanentNormalizationError(ValueError):
"""A deterministic input failure which retrying cannot repair."""
def __init__(self, reason: str) -> None:
super().__init__(f"document normalization failed: {reason}")
self.reason = reason
self.permanent = True
def _sha256(value: str) -> str:
return hashlib.sha256(value.encode("utf-8")).hexdigest()
def _decode(acquired: AcquiredDocument) -> str:
if len(acquired.content) > MAX_DOCUMENT_BYTES:
raise PermanentNormalizationError("oversized")
media_type = acquired.media_type or "text/plain"
charset = _CHARSET.search(media_type)
if charset and charset.group(1).lower().replace("_", "-") not in {
"utf-8",
"utf8",
"us-ascii",
"ascii",
}:
raise PermanentNormalizationError("unsupported_charset")
try:
return acquired.content.decode("utf-8-sig", errors="strict")
except UnicodeDecodeError as error:
raise PermanentNormalizationError("undecodable") from error
def _frontmatter(text: str) -> tuple[dict[str, JsonValue], str]:
match = _FRONTMATTER.match(text)
if match is None:
return {}, text
try:
loaded = yaml.load(match.group(1), Loader=_FrontmatterLoader)
if loaded is None:
loaded = {}
if not isinstance(loaded, Mapping):
raise TypeError("frontmatter is not a mapping")
metadata = _JSON_OBJECT.validate_python(dict(loaded))
except (TypeError, UnicodeError, ValidationError, yaml.YAMLError) as error:
raise PermanentNormalizationError("invalid_frontmatter") from error
return metadata, text[match.end() :]
def normalize(acquired: AcquiredDocument, pipeline_version: str) -> CanonicalDocument:
"""Normalize one transport result without I/O or implicit data loss."""
if not pipeline_version:
raise ValueError("pipeline_version must not be empty")
decoded = _decode(acquired)
canonical = unicodedata.normalize("NFC", decoded.replace("\r\n", "\n").replace("\r", "\n"))
frontmatter, content = _frontmatter(canonical)
source_uri = canonical_provenance_uri(acquired.source.uri)
identity = f"{acquired.source.source_id}\n{source_uri}"
media_type = (acquired.media_type or "text/plain").split(";", 1)[0].strip().lower()
metadata: dict[str, JsonValue] = {
"source": acquired.source.model_dump(mode="json")["metadata"],
"acquisition": acquired.model_dump(mode="json")["metadata"],
}
if frontmatter:
metadata["frontmatter"] = frontmatter
try:
return CanonicalDocument(
document_id=f"doc:{_sha256(identity)}",
source_id=acquired.source.source_id,
source_uri=source_uri,
source_fingerprint=acquired.source.fingerprint,
content_hash=f"sha256:{_sha256(content)}",
title=str(frontmatter.get("title", "")),
content=content,
media_type=media_type,
modified_at=acquired.source.modified_at,
pipeline_version=pipeline_version,
metadata=metadata,
)
except ValidationError as error:
if frontmatter:
raise PermanentNormalizationError("invalid_frontmatter") from error
raise
+806
View File
@@ -0,0 +1,806 @@
"""Incremental Evidence preprocessing with generation-isolated vector writes."""
from __future__ import annotations
import hashlib
import json
import re
import uuid
from collections.abc import Mapping, Sequence
from dataclasses import asdict, dataclass, field
from datetime import UTC
from pathlib import Path
from tht.corpus.chunk import ChunkPolicy, chunk
from tht.corpus.models import CanonicalChunk, CanonicalDocument, CorpusManifest
from tht.corpus.normalize import normalize
from tht.corpus.store import CorpusStore
from tht.ports.evidence import EvidenceSource, SourceObject, canonical_provenance_uri
from tht.ports.vector import VectorStore, VectorWriteRecord
from tht.vectorstore.records import VectorRecord
from tht.jobs.models import JobSpec
from tht.jobs.runner import JobContext, StageArtifacts, run_job, seal_stage_artifacts
EVIDENCE_STAGE_IDS = (
"discover",
"acquire_normalize_chunk",
"embed",
"vector_upsert",
"stage_validate",
"publish",
"retention_cleanup",
)
class PipelineError(RuntimeError):
"""Credential-free failure at the preprocessing boundary."""
@dataclass(frozen=True)
class PipelineResult:
status: str
generation: str | None
published: bool
changed: tuple[str, ...]
unchanged: tuple[str, ...]
removed: tuple[str, ...]
manifest: CorpusManifest = field(repr=False)
run_id: str | None = None
resumed_from: str | None = None
def __repr__(self) -> str:
counts = {
"changed": len(self.changed),
"unchanged": len(self.unchanged),
"removed": len(self.removed),
}
return (
f"PipelineResult(status={self.status!r}, generation={self.generation!r}, "
f"published={self.published!r}, counts={counts!r}, "
f"run_id={self.run_id!r}, resumed_from={self.resumed_from!r})"
)
def model_dump(self, mode=None):
def bounded(values: tuple[str, ...]) -> list[str]:
return [value[:200] for value in values[:100]]
return {
"status": self.status,
"generation": self.generation,
"published": self.published,
"changed": bounded(self.changed),
"unchanged": bounded(self.unchanged),
"removed": bounded(self.removed),
"counts": {
"changed": len(self.changed),
"unchanged": len(self.unchanged),
"removed": len(self.removed),
"documents": len(self.manifest.documents),
"chunks": len(self.manifest.chunks),
},
"manifest_id": self.manifest.manifest_id,
"run_id": self.run_id,
"resumed_from": self.resumed_from,
}
def _fingerprint(value) -> str:
payload = json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
return "sha256:" + hashlib.sha256(payload.encode()).hexdigest()
def _canonical_json(value):
if isinstance(value, Mapping):
return {str(key): _canonical_json(value[key]) for key in sorted(value)}
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
return [_canonical_json(child) for child in value]
return value
def _source_snapshot(discovered) -> dict[str, dict]:
snapshot = {}
for _, item in discovered:
modified_at = item.modified_at.astimezone(UTC) if item.modified_at else None
metadata = _canonical_json(item.metadata)
snapshot[item.source_id] = {
"source_id": item.source_id,
"uri": item.uri,
"fingerprint": item.fingerprint,
"modified_at": modified_at.isoformat().replace("+00:00", "Z") if modified_at else None,
"metadata": metadata,
"media_type": metadata.get("media_type"),
"size": metadata.get("size"),
}
return snapshot
class CorpusPipeline:
def __init__(
self, *, store: CorpusStore, sources: list[EvidenceSource], embedder,
vector_store: VectorStore, embedding_model: str, embedding_dimensions: int,
chunk_policy: ChunkPolicy, pipeline_version: str, retain_published_generations: int = 3,
workspace_id: str | None = None,
) -> None:
self.store = store
self.sources = sources
self.embedder = embedder
self.vector_store = vector_store
self.embedding_model = embedding_model
self.embedding_dimensions = embedding_dimensions
self.chunk_policy = chunk_policy
self.pipeline_version = pipeline_version
if isinstance(retain_published_generations, bool) or retain_published_generations < 1:
raise ValueError("retain_published_generations must be at least 1")
self.retain_published_generations = retain_published_generations
self.workspace_id = workspace_id
def _assert_workspace_binding(self) -> None:
manifest = self.store.active_manifest()
if manifest is None:
if self.workspace_id is None:
self.workspace_id = "default"
return
persisted = manifest.metadata.get("workspace_id")
if not isinstance(persisted, str) or re.fullmatch(
r"[a-z][a-z0-9_-]{0,63}", persisted
) is None:
raise PipelineError(
"corpus workspace ownership is missing or invalid; use a new corpus root or explicit rebuild"
)
if self.workspace_id is None and isinstance(persisted, str):
self.workspace_id = persisted
return
if persisted != self.workspace_id:
raise PipelineError(
"corpus belongs to a different workspace; use a new corpus root or explicit rebuild"
)
def _protected_generations(self, workspace_root: Path) -> set[str]:
protected = {value for value in (self.store.active_generation(),) if value}
runs = workspace_root / ".tht-jobs" / "evidence" / "runs"
for checkpoint in runs.glob("*/checkpoint.json") if runs.exists() else ():
try:
state = json.loads(checkpoint.read_text(encoding="utf-8"))
if state.get("status") not in {"running", "failed"}:
continue
plan = checkpoint.parent / "artifacts" / "plan.json"
generation = json.loads(plan.read_text(encoding="utf-8")).get("generation")
if isinstance(generation, str):
protected.add(generation)
except (OSError, ValueError):
continue
return protected
def gc(self, *, workspace_root: Path, dry_run: bool = False) -> dict:
with self.store.writer_lock():
return self._gc(workspace_root=workspace_root, dry_run=dry_run)
def _gc(self, *, workspace_root: Path, dry_run: bool = False) -> dict:
self._assert_workspace_binding()
published = self.store.published_generations()
list_vectors = getattr(self.vector_store, "list_evidence_generations", None)
vector_generations = set(list_vectors("evidence", self.workspace_id)) if list_vectors else set()
generations = sorted(set(published) | vector_generations)
job_protected = self._protected_generations(workspace_root)
active = self.store.active_generation()
rollback_count = self.retain_published_generations - 1
rollback = [generation for generation in published if generation != active]
keep = ({active} if active else set()) | set(rollback[-rollback_count:] if rollback_count else ())
fs_keep = keep | job_protected
vector_protected = set(fs_keep)
for generation in fs_keep:
try:
manifest = self.store.manifest(generation)
except (OSError, ValueError):
continue
vector_protected.update(
value for value in manifest.metadata.get("document_generations", {}).values()
if isinstance(value, str)
)
evicted, failures = [], []
filesystem_generations = set(self.store.list_generations())
for generation in generations:
purge_vector = generation not in vector_protected
purge_filesystem = generation in filesystem_generations and generation not in fs_keep
if not purge_vector and not purge_filesystem:
continue
if dry_run:
evicted.append(generation)
continue
if purge_vector:
try:
self.vector_store.delete_generation("evidence", generation, self.workspace_id)
except Exception:
failures.append({"generation": generation, "error": "vector cleanup failed"})
continue
try:
if purge_filesystem:
self.store.discard(generation)
evicted.append(generation)
except Exception:
failures.append({"generation": generation, "error": "filesystem cleanup failed"})
return {"status": "partial" if failures else "succeeded", "dry_run": dry_run,
"active_generation": self.store.active_generation(), "evicted": evicted,
"protected": sorted(vector_protected), "failures": failures}
def _discover(self) -> list[tuple[EvidenceSource, SourceObject]]:
discovered = []
seen = set()
for source in self.sources:
for item in source.discover():
if item.source_id in seen:
raise PipelineError("duplicate Evidence source identity")
seen.add(item.source_id)
discovered.append((source, item))
return sorted(discovered, key=lambda pair: pair[1].source_id)
def run(self, *, dry_run: bool = False, resume: str | None = None) -> PipelineResult:
with self.store.writer_lock():
self._assert_workspace_binding()
return self._run(dry_run=dry_run, resume=resume)
def run_as_job(self, **kwargs) -> PipelineResult:
self.workspace_id = kwargs["workspace_id"]
with self.store.writer_lock():
self._assert_workspace_binding()
return self._run_as_job(**kwargs)
def _run_as_job(
self,
*,
workspace_id: str,
workspace_root: Path,
config_fingerprint: str,
input_fingerprint: str,
dry_run: bool = False,
resume_run_id: str | None = None,
after_stage_return=None,
) -> PipelineResult:
"""Execute preprocessing through the durable shared job envelope."""
discovered = self._discover()
discovered_fingerprint = _fingerprint(
{item.source_id: item.fingerprint for _, item in discovered}
)
source_snapshot = _source_snapshot(discovered)
source_by_id = {item.source_id: (source, item) for source, item in discovered}
compatibility = _fingerprint({
"pipeline": self.pipeline_version,
"model": self.embedding_model,
"dimensions": self.embedding_dimensions,
"chunk_policy": asdict(self.chunk_policy),
})
job_binding = {
"config_fingerprint": config_fingerprint,
"input_fingerprint": input_fingerprint,
"compatibility_fingerprint": compatibility,
"pipeline_version": self.pipeline_version,
"chunk_policy_version": self.chunk_policy.version,
"embedding_model": self.embedding_model,
"embedding_dimensions": self.embedding_dimensions,
}
previous = self.store.active_manifest()
def document_sources(manifest: CorpusManifest) -> dict[str, dict]:
return {
document.document_id: {
"document_id": document.document_id,
"source_id": document.source_id,
"source_uri": document.source_uri,
"source_fingerprint": document.source_fingerprint,
"modified_at": (
document.modified_at.isoformat().replace("+00:00", "Z")
if document.modified_at else None
),
"source_metadata": _canonical_json(document.metadata.get("source")),
"media_type": document.media_type,
"content_hash": document.content_hash,
"pipeline_version": document.pipeline_version,
}
for document in manifest.documents
}
def active_assets_are_valid(manifest: CorpusManifest | None) -> bool:
if manifest is None or manifest.metadata.get("workspace_id") != workspace_id:
return False
actual_documents = {document.source_id: document for document in manifest.documents}
persisted_snapshot = _canonical_json(manifest.metadata.get("source_snapshot"))
if (
not isinstance(persisted_snapshot, dict)
or set(persisted_snapshot) != set(actual_documents)
or manifest.metadata.get("compatibility_fingerprint") != compatibility
or _canonical_json(manifest.metadata.get("document_sources"))
!= document_sources(manifest)
):
return False
for source_id, document in actual_documents.items():
source_payload = persisted_snapshot[source_id]
source = SourceObject.model_validate({
"source_id": source_payload["source_id"],
"uri": source_payload["uri"],
"fingerprint": source_payload["fingerprint"],
"modified_at": source_payload["modified_at"],
"metadata": source_payload["metadata"],
})
content = self.store.read_document(document.document_id, manifest.manifest_id)
expected_uri = canonical_provenance_uri(source.uri)
expected_id = "doc:" + hashlib.sha256(
f"{source.source_id}\n{expected_uri}".encode()
).hexdigest()
expected_media_type = source_payload.get("media_type")
if (
content != document.content
or document.document_id != expected_id
or document.source_id != source.source_id
or document.source_uri != expected_uri
or document.source_fingerprint != source.fingerprint
or document.modified_at != source.modified_at
or _canonical_json(document.metadata.get("source"))
!= _canonical_json(source.metadata)
or (
isinstance(expected_media_type, str)
and document.media_type != expected_media_type
)
or document.pipeline_version != self.pipeline_version
):
return False
expected_chunks = tuple(
part for document in manifest.documents for part in chunk(document, self.chunk_policy)
)
if any(document.content and not chunk(document, self.chunk_policy)
for document in manifest.documents):
return False
if _canonical_json([part.model_dump(mode="json") for part in manifest.chunks]) != (
_canonical_json([part.model_dump(mode="json") for part in expected_chunks])
):
return False
generations = manifest.metadata.get("document_generations")
if not isinstance(generations, Mapping):
return False
health = self.vector_store.health()
if (
not health.ok
or health.dimension_compatible is not True
or health.expected_dimension != self.embedding_dimensions
or health.observed_dimensions != (self.embedding_dimensions,)
):
return False
existing = self.vector_store.existing_hashes("evidence", ["evidence"])
for part in expected_chunks:
generation = generations.get(part.document_id)
if not isinstance(generation, str):
return False
record_id = f"{workspace_id}:{generation}:{part.chunk_id}"
if existing.get(record_id) != part.content_hash:
return False
return True
try:
active_assets_valid = active_assets_are_valid(previous)
except Exception:
active_assets_valid = False
reusable = (
active_assets_valid
and _canonical_json(previous.metadata.get("source_snapshot")) == source_snapshot
and _canonical_json(previous.metadata.get("job_binding")) == job_binding
)
if not dry_run and resume_run_id is None and reusable:
return PipelineResult(
"succeeded", previous.manifest_id, False, (),
tuple(sorted(item.source_id for _, item in discovered)), (), previous,
)
spec = JobSpec(
workspace_id=workspace_id,
job_type="evidence",
workspace_root=workspace_root,
spec_version="jobs-v1",
pipeline_version=self.pipeline_version,
config_fingerprint=config_fingerprint,
input_fingerprint=_fingerprint([input_fingerprint, discovered_fingerprint]),
stage_ids=EVIDENCE_STAGE_IDS,
dry_run=dry_run,
resume_run_id=resume_run_id,
)
def artifact(context: JobContext, name: str) -> Path:
root = context.run_dir / "artifacts"
root.mkdir(exist_ok=True)
return root / name
def write(context: JobContext, name: str, value) -> None:
artifact(context, name).write_text(
json.dumps(value, sort_keys=True, separators=(",", ":")), encoding="utf-8"
)
def read(context: JobContext, name: str):
try:
return json.loads(artifact(context, name).read_text(encoding="utf-8"))
except (OSError, ValueError) as error:
raise PipelineError("preprocessing checkpoint artifact is corrupt") from error
def discover_stage(context: JobContext) -> None:
previous = self.store.active_manifest()
prior = {doc.source_id: doc for doc in previous.documents} if previous else {}
fingerprints = {item.source_id: item.fingerprint for _, item in discovered}
previous_snapshot = (
_canonical_json(previous.metadata.get("source_snapshot")) if previous else {}
)
rebuild = bool(previous and not active_assets_valid)
changed = sorted(
item.source_id for _, item in discovered
if rebuild or item.source_id not in prior
or previous_snapshot.get(item.source_id) != source_snapshot[item.source_id]
)
unchanged = sorted(set(fingerprints) - set(changed))
removed = sorted(set(prior) - set(fingerprints))
write(context, "plan.json", {
"generation": f"gen:{context.run_id}",
"compatibility": compatibility,
"job_binding": job_binding,
"source_snapshot": source_snapshot,
"fingerprints": fingerprints,
"changed": changed,
"unchanged": unchanged,
"removed": removed,
"previous": previous.model_dump(mode="json") if previous else None,
})
return StageArtifacts(("plan.json",))
def acquire_stage(context: JobContext) -> None:
if context.dry_run:
return StageArtifacts()
plan = read(context, "plan.json")
previous = CorpusManifest.model_validate(plan["previous"]) if plan["previous"] else None
prior = {doc.source_id: doc for doc in previous.documents} if previous else {}
documents = [prior[source_id] for source_id in plan["unchanged"]]
for source_id in plan["changed"]:
source, item = source_by_id[source_id]
documents.append(normalize(source.acquire(item), self.pipeline_version))
documents.sort(key=lambda value: value.source_id)
chunks = [part for document in documents for part in chunk(document, self.chunk_policy)]
previous_generations = dict(previous.metadata.get("document_generations", {})) if previous else {}
changed = set(plan["changed"])
generations = {
document.document_id: (
plan["generation"] if document.source_id in changed
else previous_generations.get(document.document_id, previous.vector_generation)
) for document in documents
}
manifest = CorpusManifest(
pipeline_version=self.pipeline_version,
embedding_model=self.embedding_model,
embedding_dimensions=self.embedding_dimensions,
vector_generation=plan["generation"],
documents=tuple(documents), chunks=tuple(chunks),
metadata={
"workspace_id": self.workspace_id,
"compatibility_fingerprint": compatibility,
"job_binding": plan["job_binding"],
"source_snapshot": plan["source_snapshot"],
"document_sources": {
document.document_id: {
"document_id": document.document_id,
"source_id": document.source_id,
"source_uri": document.source_uri,
"source_fingerprint": document.source_fingerprint,
"modified_at": (
document.modified_at.isoformat().replace("+00:00", "Z")
if document.modified_at else None
),
"source_metadata": _canonical_json(
document.metadata.get("source")
),
"media_type": document.media_type,
"content_hash": document.content_hash,
"pipeline_version": document.pipeline_version,
}
for document in documents
},
"fingerprints": plan["fingerprints"],
"removed": plan["removed"],
"document_generations": generations,
},
)
write(context, "manifest.json", manifest.model_dump(mode="json"))
return StageArtifacts(("manifest.json",))
def embed_stage(context: JobContext) -> None:
if context.dry_run:
return StageArtifacts()
plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
changed_docs = {doc.document_id for doc in manifest.documents if doc.source_id in plan["changed"]}
parts = [part for part in manifest.chunks if part.document_id in changed_docs]
embeddings = self.embedder.embed_documents([part.content for part in parts])
if len(embeddings) != len(parts) or any(
len(vector) != self.embedding_dimensions for vector in embeddings
):
raise PipelineError("embedding output is incompatible")
write(context, "embeddings.json", embeddings)
return StageArtifacts(("embeddings.json",))
def records(context: JobContext):
plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
changed_docs = {doc.document_id for doc in manifest.documents if doc.source_id in plan["changed"]}
parts = [part for part in manifest.chunks if part.document_id in changed_docs]
embeddings = read(context, "embeddings.json")
return [self._vector_record(part, vector, plan["generation"], self.workspace_id)
for part, vector in zip(parts, embeddings, strict=True)]
def compensate(context: JobContext) -> None:
generation = read(context, "plan.json")["generation"]
if self.store.active_generation() != generation:
self.store.discard(generation)
try:
self.vector_store.delete_generation("evidence", generation, self.workspace_id)
except Exception:
pass
write(context, "compensated.json", {"generation": generation})
def rotate_compensated_generation(context: JobContext) -> None:
marker = artifact(context, "compensated.json")
if not marker.exists():
return
plan = read(context, "plan.json")
old = plan["generation"]
plan["generation"] = f"gen:{uuid.uuid4().hex}"
write(context, "plan.json", plan)
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
changed = set(plan["changed"])
generations = dict(manifest.metadata["document_generations"])
for document in manifest.documents:
if document.source_id in changed and generations.get(document.document_id) == old:
generations[document.document_id] = plan["generation"]
manifest_payload = manifest.model_dump(mode="json")
manifest_payload["metadata"]["document_generations"] = generations
manifest_payload["vector_generation"] = plan["generation"]
manifest = CorpusManifest.model_validate(manifest_payload)
write(context, "manifest.json", manifest.model_dump(mode="json"))
marker.unlink()
def vector_stage(context: JobContext) -> None:
if context.dry_run:
return StageArtifacts()
rotate_compensated_generation(context)
values = records(context)
write(context, "vector-intent.json", {
"generation": read(context, "plan.json")["generation"],
"records": {value.record.id: value.content_hash for value in values},
})
seal_stage_artifacts(
context, "vector_upsert",
("plan.json", "manifest.json", "vector-intent.json"), spec,
)
try:
existing = self.vector_store.existing_hashes("evidence", ["evidence"])
missing = [
value for value in values
if existing.get(value.record.id) != value.content_hash
]
if missing and self.vector_store.upsert("evidence", missing) != len(missing):
raise PipelineError("vector write count mismatch")
except Exception:
compensate(context)
raise
return StageArtifacts(("plan.json", "manifest.json", "vector-intent.json"))
def stage_stage(context: JobContext) -> None:
if context.dry_run:
return StageArtifacts()
plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
recovered = False
try:
if artifact(context, "compensated.json").exists():
recovered = True
rotate_compensated_generation(context)
values = records(context)
existing = self.vector_store.existing_hashes("evidence", ["evidence"])
missing = [value for value in values if existing.get(value.record.id) != value.content_hash]
if missing and self.vector_store.upsert("evidence", missing) != len(missing):
raise PipelineError("vector write count mismatch")
plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
if not self.store.generation_path(plan["generation"]).exists():
self.store.stage(
manifest, {doc.document_id: doc.content for doc in manifest.documents},
generation=plan["generation"],
)
self.store.manifest(plan["generation"])
except Exception:
compensate(context)
raise
return StageArtifacts(
("plan.json", "manifest.json", "vector-intent.json") if recovered else ()
)
def publish_stage(context: JobContext) -> None:
if context.dry_run:
return StageArtifacts()
if artifact(context, "compensated.json").exists():
rotate_compensated_generation(context)
values = records(context)
try:
existing = self.vector_store.existing_hashes("evidence", ["evidence"])
missing = [value for value in values if existing.get(value.record.id) != value.content_hash]
if missing and self.vector_store.upsert("evidence", missing) != len(missing):
raise PipelineError("vector write count mismatch")
except Exception:
compensate(context)
raise
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
generation = read(context, "plan.json")["generation"]
try:
if not self.store.generation_path(generation).exists():
self.store.stage(
manifest, {doc.document_id: doc.content for doc in manifest.documents},
generation=generation,
)
except Exception:
compensate(context)
raise
generation = read(context, "plan.json")["generation"]
try:
self.store.publish(generation)
except Exception:
compensate(context)
raise
return StageArtifacts(("plan.json", "manifest.json", "vector-intent.json"))
def retention_stage(context: JobContext) -> None:
if not context.dry_run:
self.gc(workspace_root=workspace_root)
report = run_job(spec, [
discover_stage, acquire_stage, embed_stage, vector_stage,
stage_stage, publish_stage, retention_stage,
], after_stage_return=after_stage_return)
run_dir = workspace_root / ".tht-jobs" / "evidence" / "runs" / report.run_id
plan = json.loads((run_dir / "artifacts" / "plan.json").read_text())
if dry_run:
manifest = self.store.active_manifest() or CorpusManifest(pipeline_version=self.pipeline_version)
generation = None
published = False
elif report.status == "succeeded":
generation = plan["generation"]
manifest = self.store.manifest(generation)
published = True
else:
generation = plan["generation"]
manifest_path = run_dir / "artifacts" / "manifest.json"
manifest = (CorpusManifest.model_validate_json(manifest_path.read_text())
if manifest_path.exists() else CorpusManifest(pipeline_version=self.pipeline_version))
published = False
return PipelineResult(
report.status, generation, published, tuple(plan["changed"]),
tuple(plan["unchanged"]), tuple(plan["removed"]), manifest,
report.run_id, report.resumed_from,
)
def _run(self, *, dry_run: bool = False, resume: str | None = None) -> PipelineResult:
generation = None
vector_written = False
previous = self.store.active_manifest()
try:
discovered = self._discover()
except Exception as error:
raise PipelineError("Evidence discovery failed") from error
prior_documents = {doc.source_id: doc for doc in previous.documents} if previous else {}
fingerprints = {item.source_id: item.fingerprint for _, item in discovered}
compatibility = _fingerprint({
"pipeline": self.pipeline_version, "model": self.embedding_model,
"dimensions": self.embedding_dimensions, "chunk_policy": asdict(self.chunk_policy),
})
previous_compatibility = previous.metadata.get("compatibility_fingerprint") if previous else None
rebuild = previous is not None and compatibility != previous_compatibility
changed = tuple(item.source_id for _, item in discovered if rebuild or prior_documents.get(item.source_id) is None or prior_documents[item.source_id].source_fingerprint != item.fingerprint)
unchanged = tuple(item.source_id for _, item in discovered if item.source_id not in changed)
removed = tuple(sorted(set(prior_documents) - set(fingerprints)))
if dry_run:
manifest = previous or CorpusManifest(pipeline_version=self.pipeline_version)
return PipelineResult("succeeded", None, False, changed, unchanged, removed, manifest)
if previous is not None and not changed and not removed:
return PipelineResult(
"succeeded", previous.manifest_id, False, changed, unchanged, removed, previous
)
documents: list[CanonicalDocument] = [prior_documents[source_id] for source_id in unchanged]
changed_set = set(changed)
try:
for source, item in discovered:
if item.source_id in changed_set:
documents.append(normalize(source.acquire(item), self.pipeline_version))
documents.sort(key=lambda document: document.source_id)
chunks: list[CanonicalChunk] = []
for document in documents:
chunks.extend(chunk(document, self.chunk_policy))
generation = resume or f"gen:{uuid.uuid4().hex}"
previous_generations = dict(previous.metadata.get("document_generations", {})) if previous else {}
document_generations = {
document.document_id: (
generation if document.source_id in changed_set
else previous_generations.get(document.document_id, previous.vector_generation)
)
for document in documents
}
manifest = CorpusManifest(
pipeline_version=self.pipeline_version,
embedding_model=self.embedding_model,
embedding_dimensions=self.embedding_dimensions,
vector_generation=generation,
documents=tuple(documents), chunks=tuple(chunks),
metadata={
"workspace_id": self.workspace_id,
"compatibility_fingerprint": compatibility,
"fingerprints": fingerprints,
"removed": list(removed),
"document_generations": document_generations,
},
)
changed_documents = {document.document_id for document in documents if document.source_id in changed_set}
changed_chunks = [part for part in chunks if part.document_id in changed_documents]
embeddings = self.embedder.embed_documents([part.content for part in changed_chunks])
if len(embeddings) != len(changed_chunks):
raise PipelineError("embedding count mismatch")
if any(len(vector) != self.embedding_dimensions for vector in embeddings):
raise PipelineError("embedding dimension mismatch")
records = [self._vector_record(part, vector, generation, self.workspace_id) for part, vector in zip(changed_chunks, embeddings, strict=True)]
if records:
written = self.vector_store.upsert("evidence", records)
vector_written = True
if written != len(records):
raise PipelineError("vector write count mismatch")
generation_path = self.store.generation_path(generation)
if resume is not None and generation_path.exists():
staged_manifest = self.store.manifest(generation)
expected = manifest.model_dump(mode="json", exclude={"created_at", "manifest_id"})
actual = staged_manifest.model_dump(mode="json", exclude={"created_at", "manifest_id"})
actual["metadata"].pop("files", None)
if actual != expected:
raise PipelineError("resume generation is incompatible")
staged = generation
else:
staged = self.store.stage(
manifest, {document.document_id: document.content for document in documents},
generation=generation,
)
self.store.publish(staged)
self.gc(workspace_root=self.store.root.parent)
except PipelineError:
self._compensate(generation, vector_written)
raise
except Exception as error:
self._compensate(generation, vector_written)
raise PipelineError("Evidence preprocessing failed") from error
return PipelineResult("succeeded", generation, True, changed, unchanged, removed, self.store.manifest(generation))
def _compensate(self, generation: str | None, vector_written: bool) -> None:
if generation is None:
return
try:
self.store.discard(generation)
except Exception:
pass
if vector_written:
try:
self.vector_store.delete_generation("evidence", generation, self.workspace_id)
except Exception:
pass
@staticmethod
def _vector_record(
chunk: CanonicalChunk, embedding: list[float], generation: str, workspace_id: str,
):
record = VectorRecord(
id=f"{workspace_id}:{generation}:{chunk.chunk_id}",
kind="evidence", ref=chunk.document_id,
title=str(chunk.metadata.get("title", "")), content=chunk.content,
metadata={
**dict(chunk.metadata), "document_id": chunk.document_id,
"workspace_id": workspace_id,
"source_uri": chunk.source_uri, "ordinal": chunk.ordinal,
"vector_generation": generation,
},
)
return VectorWriteRecord(record=record, embedding=embedding, content_hash=chunk.content_hash)
+272
View File
@@ -0,0 +1,272 @@
"""Durable immutable corpus generations and an atomic ACTIVE pointer."""
from __future__ import annotations
import json
import fcntl
import os
import re
import stat
import shutil
import uuid
import hashlib
import threading
from datetime import UTC, datetime
from pathlib import Path
from contextlib import contextmanager
from tht.corpus.models import CorpusManifest
_GENERATION = re.compile(r"^gen:[0-9a-f]{32}$")
class UnsafeCorpusPath(RuntimeError):
pass
def _atomic_write(path: Path, payload: bytes) -> None:
temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, 0o600)
try:
with os.fdopen(fd, "wb") as stream:
stream.write(payload)
stream.flush()
os.fsync(stream.fileno())
os.replace(temporary, path)
directory = os.open(path.parent, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
os.fsync(directory)
finally:
os.close(directory)
except BaseException:
temporary.unlink(missing_ok=True)
raise
class CorpusStore:
def __init__(self, root: Path) -> None:
self.root = Path(root)
self.active_path = self.root / "ACTIVE"
self._replace = os.replace
self._fsync_directory = self._sync_root
self._lock_state = threading.local()
self._ensure_root()
def _ensure_root(self) -> None:
if self.root.is_symlink():
raise UnsafeCorpusPath("corpus root must not be a symlink")
self.root.mkdir(parents=True, exist_ok=True, mode=0o700)
info = self.root.lstat()
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
raise UnsafeCorpusPath("corpus root is unsafe")
@contextmanager
def writer_lock(self):
depth = getattr(self._lock_state, "depth", 0)
if depth:
self._lock_state.depth = depth + 1
try:
yield
finally:
self._lock_state.depth -= 1
return
lock_path = self.root / ".writer.lock"
fd = os.open(lock_path, os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW | os.O_CLOEXEC, 0o600)
try:
info = os.fstat(fd)
if not stat.S_ISREG(info.st_mode) or info.st_uid != os.getuid() or info.st_nlink != 1:
raise UnsafeCorpusPath("corpus writer lock is unsafe")
fcntl.flock(fd, fcntl.LOCK_EX)
self._lock_state.depth = 1
yield
finally:
self._lock_state.depth = 0
fcntl.flock(fd, fcntl.LOCK_UN)
os.close(fd)
def generation_path(self, generation: str) -> Path:
if not _GENERATION.fullmatch(generation):
raise UnsafeCorpusPath("invalid corpus generation")
path = self.root / generation.replace(":", "-")
if path.is_symlink():
raise UnsafeCorpusPath("generation must not be a symlink")
return path
def stage(
self,
manifest: CorpusManifest,
materialized: dict[str, str],
*,
generation: str | None = None,
) -> str:
generation = generation or f"gen:{uuid.uuid4().hex}"
path = self.generation_path(generation)
try:
path.mkdir(mode=0o700)
except FileExistsError:
raise UnsafeCorpusPath("generation already exists") from None
documents = path / "documents"
documents.mkdir(mode=0o700)
files: dict[str, str] = {}
for document in manifest.documents:
relative = f"documents/{document.document_id.removeprefix('doc:')}.md"
_atomic_write(path / relative, materialized[document.document_id].encode("utf-8"))
files[document.document_id] = relative
payload = json.loads(manifest.model_dump_json())
metadata = payload["metadata"]
metadata["files"] = files
payload.update({"manifest_id": generation, "metadata": metadata})
staged = CorpusManifest.model_validate(payload)
_atomic_write(path / "manifest.json", (staged.model_dump_json(indent=2) + "\n").encode())
return generation
def publish(self, generation: str) -> str:
manifest = self.manifest(generation)
if manifest.manifest_id != generation:
raise UnsafeCorpusPath("manifest generation mismatch")
if self.active_generation() == generation:
return generation
previous = self.active_generation()
published_marker = self.generation_path(generation) / "PUBLISHED"
temporary = self.active_path.with_name(f".ACTIVE.{uuid.uuid4().hex}.tmp")
replaced = False
try:
_atomic_write(temporary, (generation + "\n").encode())
self._replace(temporary, self.active_path)
replaced = True
self._fsync_directory()
_atomic_write(
published_marker,
(datetime.now(UTC).isoformat().replace("+00:00", "Z") + "\n").encode("ascii"),
)
except BaseException:
temporary.unlink(missing_ok=True)
if replaced:
if previous is None:
self.active_path.unlink(missing_ok=True)
else:
rollback = self.active_path.with_name(f".ACTIVE.rollback.{uuid.uuid4().hex}.tmp")
_atomic_write(rollback, (previous + "\n").encode())
self._replace(rollback, self.active_path)
self._sync_root()
raise
return generation
def _sync_root(self) -> None:
directory = os.open(self.root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
os.fsync(directory)
finally:
os.close(directory)
def active_generation(self) -> str | None:
try:
if self.active_path.is_symlink():
raise UnsafeCorpusPath("ACTIVE must not be a symlink")
value = self.active_path.read_text(encoding="ascii").strip()
except FileNotFoundError:
return None
if not _GENERATION.fullmatch(value):
raise UnsafeCorpusPath("ACTIVE contains an invalid generation")
return value
def manifest(self, generation: str) -> CorpusManifest:
path = self.generation_path(generation)
manifest_path = path / "manifest.json"
if manifest_path.is_symlink():
raise UnsafeCorpusPath("manifest must not be a symlink")
return CorpusManifest.model_validate_json(manifest_path.read_text(encoding="utf-8"))
def discard(self, generation: str) -> None:
path = self.generation_path(generation)
if path.exists():
if path.is_symlink() or not stat.S_ISDIR(path.lstat().st_mode):
raise UnsafeCorpusPath("generation cleanup target is unsafe")
shutil.rmtree(path)
def active_manifest(self) -> CorpusManifest | None:
generation = self.active_generation()
return self.manifest(generation) if generation else None
def list_generations(self) -> list[str]:
values = []
for entry in self.root.iterdir():
match = re.fullmatch(r"gen-([0-9a-f]{32})", entry.name)
if match and not entry.is_symlink() and stat.S_ISDIR(entry.lstat().st_mode):
values.append(f"gen:{match.group(1)}")
return sorted(values, key=lambda value: self.generation_path(value).stat().st_mtime_ns)
def published_generations(self) -> list[str]:
active = self.active_generation()
published = []
for generation in self.list_generations():
path = self.generation_path(generation)
marker = path / "PUBLISHED"
if generation != active and not marker.is_file():
continue
try:
manifest = self.manifest(generation)
if manifest.manifest_id != generation:
continue
timestamp = marker.read_text(encoding="ascii").strip() if marker.is_file() else ""
key = (timestamp or manifest.created_at.isoformat(), generation)
published.append((key, generation))
except (OSError, ValueError):
continue
return [generation for _, generation in sorted(published)]
def resolve_document(self, document_id: str, generation: str | None = None) -> Path | None:
generation = generation or self.active_generation()
if generation is None:
return None
manifest = self.manifest(generation)
relative = manifest.metadata.get("files", {}).get(document_id)
if not isinstance(relative, str):
return None
parts = Path(relative).parts
if Path(relative).is_absolute() or parts[:1] != ("documents",) or len(parts) != 2:
raise UnsafeCorpusPath("materialized document path is unsafe")
return self.generation_path(generation) / relative
def read_document(self, document_id: str, generation: str | None = None) -> str | None:
generation = generation or self.active_generation()
if generation is None:
return None
manifest = self.manifest(generation)
path = self.resolve_document(document_id, generation)
document = next((item for item in manifest.documents if item.document_id == document_id), None)
if path is None or document is None:
return None
generation_fd = os.open(self.generation_path(generation), os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
documents_fd = fd = None
try:
documents_fd = os.open("documents", os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=generation_fd)
fd = os.open(path.name, os.O_RDONLY | os.O_NOFOLLOW | os.O_CLOEXEC, dir_fd=documents_fd)
info = os.fstat(fd)
if not stat.S_ISREG(info.st_mode) or info.st_uid != os.getuid() or info.st_nlink != 1:
raise UnsafeCorpusPath("materialized document is unsafe")
payload = os.read(fd, info.st_size + 1)
if len(payload) != info.st_size or "sha256:" + hashlib.sha256(payload).hexdigest() != document.content_hash:
raise UnsafeCorpusPath("materialized document hash mismatch")
return payload.decode("utf-8")
except (OSError, UnicodeError) as error:
raise UnsafeCorpusPath("materialized document read failed") from error
finally:
if fd is not None:
os.close(fd)
if documents_fd is not None:
os.close(documents_fd)
os.close(generation_fd)
def materialize_document(
self, document_id: str, destination: Path, generation: str | None = None,
) -> Path | None:
content = self.read_document(document_id, generation)
if content is None:
return None
destination = Path(destination)
destination.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
_atomic_write(destination, content.encode("utf-8"))
destination.chmod(0o400)
return destination
+27
View File
@@ -0,0 +1,27 @@
"""PostgreSQL execution operations used by the direct DWH adapter."""
from sqlalchemy import Engine
from tht.execute import (
ExecResult,
PlanSummary,
explain as _explain,
require_positive_int,
run_controlled,
)
DEFAULT_TIMEOUT_MS = 30_000
def run_query(engine: Engine, sql: str, *, limit: int, timeout_ms: int = DEFAULT_TIMEOUT_MS) -> ExecResult:
limit = require_positive_int(limit, name="limit")
return run_controlled(
engine,
sql,
limit=limit,
timeout_ms=timeout_ms,
)
def explain(engine: Engine, sql: str, *, timeout_ms: int = DEFAULT_TIMEOUT_MS) -> PlanSummary:
return _explain(engine, sql, timeout_ms=timeout_ms)
+57
View File
@@ -4,11 +4,68 @@ from dataclasses import dataclass
from sqlalchemy import Engine, text
from tht.config import ExamplesConfig, LshConfig
from tht.execute import require_positive_int
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]:
limit = require_positive_int(limit, name="limit")
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]:
limit = require_positive_int(limit, name="limit")
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:
max_values = require_positive_int(max_values, name="max_values")
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:
max_values = require_positive_int(max_values, name="max_values")
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:
+7
View File
@@ -26,6 +26,13 @@ class PlanSummary:
node_types: list[str]
def require_positive_int(value: object, *, name: str) -> int:
"""Return a validated positive integer, excluding booleans and numeric lookalikes."""
if type(value) is not int or value <= 0:
raise ValueError(f"{name} must be a positive integer")
return value
def _inject_limit(sql: str, limit: int) -> tuple[str, bool]:
"""Aggiunge LIMIT limit+1 se assente (il +1 serve a rilevare il troncamento).
Se la query ha gia' un suo LIMIT, lo si rispetta."""
+15
View File
@@ -0,0 +1,15 @@
"""Shared execution envelope for resumable preprocessing jobs."""
from tht.jobs.locking import JobAlreadyRunningError, WorkspaceJobLock
from tht.jobs.models import JobReport, JobRun, JobSpec
from tht.jobs.runner import JobContext, run_job
__all__ = [
"JobAlreadyRunningError",
"JobContext",
"JobReport",
"JobRun",
"JobSpec",
"WorkspaceJobLock",
"run_job",
]
File diff suppressed because it is too large Load Diff
+115
View File
@@ -0,0 +1,115 @@
"""Crash-safe interprocess locking scoped by workspace and job type."""
from __future__ import annotations
import fcntl
import hashlib
import os
import re
import stat
from pathlib import Path
from types import TracebackType
class JobAlreadyRunningError(RuntimeError):
pass
_JOB_KEY = re.compile(r"^[a-z][a-z0-9_-]{0,63}$")
def _lock_name(workspace_id: str, job_type: str) -> str:
if not _JOB_KEY.fullmatch(workspace_id) or not _JOB_KEY.fullmatch(job_type):
raise ValueError("lock identifiers must be lowercase filesystem-safe keys")
workspace_key = hashlib.sha256(workspace_id.encode("utf-8")).hexdigest()[:16]
return f"{workspace_key}-{job_type}.lock"
class WorkspaceJobLock:
"""Advisory kernel lock; the inode remains stable and is never deleted by PID."""
def __init__(self, workspace_root: Path, workspace_id: str, job_type: str) -> None:
self.path = workspace_root / ".tht-jobs" / ".locks" / _lock_name(
workspace_id, job_type
)
self._fd: int | None = None
def acquire(self) -> "WorkspaceJobLock":
if self._fd is not None:
raise RuntimeError("job lock is already held by this object")
root_fd = os.open(self.path.parents[2], os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
jobs_fd = _open_owned_directory(root_fd, ".tht-jobs")
try:
locks_fd = _open_owned_directory(jobs_fd, ".locks")
try:
fd = os.open(
self.path.name,
os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW | os.O_CLOEXEC,
0o600,
dir_fd=locks_fd,
)
try:
info = os.fstat(fd)
if (
not stat.S_ISREG(info.st_mode)
or info.st_uid != os.getuid()
or info.st_nlink != 1
):
raise OSError("unsafe job lock file")
os.fchmod(fd, 0o600)
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError as error:
raise JobAlreadyRunningError(
"this workspace job is already running"
) from error
except BaseException:
os.close(fd)
raise
finally:
os.close(locks_fd)
finally:
os.close(jobs_fd)
finally:
os.close(root_fd)
self._fd = fd
return self
def release(self) -> None:
if self._fd is None:
return
fd, self._fd = self._fd, None
try:
fcntl.flock(fd, fcntl.LOCK_UN)
finally:
os.close(fd)
def __enter__(self) -> "WorkspaceJobLock":
return self.acquire()
def __exit__(
self,
exc_type: type[BaseException] | None,
exc: BaseException | None,
traceback: TracebackType | None,
) -> None:
self.release()
def _open_owned_directory(parent_fd: int, name: str) -> int:
try:
os.mkdir(name, 0o700, dir_fd=parent_fd)
os.fsync(parent_fd)
except FileExistsError:
pass
fd = os.open(name, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=parent_fd)
try:
info = os.fstat(fd)
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
raise OSError("unsafe job lock directory")
os.fchmod(fd, 0o700)
except BaseException:
os.close(fd)
raise
return fd
+213
View File
@@ -0,0 +1,213 @@
"""Immutable, secret-free records for preprocessing execution."""
from __future__ import annotations
import re
from datetime import UTC, datetime
from pathlib import Path
from typing import Literal, Self
from pydantic import BaseModel, ConfigDict, Field, field_serializer, field_validator, model_validator
_JOB_KEY = re.compile(r"^[a-z][a-z0-9_-]{0,63}$")
_RUN_ID = re.compile(r"^[0-9a-f]{32}$")
_FINGERPRINT = re.compile(r"^sha256:[0-9a-f]{64}$")
JobStatus = Literal["pending", "running", "succeeded", "failed"]
StageStatus = Literal["pending", "running", "succeeded", "failed"]
EffectState = Literal["intent", "completed"]
def utc_now() -> datetime:
return datetime.now(UTC)
def _validate_job_key(value: str) -> str:
if not _JOB_KEY.fullmatch(value):
raise ValueError("job identifiers must be lowercase filesystem-safe keys")
return value
def _validate_run_id(value: str | None) -> str | None:
if value is not None and not _RUN_ID.fullmatch(value):
raise ValueError("run id must contain 32 lowercase hexadecimal characters")
return value
class _FrozenModel(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", validate_default=True)
def model_copy(self, *, update=None, deep: bool = False) -> Self:
data = self.model_dump(round_trip=True)
if update:
data.update(update)
return type(self).model_validate(data)
@field_serializer("*", when_used="json", check_fields=False)
def serialize_utc(self, value):
if isinstance(value, datetime):
return value.astimezone(UTC).isoformat().replace("+00:00", "Z")
return value
class JobSpec(_FrozenModel):
"""Execution input. The local root is deliberately excluded from serialization."""
workspace_id: str
job_type: str
workspace_root: Path = Field(exclude=True)
spec_version: str = Field(min_length=1, max_length=64)
pipeline_version: str = Field(min_length=1, max_length=64)
config_fingerprint: str
input_fingerprint: str
stage_ids: tuple[str, ...]
dry_run: bool = False
resume_run_id: str | None = None
_workspace_key = field_validator("workspace_id")(_validate_job_key)
_job_type_key = field_validator("job_type")(_validate_job_key)
_version_keys = field_validator("spec_version", "pipeline_version")(_validate_job_key)
_resume_id = field_validator("resume_run_id")(_validate_run_id)
_config_fingerprint = field_validator("config_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_input_fingerprint = field_validator("input_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_stage_ids = field_validator("stage_ids")(
lambda values: tuple(_validate_job_key(value) for value in values)
)
def model_copy(self, *, update=None, deep: bool = False) -> Self:
data = {
"workspace_id": self.workspace_id,
"job_type": self.job_type,
"workspace_root": self.workspace_root,
"spec_version": self.spec_version,
"pipeline_version": self.pipeline_version,
"config_fingerprint": self.config_fingerprint,
"input_fingerprint": self.input_fingerprint,
"stage_ids": self.stage_ids,
"dry_run": self.dry_run,
"resume_run_id": self.resume_run_id,
}
if update:
data.update(update)
return type(self).model_validate(data)
def with_resume(self, run_id: str) -> "JobSpec":
return self.model_copy(update={"resume_run_id": run_id})
class StageError(_FrozenModel):
category: Literal["internal"] = "internal"
code: Literal["stage_exception"] = "stage_exception"
message: Literal["stage execution failed"] = "stage execution failed"
class StageRun(_FrozenModel):
name: str
status: StageStatus = "pending"
started_at: datetime | None = None
finished_at: datetime | None = None
error: StageError | None = None
effect_state: EffectState | None = None
artifact_manifest_digest: str | None = None
artifact_files: tuple[str, ...] = ()
_name_key = field_validator("name")(_validate_job_key)
_artifact_digest = field_validator("artifact_manifest_digest")(
lambda value: value if value is None or _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
@model_validator(mode="after")
def state_shape(self) -> "StageRun":
if self.status == "pending" and any(
value is not None for value in (
self.started_at, self.finished_at, self.error, self.effect_state,
self.artifact_manifest_digest,
)
):
raise ValueError("pending stage cannot contain timestamps or error")
if self.status == "running" and (
self.started_at is None or self.finished_at is not None or self.error is not None
):
raise ValueError("running stage requires only started_at")
if self.status == "succeeded" and (
self.started_at is None or self.finished_at is None or self.error is not None
):
raise ValueError("succeeded stage requires timestamps and no error")
if self.status == "failed" and (
self.started_at is None or self.finished_at is None or self.error is None
):
raise ValueError("failed stage requires timestamps and safe error")
if (self.effect_state is None) != (self.artifact_manifest_digest is None):
raise ValueError("effect state and artifact manifest digest must be persisted together")
if self.artifact_files and self.effect_state is None:
raise ValueError("artifact files require a persisted effect state")
return self
class JobRun(_FrozenModel):
"""Durable checkpoint, persisted after every state transition."""
schema_version: Literal[1] = 1
run_id: str
compatibility_fingerprint: str
workspace_fingerprint: str
job_type: str
spec_version: str
pipeline_version: str
config_fingerprint: str
input_fingerprint: str
dry_run: bool
status: JobStatus
started_at: datetime
finished_at: datetime | None = None
resumed_from: str | None = None
stages: tuple[StageRun, ...] = ()
_run_id = field_validator("run_id")(_validate_run_id)
_compatibility = field_validator("compatibility_fingerprint", "workspace_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_input_fingerprints = field_validator("config_fingerprint", "input_fingerprint")(
lambda value: value if _FINGERPRINT.fullmatch(value) else _invalid_fingerprint()
)
_job_type = field_validator("job_type")(_validate_job_key)
_persisted_versions = field_validator("spec_version", "pipeline_version")(_validate_job_key)
_resumed_from = field_validator("resumed_from")(_validate_run_id)
@model_validator(mode="after")
def ledger_shape(self) -> "JobRun":
names = [stage.name for stage in self.stages]
if len(names) != len(set(names)):
raise ValueError("stage identifiers must be unique")
statuses = [stage.status for stage in self.stages]
first_incomplete = next(
(index for index, status in enumerate(statuses) if status != "succeeded"),
len(statuses),
)
if any(status != "pending" for status in statuses[first_incomplete + 1 :]):
raise ValueError("stage ledger must be an ordered execution prefix")
if self.status == "succeeded" and (
self.finished_at is None or any(status != "succeeded" for status in statuses)
):
raise ValueError("succeeded job requires a complete succeeded ledger")
if self.status == "failed" and (
self.finished_at is None
or first_incomplete == len(statuses)
or statuses[first_incomplete] != "failed"
):
raise ValueError("failed job requires the first incomplete stage to be failed")
if self.status == "running" and self.finished_at is not None:
raise ValueError("running job cannot have finished_at")
return self
class JobReport(JobRun):
"""Public machine-readable terminal report (contains no paths or stage outputs)."""
def _invalid_fingerprint():
raise ValueError("fingerprint must be sha256 followed by 64 lowercase hexadecimal characters")
+427
View File
@@ -0,0 +1,427 @@
"""Resumable stage runner with durable atomic checkpoints and reports."""
from __future__ import annotations
import json
import hashlib
import os
import uuid
import stat
import shutil
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from pydantic import ValidationError
from tht.jobs.locking import WorkspaceJobLock
from tht.jobs.models import JobReport, JobRun, JobSpec, StageError, StageRun, utc_now
class CorruptCheckpointError(RuntimeError):
pass
@dataclass(frozen=True)
class StageArtifacts:
required: tuple[str, ...] = ()
@dataclass(frozen=True)
class JobContext:
run_id: str
job_type: str
dry_run: bool
workspace_root: Path
run_dir: Path
_record_artifacts: Callable[[str, tuple[str, ...], str], None] | None = None
def record_artifacts(
self, stage: str, required: tuple[str, ...], effect_state: str,
) -> None:
if self._record_artifacts is None:
raise RuntimeError("artifact recorder is unavailable")
self._record_artifacts(stage, required, effect_state)
Stage = Callable[[JobContext], Any]
def _artifact_digest(path: Path) -> dict[str, Any]:
payload = path.read_bytes()
return {"sha256": hashlib.sha256(payload).hexdigest(), "size": len(payload)}
def _seal_artifacts(context: JobContext, stage: str, result: Any, spec: JobSpec) -> str:
required = result.required if isinstance(result, StageArtifacts) else ()
root = context.run_dir / "artifacts"
root.mkdir(exist_ok=True)
manifest_path = root / "artifact-manifest.json"
manifest = json.loads(manifest_path.read_text()) if manifest_path.exists() else {
"schema_version": 1,
"spec_fingerprint": _compatibility_fingerprint(spec, list(spec.stage_ids)),
"stages": {},
}
files = {}
for relative in required:
candidate = root / relative
if Path(relative).is_absolute() or ".." in Path(relative).parts or candidate.is_symlink():
raise CorruptCheckpointError("artifact path is unsafe")
if not candidate.is_file():
raise CorruptCheckpointError("required stage artifact is missing")
files[relative] = _artifact_digest(candidate)
for prior in manifest["stages"].values():
if relative in prior.get("required", []):
prior["required"].remove(relative)
prior["files"].pop(relative, None)
manifest["stages"][stage] = {"required": list(required), "files": files}
canonical = json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n"
_atomic_write(manifest_path, canonical)
return _value_fingerprint(canonical)
def seal_stage_artifacts(
context: JobContext, stage: str, required: tuple[str, ...], spec: JobSpec,
) -> None:
"""Durably record external-effect intent before a stage performs that effect."""
context.record_artifacts(stage, required, "intent")
def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str]:
root = run_dir / "artifacts"
manifest_path = root / "artifact-manifest.json"
successful = {stage.name for stage in source.stages if stage.status == "succeeded"}
if not successful and not manifest_path.exists():
return set()
try:
manifest = json.loads(manifest_path.read_text())
root_digest = _value_fingerprint(manifest_path.read_text())
if manifest["spec_fingerprint"] != _compatibility_fingerprint(spec, list(spec.stage_ids)):
raise ValueError
sealed = set(manifest["stages"])
allowed = {"artifact-manifest.json"}
for stage, record in manifest["stages"].items():
for relative in record["required"]:
candidate = root / relative
if Path(relative).is_absolute() or ".." in Path(relative).parts or candidate.is_symlink():
raise ValueError
if not candidate.is_file() or _artifact_digest(candidate) != record["files"][relative]:
raise ValueError
allowed.add(relative)
if not successful.issubset(sealed):
raise ValueError
for stage in source.stages:
if stage.effect_state is None:
continue
record = manifest["stages"].get(stage.name)
if (
stage.artifact_manifest_digest != root_digest
or record is None
or tuple(record["required"]) != stage.artifact_files
):
raise ValueError
entries = list(root.iterdir())
if any(path.is_symlink() or not path.is_file() for path in entries):
raise ValueError
actual = {path.name for path in entries}
incomplete = next((stage for stage in source.stages if stage.status != "succeeded"), None)
marker = root / "compensated.json"
if incomplete is not None and incomplete.status in {"failed", "running"} and marker.is_file():
payload = json.loads(marker.read_text())
if not isinstance(payload.get("generation"), str):
raise ValueError
allowed.add("compensated.json")
if actual != allowed:
raise ValueError
return {
stage.name for stage in source.stages
if stage.effect_state == "completed"
}
except (OSError, KeyError, TypeError, ValueError, json.JSONDecodeError) as error:
raise CorruptCheckpointError("resume artifact manifest is invalid") from error
def _atomic_write(path: Path, payload: str) -> None:
temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
try:
with os.fdopen(fd, "w", encoding="utf-8") as stream:
stream.write(payload)
stream.flush()
os.fsync(stream.fileno())
os.replace(temporary, path)
directory_fd = os.open(path.parent, os.O_RDONLY)
try:
os.fsync(directory_fd)
finally:
os.close(directory_fd)
except BaseException:
try:
temporary.unlink()
except FileNotFoundError:
pass
raise
def _persist(path: Path, run: JobRun) -> None:
_atomic_write(path, run.model_dump_json(indent=2) + "\n")
def _load_checkpoint(path: Path) -> JobRun:
try:
return JobRun.model_validate_json(path.read_text(encoding="utf-8"))
except (OSError, ValidationError, ValueError, json.JSONDecodeError) as error:
raise CorruptCheckpointError("checkpoint is invalid and cannot be resumed") from error
def _new_run(spec: JobSpec, run_id: str, stages: Sequence[Stage]) -> JobRun:
names = list(spec.stage_ids)
if len(names) != len(stages):
raise ValueError("stage_ids must identify every stage exactly once")
if len(names) != len(set(names)):
raise ValueError("stage names must be unique")
return JobRun(
run_id=run_id,
compatibility_fingerprint=_compatibility_fingerprint(spec, names),
workspace_fingerprint=_value_fingerprint(spec.workspace_id),
job_type=spec.job_type,
spec_version=spec.spec_version,
pipeline_version=spec.pipeline_version,
config_fingerprint=spec.config_fingerprint,
input_fingerprint=spec.input_fingerprint,
dry_run=spec.dry_run,
status="running",
started_at=utc_now(),
resumed_from=spec.resume_run_id,
stages=tuple(StageRun(name=name) for name in names),
)
def _resume_run(
spec: JobSpec, run_id: str, stages: Sequence[Stage], source: JobRun,
effect_completed: set[str] | None = None,
) -> JobRun:
requested_names = list(spec.stage_ids)
if len(requested_names) != len(stages) or len(requested_names) != len(set(requested_names)):
raise CorruptCheckpointError("resume checkpoint is incompatible with requested stages")
source_fingerprint = _source_compatibility_fingerprint(source)
if source.compatibility_fingerprint != source_fingerprint:
raise CorruptCheckpointError("resume checkpoint compatibility fingerprint is invalid")
expected = _compatibility_fingerprint(spec, requested_names)
if source_fingerprint != expected or [stage.name for stage in source.stages] != requested_names:
raise CorruptCheckpointError(
"resume checkpoint is incompatible; start an intentional new run without resume"
)
source_by_name = {stage.name: stage for stage in source.stages}
resumed_stages = []
for name in requested_names:
previous = source_by_name.get(name)
if previous is not None and previous.status == "succeeded":
resumed_stages.append(previous)
elif previous is not None and previous.status == "running" and name in (effect_completed or set()):
resumed_stages.append(StageRun(
name=name, status="succeeded", started_at=previous.started_at or utc_now(),
finished_at=utc_now(),
effect_state="completed",
artifact_manifest_digest=previous.artifact_manifest_digest,
artifact_files=previous.artifact_files,
))
else:
resumed_stages.append(StageRun(name=name))
return JobRun(
run_id=run_id,
compatibility_fingerprint=source.compatibility_fingerprint,
workspace_fingerprint=source.workspace_fingerprint,
job_type=spec.job_type,
spec_version=spec.spec_version,
pipeline_version=spec.pipeline_version,
config_fingerprint=spec.config_fingerprint,
input_fingerprint=spec.input_fingerprint,
dry_run=spec.dry_run,
status="running",
started_at=utc_now(),
resumed_from=source.run_id,
stages=tuple(resumed_stages),
)
def run_job(
spec: JobSpec, stages: Sequence[Stage], *,
after_stage_return: Callable[[JobContext, str], Any] | None = None,
reconcile_effects: Callable[[JobRun, Path], set[str]] | None = None,
) -> JobReport:
"""Run stages once, returning a terminal report instead of leaking stage exceptions."""
with WorkspaceJobLock(spec.workspace_root, spec.workspace_id, spec.job_type):
jobs_root = spec.workspace_root / ".tht-jobs" / spec.job_type / "runs"
if spec.resume_run_id is None:
source = None
else:
source_path = jobs_root / spec.resume_run_id / "checkpoint.json"
if not source_path.exists():
matches = list(
(spec.workspace_root / ".tht-jobs").glob(
f"*/runs/{spec.resume_run_id}/checkpoint.json"
)
)
if len(matches) == 1:
source_path = matches[0]
source = _load_checkpoint(source_path)
_validate_resume_source(spec, stages, source)
effect_completed = _validate_artifacts(source_path.parent, spec, source)
if reconcile_effects is not None:
effect_completed |= reconcile_effects(source, source_path.parent)
run_id = uuid.uuid4().hex
run_dir = jobs_root / run_id
_prepare_run_directory(spec.workspace_root, spec.job_type, run_id)
checkpoint_path = run_dir / "checkpoint.json"
if source is None:
run = _new_run(spec, run_id, stages)
else:
run = _resume_run(spec, run_id, stages, source, effect_completed)
source_artifacts = jobs_root / source.run_id / "artifacts"
if source_artifacts.exists():
shutil.copytree(source_artifacts, run_dir / "artifacts")
_persist(checkpoint_path, run)
current_index = -1
def record_artifacts(stage_name: str, required: tuple[str, ...], effect_state: str) -> None:
nonlocal run
if current_index < 0 or run.stages[current_index].name != stage_name:
raise CorruptCheckpointError("artifact producer does not match running stage")
digest = _seal_artifacts(context, stage_name, StageArtifacts(required), spec)
manifest = json.loads((run_dir / "artifacts" / "artifact-manifest.json").read_text())
updated = []
for position, value in enumerate(run.stages):
record = manifest["stages"].get(value.name)
if record is not None and (value.status == "succeeded" or position == current_index):
state = effect_state if position == current_index else value.effect_state
updated.append(value.model_copy(update={
"effect_state": state,
"artifact_manifest_digest": digest,
"artifact_files": tuple(record["required"]),
}))
else:
updated.append(value)
run = run.model_copy(update={"stages": tuple(updated)})
_persist(checkpoint_path, run)
context = JobContext(
run_id, spec.job_type, spec.dry_run, spec.workspace_root, run_dir,
record_artifacts,
)
for index, stage_callable in enumerate(stages):
current_index = index
if run.stages[index].status == "succeeded":
continue
stage = run.stages[index].model_copy(
update={"status": "running", "started_at": utc_now()}
)
run = run.model_copy(
update={"stages": run.stages[:index] + (stage,) + run.stages[index + 1 :]}
)
_persist(checkpoint_path, run)
try:
stage_result = stage_callable(context)
except Exception:
failed = stage.model_copy(
update={
"status": "failed",
"finished_at": utc_now(),
"error": StageError(),
}
)
run = run.model_copy(
update={
"status": "failed",
"finished_at": utc_now(),
"stages": run.stages[:index] + (failed,) + run.stages[index + 1 :],
}
)
_persist(checkpoint_path, run)
break
required = stage_result.required if isinstance(stage_result, StageArtifacts) else ()
context.record_artifacts(stage.name, required, "completed")
stage = run.stages[index]
if after_stage_return is not None:
after_stage_return(context, stage.name)
succeeded = stage.model_copy(update={"status": "succeeded", "finished_at": utc_now()})
run = run.model_copy(
update={"stages": run.stages[:index] + (succeeded,) + run.stages[index + 1 :]}
)
_persist(checkpoint_path, run)
else:
run = run.model_copy(update={"status": "succeeded", "finished_at": utc_now()})
_persist(checkpoint_path, run)
report = JobReport.model_validate(run.model_dump())
_atomic_write(run_dir / "report.json", report.model_dump_json(indent=2) + "\n")
return report
def _value_fingerprint(value: str) -> str:
return "sha256:" + hashlib.sha256(value.encode("utf-8")).hexdigest()
def _compatibility_fingerprint(spec: JobSpec, stage_ids: list[str]) -> str:
payload = {
"schema_version": 1,
"workspace": _value_fingerprint(spec.workspace_id),
"job_type": spec.job_type,
"dry_run": spec.dry_run,
"spec_version": spec.spec_version,
"pipeline_version": spec.pipeline_version,
"config_fingerprint": spec.config_fingerprint,
"input_fingerprint": spec.input_fingerprint,
"stage_ids": stage_ids,
}
canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"))
return _value_fingerprint(canonical)
def _source_compatibility_fingerprint(source: JobRun) -> str:
payload = {
"schema_version": source.schema_version,
"workspace": source.workspace_fingerprint,
"job_type": source.job_type,
"dry_run": source.dry_run,
"spec_version": source.spec_version,
"pipeline_version": source.pipeline_version,
"config_fingerprint": source.config_fingerprint,
"input_fingerprint": source.input_fingerprint,
"stage_ids": [stage.name for stage in source.stages],
}
canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"))
return _value_fingerprint(canonical)
def _validate_resume_source(spec: JobSpec, stages: Sequence[Stage], source: JobRun) -> None:
_resume_run(spec, "0" * 32, stages, source)
def _prepare_run_directory(workspace_root: Path, job_type: str, run_id: str) -> None:
parent_fd = os.open(workspace_root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try:
for component in (".tht-jobs", job_type, "runs", run_id):
try:
os.mkdir(component, 0o700, dir_fd=parent_fd)
os.fsync(parent_fd)
except FileExistsError:
pass
child_fd = os.open(
component,
os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW,
dir_fd=parent_fd,
)
info = os.fstat(child_fd)
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid():
os.close(child_fd)
raise OSError("unsafe job run directory")
os.fchmod(child_fd, 0o700)
os.close(parent_fd)
parent_fd = child_fd
os.fsync(parent_fd)
finally:
os.close(parent_fd)
+8 -13
View File
@@ -255,7 +255,7 @@ def memory_vector_record_for_decision(
def save_one_memory(
records: list[MemoryRecord], decision_seq: int, *, writer, embedder
records: list[MemoryRecord], decision_seq: int, *, store, embedder
) -> int:
"""Targeted one-row upsert of a promoted decision to pgvector via the writer key
(spec D11). This is NOT a full vectorstore resync: it embeds and pushes a single
@@ -267,27 +267,22 @@ def save_one_memory(
the writer's existing_vector_hashes; the embedding (Ollama round-trip) and the
upsert are skipped when the content is unchanged. Idempotent by construction.
`writer` is a VectorRestClient (writer key); `embedder` an embeddings client.
`store` is the configured writable VectorStore; `embedder` an embeddings client.
The destructive cleanup (sync's delete-stale step) is intentionally absent: it
remains a server-side-only operation via the direct vectordb connection.
"""
from tht.vectorstore.rest_writer import pack_metadata
from tht.ports.vector import VectorWriteRecord
from tht.vectorstore.store import content_hash
record = memory_vector_record_for_decision(records, decision_seq)
if record is None:
return 0
new_hash = content_hash(record.content)
existing = writer.existing_hashes("memory", ["memory"])
existing = store.existing_hashes("memory", ["memory"])
if existing.get(record.id) == new_hash:
return 0 # unchanged: skip embedding + upsert
embedding = embedder.embed_documents([record.content])[0]
row = {
"record_key": record.id,
"kind": record.kind,
"content_hash": new_hash,
"metadata": pack_metadata(record),
"embedding": embedding,
}
return writer.upsert_records("memory", [row])
return store.upsert(
"memory",
[VectorWriteRecord(record=record, embedding=embedding, content_hash=new_hash)],
)
@@ -0,0 +1,3 @@
CREATE SCHEMA IF NOT EXISTS vectors;
REVOKE ALL ON SCHEMA vectors FROM PUBLIC;
CREATE EXTENSION IF NOT EXISTS vector WITH SCHEMA vectors;
@@ -0,0 +1,32 @@
CREATE TABLE IF NOT EXISTS vectors.schema_records (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(768) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT pg_catalog.now()
);
CREATE TABLE IF NOT EXISTS vectors.evidence (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(768) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT pg_catalog.now()
);
CREATE TABLE IF NOT EXISTS vectors.memory (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vectors.vector(768) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT pg_catalog.now()
);
REVOKE ALL ON ALL TABLES IN SCHEMA vectors FROM PUBLIC;
REVOKE ALL ON ALL SEQUENCES IN SCHEMA vectors FROM PUBLIC;
@@ -0,0 +1,23 @@
DO $roles$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_catalog.pg_roles WHERE rolname = 'vector_reader') THEN
CREATE ROLE vector_reader NOLOGIN;
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_catalog.pg_roles WHERE rolname = 'vector_writer') THEN
CREATE ROLE vector_writer NOLOGIN;
END IF;
END
$roles$;
REVOKE ALL ON SCHEMA vectors FROM vector_reader, vector_writer;
REVOKE ALL ON ALL TABLES IN SCHEMA vectors FROM vector_reader, vector_writer;
REVOKE ALL ON ALL SEQUENCES IN SCHEMA vectors FROM vector_reader, vector_writer;
GRANT USAGE ON SCHEMA vectors TO vector_reader, vector_writer;
GRANT SELECT ON ALL TABLES IN SCHEMA vectors TO vector_reader;
GRANT INSERT, UPDATE
ON vectors.schema_records, vectors.evidence, vectors.memory TO vector_writer;
GRANT SELECT (record_key, kind, content_hash)
ON vectors.schema_records, vectors.evidence, vectors.memory TO vector_writer;
GRANT USAGE ON ALL SEQUENCES IN SCHEMA vectors TO vector_writer;
@@ -0,0 +1,3 @@
-- The writer owns derived-generation reconciliation but not runtime similarity reads.
GRANT SELECT (metadata) ON vectors.evidence TO vector_writer;
GRANT DELETE ON vectors.evidence TO vector_writer;
+54
View File
@@ -0,0 +1,54 @@
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING
from tht.config import ConfigError
if TYPE_CHECKING:
from tht.config import Config
@dataclass(frozen=True)
class ResolvedPaths:
workspace: Path
sessions: Path
artifacts: Path
indexes: Path
corpus: Path
def _resolve_root(path: Path, workspace: Path, name: str) -> Path:
if path.is_absolute():
# Existing deployments commonly point at an external workspace checkout. Keep
# those paths working while `tht doctor` identifies them for migration.
return path
resolved = (workspace / path).resolve()
if not resolved.is_relative_to(workspace):
raise ConfigError(f"paths.{name} resolves outside workspace root")
return resolved
def resolve_workspace_paths(
config_path: Path, cfg: Config, data_root: Path
) -> ResolvedPaths:
"""Resolve portable paths under ``<data_root>/workspaces/<config stem>``.
Relative roots are sandboxed to the logical workspace. Absolute roots are a
compatibility bridge for existing installations and are never rewritten.
"""
canonical_data_root = data_root.resolve()
workspaces_root = canonical_data_root / "workspaces"
workspace = (workspaces_root / config_path.stem).resolve()
if not workspace.is_relative_to(workspaces_root):
raise ConfigError("workspace resolves outside workspaces root")
roots = cfg.roots
return ResolvedPaths(
workspace=workspace,
sessions=_resolve_root(roots.sessions, workspace, "sessions"),
artifacts=_resolve_root(roots.artifacts, workspace, "artifacts"),
indexes=_resolve_root(roots.indexes, workspace, "indexes"),
corpus=_resolve_root(Path("corpus"), workspace, "corpus"),
)
+37
View File
@@ -0,0 +1,37 @@
"""Stable interfaces implemented by Thoth infrastructure adapters."""
from tht.ports.dwh import (
DwhAdapter,
DwhCapabilities,
DwhHealth,
DistinctValues,
UnsupportedCapability,
)
from tht.ports.vector import (
VectorCapabilities,
VectorHealth,
VectorHit,
VectorRecord,
VectorReadUnavailable,
VectorStore,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
)
__all__ = [
"DwhAdapter",
"DwhCapabilities",
"DwhHealth",
"DistinctValues",
"UnsupportedCapability",
"VectorCapabilities",
"VectorHealth",
"VectorHit",
"VectorRecord",
"VectorReadUnavailable",
"VectorStore",
"VectorStoreError",
"VectorWriteRecord",
"VectorWriteUnavailable",
]
+56
View File
@@ -0,0 +1,56 @@
"""Data-warehouse adapter contract."""
from dataclasses import dataclass
from typing import Protocol, runtime_checkable
from tht.execute import ExecResult, PlanSummary
from tht.mschema.models import PhysicalSchema
@dataclass(frozen=True)
class DwhCapabilities:
introspection: bool = True
explain: bool = True
sampling: bool = True
distinct_values: bool = True
@dataclass(frozen=True)
class DwhHealth:
ok: bool
detail: str | None = None
database: str | None = None
schema: str | None = None
endpoint: str | None = None
read_only: bool | None = None
writable_tables: tuple[str, ...] = ()
can_create: bool = False
error_kind: 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."""
@runtime_checkable
class DwhAdapter(Protocol):
@property
def capabilities(self) -> DwhCapabilities: ...
def health(self) -> DwhHealth: ...
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, *, limit: int) -> DistinctValues: ...
+210
View File
@@ -0,0 +1,210 @@
"""Credential-free port for discovering and acquiring Evidence objects."""
import re
from collections.abc import Iterable, Mapping, Sequence
from datetime import UTC, datetime
from enum import Enum
from typing import Protocol, Self, runtime_checkable
from urllib.parse import parse_qsl, urlsplit, urlunsplit
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, field_validator
class FrozenDict(dict):
"""A JSON-serializable dict whose mutation operations are disabled."""
def _immutable(self, *args, **kwargs):
raise TypeError("frozen JSON metadata cannot be mutated")
__delitem__ = _immutable
__ior__ = _immutable
__setitem__ = _immutable
clear = _immutable
pop = _immutable
popitem = _immutable
setdefault = _immutable
update = _immutable
_CAMEL_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])")
_SEPARATORS = re.compile(r"[^a-z0-9]+")
_NAMESPACED_VALUE = re.compile(r"^[a-z][a-z0-9_-]*:[A-Za-z0-9._:-]+$")
_CREDENTIAL_KEYS = {
"apikey",
"authorization",
"authtoken",
"bearertoken",
"clientsecret",
"credential",
"credentials",
"password",
"passwd",
"privatekey",
"refreshtoken",
"sessioncookie",
"xapikey",
"accesstoken",
}
_JSON_METADATA = TypeAdapter(dict[str, JsonValue])
def _normalize_key(key: str) -> str:
return _SEPARATORS.sub("", _CAMEL_BOUNDARY.sub("_", key).lower())
def _is_credential_key(key: str) -> bool:
return _normalize_key(key) in _CREDENTIAL_KEYS
def _reject_credentials(value, path: str = "metadata") -> None:
if isinstance(value, Mapping):
for key, child in value.items():
if _is_credential_key(str(key)):
raise ValueError(f"credential-like metadata key is not allowed: {path}.{key}")
_reject_credentials(child, f"{path}.{key}")
elif isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
for index, child in enumerate(value):
_reject_credentials(child, f"{path}[{index}]")
def freeze_json(value):
"""Recursively freeze a Pydantic-validated JSON value without changing its JSON shape."""
if isinstance(value, Mapping):
return FrozenDict({str(key): freeze_json(child) for key, child in value.items()})
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
return tuple(freeze_json(child) for child in value)
return value
def validate_safe_metadata(value: dict[str, JsonValue]) -> FrozenDict:
_reject_credentials(value)
return freeze_json(value)
def validate_canonical_uri(value: str) -> str:
try:
parsed = urlsplit(value)
_ = parsed.port
except ValueError as error:
raise ValueError("invalid canonical URI") from error
if not parsed.scheme:
raise ValueError("canonical URI must include a scheme")
if parsed.username is not None or parsed.password is not None:
raise ValueError("canonical URI must not contain credentials in userinfo")
for key, _ in parse_qsl(parsed.query, keep_blank_values=True):
if _is_credential_key(key):
raise ValueError("canonical URI must not contain credentials in query parameters")
return value
def canonical_provenance_uri(value: str) -> str:
"""Return only stable URI identity; transport query/fragment data is never provenance."""
try:
parsed = urlsplit(value)
_ = parsed.port
except ValueError as error:
raise ValueError("invalid canonical URI") from error
if not parsed.scheme:
raise ValueError("canonical URI must include a scheme")
if parsed.username is not None or parsed.password is not None:
raise ValueError("canonical URI must not contain credentials in userinfo")
return urlunsplit((parsed.scheme, parsed.netloc, parsed.path, "", ""))
def normalize_aware_datetime(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None or value.utcoffset() is None:
raise ValueError("datetime must be timezone-aware")
return value.astimezone(UTC)
def validate_namespaced_value(value: str) -> str:
if not _NAMESPACED_VALUE.fullmatch(value):
raise ValueError("value must be namespaced as '<kind>:<stable-value>'")
return value
class _EvidenceValue(BaseModel):
model_config = ConfigDict(
frozen=True,
extra="forbid",
revalidate_instances="always",
validate_default=True,
ser_json_bytes="base64",
val_json_bytes="base64",
)
def model_copy(self, *, update: Mapping[str, object] | None = None, deep: bool = False) -> Self:
"""Copy through validation; Pydantic's unchecked update-copy is unsafe for contracts."""
data = self.model_dump(round_trip=True)
if update:
data.update(update)
return type(self).model_validate(data)
class SourceObject(_EvidenceValue):
source_id: str = Field(min_length=1)
uri: str = Field(min_length=1)
fingerprint: str = Field(min_length=1)
modified_at: datetime | None = None
metadata: dict[str, JsonValue] = Field(default_factory=dict)
_source_id = field_validator("source_id")(validate_namespaced_value)
_fingerprint = field_validator("fingerprint")(validate_namespaced_value)
_safe_uri = field_validator("uri")(validate_canonical_uri)
_aware_modified_at = field_validator("modified_at")(normalize_aware_datetime)
_frozen_metadata = field_validator("metadata")(validate_safe_metadata)
class AcquiredDocument(_EvidenceValue):
"""Transport result; bytes use explicit base64 encoding in JSON mode."""
source: SourceObject
content: bytes
media_type: str | None = None
acquired_at: datetime | None = None
metadata: dict[str, JsonValue] = Field(default_factory=dict)
_aware_acquired_at = field_validator("acquired_at")(normalize_aware_datetime)
_frozen_metadata = field_validator("metadata")(validate_safe_metadata)
class EvidenceSourceErrorCategory(str, Enum):
TRANSIENT = "transient"
PERMANENT = "permanent"
class EvidenceSourceError(Exception):
"""Classified source failure with credential-free structured diagnostics."""
def __init__(
self,
_message: str,
*,
category: EvidenceSourceErrorCategory,
details: dict[str, JsonValue] | None = None,
) -> None:
super().__init__("evidence source operation failed")
object.__setattr__(self, "category", EvidenceSourceErrorCategory(category))
object.__setattr__(
self,
"details",
validate_safe_metadata(_JSON_METADATA.validate_python(details or {})),
)
def __setattr__(self, name: str, value) -> None:
if name in {"args", "category", "details"} and hasattr(self, name):
raise AttributeError(f"{name} is immutable")
super().__setattr__(name, value)
@property
def retryable(self) -> bool:
return self.category is EvidenceSourceErrorCategory.TRANSIENT
@runtime_checkable
class EvidenceSource(Protocol):
def discover(self) -> Iterable[SourceObject]: ...
def acquire(self, item: SourceObject) -> AcquiredDocument: ...
+98
View File
@@ -0,0 +1,98 @@
"""Transport-neutral vector-store contract and canonical vector models."""
from dataclasses import dataclass
from typing import Protocol, runtime_checkable
from tht.vectorstore.records import VectorRecord
from tht.vectorstore.store import VectorHit
@dataclass(frozen=True)
class VectorCapabilities:
search: bool = True
existing_hashes: bool = False
upsert: bool = False
metadata_filter: bool = False
delete_generation: bool = False
list_evidence_generations: bool = False
@dataclass(frozen=True)
class VectorHealth:
ok: bool
detail: str | None = None
read_configured: bool = False
read_reachable: bool | None = None
read_detail: str | None = None
write_configured: bool = False
write_reachable: bool | None = None
write_detail: str | None = None
expected_dimension: int | None = None
observed_dimensions: tuple[int, ...] = ()
dimension_compatible: bool | None = None
@dataclass(frozen=True)
class VectorWriteRecord:
"""A canonical record plus transport-neutral, precomputed vector data."""
record: VectorRecord
embedding: list[float]
content_hash: str
class VectorStoreError(Exception):
"""Base error exposed by vector adapters."""
class VectorWriteUnavailable(VectorStoreError):
"""Raised when a deployment has no vector writer credential."""
class VectorReadUnavailable(VectorStoreError):
"""Raised when a deployment has no vector reader credential."""
def require_positive_limit(limit: int) -> None:
"""Reject coercible values: vector limits are exact positive integers."""
if type(limit) is not int or limit <= 0:
raise ValueError("Vector search limit must be a positive integer")
@runtime_checkable
class VectorStore(Protocol):
@property
def capabilities(self) -> VectorCapabilities: ...
def health(self) -> VectorHealth: ...
def search(
self,
collections: list[str],
embedding: list[float],
*,
limit: int,
kinds: list[str] | None = None,
metadata_filter: dict[str, object] | None = None,
) -> list[VectorHit]: ...
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]: ...
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int: ...
def delete_generation(self, collection: str, generation: str, workspace_id: str) -> int: ...
def list_evidence_generations(self, collection: str, workspace_id: str) -> list[str]: ...
__all__ = [
"VectorCapabilities",
"VectorHealth",
"VectorHit",
"VectorRecord",
"VectorReadUnavailable",
"VectorStore",
"VectorStoreError",
"VectorWriteRecord",
"VectorWriteUnavailable",
]
+9 -1
View File
@@ -9,7 +9,14 @@ il client mantiene solo l'iniezione del LIMIT (per il rilevamento del troncament
import time
from tht.execute import ExecResult, ExecutionError, PlanSummary, _inject_limit, assert_read_only
from tht.execute import (
ExecResult,
ExecutionError,
PlanSummary,
_inject_limit,
assert_read_only,
require_positive_int,
)
from tht.rest.client import RestError
from tht.rest.explain import parse_text_plan
@@ -17,6 +24,7 @@ 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.
limit = require_positive_int(limit, name="limit")
assert_read_only(sql)
final_sql, injected = _inject_limit(sql, limit)
start = time.monotonic()
+116
View File
@@ -0,0 +1,116 @@
"""Runtime Evidence lookup bound to the atomically active corpus generation."""
import re
from tht.corpus.store import CorpusStore
class CorpusWorkspaceMismatchError(RuntimeError):
"""The configured workspace does not own the persisted corpus."""
class ActiveEvidenceSearcher:
"""Searcher facade that enforces ACTIVE generation predicates before LIMIT."""
def __init__(self, corpus: CorpusStore, delegate, expected_workspace_id: str | None = None):
self.corpus = corpus
self.delegate = delegate
self.expected_workspace_id = expected_workspace_id
def search(self, embedding, top_n=10, kinds=None, metadata_filter=None):
requested = set(kinds) if kinds is not None else {
"schema_table", "schema_column", "evidence", "memory", "solved_question",
}
include_evidence = "evidence" in requested
other_kinds = sorted(requested - {"evidence"})
with self.corpus.writer_lock():
manifest = self.corpus.active_manifest()
persisted_workspace = manifest.metadata.get("workspace_id") if manifest else None
if manifest is not None and (
not isinstance(persisted_workspace, str)
or re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", persisted_workspace) is None
):
raise CorpusWorkspaceMismatchError(
"corpus workspace ownership is missing or invalid; use a new corpus root or rebuild"
)
if manifest is not None and self.expected_workspace_id is not None and (
persisted_workspace != self.expected_workspace_id
):
raise CorpusWorkspaceMismatchError(
"corpus belongs to a different workspace; use a new corpus root or rebuild"
)
if not include_evidence:
kwargs = {"top_n": top_n, "kinds": kinds}
if metadata_filter is not None:
kwargs["metadata_filter"] = metadata_filter
return self.delegate.search(embedding, **kwargs)
hits = []
if other_kinds:
kwargs = {"top_n": top_n, "kinds": other_kinds}
if metadata_filter is not None:
kwargs["metadata_filter"] = metadata_filter
hits.extend(self.delegate.search(embedding, **kwargs))
if include_evidence:
workspace_id = manifest.metadata.get("workspace_id") if manifest else None
if manifest is not None and isinstance(workspace_id, str):
by_generation: dict[str, list[str]] = {}
mapping = dict(manifest.metadata.get("document_generations", {}))
for document in manifest.documents:
generation = mapping.get(document.document_id, manifest.vector_generation)
if generation:
by_generation.setdefault(generation, []).append(document.document_id)
for generation, document_ids in sorted(by_generation.items()):
hits.extend(self.delegate.search(
embedding, top_n=top_n, kinds=["evidence"],
metadata_filter={
"vector_generation": generation,
"document_ids": sorted(document_ids),
"workspace_id": workspace_id,
},
))
return sorted(hits, key=lambda hit: (-hit.similarity, hit.id))[:top_n]
def active_searcher(cfg, delegate, *, workspace_id: str | None = None):
corpus_root = cfg.paths.artifacts.parent / "corpus"
return ActiveEvidenceSearcher(CorpusStore(corpus_root), delegate, workspace_id)
def validate_corpus_workspace(cfg, workspace_id: str) -> None:
"""Fail before downstream retrieval setup when configured corpus ownership differs."""
corpus = CorpusStore(cfg.paths.artifacts.parent / "corpus")
with corpus.writer_lock():
manifest = corpus.active_manifest()
if manifest is None:
return
persisted = manifest.metadata.get("workspace_id")
if not isinstance(persisted, str) or re.fullmatch(
r"[a-z][a-z0-9_-]{0,63}", persisted
) is None:
raise CorpusWorkspaceMismatchError(
"corpus workspace ownership is missing or invalid; use a new corpus root or rebuild"
)
if persisted != workspace_id:
raise CorpusWorkspaceMismatchError(
"corpus belongs to a different workspace; use a new corpus root or rebuild"
)
def resolve_evidence_file(
store: CorpusStore, evidence_id: str, *, materialized_root=None,
) -> str:
with store.writer_lock():
manifest = store.active_manifest()
if manifest is None:
return ""
for document in manifest.documents:
frontmatter = document.metadata.get("frontmatter", {})
identifiers = {document.document_id, document.source_id, str(frontmatter.get("id", ""))}
if evidence_id in identifiers:
root = materialized_root or (store.root / "runtime")
filename = document.document_id.removeprefix("doc:") + ".md"
path = store.materialize_document(
document.document_id, root / filename, generation=manifest.manifest_id,
)
return str(path) if path else ""
return ""
+10
View File
@@ -5,6 +5,16 @@ from tht.session.models import SchemaLinking
def _find_evidence_file(evidence_root: Path, evidence_id: str) -> str:
# New deployments resolve only immutable materialized files from ACTIVE. Keep
# the legacy curated-tree fallback for sessions created before a corpus exists.
corpus_root = evidence_root.parent.parent / "corpus"
if corpus_root.exists():
from tht.corpus.store import CorpusStore
from tht.search.evidence import resolve_evidence_file
return resolve_evidence_file(
CorpusStore(corpus_root), evidence_id,
materialized_root=evidence_root.parent / ".materialized-evidence",
)
for match in evidence_root.rglob(f"{evidence_id}.md"):
return str(match)
return ""
+7 -10
View File
@@ -43,25 +43,22 @@ def _solved_hash(record: VectorRecord) -> str:
return content_hash(record.content + "\n" + str(record.metadata.get("sql", "")))
def save_solved_question(record: VectorRecord, *, writer, embedder) -> int:
def save_solved_question(record: VectorRecord, *, store, embedder) -> int:
"""Upsert one-row della coppia domanda->SQL via writer key (stesso pattern di
save_one_memory, spec D11): hash dedup client-side, embedding solo se domanda
o SQL sono cambiati. `writer` e' un VectorRestClient (writer key). Ritorna il
numero di righe upsertate (0 = invariata)."""
from tht.vectorstore.rest_writer import pack_metadata
from tht.ports.vector import VectorWriteRecord
new_hash = _solved_hash(record)
existing = writer.existing_hashes("memory", [SOLVED_KIND])
existing = store.existing_hashes("memory", [SOLVED_KIND])
if existing.get(record.id) == new_hash:
return 0
embedding = embedder.embed_documents([record.content])[0]
return writer.upsert_records("memory", [{
"record_key": record.id,
"kind": record.kind,
"content_hash": new_hash,
"metadata": pack_metadata(record),
"embedding": embedding,
}])
return store.upsert(
"memory",
[VectorWriteRecord(record=record, embedding=embedding, content_hash=new_hash)],
)
class SolvedIndexError(Exception):
+11 -22
View File
@@ -9,8 +9,10 @@ Entrambe mappano i `kind` sulle tabelle per-dominio dello schema `vectors`.
from sqlalchemy import Engine
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
from tht.vectorstore.rest_client import VectorRestClient
from tht.vectorstore.store import VectorHit, VectorStore, hit_from_metadata
from tht.vectorstore.store import VectorHit
# kind Thoth → tabella dello schema `vectors`.
KIND_TO_TABLE = {
@@ -30,30 +32,19 @@ def tables_for_kinds(kinds: list[str] | None) -> list[str]:
return sorted({KIND_TO_TABLE[k] for k in kinds if k in KIND_TO_TABLE})
def _merge(hits: list[VectorHit], top_n: int) -> list[VectorHit]:
return sorted(hits, key=lambda h: h.similarity, reverse=True)[:top_n]
class RestSearcher:
"""Similarity search via REST: una chiamata `search_similar` per tabella, poi fusione."""
def __init__(self, client: VectorRestClient):
self.client = client
self._store = ThothHttpVectorStore(reader=client, writer=None)
def search(
self, query_vec: list[float], top_n: int = 10, kinds: list[str] | None = None
) -> list[VectorHit]:
hits: list[VectorHit] = []
for table in tables_for_kinds(kinds):
for row in self.client.search_similar(table, query_vec, top_n, kinds=kinds):
hits.append(hit_from_metadata(row.get("similarity", 0.0), row.get("metadata")))
# Il filtro per kind avviene server-side (RPC con `kinds`); il post-filter resta
# come difesa per il fallback legacy (server pre-migrazione: 404 -> query senza
# filtro) e per parita' col path diretto (#25).
if kinds:
allowed = set(kinds)
hits = [h for h in hits if h.kind in allowed]
return _merge(hits, top_n)
return self._store.search(
tables_for_kinds(kinds), query_vec, limit=top_n, kinds=kinds
)
class DirectSearcher:
@@ -63,13 +54,11 @@ class DirectSearcher:
self.engine = engine
self.schema = schema
self.dim = dim
self._store = LegacyDirectVectorStore(engine, schema=schema, dim=dim)
def search(
self, query_vec: list[float], top_n: int = 10, kinds: list[str] | None = None
) -> list[VectorHit]:
hits: list[VectorHit] = []
for table in tables_for_kinds(kinds):
store = VectorStore(self.engine, schema=self.schema, table=table, dim=self.dim)
# passa kinds: dentro schema_records filtra schema_table vs schema_column (#25).
hits.extend(store.search(query_vec, top_n=top_n, kinds=kinds))
return _merge(hits, top_n)
return self._store.search(
tables_for_kinds(kinds), query_vec, limit=top_n, kinds=kinds
)
+52
View File
@@ -6,6 +6,7 @@ Errori in italiano e azionabili, stile `rest/client.py`.
"""
import requests
import re
from tht.config import RestConfig
@@ -60,6 +61,7 @@ class VectorRestClient:
def search_similar(
self, table_name: str, query_embedding: list[float], limit_count: int,
kinds: list[str] | None = None,
metadata_filter: dict | None = None,
) -> list[dict]:
"""Ricerca per similarità coseno su `vectors.<table_name>`: ritorna le righe
`{id, similarity, metadata}` ordinate per similarity decrescente. Con `kinds`
@@ -72,6 +74,13 @@ class VectorRestClient:
"limit_count": limit_count,
"table_name": table_name,
}
if metadata_filter is not None:
# ACTIVE corpus reads must never degrade to an unfiltered legacy RPC:
# filtering after LIMIT is incomplete and could expose stale generations.
return self._call(
"search_similar",
{**args, "kinds": kinds, "metadata_filter": metadata_filter},
) or []
if kinds is not None:
try:
return self._call("search_similar", {**args, "kinds": kinds}) or []
@@ -118,3 +127,46 @@ class VectorRestClient:
return int(payload[0]["upserted"])
return len(payload)
return len(rows)
def delete_generation(self, table_name: str, generation: str, workspace_id: str) -> int:
if table_name != "evidence" or re.fullmatch(r"gen:[0-9a-f]{32}", generation) is None:
raise ValueError("generation must be canonical")
if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None:
raise ValueError("workspace namespace must be canonical")
try:
payload = self._call(
"delete_vector_generation",
{"table_name": table_name, "kind": "evidence", "generation": generation,
"workspace_id": workspace_id},
)
except VectorRestError as error:
if "HTTP 404" in str(error):
raise VectorRestError(
"delete_vector_generation RPC is unavailable; deploy the cleanup migration"
) from None
raise
if isinstance(payload, dict):
return int(payload.get("deleted", 0))
return 0
def list_evidence_generations(self, table_name: str, workspace_id: str) -> list[str]:
if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None:
raise ValueError("workspace namespace must be canonical")
try:
rows = self._call(
"list_evidence_generations",
{"table_name": table_name, "kind": "evidence", "workspace_id": workspace_id},
) or []
except VectorRestError as error:
if "HTTP 404" in str(error):
raise VectorRestError(
"list_evidence_generations RPC is unavailable; deploy the cleanup migration"
) from None
raise
if not isinstance(rows, list) or any(
not isinstance(row, dict)
or re.fullmatch(r"gen:[0-9a-f]{32}", str(row.get("generation", ""))) is None
for row in rows
):
raise VectorRestError("list_evidence_generations returned malformed data")
return sorted({row["generation"] for row in rows})
+28 -42
View File
@@ -1,26 +1,19 @@
# Workspace ThothII (esempio). I segreti vivono SOLO in .env (${THT_*}).
# La struttura rispecchia esattamente tht/config.py:
# database + rest per il DWH; vector_rest/vector_write_rest per il pgvector (doppia key);
# vector_db per il loading diretto (server-only); embeddings + evidence + execution.
# Ogni risorsa dichiara il proprio adapter tramite `type`.
language: it # descrizioni tabelle/colonne ed evidence sono in italiano (PSD)
database:
host: ${THT_DB_HOST}
port: ${THT_DB_PORT} # es. 5437 (Postgres diretto Supabase; 5432 = pooler)
database: ${THT_DB_NAME} # es. postgres (lo schema a stella vive in `datawarehouse`)
schema: datawarehouse
user: ${THT_DB_USER}
password: ${THT_DB_PASSWORD}
transport: rest # direct (Postgres) | rest (Supabase/PostgREST)
dwh:
type: thoth_rest # postgres_direct | thoth_rest
database:
database: ${THT_DB_NAME}
schema: datawarehouse
endpoint:
base_url: ${THT_DWH_REST_URL}
api_key: ${THT_DWH_API_KEY}
ssl_ca: ${THT_SSL_CA}
# Accesso al DWH via REST (richiesto se database.transport = rest).
rest:
base_url: ${THT_DWH_REST_URL} # es. https://supabase-aritmolab.policlinicosandonato.it/dwh/
api_key: ${THT_DWH_API_KEY} # header X-API-Key, ruolo dwh_reader (read-only)
ssl_ca: ${THT_SSL_CA} # path al certificato CA (per server con CA interna)
paths:
roots:
artifacts: artifacts
indexes: indexes
sessions: sessions
@@ -50,30 +43,23 @@ embeddings:
dim: 768
batch_size: 32
# LOADING del pgvector: connessione diretta, eseguita sul server (profilo server).
# Opzionale su postazione remota (lì la lettura passa da vector_rest).
vector_db:
host: ${THT_VEC_HOST} # Postgres locale del server
port: ${THT_VEC_PORT} # es. 5437
database: postgres
schema: vectors
user: ${THT_VEC_USER}
password: ${THT_VEC_PASSWORD}
# LETTURA (similarity search) del pgvector via REST remota: rpc search_similar.
vector_rest:
base_url: ${THT_VEC_REST_URL} # es. https://host/vector/v1/
api_key: ${THT_VEC_API_KEY} # header X-API-Key, ruolo vector_reader (read-only)
ssl_ca: ${THT_SSL_CA}
# SCRITTURA controllata del pgvector via REST remota: upsert/hash via RPC allowlist,
# niente delete/clear. Usa una API key SEPARATA dalla lettura (ruolo vector_writer).
# OPZIONALE: assente o key vuota = scrittura non abilitata (solo lettura).
# Abilita tht memory save-one / vector index-schema da postazione remota.
vector_write_rest:
base_url: ${THT_VEC_REST_URL}
api_key: ${THT_VEC_WRITE_API_KEY}
ssl_ca: ${THT_SSL_CA}
vectors:
type: thoth_vector_http # pgvector_direct | thoth_vector_http
reader:
base_url: ${THT_VEC_REST_URL}
api_key: ${THT_VEC_API_KEY}
ssl_ca: ${THT_SSL_CA}
writer: # opzionale: credenziale separata dalla lettura
base_url: ${THT_VEC_REST_URL}
api_key: ${THT_VEC_WRITE_API_KEY}
ssl_ca: ${THT_SSL_CA}
direct: # opzionale: loading server-side diretto
host: ${THT_VEC_HOST}
port: ${THT_VEC_PORT}
database: postgres
schema: vectors
user: ${THT_VEC_USER}
password: ${THT_VEC_PASSWORD}
vector:
max_chunk_chars: 4000