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:
@@ -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
|
||||
|
||||
|
||||
@@ -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 $$;
|
||||
|
||||
|
||||
@@ -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"] == []
|
||||
@@ -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")
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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 == ""
|
||||
@@ -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)]
|
||||
@@ -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
|
||||
@@ -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())
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
)
|
||||
@@ -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"]
|
||||
@@ -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),
|
||||
)
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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"]
|
||||
@@ -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")
|
||||
@@ -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"]
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
)
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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", {}))
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
"""Canonical, transport-independent Evidence corpus."""
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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")
|
||||
@@ -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
@@ -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;
|
||||
@@ -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"),
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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: ...
|
||||
@@ -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: ...
|
||||
@@ -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,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()
|
||||
|
||||
@@ -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 ""
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user