Files
ThothII/harness/tht/corpus/models.py
T

164 lines
6.7 KiB
Python

"""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")