diff --git a/README.md b/README.md index 43f03562..ac37de43 100644 --- a/README.md +++ b/README.md @@ -145,7 +145,13 @@ All dispatch modes read **only** a role-specific scheduler, worker URL, WorkerLease signing key, the server-side JSON root registry (`CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON`), and an optional bounded per-file byte ceiling (`CONTEXT_ENGINE_WORKER_MAX_FILE_BYTES`, default 1 MiB, accepted only -within 1–64 MiB). Markdown files are discovered recursively. **A caller may not +within 1–64 MiB). They also require an explicit embedding provider mode and the +schema-pinned dimension (`CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER` and +`CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION`). CI uses the network-free `twin` +mode. Real deployments select `external` and supply endpoint, model, and API key +only through the corresponding `CONTEXT_ENGINE_WORKER_EMBEDDING_*` environment +variables, including a required batch size bounded to 1–256 inputs per request. +Markdown files are discovered recursively. **A caller may not supply Organization, Source, job, or token** — that is the point of the boundary. Output is limited to `dispatched` / `no_work` / `refused`. diff --git a/STATUS.md b/STATUS.md index e88fadaa..a7b49338 100644 --- a/STATUS.md +++ b/STATUS.md @@ -93,6 +93,7 @@ Follow the ADR for its exact evidence boundary. | [0059](./docs/decisions/0059-dispatch-scheduled-file-imports-through-exact-leases.md) | Dispatch scheduled File imports through exact leases | | [0060](./docs/decisions/0060-reclaim-expired-file-imports-with-bounded-retries.md) | Reclaim expired File imports with bounded retries | | [0065](./docs/decisions/0065-recurse-file-discovery-with-anchored-descriptors.md) | Recurse File discovery through anchored descriptors under one bounded byte ceiling | +| [0066](./docs/decisions/0066-embed-fragments-before-publication.md) | Embed newly published Fragments before activation through an explicit provider | ADR-0065 extends the active File Provider boundary from a flat root to deterministic recursive discovery of canonical nested Markdown paths. Each @@ -106,6 +107,15 @@ PostgreSQL evidence covers nested publication plus mixed flat/nested replay. This does **not** activate provider polling/watchers, a full-resync mechanism, new delete authority, or any non-Markdown carrier. +ADR-0066 adds one Supply-owned embedding seam to File publication. New Fragment +rows receive validated 384-dimensional float32 vectors in the same durable +publication boundary before activation; unchanged acquisitions and recovery +past preparation do not call the provider again. The partial HNSW index is a +future candidate-discovery implementation detail and has no authorization role. + +This does **not** activate vector retrieval, query embedding, historical +backfill, or any Runtime/AuthorizationKernel change. + ### Wire contract, SDK, and trusted delivery | ADR | Activates | diff --git a/adapters/embeddings.py b/adapters/embeddings.py new file mode 100644 index 00000000..728e79ae --- /dev/null +++ b/adapters/embeddings.py @@ -0,0 +1,231 @@ +"""Network-free and external adapters for the Supply embedding seam.""" + +from __future__ import annotations + +import json +from collections.abc import Callable +from contextlib import closing +from dataclasses import dataclass, field +from hashlib import shake_256 +from math import sqrt +from typing import IO, BinaryIO, cast +from urllib.error import HTTPError +from urllib.parse import urlsplit +from urllib.request import HTTPRedirectHandler, Request, build_opener + +from engine.supply.embeddings import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + EmbeddingProfile, + EmbeddingProviderUnavailable, + EmbeddingVector, + validate_embedding_batch, +) + +_MAX_EXTERNAL_RESPONSE_BYTES = 64 * 1024 * 1024 +_DEFAULT_TIMEOUT_SECONDS = 30.0 +EmbeddingTransport = Callable[[Request, float, int], bytes] + + +class _RejectRedirectHandler(HTTPRedirectHandler): + """Keep the configured endpoint as the only bearer-credential recipient.""" + + def redirect_request( + self, + request: Request, + fp: IO[bytes], + code: int, + message: str, + headers: object, + new_url: str, + ) -> Request: + del message, new_url + raise HTTPError( + request.full_url, + code, + "Embedding redirect is unavailable", + headers, # type: ignore[arg-type] + fp, + ) + + +@dataclass(frozen=True, slots=True) +class ExternalEmbeddingConfiguration: + """Environment-derived external provider configuration.""" + + endpoint: str = field(repr=False) + model: str + api_key: str = field(repr=False) + dimension: int + batch_size: int + timeout_seconds: float = _DEFAULT_TIMEOUT_SECONDS + + def __post_init__(self) -> None: + parsed = urlsplit(self.endpoint) + if ( + type(self.endpoint) is not str + or not self.endpoint + or self.endpoint != self.endpoint.strip() + or parsed.scheme != "https" + or not parsed.hostname + or parsed.username is not None + or parsed.password is not None + or bool(parsed.query) + or bool(parsed.fragment) + or type(self.model) is not str + or not self.model + or self.model != self.model.strip() + or type(self.api_key) is not str + or not self.api_key + or self.api_key != self.api_key.strip() + or type(self.timeout_seconds) not in {int, float} + or not 0 < float(self.timeout_seconds) <= 120 + or type(self.batch_size) is not int + or not 1 <= self.batch_size <= 256 + ): + raise ValueError("Embedding configuration is not available") + EmbeddingProfile(self.dimension) + + +def _default_transport(request: Request, timeout: float, maximum_bytes: int) -> bytes: + with closing( + cast( + BinaryIO, + build_opener(_RejectRedirectHandler()).open( # noqa: S310 + request, + timeout=timeout, + ), + ) + ) as response: + payload = response.read(maximum_bytes + 1) + if len(payload) > maximum_bytes: + raise OSError("embedding response exceeded the configured bound") + return payload + + +class ExternalEmbeddingProvider: + """Call one environment-configured JSON embedding endpoint.""" + + __slots__ = ("_configuration", "_transport") + + def __init__( + self, + configuration: ExternalEmbeddingConfiguration, + *, + transport: EmbeddingTransport = _default_transport, + ) -> None: + if type(configuration) is not ExternalEmbeddingConfiguration: + raise TypeError("External embedding configuration is required") + if not callable(transport): + raise TypeError("External embedding transport is required") + self._configuration = configuration + self._transport = transport + + @property + def profile(self) -> EmbeddingProfile: + return EmbeddingProfile(self._configuration.dimension) + + def embed(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: + if ( + type(inputs) is not tuple + or not inputs + or any(type(value) is not str or not value for value in inputs) + ): + raise EmbeddingProviderUnavailable("Embedding provider is unavailable") + try: + vectors: list[EmbeddingVector] = [] + for offset in range(0, len(inputs), self._configuration.batch_size): + batch = inputs[offset : offset + self._configuration.batch_size] + vectors.extend(self._embed_batch(batch)) + return tuple(vectors) + except Exception: + raise EmbeddingProviderUnavailable( + "Embedding provider is unavailable" + ) from None + + def _embed_batch(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: + body = json.dumps( + { + "dimensions": self.profile.dimension, + "encoding_format": "float", + "input": list(inputs), + "model": self._configuration.model, + }, + ensure_ascii=False, + separators=(",", ":"), + ).encode("utf-8") + request = Request( + self._configuration.endpoint, + data=body, + headers={ + "Accept": "application/json", + "Authorization": f"Bearer {self._configuration.api_key}", + "Content-Type": "application/json", + }, + method="POST", + ) + raw_response = self._transport( + request, + float(self._configuration.timeout_seconds), + _MAX_EXTERNAL_RESPONSE_BYTES, + ) + response = json.loads(raw_response) + raw_data = response["data"] + if type(raw_data) is not list or len(raw_data) != len(inputs): + raise ValueError + ordered: list[list[object] | None] = [None] * len(inputs) + for item in raw_data: + if type(item) is not dict: + raise ValueError + index = item.get("index") + vector = item.get("embedding") + if ( + type(index) is not int + or not 0 <= index < len(inputs) + or ordered[index] is not None + or type(vector) is not list + ): + raise ValueError + ordered[index] = cast(list[object], vector) + if any(vector is None for vector in ordered): + raise ValueError + return validate_embedding_batch( + inputs, + cast(list[list[object]], ordered), + self.profile, + ) + + +class DeterministicEmbeddingTwin: + """Stable content-derived vectors for tests without network egress.""" + + __slots__ = ("_profile",) + + def __init__( + self, + dimension: int = CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + ) -> None: + self._profile = EmbeddingProfile(dimension) + + @property + def profile(self) -> EmbeddingProfile: + return self._profile + + def embed(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: + if ( + type(inputs) is not tuple + or not inputs + or any(type(value) is not str or not value for value in inputs) + ): + raise EmbeddingProviderUnavailable("Embedding provider is unavailable") + vectors: list[EmbeddingVector] = [] + for value in inputs: + raw = shake_256( + b"context-engine.embedding-twin.v1\x00" + value.encode("utf-8") + ).digest(self.profile.dimension * 2) + unscaled = tuple( + (int.from_bytes(raw[offset : offset + 2], "big") - 32767.5) / 32767.5 + for offset in range(0, len(raw), 2) + ) + norm = sqrt(sum(component * component for component in unscaled)) + vectors.append(tuple(component / norm for component in unscaled)) + return validate_embedding_batch(inputs, vectors, self.profile) diff --git a/applications/worker.py b/applications/worker.py index 9d44f439..c6da7af0 100644 --- a/applications/worker.py +++ b/applications/worker.py @@ -15,6 +15,11 @@ from sqlalchemy import Engine, text from sqlalchemy.exc import SQLAlchemyError +from adapters.embeddings import ( + DeterministicEmbeddingTwin, + ExternalEmbeddingConfiguration, + ExternalEmbeddingProvider, +) from adapters.file_source import FileReadLimits, FileRootRegistry from engine import BUILD_IDENTIFIER from engine.control import FileImportReceiver, FileRootRef, SourceRef @@ -38,6 +43,8 @@ from engine.runtime import Runtime from engine.runtime.construction import required_kernel_dependencies from engine.supply import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + EmbeddingProvider, MarkdownCompilerConfig, WorkerLeaseCodec, WorkerLeaseKeyring, @@ -48,6 +55,8 @@ _FILE_DISPATCH_POLL_SECONDS = 1.0 DEFAULT_WORKER_MAX_FILE_BYTES = 1_048_576 _WORKER_MAX_FILE_BYTES_ENV = "CONTEXT_ENGINE_WORKER_MAX_FILE_BYTES" +_WORKER_EMBEDDING_PROVIDER_ENV = "CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER" +_WORKER_EMBEDDING_DIMENSION_ENV = "CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION" class WorkerNoOpCompletionAuthority(Protocol): @@ -140,6 +149,24 @@ def _required_environment(name: str) -> str: return value +def _required_bounded_integer_environment( + name: str, + *, + minimum: int, + maximum: int, +) -> int: + raw_value = _required_environment(name) + if not raw_value.isascii() or not raw_value.isdecimal(): + raise ValueError("Supply worker configuration is not available") + try: + value = int(raw_value) + except ValueError: + raise ValueError("Supply worker configuration is not available") from None + if not minimum <= value <= maximum: + raise ValueError("Supply worker configuration is not available") + return value + + def _file_read_limits() -> FileReadLimits: raw_limit = os.environ.get(_WORKER_MAX_FILE_BYTES_ENV) if raw_limit is None: @@ -152,6 +179,47 @@ def _file_read_limits() -> FileReadLimits: raise ValueError("Supply worker configuration is not available") from None +def _embedding_provider() -> EmbeddingProvider: + """Compose the explicit CI twin or one environment-only external provider.""" + + mode = _required_environment(_WORKER_EMBEDDING_PROVIDER_ENV) + raw_dimension = _required_environment(_WORKER_EMBEDDING_DIMENSION_ENV) + if not raw_dimension.isdecimal(): + raise ValueError("Supply worker configuration is not available") + try: + dimension = int(raw_dimension) + except ValueError: + raise ValueError("Supply worker configuration is not available") from None + if dimension != CONTEXT_FRAGMENT_EMBEDDING_DIMENSION: + raise ValueError("Supply worker configuration is not available") + if mode == "twin": + return DeterministicEmbeddingTwin(dimension) + if mode != "external": + raise ValueError("Supply worker configuration is not available") + raw_timeout = os.environ.get("CONTEXT_ENGINE_WORKER_EMBEDDING_TIMEOUT_SECONDS") + if raw_timeout is None: + timeout_seconds = 30.0 + else: + try: + timeout_seconds = float(raw_timeout) + except ValueError: + raise ValueError("Supply worker configuration is not available") from None + return ExternalEmbeddingProvider( + ExternalEmbeddingConfiguration( + endpoint=_required_environment("CONTEXT_ENGINE_WORKER_EMBEDDING_ENDPOINT"), + model=_required_environment("CONTEXT_ENGINE_WORKER_EMBEDDING_MODEL"), + api_key=_required_environment("CONTEXT_ENGINE_WORKER_EMBEDDING_API_KEY"), + dimension=dimension, + batch_size=_required_bounded_integer_environment( + "CONTEXT_ENGINE_WORKER_EMBEDDING_BATCH_SIZE", + minimum=1, + maximum=256, + ), + timeout_seconds=timeout_seconds, + ) + ) + + def _run_one_file_import() -> int: """Consume one exact, signed File job in the independent Supply process.""" @@ -180,6 +248,7 @@ def _run_one_file_import() -> int: ), roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=_embedding_provider(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -289,6 +358,7 @@ def _worker_database_time(engine: Engine) -> datetime: def _run_file_dispatch(*, single_cycle: bool) -> int: """Run configured autonomous File dispatch without caller routing facts.""" + embedding_provider = _embedding_provider() codec = WorkerLeaseCodec( WorkerLeaseKeyring(active_version=1, keys={1: _worker_signing_key()}) ) @@ -319,6 +389,7 @@ def worker_factory(receiver: FileImportReceiver) -> PostgreSQLFileImportWorker: receiver, roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=embedding_provider, clock=lambda: _worker_database_time(worker_engine), ) diff --git a/docs/decisions/0066-embed-fragments-before-publication.md b/docs/decisions/0066-embed-fragments-before-publication.md new file mode 100644 index 00000000..9c69c09d --- /dev/null +++ b/docs/decisions/0066-embed-fragments-before-publication.md @@ -0,0 +1,93 @@ +--- +name: adr-0066-embed-fragments-before-publication +version: "1.0.0" +description: > + Persist fixed-dimension pgvector embeddings with newly prepared Fragments + through an explicit external provider or deterministic CI twin. +--- + +# 0066. Embed Fragments before publication + +- Status: accepted +- Date: 2026-07-26 +- Refines: ADR-0009, ADR-0037, ADR-0041, ADR-0064 + +## Context + +The first dogfood retrieval slice needs a real vector-bearing corpus before a +vector CandidateIndex can be activated. Fragment rows are immutable once +prepared, and publication recovery must never activate a partially enriched +Revision. Embedding calls are fallible and can be costly, while unchanged File +acquisitions already have an authoritative no-op classification that must not +repeat derived work. + +The storage dimension, provider response shape, and migration compatibility are +hard to reverse. Allowing an implicit test provider in production composition +would also turn a network-free twin into accidental product behavior. + +## Decision + +Supply owns a small batch `EmbeddingProvider` seam. The current schema profile +pins vectors to 384 dimensions in one source constant shared by worker +composition, the deterministic twin, and response validation. The migration +pins the PostgreSQL column to the same value, with a mechanical equality test +preventing drift between the two declarations. Worker composition requires an +explicit provider mode and dimension. A dimension other than the schema profile +is rejected before work begins. + +The external adapter sends contextual Fragment text to one environment-derived +HTTPS JSON endpoint. Endpoint, model, API key, timeout, and dimension enter only +through worker environment configuration. A required 1–256 batch-size setting +bounds each request while the adapter reassembles validated batches in original +Fragment order. Secret-bearing values are excluded from representations, and +transport, status, parsing, ordering, count, +dimension, non-finite, and float32-zero-vector failures collapse to one +content-free unavailability category. Provider values are normalized to the +same IEEE-754 float32 representation pgvector persists before the nonzero check. + +The CI and integration twin is network-free. It derives a normalized vector +from a domain-separated SHAKE-256 stream over exact contextual text, giving +stable content-derived values without claiming semantic quality. + +File publication first performs the existing acquisition and unchanged-content +classification. Only an `acquired` new or replacement Revision calls the +provider. Its complete validated embedding document enters the same transaction +that inserts immutable Fragment rows and advances `acquired -> prepared`. +Provider failure records the existing acquired-boundary interruption, leaving no +Revision or Fragment rows and allowing the bounded lease-reclaim path to retry. +Recovery from `prepared` or `ready` reuses stored vectors and does not call the +provider again. Indexing and activation both reject any current Revision with a +missing or wrong-dimension vector. + +During a rolling schema change, the worker detects the installed prepare +signature: it keeps the pre-embedding publication contract while the database +is at the predecessor revision and uses the vector-bearing contract only after +0036 is active. The upgrade takes exclusive locks and rewinds any inactive, +pre-0036 `prepared` or `ready` Revision to `acquired`, removing only its staged +derived rows. The next bounded recovery lease therefore re-embeds through the +configured provider before activation; migration code never fabricates vectors. + +`context_fragment.embedding` is nullable only for historical rows because this +slice deliberately adds no backfill authority. New publication requires a +vector. One partial HNSW cosine index covers embedded rows. It narrows future +candidate discovery only and never participates in authorization. + +## Consequences + +- Unchanged acquisitions perform zero embedding calls and keep their current + active lineage. +- A newly active Revision always has one validated vector per Fragment. +- Historical rows remain readable and are embedded only by re-importing them as + a new immutable Revision. +- CI has no embedding network egress; production has no implicit twin fallback. +- The migration downgrade takes an exclusive Fragment-table lock, removes only + the derived vectors, and preserves immutable Fragment content and lineage. +- Runtime, AuthorizationKernel, projection, and served composition do not + change in this slice. + +## Revisit trigger + +Revisit before changing the dimension or embedding input profile, adding model +version lineage, backfill authority, multiple embedding profiles, query-time +embedding, or a non-PostgreSQL vector store. Query-time candidate discovery and +its recall/latency evidence belong to the next vector CandidateIndex decision. diff --git a/docs/decisions/README.md b/docs/decisions/README.md index 386083d9..19a0e62c 100644 --- a/docs/decisions/README.md +++ b/docs/decisions/README.md @@ -51,6 +51,7 @@ kernel, capability separation, and publication visibility model. | Autonomous File dispatch | [0059 — Dispatch scheduled File imports through exact leases](0059-dispatch-scheduled-file-imports-through-exact-leases.md) | A function-only scheduler login atomically claims the oldest current page-scheduled upsert and mints the existing exact first-attempt WorkerLease | Caller tenant/job routing, broad Control credentials, direct scheduler table access, retry/reclaim, delete execution, or a second queue/process | | Bounded File reclaim | [0060 — Reclaim expired File imports with bounded retries](0060-reclaim-expired-file-imports-with-bounded-retries.md) | The same function-only scheduler prefers database-timed expired work, revalidates exact current authority, and advances at most three higher WorkerLease generations through durable-boundary recovery | Caller-selected retry routing/timing/generation, terminal-failure retry, dead-letter or requeue authority, delete execution, Runtime authority, or a new process/queue | | Recursive File discovery | [0065 — Recurse File discovery with anchored descriptors](0065-recurse-file-discovery-with-anchored-descriptors.md) | Canonical nested Markdown paths are discovered through stable descriptor-relative no-follow traversal under one server-owned 1–64 MiB byte ceiling | Path-string traversal, symlink following, truncated baselines, unbounded or caller-owned read limits, non-Markdown discovery, or new polling/resync authority | +| Fragment embedding publication | [0066 — Embed Fragments before publication](0066-embed-fragments-before-publication.md) | New immutable Fragments receive one validated 384-dimensional vector before activation through an explicit external provider or network-free CI twin | Implicit twin fallback, query-time retrieval, backfill authority, missing-vector activation, or vector-based authorization | | Private delivery ingress | [0045 — Redeem private delivery evidence at ingress](0045-redeem-private-delivery-evidence-at-ingress.md) | One digest-only service/request/asker/audience/epoch-bound DeliveryEvidenceRef constructs private TrustedDeliveryContext inside the current UserActor transaction before content work | Raw trusted delivery facts on the wire, bearer persistence, application-role minting/table reads, alternate Runtime paths, or claiming later M2 carriers | | Exact Package egress | [0046 — Bind egress to one exact Package hop](0046-bind-egress-to-one-exact-package-hop.md) | One digest-only grant binds one exact audience-bound Package to one model or channel preflight hop and redeems atomically | Treating Package construction as disclosure authority, arbitrary content at egress, cross-hop reuse, or bypassing final policy | | Public OpenAPI v0 | [0047 — Freeze OpenAPI v0 through one Runtime path](0047-freeze-openapi-v0-through-one-runtime-path.md) | One public `/v0/resolve` schema and a hidden provisional v1 bridge share the same sealed Runtime; Package release lineage is read-only from the Learning-published active manifest | Two authorization compositions, caller-authored release facts, Runtime publication/fallback, or in-place mutation of historical snapshots | @@ -164,3 +165,4 @@ touched: - [0063 — Admit an explicit dogfood authentication composition](0063-admit-an-explicit-dogfood-authentication-composition.md) - [0064 — Split process ceremony along the kernel-seam boundary](0064-split-process-ceremony-along-the-kernel-seam-boundary.md) - [0065 — Recursive File discovery with anchored descriptors](0065-recurse-file-discovery-with-anchored-descriptors.md) +- [0066 — Embed Fragments before publication](0066-embed-fragments-before-publication.md) diff --git a/engine/persistence/file_imports.py b/engine/persistence/file_imports.py index d35047c6..a052edba 100644 --- a/engine/persistence/file_imports.py +++ b/engine/persistence/file_imports.py @@ -26,8 +26,12 @@ from engine.persistence.role_guard import assert_worker_role from engine.runtime.evidence import CandidateRef from engine.supply import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, FILE_IMPORT_WORKER_LEASE_OPERATION, CompilationFailure, + EmbeddingProfile, + EmbeddingProvider, + EmbeddingProviderUnavailable, MarkdownCompilerConfig, ParsedDocument, WorkerLeaseClaims, @@ -36,12 +40,18 @@ WorkerLeaseToken, WorkNotAvailable, canonicalize_parsed_document, + validate_embedding_batch, worker_lease_digest, ) from engine.supply.jobs import _require_utc _CONCURRENT_PUBLICATION_WAIT_SECONDS = 5.0 _CONCURRENT_PUBLICATION_POLL_SECONDS = 0.01 +_EMBEDDING_PREPARE_REGPROCEDURE = ( + "public.context_worker_prepare_file_publication" + "(uuid,uuid,uuid,text,text,uuid,text,jsonb,jsonb,jsonb,bigint,bigint,bytea," + "timestamp with time zone,timestamp with time zone)" +) @dataclass(frozen=True, slots=True) @@ -149,7 +159,7 @@ class FileImportRefused(FileImportUnavailable): class FilePublicationBoundary(StrEnum): - """The three explicit post-commit fault-injection boundaries.""" + """The three explicit durable publication recovery boundaries.""" ACQUIRED = "acquired" PREPARED = "prepared" @@ -157,7 +167,7 @@ class FilePublicationBoundary(StrEnum): class FileImportInterrupted(RuntimeError): - """Deterministic test interruption recorded after a durable boundary.""" + """Content-free resumable interruption recorded at a durable boundary.""" def __init__(self, boundary: FilePublicationBoundary) -> None: if type(boundary) is not FilePublicationBoundary: @@ -199,6 +209,8 @@ class PostgreSQLFileImportWorker: "_codec", "_config", "_engine", + "_embedding_profile", + "_embedding_provider", "_identity", "_interrupt_after", "_roots", @@ -213,6 +225,7 @@ def __init__( roots: FileRootRegistry, config: MarkdownCompilerConfig, *, + embedding_provider: EmbeddingProvider, clock: Callable[[], object], uuid_factory: Callable[[], UUID] = uuid4, interrupt_after: FilePublicationBoundary | None = None, @@ -225,6 +238,16 @@ def __init__( raise TypeError("File import worker requires FileRootRegistry") if type(config) is not MarkdownCompilerConfig: raise TypeError("File import worker requires MarkdownCompilerConfig") + try: + embedding_profile = embedding_provider.profile + except (AttributeError, TypeError, ValueError): + raise TypeError( + "File import worker requires an embedding provider" + ) from None + if type(embedding_profile) is not EmbeddingProfile: + raise TypeError("File import worker requires an embedding provider") + if embedding_profile.dimension != CONTEXT_FRAGMENT_EMBEDDING_DIMENSION: + raise ValueError("Embedding provider dimension does not match storage") if not callable(clock) or not callable(uuid_factory): raise TypeError("File import worker requires clock and UUID factory") if ( @@ -237,6 +260,8 @@ def __init__( self._identity = identity self._roots = roots self._config = config + self._embedding_profile = embedding_profile + self._embedding_provider = embedding_provider self._clock = clock self._uuid_factory = uuid_factory self._interrupt_after = interrupt_after @@ -302,12 +327,19 @@ def _redeem( row = connection.execute( text( """ - SELECT * FROM public.context_worker_redeem_file_import( + SELECT redeemed.source_ref, redeemed.root_ref, + redeemed.relative_path, + redeemed.acquisition_id, + to_jsonb(redeemed)->>'expected_content_sha256' + AS expected_content_sha256, + (to_jsonb(redeemed)->>'expected_content_length') + ::bigint AS expected_content_length + FROM public.context_worker_redeem_file_import( :organization_id, :job_id, :service_principal_id, :source_ref, :lease_generation, :signing_key_version, :nonce, :issued_at, :expires_at - ) + ) AS redeemed """ ), { @@ -324,12 +356,8 @@ def _redeem( ).one_or_none() if row is None or row.source_ref != claims.source_ref: raise _rejection(token) - expected_content_sha256 = row._mapping.get( - "expected_content_sha256" - ) - expected_content_length = row._mapping.get( - "expected_content_length" - ) + expected_content_sha256 = row._mapping.get("expected_content_sha256") + expected_content_length = row._mapping.get("expected_content_length") return _RedeemedFileImport( source_ref=SourceRef(UUID(row.source_ref)), root_ref=FileRootRef(row.root_ref), @@ -427,19 +455,8 @@ def _publish( self._interrupt_if_requested( token, claims, FilePublicationBoundary.ACQUIRED ) - prepared = self._execute_one( - """ - SELECT * FROM public.context_worker_prepare_file_publication( - :organization_id, :job_id, :service_principal_id, - :source_ref, :resource_ref, :revision_id, - :canonical_text, - CAST(:compilation_document AS jsonb), - CAST(:artifact_document AS jsonb), - :lease_generation, :signing_key_version, :nonce, - :issued_at, :expires_at - ) - """, - parameters, + prepared = self._prepare_publication( + token, claims, document, parameters ) if prepared is None or prepared.checkpoint != "prepared": raise _rejection(token) @@ -515,6 +532,60 @@ def _publish( effect_count=row.effect_count, ) + def _prepare_publication( + self, + token: WorkerLeaseToken, + claims: WorkerLeaseClaims, + document: ParsedDocument, + parameters: dict[str, object], + ) -> Row[tuple[object, ...]] | None: + if self._embedding_storage_active(): + parameters["embedding_document"] = self._embedding_document( + token, + claims, + document, + ) + return self._execute_one( + """ + SELECT * FROM public.context_worker_prepare_file_publication( + :organization_id, :job_id, :service_principal_id, + :source_ref, :resource_ref, :revision_id, + :canonical_text, + CAST(:compilation_document AS jsonb), + CAST(:artifact_document AS jsonb), + CAST(:embedding_document AS jsonb), + :lease_generation, :signing_key_version, :nonce, + :issued_at, :expires_at + ) + """, + parameters, + ) + return self._execute_one( + """ + SELECT * FROM public.context_worker_prepare_file_publication( + :organization_id, :job_id, :service_principal_id, + :source_ref, :resource_ref, :revision_id, + :canonical_text, + CAST(:compilation_document AS jsonb), + CAST(:artifact_document AS jsonb), + :lease_generation, :signing_key_version, :nonce, + :issued_at, :expires_at + ) + """, + parameters, + ) + + def _embedding_storage_active(self) -> bool: + with self._engine.begin() as connection: + assert_worker_role(connection) + available = connection.execute( + text( + "SELECT pg_catalog.to_regprocedure(:signature) IS NOT NULL" + ), + {"signature": _EMBEDDING_PREPARE_REGPROCEDURE}, + ).scalar_one() + return available is True + def _await_concurrent_publication( self, token: WorkerLeaseToken, @@ -563,6 +634,47 @@ def _interrupt_if_requested( ) -> None: if self._interrupt_after is not boundary: return + self._record_interruption(token, claims, boundary) + raise FileImportInterrupted(boundary) + + def _embedding_document( + self, + token: WorkerLeaseToken, + claims: WorkerLeaseClaims, + document: ParsedDocument, + ) -> str: + inputs = tuple(fragment.contextual_text for fragment in document.fragments) + try: + vectors = validate_embedding_batch( + inputs, + self._embedding_provider.embed(inputs), + self._embedding_profile, + ) + except EmbeddingProviderUnavailable: + self._record_interruption( + token, + claims, + FilePublicationBoundary.ACQUIRED, + ) + raise FileImportInterrupted(FilePublicationBoundary.ACQUIRED) from None + return json.dumps( + [ + { + "embedding": list(vector), + "fragmentRef": fragment.fragment_ref, + } + for fragment, vector in zip(document.fragments, vectors, strict=True) + ], + ensure_ascii=False, + separators=(",", ":"), + ) + + def _record_interruption( + self, + token: WorkerLeaseToken, + claims: WorkerLeaseClaims, + boundary: FilePublicationBoundary, + ) -> None: with self._engine.begin() as connection: assert_worker_role(connection) recorded = connection.execute( @@ -590,7 +702,6 @@ def _interrupt_if_requested( ).scalar_one() if recorded is not True: raise _rejection(token) - raise FileImportInterrupted(boundary) def _fail( self, diff --git a/engine/supply/__init__.py b/engine/supply/__init__.py index fdd1d2b5..9a3772fc 100644 --- a/engine/supply/__init__.py +++ b/engine/supply/__init__.py @@ -7,6 +7,14 @@ PreparedFileImport, PrepareFileImport, ) +from engine.supply.embeddings import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + EmbeddingProfile, + EmbeddingProvider, + EmbeddingProviderUnavailable, + EmbeddingVector, + validate_embedding_batch, +) from engine.supply.jobs import ( FILE_IMPORT_WORKER_LEASE_OPERATION, WORKER_LEASE_ACTOR_KIND, @@ -50,6 +58,7 @@ ) __all__ = [ + "CONTEXT_FRAGMENT_EMBEDDING_DIMENSION", "MARKDOWN_CANONICALIZATION_PROFILE", "MARKDOWN_CODE_LANGUAGE_MAX_LENGTH", "MARKDOWN_CANONICALIZATION_V1_PROFILE", @@ -66,6 +75,10 @@ "CompilationProvenance", "CompilationWarning", "CompilationWarningCode", + "EmbeddingProfile", + "EmbeddingProvider", + "EmbeddingProviderUnavailable", + "EmbeddingVector", "CompiledFragment", "FileImportAudience", "FileImportPath", @@ -92,4 +105,5 @@ "canonicalize_parsed_document", "worker_lease_digest", "worker_lease_nonce_digest", + "validate_embedding_batch", ] diff --git a/engine/supply/embeddings.py b/engine/supply/embeddings.py new file mode 100644 index 00000000..e9615bee --- /dev/null +++ b/engine/supply/embeddings.py @@ -0,0 +1,84 @@ +"""Supply-only embedding contracts for immutable Fragment publication.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from math import isfinite +from struct import Struct +from struct import error as StructError +from typing import Protocol, cast + +CONTEXT_FRAGMENT_EMBEDDING_DIMENSION = 384 +type EmbeddingVector = tuple[float, ...] +_FLOAT32 = Struct("!f") + + +@dataclass(frozen=True, slots=True) +class EmbeddingProfile: + """Closed vector shape stored by the current PostgreSQL schema.""" + + dimension: int + + def __post_init__(self) -> None: + if type(self.dimension) is not int or not 1 <= self.dimension <= 4096: + raise ValueError("Embedding dimension is not available") + + +class EmbeddingProviderUnavailable(RuntimeError): + """Content-free transient provider failure.""" + + +class EmbeddingProvider(Protocol): + """Batch embedding seam used only by Supply publication.""" + + @property + def profile(self) -> EmbeddingProfile: ... + + def embed(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: ... + + +def validate_embedding_batch( + inputs: tuple[str, ...], + vectors: Sequence[Sequence[object]], + profile: EmbeddingProfile, +) -> tuple[EmbeddingVector, ...]: + """Validate one provider response before any vector crosses persistence.""" + + try: + if ( + type(inputs) is not tuple + or not inputs + or any(type(value) is not str or not value for value in inputs) + or len(vectors) != len(inputs) + ): + raise EmbeddingProviderUnavailable("Embedding provider is unavailable") + validated: list[EmbeddingVector] = [] + for raw_vector in vectors: + if len(raw_vector) != profile.dimension: + raise EmbeddingProviderUnavailable("Embedding provider is unavailable") + vector: list[float] = [] + for raw_value in raw_vector: + if type(raw_value) not in {int, float}: + raise EmbeddingProviderUnavailable( + "Embedding provider is unavailable" + ) + value = float(cast(int | float, raw_value)) + if not isfinite(value) or abs(value) > 1.0e30: + raise EmbeddingProviderUnavailable( + "Embedding provider is unavailable" + ) + stored_value = _FLOAT32.unpack(_FLOAT32.pack(value))[0] + if not isfinite(stored_value) or abs(stored_value) > 1.0e30: + raise EmbeddingProviderUnavailable( + "Embedding provider is unavailable" + ) + vector.append(stored_value) + if not any(value != 0.0 for value in vector): + raise EmbeddingProviderUnavailable("Embedding provider is unavailable") + validated.append(tuple(vector)) + return tuple(validated) + except (TypeError, ValueError, OverflowError, StructError): + raise EmbeddingProviderUnavailable( + "Embedding provider is unavailable" + ) from None diff --git a/migrations/versions/20260726_0036_fragment_embeddings.py b/migrations/versions/20260726_0036_fragment_embeddings.py new file mode 100644 index 00000000..1a9a1c09 --- /dev/null +++ b/migrations/versions/20260726_0036_fragment_embeddings.py @@ -0,0 +1,546 @@ +"""Persist Supply-side Fragment embeddings before Revision activation. + +Revision ID: 20260726_0036 +Revises: 20260726_0035 +Create Date: 2026-07-26 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "20260726_0036" +down_revision: str | None = "20260726_0035" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +_DEFINER = "context_engine_worker_lease_definer" +_WORKER = "context_engine_worker" +_DIMENSION = 384 +_PREPARE = "context_worker_prepare_file_publication" +_OLD_PREPARE_SIGNATURE = ( + "(uuid,uuid,uuid,text,text,uuid,text,jsonb,jsonb,bigint,bigint,bytea," + "timestamp with time zone,timestamp with time zone)" +) +_NEW_PREPARE_SIGNATURE = ( + "(uuid,uuid,uuid,text,text,uuid,text,jsonb,jsonb,jsonb,bigint,bigint,bytea," + "timestamp with time zone,timestamp with time zone)" +) +_INDEX_REGPROCEDURE = ( + "context_worker_index_file_publication" + "(uuid,uuid,uuid,text,text,uuid,text,jsonb,jsonb,bigint,bigint,bytea," + "timestamp with time zone,timestamp with time zone)" +) +_ACQUIRE_REGPROCEDURE = ( + "context_worker_acquire_file_publication" + "(uuid,uuid,uuid,text,text,uuid,text,text,text,text,text,jsonb,jsonb," + "bigint,bigint,bytea,timestamp with time zone,timestamp with time zone)" +) +_CLASSIFY_REGPROCEDURE = ( + "context_worker_classify_file_import_internal" + "(uuid,uuid,uuid,text,text,text,text,text,text,bigint,bytea," + "timestamp with time zone,timestamp with time zone)" +) +_ACTIVATE_REGPROCEDURE = ( + "context_worker_activate_recoverable_file_publication" + "(uuid,uuid,uuid,text,text,uuid,bigint,bigint,bytea," + "timestamp with time zone,timestamp with time zone)" +) +_SIGNATURE_OLD = "requested_artifact_document jsonb, requested_lease_generation" +_SIGNATURE_NEW = ( + "requested_artifact_document jsonb, requested_embedding_document jsonb, " + "requested_lease_generation" +) +_ENTRY_ANCHOR = ( + "THEN RETURN; END IF;\n" + " PERFORM pg_catalog.set_config('app.organization_id'" +) +_EMBEDDING_VALIDATION = f"""THEN RETURN; END IF; + IF jsonb_typeof(requested_embedding_document) IS DISTINCT FROM 'array' + OR jsonb_array_length(requested_embedding_document) + <> jsonb_array_length(requested_artifact_document) + THEN RETURN; END IF; + IF EXISTS ( + SELECT 1 + FROM jsonb_array_elements(requested_embedding_document) + WITH ORDINALITY AS embedded(item, ordinal) + JOIN jsonb_array_elements(requested_artifact_document) + WITH ORDINALITY AS artifact(item, ordinal) + USING (ordinal) + WHERE jsonb_typeof(embedded.item) IS DISTINCT FROM 'object' + OR embedded.item->>'fragmentRef' + IS DISTINCT FROM artifact.item->>'fragmentRef' + OR jsonb_typeof(embedded.item->'embedding') + IS DISTINCT FROM 'array' + OR jsonb_array_length(embedded.item->'embedding') <> {_DIMENSION} + OR EXISTS ( + SELECT 1 + FROM jsonb_array_elements(embedded.item->'embedding') + AS component(value) + WHERE jsonb_typeof(component.value) <> 'number' + OR char_length(component.value::text) > 64 + OR component.value::text + !~ '^-?[0-9]+(\\.[0-9]+)?([eE][+-]?[0-9]+)?$' + OR CASE + WHEN jsonb_typeof(component.value) = 'number' + AND char_length(component.value::text) <= 64 + AND component.value::text + ~ '^-?[0-9]+(\\.[0-9]+)?([eE][+-]?[0-9]+)?$' + THEN abs((component.value::text)::numeric) > 1.0e30 + ELSE false + END + ) + OR NOT EXISTS ( + SELECT 1 + FROM jsonb_array_elements(embedded.item->'embedding') + AS component(value) + WHERE CASE + WHEN jsonb_typeof(component.value) = 'number' + AND char_length(component.value::text) <= 64 + AND component.value::text + ~ '^-?[0-9]+(\\.[0-9]+)?([eE][+-]?[0-9]+)?$' + THEN abs((component.value::text)::numeric) + > 7.006492321624085e-46 + ELSE false + END + ) + ) THEN RETURN; END IF; + PERFORM pg_catalog.set_config('app.organization_id'""" +_INSERT_OLD = """ INSERT INTO public.context_fragment ( + organization_id, resource_ref, revision_id, fragment_ref, + ordinal, content, projection_kind + ) + SELECT requested_organization_id, requested_resource_ref, + requested_revision_id, item.fragment->>'fragmentRef', + (item.ordinal - 1)::integer, + item.fragment->>'contextualText', 'body' + FROM jsonb_array_elements(requested_artifact_document) + WITH ORDINALITY AS item(fragment, ordinal) + ORDER BY item.ordinal;""" +_INSERT_NEW = """ INSERT INTO public.context_fragment ( + organization_id, resource_ref, revision_id, fragment_ref, + ordinal, content, projection_kind, embedding + ) + SELECT requested_organization_id, requested_resource_ref, + requested_revision_id, item.fragment->>'fragmentRef', + (item.ordinal - 1)::integer, + item.fragment->>'contextualText', 'body', + (embedded.item->'embedding')::text::public.vector + FROM jsonb_array_elements(requested_artifact_document) + WITH ORDINALITY AS item(fragment, ordinal) + JOIN jsonb_array_elements(requested_embedding_document) + WITH ORDINALITY AS embedded(item, ordinal) + USING (ordinal) + ORDER BY item.ordinal;""" +_INDEX_ANCHOR = """ IF recovery_row.job_id IS NULL + OR NOT EXISTS (""" +_NOOP_ANCHOR = """ AND NOT EXISTS ( + SELECT 1 FROM jsonb_array_elements( + requested_artifact_document + ) WITH ORDINALITY AS expected(fragment, ordinal)""" +_NOOP_WITH_EMBEDDINGS = f""" AND NOT EXISTS ( + SELECT 1 FROM public.context_fragment AS fragment + WHERE fragment.organization_id = requested_organization_id + AND fragment.resource_ref = requested_resource_ref + AND fragment.revision_id = decision.active_revision_id + AND (fragment.embedding IS NULL + OR public.vector_dims(fragment.embedding) + <> {_DIMENSION}) + ) + AND NOT EXISTS ( + SELECT 1 FROM jsonb_array_elements( + requested_artifact_document + ) WITH ORDINALITY AS expected(fragment, ordinal)""" +_CLASSIFY_ANCHOR = """ + AND snapshot.config_version = requested_config_version + AND ( + SELECT array_agg(event.state ORDER BY event.ordinal)""" +_CLASSIFY_WITH_EMBEDDINGS = f""" + AND snapshot.config_version = requested_config_version + AND NOT EXISTS ( + SELECT 1 FROM public.context_fragment AS embedded_fragment + WHERE embedded_fragment.organization_id = resource.organization_id + AND embedded_fragment.resource_ref = resource.resource_ref + AND embedded_fragment.revision_id = resource.active_revision_id + AND (embedded_fragment.embedding IS NULL + OR public.vector_dims(embedded_fragment.embedding) + <> {_DIMENSION}) + ) + AND ( + SELECT array_agg(event.state ORDER BY event.ordinal)""" +_INDEX_WITH_EMBEDDINGS = f""" IF recovery_row.job_id IS NULL + OR EXISTS ( + SELECT 1 FROM public.context_fragment AS fragment + WHERE fragment.organization_id = requested_organization_id + AND fragment.resource_ref = requested_resource_ref + AND fragment.revision_id = requested_revision_id + AND (fragment.embedding IS NULL + OR public.vector_dims(fragment.embedding) <> {_DIMENSION}) + ) + OR NOT EXISTS (""" +_ACTIVATE_ANCHOR = """ RETURN QUERY + SELECT wrapped.effect_count""" +_ACTIVATE_WITH_EMBEDDINGS = f""" IF EXISTS ( + SELECT 1 FROM public.context_fragment AS fragment + WHERE fragment.organization_id = requested_organization_id + AND fragment.resource_ref = requested_resource_ref + AND fragment.revision_id = requested_revision_id + AND (fragment.embedding IS NULL + OR public.vector_dims(fragment.embedding) <> {_DIMENSION}) + ) THEN RETURN; END IF; + RETURN QUERY + SELECT wrapped.effect_count""" + + +def _rewind_unembedded_recovery() -> None: + """Return pre-0036 staged Revisions to the provider-backed boundary.""" + + op.execute( + "LOCK TABLE public.file_publication_recovery, public.file_import_job, " + "public.file_revision_replacement_plan, public.exact_phrase_candidate, " + "public.revision_publication_event, public.file_revision_snapshot, " + "public.context_fragment, public.context_revision, " + "public.resource_access_policy, " + "public.membership_resource_field_right " + "IN ACCESS EXCLUSIVE MODE" + ) + op.execute( + "ALTER TABLE public.file_revision_replacement_plan DISABLE TRIGGER " + "file_revision_replacement_plan_immutable" + ) + op.execute( + "ALTER TABLE public.exact_phrase_candidate DISABLE TRIGGER " + "exact_phrase_candidate_immutable" + ) + op.execute( + "ALTER TABLE public.revision_publication_event DISABLE TRIGGER " + "revision_publication_event_immutable" + ) + op.execute( + "ALTER TABLE public.file_revision_snapshot DISABLE TRIGGER " + "file_revision_snapshot_immutable" + ) + op.execute( + "ALTER TABLE public.context_fragment DISABLE TRIGGER " + "context_fragment_reject_mutation" + ) + op.execute( + "ALTER TABLE public.context_revision DISABLE TRIGGER " + "context_revision_reject_mutation" + ) + op.execute( + "ALTER TABLE public.membership_resource_field_right DISABLE TRIGGER " + "membership_resource_field_right_mutation_lock" + ) + op.execute( + """ + CREATE TEMPORARY TABLE context_embedding_rewind + ON COMMIT DROP AS + SELECT recovery.organization_id, recovery.job_id, + recovery.resource_ref, recovery.revision_id, + recovery.previous_revision_id, recovery.publication_kind + FROM public.file_publication_recovery AS recovery + JOIN public.file_import_job AS job + ON job.organization_id = recovery.organization_id + AND job.job_id = recovery.job_id + AND job.resource_ref = recovery.resource_ref + AND job.revision_id = recovery.revision_id + WHERE recovery.checkpoint IN ('prepared', 'ready') + AND ( + job.state IN ('prepared', 'ready') + OR ( + job.state = 'leased' + AND job.recovery_from_state IN ('prepared', 'ready') + ) + ) + AND NOT EXISTS ( + SELECT 1 FROM public.context_resource AS resource + WHERE resource.organization_id = recovery.organization_id + AND resource.resource_ref = recovery.resource_ref + AND resource.active_revision_id = recovery.revision_id + ) + """ + ) + op.execute( + """ + DELETE FROM public.membership_resource_field_right AS field_right + USING context_embedding_rewind AS rewind + WHERE rewind.publication_kind = 'initial' + AND field_right.organization_id = rewind.organization_id + AND field_right.resource_ref = rewind.resource_ref + """ + ) + op.execute( + """ + DELETE FROM public.resource_access_policy AS access_policy + USING context_embedding_rewind AS rewind + WHERE rewind.publication_kind = 'initial' + AND access_policy.organization_id = rewind.organization_id + AND access_policy.resource_ref = rewind.resource_ref + """ + ) + op.execute( + """ + DELETE FROM public.exact_phrase_candidate AS candidate + USING context_embedding_rewind AS rewind + WHERE candidate.organization_id = rewind.organization_id + AND candidate.resource_ref = rewind.resource_ref + AND candidate.revision_id = rewind.revision_id + """ + ) + op.execute( + """ + DELETE FROM public.revision_publication_event AS event + USING context_embedding_rewind AS rewind + WHERE event.organization_id = rewind.organization_id + AND event.resource_ref = rewind.resource_ref + AND event.revision_id = rewind.revision_id + """ + ) + op.execute( + """ + DELETE FROM public.file_revision_replacement_plan AS plan + USING context_embedding_rewind AS rewind + WHERE plan.organization_id = rewind.organization_id + AND plan.resource_ref = rewind.resource_ref + AND plan.replacement_revision_id = rewind.revision_id + """ + ) + op.execute( + """ + DELETE FROM public.context_fragment AS fragment + USING context_embedding_rewind AS rewind + WHERE fragment.organization_id = rewind.organization_id + AND fragment.resource_ref = rewind.resource_ref + AND fragment.revision_id = rewind.revision_id + """ + ) + op.execute( + """ + DELETE FROM public.file_revision_snapshot AS snapshot + USING context_embedding_rewind AS rewind + WHERE snapshot.organization_id = rewind.organization_id + AND snapshot.resource_ref = rewind.resource_ref + AND snapshot.revision_id = rewind.revision_id + """ + ) + op.execute( + """ + DELETE FROM public.context_revision AS revision + USING context_embedding_rewind AS rewind + WHERE revision.organization_id = rewind.organization_id + AND revision.resource_ref = rewind.resource_ref + AND revision.revision_id = rewind.revision_id + """ + ) + op.execute( + """ + UPDATE public.file_publication_recovery AS recovery + SET checkpoint = 'acquired', updated_at = pg_catalog.statement_timestamp() + FROM context_embedding_rewind AS rewind + WHERE recovery.organization_id = rewind.organization_id + AND recovery.job_id = rewind.job_id + """ + ) + op.execute( + """ + UPDATE public.file_import_job AS job + SET state = CASE + WHEN job.state = 'leased' THEN 'leased' + ELSE 'running' + END, + recovery_from_state = CASE + WHEN job.state = 'leased' THEN 'running' + ELSE NULL + END, + fragment_ref = NULL + FROM context_embedding_rewind AS rewind + WHERE job.organization_id = rewind.organization_id + AND job.job_id = rewind.job_id + """ + ) + op.execute("SET CONSTRAINTS ALL IMMEDIATE") + op.execute( + "ALTER TABLE public.membership_resource_field_right ENABLE TRIGGER " + "membership_resource_field_right_mutation_lock" + ) + op.execute( + "ALTER TABLE public.context_revision ENABLE TRIGGER " + "context_revision_reject_mutation" + ) + op.execute( + "ALTER TABLE public.context_fragment ENABLE TRIGGER " + "context_fragment_reject_mutation" + ) + op.execute( + "ALTER TABLE public.file_revision_snapshot ENABLE TRIGGER " + "file_revision_snapshot_immutable" + ) + op.execute( + "ALTER TABLE public.revision_publication_event ENABLE TRIGGER " + "revision_publication_event_immutable" + ) + op.execute( + "ALTER TABLE public.exact_phrase_candidate ENABLE TRIGGER " + "exact_phrase_candidate_immutable" + ) + op.execute( + "ALTER TABLE public.file_revision_replacement_plan ENABLE TRIGGER " + "file_revision_replacement_plan_immutable" + ) + + +def _function_definition(regprocedure: str) -> str: + definition = ( + op.get_bind() + .execute( + sa.text( + "SELECT pg_catalog.pg_get_functiondef(" + f"'public.{regprocedure}'::regprocedure)" + ) + ) + .scalar_one() + ) + if not isinstance(definition, str): + raise RuntimeError("File publication function definition is unavailable") + return definition + + +def _replace_exact(definition: str, searched: str, replacement: str) -> str: + if definition.count(searched) != 1: + raise RuntimeError("File publication function shape was not recognized") + return definition.replace(searched, replacement) + + +def _install_definition(definition: str) -> None: + op.execute(f"GRANT CREATE ON SCHEMA public TO {_DEFINER}") + op.execute(f"SET LOCAL ROLE {_DEFINER}") + op.execute(definition) + op.execute("RESET ROLE") + op.execute(f"REVOKE CREATE ON SCHEMA public FROM {_DEFINER}") + + +def _replace_prepare(*, add_embeddings: bool) -> None: + old_signature, new_signature = ( + (_OLD_PREPARE_SIGNATURE, _NEW_PREPARE_SIGNATURE) + if add_embeddings + else (_NEW_PREPARE_SIGNATURE, _OLD_PREPARE_SIGNATURE) + ) + definition = _function_definition(f"{_PREPARE}{old_signature}") + replacements = ( + ( + (_SIGNATURE_OLD, _SIGNATURE_NEW), + (_ENTRY_ANCHOR, _EMBEDDING_VALIDATION), + (_INSERT_OLD, _INSERT_NEW), + ) + if add_embeddings + else ( + (_SIGNATURE_NEW, _SIGNATURE_OLD), + (_EMBEDDING_VALIDATION, _ENTRY_ANCHOR), + (_INSERT_NEW, _INSERT_OLD), + ) + ) + for searched, replacement in replacements: + definition = _replace_exact(definition, searched, replacement) + _install_definition(definition) + op.execute(f"SET LOCAL ROLE {_DEFINER}") + op.execute(f"REVOKE ALL ON FUNCTION public.{_PREPARE}{new_signature} FROM PUBLIC") + op.execute( + f"GRANT EXECUTE ON FUNCTION public.{_PREPARE}{new_signature} TO {_WORKER}" + ) + op.execute( + f"REVOKE EXECUTE ON FUNCTION public.{_PREPARE}{old_signature} FROM {_WORKER}" + ) + op.execute(f"DROP FUNCTION public.{_PREPARE}{old_signature}") + op.execute("RESET ROLE") + + +def _replace_guard( + regprocedure: str, + *, + add_embeddings: bool, + old: str, + new: str, +) -> None: + definition = _function_definition(regprocedure) + searched, replacement = (old, new) if add_embeddings else (new, old) + _install_definition(_replace_exact(definition, searched, replacement)) + + +def upgrade() -> None: + """Store one fixed-dimension vector with each newly published Fragment.""" + + op.add_column( + "context_fragment", + sa.Column("embedding", sa.Text(), nullable=True), + ) + op.execute( + "ALTER TABLE public.context_fragment " + f"ALTER COLUMN embedding TYPE public.vector({_DIMENSION}) " + "USING embedding::public.vector" + ) + op.execute( + "CREATE INDEX ix_context_fragment_embedding_hnsw " + "ON public.context_fragment USING hnsw " + "(embedding public.vector_cosine_ops) WHERE embedding IS NOT NULL" + ) + _rewind_unembedded_recovery() + _replace_prepare(add_embeddings=True) + _replace_guard( + _CLASSIFY_REGPROCEDURE, + add_embeddings=True, + old=_CLASSIFY_ANCHOR, + new=_CLASSIFY_WITH_EMBEDDINGS, + ) + _replace_guard( + _ACQUIRE_REGPROCEDURE, + add_embeddings=True, + old=_NOOP_ANCHOR, + new=_NOOP_WITH_EMBEDDINGS, + ) + _replace_guard( + _INDEX_REGPROCEDURE, + add_embeddings=True, + old=_INDEX_ANCHOR, + new=_INDEX_WITH_EMBEDDINGS, + ) + _replace_guard( + _ACTIVATE_REGPROCEDURE, + add_embeddings=True, + old=_ACTIVATE_ANCHOR, + new=_ACTIVATE_WITH_EMBEDDINGS, + ) + + +def downgrade() -> None: + """Restore vector-free publication while preserving Fragment content.""" + + op.execute("LOCK TABLE public.context_fragment IN ACCESS EXCLUSIVE MODE") + _replace_guard( + _ACTIVATE_REGPROCEDURE, + add_embeddings=False, + old=_ACTIVATE_ANCHOR, + new=_ACTIVATE_WITH_EMBEDDINGS, + ) + _replace_guard( + _INDEX_REGPROCEDURE, + add_embeddings=False, + old=_INDEX_ANCHOR, + new=_INDEX_WITH_EMBEDDINGS, + ) + _replace_guard( + _ACQUIRE_REGPROCEDURE, + add_embeddings=False, + old=_NOOP_ANCHOR, + new=_NOOP_WITH_EMBEDDINGS, + ) + _replace_guard( + _CLASSIFY_REGPROCEDURE, + add_embeddings=False, + old=_CLASSIFY_ANCHOR, + new=_CLASSIFY_WITH_EMBEDDINGS, + ) + _replace_prepare(add_embeddings=False) + op.drop_index("ix_context_fragment_embedding_hnsw", table_name="context_fragment") + op.drop_column("context_fragment", "embedding") diff --git a/tests/integration/test_content_schema.py b/tests/integration/test_content_schema.py index ffc10f17..7df5e344 100644 --- a/tests/integration/test_content_schema.py +++ b/tests/integration/test_content_schema.py @@ -771,6 +771,7 @@ def test_content_tables_have_force_rls_and_least_privilege_grants( "ordinal", "projection_kind", "content", + "embedding", }, } assert { diff --git a/tests/integration/test_file_change_pages.py b/tests/integration/test_file_change_pages.py index 57c89ddb..872b84d7 100644 --- a/tests/integration/test_file_change_pages.py +++ b/tests/integration/test_file_change_pages.py @@ -20,6 +20,7 @@ from sqlalchemy import Connection, Engine, event, text from sqlalchemy.exc import DBAPIError, SQLAlchemyError +from adapters.embeddings import DeterministicEmbeddingTwin from adapters.exact_phrase import PostgreSQLExactPhraseCandidateIndex from adapters.file_source import FileChangeProvider, FileReadLimits, FileRootRegistry from adapters.http.app import create_app @@ -622,6 +623,7 @@ def test_control_executes_a_nonterminal_current_delete_observation( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ) publications = [ @@ -1778,6 +1780,7 @@ def test_control_accepts_delete_observations_without_visibility_effect( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ) published_deletes = { @@ -3103,6 +3106,7 @@ def test_control_atomically_schedules_exact_accepted_file_upserts( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -3164,6 +3168,7 @@ def test_control_atomically_schedules_exact_accepted_file_upserts( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -3337,6 +3342,7 @@ def read_after_receiver_revocation( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -3671,6 +3677,7 @@ def track_content_read( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -3751,6 +3758,7 @@ def supersede_during_compile( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -3932,6 +3940,7 @@ def test_scheduled_redeem_waits_for_progress_before_offboard_job_fence( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ) @@ -4135,6 +4144,7 @@ def track_content_read( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ) with ThreadPoolExecutor(max_workers=1) as executor: @@ -4317,6 +4327,7 @@ def test_scheduled_pages_reuse_unchanged_and_replaced_publication_paths( receiver, worker_roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( diff --git a/tests/integration/test_file_dispatch.py b/tests/integration/test_file_dispatch.py index 96bbf0f1..01986339 100644 --- a/tests/integration/test_file_dispatch.py +++ b/tests/integration/test_file_dispatch.py @@ -20,6 +20,7 @@ from sqlalchemy.exc import DBAPIError from uvicorn import Config, Server +from adapters.embeddings import DeterministicEmbeddingTwin from adapters.exact_phrase import PostgreSQLExactPhraseCandidateIndex from adapters.file_source import FileChangeProvider, FileReadLimits, FileRootRegistry from adapters.http.app import create_app @@ -84,6 +85,10 @@ pytestmark = pytest.mark.integration ROOT = Path(__file__).parents[2] SIGNING_KEY = b"issue-91-file-dispatch-key-00001" +EMBEDDING_ENVIRONMENT = { + "CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER": "twin", + "CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION": "384", +} ALL_TEST_ROOTS = ("dispatch-root",) @@ -1319,6 +1324,7 @@ def test_scheduler_recovers_interrupted_publication_and_stales_old_lease( FileImportReceiver(first.service_principal_id), roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: first.issued_at, interrupt_after=boundary, ) @@ -1340,6 +1346,7 @@ def test_scheduler_recovers_interrupted_publication_and_stales_old_lease( ["context-engine-worker", "--dispatch-file-once"], env={ **os.environ, + **EMBEDDING_ENVIRONMENT, "CONTEXT_ENGINE_WORKER_LEASE_SIGNING_KEY_HEX": SIGNING_KEY.hex(), "CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON": json.dumps( {root_ref.value: str(root)} @@ -1363,6 +1370,7 @@ def test_scheduler_recovers_interrupted_publication_and_stales_old_lease( FileImportReceiver(first.service_principal_id), roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: first.issued_at, ) with pytest.raises(WorkNotAvailable): @@ -1544,6 +1552,7 @@ def test_autonomous_replacement_reclaim_is_all_old_then_all_new_over_sdk( receiver, roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: old_claim.issued_at, ).run(old_claim.redemption) with _authorize( @@ -1597,6 +1606,7 @@ def test_autonomous_replacement_reclaim_is_all_old_then_all_new_over_sdk( receiver, roots, MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: new_claim.issued_at, interrupt_after=FilePublicationBoundary.INDEXED, ).run(new_claim.redemption) @@ -1704,6 +1714,7 @@ def test_autonomous_replacement_reclaim_is_all_old_then_all_new_over_sdk( ["context-engine-worker", "--dispatch-file-once"], env={ **os.environ, + **EMBEDDING_ENVIRONMENT, "CONTEXT_ENGINE_WORKER_LEASE_SIGNING_KEY_HEX": SIGNING_KEY.hex(), "CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON": json.dumps( {source.source_version.root_ref.value: str(root)} @@ -2556,6 +2567,7 @@ def test_independent_worker_process_dispatches_and_publishes_one_job( ["context-engine-worker", "--dispatch-file-once"], env={ **os.environ, + **EMBEDDING_ENVIRONMENT, "CONTEXT_ENGINE_WORKER_LEASE_SIGNING_KEY_HEX": SIGNING_KEY.hex(), "CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON": json.dumps( {root_ref.value: str(root)} @@ -2578,6 +2590,7 @@ def test_independent_worker_process_dispatches_and_publishes_one_job( ["context-engine-worker", "--dispatch-file-once"], env={ **os.environ, + **EMBEDDING_ENVIRONMENT, "CONTEXT_ENGINE_WORKER_LEASE_SIGNING_KEY_HEX": SIGNING_KEY.hex(), "CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON": json.dumps( {root_ref.value: str(root)} @@ -2740,6 +2753,7 @@ def test_independent_worker_publishes_nested_markdown_fragment_lineage( ["context-engine-worker", "--dispatch-file-once"], env={ **os.environ, + **EMBEDDING_ENVIRONMENT, "CONTEXT_ENGINE_WORKER_LEASE_SIGNING_KEY_HEX": SIGNING_KEY.hex(), "CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON": json.dumps( {root_ref.value: str(root)} @@ -2802,6 +2816,7 @@ def test_long_running_dispatch_process_exits_cleanly_on_sigterm( ["context-engine-worker", "--dispatch-files"], env={ **os.environ, + **EMBEDDING_ENVIRONMENT, "CONTEXT_ENGINE_WORKER_LEASE_SIGNING_KEY_HEX": SIGNING_KEY.hex(), "CONTEXT_ENGINE_WORKER_FILE_ROOTS_JSON": json.dumps( {"empty-process-root": str(root)} diff --git a/tests/integration/test_file_import_tracer.py b/tests/integration/test_file_import_tracer.py index 841c62f5..7e4a5fa1 100644 --- a/tests/integration/test_file_import_tracer.py +++ b/tests/integration/test_file_import_tracer.py @@ -16,6 +16,7 @@ from sqlalchemy.engine import Connection from sqlalchemy.exc import SQLAlchemyError +from adapters.embeddings import DeterministicEmbeddingTwin from adapters.exact_phrase import PostgreSQLExactPhraseCandidateIndex from adapters.file_source import FileReadLimits, FileRootRegistry from adapters.http.app import create_app @@ -713,6 +714,7 @@ def test_registered_file_import_publishes_one_exact_authorized_http_package( limits=FileReadLimits(max_file_bytes=1024), ), MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -1030,6 +1032,7 @@ def _assert_structural_file_import_returns_coherent_authorized_units_over_http( limits=FileReadLimits(max_file_bytes=4096), ), MarkdownCompilerConfig("markdown-config-v2"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -1654,6 +1657,7 @@ def test_missing_file_after_redemption_records_terminal_zero_effect_failure( limits=FileReadLimits(max_file_bytes=1024), ), MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ) with pytest.raises(FileImportUnavailable): @@ -2095,6 +2099,7 @@ def test_invalid_markdown_records_terminal_failure_without_content_effects( limits=FileReadLimits(max_file_bytes=1024), ), MarkdownCompilerConfig("markdown-config-v1"), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ) diff --git a/tests/integration/test_fragment_embeddings.py b/tests/integration/test_fragment_embeddings.py new file mode 100644 index 00000000..f1d3263b --- /dev/null +++ b/tests/integration/test_fragment_embeddings.py @@ -0,0 +1,458 @@ +from __future__ import annotations + +import json +from datetime import UTC, datetime +from pathlib import Path +from typing import Any, cast + +import pytest +from sqlalchemy import Engine, text + +from adapters.embeddings import DeterministicEmbeddingTwin +from adapters.file_source import FileReadLimits, FileRootRegistry +from engine.persistence import ( + DatabaseConfiguration, + FileImportInterrupted, + FileImportLeaseRedemption, + FilePublicationBoundary, + PostgreSQLFileImportWorker, + PostgreSQLWorkerLeaseIssuer, + create_database_engine, +) +from engine.supply import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + EmbeddingProfile, + EmbeddingProviderUnavailable, + EmbeddingVector, + MarkdownCompilerConfig, + WorkNotAvailable, +) +from tests.support.file_imports import ( + FileImportScenario, + delete_file_import_scenario, + prepare_file_import_scenario, + prepare_repeat_file_import, +) + +pytestmark = pytest.mark.integration + + +class _RecordingEmbeddingProvider: + def __init__(self, *, available: bool = True) -> None: + self.profile = EmbeddingProfile(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION) + self.available = available + self.calls: list[tuple[str, ...]] = [] + self._twin = DeterministicEmbeddingTwin() + + def embed(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: + self.calls.append(inputs) + if not self.available: + raise EmbeddingProviderUnavailable("provider detail must not escape") + return self._twin.embed(inputs) + + +class _InvalidEmbeddingProvider(_RecordingEmbeddingProvider): + def embed(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: + self.calls.append(inputs) + return ((0.25,),) * len(inputs) + + +class _MutableProfileEmbeddingProvider(_RecordingEmbeddingProvider): + def embed(self, inputs: tuple[str, ...]) -> tuple[EmbeddingVector, ...]: + self.calls.append(inputs) + self.profile = EmbeddingProfile(1) + return ((0.25,),) * len(inputs) + + +def _worker( + scenario: FileImportScenario, + guarded_worker_engine: Engine, + provider: _RecordingEmbeddingProvider, +) -> PostgreSQLFileImportWorker: + return PostgreSQLFileImportWorker( + guarded_worker_engine, + scenario.codec, + scenario.receiver, + FileRootRegistry( + {scenario.root_ref: scenario.root}, + limits=FileReadLimits(max_file_bytes=4096), + ), + MarkdownCompilerConfig("markdown-config-v2"), + embedding_provider=provider, + clock=lambda: datetime.now(UTC).replace(microsecond=0), + ) + + +def _run( + worker: PostgreSQLFileImportWorker, + scenario: FileImportScenario, + token: Any, +) -> object: + return worker.run( + FileImportLeaseRedemption( + token, + scenario.organization_id, + scenario.prepared.job_id, + scenario.source_ref, + ) + ) + + +def _stored_vectors( + configuration: DatabaseConfiguration, + scenario: FileImportScenario, +) -> tuple[tuple[str, int], ...]: + engine = create_database_engine(configuration) + try: + with engine.connect() as connection: + rows = connection.execute( + text( + """ + SELECT embedding::text, vector_dims(embedding) + FROM context_fragment + WHERE organization_id = :organization_id + ORDER BY ordinal + """ + ), + {"organization_id": scenario.organization_id}, + ).all() + return tuple((str(row[0]), int(row[1])) for row in rows) + finally: + engine.dispose() + + +def test_real_publication_persists_deterministic_vectors_and_noop_skips_provider( + request: pytest.FixtureRequest, + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + scenario = prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + payload=b"# Handbook\n\nFirst fragment.\n\nSecond fragment.\n", + ) + request.addfinalizer( + lambda: delete_file_import_scenario( + migration_configuration, scenario.organization_id + ) + ) + assert scenario.token is not None + provider = _RecordingEmbeddingProvider() + published = cast( + Any, + _run( + _worker(scenario, guarded_worker_engine, provider), + scenario, + scenario.token, + ), + ) + + vectors = _stored_vectors(migration_configuration, scenario) + assert len(vectors) == len(published.candidate_refs) + assert all( + dimension == CONTEXT_FRAGMENT_EMBEDDING_DIMENSION + for _vector, dimension in vectors + ) + assert len(provider.calls) == 1 + engine = create_database_engine(migration_configuration) + try: + with engine.connect() as connection: + storage = connection.execute( + text( + """ + SELECT format_type(attribute.atttypid, attribute.atttypmod), + table_class.relrowsecurity, + table_class.relforcerowsecurity, + index_definition.indexdef + FROM pg_attribute AS attribute + JOIN pg_class AS table_class + ON table_class.oid = attribute.attrelid + JOIN pg_namespace AS namespace + ON namespace.oid = table_class.relnamespace + JOIN pg_indexes AS index_definition + ON index_definition.schemaname = namespace.nspname + AND index_definition.tablename = table_class.relname + AND index_definition.indexname = + 'ix_context_fragment_embedding_hnsw' + WHERE namespace.nspname = 'public' + AND table_class.relname = 'context_fragment' + AND attribute.attname = 'embedding' + """ + ) + ).one() + assert tuple(storage[:3]) == ("vector(384)", True, True) + assert "USING hnsw" in storage.indexdef + assert "vector_cosine_ops" in storage.indexdef + assert "WHERE (embedding IS NOT NULL)" in storage.indexdef + finally: + engine.dispose() + + prepared, replay_token = prepare_repeat_file_import( + scenario, + guarded_control_engine, + idempotency_key="embedding-noop-replay", + ) + replay = PostgreSQLFileImportWorker( + guarded_worker_engine, + scenario.codec, + scenario.receiver, + FileRootRegistry( + {scenario.root_ref: scenario.root}, + limits=FileReadLimits(max_file_bytes=4096), + ), + MarkdownCompilerConfig("markdown-config-v2"), + embedding_provider=provider, + clock=lambda: datetime.now(UTC).replace(microsecond=0), + ).run( + FileImportLeaseRedemption( + replay_token, + prepared.organization_id, + prepared.job_id, + prepared.source_ref, + ) + ) + + assert replay.outcome == "unchanged" + assert len(provider.calls) == 1 + assert _stored_vectors(migration_configuration, scenario) == vectors + + +def test_provider_failure_interrupts_acquired_checkpoint_and_recovers( + request: pytest.FixtureRequest, + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + scenario = prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + lease_ttl_seconds=2, + ) + request.addfinalizer( + lambda: delete_file_import_scenario( + migration_configuration, scenario.organization_id + ) + ) + assert scenario.token is not None + provider = _RecordingEmbeddingProvider(available=False) + + with pytest.raises(FileImportInterrupted) as interrupted: + _run( + _worker(scenario, guarded_worker_engine, provider), + scenario, + scenario.token, + ) + + assert interrupted.value.boundary is FilePublicationBoundary.ACQUIRED + engine = create_database_engine(migration_configuration) + try: + with engine.connect() as connection: + before = connection.execute( + text( + """ + SELECT job.state, job.revision_id, + resource.active_revision_id, + (SELECT count(*) FROM context_fragment AS fragment + WHERE fragment.organization_id = job.organization_id) + FROM file_import_job AS job + LEFT JOIN context_resource AS resource + ON resource.organization_id = job.organization_id + AND resource.resource_ref = job.resource_ref + WHERE job.organization_id = :organization_id + AND job.job_id = :job_id + """ + ), + { + "organization_id": scenario.organization_id, + "job_id": scenario.prepared.job_id, + }, + ).one() + assert tuple(before[:1]) == ("running",) + assert before[1] is not None + assert before[2] is None + assert before[3] == 0 + with engine.connect() as connection: + connection.execute(text("SELECT pg_sleep(2.1)")) + finally: + engine.dispose() + + provider.available = True + recovery_token = PostgreSQLWorkerLeaseIssuer( + guarded_control_engine, + scenario.codec, + ).issue_file_import_lease(scenario.prepared) + recovered = cast( + Any, + _run( + _worker(scenario, guarded_worker_engine, provider), + scenario, + recovery_token, + ), + ) + + assert recovered.outcome == "published" + assert len(provider.calls) == 2 + assert _stored_vectors(migration_configuration, scenario) + + +def test_worker_refuses_embedding_dimension_mismatch_at_composition( + request: pytest.FixtureRequest, + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + scenario = prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + ) + request.addfinalizer( + lambda: delete_file_import_scenario( + migration_configuration, scenario.organization_id + ) + ) + provider = _RecordingEmbeddingProvider() + provider.profile = EmbeddingProfile(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION - 1) + + with pytest.raises(ValueError, match="dimension does not match"): + _worker(scenario, guarded_worker_engine, provider) + + +def test_invalid_provider_response_interrupts_before_fragment_persistence( + request: pytest.FixtureRequest, + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + scenario = prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + ) + request.addfinalizer( + lambda: delete_file_import_scenario( + migration_configuration, scenario.organization_id + ) + ) + assert scenario.token is not None + provider = _InvalidEmbeddingProvider() + + with pytest.raises(FileImportInterrupted) as interrupted: + _run( + _worker(scenario, guarded_worker_engine, provider), + scenario, + scenario.token, + ) + + assert interrupted.value.boundary is FilePublicationBoundary.ACQUIRED + engine = create_database_engine(migration_configuration) + try: + with engine.connect() as connection: + assert connection.execute( + text( + "SELECT count(*) FROM context_fragment " + "WHERE organization_id = :organization_id" + ), + {"organization_id": scenario.organization_id}, + ).scalar_one() == 0 + finally: + engine.dispose() + + +def test_provider_cannot_change_the_composed_dimension_during_publication( + request: pytest.FixtureRequest, + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + scenario = prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + ) + request.addfinalizer( + lambda: delete_file_import_scenario( + migration_configuration, scenario.organization_id + ) + ) + assert scenario.token is not None + + with pytest.raises(FileImportInterrupted) as interrupted: + _run( + _worker( + scenario, + guarded_worker_engine, + _MutableProfileEmbeddingProvider(), + ), + scenario, + scenario.token, + ) + + assert interrupted.value.boundary is FilePublicationBoundary.ACQUIRED + assert _stored_vectors(migration_configuration, scenario) == () + + +def test_postgresql_refuses_a_vector_that_underflows_to_float32_zero( + request: pytest.FixtureRequest, + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + scenario = prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + ) + request.addfinalizer( + lambda: delete_file_import_scenario( + migration_configuration, scenario.organization_id + ) + ) + assert scenario.token is not None + + def underflow_document( + _worker: PostgreSQLFileImportWorker, + _token: object, + _claims: object, + document: object, + ) -> str: + fragments = cast(Any, document).fragments + return json.dumps( + [ + { + "embedding": [1.0e-50] + * CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + "fragmentRef": fragment.fragment_ref, + } + for fragment in fragments + ] + ) + + monkeypatch.setattr( + PostgreSQLFileImportWorker, + "_embedding_document", + underflow_document, + ) + + with pytest.raises(WorkNotAvailable): + _run( + _worker( + scenario, + guarded_worker_engine, + _RecordingEmbeddingProvider(), + ), + scenario, + scenario.token, + ) + + assert _stored_vectors(migration_configuration, scenario) == () diff --git a/tests/integration/test_migrations.py b/tests/integration/test_migrations.py index 9c7b35b1..2ad6864e 100644 --- a/tests/integration/test_migrations.py +++ b/tests/integration/test_migrations.py @@ -2304,6 +2304,86 @@ def test_recursive_file_path_revision_downgrades_and_reapplies_when_empty( assert _revision_rows(migration_configuration) == [HEAD_REVISION] +def test_fragment_embedding_revision_downgrades_and_reapplies_when_empty( + migration_configuration: DatabaseConfiguration, +) -> None: + alembic_configuration = Config(ROOT / "alembic.ini") + + try: + command.downgrade(alembic_configuration, "20260726_0035") + assert _revision_rows(migration_configuration) == ["20260726_0035"] + engine = create_database_engine(migration_configuration) + try: + with engine.connect() as connection: + assert ( + connection.execute( + text( + "SELECT to_regclass(" + "'public.ix_context_fragment_embedding_hnsw')" + ) + ).scalar_one() + is None + ) + finally: + engine.dispose() + finally: + command.upgrade(alembic_configuration, "head") + + assert _revision_rows(migration_configuration) == [HEAD_REVISION] + + +def test_fragment_embedding_revision_preserves_retained_fragments( + tmp_path: Path, + migration_configuration: DatabaseConfiguration, + guarded_control_engine: Engine, + guarded_worker_engine: Engine, +) -> None: + alembic_configuration = Config(ROOT / "alembic.ini") + scenario = _prepare_file_import_scenario( + tmp_path, + migration_configuration, + guarded_control_engine, + ) + assert scenario.token is not None + _run_file_import( + scenario, + scenario.prepared, + scenario.token, + guarded_worker_engine, + ) + try: + command.downgrade(alembic_configuration, "20260726_0035") + assert _revision_rows(migration_configuration) == ["20260726_0035"] + engine = create_database_engine(migration_configuration) + try: + with engine.connect() as connection: + retained = connection.execute( + text( + "SELECT count(*) FROM context_fragment " + "WHERE organization_id = :organization_id" + ), + {"organization_id": scenario.organization_id}, + ).scalar_one() + embedding_column = connection.execute( + text( + "SELECT count(*) FROM information_schema.columns " + "WHERE table_schema = 'public' " + "AND table_name = 'context_fragment' " + "AND column_name = 'embedding'" + ) + ).scalar_one() + assert retained > 0 + assert embedding_column == 0 + finally: + engine.dispose() + finally: + command.upgrade(alembic_configuration, "head") + _delete_issue_27_upgrade_fixture( + migration_configuration, + scenario.organization_id, + ) + + def test_recursive_file_path_revision_refuses_retained_nested_lineage( tmp_path: Path, migration_configuration: DatabaseConfiguration, @@ -2864,7 +2944,7 @@ def test_recovery_upgrade_adopts_an_existing_ready_replacement( guarded_control_engine: Engine, guarded_worker_engine: Engine, ) -> None: - """An Issue #26 ready job remains resumable after the Issue #27 upgrade.""" + """A pre-embedding ready Revision rewinds and resumes through embedding.""" scenario = _prepare_file_import_scenario( tmp_path, @@ -2977,7 +3057,7 @@ def test_recovery_upgrade_adopts_an_existing_ready_replacement( "job_id": claims.job_id, }, ).scalar_one() - == "ready" + == "acquired" ) finally: migration_engine.dispose() diff --git a/tests/integration/test_z_egress_grant_file.py b/tests/integration/test_z_egress_grant_file.py index 46f58ea1..0df1d92c 100644 --- a/tests/integration/test_z_egress_grant_file.py +++ b/tests/integration/test_z_egress_grant_file.py @@ -699,6 +699,8 @@ def _published_file_scenario( worker_signing_key = bytes(range(32)).hex() worker_environment = { **os.environ, + "CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER": "twin", + "CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION": "384", "CONTEXT_ENGINE_WORKER_DATABASE_URL": worker_database_url, "CONTEXT_ENGINE_WORKER_FILE_ROOT_PATH": str(scenario.root), "CONTEXT_ENGINE_WORKER_FILE_ROOT_REF": scenario.root_ref.value, diff --git a/tests/integration/test_zz_file_publication_recovery.py b/tests/integration/test_zz_file_publication_recovery.py index 76fe9982..4851122d 100644 --- a/tests/integration/test_zz_file_publication_recovery.py +++ b/tests/integration/test_zz_file_publication_recovery.py @@ -11,6 +11,7 @@ from sqlalchemy import Engine, text from sqlalchemy.exc import SQLAlchemyError +from adapters.embeddings import DeterministicEmbeddingTwin from adapters.file_source import FileReadLimits, FileRootRegistry from adapters.parsers.markdown import compile_markdown from engine.control import FileRootRef @@ -76,6 +77,7 @@ def _worker( limits=FileReadLimits(max_file_bytes=4096), ), MarkdownCompilerConfig(config_version), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), interrupt_after=interrupt_after, ) diff --git a/tests/support/file_imports.py b/tests/support/file_imports.py index 5ea721cd..2de5f9fe 100644 --- a/tests/support/file_imports.py +++ b/tests/support/file_imports.py @@ -15,6 +15,7 @@ from sqlalchemy import Engine, text +from adapters.embeddings import DeterministicEmbeddingTwin from adapters.file_source import FileReadLimits, FileRootRegistry from engine.control import ( ContextControl, @@ -337,6 +338,7 @@ def run_file_import( limits=FileReadLimits(max_file_bytes=4096), ), MarkdownCompilerConfig(config_version), + embedding_provider=DeterministicEmbeddingTwin(), clock=lambda: datetime.now(UTC).replace(microsecond=0), ).run( FileImportLeaseRedemption( @@ -348,6 +350,105 @@ def run_file_import( ) +def delete_file_import_scenario( + configuration: DatabaseConfiguration, + organization_id: UUID, +) -> None: + """Remove one disposable File scenario without touching sibling tenants.""" + + engine = create_database_engine(configuration) + immutable_tables = ( + ("file_source_publish_watermark", "file_source_publish_watermark_immutable"), + ( + "file_source_acquisition_checkpoint", + "file_source_acquisition_checkpoint_immutable", + ), + ("file_import_job_event", "file_import_job_event_immutable"), + ("file_revision_replacement_plan", "file_revision_replacement_plan_immutable"), + ("exact_phrase_candidate", "exact_phrase_candidate_immutable"), + ("revision_publication_event", "revision_publication_event_immutable"), + ("context_fragment", "context_fragment_reject_mutation"), + ("file_revision_snapshot", "file_revision_snapshot_immutable"), + ("context_revision", "context_revision_reject_mutation"), + ("file_acquisition_result", "file_acquisition_result_immutable"), + ("file_resource_ingestion_guard", "file_resource_ingestion_guard_immutable"), + ("file_acquisition", "file_acquisition_immutable"), + ("source_version", "source_version_immutable"), + ) + try: + with engine.begin() as connection: + for table, trigger in immutable_tables: + connection.execute( + text(f"ALTER TABLE {table} DISABLE TRIGGER {trigger}") + ) + try: + with engine.begin() as connection: + user_ids = tuple( + connection.execute( + text( + "SELECT user_id FROM membership " + "WHERE organization_id = :organization_id" + ), + {"organization_id": organization_id}, + ).scalars() + ) + for table in ( + "file_source_publish_watermark", + "file_source_acquisition_checkpoint", + "file_import_job_event", + "file_publication_recovery", + "file_revision_replacement_plan", + "file_acquisition_result", + "exact_phrase_candidate", + "revision_publication_event", + "membership_resource_field_right", + "resource_access_policy", + "context_fragment", + "file_revision_snapshot", + "context_revision", + "context_resource", + "file_resource_ingestion_guard", + "file_import_job", + "file_acquisition", + "context_source", + "source_version", + "service_principal", + "membership", + ): + connection.execute( + text( + f"DELETE FROM {table} " # noqa: S608 + "WHERE organization_id = :organization_id" + ), + {"organization_id": organization_id}, + ) + for user_id in user_ids: + connection.execute( + text( + "DELETE FROM user_account " + "WHERE user_id = :user_id AND NOT EXISTS (" + "SELECT 1 FROM membership " + "WHERE membership.user_id = user_account.user_id)" + ), + {"user_id": user_id}, + ) + connection.execute( + text( + "DELETE FROM organization " + "WHERE organization_id = :organization_id" + ), + {"organization_id": organization_id}, + ) + finally: + with engine.begin() as connection: + for table, trigger in reversed(immutable_tables): + connection.execute( + text(f"ALTER TABLE {table} ENABLE TRIGGER {trigger}") + ) + finally: + engine.dispose() + + def redeem_file_import_direct( guarded_worker_engine: Engine, claims: WorkerLeaseClaims, diff --git a/tests/support/migrations.py b/tests/support/migrations.py index c3f22264..003f743c 100644 --- a/tests/support/migrations.py +++ b/tests/support/migrations.py @@ -1,3 +1,3 @@ """Shared migration assertions for tests that require the current schema head.""" -HEAD_REVISION = "20260726_0035" +HEAD_REVISION = "20260726_0036" diff --git a/tests/unit/test_embeddings.py b/tests/unit/test_embeddings.py new file mode 100644 index 00000000..16cc8286 --- /dev/null +++ b/tests/unit/test_embeddings.py @@ -0,0 +1,287 @@ +from __future__ import annotations + +import json +from email.message import Message +from importlib import import_module +from io import BytesIO +from typing import Any, cast +from urllib.error import HTTPError +from urllib.request import Request + +import pytest + +from adapters.embeddings import ( + DeterministicEmbeddingTwin, + ExternalEmbeddingConfiguration, + ExternalEmbeddingProvider, + _RejectRedirectHandler, +) +from engine.persistence.file_imports import _EMBEDDING_PREPARE_REGPROCEDURE +from engine.supply import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + EmbeddingProfile, + EmbeddingProviderUnavailable, + validate_embedding_batch, +) + + +def test_deterministic_twin_is_content_derived_and_fixed_dimension() -> None: + provider = DeterministicEmbeddingTwin() + + first = provider.embed(("same fragment", "different fragment")) + replay = provider.embed(("same fragment", "different fragment")) + + assert first == replay + assert first[0] != first[1] + assert all(len(vector) == CONTEXT_FRAGMENT_EMBEDDING_DIMENSION for vector in first) + assert all(any(value != 0.0 for value in vector) for vector in first) + + +def test_external_provider_binds_model_dimension_and_input_without_leaking_key() -> ( + None +): + observed: list[tuple[Request, float, int]] = [] + + def transport(request: Request, timeout: float, maximum_bytes: int) -> bytes: + observed.append((request, timeout, maximum_bytes)) + inputs = json.loads(cast(bytes, request.data or b"{}"))["input"] + return json.dumps( + { + "data": [ + { + "index": index, + "embedding": [0.25] * CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + } + for index, _value in enumerate(inputs) + ] + } + ).encode("utf-8") + + configuration = ExternalEmbeddingConfiguration( + endpoint="https://embedding.invalid/v1/embeddings", + model="configured-model", + api_key="credential-value", + dimension=CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + batch_size=1, + ) + provider = ExternalEmbeddingProvider(configuration, transport=transport) + + vectors = provider.embed(("first", "second")) + + assert len(vectors) == 2 + assert len(observed) == 2 + request, timeout, maximum_bytes = observed[0] + payload = json.loads(cast(bytes, request.data or b"{}")) + assert payload == { + "dimensions": CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + "encoding_format": "float", + "input": ["first"], + "model": "configured-model", + } + assert request.get_header("Authorization") == "Bearer credential-value" + assert timeout == configuration.timeout_seconds + assert maximum_bytes > 0 + assert "credential-value" not in repr(configuration) + assert "embedding.invalid" not in repr(configuration) + assert "credential-value" not in repr(provider) + assert "embedding.invalid" not in repr(provider) + + +def test_external_provider_replaces_transport_details_with_generic_failure() -> None: + def transport(_request: Request, _timeout: float, _maximum_bytes: int) -> bytes: + raise OSError("credential-value and response content") + + provider = ExternalEmbeddingProvider( + ExternalEmbeddingConfiguration( + endpoint="https://embedding.invalid/v1/embeddings", + model="configured-model", + api_key="credential-value", + dimension=CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + batch_size=64, + ), + transport=transport, + ) + + with pytest.raises( + EmbeddingProviderUnavailable, + match="Embedding provider is unavailable", + ) as failure: + provider.embed(("content",)) + + assert "credential-value" not in str(failure.value) + assert failure.value.__cause__ is None + + +def test_external_transport_rejects_redirect_before_reusing_request_headers() -> None: + request = Request( + "https://embedding.invalid/v1/embeddings", + headers={"Authorization": "Bearer credential-value"}, + ) + + with pytest.raises(HTTPError) as rejected: + _RejectRedirectHandler().redirect_request( + request, + BytesIO(), + 307, + "redirect", + Message(), + "https://redirect.invalid/collect", + ) + + assert rejected.value.url == request.full_url + assert "redirect.invalid" not in str(rejected.value) + + +@pytest.mark.parametrize( + ("endpoint", "model", "api_key"), + [ + ("http://embedding.invalid/v1/embeddings", "model", "key"), + ("https://embedding.invalid/v1/embeddings?token=value", "model", "key"), + ("https://embedding.invalid/v1/embeddings", " model", "key"), + ("https://embedding.invalid/v1/embeddings", "model", " key"), + ], +) +def test_external_configuration_refuses_unsafe_or_ambiguous_values( + endpoint: str, + model: str, + api_key: str, +) -> None: + with pytest.raises(ValueError, match="configuration is not available"): + ExternalEmbeddingConfiguration( + endpoint=endpoint, + model=model, + api_key=api_key, + dimension=CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + batch_size=64, + ) + + +def test_external_provider_preserves_order_across_bounded_batches() -> None: + observed_inputs: list[list[str]] = [] + + def transport(request: Request, _timeout: float, _maximum_bytes: int) -> bytes: + inputs = json.loads(cast(bytes, request.data or b"{}"))["input"] + observed_inputs.append(inputs) + return json.dumps( + { + "data": [ + { + "index": index, + "embedding": [float(value)] + * CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + } + for index, value in reversed(tuple(enumerate(inputs, start=0))) + ] + } + ).encode("utf-8") + + provider = ExternalEmbeddingProvider( + ExternalEmbeddingConfiguration( + endpoint="https://embedding.invalid/v1/embeddings", + model="configured-model", + api_key="credential-value", + dimension=CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + batch_size=2, + ), + transport=transport, + ) + + vectors = provider.embed(("1", "2", "3", "4", "5")) + + assert observed_inputs == [["1", "2"], ["3", "4"], ["5"]] + assert tuple(vector[0] for vector in vectors) == (1.0, 2.0, 3.0, 4.0, 5.0) + + +def test_external_provider_collapses_later_batch_failure() -> None: + call_count = 0 + + def transport(request: Request, _timeout: float, _maximum_bytes: int) -> bytes: + nonlocal call_count + call_count += 1 + if call_count == 2: + raise OSError("response details") + inputs = json.loads(cast(bytes, request.data or b"{}"))["input"] + return json.dumps( + { + "data": [ + { + "index": index, + "embedding": [0.25] + * CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + } + for index, _value in enumerate(inputs) + ] + } + ).encode("utf-8") + + provider = ExternalEmbeddingProvider( + ExternalEmbeddingConfiguration( + endpoint="https://embedding.invalid/v1/embeddings", + model="configured-model", + api_key="credential-value", + dimension=CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + batch_size=1, + ), + transport=transport, + ) + + with pytest.raises(EmbeddingProviderUnavailable) as failure: + provider.embed(("first", "second")) + + assert call_count == 2 + assert str(failure.value) == "Embedding provider is unavailable" + + +@pytest.mark.parametrize("batch_size", [0, 257, True]) +def test_external_configuration_refuses_unbounded_batch_size( + batch_size: Any, +) -> None: + with pytest.raises(ValueError, match="configuration is not available"): + ExternalEmbeddingConfiguration( + endpoint="https://embedding.invalid/v1/embeddings", + model="configured-model", + api_key="credential-value", + dimension=CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, + batch_size=batch_size, + ) + + +@pytest.mark.parametrize( + "value", + [float("nan"), float("inf"), 1.0e31, 1.0e30, 1.0e-50, 0.0], +) +def test_embedding_validation_refuses_unstorable_or_zero_vectors( + value: float, +) -> None: + with pytest.raises(EmbeddingProviderUnavailable): + validate_embedding_batch( + ("content",), + ((value,),), + EmbeddingProfile(1), + ) + + +@pytest.mark.parametrize( + "vectors", + [ + cast(Any, (None,)), + cast(Any, ((10**10_000,),)), + ], +) +def test_embedding_validation_normalizes_malformed_provider_containers( + vectors: Any, +) -> None: + with pytest.raises(EmbeddingProviderUnavailable): + validate_embedding_batch(("content",), vectors, EmbeddingProfile(1)) + + +def test_worker_schema_probe_matches_the_embedding_migration_signature() -> None: + migration = import_module( + "migrations.versions.20260726_0036_fragment_embeddings" + ) + + assert ( + "public.context_worker_prepare_file_publication" + f"{migration._NEW_PREPARE_SIGNATURE}" + ) == _EMBEDDING_PREPARE_REGPROCEDURE + assert migration._DIMENSION == CONTEXT_FRAGMENT_EMBEDDING_DIMENSION diff --git a/tests/unit/test_file_dispatch.py b/tests/unit/test_file_dispatch.py index ecbfbb61..d95581e7 100644 --- a/tests/unit/test_file_dispatch.py +++ b/tests/unit/test_file_dispatch.py @@ -12,9 +12,11 @@ from sqlalchemy import Engine from sqlalchemy.exc import SQLAlchemyError +from adapters.embeddings import DeterministicEmbeddingTwin, ExternalEmbeddingProvider from applications.worker import ( DEFAULT_WORKER_MAX_FILE_BYTES, FileDispatchCycleResult, + _embedding_provider, _file_dispatch_roots, _file_read_limits, _worker_database_time, @@ -31,6 +33,7 @@ _database_timestamp_utc, ) from engine.supply import ( + CONTEXT_FRAGMENT_EMBEDDING_DIMENSION, WorkerLeaseCodec, WorkerLeaseKeyring, WorkerLeaseRejectionAuditReceipt, @@ -285,6 +288,95 @@ def test_worker_file_byte_limit_defaults_to_one_mib_and_is_configurable( assert _file_read_limits().max_file_bytes == 8192 +def test_worker_composes_only_explicit_fixed_dimension_embedding_twin( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER", "twin") + monkeypatch.setenv( + "CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION", + str(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION), + ) + + provider = _embedding_provider() + + assert type(provider) is DeterministicEmbeddingTwin + assert provider.profile.dimension == CONTEXT_FRAGMENT_EMBEDDING_DIMENSION + + +def test_worker_external_embedding_configuration_keeps_key_out_of_repr( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER", "external") + monkeypatch.setenv( + "CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION", + str(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION), + ) + monkeypatch.setenv( + "CONTEXT_ENGINE_WORKER_EMBEDDING_ENDPOINT", + "https://embedding.invalid/v1/embeddings", + ) + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_MODEL", "configured-model") + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_API_KEY", "credential-value") + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_BATCH_SIZE", "64") + + provider = _embedding_provider() + + assert type(provider) is ExternalEmbeddingProvider + assert "credential-value" not in repr(provider) + + +@pytest.mark.parametrize("batch_size", ["", "0", "257", "not-a-number"]) +def test_worker_refuses_missing_or_unbounded_external_embedding_batch_size( + monkeypatch: pytest.MonkeyPatch, + batch_size: str, +) -> None: + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER", "external") + monkeypatch.setenv( + "CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION", + str(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION), + ) + monkeypatch.setenv( + "CONTEXT_ENGINE_WORKER_EMBEDDING_ENDPOINT", + "https://embedding.invalid/v1/embeddings", + ) + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_MODEL", "configured-model") + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_API_KEY", "credential-value") + if batch_size: + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_BATCH_SIZE", batch_size) + else: + monkeypatch.delenv( + "CONTEXT_ENGINE_WORKER_EMBEDDING_BATCH_SIZE", + raising=False, + ) + + with pytest.raises(ValueError, match="configuration is not available"): + _embedding_provider() + + +@pytest.mark.parametrize( + ("mode", "dimension"), + [ + ("", str(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION)), + ("automatic", str(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION)), + ("twin", str(CONTEXT_FRAGMENT_EMBEDDING_DIMENSION - 1)), + ("twin", "not-a-number"), + ], +) +def test_worker_refuses_missing_unknown_or_mismatched_embedding_configuration( + monkeypatch: pytest.MonkeyPatch, + mode: str, + dimension: str, +) -> None: + if mode: + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER", mode) + else: + monkeypatch.delenv("CONTEXT_ENGINE_WORKER_EMBEDDING_PROVIDER", raising=False) + monkeypatch.setenv("CONTEXT_ENGINE_WORKER_EMBEDDING_DIMENSION", dimension) + + with pytest.raises(ValueError, match="configuration is not available"): + _embedding_provider() + + def test_worker_default_file_limit_accepts_above_legacy_ceiling_and_refuses_oversize( monkeypatch: pytest.MonkeyPatch, tmp_path: Path,