From 5b7835cb7e89f7e31f8c739d349e46eecf6de3ad Mon Sep 17 00:00:00 2001 From: salmonumbrella <182032677+salmonumbrella@users.noreply.github.com> Date: Mon, 24 Aug 2026 08:30:31 -0500 Subject: [PATCH 1/4] feat(documents): add semantic attachment search Add optional lexical, semantic, and hybrid search over indexed attachment text, with generation-based activation so incomplete vector builds never replace the active corpus. Use Docbank's released normalized-document and embedding contracts as the shared source of truth. Keep document and query egress consent separate, reject redirects, and page vector candidates until scoped retrieval is complete and accurately bounded. This also: - shares the document-vector CLI consent command and routing - stabilizes generation, endpoint, and provider identity - preserves provenance through bounded hybrid fusion - updates the Go and Nix dependency closures for Docbank's Go 1.27 requirement - makes vector builds converge safely - rejects stale normalized identity and invalid embedding policy - makes malformed provider responses terminal - makes semantic paging exhaustive and stable - decouples retired cleanup from consent - exposes vector status and coherent hybrid evidence - keeps restore and cleanup recoverable - preserves lexical search depth and store-native pagination - refreshes the generated document-search client Generated with Codex Co-authored-by: Wes McKinney Co-authored-by: Codex --- .github/workflows/ci.yml | 3 + Makefile | 2 +- api/openapi.yaml | 356 ++++ cmd/msgvault/cmd/add_synctech_sms_drive.go | 2 +- .../cmd/add_synctech_sms_drive_test.go | 4 +- cmd/msgvault/cmd/backup.go | 20 + cmd/msgvault/cmd/backup_test.go | 47 + cmd/msgvault/cmd/constants.go | 5 + cmd/msgvault/cmd/documents.go | 13 + cmd/msgvault/cmd/documents_routing_test.go | 45 + cmd/msgvault/cmd/documents_test.go | 12 + cmd/msgvault/cmd/documents_vector.go | 460 +++++ cmd/msgvault/cmd/documents_vector_runtime.go | 183 ++ cmd/msgvault/cmd/documents_vector_stub.go | 19 + .../cmd/documents_vector_stub_test.go | 30 + cmd/msgvault/cmd/documents_vector_test.go | 572 ++++++ cmd/msgvault/cmd/embed_vector_test.go | 62 +- cmd/msgvault/cmd/import_imessage.go | 4 +- cmd/msgvault/cmd/import_imessage_test.go | 10 +- cmd/msgvault/cmd/person_provider.go | 2 +- cmd/msgvault/cmd/root.go | 2 +- cmd/msgvault/cmd/serve.go | 19 + cmd/msgvault/cmd/serve_vector.go | 134 +- cmd/msgvault/cmd/serve_vector_init.go | 30 + cmd/msgvault/cmd/serve_vector_init_test.go | 143 ++ cmd/msgvault/cmd/store_adapter_test.go | 19 + cmd/msgvault/cmd/tui.go | 2 +- cmd/msgvault/cmd/vector_features.go | 16 +- go.mod | 2 +- go.sum | 4 +- internal/api/cli_allowlist_documents_test.go | 12 + internal/api/cli_handlers.go | 11 + internal/api/document_search.go | 143 +- internal/api/document_search_test.go | 50 +- internal/api/routes.go | 8 + internal/api/server.go | 3 + internal/backupapp/restore.go | 70 + internal/config/documentindex_config_test.go | 36 + internal/daemonclient/documents.go | 20 +- internal/daemonclient/documents_test.go | 42 +- internal/documentindex/config.go | 4 +- internal/documentindex/config_test.go | 51 +- internal/documentindex/worker.go | 3 + internal/mcp/catalog.go | 13 +- internal/mcp/catalog_test.go | 2 +- internal/mcp/handlers.go | 29 + internal/mcp/server_test.go | 13 + internal/scheduler/scheduler.go | 75 +- internal/scheduler/scheduler_test.go | 64 + internal/store/dialect_pg.go | 5 + internal/store/dialect_sqlite.go | 5 + internal/store/document_normalized.go | 162 ++ internal/store/document_publication.go | 51 +- internal/store/document_publication_test.go | 93 +- internal/store/document_search.go | 313 +++- internal/store/document_search_test.go | 245 ++- internal/store/document_vector_lock.go | 50 + internal/store/document_vector_operations.go | 449 +++++ internal/store/document_vectors.go | 1595 +++++++++++++++++ internal/store/document_vectors_pg_test.go | 22 + internal/store/document_vectors_test.go | 1524 ++++++++++++++++ internal/store/schema.sql | 88 + internal/store/schema_pg.sql | 80 + internal/store/sqlite_error_test.go | 3 +- internal/vector/backend.go | 6 + internal/vector/config.go | 61 +- internal/vector/document/backend.go | 97 + internal/vector/document/fusion.go | 128 ++ internal/vector/document/fusion_test.go | 127 ++ internal/vector/document/model.go | 165 ++ internal/vector/document/model_test.go | 121 ++ internal/vector/document/provider.go | 13 + internal/vector/document/reconcile.go | 252 +++ internal/vector/document/reconcile_test.go | 387 ++++ internal/vector/document/search.go | 512 ++++++ internal/vector/document/search_test.go | 537 ++++++ internal/vector/document/status.go | 45 + internal/vector/document/status_test.go | 47 + internal/vector/document/worker.go | 650 +++++++ internal/vector/document/worker_test.go | 1149 ++++++++++++ internal/vector/embed/client.go | 23 +- internal/vector/embed/document.go | 4 +- internal/vector/embed/request_gate.go | 4 +- internal/vector/embed/voyage.go | 23 +- internal/vector/errors.go | 10 + internal/vector/pgvector/document_vectors.go | 274 +++ .../vector/pgvector/document_vectors_stub.go | 28 + .../vector/pgvector/document_vectors_test.go | 327 ++++ internal/vector/pgvector/migrate.go | 38 + internal/vector/pgvector/schema.sql | 22 +- internal/vector/sqlitevec/document_vectors.go | 294 +++ .../vector/sqlitevec/document_vectors_stub.go | 28 + .../vector/sqlitevec/document_vectors_test.go | 268 +++ internal/vector/sqlitevec/migrate.go | 23 + internal/vector/sqlitevec/schema.sql | 13 +- nix/package.nix | 2 +- pkg/client/generated/client.go | 67 + pkg/client/generated/client_options.go | 44 + pkg/client/generated/client_with_response.go | 111 ++ pkg/client/generated/queries.go | 17 + pkg/client/generated/responses.go | 21 + pkg/client/generated/types.go | 288 ++- pkg/client/openapi.yaml | 346 ++++ web/src/lib/api/generated/schema.d.ts | 237 +++ 104 files changed, 14122 insertions(+), 243 deletions(-) create mode 100644 cmd/msgvault/cmd/documents_vector.go create mode 100644 cmd/msgvault/cmd/documents_vector_runtime.go create mode 100644 cmd/msgvault/cmd/documents_vector_stub.go create mode 100644 cmd/msgvault/cmd/documents_vector_stub_test.go create mode 100644 cmd/msgvault/cmd/documents_vector_test.go create mode 100644 internal/backupapp/restore.go create mode 100644 internal/store/document_normalized.go create mode 100644 internal/store/document_vector_lock.go create mode 100644 internal/store/document_vector_operations.go create mode 100644 internal/store/document_vectors.go create mode 100644 internal/store/document_vectors_pg_test.go create mode 100644 internal/store/document_vectors_test.go create mode 100644 internal/vector/document/backend.go create mode 100644 internal/vector/document/fusion.go create mode 100644 internal/vector/document/fusion_test.go create mode 100644 internal/vector/document/model.go create mode 100644 internal/vector/document/model_test.go create mode 100644 internal/vector/document/provider.go create mode 100644 internal/vector/document/reconcile.go create mode 100644 internal/vector/document/reconcile_test.go create mode 100644 internal/vector/document/search.go create mode 100644 internal/vector/document/search_test.go create mode 100644 internal/vector/document/status.go create mode 100644 internal/vector/document/status_test.go create mode 100644 internal/vector/document/worker.go create mode 100644 internal/vector/document/worker_test.go create mode 100644 internal/vector/pgvector/document_vectors.go create mode 100644 internal/vector/pgvector/document_vectors_stub.go create mode 100644 internal/vector/pgvector/document_vectors_test.go create mode 100644 internal/vector/sqlitevec/document_vectors.go create mode 100644 internal/vector/sqlitevec/document_vectors_stub.go create mode 100644 internal/vector/sqlitevec/document_vectors_test.go diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 83e681d06..08e3e6c08 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -92,6 +92,9 @@ jobs: - name: Test run: make test + - name: Test build without vector backends + run: go test -tags fts5 -run '^TestDocumentVectorStubKeepsStatusAvailableAndBuildActionable$' -count=1 ./cmd/msgvault/cmd + - name: Check OpenAPI artifacts run: make openapi-check diff --git a/Makefile b/Makefile index 08ab87361..8443c0328 100644 --- a/Makefile +++ b/Makefile @@ -45,7 +45,7 @@ PG_TEST_TAGS := fts5 sqlite_vec pgvector # in both configurations, so test-pg-both runs just these in the shipped-build # configuration. Verified by `make pg-shipped-only-check`, which re-derives the # closure from `go list`. -PG_SHIPPED_ONLY_PKGS := ./cmd/msgvault ./cmd/msgvault/cmd ./internal/api ./internal/mcp ./internal/scheduler ./internal/store ./internal/vector/chunkmatch ./internal/vector/embed ./internal/vector/hybrid ./internal/vector/pgvector ./scripts/contextual-retrieval-eval +PG_SHIPPED_ONLY_PKGS := ./cmd/msgvault ./cmd/msgvault/cmd ./internal/api ./internal/mcp ./internal/scheduler ./internal/store ./internal/vector/chunkmatch ./internal/vector/document ./internal/vector/embed ./internal/vector/hybrid ./internal/vector/pgvector ./scripts/contextual-retrieval-eval OPENAPI_ARTIFACTS := api/openapi.yaml pkg/client/openapi.yaml pkg/client/generated WEB_INSTALL_STAMP := web/node_modules/.msgvault-install-stamp diff --git a/api/openapi.yaml b/api/openapi.yaml index 113063d6e..d9bccc67a 100644 --- a/api/openapi.yaml +++ b/api/openapi.yaml @@ -2608,6 +2608,8 @@ components: DocumentSearchResponse: additionalProperties: true properties: + effective_mode: + type: string next_cursor: type: string results: @@ -2621,6 +2623,11 @@ components: type: integer truncated: type: boolean + vector_generation_fingerprint: + type: string + vector_generation_id: + format: int64 + type: integer required: - results - revision @@ -2652,6 +2659,9 @@ components: first_unit_index: format: int64 type: integer + fusion_score: + format: double + type: number heading_path: items: type: string @@ -2667,6 +2677,9 @@ components: last_unit_index: format: int64 type: integer + lexical_rank: + format: int64 + type: integer matched_signals: items: type: string @@ -2697,6 +2710,12 @@ components: rank: format: int64 type: integer + semantic_rank: + format: int64 + type: integer + semantic_score: + format: double + type: number source_id: format: int64 type: integer @@ -2706,6 +2725,20 @@ components: type: string truncated: type: boolean + vector_dimension: + format: int64 + type: integer + vector_embedding_profile: + type: string + vector_generation_fingerprint: + type: string + vector_generation_id: + format: int64 + type: integer + vector_model: + type: string + vector_token: + type: string required: - attachment_id - message_id @@ -2729,6 +2762,255 @@ components: - truncated - rank type: object + DocumentVectorConsent: + additionalProperties: true + properties: + consented_at: + format: date-time + type: string + dimension: + format: int64 + type: integer + egress_fingerprint: + type: string + embedding_profile: + type: string + fingerprint: + type: string + model: + type: string + purpose: + type: string + target_extraction_profile_id: + type: string + required: + - consented_at + - egress_fingerprint + - purpose + - fingerprint + - target_extraction_profile_id + - embedding_profile + - model + - dimension + type: object + DocumentVectorCoverage: + additionalProperties: true + properties: + ready: + format: int64 + type: integer + required: + format: int64 + type: integer + required: + - required + - ready + type: object + DocumentVectorFailureDiagnostic: + additionalProperties: true + properties: + attempt_count: + format: int64 + type: integer + error_code: + type: string + next_retry_at: + format: date-time + type: string + terminal: + type: boolean + token: + type: string + required: + - token + - attempt_count + - terminal + - error_code + type: object + DocumentVectorGeneration: + additionalProperties: true + properties: + activated_at: + format: date-time + type: string + created_at: + format: date-time + type: string + dimension: + format: int64 + type: integer + embedding_profile: + type: string + fingerprint: + type: string + id: + format: int64 + type: integer + model: + type: string + retired_at: + format: date-time + type: string + state: + type: string + target_extraction_profile_id: + type: string + required: + - id + - state + - created_at + - fingerprint + - target_extraction_profile_id + - embedding_profile + - model + - dimension + type: object + DocumentVectorGenerationSpec: + additionalProperties: true + properties: + dimension: + format: int64 + type: integer + embedding_profile: + type: string + fingerprint: + type: string + model: + type: string + target_extraction_profile_id: + type: string + required: + - fingerprint + - target_extraction_profile_id + - embedding_profile + - model + - dimension + type: object + DocumentVectorGenerationStatus: + additionalProperties: true + properties: + blocked: + type: boolean + cleanup_pending: + format: int64 + type: integer + failure_after_generation_id: + format: int64 + type: integer + failure_after_token: + type: string + failures: + items: + $ref: "#/components/schemas/DocumentVectorFailureDiagnostic" + type: + - array + - "null" + failures_exhausted: + type: boolean + generation_id: + format: int64 + type: integer + pending: + format: int64 + type: integer + ready_live: + format: int64 + type: integer + retryable: + format: int64 + type: integer + stale_obsolete: + format: int64 + type: integer + state: + type: string + terminal: + format: int64 + type: integer + required: + - generation_id + - state + - blocked + - pending + - retryable + - terminal + - ready_live + - stale_obsolete + - cleanup_pending + - failures + - failures_exhausted + type: object + DocumentVectorOperationsResponse: + additionalProperties: true + properties: + configured: + type: boolean + enabled: + type: boolean + scheduled_registration_requires_restart: + type: boolean + status: + $ref: "#/components/schemas/DocumentVectorOperationsStatus" + required: + - enabled + - configured + type: object + DocumentVectorOperationsStatus: + additionalProperties: true + properties: + active: + $ref: "#/components/schemas/DocumentVectorGeneration" + building: + $ref: "#/components/schemas/DocumentVectorGeneration" + configured_document_egress_fingerprint: + type: string + configured_query_egress_fingerprint: + type: string + configured_spec: + $ref: "#/components/schemas/DocumentVectorGenerationSpec" + coverage: + $ref: "#/components/schemas/DocumentVectorCoverage" + document_consent: + $ref: "#/components/schemas/DocumentVectorConsent" + query_consent: + $ref: "#/components/schemas/DocumentVectorConsent" + selected: + $ref: "#/components/schemas/DocumentVectorGenerationStatus" + usage: + $ref: "#/components/schemas/DocumentVectorProviderUsage" + required: + - configured_spec + - configured_document_egress_fingerprint + - configured_query_egress_fingerprint + - usage + type: object + DocumentVectorProviderUsage: + additionalProperties: true + properties: + fingerprint: + type: string + provider_calls: + format: int64 + type: integer + provider_chunks: + format: int64 + type: integer + provider_documents: + format: int64 + type: integer + provider_input_chars: + format: int64 + type: integer + updated_at: + format: date-time + type: string + required: + - fingerprint + - provider_calls + - provider_documents + - provider_chunks + - provider_input_chars + type: object DomainContextSummaryHTTPResponse: additionalProperties: true properties: @@ -13972,6 +14254,17 @@ paths: name: cursor schema: type: string + - description: "Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider" + in: query + name: mode + schema: + type: string + - description: "Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000)" + in: query + name: candidate_limit + schema: + format: int64 + type: integer responses: "200": content: @@ -14099,6 +14392,69 @@ paths: summary: Get extracted document index status tags: - API + /api/v1/documents/vectors/status: + get: + operationId: getDocumentVectorStatus + parameters: + - description: Generation whose bounded failures to inspect + in: query + name: generation_id + schema: + format: int64 + type: integer + - description: Stable failure cursor token + in: query + name: after_token + schema: + type: string + - description: Maximum failure diagnostics (default 20, max 1000) + in: query + name: limit + schema: + format: int64 + type: integer + responses: + "200": + content: + application/json: + schema: + $ref: "#/components/schemas/DocumentVectorOperationsResponse" + description: OK + "400": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + "403": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + "429": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + "503": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + default: + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + security: + - apiKey: [] + summary: Get document vector generation, consent, usage, and failure status + tags: + - API /api/v1/domains/search: post: operationId: searchDomains diff --git a/cmd/msgvault/cmd/add_synctech_sms_drive.go b/cmd/msgvault/cmd/add_synctech_sms_drive.go index c405bb276..eddd8afb7 100644 --- a/cmd/msgvault/cmd/add_synctech_sms_drive.go +++ b/cmd/msgvault/cmd/add_synctech_sms_drive.go @@ -138,7 +138,7 @@ func runConfiguredSynctechSMSSourceWithStoreDriveClient(ctx context.Context, st } var err error switch src.Backend { - case "", "local": + case "", localValue: if src.Path == "" { return fmt.Errorf("synctech-sms source %q path is required for local backend", src.Name) } diff --git a/cmd/msgvault/cmd/add_synctech_sms_drive_test.go b/cmd/msgvault/cmd/add_synctech_sms_drive_test.go index ae3312e58..04b08d4dc 100644 --- a/cmd/msgvault/cmd/add_synctech_sms_drive_test.go +++ b/cmd/msgvault/cmd/add_synctech_sms_drive_test.go @@ -318,7 +318,7 @@ func TestRunConfiguredSynctechSMSSourceLeavesManualSyncMessagesUnstamped(t *test src := config.SynctechSMSSource{ Name: "pixel-local", - Backend: "local", + Backend: localValue, Path: importDir, OwnerPhone: "+15550000001", IncludeSMS: true, @@ -354,7 +354,7 @@ func TestConfiguredSynctechSMSCompletesAfterImport(t *testing.T) { `), 0o600), "write sms fixture") src := synctechDriveTestSource() - src.Backend = "local" + src.Backend = localValue src.Path = xmlPath err := runConfiguredSynctechSMSSourceWithStore(context.Background(), f.Store, src) diff --git a/cmd/msgvault/cmd/backup.go b/cmd/msgvault/cmd/backup.go index 5db7e8975..31201c82b 100644 --- a/cmd/msgvault/cmd/backup.go +++ b/cmd/msgvault/cmd/backup.go @@ -265,13 +265,33 @@ func runBackupRestore(cmd *cobra.Command, args []string) error { PackedContent: backupRestorePackedContentTarget(looseAttachments), TargetCoordinator: targetCoordinatorOption, AuxiliaryTarget: backupapp.NewDocumentAuxiliaryTarget(), + BeforePublication: backupapp.InvalidateRestoredDocumentVectors, }) if err != nil { return fmt.Errorf("restoring snapshot: %w", err) } + if backupRestoreOverwrite { + if err := removeRestoredSQLiteVectorBackend(backupRestoreTarget); err != nil { + return err + } + } return printBackupRestoreSummary(cmd.OutOrStdout(), backupRestoreTarget, res, looseAttachments) } +func removeRestoredSQLiteVectorBackend(target string) error { + root, err := os.OpenRoot(target) + if err != nil { + return fmt.Errorf("backup restore: open restored target for vector reset: %w", err) + } + defer func() { _ = root.Close() }() + for _, name := range []string{"vectors.db", "vectors.db-wal", "vectors.db-shm", "vectors.db-journal"} { + if err := root.Remove(name); err != nil && !errors.Is(err, os.ErrNotExist) { + return fmt.Errorf("backup restore: remove excluded vector backend %s: %w", name, err) + } + } + return nil +} + func backupRestorePackedContentTarget(loose bool) backup.PackedContentTarget { limits := packstore.DefaultLimits() if loose { diff --git a/cmd/msgvault/cmd/backup_test.go b/cmd/msgvault/cmd/backup_test.go index 3f6684707..da151cb6f 100644 --- a/cmd/msgvault/cmd/backup_test.go +++ b/cmd/msgvault/cmd/backup_test.go @@ -152,6 +152,31 @@ func TestRunBackupRestorePackedDefaultAndExplicitLooseCleanup(t *testing.T) { require.NoError(os.WriteFile(loosePath, content, 0o600)) require.NoError(st.UpsertAttachment(messageID, "restore-cli.bin", "application/octet-stream", hash[:2]+"/"+hash, hash, len(content))) + profileFingerprint := strings.Repeat("d", 64) + profile := store.DocumentExtractionProfile{ + ID: "profile-" + profileFingerprint, Fingerprint: profileFingerprint, + Provider: "synthetic", Endpoint: "https://documents.example.test/v1", + Region: localValue, Model: "extract-test", RetentionPosture: "standard", + TrainingPosture: "opted-out", AllowedMediaTypes: []string{"application/pdf"}, + PolicyJSON: []byte(`{"policy":1}`), + } + _, err = st.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(err) + var vectorGenerationID int64 + require.NoError(st.DB().QueryRow(st.Rebind(` + INSERT INTO document_vector_generations + (fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state) + VALUES (?, ?, 'vector.embeddings', 'embed-test', 3, 'active') RETURNING id`), + strings.Repeat("e", 64), profile.ID).Scan(&vectorGenerationID)) + _, err = st.DB().Exec(st.Rebind(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, + token, state, created_at, updated_at) + VALUES (?, 'extraction-restore', ?, ?, 'input-restore', 1, 'chunk-restore', ?, 1, + 'token-restore', 'ready', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`), + vectorGenerationID, profile.ID, strings.Repeat("f", 64), strings.Repeat("a", 64)) + require.NoError(err) layout, err := packstore.NewLayout(attachmentsDir, packstore.LayoutOptions{ Staging: packstore.StagingSameDirectory, }) @@ -208,8 +233,13 @@ func TestRunBackupRestorePackedDefaultAndExplicitLooseCleanup(t *testing.T) { assert.Contains(packedOutput.String(), "page and blob hashes verified; manifest stats match") assert.NotContains(packedOutput.String(), "SQLite integrity_check") assertRestoredCLIBlob(t, backupRestoreTarget, hash, content, true) + assertRestoredDocumentVectorsInvalidated(t, backupRestoreTarget) backupRestoreTarget = filepath.Join(t.TempDir(), "loose-target") + require.NoError(os.MkdirAll(backupRestoreTarget, 0o700)) + staleVectorPath := filepath.Join(backupRestoreTarget, "vectors.db") + require.NoError(os.WriteFile(staleVectorPath, []byte("stale derived vectors"), 0o600)) + backupRestoreOverwrite = true backupRestoreLooseAttachments = true backupRestoreIntegrityCheck = true var looseOutput bytes.Buffer @@ -220,6 +250,8 @@ func TestRunBackupRestorePackedDefaultAndExplicitLooseCleanup(t *testing.T) { assert.Contains(looseOutput.String(), "Pack metadata cleared") assert.Contains(looseOutput.String(), "SQLite integrity_check ok") assertRestoredCLIBlob(t, backupRestoreTarget, hash, content, false) + assertRestoredDocumentVectorsInvalidated(t, backupRestoreTarget) + assert.NoFileExists(staleVectorPath, "overwrite restore removes the excluded derived vector backend") } func TestRunBackupRestoreIntoNonexistentConfiguredDataDir(t *testing.T) { @@ -514,6 +546,21 @@ func assertRestoredCLIBlob(t *testing.T, target, hash string, want []byte, packe assert.Equal(want, got) } +func assertRestoredDocumentVectorsInvalidated(t *testing.T, target string) { + t.Helper() + require := require.New(t) + assert := assert.New(t) + restored, err := store.OpenForTest(filepath.Join(target, "msgvault.db")) + require.NoError(err) + defer func() { require.NoError(restored.Close()) }() + active, err := restored.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(err) + assert.Nil(active) + var publications int64 + require.NoError(restored.DB().QueryRow(`SELECT COUNT(*) FROM document_vector_publications`).Scan(&publications)) + assert.Zero(publications) +} + func TestPrintBackupRestoreSummaryReportsPackedMixedAndLooseLayouts(t *testing.T) { tests := []struct { name string diff --git a/cmd/msgvault/cmd/constants.go b/cmd/msgvault/cmd/constants.go index 4a9fd4878..b3eec414b 100644 --- a/cmd/msgvault/cmd/constants.go +++ b/cmd/msgvault/cmd/constants.go @@ -42,6 +42,11 @@ const cmdUseList = "list" // cmdUseResume is the shared Cobra use/name for resume subcommands. const cmdUseResume = "resume" +// cmdUseConsent is the shared Cobra use/name for consent subcommands. +const cmdUseConsent = "consent" + +const localValue = "local" + // outputFormatJSON is the "json" value accepted by the --format flag. const outputFormatJSON = "json" diff --git a/cmd/msgvault/cmd/documents.go b/cmd/msgvault/cmd/documents.go index ea81445c3..4a6c1952a 100644 --- a/cmd/msgvault/cmd/documents.go +++ b/cmd/msgvault/cmd/documents.go @@ -23,6 +23,7 @@ import ( "go.kenn.io/msgvault/internal/personscope" personresolver "go.kenn.io/msgvault/internal/personscope/resolver" "go.kenn.io/msgvault/internal/store" + vectordocument "go.kenn.io/msgvault/internal/vector/document" ) const ( @@ -90,6 +91,7 @@ type documentsCommandDeps struct { openStore func() (*store.Store, func(), error) openAttachments func(*store.Store) (documentindex.DocumentAttachmentOpener, func() error, error) openReadClient func(context.Context) (documentReadClient, func(), error) + runDocumentVector func(context.Context, *store.Store, int64, int) (vectordocument.ReconcileResult, error) } type documentReadClient interface { @@ -110,6 +112,7 @@ func defaultDocumentsCommandDeps() documentsCommandDeps { validateProbeFixtures: mistral.ValidateProbeFixtures, runCapabilityProbe: mistral.RunCapabilityProbe, openStore: openWritableStoreAndInit, + runDocumentVector: runConfiguredDocumentVectorGeneration, openAttachments: openDocumentAttachments, openReadClient: func(ctx context.Context) (documentReadClient, func(), error) { client, _, err := OpenHTTPStore(ctx) @@ -135,6 +138,7 @@ func newDocumentsCmd(deps documentsCommandDeps) *cobra.Command { parent.AddCommand(newRetryDocumentCmd(deps)) parent.AddCommand(newRetireDocumentProfileCmd(deps)) parent.AddCommand(newPurgeDocumentDerivedCmd(deps)) + parent.AddCommand(newDocumentVectorsCmd(deps)) return parent } @@ -194,6 +198,8 @@ func newSearchDocumentsCmd(deps documentsCommandDeps) *cobra.Command { command.Flags().StringVar(&beforeValue, "before", "", "Only messages before YYYY-MM-DD or RFC3339") command.Flags().IntVarP(&request.PageSize, "limit", "n", 20, "Maximum results to return") command.Flags().StringVar(&request.Cursor, "cursor", "", "Opaque cursor from the previous page") + command.Flags().StringVar(&request.SearchMode, "mode", "lexical", "Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider") + command.Flags().IntVar(&request.CandidateLimit, "candidate-limit", 0, "Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000)") command.Flags().BoolVar(&jsonOutput, flagJSON, false, "Output structured JSON") return command } @@ -1081,6 +1087,13 @@ func (c localDocumentReadClient) SearchDocuments( } request.Person = &resolved.Scope } + mode, err := vectordocument.ParseSearchMode(request.SearchMode) + if err != nil { + return store.DocumentSearchResponse{}, fmt.Errorf("%w: %w", store.ErrDocumentSearchInvalidRequest, err) + } + if mode == vectordocument.SearchModeSemantic || mode == vectordocument.SearchModeHybrid { + return store.DocumentSearchResponse{}, vectordocument.ErrSemanticSearchUnavailable + } if err := reconcileDocumentOccurrencesForSearch(ctx, c.store); err != nil { return store.DocumentSearchResponse{}, err } diff --git a/cmd/msgvault/cmd/documents_routing_test.go b/cmd/msgvault/cmd/documents_routing_test.go index 39afd19e3..bf9a829b8 100644 --- a/cmd/msgvault/cmd/documents_routing_test.go +++ b/cmd/msgvault/cmd/documents_routing_test.go @@ -9,6 +9,51 @@ import ( "go.kenn.io/msgvault/internal/documentindex" ) +func TestDocumentVectorCommandsRouteWithConfiguredRemote(t *testing.T) { + const ( + apiKeyEnv = "MSGVAULT_VECTOR_TEST_KEY" + apiKey = "synthetic-vector-key" + ) + tests := []struct { + name string + args []string + wantArgs []string + wantAPIKey bool + }{ + {name: "consent", args: []string{"documents", "vectors", "consent", "--yes"}, wantArgs: []string{"documents", "vectors", "consent", "--yes"}}, + {name: "build", args: []string{"documents", "vectors", "build", "--limit", "5"}, wantArgs: []string{"documents", "vectors", "build", "--limit=5"}, wantAPIKey: true}, + {name: "resume", args: []string{"documents", "vectors", "resume", "--generation-id", "7", "--limit", "5"}, wantArgs: []string{"documents", "vectors", "resume", "--generation-id=7", "--limit=5"}, wantAPIKey: true}, + {name: "retry", args: []string{"documents", "vectors", "retry", "--generation-id", "7"}, wantArgs: []string{"documents", "vectors", "retry", "--generation-id=7"}}, + {name: "rebuild", args: []string{"documents", "vectors", "rebuild", "--generation-id", "7", "--yes"}, wantArgs: []string{"documents", "vectors", "rebuild", "--generation-id=7", "--yes"}, wantAPIKey: true}, + {name: "retire", args: []string{"documents", "vectors", "retire", "--generation-id", "7", "--yes"}, wantArgs: []string{"documents", "vectors", "retire", "--generation-id=7", "--yes"}}, + {name: "status", args: []string{"documents", "vectors", "status", "--json"}, wantArgs: []string{"documents", "vectors", "status", "--json"}}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + assert := assert.New(t) + server, requests := newDaemonCLIRunnerTestServer(t, func(req daemonCLIRunTestRequest) { + assert.Equal(test.wantArgs, req.Args) + if test.wantAPIKey { + assert.Equal(map[string]string{apiKeyEnv: apiKey}, req.Env) + } else { + assert.Empty(req.Env) + } + }, `{"type":"complete"}`) + configureRemoteDaemonForTest(t, server.URL) + cfg.Vector.Embeddings.APIKeyEnv = apiKeyEnv + t.Setenv(apiKeyEnv, apiKey) + + root := &cobra.Command{Use: "msgvault"} + root.AddCommand(newDocumentsCmd(documentsCommandDeps{})) + root.SetArgs(test.args) + + require.NoError(t, root.ExecuteContext(t.Context())) + assert.Equal(1, int(requests.Load())) + }) + } +} + func TestDocumentMutationsRouteSafelyWithConfiguredRemote(t *testing.T) { const ( apiKeyEnv = "MSGVAULT_DOCUMENT_TEST_KEY" diff --git a/cmd/msgvault/cmd/documents_test.go b/cmd/msgvault/cmd/documents_test.go index 2858eeeca..8bf48944c 100644 --- a/cmd/msgvault/cmd/documents_test.go +++ b/cmd/msgvault/cmd/documents_test.go @@ -26,6 +26,7 @@ import ( "go.kenn.io/msgvault/internal/scheduler" "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/testutil/storetest" + vectordocument "go.kenn.io/msgvault/internal/vector/document" ) func TestProbeMistralCommandWritesCompleteSanitizedManifest(t *testing.T) { @@ -406,6 +407,16 @@ func TestDocumentsSearchDoesNotRegisterUnconsentedJournalConsumer(t *testing.T) require.ErrorIs(err, store.ErrAttachmentChangeConsumerMissing) } +func TestDocumentsSearchLocalExplicitSemanticNeverMasqueradesAsLexical(t *testing.T) { + fixture := storetest.New(t) + command := newDocumentsCmd(documentsCommandDeps{ + openStore: func() (*store.Store, func(), error) { return fixture.Store, func() {}, nil }, + }) + command.SetArgs([]string{"search", "evidence", "--mode", "semantic", "--candidate-limit", "25"}) + err := command.ExecuteContext(t.Context()) + require.ErrorIs(t, err, vectordocument.ErrSemanticSearchUnavailable) +} + func TestDocumentsSearchUsesConfiguredReadClient(t *testing.T) { assert := assert.New(t) openStoreCalled := false @@ -416,6 +427,7 @@ func TestDocumentsSearchUsesConfiguredReadClient(t *testing.T) { request store.DocumentSearchRequest, ) (store.DocumentSearchResponse, error) { assert.Equal("damage report", request.Query) + assert.Zero(request.CandidateLimit) assert.Equal(int64(40), request.PersonID) assert.Equal([]personscope.Direction{personscope.FromPerson, personscope.Group}, request.Directions) require.NotNil(t, request.After) diff --git a/cmd/msgvault/cmd/documents_vector.go b/cmd/msgvault/cmd/documents_vector.go new file mode 100644 index 000000000..0158cff6b --- /dev/null +++ b/cmd/msgvault/cmd/documents_vector.go @@ -0,0 +1,460 @@ +package cmd + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/url" + "time" + + "github.com/spf13/cobra" + "go.kenn.io/msgvault/internal/store" + vectordocument "go.kenn.io/msgvault/internal/vector/document" +) + +const ( + defaultDocumentVectorOperationLimit = 100 + documentVectorsSubcommand = "vectors" +) + +func newDocumentVectorsCmd(deps documentsCommandDeps) *cobra.Command { + command := &cobra.Command{Use: documentVectorsSubcommand, Short: "Manage document attachment vectors"} + command.AddCommand( + newDocumentVectorConsentCmd(deps), + newDocumentVectorBuildCmd(deps, false), + newDocumentVectorBuildCmd(deps, true), + newDocumentVectorRetryCmd(deps), + newDocumentVectorRebuildCmd(deps), + newDocumentVectorRetireCmd(deps), + newDocumentVectorStatusCmd(deps), + ) + return command +} + +func desiredDocumentVectorSpec(ctx context.Context, st *store.Store) (store.DocumentVectorGenerationSpec, error) { + if cfg == nil || !cfg.Vector.Enabled || !cfg.Attachments.Documents.Index.Embeddings.Enabled { + return store.DocumentVectorGenerationSpec{}, errors.New("document embeddings are disabled; enable [vector] and [attachments.documents.index.embeddings]") + } + return configuredDocumentVectorSpec(ctx, st) +} + +func configuredDocumentVectorSpec(ctx context.Context, st *store.Store) (store.DocumentVectorGenerationSpec, error) { + if cfg == nil || !cfg.Attachments.Documents.Index.Embeddings.Enabled { + return store.DocumentVectorGenerationSpec{}, errors.New("document embeddings are not configured") + } + if err := cfg.Vector.Embeddings.Validate(); err != nil { + return store.DocumentVectorGenerationSpec{}, fmt.Errorf("document embedding config: %w", err) + } + if err := cfg.Vector.Validate(); err != nil { + return store.DocumentVectorGenerationSpec{}, fmt.Errorf("vector config: %w", err) + } + target, err := st.GetDocumentVectorTargetProfileID(ctx) + if err != nil { + return store.DocumentVectorGenerationSpec{}, err + } + fingerprint, err := vectordocument.Fingerprint(target, cfg.Vector) + if err != nil { + return store.DocumentVectorGenerationSpec{}, err + } + return store.DocumentVectorGenerationSpec{ + Fingerprint: fingerprint, + TargetExtractionProfileID: target, + EmbeddingProfile: cfg.Attachments.Documents.Index.Embeddings.Profile, + Model: cfg.Vector.Embeddings.Model, + Dimension: cfg.Vector.Embeddings.Dimension, + }, nil +} + +func configuredDocumentVectorConsentSpec(spec store.DocumentVectorGenerationSpec) (store.DocumentVectorConsentSpec, error) { + egressFingerprint, err := vectordocument.EgressFingerprint(spec.TargetExtractionProfileID, cfg.Vector) + if err != nil { + return store.DocumentVectorConsentSpec{}, err + } + return store.DocumentVectorConsentSpec{ + DocumentVectorGenerationSpec: spec, + EgressFingerprint: egressFingerprint, + Purpose: "document_embedding", + }, nil +} + +func configuredDocumentVectorQueryConsentSpec(spec store.DocumentVectorGenerationSpec) (store.DocumentVectorConsentSpec, error) { + egressFingerprint, err := vectordocument.QueryEgressFingerprint(spec.TargetExtractionProfileID, cfg.Vector) + if err != nil { + return store.DocumentVectorConsentSpec{}, err + } + return store.DocumentVectorConsentSpec{ + DocumentVectorGenerationSpec: spec, + EgressFingerprint: egressFingerprint, + Purpose: "query_embedding", + }, nil +} + +func withDocumentVectorStore(deps documentsCommandDeps, fn func(*store.Store) error) error { + if deps.openStore == nil { + return errors.New("document vector ledger is unavailable") + } + st, cleanup, err := deps.openStore() + if err != nil { + return err + } + defer cleanup() + return fn(st) +} + +func runDocumentVectorCommandHTTP(command *cobra.Command, args []string, forwardEmbeddingKey bool) error { + if forwardEmbeddingKey { + return runDaemonCLICommandHTTPFromCobraWithEnv(command, args, embeddingsForwardEnv()) + } + return runDaemonCLICommandHTTPFromCobra(command, args) +} + +func newDocumentVectorConsentCmd(deps documentsCommandDeps) *cobra.Command { + var yes bool + var purpose string + command := &cobra.Command{ + Use: cmdUseConsent, + Short: "Consent to hosted embedding for the configured document policy", + Args: cobra.NoArgs, + RunE: func(command *cobra.Command, args []string) error { + if !isDaemonCLISubprocess() { + return runDocumentVectorCommandHTTP(command, args, false) + } + return withDocumentVectorStore(deps, func(st *store.Store) error { + spec, err := configuredDocumentVectorSpec(command.Context(), st) + if err != nil { + return err + } + var consentSpec store.DocumentVectorConsentSpec + switch purpose { + case "documents": + consentSpec, err = configuredDocumentVectorConsentSpec(spec) + case "queries": + consentSpec, err = configuredDocumentVectorQueryConsentSpec(spec) + default: + return errors.New("document vector consent purpose must be documents or queries") + } + if err != nil { + return err + } + printDocumentVectorConsentDisclosure(command.OutOrStdout(), consentSpec) + if !yes { + return errors.New("hosted document embedding consent requires --yes after reviewing the provider disclosure") + } + consent, _, err := st.RecordDocumentVectorConsent(command.Context(), consentSpec, time.Now()) + if err != nil { + return err + } + _, _ = fmt.Fprintf(command.OutOrStdout(), "Recorded consent for document vector egress fingerprint %s. Restart the daemon to enable scheduled document vector work.\n", consent.EgressFingerprint) + return nil + }) + }, + } + command.Flags().BoolVar(&yes, "yes", false, "Confirm hosted document embedding consent") + command.Flags().StringVar(&purpose, "purpose", "documents", "Plaintext class to consent: documents or queries") + return command +} + +func printDocumentVectorConsentDisclosure(w io.Writer, spec store.DocumentVectorConsentSpec) { + authentication := "no authentication environment variable configured" + if cfg.Vector.Embeddings.APIKeyEnv != "" { + authentication = "environment variable " + cfg.Vector.Embeddings.APIKeyEnv + } + _, _ = fmt.Fprintln(w, "Hosted document embedding disclosure:") + _, _ = fmt.Fprintf(w, "Corpus fingerprint: %s\n", spec.Fingerprint) + _, _ = fmt.Fprintf(w, "Egress fingerprint: %s\n", spec.EgressFingerprint) + _, _ = fmt.Fprintf(w, "Destination: %s\n", documentVectorConsentEndpoint(cfg.Vector.Embeddings.Endpoint)) + _, _ = fmt.Fprintf(w, "Authentication: %s\n", authentication) + _, _ = fmt.Fprintf(w, "API format: %s\n", cfg.Vector.Embeddings.EffectiveAPIFormat()) + _, _ = fmt.Fprintf(w, "Model: %s\n", spec.Model) + _, _ = fmt.Fprintf(w, "Dimension: %d\n", spec.Dimension) + _, _ = fmt.Fprintf(w, "Maximum input: %d characters\n", cfg.Vector.Embeddings.MaxInputChars) + if spec.Purpose == "query_embedding" { + _, _ = fmt.Fprintln(w, "Explicit semantic or hybrid document searches will send query text to the configured destination.") + } else { + _, _ = fmt.Fprintln(w, "Docbank-prepared normalized attachment document inputs will be sent to the configured destination for embedding.") + } +} + +func documentVectorConsentEndpoint(endpoint string) string { + parsed, err := url.Parse(endpoint) + if err != nil { + return "[configured endpoint omitted: invalid URL]" + } + parsed.User = nil + parsed.RawQuery = "" + parsed.ForceQuery = false + parsed.Fragment = "" + return parsed.String() +} + +func newDocumentVectorBuildCmd(deps documentsCommandDeps, resume bool) *cobra.Command { + name, short := documentBuildSubcommand, "Build the configured document vector generation" + if resume { + name, short = cmdUseResume, "Resume a building generation or retired-generation cleanup" + } + var generationID int64 + var limit int + command := &cobra.Command{ + Use: name, Short: short, Args: cobra.NoArgs, + RunE: func(command *cobra.Command, args []string) error { + if !isDaemonCLISubprocess() { + return runDocumentVectorCommandHTTP(command, args, true) + } + return withDocumentVectorStore(deps, func(st *store.Store) error { + return st.WithDocumentVectorOperationLock(command.Context(), func() error { + if resume { + generation, err := st.GetDocumentVectorGeneration(command.Context(), generationID) + if err != nil { + return err + } + if generation.State == store.DocumentVectorGenerationRetired { + return runDocumentVectorCommand(command, deps, st, generationID, limit) + } + } + spec, err := desiredDocumentVectorSpec(command.Context(), st) + if err != nil { + return err + } + if err := requireDocumentVectorConsent(command.Context(), st, spec); err != nil { + return err + } + if resume { + generation, err := st.GetDocumentVectorGeneration(command.Context(), generationID) + if err != nil { + return err + } + if generation.State != store.DocumentVectorGenerationBuilding || generation.DocumentVectorGenerationSpec != spec { + return store.ErrDocumentVectorInvalidGenerationState + } + } else { + generation, _, err := st.EnsureDocumentVectorGeneration(command.Context(), spec) + if err != nil { + return err + } + if generation.State != store.DocumentVectorGenerationBuilding { + return errors.New("the configured generation is already active; use documents vectors rebuild for coverage drift") + } + generationID = generation.ID + } + return runDocumentVectorCommand(command, deps, st, generationID, limit) + }) + }) + }, + } + command.Flags().IntVarP(&limit, "limit", "n", defaultDocumentVectorOperationLimit, "Maximum chunks and cleanup tokens to process (1-1000)") + if resume { + command.Flags().Int64Var(&generationID, "generation-id", 0, "Building or retired generation to resume") + _ = command.MarkFlagRequired("generation-id") + } + return command +} + +func requireDocumentVectorConsent(ctx context.Context, st *store.Store, spec store.DocumentVectorGenerationSpec) error { + consented, err := hasDocumentVectorConsent(ctx, st, spec) + if err != nil { + return err + } + if !consented { + return errors.New("exact document vector policy is not consented; run `msgvault documents vectors consent --yes`") + } + return nil +} + +func hasDocumentVectorConsent(ctx context.Context, st *store.Store, spec store.DocumentVectorGenerationSpec) (bool, error) { + consentSpec, err := configuredDocumentVectorConsentSpec(spec) + if err != nil { + return false, err + } + consent, err := st.GetDocumentVectorConsent(ctx, consentSpec.EgressFingerprint) + if err != nil { + return false, err + } + return consent != nil && consent.DocumentVectorConsentSpec == consentSpec, nil +} + +func runDocumentVectorCommand(command *cobra.Command, deps documentsCommandDeps, st *store.Store, generationID int64, limit int) error { + if limit < 1 || limit > 1000 { + return errors.New("document vector operation limit must be between 1 and 1000") + } + if deps.runDocumentVector == nil { + return errors.New("document vector backend is unavailable in this binary") + } + result, err := deps.runDocumentVector(command.Context(), st, generationID, limit) + if encodeErr := json.NewEncoder(command.OutOrStdout()).Encode(result); encodeErr != nil { + return errors.Join(err, encodeErr) + } + return err +} + +func newDocumentVectorRetryCmd(deps documentsCommandDeps) *cobra.Command { + var generationID int64 + var afterToken string + var limit int + command := &cobra.Command{Use: "retry", Short: "Reset current failed publications for retry", Args: cobra.NoArgs, + RunE: func(command *cobra.Command, args []string) error { + if !isDaemonCLISubprocess() { + return runDocumentVectorCommandHTTP(command, args, false) + } + return withDocumentVectorStore(deps, func(st *store.Store) error { + return st.WithDocumentVectorOperationLock(command.Context(), func() error { + result, err := st.ResetDocumentVectorFailures(command.Context(), generationID, afterToken, limit, time.Now()) + if err != nil { + return err + } + return json.NewEncoder(command.OutOrStdout()).Encode(result) + }) + }) + }} + command.Flags().Int64Var(&generationID, "generation-id", 0, "Generation whose failures should be reset") + command.Flags().StringVar(&afterToken, "after-token", "", "Stable retry cursor token") + command.Flags().IntVarP(&limit, "limit", "n", defaultDocumentVectorOperationLimit, "Maximum failures to scan (1-1000)") + _ = command.MarkFlagRequired("generation-id") + return command +} + +func newDocumentVectorRebuildCmd(deps documentsCommandDeps) *cobra.Command { + var activeID int64 + var limit int + var yes bool + command := &cobra.Command{Use: "rebuild", Short: "Build a fresh generation while the active generation remains searchable", Args: cobra.NoArgs, + RunE: func(command *cobra.Command, args []string) error { + if !isDaemonCLISubprocess() { + return runDocumentVectorCommandHTTP(command, args, true) + } + if !yes { + return errors.New("document vector rebuild requires --yes") + } + return withDocumentVectorStore(deps, func(st *store.Store) error { + return st.WithDocumentVectorOperationLock(command.Context(), func() error { + spec, err := desiredDocumentVectorSpec(command.Context(), st) + if err != nil { + return err + } + if err := requireDocumentVectorConsent(command.Context(), st, spec); err != nil { + return err + } + generation, err := st.StartDocumentVectorRebuild(command.Context(), activeID, spec, time.Now()) + if err != nil { + return err + } + return runDocumentVectorCommand(command, deps, st, generation.ID, limit) + }) + }) + }} + command.Flags().Int64Var(&activeID, "generation-id", 0, "Active generation being replaced") + command.Flags().IntVarP(&limit, "limit", "n", defaultDocumentVectorOperationLimit, "Maximum chunks and cleanup tokens to process (1-1000)") + command.Flags().BoolVar(&yes, "yes", false, "Confirm the rebuild") + _ = command.MarkFlagRequired("generation-id") + return command +} + +func newDocumentVectorRetireCmd(deps documentsCommandDeps) *cobra.Command { + var generationID int64 + var yes bool + command := &cobra.Command{Use: cliEmbeddingsOperationRetire, Short: "Retire a document vector generation without deleting its backend ledger", Args: cobra.NoArgs, + RunE: func(command *cobra.Command, args []string) error { + if !isDaemonCLISubprocess() { + return runDocumentVectorCommandHTTP(command, args, false) + } + if !yes { + return errors.New("document vector retirement requires --yes") + } + return withDocumentVectorStore(deps, func(st *store.Store) error { + return st.WithDocumentVectorOperationLock(command.Context(), func() error { + retired, err := st.RetireDocumentVectorGeneration(command.Context(), generationID, time.Now()) + if err != nil { + return err + } + _, _ = fmt.Fprintf(command.OutOrStdout(), "retired=%t generation_id=%d; backend cleanup will resume when vector operations next run\n", retired, generationID) + return nil + }) + }) + }} + command.Flags().Int64Var(&generationID, "generation-id", 0, "Generation to retire") + command.Flags().BoolVar(&yes, "yes", false, "Confirm retirement") + _ = command.MarkFlagRequired("generation-id") + return command +} + +func newDocumentVectorStatusCmd(deps documentsCommandDeps) *cobra.Command { + var generationID int64 + var afterToken string + var limit int + var jsonOutput bool + command := &cobra.Command{Use: statusValue, Short: "Inspect document vector generations, consent, usage, and failures", Args: cobra.NoArgs, + RunE: func(command *cobra.Command, args []string) error { + if !isDaemonCLISubprocess() { + return runDocumentVectorCommandHTTP(command, args, false) + } + if cfg == nil || !cfg.Vector.Enabled || !cfg.Attachments.Documents.Index.Embeddings.Enabled { + if jsonOutput { + return json.NewEncoder(command.OutOrStdout()).Encode(map[string]bool{"enabled": false}) + } + _, _ = fmt.Fprintln(command.OutOrStdout(), "document_vectors=disabled") + return nil + } + return withDocumentVectorStore(deps, func(st *store.Store) error { + spec, err := desiredDocumentVectorSpec(command.Context(), st) + if errors.Is(err, store.ErrDocumentVectorInvalidGenerationState) { + if jsonOutput { + return json.NewEncoder(command.OutOrStdout()).Encode(map[string]bool{"enabled": true, "configured": false}) + } + _, _ = fmt.Fprintln(command.OutOrStdout(), "document_vectors=enabled configured=false") + return nil + } + if err != nil { + return err + } + documentConsentSpec, err := configuredDocumentVectorConsentSpec(spec) + if err != nil { + return err + } + queryConsentSpec, err := configuredDocumentVectorQueryConsentSpec(spec) + if err != nil { + return err + } + status, err := st.GetDocumentVectorOperationsStatus(command.Context(), spec, documentConsentSpec.EgressFingerprint, queryConsentSpec.EgressFingerprint, generationID, afterToken, limit) + if err != nil { + return err + } + if jsonOutput { + return json.NewEncoder(command.OutOrStdout()).Encode(struct { + Enabled bool `json:"enabled"` + Configured bool `json:"configured"` + Status store.DocumentVectorOperationsStatus `json:"status"` + }{Enabled: true, Configured: true, Status: status}) + } + _, _ = fmt.Fprintf(command.OutOrStdout(), "configured_fingerprint=%s document_consented=%t query_consented=%t provider_calls=%d provider_documents=%d provider_chunks=%d provider_input_chars=%d\n", status.ConfiguredSpec.Fingerprint, status.DocumentConsent != nil, status.QueryConsent != nil, status.Usage.ProviderCalls, status.Usage.ProviderDocuments, status.Usage.ProviderChunks, status.Usage.ProviderInputChars) + if status.Active != nil { + _, _ = fmt.Fprintf(command.OutOrStdout(), "active_generation=%d state=%s\n", status.Active.ID, status.Active.State) + } else { + _, _ = fmt.Fprintln(command.OutOrStdout(), "active_generation=none") + } + if status.Building != nil { + _, _ = fmt.Fprintf(command.OutOrStdout(), "building_generation=%d state=%s\n", status.Building.ID, status.Building.State) + } else { + _, _ = fmt.Fprintln(command.OutOrStdout(), "building_generation=none") + } + if status.Selected != nil { + _, _ = fmt.Fprintf(command.OutOrStdout(), "selected_generation=%d state=%s blocked=%t pending=%d retryable=%d terminal=%d ready_live=%d obsolete=%d cleanup_pending=%d\n", + status.Selected.GenerationID, status.Selected.State, status.Selected.Blocked, status.Selected.Pending, + status.Selected.Retryable, status.Selected.Terminal, status.Selected.ReadyLive, + status.Selected.Obsolete, status.Selected.CleanupPending) + } + if status.Coverage != nil { + _, _ = fmt.Fprintf(command.OutOrStdout(), "coverage_required=%d coverage_ready=%d\n", status.Coverage.Required, status.Coverage.Ready) + } + if status.DocumentConsent != nil || status.QueryConsent != nil { + _, _ = fmt.Fprintln(command.OutOrStdout(), "scheduled_registration=restart-required-after-new-consent") + } + return nil + }) + }} + command.Flags().Int64Var(&generationID, "generation-id", 0, "Generation whose bounded failures to inspect") + command.Flags().StringVar(&afterToken, "after-token", "", "Stable failure cursor token") + command.Flags().IntVarP(&limit, "limit", "n", 20, "Maximum failure diagnostics (1-1000)") + command.Flags().BoolVar(&jsonOutput, flagJSON, false, "Output structured JSON") + return command +} diff --git a/cmd/msgvault/cmd/documents_vector_runtime.go b/cmd/msgvault/cmd/documents_vector_runtime.go new file mode 100644 index 000000000..460039868 --- /dev/null +++ b/cmd/msgvault/cmd/documents_vector_runtime.go @@ -0,0 +1,183 @@ +//go:build sqlite_vec || pgvector + +package cmd + +import ( + "context" + "errors" + "time" + + "github.com/google/uuid" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/vector" + vectordocument "go.kenn.io/msgvault/internal/vector/document" +) + +func nextDocumentVectorWorkerOwner() string { + return "document-vector-" + uuid.NewString() +} + +type checkpointingDocumentVectorWorker struct { + worker vectordocument.WorkerRunner + checkpointer documentVectorBuildCheckpointer + fingerprint string + now func() time.Time +} + +type documentVectorBuildCheckpointer interface { + CheckpointDocumentVectorBuildForFingerprint(ctx context.Context, generationID int64, fingerprint string, afterChunkID int64, exhausted bool, delta store.DocumentVectorUsageDelta, now time.Time) error +} + +func (w checkpointingDocumentVectorWorker) Run(ctx context.Context, generationID vectordocument.GenerationID, limit int) (vectordocument.RunResult, error) { + result, runErr := w.worker.Run(ctx, generationID, limit) + if !result.Exhausted && result.AfterGenerationID == 0 { + return result, runErr + } + delta := store.DocumentVectorUsageDelta{ + ProviderCalls: int64(result.ProviderCalls), ProviderDocuments: int64(result.ProviderDocuments), + ProviderChunks: int64(result.ProviderChunks), ProviderInputChars: int64(result.ProviderInputChars), + } + checkpointErr := w.checkpointer.CheckpointDocumentVectorBuildForFingerprint( + ctx, int64(generationID), w.fingerprint, result.AfterChunkID, result.Exhausted, delta, w.now(), + ) + return result, errors.Join(runErr, checkpointErr) +} + +func runConfiguredDocumentVectorGeneration(ctx context.Context, st *store.Store, generationID int64, limit int) (vectordocument.ReconcileResult, error) { + if limit < 1 || limit > 1000 { + return vectordocument.ReconcileResult{}, errors.New("document vector operation limit must be between 1 and 1000") + } + vf, err := setupVectorFeatures(ctx, st, cfg.DatabaseDSN(), false) + if err != nil { + return vectordocument.ReconcileResult{}, err + } + if vf == nil || vf.DocumentBackend == nil || vf.SemanticClient == nil { + return vectordocument.ReconcileResult{}, errors.New("document vector runtime is unavailable") + } + defer func() { _ = vf.Close() }() + return runDocumentVectorWithFeatures(ctx, st, vf, generationID, limit) +} + +func runDocumentVectorWithFeatures(ctx context.Context, st *store.Store, vf *vectorFeatures, generationID int64, limit int) (vectordocument.ReconcileResult, error) { + limit = min(limit, max(1, vf.Cfg.Embeddings.BatchSize)) + generation, err := st.GetDocumentVectorGeneration(ctx, generationID) + if err != nil { + return vectordocument.ReconcileResult{}, err + } + cursor, err := st.GetDocumentVectorBuildCursor(ctx, generationID) + if err != nil { + return vectordocument.ReconcileResult{}, err + } + now := func() time.Time { return time.Now().UTC() } + var afterGenerationID vectordocument.GenerationID + if cursor > 0 { + afterGenerationID = vectordocument.GenerationID(generationID) + } + worker := vectordocument.NewWorker(vectordocument.WorkerDeps{ + Ledger: st, Provider: vf.SemanticClient, Backend: vf.DocumentBackend, + Owner: nextDocumentVectorWorkerOwner(), Dimension: generation.Dimension, + MaxInputChars: cfg.Vector.Embeddings.MaxInputChars, + ContextualDocuments: cfg.Vector.Embeddings.EffectiveAPIFormat() == vector.APIFormatVoyageContextual, + LeaseDuration: 2 * time.Minute, HeartbeatInterval: 20 * time.Second, + RetryDelay: time.Minute, MaxAttempts: 5, + AfterGenerationID: afterGenerationID, AfterChunkID: cursor, Now: now, + }) + checkpointed := checkpointingDocumentVectorWorker{ + worker: worker, checkpointer: st, fingerprint: generation.Fingerprint, now: now, + } + reconciler := vectordocument.NewReconciler(vectordocument.ReconcilerDeps{ + Ledger: st, Worker: checkpointed, Backend: vf.DocumentBackend, Now: now, + }) + return reconciler.Run(ctx, vectordocument.GenerationID(generationID), limit) +} + +func runScheduledDocumentVectorGeneration(ctx context.Context, st *store.Store, vf *vectorFeatures, limit int) error { + return st.WithDocumentVectorOperationLock(ctx, func() error { + return runScheduledDocumentVectorGenerationLocked(ctx, st, vf, limit) + }) +} + +func runScheduledDocumentVectorGenerationLocked(ctx context.Context, st *store.Store, vf *vectorFeatures, limit int) error { + retired, err := st.GetOldestRetiredDocumentVectorGeneration(ctx) + if err != nil { + return err + } + if retired != nil { + reconciler := vectordocument.NewReconciler(vectordocument.ReconcilerDeps{ + Ledger: st, Backend: vf.DocumentBackend, Now: func() time.Time { return time.Now().UTC() }, + }) + _, err := reconciler.Run(ctx, vectordocument.GenerationID(retired.ID), limit) + return err + } + if vf.SemanticClient == nil { + return nil + } + spec, err := desiredDocumentVectorSpec(ctx, st) + if errors.Is(err, store.ErrDocumentVectorInvalidGenerationState) { + return nil + } + if err != nil { + return err + } + consented, err := hasDocumentVectorConsent(ctx, st, spec) + if err != nil { + return err + } + if !consented { + return nil + } + building, err := st.GetBuildingDocumentVectorGeneration(ctx) + if err != nil { + return err + } + if building != nil && building.DocumentVectorGenerationSpec != spec { + retired, retireErr := st.RetireDocumentVectorGeneration(ctx, building.ID, time.Now()) + if retireErr != nil { + return retireErr + } + if !retired { + return store.ErrDocumentVectorInvalidGenerationState + } + // Reconcile only the obsolete generation this pass. The next bounded + // run creates the desired generation without exceeding one cleanup page. + _, reconcileErr := runDocumentVectorWithFeatures(ctx, st, vf, building.ID, limit) + return reconcileErr + } + if building == nil { + active, err := st.GetActiveDocumentVectorGeneration(ctx) + if err != nil { + return err + } + switch { + case active == nil || active.DocumentVectorGenerationSpec != spec: + generation, _, ensureErr := st.EnsureDocumentVectorGeneration(ctx, spec) + if ensureErr != nil { + return ensureErr + } + building = &generation + default: + status, statusErr := st.GetDocumentVectorGenerationStatus(ctx, active.ID, "", limit) + if statusErr != nil { + return statusErr + } + if status.CleanupPending > 0 { + _, reconcileErr := runDocumentVectorWithFeatures(ctx, st, vf, active.ID, limit) + return reconcileErr + } + coverage, coverageErr := st.GetDocumentVectorCoverage(ctx, active.ID) + if coverageErr != nil { + return coverageErr + } + if coverage.Complete() { + return nil + } + generation, rebuildErr := st.StartDocumentVectorRebuild(ctx, active.ID, spec, time.Now()) + if rebuildErr != nil { + return rebuildErr + } + building = &generation + } + } + _, err = runDocumentVectorWithFeatures(ctx, st, vf, building.ID, limit) + return err +} diff --git a/cmd/msgvault/cmd/documents_vector_stub.go b/cmd/msgvault/cmd/documents_vector_stub.go new file mode 100644 index 000000000..35c493de8 --- /dev/null +++ b/cmd/msgvault/cmd/documents_vector_stub.go @@ -0,0 +1,19 @@ +//go:build !sqlite_vec && !pgvector + +package cmd + +import ( + "context" + "errors" + + "go.kenn.io/msgvault/internal/store" + vectordocument "go.kenn.io/msgvault/internal/vector/document" +) + +func runConfiguredDocumentVectorGeneration(context.Context, *store.Store, int64, int) (vectordocument.ReconcileResult, error) { + return vectordocument.ReconcileResult{}, errors.New("document vector backend is unavailable: rebuild with sqlite_vec or pgvector support") +} + +func runScheduledDocumentVectorGeneration(context.Context, *store.Store, *vectorFeatures, int) error { + return errors.New("document vector backend is unavailable: rebuild with sqlite_vec or pgvector support") +} diff --git a/cmd/msgvault/cmd/documents_vector_stub_test.go b/cmd/msgvault/cmd/documents_vector_stub_test.go new file mode 100644 index 000000000..d685cbcb2 --- /dev/null +++ b/cmd/msgvault/cmd/documents_vector_stub_test.go @@ -0,0 +1,30 @@ +//go:build !sqlite_vec && !pgvector + +package cmd + +import ( + "bytes" + "os" + "strconv" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/config" +) + +func TestDocumentVectorStubKeepsStatusAvailableAndBuildActionable(t *testing.T) { + previous := cfg + t.Cleanup(func() { cfg = previous }) + t.Setenv(daemonCLISubprocessEnv, strconv.Itoa(os.Getppid())) + cfg = config.NewDefaultConfig() + status := newDocumentsCmd(documentsCommandDeps{}) + var output bytes.Buffer + status.SetOut(&output) + status.SetArgs([]string{documentVectorsSubcommand, statusValue, "--json"}) + require.NoError(t, status.ExecuteContext(t.Context())) + assert.JSONEq(t, `{"enabled":false}`, output.String()) + + _, err := runConfiguredDocumentVectorGeneration(t.Context(), nil, 1, 1) + require.ErrorContains(t, err, "rebuild with sqlite_vec or pgvector support") +} diff --git a/cmd/msgvault/cmd/documents_vector_test.go b/cmd/msgvault/cmd/documents_vector_test.go new file mode 100644 index 000000000..42482c179 --- /dev/null +++ b/cmd/msgvault/cmd/documents_vector_test.go @@ -0,0 +1,572 @@ +//go:build sqlite_vec || pgvector + +package cmd + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "strconv" + "strings" + "testing" + "time" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/config" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/testutil/storetest" + "go.kenn.io/msgvault/internal/vector" + vectordocument "go.kenn.io/msgvault/internal/vector/document" +) + +func TestDocumentVectorLedgerCommandsNeverOpenRuntime(t *testing.T) { + fixture, spec := documentVectorCommandFixture(t) + t.Setenv("SYNTHETIC_EMBEDDING_KEY", "secret-that-must-not-print") + runtimeCalls := 0 + deps := documentsCommandDeps{ + openStore: func() (*store.Store, func(), error) { return fixture.Store, func() {}, nil }, + runDocumentVector: func(context.Context, *store.Store, int64, int) (vectordocument.ReconcileResult, error) { + runtimeCalls++ + return vectordocument.ReconcileResult{}, nil + }, + } + + unconfirmed := newDocumentsCmd(deps) + var unconfirmedOutput bytes.Buffer + unconfirmed.SetOut(&unconfirmedOutput) + unconfirmed.SetArgs([]string{documentVectorsSubcommand, "consent"}) + require.ErrorContains(t, unconfirmed.ExecuteContext(t.Context()), "--yes") + assert.Contains(t, unconfirmedOutput.String(), "Hosted document embedding disclosure:") + assert.Contains(t, unconfirmedOutput.String(), "Destination: https://embeddings.example.test/v1") + assert.Contains(t, unconfirmedOutput.String(), "Authentication: environment variable SYNTHETIC_EMBEDDING_KEY") + assert.Contains(t, unconfirmedOutput.String(), "API format: openai") + assert.Contains(t, unconfirmedOutput.String(), "Model: embed-test") + assert.Contains(t, unconfirmedOutput.String(), "Dimension: 3") + assert.Contains(t, unconfirmedOutput.String(), "Maximum input: 4096 characters") + assert.Contains(t, unconfirmedOutput.String(), "Docbank-prepared normalized attachment document inputs will be sent") + assert.NotContains(t, unconfirmedOutput.String(), "Explicit semantic or hybrid document searches") + assert.NotContains(t, unconfirmedOutput.String(), "secret-that-must-not-print") + consentSpec, err := configuredDocumentVectorConsentSpec(spec) + require.NoError(t, err) + unconfirmedConsent, err := fixture.Store.GetDocumentVectorConsent(t.Context(), consentSpec.EgressFingerprint) + require.NoError(t, err) + assert.Nil(t, unconfirmedConsent) + + consent := newDocumentsCmd(deps) + var consentOutput bytes.Buffer + consent.SetOut(&consentOutput) + consent.SetArgs([]string{documentVectorsSubcommand, "consent", "--yes"}) + require.NoError(t, consent.ExecuteContext(t.Context())) + assert.Contains(t, consentOutput.String(), "Hosted document embedding disclosure:") + assert.Contains(t, consentOutput.String(), "Recorded consent for document vector egress fingerprint "+consentSpec.EgressFingerprint) + assert.NotContains(t, consentOutput.String(), "secret-that-must-not-print") + recorded, err := fixture.Store.GetDocumentVectorConsent(t.Context(), consentSpec.EgressFingerprint) + require.NoError(t, err) + require.NotNil(t, recorded) + assert.Equal(t, spec, recorded.DocumentVectorGenerationSpec) + assert.Equal(t, "document_embedding", recorded.Purpose) + + queryConsentSpec, err := configuredDocumentVectorQueryConsentSpec(spec) + require.NoError(t, err) + assert.NotEqual(t, consentSpec.EgressFingerprint, queryConsentSpec.EgressFingerprint) + queryConsent := newDocumentsCmd(deps) + var queryConsentOutput bytes.Buffer + queryConsent.SetOut(&queryConsentOutput) + queryConsent.SetArgs([]string{documentVectorsSubcommand, "consent", "--purpose", "queries", "--yes"}) + require.NoError(t, queryConsent.ExecuteContext(t.Context())) + assert.Contains(t, queryConsentOutput.String(), "Explicit semantic or hybrid document searches will send query text") + recordedQuery, err := fixture.Store.GetDocumentVectorConsent(t.Context(), queryConsentSpec.EgressFingerprint) + require.NoError(t, err) + require.NotNil(t, recordedQuery) + assert.Equal(t, "query_embedding", recordedQuery.Purpose) + + consentedEndpoint := cfg.Vector.Embeddings.Endpoint + cfg.Vector.Embeddings.Endpoint = "https://hosted.example.test/v1" + changedConsentSpec, err := configuredDocumentVectorConsentSpec(spec) + require.NoError(t, err) + assert.NotEqual(t, consentSpec.EgressFingerprint, changedConsentSpec.EgressFingerprint) + require.ErrorContains(t, requireDocumentVectorConsent(t.Context(), fixture.Store, spec), "not consented") + cfg.Vector.Embeddings.Endpoint = consentedEndpoint + require.NoError(t, requireDocumentVectorConsent(t.Context(), fixture.Store, spec)) + + generation, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(t, err) + status := newDocumentsCmd(deps) + var statusOutput bytes.Buffer + status.SetOut(&statusOutput) + status.SetArgs([]string{documentVectorsSubcommand, statusValue}) + require.NoError(t, status.ExecuteContext(t.Context())) + assert.Contains(t, statusOutput.String(), "building_generation=") + assert.Contains(t, statusOutput.String(), "state=building") + assert.Contains(t, statusOutput.String(), "pending=0 retryable=0 terminal=0 ready_live=0 obsolete=0 cleanup_pending=0") + assert.Contains(t, statusOutput.String(), "coverage_required=0 coverage_ready=0") + + retry := newDocumentsCmd(deps) + retry.SetOut(&bytes.Buffer{}) + retry.SetArgs([]string{documentVectorsSubcommand, "retry", "--generation-id", "999", "--limit", "1"}) + require.Error(t, retry.ExecuteContext(t.Context())) + + retire := newDocumentsCmd(deps) + retire.SetOut(&bytes.Buffer{}) + retire.SetArgs([]string{documentVectorsSubcommand, cliEmbeddingsOperationRetire, "--generation-id", fmtInt64(generation.ID), "--yes"}) + require.NoError(t, retire.ExecuteContext(t.Context())) + assert.Zero(t, runtimeCalls) +} + +func TestDocumentVectorStatusWorksWhenEmbeddingsAreDisabled(t *testing.T) { + markDaemonCLISubprocessForTest(t) + previous := cfg + t.Cleanup(func() { cfg = previous }) + cfg = config.NewDefaultConfig() + runtimeCalls := 0 + command := newDocumentsCmd(documentsCommandDeps{ + runDocumentVector: func(context.Context, *store.Store, int64, int) (vectordocument.ReconcileResult, error) { + runtimeCalls++ + return vectordocument.ReconcileResult{}, nil + }, + }) + var output bytes.Buffer + command.SetOut(&output) + command.SetArgs([]string{documentVectorsSubcommand, statusValue, "--json"}) + require.NoError(t, command.ExecuteContext(t.Context())) + assert.JSONEq(t, `{"enabled":false}`, output.String()) + assert.Zero(t, runtimeCalls) +} + +func TestConfiguredDocumentVectorSpecRejectsInvalidDisabledEmbeddingPolicy(t *testing.T) { + fixture, _ := documentVectorCommandFixture(t) + base := *cfg + tests := []struct { + name string + mutate func(*config.Config) + wantErr string + }{ + { + name: "endpoint", + mutate: func(c *config.Config) { + c.Vector.Embeddings.Endpoint = "::not-a-url" + }, + wantErr: "endpoint", + }, + { + name: "dimension", + mutate: func(c *config.Config) { + c.Vector.Embeddings.Dimension = 0 + }, + wantErr: "dimension", + }, + { + name: "input limit", + mutate: func(c *config.Config) { + c.Vector.Embeddings.MaxInputChars = -1 + }, + wantErr: "max_input_chars", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + changed := base + changed.Vector.Enabled = false + test.mutate(&changed) + cfg = &changed + var err error + require.NotPanics(t, func() { + _, err = configuredDocumentVectorSpec(t.Context(), fixture.Store) + }) + require.ErrorContains(t, err, test.wantErr) + }) + } +} + +func TestDocumentVectorStatusWorksBeforeExtractionTargetExists(t *testing.T) { + markDaemonCLISubprocessForTest(t) + previous := cfg + t.Cleanup(func() { cfg = previous }) + cfg = config.NewDefaultConfig() + cfg.Vector.Enabled = true + cfg.Vector.Embeddings.Endpoint = "https://embeddings.example.test/v1" + cfg.Vector.Embeddings.Model = "embed-test" + cfg.Vector.Embeddings.Dimension = 3 + cfg.Attachments.Documents.Index.Embeddings.Enabled = true + fixture := storetest.New(t) + command := newDocumentsCmd(documentsCommandDeps{ + openStore: func() (*store.Store, func(), error) { return fixture.Store, func() {}, nil }, + }) + var output bytes.Buffer + command.SetOut(&output) + command.SetArgs([]string{documentVectorsSubcommand, statusValue, "--json"}) + require.NoError(t, command.ExecuteContext(t.Context())) + assert.JSONEq(t, `{"enabled":true,"configured":false}`, output.String()) +} + +func TestDocumentVectorProviderCommandsUseRuntimeAndValidateBounds(t *testing.T) { + fixture, spec := documentVectorCommandFixture(t) + consentSpec, err := configuredDocumentVectorConsentSpec(spec) + require.NoError(t, err) + _, _, err = fixture.Store.RecordDocumentVectorConsent(t.Context(), consentSpec, time.Now()) + require.NoError(t, err) + var calls []int64 + deps := documentsCommandDeps{ + openStore: func() (*store.Store, func(), error) { return fixture.Store, func() {}, nil }, + runDocumentVector: func(_ context.Context, _ *store.Store, generationID int64, _ int) (vectordocument.ReconcileResult, error) { + calls = append(calls, generationID) + return vectordocument.ReconcileResult{}, nil + }, + } + + invalid := newDocumentsCmd(deps) + invalid.SetArgs([]string{documentVectorsSubcommand, documentBuildSubcommand, "--limit", "0"}) + require.ErrorContains(t, invalid.ExecuteContext(t.Context()), "limit") + assert.Empty(t, calls) + + build := newDocumentsCmd(deps) + build.SetOut(&bytes.Buffer{}) + build.SetArgs([]string{documentVectorsSubcommand, documentBuildSubcommand, "--limit", "1"}) + require.NoError(t, build.ExecuteContext(t.Context())) + require.Len(t, calls, 1) + building, err := fixture.Store.GetBuildingDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + require.NotNil(t, building) + + resume := newDocumentsCmd(deps) + resume.SetOut(&bytes.Buffer{}) + resume.SetArgs([]string{documentVectorsSubcommand, cmdUseResume, "--generation-id", fmtInt64(building.ID), "--limit", "1"}) + require.NoError(t, resume.ExecuteContext(t.Context())) + assert.Len(t, calls, 2) + + require.NoError(t, fixture.Store.ActivateDocumentVectorGeneration(t.Context(), building.ID, time.Now())) + rebuild := newDocumentsCmd(deps) + rebuild.SetOut(&bytes.Buffer{}) + rebuild.SetArgs([]string{documentVectorsSubcommand, "rebuild", "--generation-id", fmtInt64(building.ID), "--limit", "1", "--yes"}) + require.NoError(t, rebuild.ExecuteContext(t.Context())) + assert.Len(t, calls, 3) +} + +func TestDocumentVectorResumeRunsBoundedCleanupForRetiredGeneration(t *testing.T) { + fixture, spec := documentVectorCommandFixture(t) + generation, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(t, err) + token := strings.Repeat("8", 64) + _, err = fixture.Store.DB().Exec(fixture.Store.Rebind(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, token, state) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'ready')`), generation.ID, + "manual-cleanup-extraction", spec.TargetExtractionProfileID, strings.Repeat("a", 64), + "original", 1, "manual-cleanup-chunk", "manual-cleanup-checksum", 1, token) + require.NoError(t, err) + retired, err := fixture.Store.RetireDocumentVectorGeneration(t.Context(), generation.ID, time.Now()) + require.NoError(t, err) + require.True(t, retired) + + backend := &commandDocumentVectorBackend{} + deps := documentsCommandDeps{ + openStore: func() (*store.Store, func(), error) { return fixture.Store, func() {}, nil }, + runDocumentVector: func(ctx context.Context, st *store.Store, generationID int64, limit int) (vectordocument.ReconcileResult, error) { + return runDocumentVectorWithFeatures(ctx, st, &vectorFeatures{ + DocumentBackend: backend, + Cfg: cfg.Vector, + }, generationID, limit) + }, + } + command := newDocumentsCmd(deps) + var output bytes.Buffer + command.SetOut(&output) + command.SetArgs([]string{ + documentVectorsSubcommand, cmdUseResume, + "--generation-id", fmtInt64(generation.ID), "--limit", "1", + }) + + require.NoError(t, command.ExecuteContext(t.Context())) + var result vectordocument.ReconcileResult + require.NoError(t, json.Unmarshal(output.Bytes(), &result)) + assert.True(t, result.Purged) + assert.True(t, result.Converged) + assert.Equal(t, [][]string{{token}}, backend.deletes) + _, err = fixture.Store.GetDocumentVectorGeneration(t.Context(), generation.ID) + require.ErrorContains(t, err, "not found") +} + +func TestDocumentVectorWorkerCheckpointsPartialErrorResult(t *testing.T) { + wantErr := errors.New("provider partial failure") + runResult := vectordocument.RunResult{ + ProviderCalls: 1, ProviderDocuments: 2, ProviderChunks: 3, ProviderInputChars: 44, + AfterGenerationID: 7, AfterChunkID: 81, + } + worker := checkpointingDocumentVectorWorker{ + worker: fakeDocumentVectorWorkerRunner{result: runResult, err: wantErr}, + checkpointer: &fakeDocumentVectorCheckpointer{}, + fingerprint: strings.Repeat("a", 64), + now: func() time.Time { return time.Date(2026, 8, 20, 12, 0, 0, 0, time.UTC) }, + } + result, err := worker.Run(t.Context(), 7, 10) + require.ErrorIs(t, err, wantErr) + assert.Equal(t, runResult, result) + checkpoint, ok := worker.checkpointer.(*fakeDocumentVectorCheckpointer) + require.True(t, ok) + require.Len(t, checkpoint.calls, 1) + assert.Equal(t, strings.Repeat("a", 64), checkpoint.calls[0].fingerprint) + assert.Equal(t, int64(81), checkpoint.calls[0].afterChunkID) + assert.False(t, checkpoint.calls[0].exhausted) + assert.Equal(t, store.DocumentVectorUsageDelta{ + ProviderCalls: 1, ProviderDocuments: 2, ProviderChunks: 3, ProviderInputChars: 44, + }, checkpoint.calls[0].delta) +} + +func TestDocumentVectorWorkerOwnersAreUniqueWithinOneProcess(t *testing.T) { + first := nextDocumentVectorWorkerOwner() + second := nextDocumentVectorWorkerOwner() + assert.NotEqual(t, first, second) + assert.True(t, strings.HasPrefix(first, "document-vector-")) + _, err := uuid.Parse(strings.TrimPrefix(first, "document-vector-")) + require.NoError(t, err) +} + +type fakeDocumentVectorWorkerRunner struct { + result vectordocument.RunResult + err error +} + +func (r fakeDocumentVectorWorkerRunner) Run(context.Context, vectordocument.GenerationID, int) (vectordocument.RunResult, error) { + return r.result, r.err +} + +type fakeDocumentVectorCheckpoint struct { + afterChunkID int64 + exhausted bool + fingerprint string + delta store.DocumentVectorUsageDelta +} + +type fakeDocumentVectorCheckpointer struct { + calls []fakeDocumentVectorCheckpoint +} + +func (c *fakeDocumentVectorCheckpointer) CheckpointDocumentVectorBuildForFingerprint(_ context.Context, _ int64, fingerprint string, afterChunkID int64, exhausted bool, delta store.DocumentVectorUsageDelta, _ time.Time) error { + c.calls = append(c.calls, fakeDocumentVectorCheckpoint{ + fingerprint: fingerprint, afterChunkID: afterChunkID, exhausted: exhausted, delta: delta, + }) + return nil +} + +func TestScheduledDocumentVectorRotationRetiresObsoleteBuildingBeforeDesiredBuild(t *testing.T) { + fixture, desired := documentVectorCommandFixture(t) + consentSpec, err := configuredDocumentVectorConsentSpec(desired) + require.NoError(t, err) + _, _, err = fixture.Store.RecordDocumentVectorConsent(t.Context(), consentSpec, time.Now()) + require.NoError(t, err) + activeSpec := desired + activeSpec.Fingerprint = strings.Repeat("1", 64) + active, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), activeSpec) + require.NoError(t, err) + require.NoError(t, fixture.Store.ActivateDocumentVectorGeneration(t.Context(), active.ID, time.Now())) + obsoleteSpec := desired + obsoleteSpec.Fingerprint = strings.Repeat("2", 64) + obsolete, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), obsoleteSpec) + require.NoError(t, err) + for index := range 3 { + _, err = fixture.Store.DB().Exec(fixture.Store.Rebind(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, token, state) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'ready')`), obsolete.ID, + fmt.Sprintf("obsolete-%d", index), desired.TargetExtractionProfileID, strings.Repeat("a", 64), + "original", index+1, fmt.Sprintf("chunk-%d", index), fmt.Sprintf("checksum-%d", index), + 1, fmt.Sprintf("%064x", index+1)) + require.NoError(t, err) + } + client := &commandDocumentSemanticClient{} + backend := &commandDocumentVectorBackend{} + vf := &vectorFeatures{ + DocumentBackend: backend, SemanticClient: client, Cfg: cfg.Vector, + } + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 2)) + stillActive, err := fixture.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + require.NotNil(t, stillActive) + assert.Equal(t, active.ID, stillActive.ID) + retired, err := fixture.Store.GetDocumentVectorGeneration(t.Context(), obsolete.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorGenerationRetired, retired.State) + require.Len(t, backend.deletes, 1) + assert.Len(t, backend.deletes[0], 2) + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 2)) + _, err = fixture.Store.GetDocumentVectorGeneration(t.Context(), obsolete.ID) + require.ErrorContains(t, err, "not found") + require.Len(t, backend.deletes, 2) + assert.Len(t, backend.deletes[1], 1) + stillActive, err = fixture.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + assert.Equal(t, active.ID, stillActive.ID) + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 2)) + newActive, err := fixture.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + require.NotNil(t, newActive) + assert.Equal(t, desired, newActive.DocumentVectorGenerationSpec) + assert.Zero(t, client.documentCalls) +} + +func TestScheduledDocumentVectorCleansRetiredWithoutConsentOrProvider(t *testing.T) { + fixture, spec := documentVectorCommandFixture(t) + generation, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(t, err) + token := strings.Repeat("9", 64) + _, err = fixture.Store.DB().Exec(fixture.Store.Rebind(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, token, state) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'ready')`), generation.ID, + "retired-extraction", spec.TargetExtractionProfileID, strings.Repeat("a", 64), + "original", 1, "retired-chunk", "retired-checksum", 1, token) + require.NoError(t, err) + retired, err := fixture.Store.RetireDocumentVectorGeneration(t.Context(), generation.ID, time.Now()) + require.NoError(t, err) + require.True(t, retired) + backend := &commandDocumentVectorBackend{} + vf := &vectorFeatures{DocumentBackend: backend, Cfg: cfg.Vector} + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 10)) + + assert.Equal(t, [][]string{{token}}, backend.deletes) + _, err = fixture.Store.GetDocumentVectorGeneration(t.Context(), generation.ID) + require.ErrorContains(t, err, "not found") + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 10)) +} + +func TestScheduledDocumentVectorObservesConsentRecordedAfterRuntimeInitialization(t *testing.T) { + fixture, spec := documentVectorCommandFixture(t) + vf := &vectorFeatures{ + DocumentBackend: &commandDocumentVectorBackend{}, + SemanticClient: &commandDocumentSemanticClient{}, + Cfg: cfg.Vector, + } + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 10)) + active, err := fixture.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + assert.Nil(t, active) + + consentSpec, err := configuredDocumentVectorConsentSpec(spec) + require.NoError(t, err) + _, _, err = fixture.Store.RecordDocumentVectorConsent(t.Context(), consentSpec, time.Now()) + require.NoError(t, err) + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 10)) + active, err = fixture.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + require.NotNil(t, active) + assert.Equal(t, spec, active.DocumentVectorGenerationSpec) +} + +func TestScheduledDocumentVectorCleansObsoleteActiveTokensAfterCoverageIsComplete(t *testing.T) { + fixture, desired := documentVectorCommandFixture(t) + consentSpec, err := configuredDocumentVectorConsentSpec(desired) + require.NoError(t, err) + _, _, err = fixture.Store.RecordDocumentVectorConsent(t.Context(), consentSpec, time.Now()) + require.NoError(t, err) + active, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), desired) + require.NoError(t, err) + require.NoError(t, fixture.Store.ActivateDocumentVectorGeneration(t.Context(), active.ID, time.Now())) + token := strings.Repeat("c", 64) + _, err = fixture.Store.DB().Exec(fixture.Store.Rebind(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, token, state) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'ready')`), active.ID, + "deleted-extraction", desired.TargetExtractionProfileID, strings.Repeat("d", 64), + "original", 1, "deleted-chunk", "deleted-checksum", 1, token) + require.NoError(t, err) + coverage, err := fixture.Store.GetDocumentVectorCoverage(t.Context(), active.ID) + require.NoError(t, err) + assert.True(t, coverage.Complete()) + status, err := fixture.Store.GetDocumentVectorGenerationStatus(t.Context(), active.ID, "", 10) + require.NoError(t, err) + assert.Equal(t, int64(1), status.CleanupPending) + + client := &commandDocumentSemanticClient{} + backend := &commandDocumentVectorBackend{} + vf := &vectorFeatures{ + DocumentBackend: backend, SemanticClient: client, Cfg: cfg.Vector, + } + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 10)) + require.Equal(t, [][]string{{token}}, backend.deletes) + assert.Zero(t, client.documentCalls) + status, err = fixture.Store.GetDocumentVectorGenerationStatus(t.Context(), active.ID, "", 10) + require.NoError(t, err) + assert.Zero(t, status.CleanupPending) + var publications int + require.NoError(t, fixture.Store.DB().QueryRow(fixture.Store.Rebind( + `SELECT COUNT(*) FROM document_vector_publications WHERE generation_id = ?`), active.ID).Scan(&publications)) + assert.Zero(t, publications) + + require.NoError(t, runScheduledDocumentVectorGeneration(t.Context(), fixture.Store, vf, 10)) + assert.Len(t, backend.deletes, 1, "a converged replay does not re-delete finalized tokens") + assert.Zero(t, client.documentCalls) +} + +type commandDocumentSemanticClient struct{ documentCalls int } + +func (*commandDocumentSemanticClient) EmbedQuery(context.Context, string) ([]float32, error) { + return []float32{1, 0, 0}, nil +} + +func (c *commandDocumentSemanticClient) EmbedDocuments(context.Context, []vector.DocumentInput) ([][][]float32, error) { + c.documentCalls++ + return nil, nil +} + +type commandDocumentVectorBackend struct{ deletes [][]string } + +func (*commandDocumentVectorBackend) PutUnpublished(context.Context, vectordocument.GenerationID, int, []vectordocument.Embedding) error { + return nil +} + +func (b *commandDocumentVectorBackend) DeleteTokens(_ context.Context, _ vectordocument.GenerationID, tokens []string) error { + b.deletes = append(b.deletes, append([]string(nil), tokens...)) + return nil +} + +func (*commandDocumentVectorBackend) Search(context.Context, vectordocument.GenerationID, int, []float32, int) ([]vectordocument.Hit, error) { + return nil, nil +} + +func documentVectorCommandFixture(t *testing.T) (*storetest.Fixture, store.DocumentVectorGenerationSpec) { + t.Helper() + markDaemonCLISubprocessForTest(t) + previous := cfg + t.Cleanup(func() { cfg = previous }) + c := config.NewDefaultConfig() + c.Vector.Enabled = true + c.Vector.Embeddings.Endpoint = "https://embeddings.example.test/v1" + c.Vector.Embeddings.APIKeyEnv = "SYNTHETIC_EMBEDDING_KEY" + c.Vector.Embeddings.Model = "embed-test" + c.Vector.Embeddings.Dimension = 3 + c.Vector.Embeddings.MaxInputChars = 4096 + c.Attachments.Documents.Index.Embeddings.Enabled = true + c.Attachments.Documents.Index.Embeddings.Profile = "vector.embeddings" + cfg = c + fixture := storetest.New(t) + fingerprint := strings.Repeat("7", 64) + profile := store.DocumentExtractionProfile{ + ID: "profile-" + fingerprint, Fingerprint: fingerprint, Provider: "synthetic", + Endpoint: "https://documents.example.test/v1", Region: localValue, Model: "extract-test", + RetentionPosture: "standard", TrainingPosture: "opted-out", + AllowedMediaTypes: []string{"application/pdf"}, PolicyJSON: []byte(`{"policy":1}`), + } + _, err := fixture.Store.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(t, err) + _, err = fixture.Store.DB().Exec(fixture.Store.Rebind(`UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), profile.ID) + require.NoError(t, err) + spec, err := desiredDocumentVectorSpec(t.Context(), fixture.Store) + require.NoError(t, err) + return fixture, spec +} + +func fmtInt64(value int64) string { return strconv.FormatInt(value, 10) } diff --git a/cmd/msgvault/cmd/embed_vector_test.go b/cmd/msgvault/cmd/embed_vector_test.go index 1dcaee372..f96fcdad0 100644 --- a/cmd/msgvault/cmd/embed_vector_test.go +++ b/cmd/msgvault/cmd/embed_vector_test.go @@ -8,9 +8,12 @@ import ( "database/sql" "errors" "fmt" + "net/http" + "net/http/httptest" "os" "path/filepath" "strings" + "sync/atomic" "testing" "time" @@ -282,7 +285,9 @@ func TestActivateBuiltGeneration_ContextualLifecycleErrorsDoNotActivateAnotherGe }) } -func setupVectorFeaturesFixture(t *testing.T, apiFormat vector.EmbeddingAPIFormat, readOnly bool) *vectorFeatures { +func setupVectorFeaturesFixture( + t *testing.T, apiFormat vector.EmbeddingAPIFormat, readOnly bool, mutate ...func(*config.Config), +) *vectorFeatures { t.Helper() dir := t.TempDir() mainPath := filepath.Join(dir, "msgvault.db") @@ -294,15 +299,41 @@ func setupVectorFeaturesFixture(t *testing.T, apiFormat vector.EmbeddingAPIForma c.Vector.Embeddings.Model = "text-embedding-test" c.Vector.Embeddings.Dimension = 4 c.Vector.Embeddings.APIFormat = apiFormat + c.Attachments.Documents.Index.Embeddings.Enabled = true + c.Attachments.Documents.Index.Embeddings.Profile = "vector.embeddings" if apiFormat == vector.APIFormatVoyageContextual { c.Vector.Embeddings.Model = "voyage-context-4" } + for _, apply := range mutate { + apply(c) + } withTestConfig(t, c) s, err := store.Open(mainPath) require.NoError(t, err) t.Cleanup(func() { _ = s.Close() }) require.NoError(t, s.InitSchema()) + fingerprint := strings.Repeat("d", 64) + profile := store.DocumentExtractionProfile{ + ID: "profile-" + fingerprint, Fingerprint: fingerprint, Provider: "synthetic", + Endpoint: "https://documents.example.test/v1", Region: localValue, Model: "extract-test", + RetentionPosture: "standard", TrainingPosture: "opted-out", + AllowedMediaTypes: []string{"application/pdf"}, PolicyJSON: []byte(`{"policy":1}`), + } + _, err = s.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(t, err) + _, err = s.DB().Exec(s.Rebind(`UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), profile.ID) + require.NoError(t, err) + spec, err := configuredDocumentVectorSpec(t.Context(), s) + require.NoError(t, err) + documentConsent, err := configuredDocumentVectorConsentSpec(spec) + require.NoError(t, err) + queryConsent, err := configuredDocumentVectorQueryConsentSpec(spec) + require.NoError(t, err) + for _, consentSpec := range []store.DocumentVectorConsentSpec{documentConsent, queryConsent} { + _, _, err = s.RecordDocumentVectorConsent(t.Context(), consentSpec, time.Now()) + require.NoError(t, err) + } vf, err := setupVectorFeatures(t.Context(), s, mainPath, readOnly) require.NoError(t, err) @@ -311,6 +342,35 @@ func setupVectorFeaturesFixture(t *testing.T, apiFormat vector.EmbeddingAPIForma return vf } +func TestSetupVectorFeaturesDocumentClientRejectsCrossOriginRedirects(t *testing.T) { + for _, apiFormat := range []vector.EmbeddingAPIFormat{vector.APIFormatOpenAI, vector.APIFormatVoyageContextual} { + t.Run(string(apiFormat), func(t *testing.T) { + var targetRequests atomic.Int32 + target := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + targetRequests.Add(1) + })) + t.Cleanup(target.Close) + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Location", target.URL+"/outside-consent") + w.WriteHeader(http.StatusTemporaryRedirect) + })) + t.Cleanup(origin.Close) + + vf := setupVectorFeaturesFixture(t, apiFormat, false, func(c *config.Config) { + c.Vector.Embeddings.Endpoint = origin.URL + }) + _, err := vf.SemanticClient.EmbedDocuments(t.Context(), []vector.DocumentInput{{ + Chunks: []string{"private attachment text"}, + }}) + + require.ErrorIs(t, err, embed.ErrEmbeddingProviderRedirect) + _, err = vf.DocumentQueryClient.EmbedQuery(t.Context(), "private query text") + require.ErrorIs(t, err, embed.ErrEmbeddingProviderRedirect) + assert.Zero(t, targetRequests.Load(), "attachment text must not reach the redirect target") + }) + } +} + func TestSetupVectorFeatures_SelectsRunnerByAPIFormat(t *testing.T) { t.Run("implicit OpenAI", func(t *testing.T) { vf := setupVectorFeaturesFixture(t, "", false) diff --git a/cmd/msgvault/cmd/import_imessage.go b/cmd/msgvault/cmd/import_imessage.go index f6f67de8f..dd50aa0c3 100644 --- a/cmd/msgvault/cmd/import_imessage.go +++ b/cmd/msgvault/cmd/import_imessage.go @@ -356,7 +356,7 @@ func resolveImessageSource(s *store.Store) (*store.Source, error) { if err == nil && len(sources) > 0 { var best *store.Source for _, src := range sources { - if src.Identifier != "local" { + if src.Identifier != localValue { if best == nil || src.ID > best.ID { best = src } @@ -367,7 +367,7 @@ func resolveImessageSource(s *store.Store) (*store.Source, error) { } return sources[0], nil } - return s.GetOrCreateSource("apple_messages", "local") + return s.GetOrCreateSource("apple_messages", localValue) } func init() { diff --git a/cmd/msgvault/cmd/import_imessage_test.go b/cmd/msgvault/cmd/import_imessage_test.go index 0a192ea93..71fe1a2cd 100644 --- a/cmd/msgvault/cmd/import_imessage_test.go +++ b/cmd/msgvault/cmd/import_imessage_test.go @@ -42,14 +42,14 @@ func TestResolveImessageSource(t *testing.T) { { name: "no existing sources — creates local", seedSources: nil, - wantIdentifier: "local", + wantIdentifier: localValue, }, { name: "only local exists — reuses local", seedSources: []struct{ sourceType, identifier string }{ - {"apple_messages", "local"}, + {"apple_messages", localValue}, }, - wantIdentifier: "local", + wantIdentifier: localValue, }, { name: "only legacy exists — reuses legacy", @@ -61,7 +61,7 @@ func TestResolveImessageSource(t *testing.T) { { name: "both legacy and local — prefers legacy", seedSources: []struct{ sourceType, identifier string }{ - {"apple_messages", "local"}, + {"apple_messages", localValue}, {"apple_messages", "+15551234567"}, }, wantIdentifier: "+15551234567", @@ -69,7 +69,7 @@ func TestResolveImessageSource(t *testing.T) { { name: "multiple legacy — picks first non-local", seedSources: []struct{ sourceType, identifier string }{ - {"apple_messages", "local"}, + {"apple_messages", localValue}, {"apple_messages", "alice@icloud.com"}, {"apple_messages", "+15551234567"}, }, diff --git a/cmd/msgvault/cmd/person_provider.go b/cmd/msgvault/cmd/person_provider.go index 0a4483269..41a7a7798 100644 --- a/cmd/msgvault/cmd/person_provider.go +++ b/cmd/msgvault/cmd/person_provider.go @@ -165,7 +165,7 @@ func newPersonProviderConsentCommand(deps personProviderCommandDeps) *cobra.Comm var jsonOutput bool var semanticEmbeddings bool command := &cobra.Command{ - Use: "consent", + Use: cmdUseConsent, Short: "Consent to the exact people inference policy", Args: cobra.NoArgs, RunE: func(command *cobra.Command, args []string) error { diff --git a/cmd/msgvault/cmd/root.go b/cmd/msgvault/cmd/root.go index b57ca7d6c..bc8aa5a60 100644 --- a/cmd/msgvault/cmd/root.go +++ b/cmd/msgvault/cmd/root.go @@ -654,7 +654,7 @@ func init() { rootCmd.PersistentFlags().StringVar(&cfgFile, "config", "", "config file (default: ~/.msgvault/config.toml)") rootCmd.PersistentFlags().StringVar(&homeDir, "home", "", "home directory (overrides MSGVAULT_HOME)") rootCmd.PersistentFlags().BoolVarP(&verbose, "verbose", "v", false, "verbose output (implies --log-level=debug)") - rootCmd.PersistentFlags().BoolVar(&useLocal, "local", false, "use local daemon instead of configured remote") + rootCmd.PersistentFlags().BoolVar(&useLocal, localValue, false, "use local daemon instead of configured remote") rootCmd.PersistentFlags().StringVar(&logFile, "log-file", "", "override log file path (default: /logs/msgvault-YYYY-MM-DD.log)") rootCmd.PersistentFlags().StringVar(&logLevel, "log-level", "", diff --git a/cmd/msgvault/cmd/serve.go b/cmd/msgvault/cmd/serve.go index 2afbaf14f..4581bd12b 100644 --- a/cmd/msgvault/cmd/serve.go +++ b/cmd/msgvault/cmd/serve.go @@ -1163,6 +1163,7 @@ var _ api.ChangedMessageLister = (*storeAPIAdapter)(nil) var _ api.ArchiveIdentifier = (*storeAPIAdapter)(nil) var _ api.DocumentSearchStore = (*storeAPIAdapter)(nil) var _ api.DocumentStatusStore = (*storeAPIAdapter)(nil) +var _ api.DocumentVectorStatusStore = (*storeAPIAdapter)(nil) var _ api.ActivityStore = (*storeAPIAdapter)(nil) func (a *storeAPIAdapter) ContactStateContext( @@ -1282,6 +1283,24 @@ func (a *storeAPIAdapter) CountIncompleteDocumentExtractionRebuild( ) } +func (a *storeAPIAdapter) GetDocumentVectorTargetProfileID(ctx context.Context) (string, error) { + return a.store.GetDocumentVectorTargetProfileID(ctx) +} + +func (a *storeAPIAdapter) GetDocumentVectorOperationsStatus( + ctx context.Context, + configured store.DocumentVectorGenerationSpec, + documentEgressFingerprint, queryEgressFingerprint string, + generationID int64, + afterToken string, + limit int, +) (store.DocumentVectorOperationsStatus, error) { + return a.store.GetDocumentVectorOperationsStatus( + ctx, configured, documentEgressFingerprint, queryEgressFingerprint, + generationID, afterToken, limit, + ) +} + func (a *storeAPIAdapter) GetStats() (*api.StoreStats, error) { return a.store.GetStats() } diff --git a/cmd/msgvault/cmd/serve_vector.go b/cmd/msgvault/cmd/serve_vector.go index 40863aabb..e087e01e0 100644 --- a/cmd/msgvault/cmd/serve_vector.go +++ b/cmd/msgvault/cmd/serve_vector.go @@ -19,6 +19,7 @@ import ( "go.kenn.io/msgvault/internal/scheduler" "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/vector" + vectordocument "go.kenn.io/msgvault/internal/vector/document" "go.kenn.io/msgvault/internal/vector/embed" "go.kenn.io/msgvault/internal/vector/hybrid" "go.kenn.io/msgvault/internal/vector/personsearch" @@ -30,11 +31,13 @@ import ( const contextualDocumentUTF8Limit = 100_000 type embeddingRuntime struct { - Runner scheduler.EmbedRunner - QueryClient hybrid.EmbeddingClient - PersonQueryClient personsearch.QueryEmbedder - Convergence scheduler.ConvergenceChecker - PersonGate vector.SemanticPersonEmbeddingGate + Runner scheduler.EmbedRunner + QueryClient hybrid.EmbeddingClient + PersonQueryClient personsearch.QueryEmbedder + Convergence scheduler.ConvergenceChecker + PersonGate vector.SemanticPersonEmbeddingGate + SemanticClient embed.SemanticClient + QuerySemanticClient embed.SemanticClient } type embeddingRuntimeDeps struct { @@ -48,6 +51,8 @@ type embeddingRuntimeDeps struct { Progress func(embed.ProgressReport) Log *slog.Logger PersonGate vector.SemanticPersonEmbeddingGate + DocumentGate embed.BeforeRequestFunc + QueryGate embed.BeforeRequestFunc } type legacyConvergenceChecker struct { @@ -245,6 +250,14 @@ func newEmbeddingRuntime(vectorCfg vector.Config, deps embeddingRuntimeDeps) (*e Timeout: vectorCfg.Embeddings.Timeout, MaxRetries: vectorCfg.Embeddings.MaxRetries, } messageClient := embed.NewClient(clientConfig) + documentClientConfig := clientConfig + documentClientConfig.BeforeRequest = deps.DocumentGate + documentClientConfig.RejectRedirects = true + documentClient := embed.NewClient(documentClientConfig) + queryClientConfig := clientConfig + queryClientConfig.BeforeRequest = deps.QueryGate + queryClientConfig.RejectRedirects = true + queryClient := embed.NewClient(queryClientConfig) clientConfig.BeforeRequest = personGate.Check personClient := embed.NewClient(clientConfig) messageWorker := embed.NewWorker(embed.WorkerDeps{ @@ -263,7 +276,8 @@ func newEmbeddingRuntime(vectorCfg vector.Config, deps embeddingRuntimeDeps) (*e worker := embed.NewGenerationWorker(messageWorker, personWorker) return &embeddingRuntime{ Runner: worker, QueryClient: messageClient, PersonQueryClient: personClient, - Convergence: checker, PersonGate: personGate, + Convergence: checker, PersonGate: personGate, SemanticClient: documentClient, + QuerySemanticClient: queryClient, }, nil case vector.APIFormatVoyageContextual: if vectorCfg.Embeddings.Model != "voyage-context-4" { @@ -282,6 +296,14 @@ func newEmbeddingRuntime(vectorCfg vector.Config, deps embeddingRuntimeDeps) (*e MaxChunks: 16_000, MaxUTF8Bytes: contextualDocumentUTF8Limit}, } messageClient := embed.NewVoyageClient(clientConfig) + documentClientConfig := clientConfig + documentClientConfig.BeforeRequest = deps.DocumentGate + documentClientConfig.RejectRedirects = true + documentClient := embed.NewVoyageClient(documentClientConfig) + queryClientConfig := clientConfig + queryClientConfig.BeforeRequest = deps.QueryGate + queryClientConfig.RejectRedirects = true + queryClient := embed.NewVoyageClient(queryClientConfig) clientConfig.BeforeRequest = personGate.Check personClient := embed.NewVoyageClient(clientConfig) policy := embed.AssemblyPolicy{ @@ -304,7 +326,8 @@ func newEmbeddingRuntime(vectorCfg vector.Config, deps embeddingRuntimeDeps) (*e worker := embed.NewGenerationWorker(messageWorker, personWorker) return &embeddingRuntime{ Runner: worker, QueryClient: messageClient, PersonQueryClient: personClient, - Convergence: checker, PersonGate: personGate, + Convergence: checker, PersonGate: personGate, SemanticClient: documentClient, + QuerySemanticClient: queryClient, }, nil default: return nil, fmt.Errorf("unsupported embedding api format %q", vectorCfg.Embeddings.APIFormat) @@ -439,9 +462,10 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s } var ( - backend vector.Backend - vectorsDB *sql.DB - closeFn func() error + backend vector.Backend + documentBackend vectordocument.Backend + vectorsDB *sql.DB + closeFn func() error ) if store.IsPostgresURL(mainPath) { // Same database handle as the main store: pgvector embeddings @@ -466,6 +490,7 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s return nil, fmt.Errorf("open pgvector backend: %w", err) } backend = pgb + documentBackend = pgb.DocumentBackend() vectorsDB = pgb.DB() closeFn = pgb.Close } else { @@ -492,11 +517,14 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s return nil, fmt.Errorf("open vectors.db: %w", err) } backend = sb + documentBackend = sb.DocumentBackend() vectorsDB = sb.DB() closeFn = sb.Close } - features := &vectorFeatures{Backend: backend, Cfg: vecCfg, Close: closeFn} + features := &vectorFeatures{ + Backend: backend, DocumentBackend: documentBackend, Cfg: vecCfg, Close: closeFn, + } if vecCfg.Enabled { personGate := vector.NewPinnedExactSemanticPersonEmbeddingGate( vecCfg, currentSemanticPersonVectorConfigSource(), mainStore, @@ -504,7 +532,9 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s runtime, err := newEmbeddingRuntime(vecCfg, embeddingRuntimeDeps{ Backend: backend, VectorsDB: vectorsDB, MainDB: mainDB, Store: mainStore, Rebind: dialect.Rebind, LastModifiedExpr: lastModifiedExpr, Log: logger, - PersonGate: personGate, + PersonGate: personGate, + DocumentGate: documentVectorRequestGate(mainStore, vecCfg, "document_embedding"), + QueryGate: documentVectorRequestGate(mainStore, vecCfg, "query_embedding"), }) if err != nil { _ = closeFn() @@ -512,6 +542,8 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s } features.Runner = runtime.Runner features.Convergence = runtime.Convergence + features.SemanticClient = runtime.SemanticClient + features.DocumentQueryClient = runtime.QuerySemanticClient features.HybridEngine = hybrid.NewEngine(backend, mainDB, runtime.QueryClient, hybrid.Config{ ExpectedFingerprint: vecCfg.GenerationFingerprint(), RRFK: vecCfg.Search.RRFK, @@ -542,6 +574,45 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s Gate: runtime.PersonGate, }, ) + if cfg.Attachments.Documents.Index.Embeddings.Enabled { + target, targetErr := mainStore.GetDocumentVectorTargetProfileID(ctx) + if targetErr != nil && !errors.Is(targetErr, store.ErrDocumentVectorInvalidGenerationState) { + _ = closeFn() + return nil, fmt.Errorf("read document vector target profile: %w", targetErr) + } + if targetErr == nil { + generationFingerprint, fingerprintErr := vectordocument.Fingerprint(target, vecCfg) + if fingerprintErr != nil { + _ = closeFn() + return nil, fingerprintErr + } + desired := store.DocumentVectorGenerationSpec{ + Fingerprint: generationFingerprint, + TargetExtractionProfileID: target, + EmbeddingProfile: "vector.embeddings", + Model: vecCfg.Embeddings.Model, + Dimension: vecCfg.Embeddings.Dimension, + } + queryEgressFingerprint, queryErr := vectordocument.QueryEgressFingerprint(target, vecCfg) + if queryErr != nil { + _ = closeFn() + return nil, queryErr + } + queryConsent, queryErr := mainStore.GetDocumentVectorConsent(ctx, queryEgressFingerprint) + if queryErr != nil { + _ = closeFn() + return nil, fmt.Errorf("read document query consent: %w", queryErr) + } + if queryConsent != nil && queryConsent.DocumentVectorGenerationSpec == desired && + queryConsent.EgressFingerprint == queryEgressFingerprint && + queryConsent.Purpose == "query_embedding" { + features.DocumentSearch = vectordocument.NewSearchService(vectordocument.SearchDeps{ + Ledger: mainStore, Embedder: runtime.QuerySemanticClient, Backend: documentBackend, + ExpectedFingerprint: desired.Fingerprint, + }) + } + } + } } if vecCfg.Multimodal.Enabled && !readOnly { if len(openers) == 0 || openers[0] == nil { @@ -567,6 +638,45 @@ func setupVectorFeatures(ctx context.Context, mainStore *store.Store, mainPath s return features, nil } +func documentVectorRequestGate(st *store.Store, vectorCfg vector.Config, purpose string) embed.BeforeRequestFunc { + return func(ctx context.Context) error { + target, err := st.GetDocumentVectorTargetProfileID(ctx) + if err != nil { + return err + } + generationFingerprint, err := vectordocument.Fingerprint(target, vectorCfg) + if err != nil { + return err + } + spec := store.DocumentVectorGenerationSpec{ + Fingerprint: generationFingerprint, TargetExtractionProfileID: target, + EmbeddingProfile: "vector.embeddings", Model: vectorCfg.Embeddings.Model, + Dimension: vectorCfg.Embeddings.Dimension, + } + var fingerprint string + switch purpose { + case "document_embedding": + fingerprint, err = vectordocument.EgressFingerprint(target, vectorCfg) + case "query_embedding": + fingerprint, err = vectordocument.QueryEgressFingerprint(target, vectorCfg) + default: + return errors.New("document vector request purpose is invalid") + } + if err != nil { + return err + } + consent, err := st.GetDocumentVectorConsent(ctx, fingerprint) + if err != nil { + return err + } + if consent == nil || consent.DocumentVectorGenerationSpec != spec || + consent.EgressFingerprint != fingerprint || consent.Purpose != purpose { + return errors.New("exact document vector egress is not consented") + } + return nil + } +} + func newVisualRuntime(ctx context.Context, vecCfg vector.Config, mainStore *store.Store, backend vector.Backend, opener visual.StreamOpener) (*visualFeatures, error) { fingerprint := vecCfg.MultimodalGenerationFingerprint() var visualBackend visual.Backend diff --git a/cmd/msgvault/cmd/serve_vector_init.go b/cmd/msgvault/cmd/serve_vector_init.go index 4d903ba3b..4055fbac9 100644 --- a/cmd/msgvault/cmd/serve_vector_init.go +++ b/cmd/msgvault/cmd/serve_vector_init.go @@ -141,6 +141,7 @@ func startVectorInit( apiServer.SetVectorFeatures( vf.HybridEngine, vf.PersonSearchEngine, vf.Backend, vf.Cfg, ) + apiServer.SetDocumentSearchService(vf.DocumentSearch) // Preflight drift detection: vector-search requests re-resolve the // durable account scope (throttled) so drift latches index_stale // even on daemons whose embed job never runs (empty cron, @@ -217,6 +218,9 @@ func startVectorInit( logger.Error("register multimodal job failed", "error", err) } } + if err := registerDocumentVectorJob(sched, vf, s); err != nil { + logger.Error("register document vector job failed", "error", err) + } logger.Info("daemon startup step complete", "step", "init_vector_backend") }() return h @@ -465,6 +469,32 @@ type embedJobRegistrar interface { SetEmbedJob(job *scheduler.EmbedJob, schedule string, runAfterSync bool) error } +type documentVectorJobRegistrar interface { + SetDocumentVectorJob(job func(context.Context) error, schedule string, runAfterSync bool) error +} + +func registerDocumentVectorJob(sched documentVectorJobRegistrar, vf *vectorFeatures, st *store.Store) error { + if vf == nil || vf.DocumentBackend == nil { + return nil + } + limit := vf.Cfg.Embeddings.BatchSize + if limit < 1 { + limit = defaultDocumentVectorOperationLimit + } + if limit > 1000 { + limit = 1000 + } + job := func(ctx context.Context) error { + return runScheduledDocumentVectorGeneration(ctx, st, vf, limit) + } + if err := sched.SetDocumentVectorJob(job, cfg.Vector.Embed.Schedule.Cron, cfg.Vector.Embed.Schedule.RunAfterSync); err != nil { + return fmt.Errorf("register document vector job: %w", err) + } + logger.Info("document vectors scheduled", "cron", cfg.Vector.Embed.Schedule.Cron, + "run_after_sync", cfg.Vector.Embed.Schedule.RunAfterSync) + return nil +} + func registerEmbedJob(sched embedJobRegistrar, vf *vectorFeatures, s *store.Store, apiServer *api.Server) error { embedJob := newSchedulerEmbedJob(vf, s) embedJob.ResolveBuildScope = func() (vector.BuildScope, error) { diff --git a/cmd/msgvault/cmd/serve_vector_init_test.go b/cmd/msgvault/cmd/serve_vector_init_test.go index 59d2e6a11..c015ca03a 100644 --- a/cmd/msgvault/cmd/serve_vector_init_test.go +++ b/cmd/msgvault/cmd/serve_vector_init_test.go @@ -18,6 +18,7 @@ import ( "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/testutil" "go.kenn.io/msgvault/internal/vector" + vectordocument "go.kenn.io/msgvault/internal/vector/document" "go.kenn.io/msgvault/internal/vector/embed" "go.kenn.io/msgvault/internal/vector/personsearch" "go.kenn.io/msgvault/internal/vector/visual" @@ -213,6 +214,148 @@ func TestStartVectorInitInstallsFeaturesOnSuccess(t *testing.T) { assertions.True(closed, "CloseFeatures must close the opened backend") } +type registeredDocumentVectorJobCapture struct { + calls int + job func(context.Context) error + schedule string + runAfterSync bool +} + +func (c *registeredDocumentVectorJobCapture) SetDocumentVectorJob( + job func(context.Context) error, schedule string, runAfterSync bool, +) error { + c.calls++ + c.job = job + c.schedule = schedule + c.runAfterSync = runAfterSync + return nil +} + +type startupDocumentSearchLedger struct{} + +func (startupDocumentSearchLedger) SearchDocuments(context.Context, store.DocumentSearchRequest) (store.DocumentSearchResponse, error) { + return store.DocumentSearchResponse{}, nil +} + +func (startupDocumentSearchLedger) GetDocumentIndexRevision(context.Context) (int64, error) { + return 1, nil +} + +func (startupDocumentSearchLedger) GetActiveDocumentVectorGeneration(context.Context) (*store.DocumentVectorGeneration, error) { + return &store.DocumentVectorGeneration{ + ID: 1, + Fingerprint: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + Dimension: 3, TargetExtractionProfileID: "profile-a", + }, nil +} + +func (startupDocumentSearchLedger) GetDocumentVectorTargetProfileID(context.Context) (string, error) { + return "profile-a", nil +} + +func (startupDocumentSearchLedger) ResolveDocumentVectorSearchOccurrences( + context.Context, int64, []store.DocumentVectorSearchHit, store.DocumentSearchRequest, int, +) ([]store.DocumentSearchResult, bool, error) { + return nil, false, nil +} + +type startupDocumentSemanticClient struct{} + +func (startupDocumentSemanticClient) EmbedQuery(context.Context, string) ([]float32, error) { + return []float32{1, 0, 0}, nil +} + +func (startupDocumentSemanticClient) EmbedDocuments(context.Context, []vector.DocumentInput) ([][][]float32, error) { + return nil, nil +} + +type startupDocumentBackend struct{} + +func (startupDocumentBackend) PutUnpublished(context.Context, vectordocument.GenerationID, int, []vectordocument.Embedding) error { + return nil +} + +func (startupDocumentBackend) DeleteTokens(context.Context, vectordocument.GenerationID, []string) error { + return nil +} + +func (startupDocumentBackend) Search(context.Context, vectordocument.GenerationID, int, []float32, int) ([]vectordocument.Hit, error) { + return nil, nil +} + +func (startupDocumentBackend) SearchPage(context.Context, vectordocument.GenerationID, int, []float32, string, int) (vectordocument.HitPage, error) { + return vectordocument.HitPage{Exhausted: true}, nil +} + +func TestRegisterDocumentVectorJobRequiresBackend(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + c := config.NewDefaultConfig() + c.Vector.Embed.Schedule.Cron = "*/7 * * * *" + c.Vector.Embed.Schedule.RunAfterSync = true + withTestConfig(t, c) + available := &vectorFeatures{DocumentBackend: startupDocumentBackend{}, Cfg: c.Vector} + + capture := ®isteredDocumentVectorJobCapture{} + requirements.NoError(registerDocumentVectorJob(capture, available, nil)) + assertions.Equal(1, capture.calls) + assertions.NotNil(capture.job) + assertions.Equal("*/7 * * * *", capture.schedule) + assertions.True(capture.runAfterSync) + + for name, features := range map[string]*vectorFeatures{ + "disabled": nil, + "nil backend": {Cfg: c.Vector}, + } { + t.Run(name, func(t *testing.T) { + unregistered := ®isteredDocumentVectorJobCapture{} + require.NoError(t, registerDocumentVectorJob(unregistered, features, nil)) + assert.Zero(t, unregistered.calls) + }) + } +} + +func TestStartVectorInitInstallsOnlyConsentedDocumentSearch(t *testing.T) { + for _, test := range []struct { + name string + service *vectordocument.SearchService + wantStatus int + wantPayload string + }{ + {name: "unconsented", wantStatus: http.StatusServiceUnavailable, wantPayload: "semantic_search_unavailable"}, + {name: "consented", service: vectordocument.NewSearchService(vectordocument.SearchDeps{ + Ledger: startupDocumentSearchLedger{}, Embedder: startupDocumentSemanticClient{}, Backend: startupDocumentBackend{}, + }), wantStatus: http.StatusOK, wantPayload: `"effective_mode":"semantic"`}, + } { + t.Run(test.name, func(t *testing.T) { + c := config.NewDefaultConfig() + c.Vector.Enabled = true + withTestConfig(t, c) + overrideSetupVectorFeatures(t, func(context.Context, *store.Store, string, bool) (*vectorFeatures, error) { + return &vectorFeatures{ + Backend: &fakeCmdVectorBackend{}, DocumentSearch: test.service, + DocumentBackend: startupDocumentBackend{}, SemanticClient: startupDocumentSemanticClient{}, + Cfg: c.Vector, Close: func() error { return nil }, + }, nil + }) + mainStore := testutil.NewTestStore(t) + srv := api.NewServerWithOptions(api.ServerOptions{ + Config: c, Store: &storeAPIAdapter{store: mainStore}, Logger: slog.New(slog.DiscardHandler), + VectorStatus: api.VectorStatusInitializing, + }) + h := startVectorInit(t.Context(), mainStore, "/tmp/msgvault.db", nil, srv, scheduler.New(nil)) + require.True(t, h.WaitTimeout(5*time.Second)) + + request := httptest.NewRequest(http.MethodGet, "/api/v1/documents/search?q=bounded&mode=semantic&candidate_limit=10", nil) + response := httptest.NewRecorder() + srv.Router().ServeHTTP(response, request) + assert.Equal(t, test.wantStatus, response.Code, response.Body.String()) + assert.Contains(t, response.Body.String(), test.wantPayload) + h.CloseFeatures() + }) + } +} + func TestStartVectorInitFlagsStaleIndex(t *testing.T) { c := config.NewDefaultConfig() c.Vector.Enabled = true diff --git a/cmd/msgvault/cmd/store_adapter_test.go b/cmd/msgvault/cmd/store_adapter_test.go index bb2cf0642..6e9a69347 100644 --- a/cmd/msgvault/cmd/store_adapter_test.go +++ b/cmd/msgvault/cmd/store_adapter_test.go @@ -45,6 +45,25 @@ func TestStoreAPIAdapterServesAttributeDefinitions(t *testing.T) { require.Equal(t, http.StatusOK, response.Code, response.Body.String()) } +func TestStoreAPIAdapterServesDocumentVectorStatus(t *testing.T) { + st := testutil.NewTestStore(t) + c := config.NewDefaultConfig() + c.Vector.Enabled = true + c.Attachments.Documents.Index.Embeddings.Enabled = true + srv := api.NewServerWithOptions(api.ServerOptions{ + Config: c, + Store: &storeAPIAdapter{store: st}, + Logger: slog.New(slog.DiscardHandler), + }) + + request := httptest.NewRequest(http.MethodGet, "/api/v1/documents/vectors/status", nil) + response := httptest.NewRecorder() + srv.Router().ServeHTTP(response, request) + + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + require.JSONEq(t, `{"enabled":true,"configured":false}`, response.Body.String()) +} + var _ api.CtxMessageStore = (*storeAPIAdapter)(nil) var _ api.MessageIdentityStore = (*storeAPIAdapter)(nil) var _ api.MeetingImporter = (*storeAPIAdapter)(nil) diff --git a/cmd/msgvault/cmd/tui.go b/cmd/msgvault/cmd/tui.go index f956b52ba..7a6559f0f 100644 --- a/cmd/msgvault/cmd/tui.go +++ b/cmd/msgvault/cmd/tui.go @@ -279,7 +279,7 @@ func openTUIBackend(ctx context.Context) (*tuiBackend, error) { func init() { rootCmd.AddCommand(tuiCmd) - tuiCmd.Flags().BoolVar(&forceLocalTUI, "local", false, "Use the local daemon instead of the configured remote server") + tuiCmd.Flags().BoolVar(&forceLocalTUI, localValue, false, "Use the local daemon instead of the configured remote server") tuiCmd.Flags().BoolVar(&deprecatedTUIForceSQL, "force-sql", false, "Deprecated in 0.17.0: set [analytics].engine = \"sql\" in config.toml") tuiCmd.Flags().BoolVar(&deprecatedTUISkipCacheBuild, "no-cache-build", false, "Deprecated in 0.17.0: set [analytics].auto_build_cache = false in config.toml") tuiCmd.Flags().BoolVar(&deprecatedTUINoSQLiteScanner, "no-sqlite-scanner", false, "Deprecated in 0.17.0: cache engine selection is daemon-managed") diff --git a/cmd/msgvault/cmd/vector_features.go b/cmd/msgvault/cmd/vector_features.go index b9472f41d..5ffa89ae2 100644 --- a/cmd/msgvault/cmd/vector_features.go +++ b/cmd/msgvault/cmd/vector_features.go @@ -6,6 +6,8 @@ import ( "go.kenn.io/msgvault/internal/scheduler" "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/vector" + vectordocument "go.kenn.io/msgvault/internal/vector/document" + "go.kenn.io/msgvault/internal/vector/embed" "go.kenn.io/msgvault/internal/vector/hybrid" "go.kenn.io/msgvault/internal/vector/personsearch" "go.kenn.io/msgvault/internal/vector/visual" @@ -42,11 +44,15 @@ type vectorFeatures struct { HybridEngine *hybrid.Engine // PersonSearchEngine shares Backend and Cfg with HybridEngine, but searches // only the person-owned corpus through a separately gated provider client. - PersonSearchEngine *personsearch.Engine - Runner scheduler.EmbedRunner - Convergence scheduler.ConvergenceChecker - Cfg vector.Config - Visual *visualFeatures + PersonSearchEngine *personsearch.Engine + Runner scheduler.EmbedRunner + Convergence scheduler.ConvergenceChecker + DocumentBackend vectordocument.Backend + DocumentSearch *vectordocument.SearchService + SemanticClient embed.SemanticClient + DocumentQueryClient embed.SemanticClient + Cfg vector.Config + Visual *visualFeatures // Close releases the backend's resources: on SQLite it closes the // vectors.db handle (so WAL checkpoints complete); on PostgreSQL it is // a no-op because the pgvector backend shares the main store's handle, diff --git a/go.mod b/go.mod index f3933071f..323ad18ee 100644 --- a/go.mod +++ b/go.mod @@ -39,7 +39,7 @@ require ( github.com/spf13/cobra v1.10.2 github.com/spf13/pflag v1.0.10 github.com/stretchr/testify v1.11.1 - go.kenn.io/docbank v0.13.1-0.20260820025223-11af138586a2 + go.kenn.io/docbank v0.14.0 go.kenn.io/kit v0.18.1 go.opentelemetry.io/otel v1.44.0 go.opentelemetry.io/otel/sdk v1.44.0 diff --git a/go.sum b/go.sum index 7a231f2ca..b011e175c 100644 --- a/go.sum +++ b/go.sum @@ -297,8 +297,8 @@ github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ= github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0= github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs= github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s= -go.kenn.io/docbank v0.13.1-0.20260820025223-11af138586a2 h1:WatIwVvxbCGe/bqbRCo+sLdgYcN7yHn7VIaKxlAU7L8= -go.kenn.io/docbank v0.13.1-0.20260820025223-11af138586a2/go.mod h1:FhKONSj2UYowwhgK2Lq3pIkYiIF00fSQt7PDBqqUbNc= +go.kenn.io/docbank v0.14.0 h1:qzkCnFIAKucpZIvblMouEV948MyjX+p49F+pqPtB5bk= +go.kenn.io/docbank v0.14.0/go.mod h1:sd4cATphlFCUU/Ea6VlB+LnFKFt18mbPvkDijhG42h4= go.kenn.io/kit v0.18.1 h1:dHEwtnXApwzbhqgvkj48Ylvdhes9QeWm3Q2wf9hq6Jc= go.kenn.io/kit v0.18.1/go.mod h1:SXD5SAxpYlH65wmBnfNqr5Za3NVFzKWJph+Sobx/ZMA= go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= diff --git a/internal/api/cli_allowlist_documents_test.go b/internal/api/cli_allowlist_documents_test.go index ce9f02b38..db265bae4 100644 --- a/internal/api/cli_allowlist_documents_test.go +++ b/internal/api/cli_allowlist_documents_test.go @@ -20,4 +20,16 @@ func TestCLIRunCommandAllowedDocumentMutations(t *testing.T) { assert.False(t, cliRunCommandAllowed([]string{"documents", subcommand})) }) } + + for _, subcommand := range []string{"build", "consent", "rebuild", "resume", "retire", "retry", "status"} { + t.Run("vectors/"+subcommand, func(t *testing.T) { + assert.True(t, cliRunCommandAllowed([]string{"documents", "vectors", subcommand})) + }) + } + for _, args := range [][]string{ + {"documents", "vectors"}, + {"documents", "vectors", "unknown"}, + } { + assert.False(t, cliRunCommandAllowed(args)) + } } diff --git a/internal/api/cli_handlers.go b/internal/api/cli_handlers.go index 25ef15879..1f6c04832 100644 --- a/internal/api/cli_handlers.go +++ b/internal/api/cli_handlers.go @@ -1426,6 +1426,17 @@ func cliRunCommandAllowed(args []string) bool { if len(args) < 2 { return false } + if args[1] == "vectors" { + if len(args) < 3 { + return false + } + switch args[2] { + case "build", "consent", "rebuild", "resume", "retire", "retry", "status": + return true + default: + return false + } + } switch args[1] { case "build", "consent-mistral", "purge-derived", "resume", "retire", "retry": return true diff --git a/internal/api/document_search.go b/internal/api/document_search.go index a404d9c83..d4a6f4e80 100644 --- a/internal/api/document_search.go +++ b/internal/api/document_search.go @@ -3,6 +3,7 @@ package api import ( "context" "errors" + "fmt" "net/http" "strings" @@ -10,6 +11,7 @@ import ( "go.kenn.io/msgvault/internal/personscope" "go.kenn.io/msgvault/internal/personscope/resolver" "go.kenn.io/msgvault/internal/store" + vectordocument "go.kenn.io/msgvault/internal/vector/document" ) const ( @@ -48,6 +50,18 @@ type DocumentStatusStore interface { ) (int64, error) } +type DocumentVectorStatusStore interface { + GetDocumentVectorTargetProfileID(ctx context.Context) (string, error) + GetDocumentVectorOperationsStatus(ctx context.Context, configured store.DocumentVectorGenerationSpec, documentEgressFingerprint, queryEgressFingerprint string, generationID int64, afterToken string, limit int) (store.DocumentVectorOperationsStatus, error) +} + +type DocumentVectorOperationsResponse struct { + Enabled bool `json:"enabled"` + Configured bool `json:"configured"` + ScheduledRegistrationRequiresRestart bool `json:"scheduled_registration_requires_restart,omitempty"` + Status *store.DocumentVectorOperationsStatus `json:"status,omitempty"` +} + type documentOccurrenceStatusReconciler interface { ReconcileDocumentOccurrences(ctx context.Context) error } @@ -70,6 +84,82 @@ func (s *Server) registerDocumentSearchRoute(api huma.API) { http.StatusBadRequest, http.StatusForbidden, http.StatusTooManyRequests, http.StatusServiceUnavailable, ) + registerAPIV1RawHumaJSONRouteWithErrors[DocumentVectorOperationsResponse]( + api, "getDocumentVectorStatus", http.MethodGet, "/documents/vectors/status", + "Get document vector generation, consent, usage, and failure status", + s.documentSearchGuard("document vector status", s.handleDocumentVectorStatus), + http.StatusBadRequest, http.StatusForbidden, http.StatusTooManyRequests, + http.StatusServiceUnavailable, + ) +} + +func (s *Server) handleDocumentVectorStatus(w http.ResponseWriter, r *http.Request) { + if !s.cfg.Vector.Enabled || !s.cfg.Attachments.Documents.Index.Embeddings.Enabled { + writeJSON(w, http.StatusOK, DocumentVectorOperationsResponse{Enabled: false}) + return + } + statusStore, ok := s.store.(DocumentVectorStatusStore) + if !ok { + writeError(w, http.StatusServiceUnavailable, "document_vector_status_unavailable", "Document vector status is unavailable") + return + } + generationID, _, err := queryInt64(r, "generation_id") + if err != nil { + s.rejectBadParam(w, err) + return + } + limit, found, err := queryInt(r, "limit") + if err != nil { + s.rejectBadParam(w, err) + return + } + if !found { + limit = 20 + } + if limit < 1 || limit > 1000 { + s.rejectBadParam(w, errors.New("limit must be between 1 and 1000")) + return + } + target, err := statusStore.GetDocumentVectorTargetProfileID(r.Context()) + if errors.Is(err, store.ErrDocumentVectorInvalidGenerationState) { + writeJSON(w, http.StatusOK, DocumentVectorOperationsResponse{Enabled: true, Configured: false}) + return + } + if err != nil { + s.writeDocumentSearchError(w, err) + return + } + generationFingerprint, err := vectordocument.Fingerprint(target, s.cfg.Vector) + if err != nil { + s.writeDocumentSearchError(w, err) + return + } + spec := store.DocumentVectorGenerationSpec{ + Fingerprint: generationFingerprint, TargetExtractionProfileID: target, + EmbeddingProfile: s.cfg.Attachments.Documents.Index.Embeddings.Profile, + Model: s.cfg.Vector.Embeddings.Model, Dimension: s.cfg.Vector.Embeddings.Dimension, + } + documentEgressFingerprint, err := vectordocument.EgressFingerprint(target, s.cfg.Vector) + if err != nil { + s.writeDocumentSearchError(w, err) + return + } + queryEgressFingerprint, err := vectordocument.QueryEgressFingerprint(target, s.cfg.Vector) + if err != nil { + s.writeDocumentSearchError(w, err) + return + } + status, err := statusStore.GetDocumentVectorOperationsStatus(r.Context(), spec, documentEgressFingerprint, queryEgressFingerprint, generationID, r.URL.Query().Get("after_token"), limit) + if err != nil { + s.writeDocumentSearchError(w, err) + return + } + s.documentSearchMu.RLock() + restartRequired := status.QueryConsent != nil && s.documentSearch == nil + s.documentSearchMu.RUnlock() + writeJSON(w, http.StatusOK, DocumentVectorOperationsResponse{ + Enabled: true, Configured: true, ScheduledRegistrationRequiresRestart: restartRequired, Status: &status, + }) } // documentSearchGuard protects the document reads that reconcile attachment @@ -189,7 +279,23 @@ func (s *Server) handleDocumentSearch(w http.ResponseWriter, r *http.Request) { } request.Person = &resolved.Scope } - response, err := searcher.SearchDocuments(r.Context(), request) + s.documentSearchMu.RLock() + service := s.documentSearch + s.documentSearchMu.RUnlock() + var response store.DocumentSearchResponse + if service != nil { + if reconciler, ok := s.store.(documentOccurrenceStatusReconciler); ok { + if err := reconciler.ReconcileDocumentOccurrences(r.Context()); err != nil { + s.writeDocumentSearchError(w, err) + return + } + } + response, err = service.Search(r.Context(), request) + } else if request.SearchMode == string(vectordocument.SearchModeSemantic) || request.SearchMode == string(vectordocument.SearchModeHybrid) { + err = vectordocument.ErrSemanticSearchUnavailable + } else { + response, err = searcher.SearchDocuments(r.Context(), request) + } if err != nil { s.writeDocumentSearchError(w, err) return @@ -211,6 +317,9 @@ func (s *Server) writeDocumentSearchError(w http.ResponseWriter, err error) { case errors.Is(err, store.ErrDocumentSearchUnavailable): writeError(w, http.StatusServiceUnavailable, "document_search_unavailable", "Document search requires full-text search support") + case errors.Is(err, vectordocument.ErrSemanticSearchUnavailable): + writeError(w, http.StatusServiceUnavailable, "semantic_search_unavailable", + "Semantic document search is unavailable") default: s.logger.Error("document search failed", "error", err) writeError(w, http.StatusInternalServerError, "internal_error", "Document search failed") @@ -219,8 +328,15 @@ func (s *Server) writeDocumentSearchError(w http.ResponseWriter, err error) { func parseDocumentSearchRequest(r *http.Request) (store.DocumentSearchRequest, error) { request := store.DocumentSearchRequest{ - Query: r.URL.Query().Get("q"), - Cursor: r.URL.Query().Get("cursor"), + Query: r.URL.Query().Get("q"), Cursor: r.URL.Query().Get("cursor"), + SearchMode: r.URL.Query().Get("mode"), + } + if request.SearchMode != "" { + mode, err := vectordocument.ParseSearchMode(request.SearchMode) + if err != nil { + return request, err + } + request.SearchMode = string(mode) } var err error request.SourceIDs, _, err = queryInt64s(r, "source_id") @@ -235,6 +351,18 @@ func parseDocumentSearchRequest(r *http.Request) (store.DocumentSearchRequest, e if request.PageSize, _, err = queryInt(r, "limit"); err != nil { return request, err } + var candidateFound bool + if request.CandidateLimit, candidateFound, err = queryInt(r, "candidate_limit"); err != nil { + return request, err + } + maxCandidateLimit := store.MaxLexicalDocumentSearchCandidateLimit + if request.SearchMode == string(vectordocument.SearchModeSemantic) || + request.SearchMode == string(vectordocument.SearchModeHybrid) { + maxCandidateLimit = store.MaxDocumentSearchCandidateLimit + } + if candidateFound && (request.CandidateLimit < 1 || request.CandidateLimit > maxCandidateLimit) { + return request, fmt.Errorf("candidate_limit must be between 1 and %d for this mode", maxCandidateLimit) + } if request.AttachmentID, _, err = queryInt64(r, "attachment_id"); err != nil { return request, err } @@ -284,3 +412,12 @@ func parseDocumentSearchRequest(r *http.Request) (store.DocumentSearchRequest, e } return request, nil } + +// SetDocumentSearchService installs semantic document retrieval after the +// optional vector runtime is ready. Before installation auto/lexical remain +// available and explicit semantic/hybrid requests return a stable 503. +func (s *Server) SetDocumentSearchService(service *vectordocument.SearchService) { + s.documentSearchMu.Lock() + s.documentSearch = service + s.documentSearchMu.Unlock() +} diff --git a/internal/api/document_search_test.go b/internal/api/document_search_test.go index 52ec9fc42..e0ba05ac2 100644 --- a/internal/api/document_search_test.go +++ b/internal/api/document_search_test.go @@ -4,8 +4,10 @@ import ( "context" "encoding/json" "log/slog" + "maps" "net/http" "net/http/httptest" + "slices" "testing" "time" @@ -31,6 +33,8 @@ func TestDocumentSearchHTTPPreservesDedicatedContract(t *testing.T) { assert.Equal(int64(41), request.AttachmentID) assert.Equal(int64(42), request.MessageID) assert.Equal(7, request.PageSize) + assert.Equal("lexical", request.SearchMode) + assert.Equal(1001, request.CandidateLimit) assert.Equal("opaque", request.Cursor) return store.DocumentSearchResponse{ Revision: 12, NextCursor: "next", @@ -41,7 +45,7 @@ func TestDocumentSearchHTTPPreservesDedicatedContract(t *testing.T) { }, nil } request := httptest.NewRequest(http.MethodGet, - "/api/v1/documents/search?q=shipping+damage&source_id=4,9&message_type=email,mms&attachment_id=41&message_id=42&limit=7&cursor=opaque", + "/api/v1/documents/search?q=shipping+damage&source_id=4,9&message_type=email,mms&attachment_id=41&message_id=42&limit=7&cursor=opaque&mode=lexical&candidate_limit=1001", nil) response := httptest.NewRecorder() server.Router().ServeHTTP(response, request) @@ -80,6 +84,21 @@ func TestDocumentSearchHTTPResolvesDurablePersonScope(t *testing.T) { require.Equal(t, http.StatusOK, response.Code, response.Body.String()) } +func TestDocumentSearchHTTPExplicitSemanticUnavailableDoesNotFallBack(t *testing.T) { + server, catalog := newTestServerWithMockStore(t) + calls := 0 + catalog.documentSearchFunc = func(context.Context, store.DocumentSearchRequest) (store.DocumentSearchResponse, error) { + calls++ + return store.DocumentSearchResponse{}, nil + } + request := httptest.NewRequest(http.MethodGet, "/api/v1/documents/search?q=evidence&mode=semantic&candidate_limit=25", nil) + response := httptest.NewRecorder() + server.Router().ServeHTTP(response, request) + assert.Equal(t, http.StatusServiceUnavailable, response.Code) + assert.Contains(t, response.Body.String(), "semantic_search_unavailable") + assert.Zero(t, calls) +} + func TestDocumentSearchHTTPMapsCursorRevisionConflict(t *testing.T) { server, catalog := newTestServerWithMockStore(t) catalog.documentSearchFunc = func( @@ -230,10 +249,37 @@ func TestOpenAPIDocumentSearchParameters(t *testing.T) { names = append(names, parameter.Name) } assert.ElementsMatch(t, - []string{"q", "source_id", "message_type", "attachment_id", "message_id", "person_id", "participant_id", "direction", "after", "before", "limit", "cursor"}, + []string{"q", "source_id", "message_type", "attachment_id", "message_id", "person_id", "participant_id", "direction", "after", "before", "limit", "cursor", "mode", "candidate_limit"}, names) } +func TestDocumentVectorStatusHTTPIsUsefulBeforeTargetConfiguration(t *testing.T) { + fixture := storetest.New(t) + c := config.NewDefaultConfig() + c.Vector.Enabled = true + c.Attachments.Documents.Index.Embeddings.Enabled = true + server := NewServerWithOptions(ServerOptions{Config: c, Store: fixture.Store, Logger: slog.New(slog.DiscardHandler)}) + request := httptest.NewRequest(http.MethodGet, "/api/v1/documents/vectors/status", nil) + response := httptest.NewRecorder() + server.Router().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + assert.JSONEq(t, `{"enabled":true,"configured":false}`, response.Body.String()) +} + +func TestOpenAPIDocumentVectorStatusUsesSnakeCaseCoverage(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + document := OpenAPIDocument() + operation := document.Paths["/api/v1/documents/vectors/status"].Get + requirements.NotNil(operation) + coverage := document.Components.Schemas.Map()["DocumentVectorCoverage"] + requirements.NotNil(coverage) + assertions.ElementsMatch([]string{"required", "ready"}, slices.Collect(maps.Keys(coverage.Properties))) + encoded, err := json.Marshal(store.DocumentVectorCoverage{Required: 4, Ready: 3}) + requirements.NoError(err) + assertions.JSONEq(`{"required":4,"ready":3}`, string(encoded)) +} + func TestDocumentStatusHTTPPreservesScopedContract(t *testing.T) { assert := assert.New(t) server, catalog := newTestServerWithMockStore(t) diff --git a/internal/api/routes.go b/internal/api/routes.go index 32ae829ec..39a91e938 100644 --- a/internal/api/routes.go +++ b/internal/api/routes.go @@ -632,6 +632,8 @@ func rawRouteParameters(operationID string) []*huma.Param { queryStringParam("before", "Only messages before an RFC3339 or YYYY-MM-DD date", false), queryIntegerParam(limitParam, "Maximum results to return (default 20, max 100)"), queryStringParam("cursor", "Opaque cursor from the previous document search page", false), + queryStringParam("mode", "Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider", false), + queryIntegerParam("candidate_limit", "Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000)"), } case "getDocumentIndexStatus": mediaTypes := queryRefArrayParam("media_type", "Allowed document media types") @@ -642,6 +644,12 @@ func rawRouteParameters(operationID string) []*huma.Param { mediaTypes, queryRefArrayParam("message_type", "Allowed message types"), } + case "getDocumentVectorStatus": + return []*huma.Param{ + queryIntegerParam("generation_id", "Generation whose bounded failures to inspect"), + queryStringParam("after_token", "Stable failure cursor token", false), + queryIntegerParam("limit", "Maximum failure diagnostics (default 20, max 1000)"), + } case "getCLIMessage", "getCLIMessageRaw": return []*huma.Param{queryStringParam("id", "Message numeric ID or source message ID", true)} case "getCLIAttachment": diff --git a/internal/api/server.go b/internal/api/server.go index 5073d2bb6..6ddd089a7 100644 --- a/internal/api/server.go +++ b/internal/api/server.go @@ -31,6 +31,7 @@ import ( "go.kenn.io/msgvault/internal/taskclient" "go.kenn.io/msgvault/internal/tasklinks" "go.kenn.io/msgvault/internal/vector" + vectordocument "go.kenn.io/msgvault/internal/vector/document" "go.kenn.io/msgvault/internal/vector/hybrid" "go.kenn.io/msgvault/internal/vector/visual" webapp "go.kenn.io/msgvault/internal/web" @@ -338,6 +339,8 @@ type Server struct { // the ordinary refresh would flip straight back to ready. Only a // successful reinit (SetVectorFeatures) clears the latch. vectorStaleLatch bool + documentSearchMu sync.RWMutex + documentSearch *vectordocument.SearchService // vectorScopeCheck re-resolves the durable embedding scope on the // vector-search preflight path (throttled by vectorScopeNextCheck) so // drift is detected even when no embed job ever runs. Wired by the diff --git a/internal/backupapp/restore.go b/internal/backupapp/restore.go new file mode 100644 index 000000000..ca7c4f649 --- /dev/null +++ b/internal/backupapp/restore.go @@ -0,0 +1,70 @@ +package backupapp + +import ( + "context" + "database/sql" + "errors" + "fmt" + + "go.kenn.io/kit/backup" + "go.kenn.io/msgvault/internal/sqliteutil" +) + +// InvalidateRestoredDocumentVectors removes document-vector authority because vectors.db +// is deliberately excluded from snapshots. Consent and historical provider +// usage remain valid archive records; the next vector run rebuilds the derived +// generation from normalized document evidence. +func InvalidateRestoredDocumentVectors(ctx context.Context, target backup.RestorePublicationTarget) (err error) { + if target.DBPath == "" { + return errors.New("backupapp: restored database path is required") + } + db, err := sql.Open(sqliteutil.DriverName(), target.DBPath) + if err != nil { + return fmt.Errorf("backupapp: open restored database: %w", err) + } + db.SetMaxOpenConns(1) + defer func() { + if closeErr := db.Close(); closeErr != nil { + err = errors.Join(err, fmt.Errorf("backupapp: close restored database: %w", closeErr)) + } + }() + for _, statement := range []string{ + `PRAGMA busy_timeout = 30000`, + `PRAGMA foreign_keys = ON`, + `PRAGMA journal_mode = DELETE`, + } { + if _, err := db.ExecContext(ctx, statement); err != nil { + return fmt.Errorf("backupapp: configure restored database: %w", err) + } + } + + var vectorAuthorityExists bool + if err := db.QueryRowContext(ctx, `SELECT EXISTS ( + SELECT 1 FROM sqlite_master + WHERE type = 'table' AND name = 'document_vector_generations' + )`).Scan(&vectorAuthorityExists); err != nil { + return fmt.Errorf("backupapp: inspect restored document vector authority: %w", err) + } + if !vectorAuthorityExists { + return nil + } + + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("backupapp: begin restored document vector invalidation: %w", err) + } + defer func() { _ = tx.Rollback() }() + for _, statement := range []string{ + `DELETE FROM document_vector_publications`, + `DELETE FROM document_vector_build_progress`, + `DELETE FROM document_vector_generations`, + } { + if _, err := tx.ExecContext(ctx, statement); err != nil { + return fmt.Errorf("backupapp: invalidate restored document vectors: %w", err) + } + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("backupapp: commit restored document vector invalidation: %w", err) + } + return nil +} diff --git a/internal/config/documentindex_config_test.go b/internal/config/documentindex_config_test.go index 74dd103f4..1b206e0cc 100644 --- a/internal/config/documentindex_config_test.go +++ b/internal/config/documentindex_config_test.go @@ -42,6 +42,10 @@ message_types = ["EMAIL", "chat", "email"] [attachments.documents.index] lexical = true store_chunk_text = true + +[attachments.documents.index.embeddings] +enabled = true +profile = "vector.embeddings" `) require.NoError(t, os.WriteFile(path, content, 0o600)) @@ -57,6 +61,38 @@ store_chunk_text = true assert.Equal([]string{"chat", "email"}, documents.Scope.MessageTypes) assert.True(documents.LexicalEnabled()) assert.True(documents.StoresChunkText()) + assert.True(documents.Index.Embeddings.Enabled) + assert.Equal("vector.embeddings", documents.Index.Embeddings.Profile) +} + +func TestLoadRejectsUnknownDocumentEmbeddingProfile(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "config.toml") + content := []byte(` +[attachments.documents.index.embeddings] +enabled = true +profile = "other.embeddings" +`) + require.NoError(t, os.WriteFile(path, content, 0o600)) + + _, err := Load(path, "") + assert.ErrorContains(t, err, "profile must be \"vector.embeddings\"") +} + +func TestLoadDefaultsOmittedEnabledDocumentEmbeddingProfile(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + dir := t.TempDir() + path := filepath.Join(dir, "config.toml") + requirements.NoError(os.WriteFile(path, []byte(` +[attachments.documents.index.embeddings] +enabled = true +`), 0o600)) + + loaded, err := Load(path, "") + requirements.NoError(err) + assertions.True(loaded.Attachments.Documents.Index.Embeddings.Enabled) + assertions.Equal("vector.embeddings", loaded.Attachments.Documents.Index.Embeddings.Profile) } func TestLoadRejectsInvalidDocumentAttachmentConfigWithoutResolvingKey(t *testing.T) { diff --git a/internal/daemonclient/documents.go b/internal/daemonclient/documents.go index 99ba19311..d3ca9e3c3 100644 --- a/internal/daemonclient/documents.go +++ b/internal/daemonclient/documents.go @@ -26,6 +26,7 @@ func (c *Client) SearchDocuments( After: optionalTimeRFC3339(request.After), Before: optionalTimeRFC3339(request.Before), Limit: optionalPositiveInt64(request.PageSize), Cursor: optionalString(request.Cursor), + Mode: optionalString(request.SearchMode), CandidateLimit: optionalPositiveInt64(request.CandidateLimit), }, }) }) @@ -66,8 +67,10 @@ func documentSearchFromGenerated(response *generated.DocumentSearchResponse) sto return store.DocumentSearchResponse{Results: []store.DocumentSearchResult{}} } result := store.DocumentSearchResponse{ - Revision: response.Revision, - Results: make([]store.DocumentSearchResult, len(response.Results)), + Revision: response.Revision, EffectiveMode: stringValue(response.EffectiveMode), + VectorGenerationID: int64Value(response.VectorGenerationID), + VectorGenerationFingerprint: stringValue(response.VectorGenerationFingerprint), + Results: make([]store.DocumentSearchResult, len(response.Results)), } if response.Truncated != nil { result.Truncated = *response.Truncated @@ -90,6 +93,12 @@ func documentSearchFromGenerated(response *generated.DocumentSearchResponse) sto ProfileID: row.ProfileID, ExtractionID: row.ExtractionID, Provider: row.Provider, Model: row.Model, MatchedSignals: row.MatchedSignals, Truncated: row.Truncated, Rank: int(row.Rank), + LexicalRank: int(int64Value(row.LexicalRank)), SemanticRank: int(int64Value(row.SemanticRank)), + SemanticScore: float64Value(row.SemanticScore), FusionScore: float64Value(row.FusionScore), + VectorToken: stringValue(row.VectorToken), VectorGenerationID: int64Value(row.VectorGenerationID), + VectorGenerationFingerprint: stringValue(row.VectorGenerationFingerprint), + VectorEmbeddingProfile: stringValue(row.VectorEmbeddingProfile), VectorModel: stringValue(row.VectorModel), + VectorDimension: int(int64Value(row.VectorDimension)), } if row.PersonProvenance != nil { result.Results[index].PersonProvenance = &personscope.Provenance{ @@ -119,6 +128,13 @@ func documentDirectionStrings(directions []personscope.Direction) []string { return result } +func float64Value(value *float64) float64 { + if value == nil { + return 0 + } + return *value +} + func documentIndexStatusFromGenerated( response *generated.DocumentIndexStatusResponse, ) store.DocumentIndexStatusResponse { diff --git a/internal/daemonclient/documents_test.go b/internal/daemonclient/documents_test.go index 10af68f78..a85371dd0 100644 --- a/internal/daemonclient/documents_test.go +++ b/internal/daemonclient/documents_test.go @@ -27,21 +27,28 @@ func TestSearchDocumentsUsesGeneratedDaemonContract(t *testing.T) { assert.Equal("2026-08-01T00:00:00Z", r.URL.Query().Get("after")) assert.Equal("2026-08-20T00:00:00Z", r.URL.Query().Get("before")) assert.Equal("5", r.URL.Query().Get("limit")) + assert.Equal("hybrid", r.URL.Query().Get("mode")) + assert.Equal("88", r.URL.Query().Get("candidate_limit")) w.Header().Set("Content-Type", "application/json") occurredAt := time.Date(2026, 8, 10, 0, 0, 0, 0, time.UTC) assert.NoError(json.NewEncoder(w).Encode(store.DocumentSearchResponse{ - Revision: 9, Truncated: true, + Revision: 9, Truncated: true, EffectiveMode: "hybrid", + VectorGenerationID: 31, VectorGenerationFingerprint: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", Results: []store.DocumentSearchResult{{ AttachmentID: 22, MessageID: 23, ConversationID: 24, SourceID: 3, SourceMessageID: "synthetic-message", OccurredAt: &occurredAt, OccurrenceKey: "occurrence", CanonicalBlobHash: "hash", ChunkKey: "chunk", Filename: "claim.docx", Excerpt: "damaged carton", ProfileID: "profile", ExtractionID: "extraction", Provider: "mistral", Model: "ocr", - MatchedSignals: []string{"content"}, Rank: 1, + MatchedSignals: []string{"content", "semantic"}, Rank: 1, PersonProvenance: &personscope.Provenance{ ParticipantIDs: []int64{4}, Roles: []personscope.Role{personscope.RoleFrom}, Directions: []personscope.Direction{personscope.FromPerson}, }, + LexicalRank: 2, SemanticRank: 1, SemanticScore: 0.875, FusionScore: 0.032, + VectorToken: "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", + VectorGenerationID: 31, VectorGenerationFingerprint: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + VectorEmbeddingProfile: "vector.embeddings", VectorModel: "embed-test", VectorDimension: 3, }}, })) })) @@ -54,19 +61,34 @@ func TestSearchDocumentsUsesGeneratedDaemonContract(t *testing.T) { Query: "damaged carton", SourceIDs: []int64{3, 7}, AttachmentID: 22, PageSize: 5, PersonID: 40, Directions: []personscope.Direction{personscope.FromPerson, personscope.Group}, After: &after, Before: &before, + SearchMode: "hybrid", CandidateLimit: 88, }) require.NoError(err) assert.Equal(int64(9), response.Revision) assert.True(response.Truncated) + assert.Equal("hybrid", response.EffectiveMode) + assert.Equal(int64(31), response.VectorGenerationID) + assert.Equal("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", response.VectorGenerationFingerprint) require.Len(response.Results, 1) - assert.Equal("claim.docx", response.Results[0].Filename) - assert.Equal(int64(24), response.Results[0].ConversationID) - assert.Equal("synthetic-message", response.Results[0].SourceMessageID) - require.NotNil(response.Results[0].OccurredAt) - assert.Equal(time.Date(2026, 8, 10, 0, 0, 0, 0, time.UTC), *response.Results[0].OccurredAt) - assert.Equal([]string{"content"}, response.Results[0].MatchedSignals) - require.NotNil(response.Results[0].PersonProvenance) - assert.Equal([]personscope.Role{personscope.RoleFrom}, response.Results[0].PersonProvenance.Roles) + row := response.Results[0] + assert.Equal("claim.docx", row.Filename) + assert.Equal(int64(24), row.ConversationID) + assert.Equal("synthetic-message", row.SourceMessageID) + require.NotNil(row.OccurredAt) + assert.Equal(time.Date(2026, 8, 10, 0, 0, 0, 0, time.UTC), *row.OccurredAt) + assert.Equal([]string{"content", "semantic"}, row.MatchedSignals) + assert.Equal(2, row.LexicalRank) + assert.Equal(1, row.SemanticRank) + assert.InDelta(0.875, row.SemanticScore, 0.0001) + assert.InDelta(0.032, row.FusionScore, 0.0001) + assert.Equal("bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", row.VectorToken) + assert.Equal(int64(31), row.VectorGenerationID) + assert.Equal("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", row.VectorGenerationFingerprint) + assert.Equal("vector.embeddings", row.VectorEmbeddingProfile) + assert.Equal("embed-test", row.VectorModel) + assert.Equal(3, row.VectorDimension) + require.NotNil(row.PersonProvenance) + assert.Equal([]personscope.Role{personscope.RoleFrom}, row.PersonProvenance.Roles) } func TestDocumentIndexStatusUsesGeneratedDaemonContract(t *testing.T) { diff --git a/internal/documentindex/config.go b/internal/documentindex/config.go index 2e7ee7447..746aafbb5 100644 --- a/internal/documentindex/config.go +++ b/internal/documentindex/config.go @@ -278,8 +278,8 @@ func (c *DocumentsConfig) Validate() error { if !c.LexicalEnabled() || !c.StoresChunkText() { return errors.New("attachments.documents.index.lexical and store_chunk_text must both be true until document vector publication lands") } - if c.Index.Embeddings.Enabled { - return errors.New("attachments.documents.index.embeddings.enabled is not available until document vector publication lands") + if c.Index.Embeddings.Enabled && c.Index.Embeddings.Profile != "vector.embeddings" { + return fmt.Errorf("attachments.documents.index.embeddings.profile must be %q when embeddings are enabled", "vector.embeddings") } if slices.Contains(c.Scope.MessageTypes, "") { return errors.New("attachments.documents.scope.message_types contains an empty value") diff --git a/internal/documentindex/config_test.go b/internal/documentindex/config_test.go index 686ceaa4f..4654f2a4a 100644 --- a/internal/documentindex/config_test.go +++ b/internal/documentindex/config_test.go @@ -51,6 +51,24 @@ func TestDocumentsConfigRejectsUnavailableIndexOptOut(t *testing.T) { assert.False(t, config.StoresChunkText()) } +func TestDocumentsConfigAllowsNamedEmbeddingProfileWithLexicalFallback(t *testing.T) { + config := DefaultDocumentsConfig() + config.Index.Embeddings.Enabled = true + config.Index.Embeddings.Profile = "vector.embeddings" + + require.NoError(t, config.Validate()) + assert.True(t, config.LexicalEnabled()) + assert.True(t, config.StoresChunkText()) +} + +func TestDocumentsConfigRejectsUnknownEmbeddingProfile(t *testing.T) { + config := DefaultDocumentsConfig() + config.Index.Embeddings.Enabled = true + config.Index.Embeddings.Profile = "other.embeddings" + + require.ErrorContains(t, config.Validate(), "profile must be \"vector.embeddings\"") +} + func TestDocumentsConfigRejectsUnsafePolicy(t *testing.T) { tests := []struct { name string @@ -74,7 +92,10 @@ func TestDocumentsConfigRejectsUnsafePolicy(t *testing.T) { {name: "training", mutate: func(c *DocumentsConfig) { c.TrainingPosture = "never" }, want: "training_posture"}, {name: "lexical without text", mutate: func(c *DocumentsConfig) { disabled := false; c.Index.StoreChunkText = &disabled }, want: "must both be true"}, {name: "stored text without lexical", mutate: func(c *DocumentsConfig) { disabled := false; c.Index.Lexical = &disabled }, want: "must both be true"}, - {name: "premature vectors", mutate: func(c *DocumentsConfig) { c.Index.Embeddings.Enabled = true }, want: "not available"}, + {name: "unknown embedding profile", mutate: func(c *DocumentsConfig) { + c.Index.Embeddings.Enabled = true + c.Index.Embeddings.Profile = "other.embeddings" + }, want: "profile must be \"vector.embeddings\""}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { @@ -155,6 +176,32 @@ func TestDocumentsProfileFingerprintIsDeterministicAndPolicyBound(t *testing.T) assert.NotEqual(first, seventh) } +func TestDocumentsProfileFingerprintExcludesEmbeddingOptIn(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + config := DefaultDocumentsConfig() + config.RetentionPosture = RetentionZDR + config.TrainingPosture = TrainingOptedOut + policy, err := config.MistralPolicy() + requirements.NoError(err) + manifest := testCapabilityManifest(t, policy) + + before, err := config.ProfileFingerprint(manifest, []string{"application/pdf"}) + requirements.NoError(err) + beforePolicy, err := config.ProfilePolicyJSON(manifest, []string{"application/pdf"}) + requirements.NoError(err) + + config.Index.Embeddings.Enabled = true + config.Index.Embeddings.Profile = "vector.embeddings" + after, err := config.ProfileFingerprint(manifest, []string{"application/pdf"}) + requirements.NoError(err) + afterPolicy, err := config.ProfilePolicyJSON(manifest, []string{"application/pdf"}) + requirements.NoError(err) + + assertions.Equal(before, after) + assertions.Equal(string(beforePolicy), string(afterPolicy)) +} + func TestDocumentsProfilePolicyJSONRemainsByteStable(t *testing.T) { require := require.New(t) assert := assert.New(t) @@ -169,7 +216,7 @@ func TestDocumentsProfilePolicyJSONRemainsByteStable(t *testing.T) { policyJSON, err := config.ProfilePolicyJSON(manifest, []string{"text/csv", "application/pdf", "text/csv"}) require.NoError(err) - expected := `{"version":1,"provider":"mistral","endpoint":"https://api.eu.mistral.ai/v1/ocr","model":"mistral-ocr-4-0","retention":"zdr","training":"opted-out","max_file_bytes":52428800,"max_pages_per_document":500,"max_response_bytes":67108864,"max_normalized_chars":25000000,"max_spool_bytes":536870912,"min_free_space_bytes":1073741824,"request_timeout_nanos":300000000000,"max_retries":3,"max_pages_per_run":10000,"max_estimated_cost_usd_per_run":50,"message_types":["chat","email"],"allowed_media_types":["application/pdf","text/csv"],"document_policy_fingerprint":"2614ee7d1bcac019ca6a7b78147be1954af97ebb22600ce0c22a36384f0813cd","lexical":true,"store_chunk_text":true,"extract_header":true,"extract_footer":true,"normalization_version":2,"max_unit_chars":1000000,"max_source_unit_bytes":4000000,"max_metadata_source_bytes":65536,"max_link_chars":2048,"max_chunk_runes":4000,"chunk_overlap":200,"max_chunks":20000}` + expected := `{"version":1,"provider":"mistral","endpoint":"https://api.eu.mistral.ai/v1/ocr","model":"mistral-ocr-4-0","retention":"zdr","training":"opted-out","max_file_bytes":52428800,"max_pages_per_document":500,"max_response_bytes":67108864,"max_normalized_chars":25000000,"max_spool_bytes":536870912,"min_free_space_bytes":1073741824,"request_timeout_nanos":300000000000,"max_retries":3,"max_pages_per_run":10000,"max_estimated_cost_usd_per_run":50,"message_types":["chat","email"],"allowed_media_types":["application/pdf","text/csv"],"document_policy_fingerprint":"466816abfedf47e64d25db7b38262b15d3a52077fd0856340fedd0227b600843","lexical":true,"store_chunk_text":true,"extract_header":true,"extract_footer":true,"normalization_version":3,"max_unit_chars":1000000,"max_source_unit_bytes":4000000,"max_metadata_source_bytes":65536,"max_link_chars":2048,"max_chunk_runes":4000,"chunk_overlap":200,"max_chunks":20000}` assert.JSONEq(expected, string(policyJSON)) fingerprint, err := config.ProfileFingerprint(manifest, []string{"application/pdf", "text/csv"}) diff --git a/internal/documentindex/worker.go b/internal/documentindex/worker.go index e6b98de8c..c170cc4a1 100644 --- a/internal/documentindex/worker.go +++ b/internal/documentindex/worker.go @@ -400,6 +400,8 @@ func publicationFromNormalized( LeaseOwner: claim.LeaseOwner, LeaseFence: claim.LeaseFence, ReturnedModel: providerResult.ReturnedModel, ProviderBytes: providerResult.ProviderBytes, UnitsProcessed: providerResult.UnitsProcessed, ManifestChecksum: normalized.Checksum, + NormalizationVersion: normalized.PolicyVersion, DocumentFamily: normalized.Family, + UnitKind: normalized.UnitKind, NormalizedTruncated: normalized.Truncated, RequestCount: providerResult.Metrics.Requests, RetryCount: providerResult.Metrics.Retries, ProviderLatencyMS: requestLatencyMillis(providerResult.Metrics.Latency), Units: make([]store.DocumentPublishedUnit, len(normalized.Units)), @@ -410,6 +412,7 @@ func publicationFromNormalized( Index: unit.Index, Kind: unit.Kind, Text: unit.Text, Header: unit.Header, Footer: unit.Footer, Width: unit.Dimensions.Width, Height: unit.Dimensions.Height, DPI: unit.Dimensions.DPI, Checksum: unit.Checksum, CharCount: unit.CharCount, Truncated: unit.Truncated, + HeadingMarks: unit.HeadingMarks, } } for i, chunk := range normalized.Chunks { diff --git a/internal/mcp/catalog.go b/internal/mcp/catalog.go index 14cd56192..05f0d1260 100644 --- a/internal/mcp/catalog.go +++ b/internal/mcp/catalog.go @@ -677,6 +677,9 @@ func findSimilarMessagesDefinition(_ *handlers) toolDefinition { func searchDocumentsDefinition(_ *handlers) toolDefinition { limit := boundedIntegerSchema("Maximum results to return (default 20, max 100)", 1, 100) limit.Default = json.RawMessage("20") + mode := stringSchema("Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider", + "auto", "lexical", "semantic", "hybrid") + mode.Default = json.RawMessage(`"lexical"`) direction := stringSchema("How the owning message relates to the person", "from_person", "to_person", "group") definition := readDefinition( @@ -700,10 +703,12 @@ func searchDocumentsDefinition(_ *handlers) toolDefinition { Type: "array", Description: "Optional union of from_person, to_person, and group; requires a person reference", Items: direction, }, - toolArgAfter: stringSchema("Only messages on or after YYYY-MM-DD"), - toolArgBefore: stringSchema("Only messages before YYYY-MM-DD"), - toolArgLimit: limit, - toolArgCursor: stringSchema("Opaque cursor from the previous page"), + toolArgAfter: stringSchema("Only messages on or after YYYY-MM-DD"), + toolArgBefore: stringSchema("Only messages before YYYY-MM-DD"), + toolArgLimit: limit, + toolArgCursor: stringSchema("Opaque cursor from the previous page"), + toolArgMode: mode, + "candidate_limit": boundedIntegerSchema("Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000)", 1, store.MaxLexicalDocumentSearchCandidateLimit), }, toolArgQuery), outputSchemaFor[store.DocumentSearchResponse](), (*handlers).searchDocuments, diff --git a/internal/mcp/catalog_test.go b/internal/mcp/catalog_test.go index ea8ede187..cc36f1fe8 100644 --- a/internal/mcp/catalog_test.go +++ b/internal/mcp/catalog_test.go @@ -317,7 +317,7 @@ func TestCatalogSchemas(t *testing.T) { } if shape.document { checks.Equal( - []string{"after", "attachment_id", "before", "cursor", "directions", "limit", "message_id", "message_types", "participant_id", "person_id", "query", "source_ids"}, + []string{"after", "attachment_id", "before", "candidate_limit", "cursor", "directions", "limit", "message_id", "message_types", "mode", "participant_id", "person_id", "query", "source_ids"}, toolPropertyNames(t, byName[ToolSearchDocuments]), ) } diff --git a/internal/mcp/handlers.go b/internal/mcp/handlers.go index 122d2e564..0be15bbad 100644 --- a/internal/mcp/handlers.go +++ b/internal/mcp/handlers.go @@ -23,6 +23,7 @@ import ( "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/vector" "go.kenn.io/msgvault/internal/vector/chunkmatch" + vectordocument "go.kenn.io/msgvault/internal/vector/document" "go.kenn.io/msgvault/internal/vector/embed" "go.kenn.io/msgvault/internal/vector/hybrid" "go.kenn.io/msgvault/internal/vector/visual" @@ -761,11 +762,39 @@ func (h *handlers) searchDocuments(ctx context.Context, req toolRequest) (*toolR limit = int(parsedLimit) } cursor, _ := args[toolArgCursor].(string) + mode, _ := args[toolArgMode].(string) + parsedMode := vectordocument.SearchModeAuto + if mode != "" { + parsed, parseErr := vectordocument.ParseSearchMode(mode) + if parseErr != nil { + return toolErrorResult(parseErr.Error()), nil + } + parsedMode = parsed + mode = string(parsed) + } + candidateLimit := 0 + if _, found := args["candidate_limit"]; found { + parsed, parseErr := positiveInt64Arg(args, "candidate_limit") + if parseErr != nil { + return toolErrorResult(parseErr.Error()), nil + } + maxCandidateLimit := store.MaxLexicalDocumentSearchCandidateLimit + if parsedMode == vectordocument.SearchModeSemantic || parsedMode == vectordocument.SearchModeHybrid { + maxCandidateLimit = store.MaxDocumentSearchCandidateLimit + } + if parsed > int64(maxCandidateLimit) { + return toolErrorResult(fmt.Sprintf( + "candidate_limit must be an integer between 1 and %d for this mode", maxCandidateLimit, + )), nil + } + candidateLimit = int(parsed) + } response, err := h.documentSearcher.SearchDocuments(ctx, store.DocumentSearchRequest{ Query: queryText, SourceIDs: sourceIDs, MessageTypes: messageTypes, AttachmentID: attachmentID, MessageID: messageID, PersonID: personID, ParticipantID: participantID, Directions: directions, After: after, Before: before, PageSize: limit, Cursor: cursor, + SearchMode: mode, CandidateLimit: candidateLimit, }) if err != nil { return toolErrorResult(fmt.Sprintf("document search failed: %v", err)), nil diff --git a/internal/mcp/server_test.go b/internal/mcp/server_test.go index 347c386b7..ecb6e998c 100644 --- a/internal/mcp/server_test.go +++ b/internal/mcp/server_test.go @@ -427,6 +427,7 @@ func TestSearchDocumentAttachmentsPreservesScopeAndProvenance(t *testing.T) { "directions": []any{"from_person", "group"}, "after": "2026-08-01", "before": "2026-08-20", "limit": float64(5), "cursor": "opaque", + "mode": "lexical", "candidate_limit": float64(1001), }) assert.Equal("carton damage", searcher.request.Query) assert.Equal([]int64{3, 7}, searcher.request.SourceIDs) @@ -439,6 +440,8 @@ func TestSearchDocumentAttachmentsPreservesScopeAndProvenance(t *testing.T) { require.NotNil(searcher.request.Before) assert.Equal(5, searcher.request.PageSize) assert.Equal("opaque", searcher.request.Cursor) + assert.Equal("lexical", searcher.request.SearchMode) + assert.Equal(1001, searcher.request.CandidateLimit) require.Len(response.Results, 1) assert.Equal("inspection.xlsx", response.Results[0].Filename) assert.Equal("mistral", response.Results[0].Provider) @@ -456,6 +459,16 @@ func TestSearchDocumentAttachmentsRejectsOutOfRangeExactID(t *testing.T) { assert.Empty(t, searcher.request.Query, "an invalid exact filter must never be silently dropped") } +func TestSearchDocumentAttachmentsRejectsSemanticCandidateLimitAboveBound(t *testing.T) { + searcher := &recordingDocumentSearcher{} + h := &handlers{documentSearcher: searcher} + result := runToolExpectError(t, ToolSearchDocuments, h.searchDocuments, map[string]any{ + "query": "carton damage", "mode": "semantic", "candidate_limit": float64(1001), + }) + assert.Contains(t, resultText(t, result), "candidate_limit must be an integer between 1 and 1000") + assert.Empty(t, searcher.request.Query) +} + func TestSearchRejectsInvalidQueryBeforeDispatch(t *testing.T) { queries := []struct { name string diff --git a/internal/scheduler/scheduler.go b/internal/scheduler/scheduler.go index 4ada2551c..2c95d054f 100644 --- a/internal/scheduler/scheduler.go +++ b/internal/scheduler/scheduler.go @@ -89,13 +89,17 @@ type Scheduler struct { // Embed job state (optional). Set via SetEmbedJob; cron.EntryID 0 // may be valid, so embedEntrySet tracks whether an entry exists. - embedJob *EmbedJob - embedEntry cron.EntryID - embedEntrySet bool - runEmbedAfterSync bool - visualPostSync func(context.Context) error - visualPostRunning bool - visualPostPending bool + embedJob *EmbedJob + embedEntry cron.EntryID + embedEntrySet bool + runEmbedAfterSync bool + visualPostSync func(context.Context) error + visualPostRunning bool + visualPostPending bool + documentVectorJob func(context.Context) error + documentVectorEntry cron.EntryID + documentVectorEntrySet bool + runDocumentVectorAfterSync bool ctx context.Context // cancelled on Stop cancel context.CancelFunc // cancels ctx @@ -318,6 +322,50 @@ func (s *Scheduler) SetEmbedJob(job *EmbedJob, schedule string, runAfterSync boo return nil } +// SetDocumentVectorJob installs the bounded document-vector convergence job +// on the same cron/post-sync policy used by message embeddings. +func (s *Scheduler) SetDocumentVectorJob(job func(context.Context) error, schedule string, runAfterSync bool) error { + if job != nil && schedule != "" { + if err := ValidateCronExpr(schedule); err != nil { + return fmt.Errorf("invalid document vector cron expression %q: %w", schedule, err) + } + } + s.mu.Lock() + defer s.mu.Unlock() + if s.documentVectorEntrySet { + s.cron.Remove(s.documentVectorEntry) + s.documentVectorEntrySet = false + } + s.documentVectorJob = job + s.runDocumentVectorAfterSync = runAfterSync && job != nil + if job == nil || schedule == "" { + return nil + } + entry, err := s.cron.AddFunc(schedule, func() { + if s.isStopped() { + return + } + done, ok := s.beginWork() + if !ok { + return + } + defer done() + runCtx, endRun := s.jobContext() + defer endRun() + if runErr := job(runCtx); runErr != nil { + s.logger.Error("scheduled document vector reconciliation failed", "error", runErr) + } + }) + if err != nil { + s.documentVectorJob = nil + s.runDocumentVectorAfterSync = false + return fmt.Errorf("register document vector cron: %w", err) + } + s.documentVectorEntry = entry + s.documentVectorEntrySet = true + return nil +} + // isStopped reports s.stopped under a read lock. Used by cron // callbacks that only need to abort on shutdown. func (s *Scheduler) isStopped() bool { @@ -446,6 +494,19 @@ func (s *Scheduler) runSync(email string) { postSync.Run(embedCtx) endEmbed() } + var documentVectorPostSync func(context.Context) error + s.mu.RLock() + if s.runDocumentVectorAfterSync && s.documentVectorJob != nil && !s.stopped { + documentVectorPostSync = s.documentVectorJob + } + s.mu.RUnlock() + if documentVectorPostSync != nil { + documentCtx, endDocument := s.jobContext() + if documentErr := documentVectorPostSync(documentCtx); documentErr != nil { + s.logger.Error("post-sync document vector reconciliation failed", "error", documentErr) + } + endDocument() + } s.startVisualPostSync() } diff --git a/internal/scheduler/scheduler_test.go b/internal/scheduler/scheduler_test.go index 0c1c0b817..2cb33d4cb 100644 --- a/internal/scheduler/scheduler_test.go +++ b/internal/scheduler/scheduler_test.go @@ -1831,6 +1831,70 @@ func TestEmbedJob_Run_NilSafe(t *testing.T) { // ---------- SetEmbedJob tests ---------- +func TestSchedulerSetDocumentVectorJobUsesEmbeddingSchedulePolicy(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + s := New(func(context.Context, string) error { return nil }) + t.Cleanup(func() { <-s.Stop().Done() }) + called := 0 + requirements.NoError(s.SetDocumentVectorJob(func(context.Context) error { + called++ + return nil + }, "*/5 * * * *", true)) + assertions.True(s.documentVectorEntrySet) + assertions.True(s.runDocumentVectorAfterSync) + requirements.ErrorContains(s.SetDocumentVectorJob(func(context.Context) error { return nil }, "invalid", false), "invalid") + assertions.True(s.documentVectorEntrySet, "invalid replacement preserves the prior job") + assertions.Zero(called) +} + +func TestSchedulerDocumentVectorJobRunsOnlyAfterSuccessfulSync(t *testing.T) { + for _, test := range []struct { + name string + syncErr error + wantRuns int64 + }{ + {name: "successful sync", wantRuns: 1}, + {name: "failed sync", syncErr: errors.New("sync failed")}, + } { + t.Run(test.name, func(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + syncDone := make(chan struct{}) + s := New(func(context.Context, string) error { + close(syncDone) + return test.syncErr + }) + var runs atomic.Int64 + documentDone := make(chan struct{}, 1) + requirements.NoError(s.SetDocumentVectorJob(func(context.Context) error { + runs.Add(1) + documentDone <- struct{}{} + return nil + }, "", true)) + requirements.NoError(s.AddAccount("test@example.test", "0 0 1 1 *")) + s.Start() + t.Cleanup(func() { <-s.Stop().Done() }) + requirements.NoError(s.TriggerSync("test@example.test")) + select { + case <-syncDone: + case <-time.After(time.Second): + requirements.Fail("sync did not complete") + } + if test.wantRuns == 1 { + select { + case <-documentDone: + case <-time.After(time.Second): + requirements.Fail("document vector job did not run") + } + } else { + assertions.Never(func() bool { return runs.Load() != 0 }, 100*time.Millisecond, 5*time.Millisecond) + } + assertions.Equal(test.wantRuns, runs.Load()) + }) + } +} + func TestScheduler_SetEmbedJob_AddsCronEntry(t *testing.T) { require := require.New(t) assert := assert.New(t) diff --git a/internal/store/dialect_pg.go b/internal/store/dialect_pg.go index ebbd2f443..33303cb5e 100644 --- a/internal/store/dialect_pg.go +++ b/internal/store/dialect_pg.go @@ -659,6 +659,11 @@ func (d *PostgreSQLDialect) LegacyColumnMigrations() []ColumnMigration { {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS request_count INTEGER NOT NULL DEFAULT 0 CHECK (request_count >= 0)`, "document_extractions.request_count"}, {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS retry_count INTEGER NOT NULL DEFAULT 0 CHECK (retry_count >= 0 AND retry_count <= request_count)`, "document_extractions.retry_count"}, {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS provider_latency_ms BIGINT NOT NULL DEFAULT 0 CHECK (provider_latency_ms >= 0)`, "document_extractions.provider_latency_ms"}, + {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS normalization_version INTEGER`, "document_extractions.normalization_version"}, + {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS document_family TEXT`, "document_extractions.document_family"}, + {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS unit_kind TEXT`, "document_extractions.unit_kind"}, + {`ALTER TABLE document_extractions ADD COLUMN IF NOT EXISTS normalized_truncated BOOLEAN NOT NULL DEFAULT FALSE`, "document_extractions.normalized_truncated"}, + {`ALTER TABLE document_units ADD COLUMN IF NOT EXISTS heading_marks JSONB NOT NULL DEFAULT '[]'::jsonb`, "document_units.heading_marks"}, {`ALTER TABLE document_index_state ADD COLUMN IF NOT EXISTS target_profile_id TEXT`, "document_index_state.target_profile_id"}, {`ALTER TABLE attachments ADD COLUMN IF NOT EXISTS attachment_state TEXT`, "attachments.attachment_state"}, {`ALTER TABLE attachments ADD COLUMN IF NOT EXISTS attachment_skip_reason TEXT`, "attachments.attachment_skip_reason"}, diff --git a/internal/store/dialect_sqlite.go b/internal/store/dialect_sqlite.go index c545a6e4a..4a43322e1 100644 --- a/internal/store/dialect_sqlite.go +++ b/internal/store/dialect_sqlite.go @@ -1614,6 +1614,11 @@ func (d *SQLiteDialect) LegacyColumnMigrations() []ColumnMigration { {`ALTER TABLE document_extractions ADD COLUMN request_count INTEGER NOT NULL DEFAULT 0 CHECK (request_count >= 0)`, "document_extractions.request_count"}, {`ALTER TABLE document_extractions ADD COLUMN retry_count INTEGER NOT NULL DEFAULT 0 CHECK (retry_count >= 0 AND retry_count <= request_count)`, "document_extractions.retry_count"}, {`ALTER TABLE document_extractions ADD COLUMN provider_latency_ms INTEGER NOT NULL DEFAULT 0 CHECK (provider_latency_ms >= 0)`, "document_extractions.provider_latency_ms"}, + {`ALTER TABLE document_extractions ADD COLUMN normalization_version INTEGER`, "document_extractions.normalization_version"}, + {`ALTER TABLE document_extractions ADD COLUMN document_family TEXT`, "document_extractions.document_family"}, + {`ALTER TABLE document_extractions ADD COLUMN unit_kind TEXT`, "document_extractions.unit_kind"}, + {`ALTER TABLE document_extractions ADD COLUMN normalized_truncated BOOLEAN NOT NULL DEFAULT FALSE`, "document_extractions.normalized_truncated"}, + {`ALTER TABLE document_units ADD COLUMN heading_marks JSON NOT NULL DEFAULT '[]'`, "document_units.heading_marks"}, {`ALTER TABLE document_index_state ADD COLUMN target_profile_id TEXT`, "document_index_state.target_profile_id"}, {`ALTER TABLE attachments ADD COLUMN attachment_state TEXT`, "attachments.attachment_state"}, {`ALTER TABLE attachments ADD COLUMN attachment_skip_reason TEXT`, "attachments.attachment_skip_reason"}, diff --git a/internal/store/document_normalized.go b/internal/store/document_normalized.go new file mode 100644 index 000000000..ecbc5e2a5 --- /dev/null +++ b/internal/store/document_normalized.go @@ -0,0 +1,162 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "strings" + + "go.kenn.io/docbank/document" +) + +var ErrDocumentNormalizedIdentityUnavailable = errors.New("document normalized identity is unavailable") + +func currentDocumentNormalizationVersion() (int, error) { + policy, err := document.NewNormalizePolicy(1) + if err != nil { + return 0, fmt.Errorf("load current document normalization policy: %w", err) + } + return policy.Identity().Version, nil +} + +func validateDocumentNormalizedIdentity(version int, family, unitKind string) error { + currentVersion, err := currentDocumentNormalizationVersion() + if err != nil { + return err + } + if version != currentVersion || strings.TrimSpace(family) == "" || strings.TrimSpace(unitKind) == "" { + return fmt.Errorf("%w: require normalization version %d, document family, and unit kind", + ErrDocumentNormalizedIdentityUnavailable, currentVersion) + } + return nil +} + +func documentNormalizedIdentityRebuildError(subject string, cause error) error { + return fmt.Errorf("%s: %w; run `msgvault documents build --full-rebuild --capabilities PATH --yes` before document vector consent or build", + subject, cause) +} + +// LoadNormalizedDocument reconstructs the complete immutable Docbank evidence +// published for one extraction. The stored v3 identities are validated before +// the document can be used to prepare provider inputs. +func (s *Store) LoadNormalizedDocument(ctx context.Context, extractionID string) (document.NormalizedDocument, error) { + if extractionID == "" { + return document.NormalizedDocument{}, errors.New("document extraction id is required") + } + var normalized document.NormalizedDocument + var normalizationVersion sql.NullInt64 + var documentFamily, unitKind sql.NullString + err := s.db.QueryRowContext(ctx, s.Rebind(` + SELECT normalization_version, document_family, unit_kind, + manifest_checksum, normalized_truncated + FROM document_extractions WHERE id = ? AND state = 'ready'`), extractionID).Scan( + &normalizationVersion, &documentFamily, &unitKind, + &normalized.Checksum, &normalized.Truncated, + ) + if errors.Is(err, sql.ErrNoRows) { + return document.NormalizedDocument{}, fmt.Errorf("ready document extraction %q not found", extractionID) + } + if err != nil { + return document.NormalizedDocument{}, fmt.Errorf("read normalized document identity: %w", err) + } + if !normalizationVersion.Valid || !documentFamily.Valid || !unitKind.Valid { + return document.NormalizedDocument{}, documentNormalizedIdentityRebuildError( + fmt.Sprintf("document extraction %q", extractionID), ErrDocumentNormalizedIdentityUnavailable, + ) + } + normalized.PolicyVersion = int(normalizationVersion.Int64) + normalized.Family = documentFamily.String + normalized.UnitKind = unitKind.String + if err := validateDocumentNormalizedIdentity(normalized.PolicyVersion, normalized.Family, normalized.UnitKind); err != nil { + return document.NormalizedDocument{}, documentNormalizedIdentityRebuildError( + fmt.Sprintf("document extraction %q", extractionID), err, + ) + } + unitRows, err := s.db.QueryContext(ctx, s.Rebind(` + SELECT unit_index, unit_kind, text, COALESCE(header_text, ''), COALESCE(footer_text, ''), + COALESCE(width, 0), COALESCE(height, 0), COALESCE(dpi, 0), checksum, + char_count, truncated, CAST(heading_marks AS TEXT) + FROM document_units WHERE extraction_id = ? ORDER BY unit_index`), extractionID) + if err != nil { + return document.NormalizedDocument{}, fmt.Errorf("read normalized document units: %w", err) + } + for unitRows.Next() { + var unit document.NormalizedUnit + var headingMarks string + if err := unitRows.Scan(&unit.Index, &unit.Kind, &unit.Text, &unit.Header, &unit.Footer, + &unit.Dimensions.Width, &unit.Dimensions.Height, &unit.Dimensions.DPI, + &unit.Checksum, &unit.CharCount, &unit.Truncated, &headingMarks); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("scan normalized document unit: %w", err) + } + unit.SourceKey = fmt.Sprintf("%s:%06d", normalized.UnitKind, unit.Index) + if err := json.Unmarshal([]byte(headingMarks), &unit.HeadingMarks); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("decode normalized document unit headings: %w", err) + } + normalized.Units = append(normalized.Units, unit) + } + if err := unitRows.Err(); err != nil { + _ = unitRows.Close() + return document.NormalizedDocument{}, fmt.Errorf("iterate normalized document units: %w", err) + } + if err := unitRows.Close(); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("close normalized document units: %w", err) + } + + chunkRows, err := s.db.QueryContext(ctx, s.Rebind(` + SELECT chunk_key, ordinal, text, CAST(heading_path AS TEXT), checksum, + char_count, truncated + FROM document_chunks WHERE extraction_id = ? ORDER BY ordinal`), extractionID) + if err != nil { + return document.NormalizedDocument{}, fmt.Errorf("read normalized document chunks: %w", err) + } + for chunkRows.Next() { + var chunk document.Chunk + var headingPath string + if err := chunkRows.Scan(&chunk.Key, &chunk.Ordinal, &chunk.Text, &headingPath, + &chunk.Checksum, &chunk.CharCount, &chunk.Truncated); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("scan normalized document chunk: %w", err) + } + if err := json.Unmarshal([]byte(headingPath), &chunk.HeadingPath); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("decode normalized document chunk headings: %w", err) + } + normalized.Chunks = append(normalized.Chunks, chunk) + } + if err := chunkRows.Err(); err != nil { + _ = chunkRows.Close() + return document.NormalizedDocument{}, fmt.Errorf("iterate normalized document chunks: %w", err) + } + if err := chunkRows.Close(); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("close normalized document chunks: %w", err) + } + for index := range normalized.Chunks { + chunk := &normalized.Chunks[index] + spanRows, err := s.db.QueryContext(ctx, s.Rebind(` + SELECT unit_index, start_char, end_char + FROM document_chunk_spans + WHERE extraction_id = ? AND chunk_key = ? ORDER BY span_ordinal`), extractionID, chunk.Key) + if err != nil { + return document.NormalizedDocument{}, fmt.Errorf("read normalized document chunk spans: %w", err) + } + for spanRows.Next() { + var span document.ChunkSpan + if err := spanRows.Scan(&span.UnitIndex, &span.CharStart, &span.CharEnd); err != nil { + _ = spanRows.Close() + return document.NormalizedDocument{}, fmt.Errorf("scan normalized document chunk span: %w", err) + } + chunk.Spans = append(chunk.Spans, span) + } + if err := spanRows.Err(); err != nil { + _ = spanRows.Close() + return document.NormalizedDocument{}, fmt.Errorf("iterate normalized document chunk spans: %w", err) + } + if err := spanRows.Close(); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("close normalized document chunk spans: %w", err) + } + } + if err := document.ValidateNormalizedDocument(normalized); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("validate stored normalized document: %w", err) + } + return normalized, nil +} diff --git a/internal/store/document_publication.go b/internal/store/document_publication.go index 871aa89b2..dabfb3b55 100644 --- a/internal/store/document_publication.go +++ b/internal/store/document_publication.go @@ -8,6 +8,8 @@ import ( "fmt" "time" "unicode/utf8" + + "go.kenn.io/docbank/document" ) var ( @@ -39,17 +41,18 @@ type DocumentExtractionClaim struct { } type DocumentPublishedUnit struct { - Index int - Kind string - Text string - Header string - Footer string - Width int - Height int - DPI int - Checksum string - CharCount int - Truncated bool + Index int + Kind string + Text string + Header string + Footer string + Width int + Height int + DPI int + Checksum string + CharCount int + Truncated bool + HeadingMarks []document.HeadingMark } type DocumentPublishedSpan struct { @@ -92,6 +95,10 @@ type DocumentExtractionPublication struct { RetryCount int ProviderLatencyMS int64 ManifestChecksum string + NormalizationVersion int + DocumentFamily string + UnitKind string + NormalizedTruncated bool Units []DocumentPublishedUnit Chunks []DocumentPublishedChunk } @@ -439,15 +446,19 @@ func (s *Store) PublishDocumentExtraction( return fmt.Errorf("clear staged document units: %w", err) } for _, unit := range publication.Units { + headingMarks, err := json.Marshal(unit.HeadingMarks) + if err != nil { + return fmt.Errorf("encode document unit heading marks: %w", err) + } if _, err := q.Exec(` INSERT INTO document_units (extraction_id, unit_index, unit_kind, text, header_text, - footer_text, width, height, dpi, checksum, char_count, truncated) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + footer_text, width, height, dpi, checksum, char_count, truncated, heading_marks) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, `+s.dialect.JSONBindExpr()+`)`, publication.ExtractionID, unit.Index, unit.Kind, unit.Text, nullIfEmpty(unit.Header), nullIfEmpty(unit.Footer), nullIfZero(int64(unit.Width)), nullIfZero(int64(unit.Height)), nullIfZero(int64(unit.DPI)), unit.Checksum, - unit.CharCount, unit.Truncated, + unit.CharCount, unit.Truncated, string(headingMarks), ); err != nil { return fmt.Errorf("publish document unit %d: %w", unit.Index, err) } @@ -492,13 +503,16 @@ func (s *Store) PublishDocumentExtraction( SET state = 'ready', lease_owner = NULL, lease_until = NULL, provider_bytes = ?, units_processed = ?, returned_model = ?, request_count = ?, retry_count = ?, provider_latency_ms = ?, - manifest_checksum = ?, source_sequence = ?, + manifest_checksum = ?, normalization_version = ?, document_family = ?, + unit_kind = ?, normalized_truncated = ?, source_sequence = ?, updated_at = `+s.dialect.Now()+`, published_at = `+s.dialect.Now()+` WHERE id = ? AND profile_id = ? AND canonical_blob_hash = ? AND extraction_input_key = ? AND state = 'staging'`, providerBytes, publication.UnitsProcessed, publication.ReturnedModel, publication.RequestCount, publication.RetryCount, publication.ProviderLatencyMS, - publication.ManifestChecksum, sourceSequence, publication.ExtractionID, + publication.ManifestChecksum, publication.NormalizationVersion, + publication.DocumentFamily, publication.UnitKind, publication.NormalizedTruncated, + sourceSequence, publication.ExtractionID, publication.ProfileID, publication.CanonicalBlobHash, publication.ExtractionInputKey, ) @@ -578,6 +592,11 @@ func validateDocumentPublication(publication DocumentExtractionPublication) erro publication.RetryCount > publication.RequestCount || publication.ProviderLatencyMS < 0 { return errors.New("document extraction publication has invalid provider request accounting") } + if err := validateDocumentNormalizedIdentity( + publication.NormalizationVersion, publication.DocumentFamily, publication.UnitKind, + ); err != nil { + return fmt.Errorf("document extraction publication has invalid normalized identity: %w", err) + } for index, unit := range publication.Units { if unit.Index != index || unit.Kind == "" || !utf8.ValidString(unit.Text) || unit.CharCount != utf8.RuneCountInString(unit.Text) || !validLowerSHA256(unit.Checksum) || diff --git a/internal/store/document_publication_test.go b/internal/store/document_publication_test.go index 696f79efa..0c2b0c920 100644 --- a/internal/store/document_publication_test.go +++ b/internal/store/document_publication_test.go @@ -7,10 +7,61 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + docbankdocument "go.kenn.io/docbank/document" "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/testutil/storetest" ) +func TestDocumentExtractionPublicationRoundTripsNormalizedV3Identity(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + claim, err := f.Store.ClaimDocumentExtraction(t.Context(), documentClaimInputForHash(t, f, store.DocumentExtractionClaimInput{ + ExtractionID: "extraction-normalized-v3", ProfileID: profile.ID, + CanonicalBlobHash: hash, ExtractionInputKey: "original", + LeaseOwner: "worker-normalized-v3", LeaseUntil: time.Now().UTC().Add(10 * time.Minute), + LocalBytes: 128, SourceSequence: 1, + })) + require.NoError(err) + policy, err := docbankdocument.NewNormalizePolicy(10_000) + require.NoError(err) + normalized, err := docbankdocument.NormalizeDocument(docbankdocument.SourceDocument{ + Family: "pdf", UnitKind: "page", Units: []docbankdocument.SourceUnit{{ + Index: 0, Markdown: "# Evidence\n\nStored identity", + }}, + }, policy) + require.NoError(err) + publication := publicationFor(t, claim, normalized.Chunks[0].Text, normalized.Chunks[0].Checksum) + publication.ManifestChecksum = normalized.Checksum + publication.NormalizationVersion = normalized.PolicyVersion + publication.DocumentFamily = normalized.Family + publication.UnitKind = normalized.UnitKind + publication.NormalizedTruncated = normalized.Truncated + publication.Units[0] = store.DocumentPublishedUnit{ + Index: 0, Kind: normalized.Units[0].Kind, Text: normalized.Units[0].Text, + Header: normalized.Units[0].Header, Footer: normalized.Units[0].Footer, + Width: normalized.Units[0].Dimensions.Width, Height: normalized.Units[0].Dimensions.Height, + DPI: normalized.Units[0].Dimensions.DPI, Checksum: normalized.Units[0].Checksum, + CharCount: normalized.Units[0].CharCount, Truncated: normalized.Units[0].Truncated, + HeadingMarks: normalized.Units[0].HeadingMarks, + } + publication.Chunks[0].Key = normalized.Chunks[0].Key + publication.Chunks[0].HeadingPath = normalized.Chunks[0].HeadingPath + publication.Chunks[0].Truncated = normalized.Chunks[0].Truncated + publication.Chunks[0].Spans[0] = store.DocumentPublishedSpan{ + UnitIndex: normalized.Chunks[0].Spans[0].UnitIndex, + CharStart: normalized.Chunks[0].Spans[0].CharStart, + CharEnd: normalized.Chunks[0].Spans[0].CharEnd, + } + require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publication)) + + loaded, err := f.Store.LoadNormalizedDocument(t.Context(), claim.ExtractionID) + + require.NoError(err) + assert.Equal(normalized, loaded) +} + func TestDocumentExtractionPublicationKeepsOldHeadUntilAtomicSwitch(t *testing.T) { require := require.New(t) assert := assert.New(t) @@ -42,7 +93,7 @@ func TestDocumentExtractionPublicationKeepsOldHeadUntilAtomicSwitch(t *testing.T wrongFence.LeaseFence++ err = f.Store.RenewDocumentExtractionClaim(t.Context(), wrongFence, time.Now().UTC().Add(15*time.Minute)) require.ErrorIs(err, store.ErrDocumentExtractionFenceLost) - require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publicationFor( + require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publicationFor(t, firstClaim, "old searchable quasar evidence", strings.Repeat("d", 64), ))) assert.Equal(1, documentFTSMatchCount(t, f.Store, "quasar")) @@ -58,7 +109,7 @@ func TestDocumentExtractionPublicationKeepsOldHeadUntilAtomicSwitch(t *testing.T assert.Equal([]string{"old searchable quasar evidence"}, currentDocumentTexts(t, f, profile.ID, hash), "a staging replacement must not hide the ready head") - require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publicationFor( + require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publicationFor(t, secondClaim, "new searchable nebula evidence", strings.Repeat("e", 64), ))) assert.Equal([]string{"new searchable nebula evidence"}, currentDocumentTexts(t, f, profile.ID, hash)) @@ -82,7 +133,7 @@ func TestDocumentExtractionPublicationRejectsInvalidSpanBeforeMutation(t *testin LocalBytes: 128, SourceSequence: 1, })) require.NoError(err) - publication := publicationFor(claim, "short text", strings.Repeat("f", 64)) + publication := publicationFor(t, claim, "short text", strings.Repeat("f", 64)) publication.Chunks[0].Spans[0].CharEnd = 100 require.ErrorContains(f.Store.PublishDocumentExtraction(t.Context(), publication), "span 0 is invalid") @@ -96,6 +147,31 @@ func TestDocumentExtractionPublicationRejectsInvalidSpanBeforeMutation(t *testin assert.Equal(t, "staging", state) } +func TestDocumentExtractionPublicationRequiresNormalizedIdentity(t *testing.T) { + requirements := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + claim, err := f.Store.ClaimDocumentExtraction(t.Context(), documentClaimInputForHash(t, f, store.DocumentExtractionClaimInput{ + ExtractionID: "extraction-missing-normalized-identity", ProfileID: profile.ID, + CanonicalBlobHash: hash, ExtractionInputKey: "original", + LeaseOwner: "worker-missing-normalized-identity", LeaseUntil: time.Now().UTC().Add(10 * time.Minute), + LocalBytes: 128, SourceSequence: 1, + })) + requirements.NoError(err) + + for name, clearIdentity := range map[string]func(*store.DocumentExtractionPublication){ + "normalization version": func(publication *store.DocumentExtractionPublication) { publication.NormalizationVersion = 0 }, + "document family": func(publication *store.DocumentExtractionPublication) { publication.DocumentFamily = "" }, + "unit kind": func(publication *store.DocumentExtractionPublication) { publication.UnitKind = "" }, + } { + t.Run(name, func(t *testing.T) { + publication := publicationFor(t, claim, "normalized identity", strings.Repeat("f", 64)) + clearIdentity(&publication) + require.ErrorContains(t, f.Store.PublishDocumentExtraction(t.Context(), publication), "normalized identity") + }) + } +} + func TestDocumentExtractionClaimRequiresAuthoritativeRoleProvenance(t *testing.T) { tests := []struct { name string @@ -170,7 +246,7 @@ func TestDocumentExtractionPublicationRechecksClaimedOccurrenceScope(t *testing. _, err = f.Store.DB().Exec(f.Store.Rebind(test.update), claim.OccurrenceAttachmentID) require.NoError(err) - err = f.Store.PublishDocumentExtraction(t.Context(), publicationFor( + err = f.Store.PublishDocumentExtraction(t.Context(), publicationFor(t, claim, "must not publish", strings.Repeat("f", 64), )) require.ErrorContains(err, "claimed occurrence is no longer eligible") @@ -194,7 +270,7 @@ func TestDocumentExtractionPublicationAcceptsTrustedHashlessCASAlias(t *testing. claim, err := f.Store.ClaimDocumentExtraction(t.Context(), input) require.NoError(err) - require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publicationFor( + require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publicationFor(t, claim, "hashless alias evidence", strings.Repeat("a", 64), ))) } @@ -406,10 +482,14 @@ func seedDocumentPublicationAuthority( } func publicationFor( + t *testing.T, claim store.DocumentExtractionClaim, text string, checksum string, ) store.DocumentExtractionPublication { + t.Helper() + policy, err := docbankdocument.NewNormalizePolicy(1) + require.NoError(t, err) return store.DocumentExtractionPublication{ ExtractionID: claim.ExtractionID, ProfileID: claim.ProfileID, CanonicalBlobHash: claim.CanonicalBlobHash, ExtractionInputKey: claim.ExtractionInputKey, @@ -419,7 +499,8 @@ func publicationFor( LeaseOwner: claim.LeaseOwner, LeaseFence: claim.LeaseFence, ReturnedModel: "mistral-ocr-4-0", UnitsProcessed: 1, RequestCount: 1, ProviderLatencyMS: 25, - ManifestChecksum: strings.Repeat("c", 64), + ManifestChecksum: strings.Repeat("c", 64), NormalizationVersion: policy.Identity().Version, + DocumentFamily: "pdf", UnitKind: "page", Units: []store.DocumentPublishedUnit{{ Index: 0, Kind: "page", Text: text, Checksum: checksum, CharCount: len([]rune(text)), }}, diff --git a/internal/store/document_search.go b/internal/store/document_search.go index e8a6b58b7..66047cc86 100644 --- a/internal/store/document_search.go +++ b/internal/store/document_search.go @@ -17,18 +17,22 @@ import ( "unicode" "unicode/utf8" + docembedding "go.kenn.io/docbank/document/embedding" "go.kenn.io/msgvault/internal/personscope" ) const ( - documentSearchCursorVersion = 2 - documentSearchRRFConstant = 60.0 - maxDocumentSearchQueryBytes = 1_024 - maxDocumentSearchTerms = 20 - maxDocumentSearchPageSize = 100 - maxDocumentSearchOffset = 10_000 - maxDocumentSearchCandidates = 10_000 - maxDocumentSearchExcerpt = 320 + documentSearchCursorVersion = 2 + documentSearchRRFConstant = float64(docembedding.DefaultReciprocalRankConstant) + DefaultDocumentSearchCandidateLimit = docembedding.DefaultCandidateLimit + MaxDocumentSearchCandidateLimit = docembedding.MaxCandidateLimit + DefaultLexicalDocumentSearchCandidateLimit = 10_000 + MaxLexicalDocumentSearchCandidateLimit = 10_000 + maxDocumentSearchQueryBytes = 1_024 + maxDocumentSearchTerms = 20 + maxDocumentSearchPageSize = 100 + maxDocumentSearchOffset = 10_000 + maxDocumentSearchExcerpt = 320 ) var ( @@ -39,58 +43,81 @@ var ( ) type DocumentSearchRequest struct { - Query string - SourceIDs []int64 - MessageTypes []string - AttachmentID int64 - MessageID int64 - PageSize int - Cursor string - After *time.Time - Before *time.Time - PersonID int64 - ParticipantID int64 - Directions []personscope.Direction - Person *personscope.Scope + Query string `json:"query"` + SourceIDs []int64 `json:"source_ids,omitempty"` + MessageTypes []string `json:"message_types,omitempty"` + AttachmentID int64 `json:"attachment_id,omitempty"` + MessageID int64 `json:"message_id,omitempty"` + PageSize int `json:"page_size"` + Cursor string `json:"cursor,omitempty"` + SearchMode string `json:"search_mode,omitempty"` + CandidateLimit int `json:"candidate_limit,omitempty"` + After *time.Time `json:"after,omitempty"` + Before *time.Time `json:"before,omitempty"` + PersonID int64 `json:"person_id,omitempty"` + ParticipantID int64 `json:"participant_id,omitempty"` + Directions []personscope.Direction `json:"directions,omitempty"` + Person *personscope.Scope `json:"person,omitempty"` } type DocumentSearchResponse struct { - Results []DocumentSearchResult `json:"results"` - NextCursor string `json:"next_cursor,omitempty"` - Revision int64 `json:"revision"` - Truncated bool `json:"truncated,omitempty"` + Results []DocumentSearchResult `json:"results"` + NextCursor string `json:"next_cursor,omitempty"` + Revision int64 `json:"revision"` + Truncated bool `json:"truncated,omitempty"` + EffectiveMode string `json:"effective_mode,omitempty"` + VectorGenerationID int64 `json:"vector_generation_id,omitempty"` + VectorGenerationFingerprint string `json:"vector_generation_fingerprint,omitempty"` } type DocumentSearchResult struct { - AttachmentID int64 `json:"attachment_id"` - MessageID int64 `json:"message_id"` - ConversationID int64 `json:"conversation_id"` - SourceID int64 `json:"source_id"` - SourceMessageID string `json:"source_message_id,omitempty"` - OccurredAt *time.Time `json:"occurred_at,omitempty"` - OccurrenceKey string `json:"occurrence_key"` - SourcePartKey string `json:"source_part_key,omitempty"` - Filename string `json:"filename,omitempty"` - ContainingTitle string `json:"containing_title,omitempty"` - MIMEType string `json:"mime_type,omitempty"` - CanonicalBlobHash string `json:"canonical_blob_hash"` - OtherLiveCopies int `json:"other_live_copies"` - ChunkKey string `json:"chunk_key"` - ChunkOrdinal int `json:"chunk_ordinal"` - HeadingPath []string `json:"heading_path,omitempty"` - FirstUnitIndex int `json:"first_unit_index"` - LastUnitIndex int `json:"last_unit_index"` - Excerpt string `json:"excerpt"` - HighlightStart int `json:"highlight_start"` - HighlightEnd int `json:"highlight_end"` - ProfileID string `json:"profile_id"` - ExtractionID string `json:"extraction_id"` - Provider string `json:"provider"` - Model string `json:"model"` - MatchedSignals []string `json:"matched_signals"` - Truncated bool `json:"truncated"` - Rank int `json:"rank"` - PersonProvenance *personscope.Provenance `json:"person_provenance,omitempty"` + AttachmentID int64 `json:"attachment_id"` + MessageID int64 `json:"message_id"` + ConversationID int64 `json:"conversation_id"` + SourceID int64 `json:"source_id"` + SourceMessageID string `json:"source_message_id,omitempty"` + OccurredAt *time.Time `json:"occurred_at,omitempty"` + OccurrenceKey string `json:"occurrence_key"` + SourcePartKey string `json:"source_part_key,omitempty"` + Filename string `json:"filename,omitempty"` + ContainingTitle string `json:"containing_title,omitempty"` + MIMEType string `json:"mime_type,omitempty"` + CanonicalBlobHash string `json:"canonical_blob_hash"` + OtherLiveCopies int `json:"other_live_copies"` + ChunkKey string `json:"chunk_key"` + ChunkOrdinal int `json:"chunk_ordinal"` + HeadingPath []string `json:"heading_path,omitempty"` + FirstUnitIndex int `json:"first_unit_index"` + LastUnitIndex int `json:"last_unit_index"` + Excerpt string `json:"excerpt"` + HighlightStart int `json:"highlight_start"` + HighlightEnd int `json:"highlight_end"` + ProfileID string `json:"profile_id"` + ExtractionID string `json:"extraction_id"` + Provider string `json:"provider"` + Model string `json:"model"` + MatchedSignals []string `json:"matched_signals"` + Truncated bool `json:"truncated"` + Rank int `json:"rank"` + PersonProvenance *personscope.Provenance `json:"person_provenance,omitempty"` + LexicalRank int `json:"lexical_rank,omitempty"` + SemanticRank int `json:"semantic_rank,omitempty"` + SemanticScore float64 `json:"semantic_score,omitempty"` + FusionScore float64 `json:"fusion_score,omitempty"` + VectorToken string `json:"vector_token,omitempty"` + VectorGenerationID int64 `json:"vector_generation_id,omitempty"` + VectorGenerationFingerprint string `json:"vector_generation_fingerprint,omitempty"` + VectorEmbeddingProfile string `json:"vector_embedding_profile,omitempty"` + VectorModel string `json:"vector_model,omitempty"` + VectorDimension int `json:"vector_dimension,omitempty"` +} + +// DocumentVectorSearchHit is an opaque backend hit presented for authoritative +// occurrence expansion. Rank is one-based in backend order. +type DocumentVectorSearchHit struct { + Token string + Score float64 + Rank int } type documentSearchCursor struct { @@ -145,7 +172,7 @@ func (s *Store) SearchDocuments( moreCandidates := contentMore || filenameMore || fusionMore response := DocumentSearchResponse{ Revision: revision, - Truncated: moreCandidates && candidateLimit == maxDocumentSearchCandidates, + Truncated: moreCandidates, } if offset >= len(rows) { return response, nil @@ -177,6 +204,160 @@ func (s *Store) SearchDocuments( return response, nil } +// ResolveDocumentVectorSearchOccurrences filters backend hits through the +// active generation's exact publication snapshot and expands surviving chunks +// to scoped live attachment occurrences. The limit applies after occurrence +// deduplication, so duplicate chunks and heavily reused blobs stay bounded. +func (s *Store) ResolveDocumentVectorSearchOccurrences( + ctx context.Context, + generationID int64, + hits []DocumentVectorSearchHit, + request DocumentSearchRequest, + limit int, +) ([]DocumentSearchResult, bool, error) { + if generationID <= 0 { + return nil, false, fmt.Errorf("%w: vector generation must be positive", ErrDocumentSearchInvalidRequest) + } + if limit < 1 || limit > maxDocumentVectorCandidateLimit+1 || len(hits) > maxDocumentVectorCandidateLimit { + return nil, false, fmt.Errorf("%w: semantic candidate bounds are invalid", ErrDocumentSearchInvalidRequest) + } + if len(hits) == 0 { + return []DocumentSearchResult{}, false, nil + } + var err error + request.SourceIDs, err = sortedUniquePositive(request.SourceIDs) + if err != nil { + return nil, false, err + } + request.MessageTypes, err = sortedUniqueNonempty(request.MessageTypes) + if err != nil { + return nil, false, err + } + if request.AttachmentID < 0 || request.MessageID < 0 { + return nil, false, fmt.Errorf("%w: request scope has invalid bounds", ErrDocumentSearchInvalidRequest) + } + seenTokens := make(map[string]struct{}, len(hits)) + values := make([]string, 0, len(hits)) + args := make([]any, 0, len(hits)*3+6) + for _, hit := range hits { + if !documentVectorFingerprintPattern.MatchString(hit.Token) || hit.Rank < 1 || math.IsNaN(hit.Score) || math.IsInf(hit.Score, 0) { + return nil, false, fmt.Errorf("%w: semantic hit is invalid", ErrDocumentSearchInvalidRequest) + } + if _, exists := seenTokens[hit.Token]; exists { + return nil, false, fmt.Errorf("%w: semantic hit token is duplicated", ErrDocumentSearchInvalidRequest) + } + seenTokens[hit.Token] = struct{}{} + // Explicit parameter casts keep PostgreSQL from inferring a parameter-only + // VALUES column as text (which would order rank 10 before rank 2). + values = append(values, "(CAST(? AS TEXT), CAST(? AS INTEGER), CAST(? AS DOUBLE PRECISION))") + args = append(args, hit.Token, hit.Rank, hit.Score) + } + conditions, scopeArgs := documentSearchScope(request, "m", "a", "cv") + args = append(args, generationID, string(DocumentVectorGenerationActive)) + args = append(args, scopeArgs...) + args = append(args, limit+1) + query := ` + WITH requested(token, semantic_rank, semantic_score) AS ( + VALUES ` + strings.Join(values, ", ") + ` + ), ranked AS ( + SELECT ` + documentSearchRankedSelectColumns + `, + v.token AS vector_token, + requested.semantic_rank AS semantic_rank, + requested.semantic_score AS semantic_score, + g.id AS vector_generation_id, + g.fingerprint AS vector_generation_fingerprint, + g.embedding_profile AS vector_embedding_profile, + g.model AS vector_model, + g.dimension AS vector_dimension, + ROW_NUMBER() OVER ( + PARTITION BY o.occurrence_key + ORDER BY requested.semantic_rank, requested.token + ) AS occurrence_rank + FROM requested + JOIN document_vector_publications v ON v.token = requested.token + JOIN document_vector_generations g ON g.id = v.generation_id + JOIN document_index_state ds ON ds.singleton = 1 + JOIN document_extraction_heads h + ON h.extraction_id = v.extraction_id + AND h.profile_id = v.extraction_profile_id + AND h.canonical_blob_hash = v.canonical_blob_hash + AND h.extraction_input_key = v.extraction_input_key + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc + ON dc.id = v.chunk_id AND dc.extraction_id = v.extraction_id + AND dc.chunk_key = v.chunk_key AND dc.checksum = v.chunk_checksum + JOIN document_occurrences o ON o.canonical_blob_hash = h.canonical_blob_hash + JOIN attachments a ON a.id = o.attachment_id + JOIN messages m ON m.id = o.message_id + JOIN conversations cv ON cv.id = m.conversation_id + WHERE g.id = ? AND g.state = ? + AND ds.target_profile_id = g.target_extraction_profile_id + AND v.state = 'ready' + AND e.source_sequence = v.source_sequence + AND h.source_sequence = v.source_sequence + AND ` + documentSearchValidity() + conditions + ` + ) + SELECT ` + documentSearchOuterColumns + `, + vector_token, semantic_rank, semantic_score, + vector_generation_id, vector_generation_fingerprint, + vector_embedding_profile, vector_model, vector_dimension + FROM ranked + WHERE occurrence_rank = 1 + ORDER BY semantic_rank, occurrence_key + LIMIT ?` + rows, err := s.db.QueryContext(ctx, s.Rebind(query), args...) + if err != nil { + return nil, false, fmt.Errorf("resolve document vector search occurrences: %w", err) + } + defer func() { _ = rows.Close() }() + results := make([]DocumentSearchResult, 0, limit+1) + for rows.Next() { + var result DocumentSearchResult + var headingJSON, text string + var occurredAt nullableTimestamp + if err := rows.Scan( + &result.AttachmentID, &result.MessageID, &result.ConversationID, &result.SourceID, + &result.SourceMessageID, &occurredAt, &result.OccurrenceKey, + &result.SourcePartKey, &result.Filename, &result.ContainingTitle, &result.MIMEType, + &result.CanonicalBlobHash, &result.ChunkKey, &result.ChunkOrdinal, &headingJSON, + &result.FirstUnitIndex, &result.LastUnitIndex, &text, + &result.ProfileID, &result.ExtractionID, &result.Provider, &result.Model, + &result.Truncated, &result.VectorToken, &result.SemanticRank, &result.SemanticScore, + &result.VectorGenerationID, &result.VectorGenerationFingerprint, + &result.VectorEmbeddingProfile, &result.VectorModel, &result.VectorDimension, + ); err != nil { + return nil, false, fmt.Errorf("scan document vector search occurrence: %w", err) + } + if occurredAt.Valid { + result.OccurredAt = &occurredAt.Time + } + if err := json.Unmarshal([]byte(headingJSON), &result.HeadingPath); err != nil { + return nil, false, fmt.Errorf("decode document vector search heading path: %w", err) + } + result.Excerpt, result.HighlightStart, result.HighlightEnd = documentSearchExcerpt(text, nil) + result.MatchedSignals = []string{"semantic"} + results = append(results, result) + } + if err := rows.Err(); err != nil { + return nil, false, fmt.Errorf("iterate document vector search occurrences: %w", err) + } + truncated := len(results) > limit + if truncated { + results = results[:limit] + } + if err := s.populateDocumentLiveCopyCounts(ctx, results); err != nil { + return nil, false, err + } + if request.Person != nil { + if err := s.populateDocumentPersonProvenance(ctx, results, *request.Person); err != nil { + return nil, false, err + } + } + return results, truncated, nil +} + func (s *Store) prepareDocumentSearch( ctx context.Context, request DocumentSearchRequest, @@ -195,7 +376,11 @@ func (s *Store) prepareDocumentSearch( if request.PageSize == 0 { request.PageSize = 20 } + if request.CandidateLimit == 0 { + request.CandidateLimit = DefaultLexicalDocumentSearchCandidateLimit + } if request.PageSize < 1 || request.PageSize > maxDocumentSearchPageSize || + request.CandidateLimit < 1 || request.CandidateLimit > MaxLexicalDocumentSearchCandidateLimit || request.AttachmentID < 0 || request.MessageID < 0 || request.PersonID < 0 || request.ParticipantID < 0 || (request.PersonID > 0 && request.ParticipantID > 0) || ((request.PersonID > 0 || request.ParticipantID > 0) && request.Person == nil) { @@ -230,7 +415,7 @@ func (s *Store) prepareDocumentSearch( // RRF ranks are meaningful only for one fixed candidate set. Every page // therefore evaluates the same bounded set instead of widening it between // cursors, which could reorder earlier results and cause skips or repeats. - candidateLimit := maxDocumentSearchCandidates + candidateLimit := request.CandidateLimit if request.Cursor != "" { cursor, decodeErr := decodeDocumentSearchCursor(request.Cursor) if decodeErr != nil { @@ -258,7 +443,7 @@ func (s *Store) searchDocumentContent( ) ([]documentSearchRow, bool, error) { ftsArg := s.dialect.BuildFTSArg(terms) conditions, scopeArgs := documentSearchScope(request, "m", "a", "cv") - validity := documentSearchValidity("p", "c", "h", "o", "a", "m", "ds") + validity := documentSearchValidity() var query string args := make([]any, 0, len(scopeArgs)+2) if s.IsPostgreSQL() { @@ -355,7 +540,7 @@ func (s *Store) searchDocumentFilenames( JOIN messages m ON m.id = o.message_id JOIN conversations cv ON cv.id = m.conversation_id CROSS JOIN document_index_state ds - WHERE ` + documentSearchValidity("p", "c", "h", "o", "a", "m", "ds") + conditions + ` + WHERE ` + documentSearchValidity() + conditions + ` ORDER BY LOWER(COALESCE(o.filename, '')), o.occurrence_key LIMIT ?` args = append(args, limit+1) @@ -401,15 +586,17 @@ const documentSearchOuterColumns = ` first_unit_index, last_unit_index, chunk_text, profile_id, extraction_id, provider, model, truncated` -func documentSearchValidity(profile, consent, head, occurrence, attachment, message, state string) string { +func documentSearchValidity() string { + const profile, consent, head, attachment, message, state = "p", "c", "h", "a", "m", "ds" + return profile + `.enabled = TRUE AND ` + profile + `.retired_at IS NULL AND ` + consent + `.profile_fingerprint = ` + profile + `.fingerprint AND ` + consent + `.retention_posture = ` + profile + `.retention_posture AND ` + consent + `.training_posture = ` + profile + `.training_posture - AND ` + occurrence + `.attachment_role = 'standalone' + AND o.attachment_role = 'standalone' AND ` + attachment + `.attachment_role = 'standalone' - AND ` + occurrence + `.role_source IN ('mime_disposition', 'provider_explicit', 'importer_semantics', 'raw_mime_repair') + AND o.role_source IN ('mime_disposition', 'provider_explicit', 'importer_semantics', 'raw_mime_repair') AND ` + attachment + `.role_source IN ('mime_disposition', 'provider_explicit', 'importer_semantics', 'raw_mime_repair') AND ` + LiveMessagesWhere(message, true) + ` AND (` + attachment + `.content_hash = ` + head + `.canonical_blob_hash @@ -593,7 +780,7 @@ func (s *Store) populateDocumentLiveCopyCounts( JOIN messages m ON m.id = o.message_id CROSS JOIN document_index_state ds WHERE h.canonical_blob_hash IN (` + documentPlaceholders(len(hashes)) + `) - AND ` + documentSearchValidity("p", "c", "h", "o", "a", "m", "ds") + ` + AND ` + documentSearchValidity() + ` GROUP BY h.canonical_blob_hash` args := make([]any, len(hashes)) for index := range hashes { @@ -740,7 +927,7 @@ func decodeDocumentSearchCursor(value string) (documentSearchCursor, error) { if err := decoder.Decode(&cursor); err != nil || cursor.Version != documentSearchCursorVersion || !validLowerSHA256(cursor.RequestHash) || cursor.Revision < 0 || cursor.Offset <= 0 || cursor.Offset > maxDocumentSearchOffset || - cursor.CandidateLimit < 1 || cursor.CandidateLimit > maxDocumentSearchCandidates { + cursor.CandidateLimit < 1 || cursor.CandidateLimit > MaxLexicalDocumentSearchCandidateLimit { return documentSearchCursor{}, ErrDocumentSearchInvalidCursor } if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { diff --git a/internal/store/document_search_test.go b/internal/store/document_search_test.go index f2fb722af..34649421d 100644 --- a/internal/store/document_search_test.go +++ b/internal/store/document_search_test.go @@ -497,13 +497,35 @@ func TestSearchDocumentsAppliesCandidateLimitAfterOccurrenceDeduplication(t *tes assert.Contains(t, []int64{response.Results[0].AttachmentID, response.Results[1].AttachmentID}, secondAttachmentID) } +func TestSearchDocumentsHonorsExplicitCandidateLimit(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, "bounded nebula evidence", "search-bounded") + messageID := f.CreateMessage("document-search-bounded-copy") + attachmentID := addSearchAttachment( + t, f, messageID, hash, "bounded-copy.pdf", "provider:bounded-copy", + ) + _, eligible, err := f.Store.ReconcileDocumentOccurrence(t.Context(), attachmentID, 2) + requirements.NoError(err) + requirements.True(eligible) + + response, err := f.Store.SearchDocuments(t.Context(), store.DocumentSearchRequest{ + Query: "nebula", CandidateLimit: 1, + }) + requirements.NoError(err) + assertions.Len(response.Results, 1) + assertions.True(response.Truncated) +} + func TestSearchDocumentsPaginationUsesStableRankingSet(t *testing.T) { require := require.New(t) f := storetest.New(t) profile, hash := seedDocumentPublicationAuthority(t, f) publishSearchDocument(t, f, profile, hash, "window nebula evidence", "search-window") - const copies = 201 + const copies = 1001 for index := 1; index < copies; index++ { messageID := f.CreateMessage("document-search-window-" + strconv.Itoa(index)) attachmentID := addSearchAttachment( @@ -520,7 +542,7 @@ func TestSearchDocumentsPaginationUsesStableRankingSet(t *testing.T) { nextRank := 1 for { response, err := f.Store.SearchDocuments(t.Context(), store.DocumentSearchRequest{ - Query: "nebula", PageSize: 10, Cursor: cursor, + Query: "nebula", PageSize: 100, Cursor: cursor, }) require.NoError(err) for _, result := range response.Results { @@ -539,6 +561,221 @@ func TestSearchDocumentsPaginationUsesStableRankingSet(t *testing.T) { require.Len(seen, copies, "pagination must cover the fixed ranked candidate set") } +func TestResolveDocumentVectorSearchOccurrencesExpandsAndBoundsAfterOccurrenceDeduplication(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 2) + now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 2) + + copyMessageID := f.CreateMessage("semantic-search-copy") + copyAttachmentID := addSearchAttachment( + t, f, copyMessageID, claims[0].CanonicalBlobHash, "semantic-copy.pdf", "provider:semantic-copy", + ) + _, err := f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET sent_at = ? WHERE id = ?`), now, copyMessageID) + require.NoError(err) + _, eligible, err := f.Store.ReconcileDocumentOccurrence(t.Context(), copyAttachmentID, 2) + require.NoError(err) + require.True(eligible) + require.NoError(f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + + hits := []store.DocumentVectorSearchHit{ + {Token: claims[0].Token, Score: .8, Rank: 10}, + {Token: claims[1].Token, Score: .9, Rank: 2}, + } + results, resultsMore, err := f.Store.ResolveDocumentVectorSearchOccurrences( + t.Context(), generation.ID, hits, store.DocumentSearchRequest{}, 10, + ) + require.NoError(err) + assert.False(resultsMore) + require.Len(results, 2, "two chunks must expand to each occurrence, then collapse by occurrence") + assert.Less(results[0].OccurrenceKey, results[1].OccurrenceKey) + for _, result := range results { + assert.Equal(claims[1].Token, result.VectorToken) + assert.Equal(2, result.SemanticRank) + assert.InDelta(.9, result.SemanticScore, 1e-12) + assert.Equal(generation.ID, result.VectorGenerationID) + assert.Equal(generation.Fingerprint, result.VectorGenerationFingerprint) + assert.Equal(generation.EmbeddingProfile, result.VectorEmbeddingProfile) + assert.Equal(generation.Model, result.VectorModel) + assert.Equal(generation.Dimension, result.VectorDimension) + assert.Equal(claims[1].ChunkKey, result.ChunkKey) + assert.Equal(claims[1].ChunkOrdinal, result.ChunkOrdinal) + assert.Equal(claims[1].ExtractionID, result.ExtractionID) + assert.Equal(claims[1].ExtractionProfileID, result.ProfileID) + assert.Equal([]string{"semantic"}, result.MatchedSignals) + if result.AttachmentID == copyAttachmentID { + require.NotNil(result.OccurredAt) + assert.True(now.Equal(*result.OccurredAt)) + } + } + + bounded, boundedMore, err := f.Store.ResolveDocumentVectorSearchOccurrences( + t.Context(), generation.ID, hits, store.DocumentSearchRequest{}, 1, + ) + require.NoError(err) + require.Len(bounded, 1) + assert.True(boundedMore) + assert.Equal(results[0].OccurrenceKey, bounded[0].OccurrenceKey) + + scoped, scopedMore, err := f.Store.ResolveDocumentVectorSearchOccurrences( + t.Context(), generation.ID, hits, store.DocumentSearchRequest{AttachmentID: copyAttachmentID}, 10, + ) + require.NoError(err) + require.Len(scoped, 1) + assert.False(scopedMore) + assert.Equal(copyAttachmentID, scoped[0].AttachmentID) + + participantID := f.EnsureParticipant("semantic@example.test", "Semantic", "example.test") + var originalMessageID int64 + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind( + `SELECT message_id FROM attachments WHERE content_hash = ? AND id <> ?`), + claims[0].CanonicalBlobHash, copyAttachmentID).Scan(&originalMessageID)) + _, err = f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET sender_id = ? WHERE id = ?`), participantID, originalMessageID) + require.NoError(err) + personResults, personMore, err := f.Store.ResolveDocumentVectorSearchOccurrences( + t.Context(), generation.ID, hits, store.DocumentSearchRequest{Person: &personscope.Scope{ + ParticipantIDs: []int64{participantID}, Directions: []personscope.Direction{personscope.FromPerson}, + }}, 10, + ) + require.NoError(err) + require.Len(personResults, 1) + assert.False(personMore) + assert.Equal(&personscope.Provenance{ + ParticipantIDs: []int64{participantID}, Roles: []personscope.Role{personscope.RoleFrom}, + Directions: []personscope.Direction{personscope.FromPerson}, + }, personResults[0].PersonProvenance) +} + +func TestResolveDocumentVectorSearchOccurrencesBoundsUnicodeExcerpt(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, strings.Repeat("界", 400), "semantic-excerpt") + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("f", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 3, + }) + require.NoError(err) + now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 1) + require.NoError(f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + + results, truncated, err := f.Store.ResolveDocumentVectorSearchOccurrences(t.Context(), generation.ID, []store.DocumentVectorSearchHit{ + {Token: claims[0].Token, Score: .9, Rank: 1}, + }, store.DocumentSearchRequest{}, 10) + require.NoError(err) + require.Len(results, 1) + assert.False(truncated) + assert.Equal(strings.Repeat("界", 320), results[0].Excerpt) + assert.Zero(results[0].HighlightStart) + assert.Zero(results[0].HighlightEnd) +} + +func TestResolveDocumentVectorSearchOccurrencesHidesStaleAuthority(t *testing.T) { + tests := []struct { + name string + mutate func(*testing.T, *storetest.Fixture, store.DocumentVectorChunkClaim) + }{ + { + name: "attachment replacement", + mutate: func(t *testing.T, f *storetest.Fixture, claim store.DocumentVectorChunkClaim) { + t.Helper() + attachmentID := documentVectorAttachmentID(t, f, claim.CanonicalBlobHash) + var messageID int64 + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind( + `SELECT message_id FROM attachments WHERE id = ?`), attachmentID).Scan(&messageID)) + require.NoError(t, f.Store.UpsertAttachmentRecord(t.Context(), messageID, store.AttachmentWrite{ + Filename: "replacement.pdf", MIMEType: "application/pdf", Size: 128, + StoragePath: "ee/" + strings.Repeat("e", 64), ContentHash: strings.Repeat("e", 64), + Role: store.AttachmentRoleStandalone, RoleSource: store.AttachmentRoleSourceImporterSemantics, + SourcePartKey: "mime:1.2", + })) + }, + }, + { + name: "occurrence deletion", + mutate: func(t *testing.T, f *storetest.Fixture, claim store.DocumentVectorChunkClaim) { + t.Helper() + attachmentID := documentVectorAttachmentID(t, f, claim.CanonicalBlobHash) + _, err := f.Store.DB().Exec(f.Store.Rebind(`DELETE FROM attachments WHERE id = ?`), attachmentID) + require.NoError(t, err) + }, + }, + { + name: "role change", + mutate: func(t *testing.T, f *storetest.Fixture, claim store.DocumentVectorChunkClaim) { + t.Helper() + attachmentID := documentVectorAttachmentID(t, f, claim.CanonicalBlobHash) + var messageID int64 + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind( + `SELECT message_id FROM attachments WHERE id = ?`), attachmentID).Scan(&messageID)) + require.NoError(t, f.Store.UpsertAttachmentRecord(t.Context(), messageID, store.AttachmentWrite{ + Filename: "inline.pdf", MIMEType: "application/pdf", Size: 128, + StoragePath: claim.CanonicalBlobHash[:2] + "/" + claim.CanonicalBlobHash, + ContentHash: claim.CanonicalBlobHash, Role: store.AttachmentRoleInline, + RoleSource: store.AttachmentRoleSourceMIMEDisposition, SourcePartKey: "mime:1.2", + })) + }, + }, + { + name: "message lifecycle deletion", + mutate: func(t *testing.T, f *storetest.Fixture, claim store.DocumentVectorChunkClaim) { + t.Helper() + var messageID int64 + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT message_id FROM document_occurrences WHERE canonical_blob_hash = ?`), + claim.CanonicalBlobHash).Scan(&messageID)) + _, err := f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET deleted_from_source_at = CURRENT_TIMESTAMP WHERE id = ?`), messageID) + require.NoError(t, err) + }, + }, + { + name: "target profile rotation", + mutate: func(t *testing.T, f *storetest.Fixture, _ store.DocumentVectorChunkClaim) { + t.Helper() + profile := rotatedDocumentVectorProfile() + _, err := f.Store.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(t, err) + require.NoError(t, f.Store.RecordDocumentProviderConsent(t.Context(), store.DocumentProviderConsent{ + ProfileID: profile.ID, ProfileFingerprint: profile.Fingerprint, + RetentionPosture: profile.RetentionPosture, TrainingPosture: profile.TrainingPosture, + })) + }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 1) + require.NoError(f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + test.mutate(t, f, claims[0]) + + results, truncated, err := f.Store.ResolveDocumentVectorSearchOccurrences(t.Context(), generation.ID, []store.DocumentVectorSearchHit{ + {Token: claims[0].Token, Score: .9, Rank: 1}, + }, store.DocumentSearchRequest{}, 10) + require.NoError(err) + assert.Empty(results) + assert.False(truncated) + var publications int + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT COUNT(*) FROM document_vector_publications WHERE generation_id = ? AND token = ?`), + generation.ID, claims[0].Token).Scan(&publications)) + assert.Equal(1, publications, "authority changes hide but do not erase the token ledger") + }) + } +} + func publishSearchDocument( t *testing.T, f *storetest.Fixture, @@ -555,7 +792,7 @@ func publishSearchDocument( LocalBytes: 128, SourceSequence: 1, })) require.NoError(t, err) - require.NoError(t, f.Store.PublishDocumentExtraction(t.Context(), publicationFor( + require.NoError(t, f.Store.PublishDocumentExtraction(t.Context(), publicationFor(t, claim, text, strings.Repeat("d", 64), ))) } @@ -594,7 +831,7 @@ func publishManySearchChunks( })) require.NoError(t, err) text := "crowded nebula evidence" - publication := publicationFor(claim, text, strings.Repeat("d", 64)) + publication := publicationFor(t, claim, text, strings.Repeat("d", 64)) publication.Chunks = make([]store.DocumentPublishedChunk, count) for index := range count { publication.Chunks[index] = store.DocumentPublishedChunk{ diff --git a/internal/store/document_vector_lock.go b/internal/store/document_vector_lock.go new file mode 100644 index 000000000..3654465a9 --- /dev/null +++ b/internal/store/document_vector_lock.go @@ -0,0 +1,50 @@ +package store + +import ( + "context" + "errors" + "fmt" +) + +const documentVectorOperationLockSQL = `hashtextextended( + current_database() || ':' || current_schema() || ':msgvault.document_vectors', 0)` + +// WithDocumentVectorOperationLock serializes the document-vector writer across +// every process using the same PostgreSQL archive schema. SQLite writers are +// already serialized by the archive ownership lock. +func (s *Store) WithDocumentVectorOperationLock(ctx context.Context, operation func() error) (retErr error) { + if operation == nil { + return errors.New("document vector operation is required") + } + if !s.IsPostgreSQL() { + return operation() + } + conn, err := s.db.Conn(ctx) + if err != nil { + return fmt.Errorf("acquire document vector operation connection: %w", err) + } + locked := false + defer func() { + if locked { + cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), manualTransactionCleanupTimeout) + defer cancel() + var unlocked bool + unlockErr := conn.QueryRowContext( + cleanupCtx, `SELECT pg_advisory_unlock(`+documentVectorOperationLockSQL+`)`, + ).Scan(&unlocked) + if unlockErr != nil { + retErr = errors.Join(retErr, fmt.Errorf("release document vector operation lock: %w", unlockErr)) + } else if !unlocked { + retErr = errors.Join(retErr, errors.New("document vector operation lock was not held during release")) + } + } + if closeErr := conn.Close(); closeErr != nil { + retErr = errors.Join(retErr, fmt.Errorf("close document vector operation connection: %w", closeErr)) + } + }() + if _, err := conn.ExecContext(ctx, `SELECT pg_advisory_lock(`+documentVectorOperationLockSQL+`)`); err != nil { + return fmt.Errorf("acquire document vector operation lock: %w", err) + } + locked = true + return operation() +} diff --git a/internal/store/document_vector_operations.go b/internal/store/document_vector_operations.go new file mode 100644 index 000000000..899177c23 --- /dev/null +++ b/internal/store/document_vector_operations.go @@ -0,0 +1,449 @@ +package store + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" +) + +// DocumentVectorConsentSpec binds hosted-processing consent to both the +// reusable corpus policy and the exact canonical egress destination. +type DocumentVectorConsentSpec struct { + DocumentVectorGenerationSpec + + EgressFingerprint string `json:"egress_fingerprint"` + Purpose string `json:"purpose"` +} + +// DocumentVectorConsent records operator consent for one exact egress policy. +// It never contains credentials or raw provider endpoint data. +type DocumentVectorConsent struct { + DocumentVectorConsentSpec + + ConsentedAt time.Time `json:"consented_at"` +} + +// DocumentVectorUsageDelta is locally observed provider work. It intentionally +// excludes token usage because the provider contract does not report it. +type DocumentVectorUsageDelta struct { + ProviderCalls int64 `json:"provider_calls"` + ProviderDocuments int64 `json:"provider_documents"` + ProviderChunks int64 `json:"provider_chunks"` + ProviderInputChars int64 `json:"provider_input_chars"` +} + +// DocumentVectorProviderUsage is cumulative observed provider work for an +// immutable generation fingerprint. +type DocumentVectorProviderUsage struct { + DocumentVectorUsageDelta + + Fingerprint string `json:"fingerprint"` + UpdatedAt time.Time `json:"updated_at,omitzero"` +} + +// DocumentVectorOperationsStatus is the bounded operator view. When no +// generation is requested, failures deterministically come from the building +// generation first, then the active generation. +type DocumentVectorOperationsStatus struct { + ConfiguredSpec DocumentVectorGenerationSpec `json:"configured_spec"` + ConfiguredDocumentEgressFingerprint string `json:"configured_document_egress_fingerprint"` + ConfiguredQueryEgressFingerprint string `json:"configured_query_egress_fingerprint"` + DocumentConsent *DocumentVectorConsent `json:"document_consent,omitempty"` + QueryConsent *DocumentVectorConsent `json:"query_consent,omitempty"` + Usage DocumentVectorProviderUsage `json:"usage"` + Active *DocumentVectorGeneration `json:"active,omitempty"` + Building *DocumentVectorGeneration `json:"building,omitempty"` + Selected *DocumentVectorGenerationStatus `json:"selected,omitempty"` + Coverage *DocumentVectorCoverage `json:"coverage,omitempty"` +} + +// GetDocumentVectorTargetProfileID returns the configured extraction target +// without resolving any provider credentials. +func (s *Store) GetDocumentVectorTargetProfileID(ctx context.Context) (string, error) { + var target sql.NullString + if err := s.db.QueryRowContext(ctx, `SELECT target_profile_id FROM document_index_state WHERE singleton = 1`).Scan(&target); err != nil { + return "", fmt.Errorf("read document vector target profile: %w", err) + } + if !target.Valid || target.String == "" { + return "", ErrDocumentVectorInvalidGenerationState + } + currentVersion, err := currentDocumentNormalizationVersion() + if err != nil { + return "", err + } + var incompatible bool + if err := s.db.QueryRowContext(ctx, s.Rebind(` + SELECT EXISTS ( + SELECT 1 + FROM document_extraction_heads h + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_index_state ds ON ds.singleton = 1 + WHERE h.profile_id = ? + AND (e.normalization_version IS NULL OR e.normalization_version <> ? + OR e.document_family IS NULL OR TRIM(e.document_family) = '' + OR e.unit_kind IS NULL OR TRIM(e.unit_kind) = '') + AND `+documentVectorLiveAuthoritySQL()+` + )`), target.String, currentVersion).Scan(&incompatible); err != nil { + return "", fmt.Errorf("check document vector target normalized identity: %w", err) + } + if incompatible { + return "", documentNormalizedIdentityRebuildError( + fmt.Sprintf("document vector target profile %q contains an older extraction", target.String), + ErrDocumentNormalizedIdentityUnavailable, + ) + } + return target.String, nil +} + +func (s *Store) RecordDocumentVectorConsent(ctx context.Context, spec DocumentVectorConsentSpec, now time.Time) (DocumentVectorConsent, bool, error) { + if err := validateDocumentVectorGenerationSpec(spec.DocumentVectorGenerationSpec); err != nil { + return DocumentVectorConsent{}, false, err + } + if !documentVectorFingerprintPattern.MatchString(spec.EgressFingerprint) { + return DocumentVectorConsent{}, false, errors.New("document vector consent egress fingerprint is invalid") + } + if spec.Purpose != "document_embedding" && spec.Purpose != "query_embedding" { + return DocumentVectorConsent{}, false, errors.New("document vector consent purpose is invalid") + } + now = normalizeDocumentVectorTime(now) + if now.IsZero() { + return DocumentVectorConsent{}, false, errors.New("document vector consent time is required") + } + var consent DocumentVectorConsent + var created bool + err := s.withTxContext(ctx, func(tx *loggedTx) error { + result, err := tx.ExecContext(ctx, s.Rebind(` + INSERT INTO document_vector_consents + (egress_fingerprint, purpose, generation_fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, consented_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT (egress_fingerprint) DO NOTHING`), + spec.EgressFingerprint, spec.Purpose, spec.Fingerprint, spec.TargetExtractionProfileID, spec.EmbeddingProfile, spec.Model, + spec.Dimension, s.dialect.TimestampParam(now)) + if err != nil { + return fmt.Errorf("record document vector consent: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector consent result: %w", err) + } + created = rows == 1 + return tx.QueryRowContext(ctx, s.Rebind(` + SELECT egress_fingerprint, purpose, generation_fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, consented_at + FROM document_vector_consents WHERE egress_fingerprint = ?`), spec.EgressFingerprint).Scan( + &consent.EgressFingerprint, &consent.Purpose, &consent.Fingerprint, &consent.TargetExtractionProfileID, &consent.EmbeddingProfile, + &consent.Model, &consent.Dimension, &consent.ConsentedAt) + }) + if err != nil { + return DocumentVectorConsent{}, false, err + } + consent.ConsentedAt = normalizeDocumentVectorTime(consent.ConsentedAt) + if consent.DocumentVectorConsentSpec != spec { + return DocumentVectorConsent{}, false, fmt.Errorf("document vector consent egress fingerprint %q collides with a different immutable specification", spec.EgressFingerprint) + } + return consent, created, nil +} + +func (s *Store) GetDocumentVectorConsent(ctx context.Context, egressFingerprint string) (*DocumentVectorConsent, error) { + if !documentVectorFingerprintPattern.MatchString(egressFingerprint) { + return nil, errors.New("document vector consent egress fingerprint is invalid") + } + var consent DocumentVectorConsent + err := s.db.QueryRowContext(ctx, s.Rebind(` + SELECT egress_fingerprint, purpose, generation_fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, consented_at + FROM document_vector_consents WHERE egress_fingerprint = ?`), egressFingerprint).Scan( + &consent.EgressFingerprint, &consent.Purpose, &consent.Fingerprint, &consent.TargetExtractionProfileID, &consent.EmbeddingProfile, + &consent.Model, &consent.Dimension, &consent.ConsentedAt) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil //nolint:nilnil // Absence is a valid optional consent lookup result. + } + if err != nil { + return nil, fmt.Errorf("read document vector consent: %w", err) + } + consent.ConsentedAt = normalizeDocumentVectorTime(consent.ConsentedAt) + return &consent, nil +} + +// CheckpointDocumentVectorBuild atomically advances the durable scan cursor +// and accumulates the observed provider work returned by one completed worker +// invocation. It is not provider billing or an exactly-once crash counter. +func (s *Store) CheckpointDocumentVectorBuild(ctx context.Context, generationID, afterChunkID int64, exhausted bool, delta DocumentVectorUsageDelta, now time.Time) error { + return s.checkpointDocumentVectorBuild(ctx, generationID, "", afterChunkID, exhausted, delta, now) +} + +// CheckpointDocumentVectorBuildForFingerprint preserves observed usage even +// if a concurrently retired generation was purged after Worker.Run returned. +// The cursor remains generation-scoped and is never written after lifecycle loss. +func (s *Store) CheckpointDocumentVectorBuildForFingerprint(ctx context.Context, generationID int64, fingerprint string, afterChunkID int64, exhausted bool, delta DocumentVectorUsageDelta, now time.Time) error { + if !documentVectorFingerprintPattern.MatchString(fingerprint) { + return errors.New("document vector build checkpoint fingerprint is invalid") + } + return s.checkpointDocumentVectorBuild(ctx, generationID, fingerprint, afterChunkID, exhausted, delta, now) +} + +func (s *Store) checkpointDocumentVectorBuild(ctx context.Context, generationID int64, expectedFingerprint string, afterChunkID int64, exhausted bool, delta DocumentVectorUsageDelta, now time.Time) error { + if generationID <= 0 || afterChunkID < 0 || exhausted != (afterChunkID == 0) { + return errors.New("document vector build checkpoint is invalid") + } + if delta.ProviderCalls < 0 || delta.ProviderDocuments < 0 || delta.ProviderChunks < 0 || delta.ProviderInputChars < 0 { + return errors.New("document vector provider usage must be nonnegative") + } + now = normalizeDocumentVectorTime(now) + if now.IsZero() { + return errors.New("document vector build checkpoint time is required") + } + var lifecycleErr error + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return err + } + fingerprint := expectedFingerprint + if found { + if err := q.QueryRow(`SELECT fingerprint FROM document_vector_generations WHERE id = ?`, generationID).Scan(&fingerprint); err != nil { + return fmt.Errorf("read document vector build fingerprint: %w", err) + } + if expectedFingerprint != "" && fingerprint != expectedFingerprint { + return errors.New("document vector build checkpoint fingerprint does not match generation") + } + } else if fingerprint == "" { + return ErrDocumentVectorInvalidGenerationState + } + if _, err := q.Exec(` + INSERT INTO document_vector_provider_usage + (fingerprint, provider_calls, provider_documents, provider_chunks, provider_input_chars, updated_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT (fingerprint) DO UPDATE SET + provider_calls = document_vector_provider_usage.provider_calls + excluded.provider_calls, + provider_documents = document_vector_provider_usage.provider_documents + excluded.provider_documents, + provider_chunks = document_vector_provider_usage.provider_chunks + excluded.provider_chunks, + provider_input_chars = document_vector_provider_usage.provider_input_chars + excluded.provider_input_chars, + updated_at = excluded.updated_at`, fingerprint, delta.ProviderCalls, delta.ProviderDocuments, + delta.ProviderChunks, delta.ProviderInputChars, s.dialect.TimestampParam(now)); err != nil { + return fmt.Errorf("checkpoint document vector provider usage: %w", err) + } + if !found || state != DocumentVectorGenerationBuilding || !currentTarget { + // The provider work already happened. Commit its observed usage even + // when activation, retirement, or target rotation won the lifecycle + // race; only the resumable cursor remains building/current scoped. + lifecycleErr = ErrDocumentVectorInvalidGenerationState + return nil + } + if exhausted { + if _, err := q.Exec(`DELETE FROM document_vector_build_progress WHERE generation_id = ?`, generationID); err != nil { + return fmt.Errorf("reset document vector build cursor: %w", err) + } + return nil + } + if _, err := q.Exec(` + INSERT INTO document_vector_build_progress(generation_id, after_chunk_id, updated_at) + VALUES (?, ?, ?) + ON CONFLICT (generation_id) DO UPDATE SET after_chunk_id = excluded.after_chunk_id, updated_at = excluded.updated_at`, + generationID, afterChunkID, s.dialect.TimestampParam(now)); err != nil { + return fmt.Errorf("checkpoint document vector build cursor: %w", err) + } + return nil + }) + if err != nil { + return err + } + return lifecycleErr +} + +func (s *Store) GetDocumentVectorBuildCursor(ctx context.Context, generationID int64) (int64, error) { + if generationID <= 0 { + return 0, errors.New("document vector generation id must be positive") + } + var after int64 + err := s.db.QueryRowContext(ctx, s.Rebind(`SELECT after_chunk_id FROM document_vector_build_progress WHERE generation_id = ?`), generationID).Scan(&after) + if errors.Is(err, sql.ErrNoRows) { + return 0, nil + } + if err != nil { + return 0, fmt.Errorf("read document vector build cursor: %w", err) + } + return after, nil +} + +func (s *Store) GetDocumentVectorProviderUsage(ctx context.Context, fingerprint string) (DocumentVectorProviderUsage, error) { + if !documentVectorFingerprintPattern.MatchString(fingerprint) { + return DocumentVectorProviderUsage{}, errors.New("document vector usage fingerprint is invalid") + } + usage := DocumentVectorProviderUsage{Fingerprint: fingerprint} + err := s.db.QueryRowContext(ctx, s.Rebind(` + SELECT provider_calls, provider_documents, provider_chunks, provider_input_chars, updated_at + FROM document_vector_provider_usage WHERE fingerprint = ?`), fingerprint).Scan( + &usage.ProviderCalls, &usage.ProviderDocuments, &usage.ProviderChunks, &usage.ProviderInputChars, &usage.UpdatedAt) + if errors.Is(err, sql.ErrNoRows) { + return usage, nil + } + if err != nil { + return DocumentVectorProviderUsage{}, fmt.Errorf("read document vector provider usage: %w", err) + } + usage.UpdatedAt = normalizeDocumentVectorTime(usage.UpdatedAt) + return usage, nil +} + +func (s *Store) GetDocumentVectorOperationsStatus(ctx context.Context, configured DocumentVectorGenerationSpec, documentEgressFingerprint, queryEgressFingerprint string, generationID int64, afterToken string, limit int) (DocumentVectorOperationsStatus, error) { + if err := validateDocumentVectorGenerationSpec(configured); err != nil { + return DocumentVectorOperationsStatus{}, err + } + if generationID < 0 { + return DocumentVectorOperationsStatus{}, errors.New("document vector status generation id is invalid") + } + if !documentVectorFingerprintPattern.MatchString(documentEgressFingerprint) || + !documentVectorFingerprintPattern.MatchString(queryEgressFingerprint) { + return DocumentVectorOperationsStatus{}, errors.New("document vector status egress fingerprint is invalid") + } + result := DocumentVectorOperationsStatus{ + ConfiguredSpec: configured, + ConfiguredDocumentEgressFingerprint: documentEgressFingerprint, + ConfiguredQueryEgressFingerprint: queryEgressFingerprint, + } + var err error + result.DocumentConsent, err = s.GetDocumentVectorConsent(ctx, documentEgressFingerprint) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + result.QueryConsent, err = s.GetDocumentVectorConsent(ctx, queryEgressFingerprint) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + result.Usage, err = s.GetDocumentVectorProviderUsage(ctx, configured.Fingerprint) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + result.Active, err = s.GetActiveDocumentVectorGeneration(ctx) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + result.Building, err = s.GetBuildingDocumentVectorGeneration(ctx) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + selectedID := generationID + if selectedID == 0 && result.Building != nil { + selectedID = result.Building.ID + } else if selectedID == 0 && result.Active != nil { + selectedID = result.Active.ID + } + if selectedID != 0 { + status, err := s.GetDocumentVectorGenerationStatus(ctx, selectedID, afterToken, limit) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + result.Selected = &status + selectedUsesConfiguredTarget := result.Active != nil && result.Active.ID == selectedID && + result.Active.TargetExtractionProfileID == configured.TargetExtractionProfileID + selectedUsesConfiguredTarget = selectedUsesConfiguredTarget || result.Building != nil && result.Building.ID == selectedID && + result.Building.TargetExtractionProfileID == configured.TargetExtractionProfileID + if status.State != DocumentVectorGenerationRetired && selectedUsesConfiguredTarget { + coverage, err := s.GetDocumentVectorCoverage(ctx, selectedID) + if err != nil { + return DocumentVectorOperationsStatus{}, err + } + result.Coverage = &coverage + } + } + return result, nil +} + +// GetOldestRetiredDocumentVectorGeneration returns the next durable cleanup +// ledger. Scheduled convergence handles at most one bounded page per tick. +func (s *Store) GetOldestRetiredDocumentVectorGeneration(ctx context.Context) (*DocumentVectorGeneration, error) { + generation, found, err := scanDocumentVectorGeneration(s.db.QueryRowContext(ctx, s.Rebind(` + SELECT id, fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, + state, created_at, activated_at, retired_at + FROM document_vector_generations WHERE state = ? ORDER BY id LIMIT 1`), string(DocumentVectorGenerationRetired))) + if err != nil { + return nil, err + } + if !found { + return nil, nil //nolint:nilnil // Absence means there is no retired generation awaiting cleanup. + } + return &generation, nil +} + +// StartDocumentVectorRebuild creates a fresh building generation for the +// currently configured exact policy while leaving the active generation live. +func (s *Store) StartDocumentVectorRebuild(ctx context.Context, activeGenerationID int64, desired DocumentVectorGenerationSpec, now time.Time) (DocumentVectorGeneration, error) { + if activeGenerationID <= 0 { + return DocumentVectorGeneration{}, errors.New("active document vector generation id must be positive") + } + if err := validateDocumentVectorGenerationSpec(desired); err != nil { + return DocumentVectorGeneration{}, err + } + now = normalizeDocumentVectorTime(now) + if now.IsZero() { + return DocumentVectorGeneration{}, errors.New("document vector rebuild time is required") + } + var result DocumentVectorGeneration + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, _, err := s.lockDocumentVectorGeneration(q, activeGenerationID) + if err != nil { + return err + } + if state != DocumentVectorGenerationActive { + return ErrDocumentVectorInvalidGenerationState + } + var target string + if err := q.QueryRow(`SELECT target_profile_id FROM document_index_state WHERE singleton = 1`).Scan(&target); err != nil { + return fmt.Errorf("read document vector target profile: %w", err) + } + if desired.TargetExtractionProfileID != target { + return ErrDocumentVectorInvalidGenerationState + } + var building int + if err := q.QueryRow(`SELECT COUNT(*) FROM document_vector_generations WHERE state = ?`, string(DocumentVectorGenerationBuilding)).Scan(&building); err != nil { + return fmt.Errorf("check building document vector generation: %w", err) + } + if building != 0 { + return errors.New("a document vector generation is already building") + } + rows, err := tx.QueryContext(ctx, s.Rebind(` + SELECT target_extraction_profile_id, embedding_profile, model, dimension + FROM document_vector_generations WHERE fingerprint = ?`), desired.Fingerprint) + if err != nil { + return fmt.Errorf("check document vector rebuild fingerprint: %w", err) + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var existing DocumentVectorGenerationSpec + existing.Fingerprint = desired.Fingerprint + if err := rows.Scan(&existing.TargetExtractionProfileID, &existing.EmbeddingProfile, &existing.Model, &existing.Dimension); err != nil { + return fmt.Errorf("scan document vector rebuild fingerprint: %w", err) + } + if existing != desired { + return fmt.Errorf("document vector generation fingerprint %q collides with a different immutable specification", desired.Fingerprint) + } + } + if err := rows.Err(); err != nil { + return err + } + var id int64 + if err := q.QueryRow(`INSERT INTO document_vector_generations + (fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?) RETURNING id`, desired.Fingerprint, desired.TargetExtractionProfileID, + desired.EmbeddingProfile, desired.Model, desired.Dimension, string(DocumentVectorGenerationBuilding), + s.dialect.TimestampParam(now)).Scan(&id); err != nil { + return fmt.Errorf("start document vector rebuild: %w", err) + } + generation, found, err := scanDocumentVectorGeneration(q.QueryRow(` + SELECT id, fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, + state, created_at, activated_at, retired_at + FROM document_vector_generations WHERE id = ?`, id)) + if err != nil { + return err + } + if !found { + return fmt.Errorf("created document vector rebuild %d was not found", id) + } + result = generation + return nil + }) + return result, err +} diff --git a/internal/store/document_vectors.go b/internal/store/document_vectors.go new file mode 100644 index 000000000..d2f34b179 --- /dev/null +++ b/internal/store/document_vectors.go @@ -0,0 +1,1595 @@ +package store + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/binary" + "encoding/hex" + "errors" + "fmt" + "regexp" + "strings" + "time" + + docembedding "go.kenn.io/docbank/document/embedding" +) + +const maxDocumentVectorCandidateLimit = docembedding.MaxCandidateLimit + +var documentVectorFingerprintPattern = regexp.MustCompile(`^[0-9a-f]{64}$`) + +var ( + ErrDocumentVectorClaimLost = errors.New("document vector chunk claim lost") + ErrDocumentVectorSourceChanged = errors.New("document vector chunk source changed") + ErrDocumentVectorInvalidGenerationState = errors.New("document vector generation state is invalid for this operation") + ErrDocumentVectorCoverageIncomplete = errors.New("document vector generation coverage is incomplete") + ErrDocumentVectorGenerationBlocked = errors.New("document vector generation is blocked by terminal failures") + ErrDocumentVectorCleanupIncomplete = errors.New("document vector generation backend cleanup is incomplete") +) + +type DocumentVectorGenerationState string + +const ( + DocumentVectorGenerationBuilding DocumentVectorGenerationState = "building" + DocumentVectorGenerationActive DocumentVectorGenerationState = "active" + DocumentVectorGenerationRetired DocumentVectorGenerationState = "retired" +) + +type DocumentVectorGenerationSpec struct { + Fingerprint string `json:"fingerprint"` + TargetExtractionProfileID string `json:"target_extraction_profile_id"` + EmbeddingProfile string `json:"embedding_profile"` + Model string `json:"model"` + Dimension int `json:"dimension"` +} + +type DocumentVectorGeneration struct { + DocumentVectorGenerationSpec + + ID int64 `json:"id"` + State DocumentVectorGenerationState `json:"state"` + CreatedAt time.Time `json:"created_at"` + ActivatedAt *time.Time `json:"activated_at,omitempty"` + RetiredAt *time.Time `json:"retired_at,omitempty"` +} + +type DocumentVectorChunkCandidate struct { + GenerationID, ChunkID int64 + ExtractionID, ExtractionProfileID, CanonicalBlobHash string + ExtractionInputKey, ChunkKey, ChunkChecksum, Text string + ChunkOrdinal int + SourceSequence int64 +} + +type DocumentVectorChunkClaim struct { + DocumentVectorChunkCandidate + + Token string + LeaseOwner string + LeaseFence int64 + LeaseUntil time.Time + AttemptCount int +} + +// DocumentVectorCoverage compares the current live corpus with ready publications. +type DocumentVectorCoverage struct { + Required int64 `json:"required"` + Ready int64 `json:"ready"` +} + +// Complete reports whether every currently served chunk has a ready publication. +func (c DocumentVectorCoverage) Complete() bool { return c.Required == c.Ready } + +// DocumentVectorLivePublication is a ready token resolved to its current chunk snapshot. +type DocumentVectorLivePublication struct { + DocumentVectorChunkCandidate + + Token string +} + +// DocumentVectorCleanupToken identifies one opaque backend row to delete. +type DocumentVectorCleanupToken struct { + GenerationID int64 + Token string +} + +// DocumentVectorCleanupPage is one durable cleanup-parking page. Exhausted +// pages atomically reset the restorable generation/token cursor pair. +type DocumentVectorCleanupPage struct { + Tokens []DocumentVectorCleanupToken + AfterGenerationID int64 + AfterToken string + Exhausted bool +} + +// DocumentVectorFailureDiagnostic is the bounded, non-PII failure surface. +type DocumentVectorFailureDiagnostic struct { + Token string `json:"token"` + AttemptCount int `json:"attempt_count"` + NextRetryAt *time.Time `json:"next_retry_at,omitempty"` + Terminal bool `json:"terminal"` + ErrorCode string `json:"error_code"` +} + +// DocumentVectorGenerationStatus counts mutually exclusive current +// publication states. Obsolete counts retired or noncurrent snapshots; +// CleanupPending is the uncleaned subset of Obsolete. Failures may include +// obsolete rows so operators can inspect durable failure history. +type DocumentVectorGenerationStatus struct { + GenerationID int64 `json:"generation_id"` + State DocumentVectorGenerationState `json:"state"` + Blocked bool `json:"blocked"` + Pending int64 `json:"pending"` + Retryable int64 `json:"retryable"` + Terminal int64 `json:"terminal"` + ReadyLive int64 `json:"ready_live"` + Obsolete int64 `json:"stale_obsolete"` + CleanupPending int64 `json:"cleanup_pending"` + Failures []DocumentVectorFailureDiagnostic `json:"failures"` + FailureAfterGenerationID int64 `json:"failure_after_generation_id,omitempty"` + FailureAfterToken string `json:"failure_after_token,omitempty"` + FailuresExhausted bool `json:"failures_exhausted"` +} + +// DocumentVectorFailureResetResult reports one bounded stable-token scan. +type DocumentVectorFailureResetResult struct { + Scanned int `json:"scanned"` + Reset int `json:"reset"` + AfterGenerationID int64 `json:"after_generation_id,omitempty"` + AfterToken string `json:"after_token,omitempty"` + Exhausted bool `json:"exhausted"` +} + +func (s *Store) EnsureDocumentVectorGeneration(ctx context.Context, spec DocumentVectorGenerationSpec) (DocumentVectorGeneration, bool, error) { + if err := validateDocumentVectorGenerationSpec(spec); err != nil { + return DocumentVectorGeneration{}, false, err + } + var result DocumentVectorGeneration + created := false + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + // Serialize target rotation before consulting or creating a generation. + // This follows the same index-state -> generation lock order as claim and + // activation, and keeps the target check in the creation transaction. + if _, err := q.Exec(`UPDATE document_index_state SET revision = revision WHERE singleton = 1`); err != nil { + return fmt.Errorf("lock document vector index state: %w", err) + } + var currentTarget sql.NullString + if err := q.QueryRow(`SELECT target_profile_id FROM document_index_state WHERE singleton = 1`).Scan(¤tTarget); err != nil { + return fmt.Errorf("read document vector target profile: %w", err) + } + if !currentTarget.Valid || currentTarget.String != spec.TargetExtractionProfileID { + return ErrDocumentVectorInvalidGenerationState + } + rows, err := tx.QueryContext(ctx, ` + SELECT id, fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, + state, created_at, activated_at, retired_at + FROM document_vector_generations WHERE fingerprint = ? ORDER BY id DESC`, spec.Fingerprint) + if err != nil { + return fmt.Errorf("find document vector generation fingerprint %q: %w", spec.Fingerprint, err) + } + defer func() { _ = rows.Close() }() + var matching *DocumentVectorGeneration + for rows.Next() { + existing, _, scanErr := scanDocumentVectorGeneration(rows) + if scanErr != nil { + return scanErr + } + if existing.DocumentVectorGenerationSpec != spec { + return fmt.Errorf("document vector generation fingerprint %q collides with a different immutable specification", spec.Fingerprint) + } + if matching == nil && existing.State != DocumentVectorGenerationRetired { + matching = &existing + } + } + if err := rows.Err(); err != nil { + return fmt.Errorf("iterate document vector generations: %w", err) + } + if matching != nil { + result = *matching + return nil + } + var building int + if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM document_vector_generations WHERE state = ?`, string(DocumentVectorGenerationBuilding)).Scan(&building); err != nil { + return fmt.Errorf("check building document vector generation: %w", err) + } + if building != 0 { + return errors.New("a document vector generation is already building") + } + var id int64 + if err := tx.QueryRowContext(ctx, `INSERT INTO document_vector_generations + (fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state) + VALUES (?, ?, ?, ?, ?, ?) RETURNING id`, spec.Fingerprint, spec.TargetExtractionProfileID, spec.EmbeddingProfile, spec.Model, spec.Dimension, string(DocumentVectorGenerationBuilding)).Scan(&id); err != nil { + return fmt.Errorf("create document vector generation: %w", err) + } + generation, found, err := scanDocumentVectorGeneration(tx.QueryRowContext(ctx, ` + SELECT id, fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, + state, created_at, activated_at, retired_at + FROM document_vector_generations WHERE id = ?`, id)) + if err != nil { + return err + } + if !found { + return fmt.Errorf("created document vector generation %d was not found", id) + } + result, created = generation, true + return nil + }) + return result, created, err +} + +func validateDocumentVectorGenerationSpec(spec DocumentVectorGenerationSpec) error { + if !documentVectorFingerprintPattern.MatchString(spec.Fingerprint) || strings.TrimSpace(spec.TargetExtractionProfileID) == "" || strings.TrimSpace(spec.Model) == "" { + return errors.New("document vector generation immutable fields are required") + } + if spec.EmbeddingProfile != "vector.embeddings" { + return errors.New("document vector generation embedding profile must be vector.embeddings") + } + if spec.Dimension <= 0 { + return errors.New("document vector generation dimension must be positive") + } + return nil +} + +func (s *Store) GetDocumentVectorGeneration(ctx context.Context, id int64) (DocumentVectorGeneration, error) { + if id <= 0 { + return DocumentVectorGeneration{}, errors.New("document vector generation id must be positive") + } + g, found, err := scanDocumentVectorGeneration(s.db.QueryRowContext(ctx, s.Rebind(`SELECT id, fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state, created_at, activated_at, retired_at FROM document_vector_generations WHERE id = ?`), id)) + if err != nil { + return DocumentVectorGeneration{}, err + } + if !found { + return DocumentVectorGeneration{}, fmt.Errorf("document vector generation %d not found", id) + } + return g, nil +} + +func (s *Store) GetBuildingDocumentVectorGeneration(ctx context.Context) (*DocumentVectorGeneration, error) { + return s.getDocumentVectorGenerationByState(ctx, DocumentVectorGenerationBuilding) +} +func (s *Store) GetActiveDocumentVectorGeneration(ctx context.Context) (*DocumentVectorGeneration, error) { + return s.getDocumentVectorGenerationByState(ctx, DocumentVectorGenerationActive) +} +func (s *Store) getDocumentVectorGenerationByState(ctx context.Context, state DocumentVectorGenerationState) (*DocumentVectorGeneration, error) { + g, found, err := scanDocumentVectorGeneration(s.db.QueryRowContext(ctx, s.Rebind(`SELECT id, fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state, created_at, activated_at, retired_at FROM document_vector_generations WHERE state = ?`), string(state))) + if err != nil { + return nil, err + } + if !found { + return nil, nil //nolint:nilnil // Absence is a valid optional state lookup result. + } + return &g, nil +} + +func (s *Store) ListDocumentVectorChunkCandidates(ctx context.Context, generationID, afterChunkID int64, limit int) ([]DocumentVectorChunkCandidate, error) { + if generationID <= 0 || afterChunkID < 0 { + return nil, errors.New("document vector candidate generation and after chunk id have invalid bounds") + } + if limit < 1 || limit > maxDocumentVectorCandidateLimit { + return nil, fmt.Errorf("document vector candidate limit must be between 1 and %d", maxDocumentVectorCandidateLimit) + } + query := `SELECT dc.id, h.extraction_id, h.profile_id, h.canonical_blob_hash, h.extraction_input_key, + dc.chunk_key, dc.checksum, dc.ordinal, dc.text, e.source_sequence + FROM (SELECT id, target_extraction_profile_id AS target_profile_id, state + FROM document_vector_generations WHERE id = ?) g + JOIN document_index_state ds ON ds.singleton = 1 + JOIN document_extraction_heads h ON h.profile_id = g.target_profile_id + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc ON dc.extraction_id = h.extraction_id + WHERE g.state <> ? AND ds.target_profile_id = g.target_profile_id AND dc.id > ? + AND ` + documentVectorLiveAuthoritySQL() + ` + ORDER BY dc.id LIMIT ?` + rows, err := s.db.QueryContext(ctx, s.Rebind(query), generationID, string(DocumentVectorGenerationRetired), afterChunkID, limit) + if err != nil { + return nil, fmt.Errorf("list document vector candidates: %w", err) + } + defer func() { _ = rows.Close() }() + var out []DocumentVectorChunkCandidate + for rows.Next() { + var c DocumentVectorChunkCandidate + c.GenerationID = generationID + if err := rows.Scan(&c.ChunkID, &c.ExtractionID, &c.ExtractionProfileID, &c.CanonicalBlobHash, &c.ExtractionInputKey, &c.ChunkKey, &c.ChunkChecksum, &c.ChunkOrdinal, &c.Text, &c.SourceSequence); err != nil { + return nil, fmt.Errorf("scan document vector candidate: %w", err) + } + out = append(out, c) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate document vector candidates: %w", err) + } + return out, nil +} + +func (s *Store) ClaimDocumentVectorChunk( + ctx context.Context, + generationID, afterChunkID int64, + scanLimit int, + owner string, + now time.Time, + leaseDuration time.Duration, +) (*DocumentVectorChunkClaim, error) { + if err := validateDocumentVectorClaimRequest(generationID, afterChunkID, scanLimit, owner, now, leaseDuration); err != nil { + return nil, err + } + leaseUntil := normalizeDocumentVectorTime(now.Add(leaseDuration)) + now = normalizeDocumentVectorTime(now) + if !leaseUntil.After(now) { + return nil, errors.New("document vector lease duration has no effective database precision") + } + var claimed *DocumentVectorChunkClaim + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, err := s.lockDocumentVectorGeneration(q, generationID) + if err != nil { + return err + } + if state != DocumentVectorGenerationBuilding || !currentTarget { + return ErrDocumentVectorInvalidGenerationState + } + candidates, err := s.listDocumentVectorChunkCandidatesTx(ctx, tx, generationID, afterChunkID, scanLimit) + if err != nil { + return err + } + for _, candidate := range candidates { + token := documentVectorToken(candidate) + result, err := q.Exec(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, + token, state, lease_owner, lease_fence, lease_until, attempt_count, + created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending', ?, 1, ?, 1, ?, ?) + ON CONFLICT (generation_id, extraction_id, chunk_id) DO NOTHING`, + candidate.GenerationID, candidate.ExtractionID, candidate.ExtractionProfileID, + candidate.CanonicalBlobHash, candidate.ExtractionInputKey, candidate.ChunkID, + candidate.ChunkKey, candidate.ChunkChecksum, candidate.SourceSequence, token, + owner, s.dialect.TimestampParam(leaseUntil), s.dialect.TimestampParam(now), + s.dialect.TimestampParam(now)) + if err != nil { + return fmt.Errorf("create document vector chunk claim: %w", err) + } + inserted, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector chunk claim result: %w", err) + } + if inserted == 1 { + claimed = documentVectorChunkClaim(candidate, token, owner, 1, leaseUntil, 1) + return nil + } + + publication, found, err := s.getDocumentVectorPublicationForClaim(q, candidate) + if err != nil { + return err + } + if !found { + continue + } + if publication.token != token { + return errors.New("document vector publication token does not match its immutable identity") + } + switch publication.state { + case "ready": + continue + case "pending": + if publication.leaseUntil.Valid && publication.leaseUntil.Time.After(now) { + if publication.leaseOwner.String == owner { + claimed = documentVectorChunkClaim(candidate, token, owner, publication.leaseFence, publication.leaseUntil.Time, publication.attemptCount) + return nil + } + continue + } + case "failed": + if !publication.nextRetryAt.Valid || publication.nextRetryAt.Time.After(now) { + continue + } + default: + return fmt.Errorf("document vector publication %q has invalid state %q", token, publication.state) + } + + result, err = q.Exec(` + UPDATE document_vector_publications + SET state = 'pending', lease_owner = ?, lease_fence = lease_fence + 1, + lease_until = ?, attempt_count = attempt_count + 1, + next_retry_at = NULL, error_code = NULL, updated_at = ? + WHERE generation_id = ? AND extraction_id = ? AND chunk_id = ? + AND lease_fence = ? AND ( + (state = 'pending' AND (lease_until IS NULL OR lease_until <= ?)) + OR (state = 'failed' AND next_retry_at IS NOT NULL AND next_retry_at <= ?) + )`, owner, s.dialect.TimestampParam(leaseUntil), s.dialect.TimestampParam(now), + generationID, candidate.ExtractionID, candidate.ChunkID, publication.leaseFence, + s.dialect.TimestampParam(now), s.dialect.TimestampParam(now)) + if err != nil { + return fmt.Errorf("take over document vector chunk claim: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector chunk takeover result: %w", err) + } + if updated == 1 { + claimed = documentVectorChunkClaim(candidate, token, owner, publication.leaseFence+1, leaseUntil, publication.attemptCount+1) + return nil + } + } + return nil + }) + return claimed, err +} + +func (s *Store) RenewDocumentVectorChunkClaim( + ctx context.Context, + generationID int64, + token, owner string, + fence int64, + now time.Time, + leaseDuration time.Duration, +) (time.Time, error) { + if err := validateDocumentVectorClaimReference(generationID, token, owner, fence, now); err != nil { + return time.Time{}, err + } + if leaseDuration <= 0 { + return time.Time{}, errors.New("document vector lease duration must be positive") + } + leaseUntil := normalizeDocumentVectorTime(now.Add(leaseDuration)) + now = normalizeDocumentVectorTime(now) + if !leaseUntil.After(now) { + return time.Time{}, errors.New("document vector lease duration has no effective database precision") + } + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, err := s.lockDocumentVectorGeneration(q, generationID) + if err != nil { + return err + } + if state != DocumentVectorGenerationBuilding || !currentTarget { + return ErrDocumentVectorInvalidGenerationState + } + result, err := q.Exec(` + UPDATE document_vector_publications SET lease_until = ?, updated_at = ? + WHERE generation_id = ? AND token = ? AND state = 'pending' + AND lease_owner = ? AND lease_fence = ? AND lease_until > ?`, + s.dialect.TimestampParam(leaseUntil), s.dialect.TimestampParam(now), generationID, + token, owner, fence, s.dialect.TimestampParam(now)) + if err != nil { + return fmt.Errorf("renew document vector chunk claim: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector chunk renewal result: %w", err) + } + if updated != 1 { + return ErrDocumentVectorClaimLost + } + return nil + }) + if err != nil { + return time.Time{}, err + } + return leaseUntil, nil +} + +func (s *Store) CommitDocumentVectorPublication( + ctx context.Context, + generationID int64, + token, owner string, + fence int64, + now time.Time, +) error { + if err := validateDocumentVectorClaimReference(generationID, token, owner, fence, now); err != nil { + return err + } + now = normalizeDocumentVectorTime(now) + sourceChanged := false + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, err := s.lockDocumentVectorGeneration(q, generationID) + if err != nil { + return err + } + publication, found, err := s.getDocumentVectorPublicationByToken(q, generationID, token) + if err != nil { + return err + } + if !found || publication.leaseOwner.String != owner || publication.leaseFence != fence { + return ErrDocumentVectorClaimLost + } + if publication.state == "ready" { + return nil + } + if publication.state != "pending" || !publication.leaseUntil.Valid || !publication.leaseUntil.Time.After(now) { + return ErrDocumentVectorClaimLost + } + if state != DocumentVectorGenerationBuilding || !currentTarget { + return ErrDocumentVectorInvalidGenerationState + } + current, err := s.isDocumentVectorPublicationCurrent(q, publication) + if err != nil { + return err + } + if !current { + result, err := q.Exec(` + UPDATE document_vector_publications + SET state = 'failed', lease_owner = NULL, lease_until = NULL, + next_retry_at = NULL, error_code = 'source_changed', + backend_cleaned_at = NULL, updated_at = ? + WHERE generation_id = ? AND token = ? AND state = 'pending' + AND lease_owner = ? AND lease_fence = ?`, + s.dialect.TimestampParam(now), generationID, token, owner, fence) + if err != nil { + return fmt.Errorf("record changed document vector source: %w", err) + } + updated, err := result.RowsAffected() + if err != nil || updated != 1 { + return ErrDocumentVectorClaimLost + } + sourceChanged = true + return nil + } + result, err := q.Exec(` + UPDATE document_vector_publications + SET state = 'ready', lease_until = NULL, next_retry_at = NULL, + error_code = NULL, updated_at = ? + WHERE generation_id = ? AND token = ? AND state = 'pending' + AND lease_owner = ? AND lease_fence = ?`, + s.dialect.TimestampParam(now), generationID, token, owner, fence) + if err != nil { + return fmt.Errorf("commit document vector publication: %w", err) + } + updated, err := result.RowsAffected() + if err != nil || updated != 1 { + return ErrDocumentVectorClaimLost + } + return nil + }) + if err != nil { + return err + } + if sourceChanged { + return ErrDocumentVectorSourceChanged + } + return nil +} + +func (s *Store) FailDocumentVectorChunk( + ctx context.Context, + generationID int64, + token, owner string, + fence int64, + now time.Time, + nextRetryAt *time.Time, + terminal bool, + errorCode string, +) error { + if err := validateDocumentVectorClaimReference(generationID, token, owner, fence, now); err != nil { + return err + } + now = normalizeDocumentVectorTime(now) + var normalizedRetryAt *time.Time + if nextRetryAt != nil { + retryAt := normalizeDocumentVectorTime(*nextRetryAt) + normalizedRetryAt = &retryAt + } + if err := validateDocumentVectorFailure(now, normalizedRetryAt, terminal, errorCode); err != nil { + return err + } + var retryAt any + if normalizedRetryAt != nil { + retryAt = s.dialect.TimestampParam(*normalizedRetryAt) + } + return s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, err := s.lockDocumentVectorGeneration(q, generationID) + if err != nil { + return err + } + if state != DocumentVectorGenerationBuilding || !currentTarget { + return ErrDocumentVectorInvalidGenerationState + } + result, err := q.Exec(` + UPDATE document_vector_publications + SET state = 'failed', lease_owner = NULL, lease_until = NULL, + next_retry_at = ?, error_code = ?, updated_at = ? + WHERE generation_id = ? AND token = ? AND state = 'pending' + AND lease_owner = ? AND lease_fence = ? AND lease_until > ?`, + retryAt, errorCode, s.dialect.TimestampParam(now), generationID, token, + owner, fence, s.dialect.TimestampParam(now)) + if err != nil { + return fmt.Errorf("record document vector chunk failure: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector chunk failure result: %w", err) + } + if updated != 1 { + return ErrDocumentVectorClaimLost + } + return nil + }) +} + +// GetDocumentVectorGenerationStatus returns exact ledger counts plus one +// bounded page of stable-token failure diagnostics. +func (s *Store) GetDocumentVectorGenerationStatus( + ctx context.Context, generationID int64, afterToken string, limit int, +) (DocumentVectorGenerationStatus, error) { + if generationID <= 0 { + return DocumentVectorGenerationStatus{}, errors.New("document vector generation id must be positive") + } + if afterToken != "" && !documentVectorFingerprintPattern.MatchString(afterToken) { + return DocumentVectorGenerationStatus{}, errors.New("document vector failure cursor must be a lowercase SHA-256 value") + } + if limit < 1 || limit > maxDocumentVectorCandidateLimit { + return DocumentVectorGenerationStatus{}, fmt.Errorf("document vector failure limit must be between 1 and %d", maxDocumentVectorCandidateLimit) + } + status := DocumentVectorGenerationStatus{GenerationID: generationID} + err := s.withReadSnapshotContext(ctx, func(tx *loggedTx) error { + current := documentVectorPublicationSnapshotCurrentSQL("v", "g") + parked := `(v.state = 'failed' AND v.error_code = 'source_changed')` + query := `SELECT g.state, + COALESCE(SUM(CASE WHEN v.token IS NOT NULL AND g.state <> 'retired' AND NOT (` + parked + `) AND ` + current + ` AND v.state = 'pending' THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN v.token IS NOT NULL AND g.state <> 'retired' AND NOT (` + parked + `) AND ` + current + ` AND v.state = 'failed' AND v.next_retry_at IS NOT NULL THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN v.token IS NOT NULL AND g.state <> 'retired' AND NOT (` + parked + `) AND ` + current + ` AND v.state = 'failed' AND v.next_retry_at IS NULL THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN v.token IS NOT NULL AND g.state <> 'retired' AND NOT (` + parked + `) AND ` + current + ` AND v.state = 'ready' THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN v.token IS NOT NULL AND (g.state = 'retired' OR ` + parked + ` OR NOT (` + current + `)) THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN v.token IS NOT NULL AND v.backend_cleaned_at IS NULL AND (g.state = 'retired' OR ` + parked + ` OR NOT (` + current + `)) THEN 1 ELSE 0 END), 0) + FROM document_vector_generations g + LEFT JOIN document_vector_publications v ON v.generation_id = g.id + WHERE g.id = ? GROUP BY g.state` + if err := tx.QueryRowContext(ctx, query, generationID).Scan( + &status.State, &status.Pending, &status.Retryable, &status.Terminal, + &status.ReadyLive, &status.Obsolete, &status.CleanupPending, + ); errors.Is(err, sql.ErrNoRows) { + return ErrDocumentVectorInvalidGenerationState + } else if err != nil { + return fmt.Errorf("count document vector generation status: %w", err) + } + rows, err := tx.QueryContext(ctx, ` + SELECT token, attempt_count, next_retry_at, error_code + FROM document_vector_publications + WHERE generation_id = ? AND state = 'failed' AND token > ? + ORDER BY token LIMIT ?`, generationID, afterToken, limit) + if err != nil { + return fmt.Errorf("list document vector failure diagnostics: %w", err) + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var diagnostic DocumentVectorFailureDiagnostic + var retryAt sql.NullTime + var errorCode sql.NullString + if err := rows.Scan(&diagnostic.Token, &diagnostic.AttemptCount, &retryAt, &errorCode); err != nil { + return fmt.Errorf("scan document vector failure diagnostic: %w", err) + } + if retryAt.Valid { + normalized := normalizeDocumentVectorTime(retryAt.Time) + diagnostic.NextRetryAt = &normalized + } + diagnostic.Terminal = !retryAt.Valid + diagnostic.ErrorCode = boundedDocumentVectorErrorCode(errorCode.String) + status.Failures = append(status.Failures, diagnostic) + } + if err := rows.Err(); err != nil { + return fmt.Errorf("iterate document vector failure diagnostics: %w", err) + } + if len(status.Failures) < limit { + status.FailuresExhausted = true + status.FailureAfterGenerationID = 0 + status.FailureAfterToken = "" + } else { + status.FailureAfterGenerationID = generationID + status.FailureAfterToken = status.Failures[len(status.Failures)-1].Token + } + status.Blocked = status.State == DocumentVectorGenerationBuilding && status.Terminal > 0 + return nil + }) + return status, err +} + +// ResetDocumentVectorFailures makes one bounded page of current failed +// snapshots claimable again. Source-changed and noncurrent rows remain parked. +func (s *Store) ResetDocumentVectorFailures( + ctx context.Context, generationID int64, afterToken string, limit int, now time.Time, +) (DocumentVectorFailureResetResult, error) { + if generationID <= 0 || now.IsZero() { + return DocumentVectorFailureResetResult{}, errors.New("document vector failure reset requires a generation and time") + } + if afterToken != "" && !documentVectorFingerprintPattern.MatchString(afterToken) { + return DocumentVectorFailureResetResult{}, errors.New("document vector failure reset cursor must be a lowercase SHA-256 value") + } + if limit < 1 || limit > maxDocumentVectorCandidateLimit { + return DocumentVectorFailureResetResult{}, fmt.Errorf("document vector failure reset limit must be between 1 and %d", maxDocumentVectorCandidateLimit) + } + now = normalizeDocumentVectorTime(now) + var reset DocumentVectorFailureResetResult + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, err := s.lockDocumentVectorGeneration(q, generationID) + if err != nil { + return err + } + if state != DocumentVectorGenerationBuilding || !currentTarget { + return ErrDocumentVectorInvalidGenerationState + } + rows, err := tx.QueryContext(ctx, ` + SELECT token, error_code FROM document_vector_publications + WHERE generation_id = ? AND state = 'failed' AND token > ? + ORDER BY token LIMIT ?`, generationID, afterToken, limit) + if err != nil { + return fmt.Errorf("list document vector failures to reset: %w", err) + } + type failure struct{ token, errorCode string } + var failures []failure + for rows.Next() { + var item failure + if err := rows.Scan(&item.token, &item.errorCode); err != nil { + _ = rows.Close() + return fmt.Errorf("scan document vector failure to reset: %w", err) + } + failures = append(failures, item) + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return fmt.Errorf("iterate document vector failures to reset: %w", err) + } + if err := rows.Close(); err != nil { + return fmt.Errorf("close document vector failures to reset: %w", err) + } + reset.Scanned = len(failures) + for _, item := range failures { + if item.errorCode == "source_changed" { + continue + } + publication, found, err := s.getDocumentVectorPublicationByToken(q, generationID, item.token) + if err != nil { + return err + } + if !found || publication.state != "failed" { + continue + } + current, err := s.isDocumentVectorPublicationCurrent(q, publication) + if err != nil { + return err + } + if !current { + continue + } + result, err := q.Exec(` + UPDATE document_vector_publications + SET state = 'pending', lease_owner = NULL, lease_until = NULL, + attempt_count = 0, next_retry_at = NULL, error_code = NULL, + backend_cleaned_at = NULL, updated_at = ? + WHERE generation_id = ? AND token = ? AND state = 'failed'`, + s.dialect.TimestampParam(now), generationID, item.token) + if err != nil { + return fmt.Errorf("reset document vector failure: %w", err) + } + changed, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector failure reset result: %w", err) + } + reset.Reset += int(changed) + } + if len(failures) < limit { + reset.Exhausted = true + reset.AfterGenerationID = 0 + reset.AfterToken = "" + } else { + reset.AfterGenerationID = generationID + reset.AfterToken = failures[len(failures)-1].token + } + return nil + }) + return reset, err +} + +// GetDocumentVectorCoverage counts the current live corpus and its exact ready publications. +func (s *Store) GetDocumentVectorCoverage(ctx context.Context, generationID int64) (DocumentVectorCoverage, error) { + if generationID <= 0 { + return DocumentVectorCoverage{}, errors.New("document vector generation id must be positive") + } + var coverage DocumentVectorCoverage + err := s.withReadSnapshotContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, found, err := getDocumentVectorGenerationState(q, generationID) + if err != nil { + return err + } + if !found || state == DocumentVectorGenerationRetired { + return ErrDocumentVectorInvalidGenerationState + } + coverage, err = s.getDocumentVectorCoverage(q, generationID) + return err + }) + return coverage, err +} + +// ActivateDocumentVectorGeneration atomically swaps in one completely built generation. +func (s *Store) ActivateDocumentVectorGeneration(ctx context.Context, generationID int64, now time.Time) error { + if generationID <= 0 || now.IsZero() { + return errors.New("document vector activation requires a generation and time") + } + now = normalizeDocumentVectorTime(now) + return s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, currentTarget, err := s.lockDocumentVectorGeneration(q, generationID) + if err != nil { + return err + } + if state != DocumentVectorGenerationBuilding || !currentTarget { + return ErrDocumentVectorInvalidGenerationState + } + coverage, err := s.getDocumentVectorCoverage(q, generationID) + if err != nil { + return err + } + if !coverage.Complete() { + return ErrDocumentVectorCoverageIncomplete + } + if _, err := q.Exec(` + UPDATE document_vector_generations + SET state = 'retired', retired_at = ? + WHERE state = 'active' AND id <> ?`, + s.dialect.TimestampParam(now), generationID); err != nil { + return fmt.Errorf("retire prior active document vector generation: %w", err) + } + result, err := q.Exec(` + UPDATE document_vector_generations + SET state = 'active', activated_at = ?, retired_at = NULL + WHERE id = ? AND state = 'building'`, + s.dialect.TimestampParam(now), generationID) + if err != nil { + return fmt.Errorf("activate document vector generation: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector activation result: %w", err) + } + if updated != 1 { + return ErrDocumentVectorInvalidGenerationState + } + return bumpDocumentIndexRevision(q) + }) +} + +// RetireDocumentVectorGeneration removes a building or active generation from service. +func (s *Store) RetireDocumentVectorGeneration(ctx context.Context, generationID int64, now time.Time) (bool, error) { + if generationID <= 0 || now.IsZero() { + return false, errors.New("document vector retirement requires a generation and time") + } + now = normalizeDocumentVectorTime(now) + changed := false + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, _, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return err + } + if !found || state == DocumentVectorGenerationRetired { + return nil + } + if state != DocumentVectorGenerationActive && state != DocumentVectorGenerationBuilding { + return ErrDocumentVectorInvalidGenerationState + } + result, err := q.Exec(` + UPDATE document_vector_generations + SET state = 'retired', retired_at = ? + WHERE id = ? AND state = ?`, s.dialect.TimestampParam(now), generationID, string(state)) + if err != nil { + return fmt.Errorf("retire document vector generation: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector retirement result: %w", err) + } + if updated != 1 { + return ErrDocumentVectorInvalidGenerationState + } + changed = true + if state == DocumentVectorGenerationActive { + return bumpDocumentIndexRevision(q) + } + return nil + }) + return changed, err +} + +// ResolveLiveDocumentVectorPublications filters tokens through current archive authority. +func (s *Store) ResolveLiveDocumentVectorPublications( + ctx context.Context, generationID int64, tokens []string, +) ([]DocumentVectorLivePublication, error) { + if generationID <= 0 { + return nil, errors.New("document vector generation id must be positive") + } + if len(tokens) > maxDocumentVectorCandidateLimit { + return nil, fmt.Errorf("document vector token limit must not exceed %d", maxDocumentVectorCandidateLimit) + } + if len(tokens) == 0 { + return []DocumentVectorLivePublication{}, nil + } + unique := make([]string, 0, len(tokens)) + seen := make(map[string]struct{}, len(tokens)) + for _, token := range tokens { + if !documentVectorFingerprintPattern.MatchString(token) { + return nil, errors.New("document vector token must be a lowercase SHA-256 value") + } + if _, exists := seen[token]; exists { + continue + } + seen[token] = struct{}{} + unique = append(unique, token) + } + args := make([]any, 0, len(unique)+2) + args = append(args, generationID, string(DocumentVectorGenerationActive)) + for _, token := range unique { + args = append(args, token) + } + query := `SELECT v.token, dc.id, v.extraction_id, v.extraction_profile_id, + v.canonical_blob_hash, v.extraction_input_key, v.chunk_key, v.chunk_checksum, + dc.ordinal, dc.text, v.source_sequence + FROM document_vector_publications v + JOIN document_vector_generations g ON g.id = v.generation_id + JOIN document_index_state ds ON ds.singleton = 1 + JOIN document_extraction_heads h + ON h.extraction_id = v.extraction_id + AND h.profile_id = v.extraction_profile_id + AND h.canonical_blob_hash = v.canonical_blob_hash + AND h.extraction_input_key = v.extraction_input_key + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc + ON dc.id = v.chunk_id AND dc.extraction_id = v.extraction_id + AND dc.chunk_key = v.chunk_key AND dc.checksum = v.chunk_checksum + WHERE g.id = ? AND g.state = ? + AND ds.target_profile_id = g.target_extraction_profile_id + AND v.state = 'ready' AND e.source_sequence = v.source_sequence + AND h.source_sequence = v.source_sequence + AND v.token IN (` + documentPlaceholders(len(unique)) + `) + AND ` + documentVectorLiveAuthoritySQL() + rows, err := s.db.QueryContext(ctx, s.Rebind(query), args...) + if err != nil { + return nil, fmt.Errorf("resolve live document vector publications: %w", err) + } + defer func() { _ = rows.Close() }() + byToken := make(map[string]DocumentVectorLivePublication, len(unique)) + for rows.Next() { + var publication DocumentVectorLivePublication + publication.GenerationID = generationID + if err := rows.Scan(&publication.Token, &publication.ChunkID, &publication.ExtractionID, + &publication.ExtractionProfileID, &publication.CanonicalBlobHash, + &publication.ExtractionInputKey, &publication.ChunkKey, &publication.ChunkChecksum, + &publication.ChunkOrdinal, &publication.Text, &publication.SourceSequence); err != nil { + return nil, fmt.Errorf("scan live document vector publication: %w", err) + } + byToken[publication.Token] = publication + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate live document vector publications: %w", err) + } + result := make([]DocumentVectorLivePublication, 0, len(byToken)) + for _, token := range unique { + if publication, found := byToken[token]; found { + result = append(result, publication) + } + } + return result, nil +} + +// ListDocumentVectorCleanupTokens pages uncleaned backend tokens from a retired generation. +func (s *Store) ListDocumentVectorCleanupTokens( + ctx context.Context, generationID int64, afterToken string, limit int, +) ([]DocumentVectorCleanupToken, error) { + if generationID <= 0 { + return nil, errors.New("document vector generation id must be positive") + } + if afterToken != "" && !documentVectorFingerprintPattern.MatchString(afterToken) { + return nil, errors.New("document vector cleanup cursor must be a lowercase SHA-256 value") + } + if limit < 1 || limit > maxDocumentVectorCandidateLimit { + return nil, fmt.Errorf("document vector cleanup limit must be between 1 and %d", maxDocumentVectorCandidateLimit) + } + q := boundQuerier{ctx: ctx, q: s.db} + state, found, err := getDocumentVectorGenerationState(q, generationID) + if err != nil { + return nil, err + } + if !found || state != DocumentVectorGenerationRetired { + return nil, ErrDocumentVectorInvalidGenerationState + } + rows, err := s.db.QueryContext(ctx, s.Rebind(` + SELECT generation_id, token + FROM document_vector_publications + WHERE generation_id = ? AND backend_cleaned_at IS NULL AND token > ? + ORDER BY token LIMIT ?`), generationID, afterToken, limit) + if err != nil { + return nil, fmt.Errorf("list document vector cleanup tokens: %w", err) + } + defer func() { _ = rows.Close() }() + var result []DocumentVectorCleanupToken + for rows.Next() { + var token DocumentVectorCleanupToken + if err := rows.Scan(&token.GenerationID, &token.Token); err != nil { + return nil, fmt.Errorf("scan document vector cleanup token: %w", err) + } + result = append(result, token) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate document vector cleanup tokens: %w", err) + } + return result, nil +} + +// ParkObsoleteDocumentVectorTokens durably parks nonretired obsolete rows +// before their backend vectors are deleted. Already parked rows remain in the +// page even if source authority later returns, making crash replay safe. +func (s *Store) ParkObsoleteDocumentVectorTokens( + ctx context.Context, generationID int64, afterToken string, limit int, now time.Time, +) (DocumentVectorCleanupPage, error) { + if generationID <= 0 || now.IsZero() { + return DocumentVectorCleanupPage{}, errors.New("document vector cleanup parking requires a generation and time") + } + if afterToken != "" && !documentVectorFingerprintPattern.MatchString(afterToken) { + return DocumentVectorCleanupPage{}, errors.New("document vector cleanup parking cursor must be a lowercase SHA-256 value") + } + if limit < 1 || limit > maxDocumentVectorCandidateLimit { + return DocumentVectorCleanupPage{}, fmt.Errorf("document vector cleanup parking limit must be between 1 and %d", maxDocumentVectorCandidateLimit) + } + now = normalizeDocumentVectorTime(now) + var page DocumentVectorCleanupPage + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, _, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return err + } + if !found { + return ErrDocumentVectorInvalidGenerationState + } + current := documentVectorPublicationSnapshotCurrentSQL("v", "g") + query := `SELECT v.generation_id, v.token + FROM document_vector_publications v + JOIN document_vector_generations g ON g.id = v.generation_id + WHERE v.generation_id = ? AND v.backend_cleaned_at IS NULL AND v.token > ? + AND (v.state <> 'pending' OR v.lease_until IS NULL OR v.lease_until <= ?) + AND (g.state = 'retired' + OR (v.state = 'failed' AND v.error_code = 'source_changed') + OR NOT (` + current + `)) + ORDER BY v.token LIMIT ?` + rows, err := tx.QueryContext(ctx, query, generationID, afterToken, s.dialect.TimestampParam(now), limit) + if err != nil { + return fmt.Errorf("list document vector tokens to park: %w", err) + } + for rows.Next() { + var token DocumentVectorCleanupToken + if err := rows.Scan(&token.GenerationID, &token.Token); err != nil { + _ = rows.Close() + return fmt.Errorf("scan document vector token to park: %w", err) + } + page.Tokens = append(page.Tokens, token) + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return fmt.Errorf("iterate document vector tokens to park: %w", err) + } + if err := rows.Close(); err != nil { + return fmt.Errorf("close document vector tokens to park: %w", err) + } + if state != DocumentVectorGenerationRetired { + for _, token := range page.Tokens { + if _, err := q.Exec(` + UPDATE document_vector_publications + SET state = 'failed', lease_owner = NULL, lease_until = NULL, + next_retry_at = NULL, error_code = 'source_changed', + backend_cleaned_at = NULL, updated_at = ? + WHERE generation_id = ? AND token = ? + AND NOT (state = 'failed' AND error_code = 'source_changed')`, + s.dialect.TimestampParam(now), generationID, token.Token); err != nil { + return fmt.Errorf("park obsolete document vector token: %w", err) + } + } + } + if len(page.Tokens) < limit { + page.Exhausted = true + } else { + page.AfterGenerationID = generationID + page.AfterToken = page.Tokens[len(page.Tokens)-1].Token + } + return nil + }) + return page, err +} + +// FinalizeObsoleteDocumentVectorToken acknowledges one backend deletion. +// Nonretired parked rows are removed so a current candidate can be rebuilt; +// retired rows retain their cleanup marker until generation purge. +func (s *Store) FinalizeObsoleteDocumentVectorToken( + ctx context.Context, generationID int64, token string, now time.Time, +) (bool, error) { + if generationID <= 0 || !documentVectorFingerprintPattern.MatchString(token) || now.IsZero() { + return false, errors.New("document vector cleanup finalization is invalid") + } + now = normalizeDocumentVectorTime(now) + changed := false + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, _, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return err + } + if !found { + return nil + } + var result sql.Result + if state == DocumentVectorGenerationRetired { + result, err = q.Exec(` + UPDATE document_vector_publications SET backend_cleaned_at = ?, updated_at = ? + WHERE generation_id = ? AND token = ? AND backend_cleaned_at IS NULL`, + s.dialect.TimestampParam(now), s.dialect.TimestampParam(now), generationID, token) + } else { + result, err = q.Exec(` + DELETE FROM document_vector_publications + WHERE generation_id = ? AND token = ? + AND state = 'failed' AND error_code = 'source_changed'`, generationID, token) + } + if err != nil { + return fmt.Errorf("finalize obsolete document vector token: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read obsolete document vector finalization result: %w", err) + } + changed = updated == 1 + return nil + }) + return changed, err +} + +// MarkDocumentVectorTokenCleaned records one successful backend deletion idempotently. +func (s *Store) MarkDocumentVectorTokenCleaned( + ctx context.Context, generationID int64, token string, now time.Time, +) (bool, error) { + if generationID <= 0 || !documentVectorFingerprintPattern.MatchString(token) || now.IsZero() { + return false, errors.New("document vector cleanup marker is invalid") + } + now = normalizeDocumentVectorTime(now) + changed := false + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, _, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return err + } + if !found || state != DocumentVectorGenerationRetired { + return ErrDocumentVectorInvalidGenerationState + } + result, err := q.Exec(` + UPDATE document_vector_publications SET backend_cleaned_at = ?, updated_at = ? + WHERE generation_id = ? AND token = ? AND backend_cleaned_at IS NULL`, + s.dialect.TimestampParam(now), s.dialect.TimestampParam(now), generationID, token) + if err != nil { + return fmt.Errorf("mark document vector token cleaned: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector cleanup marker result: %w", err) + } + changed = updated == 1 + return nil + }) + return changed, err +} + +// PurgeRetiredDocumentVectorGeneration removes a retired ledger after backend cleanup. +func (s *Store) PurgeRetiredDocumentVectorGeneration(ctx context.Context, generationID int64) (bool, error) { + if generationID <= 0 { + return false, errors.New("document vector generation id must be positive") + } + purged := false + err := s.withTxContext(ctx, func(tx *loggedTx) error { + q := boundQuerier{ctx: ctx, q: tx} + state, _, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return err + } + if !found { + return nil + } + if state != DocumentVectorGenerationRetired { + return ErrDocumentVectorInvalidGenerationState + } + var uncleaned int64 + if err := q.QueryRow(` + SELECT COUNT(*) FROM document_vector_publications + WHERE generation_id = ? AND backend_cleaned_at IS NULL`, generationID).Scan(&uncleaned); err != nil { + return fmt.Errorf("count uncleaned document vector tokens: %w", err) + } + if uncleaned != 0 { + return ErrDocumentVectorCleanupIncomplete + } + if _, err := q.Exec(`DELETE FROM document_vector_publications WHERE generation_id = ?`, generationID); err != nil { + return fmt.Errorf("delete cleaned document vector publications: %w", err) + } + result, err := q.Exec(`DELETE FROM document_vector_generations WHERE id = ? AND state = 'retired'`, generationID) + if err != nil { + return fmt.Errorf("purge retired document vector generation: %w", err) + } + deleted, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("read document vector generation purge result: %w", err) + } + if deleted != 1 { + return ErrDocumentVectorInvalidGenerationState + } + purged = true + return nil + }) + return purged, err +} + +type documentVectorPublication struct { + DocumentVectorChunkCandidate + + token string + state string + leaseOwner sql.NullString + leaseFence int64 + leaseUntil sql.NullTime + attemptCount int + nextRetryAt sql.NullTime +} + +func (s *Store) listDocumentVectorChunkCandidatesTx( + ctx context.Context, tx *loggedTx, generationID, afterChunkID int64, limit int, +) ([]DocumentVectorChunkCandidate, error) { + query := `SELECT dc.id, h.extraction_id, h.profile_id, h.canonical_blob_hash, h.extraction_input_key, + dc.chunk_key, dc.checksum, dc.ordinal, dc.text, e.source_sequence + FROM document_vector_generations g + JOIN document_index_state ds ON ds.singleton = 1 + JOIN document_extraction_heads h ON h.profile_id = g.target_extraction_profile_id + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc ON dc.extraction_id = h.extraction_id + WHERE g.id = ? AND g.state = ? AND ds.target_profile_id = g.target_extraction_profile_id + AND dc.id > ? AND ` + documentVectorLiveAuthoritySQL() + ` + ORDER BY dc.id LIMIT ?` + rows, err := tx.QueryContext(ctx, query, generationID, string(DocumentVectorGenerationBuilding), afterChunkID, limit) + if err != nil { + return nil, fmt.Errorf("list claimable document vector candidates: %w", err) + } + defer func() { _ = rows.Close() }() + var candidates []DocumentVectorChunkCandidate + for rows.Next() { + var candidate DocumentVectorChunkCandidate + candidate.GenerationID = generationID + if err := rows.Scan(&candidate.ChunkID, &candidate.ExtractionID, &candidate.ExtractionProfileID, + &candidate.CanonicalBlobHash, &candidate.ExtractionInputKey, &candidate.ChunkKey, + &candidate.ChunkChecksum, &candidate.ChunkOrdinal, &candidate.Text, &candidate.SourceSequence); err != nil { + return nil, fmt.Errorf("scan claimable document vector candidate: %w", err) + } + candidates = append(candidates, candidate) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate claimable document vector candidates: %w", err) + } + return candidates, nil +} + +func (s *Store) getDocumentVectorCoverage(q boundQuerier, generationID int64) (DocumentVectorCoverage, error) { + query := `SELECT COUNT(*), COALESCE(SUM(CASE WHEN EXISTS ( + SELECT 1 FROM document_vector_publications v + WHERE v.generation_id = g.id AND v.state = 'ready' + AND v.extraction_id = h.extraction_id + AND v.extraction_profile_id = h.profile_id + AND v.canonical_blob_hash = h.canonical_blob_hash + AND v.extraction_input_key = h.extraction_input_key + AND v.chunk_id = dc.id AND v.chunk_key = dc.chunk_key + AND v.chunk_checksum = dc.checksum + AND v.source_sequence = e.source_sequence + AND v.source_sequence = h.source_sequence + ) THEN 1 ELSE 0 END), 0) + FROM document_vector_generations g + JOIN document_index_state ds ON ds.singleton = 1 + JOIN document_extraction_heads h ON h.profile_id = g.target_extraction_profile_id + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc ON dc.extraction_id = h.extraction_id + WHERE g.id = ? AND ds.target_profile_id = g.target_extraction_profile_id + AND ` + documentVectorLiveAuthoritySQL() + var coverage DocumentVectorCoverage + if err := q.QueryRow(query, generationID).Scan(&coverage.Required, &coverage.Ready); err != nil { + return DocumentVectorCoverage{}, fmt.Errorf("count document vector coverage: %w", err) + } + return coverage, nil +} + +func getDocumentVectorGenerationState( + q boundQuerier, generationID int64, +) (DocumentVectorGenerationState, bool, error) { + var state DocumentVectorGenerationState + err := q.QueryRow(`SELECT state FROM document_vector_generations WHERE id = ?`, generationID).Scan(&state) + if errors.Is(err, sql.ErrNoRows) { + return "", false, nil + } + if err != nil { + return "", false, fmt.Errorf("read document vector generation state: %w", err) + } + return state, true, nil +} + +func (s *Store) lockDocumentVectorGeneration( + q boundQuerier, generationID int64, +) (DocumentVectorGenerationState, bool, error) { + state, currentTarget, found, err := s.lockDocumentVectorGenerationIfExists(q, generationID) + if err != nil { + return "", false, err + } + if !found { + return "", false, ErrDocumentVectorInvalidGenerationState + } + return state, currentTarget, nil +} + +func (s *Store) lockDocumentVectorGenerationIfExists( + q boundQuerier, generationID int64, +) (DocumentVectorGenerationState, bool, bool, error) { + // The index-state row is the serialization point for source/profile changes. + // Taking it before the generation/publication row gives claim, commit, and + // the later activation path one lock order on both database backends. + if _, err := q.Exec(`UPDATE document_index_state SET revision = revision WHERE singleton = 1`); err != nil { + return "", false, false, fmt.Errorf("lock document vector index state: %w", err) + } + result, err := q.Exec(`UPDATE document_vector_generations SET state = state WHERE id = ?`, generationID) + if err != nil { + return "", false, false, fmt.Errorf("lock document vector generation: %w", err) + } + updated, err := result.RowsAffected() + if err != nil { + return "", false, false, fmt.Errorf("read document vector generation lock result: %w", err) + } + if updated != 1 { + return "", false, false, nil + } + var state DocumentVectorGenerationState + var currentTarget bool + err = q.QueryRow(` + SELECT g.state, COALESCE(ds.target_profile_id = g.target_extraction_profile_id, FALSE) + FROM document_vector_generations g + JOIN document_index_state ds ON ds.singleton = 1 + WHERE g.id = ?`, generationID).Scan(&state, ¤tTarget) + if errors.Is(err, sql.ErrNoRows) { + return "", false, false, nil + } + if err != nil { + return "", false, false, fmt.Errorf("check document vector generation state: %w", err) + } + return state, currentTarget, true, nil +} + +func (s *Store) getDocumentVectorPublicationForClaim( + q boundQuerier, candidate DocumentVectorChunkCandidate, +) (documentVectorPublication, bool, error) { + query := `SELECT token, state, lease_owner, lease_fence, lease_until, attempt_count, next_retry_at + FROM document_vector_publications + WHERE generation_id = ? AND extraction_id = ? AND chunk_id = ?` + if s.dialect.DriverName() == postgresDriverName { + query += ` FOR UPDATE` + } + var publication documentVectorPublication + publication.DocumentVectorChunkCandidate = candidate + err := q.QueryRow(query, candidate.GenerationID, candidate.ExtractionID, candidate.ChunkID).Scan( + &publication.token, &publication.state, &publication.leaseOwner, &publication.leaseFence, + &publication.leaseUntil, &publication.attemptCount, &publication.nextRetryAt) + if errors.Is(err, sql.ErrNoRows) { + return documentVectorPublication{}, false, nil + } + if err != nil { + return documentVectorPublication{}, false, fmt.Errorf("read document vector publication claim: %w", err) + } + return publication, true, nil +} + +func (s *Store) getDocumentVectorPublicationByToken( + q boundQuerier, generationID int64, token string, +) (documentVectorPublication, bool, error) { + query := `SELECT extraction_id, extraction_profile_id, canonical_blob_hash, extraction_input_key, + chunk_id, chunk_key, chunk_checksum, source_sequence, token, state, lease_owner, + lease_fence, lease_until, attempt_count, next_retry_at + FROM document_vector_publications WHERE generation_id = ? AND token = ?` + if s.dialect.DriverName() == postgresDriverName { + query += ` FOR UPDATE` + } + var publication documentVectorPublication + publication.GenerationID = generationID + err := q.QueryRow(query, generationID, token).Scan( + &publication.ExtractionID, &publication.ExtractionProfileID, &publication.CanonicalBlobHash, + &publication.ExtractionInputKey, &publication.ChunkID, &publication.ChunkKey, + &publication.ChunkChecksum, &publication.SourceSequence, &publication.token, + &publication.state, &publication.leaseOwner, &publication.leaseFence, + &publication.leaseUntil, &publication.attemptCount, &publication.nextRetryAt) + if errors.Is(err, sql.ErrNoRows) { + return documentVectorPublication{}, false, nil + } + if err != nil { + return documentVectorPublication{}, false, fmt.Errorf("read document vector publication: %w", err) + } + return publication, true, nil +} + +func (s *Store) isDocumentVectorPublicationCurrent(q boundQuerier, publication documentVectorPublication) (bool, error) { + query := `SELECT EXISTS ( + SELECT 1 + FROM document_vector_generations g + JOIN document_index_state ds ON ds.singleton = 1 + JOIN document_extraction_heads h ON h.extraction_id = ? + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc ON dc.extraction_id = h.extraction_id + WHERE g.id = ? AND ds.target_profile_id = g.target_extraction_profile_id + AND h.profile_id = ? AND h.canonical_blob_hash = ? AND h.extraction_input_key = ? + AND e.source_sequence = ? AND dc.id = ? AND dc.chunk_key = ? AND dc.checksum = ? + AND ` + documentVectorLiveAuthoritySQL() + ` + )` + var current bool + err := q.QueryRow(query, publication.ExtractionID, publication.GenerationID, + publication.ExtractionProfileID, + publication.CanonicalBlobHash, publication.ExtractionInputKey, publication.SourceSequence, + publication.ChunkID, publication.ChunkKey, publication.ChunkChecksum).Scan(¤t) + if err != nil { + return false, fmt.Errorf("recheck document vector publication source: %w", err) + } + return current, nil +} + +func documentVectorPublicationSnapshotCurrentSQL(publication, generation string) string { + return `EXISTS ( + SELECT 1 + FROM document_index_state ds + JOIN document_extraction_heads h ON h.extraction_id = ` + publication + `.extraction_id + JOIN document_extractions e ON e.id = h.extraction_id + JOIN document_extraction_profiles p ON p.id = h.profile_id + JOIN document_provider_consents c ON c.profile_id = p.id + JOIN document_chunks dc ON dc.id = ` + publication + `.chunk_id AND dc.extraction_id = h.extraction_id + WHERE ds.singleton = 1 + AND ds.target_profile_id = ` + generation + `.target_extraction_profile_id + AND h.profile_id = ` + publication + `.extraction_profile_id + AND h.canonical_blob_hash = ` + publication + `.canonical_blob_hash + AND h.extraction_input_key = ` + publication + `.extraction_input_key + AND e.source_sequence = ` + publication + `.source_sequence + AND h.source_sequence = ` + publication + `.source_sequence + AND dc.chunk_key = ` + publication + `.chunk_key + AND dc.checksum = ` + publication + `.chunk_checksum + AND ` + documentVectorLiveAuthoritySQL() + ` + )` +} + +func documentVectorLiveAuthoritySQL() string { + const head = "h" + + return `EXISTS ( + SELECT 1 FROM document_occurrences o + JOIN attachments a ON a.id = o.attachment_id + JOIN messages m ON m.id = o.message_id + WHERE o.canonical_blob_hash = ` + head + `.canonical_blob_hash + AND ` + documentSearchValidity() + ` + )` +} + +func boundedDocumentVectorErrorCode(value string) string { + if len(value) == 0 || len(value) > 64 { + return "unknown" + } + for _, character := range value { + if character != '_' && character != '-' && (character < 'a' || character > 'z') && + (character < '0' || character > '9') { + return "unknown" + } + } + return value +} + +func documentVectorToken(candidate DocumentVectorChunkCandidate) string { + hasher := sha256.New() + _, _ = hasher.Write([]byte("msgvault-document-vector-token-v1")) + var size [8]byte + binary.BigEndian.PutUint64(size[:], uint64(candidate.GenerationID)) //nolint:gosec // Generation IDs are positive database keys. + _, _ = hasher.Write(size[:]) + for _, field := range []string{candidate.ExtractionID, candidate.ChunkKey, candidate.ChunkChecksum} { + binary.BigEndian.PutUint64(size[:], uint64(len(field))) + _, _ = hasher.Write(size[:]) + _, _ = hasher.Write([]byte(field)) + } + return hex.EncodeToString(hasher.Sum(nil)) +} + +func documentVectorChunkClaim( + candidate DocumentVectorChunkCandidate, + token, owner string, + fence int64, + leaseUntil time.Time, + attemptCount int, +) *DocumentVectorChunkClaim { + return &DocumentVectorChunkClaim{ + DocumentVectorChunkCandidate: candidate, + Token: token, LeaseOwner: owner, LeaseFence: fence, + LeaseUntil: normalizeDocumentVectorTime(leaseUntil), AttemptCount: attemptCount, + } +} + +func normalizeDocumentVectorTime(value time.Time) time.Time { + return value.UTC().Truncate(time.Millisecond) +} + +func validateDocumentVectorClaimRequest( + generationID, afterChunkID int64, + scanLimit int, + owner string, + now time.Time, + leaseDuration time.Duration, +) error { + if generationID <= 0 || afterChunkID < 0 { + return errors.New("document vector claim generation and scan position have invalid bounds") + } + if scanLimit < 1 || scanLimit > maxDocumentVectorCandidateLimit { + return fmt.Errorf("document vector claim scan limit must be between 1 and %d", maxDocumentVectorCandidateLimit) + } + if strings.TrimSpace(owner) == "" || now.IsZero() || leaseDuration <= 0 { + return errors.New("document vector claim owner, time, and lease duration are required") + } + return nil +} + +func validateDocumentVectorClaimReference(generationID int64, token, owner string, fence int64, now time.Time) error { + if generationID <= 0 || !documentVectorFingerprintPattern.MatchString(token) || + strings.TrimSpace(owner) == "" || fence <= 0 || now.IsZero() { + return errors.New("document vector claim reference is invalid") + } + return nil +} + +func validateDocumentVectorFailure(now time.Time, nextRetryAt *time.Time, terminal bool, errorCode string) error { + if len(errorCode) == 0 || len(errorCode) > 64 { + return errors.New("document vector failure error code is invalid") + } + for _, character := range errorCode { + if character != '_' && character != '-' && (character < 'a' || character > 'z') && + (character < '0' || character > '9') { + return errors.New("document vector failure error code is invalid") + } + } + if terminal { + if nextRetryAt != nil { + return errors.New("terminal document vector failure cannot retry") + } + return nil + } + if nextRetryAt == nil || !nextRetryAt.After(now) { + return errors.New("retryable document vector failure requires a future retry time") + } + return nil +} + +func scanDocumentVectorGeneration(row scanner) (DocumentVectorGeneration, bool, error) { + var g DocumentVectorGeneration + var activated, retired sql.NullTime + err := row.Scan(&g.ID, &g.Fingerprint, &g.TargetExtractionProfileID, &g.EmbeddingProfile, &g.Model, &g.Dimension, &g.State, &g.CreatedAt, &activated, &retired) + if errors.Is(err, sql.ErrNoRows) { + return DocumentVectorGeneration{}, false, nil + } + if err != nil { + return DocumentVectorGeneration{}, false, fmt.Errorf("scan document vector generation: %w", err) + } + g.CreatedAt = normalizeDocumentVectorTime(g.CreatedAt) + if activated.Valid { + normalized := normalizeDocumentVectorTime(activated.Time) + g.ActivatedAt = &normalized + } + if retired.Valid { + normalized := normalizeDocumentVectorTime(retired.Time) + g.RetiredAt = &normalized + } + return g, true, nil +} diff --git a/internal/store/document_vectors_pg_test.go b/internal/store/document_vectors_pg_test.go new file mode 100644 index 000000000..5d4738f8d --- /dev/null +++ b/internal/store/document_vectors_pg_test.go @@ -0,0 +1,22 @@ +package store_test + +import ( + "os" + "testing" + + "go.kenn.io/msgvault/internal/store" +) + +func TestDocumentVectorChunkLifecyclePostgreSQLContract(t *testing.T) { + if !store.IsPostgresURL(os.Getenv("MSGVAULT_TEST_DB")) { + t.Skip("PostgreSQL contract runs when MSGVAULT_TEST_DB selects PostgreSQL") + } + runDocumentVectorChunkLifecycleContract(t) +} + +func TestDocumentVectorGenerationLifecyclePostgreSQLContract(t *testing.T) { + if !store.IsPostgresURL(os.Getenv("MSGVAULT_TEST_DB")) { + t.Skip("PostgreSQL contract runs when MSGVAULT_TEST_DB selects PostgreSQL") + } + runDocumentVectorGenerationLifecycleContract(t) +} diff --git a/internal/store/document_vectors_test.go b/internal/store/document_vectors_test.go new file mode 100644 index 000000000..5677303fd --- /dev/null +++ b/internal/store/document_vectors_test.go @@ -0,0 +1,1524 @@ +package store_test + +import ( + "database/sql" + "os" + "sort" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/testutil/storetest" +) + +func TestDocumentVectorChunkLifecycleSQLiteContract(t *testing.T) { + if store.IsPostgresURL(os.Getenv("MSGVAULT_TEST_DB")) { + t.Skip("SQLite contract runs without MSGVAULT_TEST_DB") + } + runDocumentVectorChunkLifecycleContract(t) +} + +func TestDocumentVectorGenerationLifecycleSQLiteContract(t *testing.T) { + if store.IsPostgresURL(os.Getenv("MSGVAULT_TEST_DB")) { + t.Skip("SQLite contract runs without MSGVAULT_TEST_DB") + } + runDocumentVectorGenerationLifecycleContract(t) +} + +func TestDocumentVectorOperationLockSerializesPostgresWriters(t *testing.T) { + requirements := require.New(t) + if !store.IsPostgresURL(os.Getenv("MSGVAULT_TEST_DB")) { + t.Skip("PostgreSQL-only cross-process writer lock") + } + f := storetest.New(t) + firstEntered := make(chan struct{}) + releaseFirst := make(chan struct{}) + firstDone := make(chan error, 1) + go func() { + firstDone <- f.Store.WithDocumentVectorOperationLock(t.Context(), func() error { + close(firstEntered) + <-releaseFirst + return nil + }) + }() + <-firstEntered + + secondEntered := make(chan struct{}) + secondDone := make(chan error, 1) + go func() { + secondDone <- f.Store.WithDocumentVectorOperationLock(t.Context(), func() error { + close(secondEntered) + return nil + }) + }() + select { + case <-secondEntered: + requirements.FailNow("second document-vector writer entered while the first held the lock") + case <-time.After(100 * time.Millisecond): + } + close(releaseFirst) + requirements.NoError(<-firstDone) + select { + case <-secondEntered: + case <-time.After(2 * time.Second): + requirements.FailNow("second document-vector writer did not acquire the released lock") + } + requirements.NoError(<-secondDone) +} + +func runDocumentVectorGenerationLifecycleContract(t *testing.T) { + t.Helper() + t.Run("consent usage progress and same-policy rebuild", testDocumentVectorOperationsState) + t.Run("coverage activation and rollback", testDocumentVectorCoverageAndActivation) + t.Run("atomic generation swap and retirement", testDocumentVectorActivationSwapAndRetirement) + t.Run("live token ordering and invalidation", testDocumentVectorLiveResolution) + t.Run("derivative garbage collection preserves cleanup token", testDocumentVectorDerivativeGarbageCollection) + t.Run("derivative purge preserves cleanup token", testDocumentVectorDerivativePurge) + t.Run("retired token cleanup and purge", testDocumentVectorCleanupAndPurge) + t.Run("obsolete cleanup rechecks live snapshot", testDocumentVectorObsoleteCleanup) +} + +func testDocumentVectorOperationsState(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 12, 34, 56, 987654321, time.FixedZone("consent", 2*60*60)) + require.NoError(t, f.Store.InitSchema(), "the document-vector base schema is replay-safe") + if f.Store.IsPostgreSQL() { + var indexDefinition string + require.NoError(t, f.Store.DB().QueryRow(` + SELECT indexdef FROM pg_indexes + WHERE schemaname = current_schema() AND indexname = 'idx_document_vector_generations_live_fingerprint'`).Scan(&indexDefinition)) + assert.NotContains(t, indexDefinition, "UNIQUE") + } else { + rows, err := f.Store.DB().Query(`PRAGMA index_list(document_vector_generations)`) + require.NoError(t, err) + defer func() { _ = rows.Close() }() + found := false + for rows.Next() { + var unique int + var sequence, partial int + var name, origin string + require.NoError(t, rows.Scan(&sequence, &name, &unique, &origin, &partial)) + if name == "idx_document_vector_generations_live_fingerprint" { + found = true + assert.Zero(t, unique) + break + } + } + require.NoError(t, rows.Err()) + require.True(t, found, "document vector generation fingerprint index") + } + + egressFingerprint := strings.Repeat("e", 64) + consentSpec := store.DocumentVectorConsentSpec{ + DocumentVectorGenerationSpec: generation.DocumentVectorGenerationSpec, + EgressFingerprint: egressFingerprint, + Purpose: "document_embedding", + } + consent, created, err := f.Store.RecordDocumentVectorConsent(t.Context(), consentSpec, now) + require.NoError(t, err) + assert.True(t, created) + assert.Equal(t, generation.DocumentVectorGenerationSpec, consent.DocumentVectorGenerationSpec) + assert.Equal(t, now.UTC().Truncate(time.Millisecond), consent.ConsentedAt) + + gotConsent, err := f.Store.GetDocumentVectorConsent(t.Context(), egressFingerprint) + require.NoError(t, err) + require.NotNil(t, gotConsent) + assert.Equal(t, consent, *gotConsent) + + collision := consentSpec + collision.Model = "different-model" + _, _, err = f.Store.RecordDocumentVectorConsent(t.Context(), collision, now) + require.ErrorContains(t, err, "fingerprint") + + err = f.Store.CheckpointDocumentVectorBuild(t.Context(), generation.ID, 42, false, store.DocumentVectorUsageDelta{ + ProviderCalls: 1, ProviderDocuments: 2, ProviderChunks: 3, ProviderInputChars: 123, + }, now) + require.NoError(t, err) + cursor, err := f.Store.GetDocumentVectorBuildCursor(t.Context(), generation.ID) + require.NoError(t, err) + assert.Equal(t, int64(42), cursor) + usage, err := f.Store.GetDocumentVectorProviderUsage(t.Context(), generation.Fingerprint) + require.NoError(t, err) + assert.Equal(t, int64(1), usage.ProviderCalls) + assert.Equal(t, int64(2), usage.ProviderDocuments) + assert.Equal(t, int64(3), usage.ProviderChunks) + assert.Equal(t, int64(123), usage.ProviderInputChars) + queryEgressFingerprint := strings.Repeat("d", 64) + operations, err := f.Store.GetDocumentVectorOperationsStatus(t.Context(), generation.DocumentVectorGenerationSpec, egressFingerprint, queryEgressFingerprint, 0, "", 10) + require.NoError(t, err) + require.NotNil(t, operations.DocumentConsent) + assert.Nil(t, operations.QueryConsent) + require.NotNil(t, operations.Coverage) + assert.Equal(t, store.DocumentVectorCoverage{Required: 1, Ready: 0}, *operations.Coverage) + + require.NoError(t, f.Store.CheckpointDocumentVectorBuild(t.Context(), generation.ID, 0, true, store.DocumentVectorUsageDelta{}, now.Add(time.Second))) + cursor, err = f.Store.GetDocumentVectorBuildCursor(t.Context(), generation.ID) + require.NoError(t, err) + assert.Zero(t, cursor) + require.ErrorContains(t, f.Store.CheckpointDocumentVectorBuild(t.Context(), generation.ID, 1, false, store.DocumentVectorUsageDelta{}, time.Time{}), "time") + + readyAllDocumentVectorChunks(t, f, generation, now) + require.NoError(t, f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + err = f.Store.CheckpointDocumentVectorBuild(t.Context(), generation.ID, 99, false, store.DocumentVectorUsageDelta{ + ProviderCalls: 2, ProviderDocuments: 3, ProviderChunks: 4, ProviderInputChars: 456, + }, now.Add(1500*time.Millisecond)) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + usage, err = f.Store.GetDocumentVectorProviderUsage(t.Context(), generation.Fingerprint) + require.NoError(t, err) + assert.Equal(t, int64(3), usage.ProviderCalls, "completed worker usage survives concurrent activation") + assert.Equal(t, int64(5), usage.ProviderDocuments) + assert.Equal(t, int64(7), usage.ProviderChunks) + assert.Equal(t, int64(579), usage.ProviderInputChars) + cursor, err = f.Store.GetDocumentVectorBuildCursor(t.Context(), generation.ID) + require.NoError(t, err) + assert.Zero(t, cursor, "an active generation cursor is not advanced") + _, err = f.Store.StartDocumentVectorRebuild(t.Context(), generation.ID, generation.DocumentVectorGenerationSpec, time.Time{}) + require.ErrorContains(t, err, "time") + rebuild, err := f.Store.StartDocumentVectorRebuild(t.Context(), generation.ID, generation.DocumentVectorGenerationSpec, now.Add(2*time.Second)) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorGenerationBuilding, rebuild.State) + assert.Equal(t, generation.Fingerprint, rebuild.Fingerprint) + assert.Equal(t, now.Add(2*time.Second).UTC().Truncate(time.Millisecond), rebuild.CreatedAt) + resumed, created, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), generation.DocumentVectorGenerationSpec) + require.NoError(t, err) + assert.False(t, created) + assert.Equal(t, rebuild.ID, resumed.ID) + active, err := f.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + require.NotNil(t, active) + assert.Equal(t, generation.ID, active.ID) + + t.Run("configured target rotation builds beside the old active generation", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + rotatedFixture, old := seedDocumentVectorGenerationWithChunks(t, 1) + readyAllDocumentVectorChunks(t, rotatedFixture, old, now.Add(-time.Second)) + require.NoError(rotatedFixture.Store.ActivateDocumentVectorGeneration(t.Context(), old.ID, now)) + profile := rotatedDocumentVectorProfile() + _, err := rotatedFixture.Store.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(err) + require.NoError(rotatedFixture.Store.RecordDocumentProviderConsent(t.Context(), store.DocumentProviderConsent{ + ProfileID: profile.ID, ProfileFingerprint: profile.Fingerprint, + RetentionPosture: profile.RetentionPosture, TrainingPosture: profile.TrainingPosture, + })) + _, err = rotatedFixture.Store.DB().Exec(rotatedFixture.Store.Rebind(` + UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), profile.ID) + require.NoError(err) + desired := old.DocumentVectorGenerationSpec + desired.Fingerprint = strings.Repeat("3", 64) + desired.TargetExtractionProfileID = profile.ID + operations, err := rotatedFixture.Store.GetDocumentVectorOperationsStatus(t.Context(), desired, strings.Repeat("d", 64), strings.Repeat("e", 64), 0, "", 10) + require.NoError(err) + require.NotNil(operations.Selected) + assert.Equal(old.ID, operations.Selected.GenerationID) + assert.Nil(operations.Coverage, "coverage is undefined for the old target") + building, err := rotatedFixture.Store.StartDocumentVectorRebuild(t.Context(), old.ID, desired, now.Add(time.Second)) + require.NoError(err) + assert.Equal(desired, building.DocumentVectorGenerationSpec) + stillActive, err := rotatedFixture.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(err) + require.NotNil(stillActive) + assert.Equal(old.ID, stillActive.ID) + }) + + t.Run("observed usage survives target rotation and retirement without advancing cursor", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + rotatedFixture, building := seedDocumentVectorGenerationWithChunks(t, 1) + require.NoError(rotatedFixture.Store.CheckpointDocumentVectorBuild( + t.Context(), building.ID, 7, false, store.DocumentVectorUsageDelta{}, now)) + profile := rotatedDocumentVectorProfile() + _, err := rotatedFixture.Store.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(err) + _, err = rotatedFixture.Store.DB().Exec(rotatedFixture.Store.Rebind( + `UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), profile.ID) + require.NoError(err) + err = rotatedFixture.Store.CheckpointDocumentVectorBuild(t.Context(), building.ID, 99, false, + store.DocumentVectorUsageDelta{ProviderCalls: 1, ProviderChunks: 2, ProviderInputChars: 30}, now.Add(time.Second)) + require.ErrorIs(err, store.ErrDocumentVectorInvalidGenerationState) + cursor, err := rotatedFixture.Store.GetDocumentVectorBuildCursor(t.Context(), building.ID) + require.NoError(err) + assert.Equal(int64(7), cursor) + retired, err := rotatedFixture.Store.RetireDocumentVectorGeneration(t.Context(), building.ID, now.Add(2*time.Second)) + require.NoError(err) + assert.True(retired) + err = rotatedFixture.Store.CheckpointDocumentVectorBuild(t.Context(), building.ID, 101, false, + store.DocumentVectorUsageDelta{ProviderCalls: 2, ProviderDocuments: 1, ProviderInputChars: 40}, now.Add(3*time.Second)) + require.ErrorIs(err, store.ErrDocumentVectorInvalidGenerationState) + usage, err := rotatedFixture.Store.GetDocumentVectorProviderUsage(t.Context(), building.Fingerprint) + require.NoError(err) + assert.Equal(int64(3), usage.ProviderCalls) + assert.Equal(int64(1), usage.ProviderDocuments) + assert.Equal(int64(2), usage.ProviderChunks) + assert.Equal(int64(70), usage.ProviderInputChars) + cursor, err = rotatedFixture.Store.GetDocumentVectorBuildCursor(t.Context(), building.ID) + require.NoError(err) + assert.Equal(int64(7), cursor) + purged, err := rotatedFixture.Store.PurgeRetiredDocumentVectorGeneration(t.Context(), building.ID) + require.NoError(err) + assert.True(purged) + err = rotatedFixture.Store.CheckpointDocumentVectorBuildForFingerprint(t.Context(), building.ID, building.Fingerprint, + 0, true, store.DocumentVectorUsageDelta{ProviderCalls: 1, ProviderInputChars: 5}, now.Add(4*time.Second)) + require.ErrorIs(err, store.ErrDocumentVectorInvalidGenerationState) + usage, err = rotatedFixture.Store.GetDocumentVectorProviderUsage(t.Context(), building.Fingerprint) + require.NoError(err) + assert.Equal(int64(4), usage.ProviderCalls, "purged-generation completion retains observed usage by fingerprint") + assert.Equal(int64(75), usage.ProviderInputChars) + }) +} + +func testDocumentVectorObsoleteCleanup(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 23, 15, 0, 654321000, time.FixedZone("cleanup-offset", -2*60*60)) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(t, claims, 1) + claim := claims[0] + var messageID int64 + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT message_id FROM document_occurrences WHERE canonical_blob_hash = ?`), claim.CanonicalBlobHash).Scan(&messageID)) + _, err := f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET deleted_from_source_at = ? WHERE id = ?`), now, messageID) + require.NoError(t, err) + page, err := f.Store.ParkObsoleteDocumentVectorTokens(t.Context(), generation.ID, "", 10, now.Add(time.Second)) + require.NoError(t, err) + assert.Equal(t, []store.DocumentVectorCleanupToken{{GenerationID: generation.ID, Token: claim.Token}}, page.Tokens) + assert.True(t, page.Exhausted) + assert.Zero(t, page.AfterGenerationID) + assert.Empty(t, page.AfterToken) + + _, err = f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET deleted_from_source_at = NULL WHERE id = ?`), messageID) + require.NoError(t, err) + status, err := f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 10) + require.NoError(t, err) + assert.Equal(t, int64(1), status.Obsolete) + assert.Equal(t, int64(1), status.CleanupPending) + assert.Zero(t, status.ReadyLive, "a parked row cannot become live when authority returns") + replayed, err := f.Store.ParkObsoleteDocumentVectorTokens(t.Context(), generation.ID, "", 10, now.Add(2*time.Second)) + require.NoError(t, err) + assert.Equal(t, page.Tokens, replayed.Tokens, "a crash after parking must replay the durable token") + + finalized, err := f.Store.FinalizeObsoleteDocumentVectorToken(t.Context(), generation.ID, claim.Token, now.Add(3*time.Second)) + require.NoError(t, err) + assert.True(t, finalized) + finalized, err = f.Store.FinalizeObsoleteDocumentVectorToken(t.Context(), generation.ID, claim.Token, now.Add(4*time.Second)) + require.NoError(t, err) + assert.False(t, finalized) + reclaimed, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "rebuild-worker", now.Add(5*time.Second), time.Minute) + require.NoError(t, err) + require.NotNil(t, reclaimed) + assert.Equal(t, claim.Token, reclaimed.Token) + assert.Equal(t, 1, reclaimed.AttemptCount) + + t.Run("retired makes every token obsolete", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 1) + retired, err := f.Store.RetireDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second)) + require.NoError(err) + require.True(retired) + page, err := f.Store.ParkObsoleteDocumentVectorTokens(t.Context(), generation.ID, "", 10, now.Add(2*time.Second)) + require.NoError(err) + assert.Equal([]store.DocumentVectorCleanupToken{{GenerationID: generation.ID, Token: claims[0].Token}}, page.Tokens) + finalized, err := f.Store.FinalizeObsoleteDocumentVectorToken(t.Context(), generation.ID, claims[0].Token, now.Add(3*time.Second)) + require.NoError(err) + assert.True(finalized) + var cleanedAt sql.NullTime + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT backend_cleaned_at FROM document_vector_publications + WHERE generation_id = ? AND token = ?`), generation.ID, claims[0].Token).Scan(&cleanedAt)) + assert.True(cleanedAt.Valid, "retired finalization preserves the ledger for purge") + }) + + t.Run("retired waits for a live publication lease", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "retiring-worker", now, time.Minute) + require.NoError(err) + require.NotNil(claim) + renewedUntil, err := f.Store.RenewDocumentVectorChunkClaim( + t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, + now.Add(59*time.Second), time.Minute, + ) + require.NoError(err) + assert.Equal(now.Add(119*time.Second).UTC().Truncate(time.Millisecond), renewedUntil) + retired, err := f.Store.RetireDocumentVectorGeneration(t.Context(), generation.ID, now.Add(60*time.Second)) + require.NoError(err) + require.True(retired) + beforeExpiry, err := f.Store.ParkObsoleteDocumentVectorTokens(t.Context(), generation.ID, "", 10, renewedUntil.Add(-time.Millisecond)) + require.NoError(err) + assert.Empty(beforeExpiry.Tokens) + atExpiry, err := f.Store.ParkObsoleteDocumentVectorTokens(t.Context(), generation.ID, "", 10, renewedUntil) + require.NoError(err) + assert.Equal([]store.DocumentVectorCleanupToken{{GenerationID: generation.ID, Token: claim.Token}}, atExpiry.Tokens) + }) + + t.Run("active generation parks before delete", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 1) + require.NoError(f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + var messageID int64 + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT message_id FROM document_occurrences WHERE canonical_blob_hash = ?`), claims[0].CanonicalBlobHash).Scan(&messageID)) + _, err := f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET deleted_from_source_at = ? WHERE id = ?`), now, messageID) + require.NoError(err) + page, err := f.Store.ParkObsoleteDocumentVectorTokens(t.Context(), generation.ID, "", 10, now.Add(2*time.Second)) + require.NoError(err) + assert.Equal([]store.DocumentVectorCleanupToken{{GenerationID: generation.ID, Token: claims[0].Token}}, page.Tokens) + _, err = f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET deleted_from_source_at = NULL WHERE id = ?`), messageID) + require.NoError(err) + live, err := f.Store.ResolveLiveDocumentVectorPublications(t.Context(), generation.ID, []string{claims[0].Token}) + require.NoError(err) + assert.Empty(live) + finalized, err := f.Store.FinalizeObsoleteDocumentVectorToken(t.Context(), generation.ID, claims[0].Token, now.Add(3*time.Second)) + require.NoError(err) + assert.True(finalized) + var publications int + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT COUNT(*) FROM document_vector_publications WHERE generation_id = ?`), generation.ID).Scan(&publications)) + assert.Zero(publications) + }) +} + +func testDocumentVectorDerivativeGarbageCollection(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 22, 0, 0, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(t, claims, 1) + require.NoError(t, f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + claim := claims[0] + + var attachmentID, messageID int64 + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT attachment_id, message_id FROM document_occurrences + WHERE canonical_blob_hash = ?`), claim.CanonicalBlobHash).Scan(&attachmentID, &messageID)) + _, err := f.Store.DB().Exec(f.Store.Rebind( + `UPDATE messages SET deleted_from_source_at = CURRENT_TIMESTAMP WHERE id = ?`), messageID) + require.NoError(t, err) + _, eligible, err := f.Store.ReconcileDocumentOccurrence(t.Context(), attachmentID, claim.SourceSequence+1) + require.NoError(t, err) + assert.False(t, eligible) + _, err = f.Store.DB().Exec(f.Store.Rebind( + `UPDATE document_extractions SET updated_at = ? WHERE id = ?`), now.Add(-48*time.Hour), claim.ExtractionID) + require.NoError(t, err) + + result, err := f.Store.GarbageCollectDocumentDerivatives(t.Context(), now.Add(-24*time.Hour), 10) + require.NoError(t, err) + assert.Equal(t, store.DocumentDerivativeGCResult{ExtractionsRemoved: 1, CurrentHeadsRemoved: 1}, result) + assertDocumentVectorPublicationSurvivesSourceRemoval(t, f, generation.ID, claim.Token) +} + +func testDocumentVectorDerivativePurge(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 22, 30, 0, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(t, claims, 1) + require.NoError(t, f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + claim := claims[0] + + result, err := f.Store.PurgeDocumentDerivedByHash(t.Context(), claim.CanonicalBlobHash) + require.NoError(t, err) + assert.Equal(t, store.DocumentDerivedPurgeResult{ExtractionsRemoved: 1, HeadsRemoved: 1}, result) + assertDocumentVectorPublicationSurvivesSourceRemoval(t, f, generation.ID, claim.Token) +} + +func assertDocumentVectorPublicationSurvivesSourceRemoval( + t *testing.T, f *storetest.Fixture, generationID int64, token string, +) { + t.Helper() + var storedToken, state string + var cleanedAt sql.NullTime + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT token, state, backend_cleaned_at + FROM document_vector_publications + WHERE generation_id = ? AND token = ?`), generationID, token).Scan( + &storedToken, &state, &cleanedAt)) + assert.Equal(t, token, storedToken) + assert.Equal(t, "ready", state) + assert.False(t, cleanedAt.Valid) + resolved, err := f.Store.ResolveLiveDocumentVectorPublications(t.Context(), generationID, []string{token}) + require.NoError(t, err) + assert.Empty(t, resolved) +} + +func testDocumentVectorCoverageAndActivation(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 2) + now := time.Date(2026, time.August, 20, 18, 19, 20, 456789000, time.FixedZone("coverage-offset", -3*60*60)) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "coverage-worker", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + require.NoError(t, f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(time.Second))) + + coverage, err := f.Store.GetDocumentVectorCoverage(t.Context(), generation.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorCoverage{Required: 2, Ready: 1}, coverage) + assert.False(t, coverage.Complete()) + revisionBefore, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + err = f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(2*time.Second)) + require.ErrorIs(t, err, store.ErrDocumentVectorCoverageIncomplete) + afterFailedActivation, err := f.Store.GetDocumentVectorGeneration(t.Context(), generation.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorGenerationBuilding, afterFailedActivation.State) + assert.Nil(t, afterFailedActivation.ActivatedAt) + revisionAfterFailure, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + assert.Equal(t, revisionBefore, revisionAfterFailure) + + remaining := readyAllDocumentVectorChunks(t, f, generation, now.Add(3*time.Second)) + require.Len(t, remaining, 1) + coverage, err = f.Store.GetDocumentVectorCoverage(t.Context(), generation.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorCoverage{Required: 2, Ready: 2}, coverage) + assert.True(t, coverage.Complete()) + + activationTime := now.Add(4 * time.Second) + require.NoError(t, f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, activationTime)) + active, err := f.Store.GetDocumentVectorGeneration(t.Context(), generation.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorGenerationActive, active.State) + require.NotNil(t, active.ActivatedAt) + assert.Equal(t, activationTime.UTC().Truncate(time.Millisecond), active.ActivatedAt.UTC()) + revisionAfterActivation, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + assert.Equal(t, revisionBefore+1, revisionAfterActivation) + + empty := storetest.New(t) + profile, _ := seedDocumentPublicationAuthority(t, empty) + emptyGeneration, _, err := empty.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("0", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.NoError(t, err) + emptyCoverage, err := empty.Store.GetDocumentVectorCoverage(t.Context(), emptyGeneration.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorCoverage{}, emptyCoverage) + assert.True(t, emptyCoverage.Complete()) + require.NoError(t, empty.Store.ActivateDocumentVectorGeneration(t.Context(), emptyGeneration.ID, activationTime)) +} + +func testDocumentVectorActivationSwapAndRetirement(t *testing.T) { + f, oldGeneration := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 19, 20, 21, 987654000, time.UTC) + oldClaims := readyAllDocumentVectorChunks(t, f, oldGeneration, now) + require.Len(t, oldClaims, 1) + require.NoError(t, f.Store.ActivateDocumentVectorGeneration(t.Context(), oldGeneration.ID, now.Add(time.Second))) + + newSpec := oldGeneration.DocumentVectorGenerationSpec + newSpec.Fingerprint = strings.Repeat("1", 64) + newSpec.Model = "embed-v2" + newGeneration, created, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), newSpec) + require.NoError(t, err) + require.True(t, created) + err = f.Store.ActivateDocumentVectorGeneration(t.Context(), newGeneration.ID, now.Add(2*time.Second)) + require.ErrorIs(t, err, store.ErrDocumentVectorCoverageIncomplete) + stillActive, err := f.Store.GetActiveDocumentVectorGeneration(t.Context()) + require.NoError(t, err) + require.NotNil(t, stillActive) + assert.Equal(t, oldGeneration.ID, stillActive.ID, "incomplete activation must not retire the serving generation") + newClaims := readyAllDocumentVectorChunks(t, f, newGeneration, now.Add(2*time.Second)) + require.Len(t, newClaims, 1) + revisionBeforeSwap, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + swapTime := now.Add(3 * time.Second) + require.NoError(t, f.Store.ActivateDocumentVectorGeneration(t.Context(), newGeneration.ID, swapTime)) + + oldAfter, err := f.Store.GetDocumentVectorGeneration(t.Context(), oldGeneration.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorGenerationRetired, oldAfter.State) + require.NotNil(t, oldAfter.RetiredAt) + assert.Equal(t, swapTime.UTC().Truncate(time.Millisecond), oldAfter.RetiredAt.UTC()) + newAfter, err := f.Store.GetDocumentVectorGeneration(t.Context(), newGeneration.ID) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorGenerationActive, newAfter.State) + revisionAfterSwap, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + assert.Equal(t, revisionBeforeSwap+1, revisionAfterSwap) + + resolved, err := f.Store.ResolveLiveDocumentVectorPublications(t.Context(), oldGeneration.ID, []string{oldClaims[0].Token}) + require.NoError(t, err) + assert.Empty(t, resolved) + resolved, err = f.Store.ResolveLiveDocumentVectorPublications(t.Context(), newGeneration.ID, []string{newClaims[0].Token}) + require.NoError(t, err) + require.Len(t, resolved, 1) + + revisionBeforeRetire, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + retired, err := f.Store.RetireDocumentVectorGeneration(t.Context(), newGeneration.ID, now.Add(4*time.Second)) + require.NoError(t, err) + assert.True(t, retired) + revisionAfterRetire, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + assert.Equal(t, revisionBeforeRetire+1, revisionAfterRetire) + retired, err = f.Store.RetireDocumentVectorGeneration(t.Context(), newGeneration.ID, now.Add(5*time.Second)) + require.NoError(t, err) + assert.False(t, retired) + revisionAfterRepeat, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + assert.Equal(t, revisionAfterRetire, revisionAfterRepeat) + resolved, err = f.Store.ResolveLiveDocumentVectorPublications(t.Context(), newGeneration.ID, []string{newClaims[0].Token}) + require.NoError(t, err) + assert.Empty(t, resolved) + _, err = f.Store.GetDocumentVectorCoverage(t.Context(), newGeneration.ID) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) +} + +func testDocumentVectorLiveResolution(t *testing.T) { + t.Run("input ordering dedupe and bounds", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 2) + now := time.Date(2026, time.August, 20, 20, 21, 22, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 2) + require.NoError(f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + resolved, err := f.Store.ResolveLiveDocumentVectorPublications(t.Context(), generation.ID, []string{strings.Repeat("8", 64), claims[1].Token, claims[0].Token, claims[1].Token}) + require.NoError(err) + require.Len(resolved, 2) + assert.Equal(claims[1].Token, resolved[0].Token) + assert.Equal(claims[0].Token, resolved[1].Token) + assert.Equal(claims[1].DocumentVectorChunkCandidate, resolved[0].DocumentVectorChunkCandidate) + _, err = f.Store.ResolveLiveDocumentVectorPublications(t.Context(), generation.ID, []string{"invalid"}) + require.Error(err) + _, err = f.Store.ResolveLiveDocumentVectorPublications(t.Context(), generation.ID, make([]string, 1001)) + require.Error(err) + }) + + for _, mutation := range []string{"attachment replacement", "occurrence deletion", "role change", "message deletion", "target rotation"} { + t.Run(mutation, func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 21, 22, 23, 0, time.UTC) + claims := readyAllDocumentVectorChunks(t, f, generation, now) + require.Len(claims, 1) + require.NoError(f.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(time.Second))) + attachmentID := documentVectorAttachmentID(t, f, claims[0].CanonicalBlobHash) + file, err := f.Store.GetFileMetadata(t.Context(), attachmentID) + require.NoError(err) + require.NotNil(file) + + switch mutation { + case "attachment replacement": + newHash := strings.Repeat("c", 64) + require.NoError(f.Store.UpsertAttachmentRecord(t.Context(), file.MessageID, store.AttachmentWrite{ + Filename: file.Filename, MIMEType: file.MimeType, Size: file.Size, + StoragePath: newHash[:2] + "/" + newHash, ContentHash: newHash, + Role: file.AttachmentRole, RoleSource: file.RoleSource, SourcePartKey: file.SourcePartKey, + })) + case "occurrence deletion": + _, err = f.Store.DB().Exec(f.Store.Rebind(`DELETE FROM document_occurrences WHERE attachment_id = ?`), attachmentID) + require.NoError(err) + case "role change": + require.NoError(f.Store.UpsertAttachmentRecord(t.Context(), file.MessageID, store.AttachmentWrite{ + Filename: file.Filename, MIMEType: file.MimeType, Size: file.Size, + StoragePath: file.StoragePath, ContentHash: file.ContentHash, + Role: store.AttachmentRoleInline, RoleSource: store.AttachmentRoleSourceMIMEDisposition, + SourcePartKey: file.SourcePartKey, + })) + case "message deletion": + require.NoError(f.Store.MarkMessageDeleted(file.SourceID, file.SourceMessageID)) + case "target rotation": + profile := rotatedDocumentVectorProfile() + _, err = f.Store.EnsureDocumentExtractionProfile(t.Context(), profile) + require.NoError(err) + require.NoError(f.Store.RecordDocumentProviderConsent(t.Context(), store.DocumentProviderConsent{ + ProfileID: profile.ID, ProfileFingerprint: profile.Fingerprint, + RetentionPosture: profile.RetentionPosture, TrainingPosture: profile.TrainingPosture, + })) + } + + resolved, err := f.Store.ResolveLiveDocumentVectorPublications(t.Context(), generation.ID, []string{claims[0].Token}) + require.NoError(err) + assert.Empty(resolved) + var state string + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(`SELECT state FROM document_vector_publications WHERE generation_id = ? AND token = ?`), generation.ID, claims[0].Token).Scan(&state)) + assert.Equal("ready", state, "source invalidation must hide without deleting the ledger row") + }) + } +} + +func testDocumentVectorCleanupAndPurge(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 3) + now := time.Date(2026, time.August, 20, 22, 23, 24, 654321000, time.FixedZone("cleanup-offset", 4*60*60)) + ready, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "cleanup-worker", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, ready) + require.NoError(t, f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, ready.Token, ready.LeaseOwner, ready.LeaseFence, now.Add(time.Second))) + pending, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, ready.ChunkID, 1, "cleanup-worker", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, pending) + failed, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, pending.ChunkID, 1, "cleanup-worker", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, failed) + require.NoError(t, f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, failed.Token, failed.LeaseOwner, failed.LeaseFence, now.Add(time.Second), nil, true, "provider_rejected")) + + _, err = f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, "", 2) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + changed, err := f.Store.MarkDocumentVectorTokenCleaned(t.Context(), generation.ID, ready.Token, now) + assert.False(t, changed) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + _, err = f.Store.PurgeRetiredDocumentVectorGeneration(t.Context(), generation.ID) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + revisionBefore, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + retired, err := f.Store.RetireDocumentVectorGeneration(t.Context(), generation.ID, now.Add(2*time.Second)) + require.NoError(t, err) + assert.True(t, retired) + retiredGeneration, err := f.Store.GetDocumentVectorGeneration(t.Context(), generation.ID) + require.NoError(t, err) + require.NotNil(t, retiredGeneration.RetiredAt) + assert.Equal(t, now.Add(2*time.Second).UTC().Truncate(time.Millisecond), retiredGeneration.RetiredAt.UTC()) + revisionAfter, err := f.Store.GetDocumentIndexRevision(t.Context()) + require.NoError(t, err) + assert.Equal(t, revisionBefore, revisionAfter, "retiring a building generation need not invalidate search") + + wantTokens := []string{ready.Token, pending.Token, failed.Token} + sort.Strings(wantTokens) + first, err := f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, "", 2) + require.NoError(t, err) + require.Len(t, first, 2) + gotTokens := []string{first[0].Token, first[1].Token} + assert.Equal(t, wantTokens[:2], gotTokens) + second, err := f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, first[1].Token, 2) + require.NoError(t, err) + require.Len(t, second, 1) + assert.Equal(t, wantTokens[2], second[0].Token) + _, err = f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, "", 0) + require.ErrorContains(t, err, "limit") + _, err = f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, "", 1001) + require.ErrorContains(t, err, "limit") + _, err = f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, "not-a-token", 1) + require.ErrorContains(t, err, "cursor") + + _, err = f.Store.PurgeRetiredDocumentVectorGeneration(t.Context(), generation.ID) + require.ErrorIs(t, err, store.ErrDocumentVectorCleanupIncomplete) + var generationCount, publicationCount int + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(`SELECT COUNT(*) FROM document_vector_generations WHERE id = ?`), generation.ID).Scan(&generationCount)) + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(`SELECT COUNT(*) FROM document_vector_publications WHERE generation_id = ?`), generation.ID).Scan(&publicationCount)) + assert.Equal(t, 1, generationCount) + assert.Equal(t, 3, publicationCount) + + cleanTime := now.Add(3 * time.Second) + for _, token := range wantTokens { + changed, err := f.Store.MarkDocumentVectorTokenCleaned(t.Context(), generation.ID, token, cleanTime) + require.NoError(t, err) + assert.True(t, changed) + } + var cleanedAt time.Time + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(`SELECT backend_cleaned_at FROM document_vector_publications WHERE generation_id = ? AND token = ?`), generation.ID, wantTokens[0]).Scan(&cleanedAt)) + assert.Equal(t, cleanTime.UTC().Truncate(time.Millisecond), cleanedAt.UTC()) + changed, err = f.Store.MarkDocumentVectorTokenCleaned(t.Context(), generation.ID, wantTokens[0], cleanTime.Add(time.Second)) + require.NoError(t, err) + assert.False(t, changed) + changed, err = f.Store.MarkDocumentVectorTokenCleaned(t.Context(), generation.ID, strings.Repeat("9", 64), cleanTime) + require.NoError(t, err) + assert.False(t, changed) + remaining, err := f.Store.ListDocumentVectorCleanupTokens(t.Context(), generation.ID, "", 10) + require.NoError(t, err) + assert.Empty(t, remaining) + + purged, err := f.Store.PurgeRetiredDocumentVectorGeneration(t.Context(), generation.ID) + require.NoError(t, err) + assert.True(t, purged) + purged, err = f.Store.PurgeRetiredDocumentVectorGeneration(t.Context(), generation.ID) + require.NoError(t, err) + assert.False(t, purged) +} + +func runDocumentVectorChunkLifecycleContract(t *testing.T) { + t.Helper() + t.Run("claim resume bounds and stable token", testDocumentVectorClaimResumeBounds) + t.Run("renew and stale fence", testDocumentVectorRenewAndStaleFence) + t.Run("retry and terminal failure", testDocumentVectorFailureLifecycle) + t.Run("failure status and terminal reset", testDocumentVectorFailureStatusAndReset) + t.Run("source changed failure stays parked", testDocumentVectorSourceChangedFailureStaysParked) + t.Run("commit idempotence and source race", testDocumentVectorCommitLifecycle) + t.Run("invalid generation state", testDocumentVectorInvalidGenerationState) + t.Run("timestamp precision parity", testDocumentVectorTimestampPrecisionParity) + t.Run("ineffective sub-millisecond boundaries", testDocumentVectorRejectsIneffectiveTimeBoundaries) +} + +func testDocumentVectorFailureStatusAndReset(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 2) + now := time.Date(2026, time.August, 20, 12, 30, 0, 456789000, time.FixedZone("reset-offset", 2*60*60)) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + require.NoError(t, f.Store.FailDocumentVectorChunk( + t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, + now.Add(time.Second), nil, true, "provider_rejected", + )) + retryable, err := f.Store.ClaimDocumentVectorChunk( + t.Context(), generation.ID, claim.ChunkID, 1, "worker-a", now, time.Minute, + ) + require.NoError(t, err) + require.NotNil(t, retryable) + retryAt := now.Add(time.Hour).UTC().Truncate(time.Millisecond) + require.NoError(t, f.Store.FailDocumentVectorChunk( + t.Context(), generation.ID, retryable.Token, retryable.LeaseOwner, retryable.LeaseFence, + now.Add(time.Second), &retryAt, false, "provider_busy", + )) + firstPage, err := f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 1) + require.NoError(t, err) + require.Len(t, firstPage.Failures, 1) + assert.False(t, firstPage.FailuresExhausted) + assert.Equal(t, generation.ID, firstPage.FailureAfterGenerationID) + assert.NotEmpty(t, firstPage.FailureAfterToken) + secondPage, err := f.Store.GetDocumentVectorGenerationStatus( + t.Context(), generation.ID, firstPage.FailureAfterToken, 1, + ) + require.NoError(t, err) + require.Len(t, secondPage.Failures, 1) + assert.False(t, secondPage.FailuresExhausted) + lastPage, err := f.Store.GetDocumentVectorGenerationStatus( + t.Context(), generation.ID, secondPage.FailureAfterToken, 1, + ) + require.NoError(t, err) + assert.Empty(t, lastPage.Failures) + assert.True(t, lastPage.FailuresExhausted) + assert.Zero(t, lastPage.FailureAfterGenerationID) + assert.Empty(t, lastPage.FailureAfterToken) + + status, err := f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 10) + require.NoError(t, err) + assert.True(t, status.Blocked) + assert.Equal(t, int64(1), status.Terminal) + assert.Zero(t, status.Pending) + assert.Equal(t, int64(1), status.Retryable) + assert.Zero(t, status.ReadyLive) + assert.Zero(t, status.Obsolete) + assert.Zero(t, status.CleanupPending) + require.Len(t, status.Failures, 2) + diagnostics := make(map[string]store.DocumentVectorFailureDiagnostic, len(status.Failures)) + for _, diagnostic := range status.Failures { + diagnostics[diagnostic.Token] = diagnostic + } + assert.Equal(t, store.DocumentVectorFailureDiagnostic{ + Token: claim.Token, AttemptCount: 1, Terminal: true, ErrorCode: "provider_rejected", + }, diagnostics[claim.Token]) + assert.Equal(t, store.DocumentVectorFailureDiagnostic{ + Token: retryable.Token, AttemptCount: 1, NextRetryAt: &retryAt, ErrorCode: "provider_busy", + }, diagnostics[retryable.Token]) + assert.True(t, status.FailuresExhausted) + assert.Empty(t, status.FailureAfterToken) + _, err = f.Store.DB().Exec(f.Store.Rebind(` + UPDATE document_vector_publications SET error_code = ? + WHERE generation_id = ? AND token = ?`), "Provider returned private response text", generation.ID, retryable.Token) + require.NoError(t, err) + sanitized, err := f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 10) + require.NoError(t, err) + for _, diagnostic := range sanitized.Failures { + if diagnostic.Token == retryable.Token { + assert.Equal(t, "unknown", diagnostic.ErrorCode) + } + } + + reset, err := f.Store.ResetDocumentVectorFailures(t.Context(), generation.ID, "", 1, now.Add(2*time.Second)) + require.NoError(t, err) + assert.Equal(t, 1, reset.Scanned) + assert.Equal(t, 1, reset.Reset) + assert.False(t, reset.Exhausted) + assert.Equal(t, generation.ID, reset.AfterGenerationID) + assert.NotEmpty(t, reset.AfterToken) + secondReset, err := f.Store.ResetDocumentVectorFailures( + t.Context(), generation.ID, reset.AfterToken, 1, now.Add(2*time.Second), + ) + require.NoError(t, err) + assert.Equal(t, 1, secondReset.Scanned) + assert.Equal(t, 1, secondReset.Reset) + assert.False(t, secondReset.Exhausted) + finalReset, err := f.Store.ResetDocumentVectorFailures( + t.Context(), generation.ID, secondReset.AfterToken, 1, now.Add(2*time.Second), + ) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorFailureResetResult{Exhausted: true}, finalReset) + resetStatus, err := f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 10) + require.NoError(t, err) + assert.False(t, resetStatus.Blocked) + + reclaimed, err := f.Store.ClaimDocumentVectorChunk( + t.Context(), generation.ID, 0, 1, "worker-b", now.Add(3*time.Second), time.Minute, + ) + require.NoError(t, err) + require.NotNil(t, reclaimed) + assert.Equal(t, claim.Token, reclaimed.Token) + assert.Equal(t, claim.LeaseFence+1, reclaimed.LeaseFence) + assert.Equal(t, 1, reclaimed.AttemptCount) + _, err = f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 0) + require.ErrorContains(t, err, "limit") + _, err = f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "not-a-token", 1) + require.ErrorContains(t, err, "cursor") + _, err = f.Store.ResetDocumentVectorFailures(t.Context(), generation.ID, "", 1001, now) + require.ErrorContains(t, err, "limit") +} + +func testDocumentVectorSourceChangedFailureStaysParked(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 12, 45, 0, 0, time.UTC) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + _, err = f.Store.DB().Exec(f.Store.Rebind( + `DELETE FROM document_occurrences WHERE canonical_blob_hash = ?`), claim.CanonicalBlobHash) + require.NoError(t, err) + err = f.Store.CommitDocumentVectorPublication( + t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(time.Second), + ) + require.ErrorIs(t, err, store.ErrDocumentVectorSourceChanged) + + reset, err := f.Store.ResetDocumentVectorFailures(t.Context(), generation.ID, "", 10, now.Add(2*time.Second)) + require.NoError(t, err) + assert.Equal(t, store.DocumentVectorFailureResetResult{Scanned: 1, Exhausted: true}, reset) + status, err := f.Store.GetDocumentVectorGenerationStatus(t.Context(), generation.ID, "", 10) + require.NoError(t, err) + assert.Zero(t, status.Terminal, "obsolete failures are counted in Obsolete instead") + assert.Equal(t, int64(1), status.Obsolete) + assert.Equal(t, int64(1), status.CleanupPending) + require.Len(t, status.Failures, 1) + assert.Equal(t, "source_changed", status.Failures[0].ErrorCode) +} + +func testDocumentVectorClaimResumeBounds(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 3) + now := time.Date(2026, time.August, 20, 10, 11, 12, 0, time.UTC) + + first, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, 2*time.Minute) + require.NoError(t, err) + require.NotNil(t, first) + assert.Equal(t, int64(1), first.LeaseFence) + assert.Equal(t, 1, first.AttemptCount) + assert.Len(t, first.Token, 64) + assert.Equal(t, now.Add(2*time.Minute), first.LeaseUntil) + + observed, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now.Add(time.Second), 2*time.Minute) + require.NoError(t, err) + require.NotNil(t, observed) + assert.Equal(t, first.Token, observed.Token) + assert.Equal(t, first.LeaseFence, observed.LeaseFence) + assert.Equal(t, first.AttemptCount, observed.AttemptCount) + assert.Equal(t, first.LeaseUntil, observed.LeaseUntil) + + blocked, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", now.Add(time.Second), 2*time.Minute) + require.NoError(t, err) + assert.Nil(t, blocked, "bounded scan must not jump past its first busy candidate") + + second, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 2, "worker-b", now.Add(time.Second), 2*time.Minute) + require.NoError(t, err) + require.NotNil(t, second) + assert.Greater(t, second.ChunkID, first.ChunkID) + + takeover, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", first.LeaseUntil, 2*time.Minute) + require.NoError(t, err) + require.NotNil(t, takeover) + assert.Equal(t, first.Token, takeover.Token) + assert.Equal(t, first.LeaseFence+1, takeover.LeaseFence) + assert.Equal(t, first.AttemptCount+1, takeover.AttemptCount) +} + +func testDocumentVectorRenewAndStaleFence(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 11, 12, 13, 0, time.UTC) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + + renewedUntil, err := f.Store.RenewDocumentVectorChunkClaim(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(30*time.Second), 3*time.Minute) + require.NoError(t, err) + assert.Equal(t, now.Add(3*time.Minute+30*time.Second), renewedUntil) + + takeover, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", renewedUntil, time.Minute) + require.NoError(t, err) + require.NotNil(t, takeover) + _, err = f.Store.RenewDocumentVectorChunkClaim(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, renewedUntil, time.Minute) + require.ErrorIs(t, err, store.ErrDocumentVectorClaimLost) + assert.NotEqual(t, claim.LeaseFence, takeover.LeaseFence) +} + +func testDocumentVectorFailureLifecycle(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 2) + now := time.Date(2026, time.August, 20, 12, 13, 14, 0, time.UTC) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + retryAt := now.Add(5 * time.Minute) + require.NoError(t, f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(time.Second), &retryAt, false, "provider_busy")) + + beforeRetry, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", retryAt.Add(-time.Second), time.Minute) + require.NoError(t, err) + assert.Nil(t, beforeRetry) + retry, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", retryAt, time.Minute) + require.NoError(t, err) + require.NotNil(t, retry) + assert.Equal(t, claim.Token, retry.Token) + assert.Equal(t, claim.LeaseFence+1, retry.LeaseFence) + assert.Equal(t, claim.AttemptCount+1, retry.AttemptCount) + + require.Error(t, f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, retry.Token, retry.LeaseOwner, retry.LeaseFence, retryAt, nil, false, "provider_busy")) + require.Error(t, f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, retry.Token, retry.LeaseOwner, retry.LeaseFence, retryAt, &retryAt, true, "provider_busy")) + require.Error(t, f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, retry.Token, retry.LeaseOwner, retry.LeaseFence, retryAt, nil, true, strings.Repeat("x", 65))) + require.NoError(t, f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, retry.Token, retry.LeaseOwner, retry.LeaseFence, retryAt, nil, true, "provider_rejected")) + + terminal, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-c", retryAt.Add(time.Hour), time.Minute) + require.NoError(t, err) + assert.Nil(t, terminal) +} + +func testDocumentVectorCommitLifecycle(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 2) + now := time.Date(2026, time.August, 20, 13, 14, 15, 0, time.UTC) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + require.NoError(t, f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(time.Second))) + require.NoError(t, f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(2*time.Second))) + err = f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, claim.Token, "stale-worker", claim.LeaseFence, now.Add(2*time.Second)) + require.ErrorIs(t, err, store.ErrDocumentVectorClaimLost) + err = f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence+1, now.Add(2*time.Second)) + require.ErrorIs(t, err, store.ErrDocumentVectorClaimLost) + + tracing, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, claim.ChunkID, 1, "worker-b", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, tracing) + _, err = f.Store.DB().Exec(f.Store.Rebind(`DELETE FROM document_occurrences WHERE canonical_blob_hash = ?`), tracing.CanonicalBlobHash) + require.NoError(t, err) + err = f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, tracing.Token, tracing.LeaseOwner, tracing.LeaseFence, now.Add(time.Second)) + require.ErrorIs(t, err, store.ErrDocumentVectorSourceChanged) + + var state, errorCode string + var owner, leaseUntil, cleanedAt sql.NullString + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(`SELECT state, error_code, lease_owner, lease_until, backend_cleaned_at FROM document_vector_publications WHERE generation_id = ? AND token = ?`), generation.ID, tracing.Token).Scan(&state, &errorCode, &owner, &leaseUntil, &cleanedAt)) + assert.Equal(t, "failed", state) + assert.Equal(t, "source_changed", errorCode) + assert.False(t, owner.Valid) + assert.False(t, leaseUntil.Valid) + assert.False(t, cleanedAt.Valid) +} + +func testDocumentVectorInvalidGenerationState(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 14, 15, 16, 0, time.UTC) + _, err := f.Store.DB().Exec(f.Store.Rebind(`UPDATE document_vector_generations SET state = 'active' WHERE id = ?`), generation.ID) + require.NoError(t, err) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + assert.Nil(t, claim) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + + _, err = f.Store.RenewDocumentVectorChunkClaim(t.Context(), generation.ID, strings.Repeat("0", 64), "worker-a", 1, now, time.Minute) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + + f, generation = seedDocumentVectorGenerationWithChunks(t, 1) + _, err = f.Store.DB().Exec(`UPDATE document_index_state SET target_profile_id = NULL WHERE singleton = 1`) + require.NoError(t, err) + claim, err = f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + assert.Nil(t, claim) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) +} + +func testDocumentVectorTimestampPrecisionParity(t *testing.T) { + t.Run("claim", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + rawNow := time.Date(2026, time.August, 20, 15, 16, 17, 123456789, time.FixedZone("test-offset", 2*60*60)) + leaseDuration := 2*time.Second + 400*time.Microsecond + wantUntil := rawNow.Add(leaseDuration).UTC().Truncate(time.Millisecond) + + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", rawNow, leaseDuration) + require.NoError(err) + require.NotNil(claim) + assert.Equal(wantUntil, claim.LeaseUntil) + assert.Equal(wantUntil, documentVectorLeaseUntil(t, f, generation.ID, claim.Token)) + + before, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", wantUntil.Add(-time.Nanosecond), time.Second) + require.NoError(err) + assert.Nil(before) + atDeadline, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", wantUntil, time.Second) + require.NoError(err) + require.NotNil(atDeadline) + assert.Equal(claim.Token, atDeadline.Token) + assert.Equal(claim.LeaseFence+1, atDeadline.LeaseFence) + }) + + t.Run("renew", func(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + claimNow := time.Date(2026, time.August, 20, 16, 17, 18, 0, time.UTC) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", claimNow, time.Minute) + require.NoError(err) + require.NotNil(claim) + + rawNow := claimNow.Add(5*time.Second + 456789*time.Nanosecond) + leaseDuration := 2*time.Second + 400*time.Microsecond + wantUntil := rawNow.Add(leaseDuration).UTC().Truncate(time.Millisecond) + renewedUntil, err := f.Store.RenewDocumentVectorChunkClaim(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, rawNow, leaseDuration) + require.NoError(err) + assert.Equal(wantUntil, renewedUntil) + assert.Equal(wantUntil, documentVectorLeaseUntil(t, f, generation.ID, claim.Token)) + + before, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", wantUntil.Add(-time.Nanosecond), time.Second) + require.NoError(err) + assert.Nil(before) + atDeadline, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-b", wantUntil, time.Second) + require.NoError(err) + require.NotNil(atDeadline) + assert.Equal(claim.Token, atDeadline.Token) + }) +} + +func testDocumentVectorRejectsIneffectiveTimeBoundaries(t *testing.T) { + f, generation := seedDocumentVectorGenerationWithChunks(t, 1) + now := time.Date(2026, time.August, 20, 17, 18, 19, 123000000, time.UTC) + + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, 500*time.Microsecond) + assert.Nil(t, claim) + require.ErrorContains(t, err, "lease duration") + + claim, err = f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "worker-a", now, time.Minute) + require.NoError(t, err) + require.NotNil(t, claim) + _, err = f.Store.RenewDocumentVectorChunkClaim(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now, 500*time.Microsecond) + require.ErrorContains(t, err, "lease duration") + + retryAt := now.Add(500 * time.Microsecond) + err = f.Store.FailDocumentVectorChunk(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now, &retryAt, false, "provider_busy") + require.ErrorContains(t, err, "future retry time") +} + +func documentVectorLeaseUntil(t *testing.T, f *storetest.Fixture, generationID int64, token string) time.Time { + t.Helper() + var leaseUntil time.Time + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(`SELECT lease_until FROM document_vector_publications WHERE generation_id = ? AND token = ?`), generationID, token).Scan(&leaseUntil)) + return leaseUntil.UTC() +} + +func readyAllDocumentVectorChunks( + t *testing.T, + f *storetest.Fixture, + generation store.DocumentVectorGeneration, + now time.Time, +) []store.DocumentVectorChunkClaim { + t.Helper() + var claims []store.DocumentVectorChunkClaim + for len(claims) < 1000 { + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1000, "ready-worker", now, time.Minute) + require.NoError(t, err) + if claim == nil { + return claims + } + require.NoError(t, f.Store.CommitDocumentVectorPublication(t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(time.Second))) + claims = append(claims, *claim) + } + require.Fail(t, "document vector ready helper exceeded its safety bound") + return nil +} + +func documentVectorAttachmentID(t *testing.T, f *storetest.Fixture, hash string) int64 { + t.Helper() + var attachmentID int64 + require.NoError(t, f.Store.DB().QueryRow(f.Store.Rebind(`SELECT id FROM attachments WHERE content_hash = ?`), hash).Scan(&attachmentID)) + return attachmentID +} + +func rotatedDocumentVectorProfile() store.DocumentExtractionProfile { + fingerprint := strings.Repeat("2", 64) + return store.DocumentExtractionProfile{ + ID: "profile-" + fingerprint, Fingerprint: fingerprint, + Provider: "mistral", Endpoint: "https://api.mistral.ai/v1/ocr", + Region: "eu", Model: "mistral-ocr-5-0", + RetentionPosture: "standard", TrainingPosture: "opted-out", + AllowedMediaTypes: []string{"application/pdf"}, PolicyJSON: []byte(`{"policy":2}`), + } +} + +func seedDocumentVectorGenerationWithChunks(t *testing.T, chunkCount int) (*storetest.Fixture, store.DocumentVectorGeneration) { + t.Helper() + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishManySearchChunks(t, f, profile, hash, chunkCount) + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("f", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.NoError(t, err) + return f, generation +} + +func TestDocumentVectorGenerationCreateResumeCollisionAndBounds(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + profile, _ := seedDocumentPublicationAuthority(t, f) + spec := store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("1", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + } + + first, created, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(err) + assert.True(created) + assert.Equal(store.DocumentVectorGenerationBuilding, first.State) + assert.NotZero(first.ID) + + second, created, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(err) + assert.False(created) + assert.Equal(first.ID, second.ID) + + collision := spec + collision.Model = "embed-v2" + _, _, err = f.Store.EnsureDocumentVectorGeneration(t.Context(), collision) + require.ErrorContains(err, "fingerprint") + + other := spec + other.Fingerprint = strings.Repeat("2", 64) + _, _, err = f.Store.EnsureDocumentVectorGeneration(t.Context(), other) + require.ErrorContains(err, "building") + + invalid := spec + invalid.Fingerprint = strings.Repeat("3", 64) + invalid.Dimension = 0 + _, _, err = f.Store.EnsureDocumentVectorGeneration(t.Context(), invalid) + require.ErrorContains(err, "dimension") +} + +func TestDocumentVectorGenerationRejectsNonCurrentTarget(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + current, _ := seedDocumentPublicationAuthority(t, f) + nonCurrent := rotatedDocumentVectorProfile() + _, err := f.Store.EnsureDocumentExtractionProfile(t.Context(), nonCurrent) + require.NoError(err) + require.NoError(f.Store.RecordDocumentProviderConsent(t.Context(), store.DocumentProviderConsent{ + ProfileID: nonCurrent.ID, ProfileFingerprint: nonCurrent.Fingerprint, + RetentionPosture: nonCurrent.RetentionPosture, TrainingPosture: nonCurrent.TrainingPosture, + })) + _, err = f.Store.DB().Exec(f.Store.Rebind(` + UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), current.ID) + require.NoError(err) + + _, _, err = f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("3", 64), TargetExtractionProfileID: nonCurrent.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.ErrorIs(err, store.ErrDocumentVectorInvalidGenerationState) + var rows int + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT COUNT(*) FROM document_vector_generations WHERE target_extraction_profile_id = ?`), + nonCurrent.ID).Scan(&rows)) + assert.Zero(rows) + + generation, created, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("4", 64), TargetExtractionProfileID: current.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.NoError(err) + assert.True(created) + assert.Equal(current.ID, generation.TargetExtractionProfileID) +} + +func TestDocumentVectorBaseSchemaPreservesCleanupAuthority(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + if f.Store.IsPostgreSQL() { + rows, err := f.Store.DB().Query(` + SELECT confrelid::regclass::text + FROM pg_constraint + WHERE conrelid = 'document_vector_publications'::regclass + AND contype = 'f' + ORDER BY confrelid::regclass::text`) + require.NoError(err) + defer func() { _ = rows.Close() }() + var references []string + for rows.Next() { + var table string + require.NoError(rows.Scan(&table)) + references = append(references, table) + } + require.NoError(rows.Err()) + assert.Equal([]string{"document_vector_generations"}, references) + var definition string + require.NoError(f.Store.DB().QueryRow(` + SELECT pg_get_indexdef('idx_document_vector_publications_cleanup'::regclass)`).Scan(&definition)) + assert.Contains(definition, "(generation_id, backend_cleaned_at, token)") + return + } + + rows, err := f.Store.DB().Query(`PRAGMA foreign_key_list('document_vector_publications')`) + require.NoError(err) + defer func() { _ = rows.Close() }() + references := map[string]struct{}{} + for rows.Next() { + var id, sequence int + var table, from, to, onUpdate, onDelete, match string + require.NoError(rows.Scan(&id, &sequence, &table, &from, &to, &onUpdate, &onDelete, &match)) + references[table] = struct{}{} + } + require.NoError(rows.Err()) + assert.Equal(map[string]struct{}{"document_vector_generations": {}}, references) + indexRows, err := f.Store.DB().Query(`PRAGMA index_info('idx_document_vector_publications_cleanup')`) + require.NoError(err) + defer func() { _ = indexRows.Close() }() + var columns []string + for indexRows.Next() { + var sequence, columnID int + var column string + require.NoError(indexRows.Scan(&sequence, &columnID, &column)) + columns = append(columns, column) + } + require.NoError(indexRows.Err()) + assert.Equal([]string{"generation_id", "backend_cleaned_at", "token"}, columns) +} + +func TestDocumentVectorGenerationRetirementAndActiveConstraints(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + profile, _ := seedDocumentPublicationAuthority(t, f) + spec := store.DocumentVectorGenerationSpec{Fingerprint: strings.Repeat("4", 64), TargetExtractionProfileID: profile.ID, EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768} + first, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(err) + _, err = f.Store.DB().Exec(f.Store.Rebind(`UPDATE document_vector_generations SET state = 'retired', retired_at = CURRENT_TIMESTAMP WHERE id = ?`), first.ID) + require.NoError(err) + fresh, created, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(err) + assert.True(created) + assert.NotEqual(first.ID, fresh.ID) + + _, err = f.Store.DB().Exec(f.Store.Rebind(`UPDATE document_vector_generations SET state = 'active' WHERE id = ?`), fresh.ID) + require.NoError(err) + _, err = f.Store.DB().Exec(f.Store.Rebind(`INSERT INTO document_vector_generations (fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state) VALUES (?, ?, ?, ?, ?, 'active')`), strings.Repeat("5", 64), profile.ID, "vector.embeddings", "embed-v1", 768) + require.Error(err) +} + +func TestDocumentVectorGenerationRejectsHistoricalSpecCollision(t *testing.T) { + f := storetest.New(t) + profile, _ := seedDocumentPublicationAuthority(t, f) + spec := store.DocumentVectorGenerationSpec{Fingerprint: strings.Repeat("c", 64), TargetExtractionProfileID: profile.ID, EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768} + _, err := f.Store.DB().Exec(f.Store.Rebind(`INSERT INTO document_vector_generations (fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state, retired_at) VALUES (?, ?, ?, ?, ?, 'retired', CURRENT_TIMESTAMP)`), spec.Fingerprint, profile.ID, spec.EmbeddingProfile, "embed-old", spec.Dimension) + require.NoError(t, err) + _, err = f.Store.DB().Exec(f.Store.Rebind(`INSERT INTO document_vector_generations (fingerprint, target_extraction_profile_id, embedding_profile, model, dimension, state) VALUES (?, ?, ?, ?, ?, 'building')`), spec.Fingerprint, profile.ID, spec.EmbeddingProfile, spec.Model, spec.Dimension) + require.NoError(t, err) + _, _, err = f.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.ErrorContains(t, err, "collides") +} + +func TestDocumentVectorGenerationRejectsInvalidFingerprintAndProfile(t *testing.T) { + f := storetest.New(t) + profile, _ := seedDocumentPublicationAuthority(t, f) + for _, spec := range []store.DocumentVectorGenerationSpec{ + {Fingerprint: "not-a-digest", TargetExtractionProfileID: profile.ID, EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768}, + {Fingerprint: strings.Repeat("6", 64), TargetExtractionProfileID: profile.ID, EmbeddingProfile: "other.embeddings", Model: "embed-v1", Dimension: 768}, + } { + _, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.Error(t, err) + } +} + +func TestDocumentVectorTargetRejectsLegacyNormalizedIdentity(t *testing.T) { + requirements := require.New(t) + f, _ := seedDocumentVectorGenerationWithChunks(t, 1) + var extractionID string + requirements.NoError(f.Store.DB().QueryRow(`SELECT extraction_id FROM document_extraction_heads LIMIT 1`).Scan(&extractionID)) + _, err := f.Store.DB().Exec(f.Store.Rebind(` + UPDATE document_extractions + SET normalization_version = NULL, document_family = NULL, unit_kind = NULL + WHERE id = ?`), extractionID) + requirements.NoError(err) + + _, err = f.Store.GetDocumentVectorTargetProfileID(t.Context()) + requirements.ErrorIs(err, store.ErrDocumentNormalizedIdentityUnavailable) + requirements.ErrorContains(err, "documents build --full-rebuild") + _, err = f.Store.LoadNormalizedDocument(t.Context(), extractionID) + requirements.ErrorIs(err, store.ErrDocumentNormalizedIdentityUnavailable) + requirements.ErrorContains(err, "documents build --full-rebuild") +} + +func TestDocumentVectorTargetIgnoresDeadLegacyNormalizedIdentity(t *testing.T) { + requirements := require.New(t) + f, _ := seedDocumentVectorGenerationWithChunks(t, 1) + var extractionID, profileID, hash string + requirements.NoError(f.Store.DB().QueryRow(` + SELECT extraction_id, profile_id, canonical_blob_hash + FROM document_extraction_heads LIMIT 1`).Scan(&extractionID, &profileID, &hash)) + _, err := f.Store.DB().Exec(f.Store.Rebind(` + UPDATE document_extractions + SET normalization_version = NULL, document_family = NULL, unit_kind = NULL + WHERE id = ?`), extractionID) + requirements.NoError(err) + _, err = f.Store.DB().Exec(f.Store.Rebind(` + DELETE FROM document_occurrences WHERE canonical_blob_hash = ?`), hash) + requirements.NoError(err) + + target, err := f.Store.GetDocumentVectorTargetProfileID(t.Context()) + requirements.NoError(err) + requirements.Equal(profileID, target) +} + +func TestDocumentVectorChunkCandidatesUseCurrentLiveAuthority(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, "vector candidate evidence", "vector-candidate") + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("7", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.NoError(err) + + candidates, err := f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 10) + require.NoError(err) + require.Len(candidates, 1) + assert.Equal(generation.ID, candidates[0].GenerationID) + assert.Equal(profile.ID, candidates[0].ExtractionProfileID) + assert.Equal(hash, candidates[0].CanonicalBlobHash) + assert.Equal("vector candidate evidence", candidates[0].Text) + assert.Equal("vector-candidate", candidates[0].ExtractionID) + assert.NotEmpty(candidates[0].ChunkKey) + assert.NotEmpty(candidates[0].ChunkChecksum) + + _, err = f.Store.DB().Exec(f.Store.Rebind(`DELETE FROM document_occurrences WHERE canonical_blob_hash = ?`), hash) + require.NoError(err) + candidates, err = f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 10) + require.NoError(err) + assert.Empty(candidates) +} + +func TestDocumentVectorChunkCandidatesHideAfterTargetProfileRotation(t *testing.T) { + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, "rotation candidate evidence", "vector-rotation") + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{Fingerprint: strings.Repeat("8", 64), TargetExtractionProfileID: profile.ID, EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768}) + require.NoError(err) + _, err = f.Store.DB().Exec(f.Store.Rebind(`UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), "rotated-profile") + require.NoError(err) + candidates, err := f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 10) + require.NoError(err) + assert.Empty(t, candidates) +} + +func TestDocumentVectorPublicationCommitRejectsInconsistentImmutableSnapshot(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, "integrity candidate evidence", "vector-integrity") + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{Fingerprint: strings.Repeat("9", 64), TargetExtractionProfileID: profile.ID, EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768}) + require.NoError(err) + now := time.Date(2026, time.August, 20, 23, 0, 0, 0, time.UTC) + claim, err := f.Store.ClaimDocumentVectorChunk(t.Context(), generation.ID, 0, 1, "integrity-worker", now, time.Minute) + require.NoError(err) + require.NotNil(claim) + _, err = f.Store.DB().Exec(f.Store.Rebind(` + UPDATE document_vector_publications SET canonical_blob_hash = ? + WHERE generation_id = ? AND token = ?`), strings.Repeat("f", 64), generation.ID, claim.Token) + require.NoError(err) + err = f.Store.CommitDocumentVectorPublication( + t.Context(), generation.ID, claim.Token, claim.LeaseOwner, claim.LeaseFence, now.Add(time.Second)) + require.ErrorIs(err, store.ErrDocumentVectorSourceChanged) + var state, errorCode string + require.NoError(f.Store.DB().QueryRow(f.Store.Rebind(` + SELECT state, error_code FROM document_vector_publications + WHERE generation_id = ? AND token = ?`), generation.ID, claim.Token).Scan(&state, &errorCode)) + assert.Equal("failed", state) + assert.Equal("source_changed", errorCode) +} + +func TestDocumentVectorGenerationCannotDeletePublicationBeforeBackendCleanup(t *testing.T) { + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, "retained token evidence", "vector-retained") + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{Fingerprint: strings.Repeat("d", 64), TargetExtractionProfileID: profile.ID, EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768}) + require.NoError(err) + candidates, err := f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 1) + require.NoError(err) + require.Len(candidates, 1) + candidate := candidates[0] + _, err = f.Store.DB().Exec(f.Store.Rebind(`INSERT INTO document_vector_publications (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, token, state) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending')`), generation.ID, candidate.ExtractionID, candidate.ExtractionProfileID, candidate.CanonicalBlobHash, candidate.ExtractionInputKey, candidate.ChunkID, candidate.ChunkKey, candidate.ChunkChecksum, candidate.SourceSequence, "opaque-retained-token") + require.NoError(err) + _, err = f.Store.DB().Exec(f.Store.Rebind(`DELETE FROM document_vector_generations WHERE id = ?`), generation.ID) + require.Error(err) +} + +func TestDocumentVectorChunkCandidatesRequireConfiguredExtractionProfile(t *testing.T) { + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishSearchDocument(t, f, profile, hash, "fallback vector evidence", "vector-fallback") + + target := profile + target.Fingerprint = strings.Repeat("e", 64) + target.ID = "profile-" + target.Fingerprint + target.Model = "mistral-ocr-5-0" + target.PolicyJSON = []byte(`{"policy":2}`) + _, err := f.Store.EnsureDocumentExtractionProfile(t.Context(), target) + require.NoError(err) + require.NoError(f.Store.RecordDocumentProviderConsent(t.Context(), store.DocumentProviderConsent{ + ProfileID: target.ID, ProfileFingerprint: target.Fingerprint, + RetentionPosture: target.RetentionPosture, TrainingPosture: target.TrainingPosture, + })) + _, err = f.Store.DB().Exec(f.Store.Rebind(`UPDATE document_index_state SET target_profile_id = ? WHERE singleton = 1`), target.ID) + require.NoError(err) + + _, err = f.Store.DB().Exec(f.Store.Rebind(`UPDATE document_vector_generations SET state = 'active' WHERE state = 'building'`)) + require.NoError(err) + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("a", 64), TargetExtractionProfileID: target.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.NoError(err) + + candidates, err := f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 10) + require.NoError(err) + assert.Empty(t, candidates, "an older fallback extraction does not satisfy the configured Docbank identity") + + publishSearchDocument(t, f, target, hash, "target vector evidence", "vector-target") + candidates, err = f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 10) + require.NoError(err) + require.Len(candidates, 1) + assert.Equal(t, target.ID, candidates[0].ExtractionProfileID) +} + +func TestDocumentVectorChunkCandidatesHaveStableBoundedPagination(t *testing.T) { + require := require.New(t) + f := storetest.New(t) + profile, hash := seedDocumentPublicationAuthority(t, f) + publishManySearchChunks(t, f, profile, hash, 3) + generation, _, err := f.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("b", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-v1", Dimension: 768, + }) + require.NoError(err) + + first, err := f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 2) + require.NoError(err) + require.Len(first, 2) + second, err := f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, first[1].ChunkID, 2) + require.NoError(err) + require.Len(second, 1) + assert.Greater(t, second[0].ChunkID, first[1].ChunkID) + + _, err = f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 0) + require.ErrorContains(err, "limit") + _, err = f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, 0, 1001) + require.ErrorContains(err, "limit") + _, err = f.Store.ListDocumentVectorChunkCandidates(t.Context(), generation.ID, -1, 1) + require.ErrorContains(err, "after") +} diff --git a/internal/store/schema.sql b/internal/store/schema.sql index 8eb1d7fbc..2a1d6b9c9 100644 --- a/internal/store/schema.sql +++ b/internal/store/schema.sql @@ -2195,6 +2195,10 @@ CREATE TABLE IF NOT EXISTS document_extractions ( units_processed INTEGER, returned_model TEXT, manifest_checksum TEXT, + normalization_version INTEGER, + document_family TEXT, + unit_kind TEXT, + normalized_truncated BOOLEAN NOT NULL DEFAULT FALSE, terminal_reason TEXT, source_sequence INTEGER NOT NULL DEFAULT 0, created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, @@ -2209,6 +2213,8 @@ CREATE INDEX IF NOT EXISTS idx_document_extractions_owner ON document_extractions(profile_id, canonical_blob_hash, extraction_input_key, state); CREATE INDEX IF NOT EXISTS idx_document_extractions_lease ON document_extractions(state, lease_until); +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_extractions_vector_identity + ON document_extractions(id, profile_id, canonical_blob_hash, extraction_input_key, source_sequence); -- One renewable claim per stable content owner. The monotonic fence prevents -- an expired worker from publishing after a later worker has taken ownership. @@ -2249,6 +2255,7 @@ CREATE TABLE IF NOT EXISTS document_units ( checksum TEXT NOT NULL, char_count INTEGER NOT NULL, truncated BOOLEAN NOT NULL DEFAULT FALSE, + heading_marks JSON NOT NULL DEFAULT '[]', PRIMARY KEY (extraction_id, unit_index), CHECK (unit_index >= 0), CHECK (char_count >= 0) @@ -2275,6 +2282,8 @@ CREATE TABLE IF NOT EXISTS document_chunks ( CHECK (first_unit_index >= 0 AND last_unit_index >= first_unit_index), CHECK (synthetic_prefix_len >= 0 AND char_count >= 0) ); +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_chunks_vector_identity + ON document_chunks(id, extraction_id, chunk_key, checksum); CREATE TABLE IF NOT EXISTS document_chunk_spans ( extraction_id TEXT NOT NULL, @@ -2326,6 +2335,85 @@ CREATE TABLE IF NOT EXISTS document_index_state ( ); INSERT OR IGNORE INTO document_index_state(singleton, revision) VALUES (1, 0); +-- Document vectors are a corpus separate from message embeddings. The main +-- archive database owns generation and publication authority; vector backends +-- only store the opaque token below. +CREATE TABLE IF NOT EXISTS document_vector_generations ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + fingerprint TEXT NOT NULL, + target_extraction_profile_id TEXT NOT NULL REFERENCES document_extraction_profiles(id) ON DELETE RESTRICT, + embedding_profile TEXT NOT NULL, + model TEXT NOT NULL, + dimension INTEGER NOT NULL CHECK (dimension > 0), + state TEXT NOT NULL CHECK (state IN ('building', 'active', 'retired')), + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + activated_at DATETIME, + retired_at DATETIME +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_vector_generations_building + ON document_vector_generations(state) WHERE state = 'building'; +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_vector_generations_active + ON document_vector_generations(state) WHERE state = 'active'; +CREATE INDEX IF NOT EXISTS idx_document_vector_generations_live_fingerprint + ON document_vector_generations(fingerprint) WHERE state <> 'retired'; + +-- Hosted embedding consent is bound to both the reusable generation policy +-- and a separate canonical egress destination fingerprint. It deliberately +-- contains no credentials, raw endpoints, or provider payloads. +CREATE TABLE IF NOT EXISTS document_vector_consents ( + egress_fingerprint TEXT PRIMARY KEY, + purpose TEXT NOT NULL CHECK (purpose IN ('document_embedding', 'query_embedding')), + generation_fingerprint TEXT NOT NULL, + target_extraction_profile_id TEXT NOT NULL REFERENCES document_extraction_profiles(id) ON DELETE RESTRICT, + embedding_profile TEXT NOT NULL, + model TEXT NOT NULL, + dimension INTEGER NOT NULL CHECK (dimension > 0), + consented_at DATETIME NOT NULL +); + +CREATE TABLE IF NOT EXISTS document_vector_provider_usage ( + fingerprint TEXT PRIMARY KEY, + provider_calls INTEGER NOT NULL DEFAULT 0 CHECK (provider_calls >= 0), + provider_documents INTEGER NOT NULL DEFAULT 0 CHECK (provider_documents >= 0), + provider_chunks INTEGER NOT NULL DEFAULT 0 CHECK (provider_chunks >= 0), + provider_input_chars INTEGER NOT NULL DEFAULT 0 CHECK (provider_input_chars >= 0), + updated_at DATETIME NOT NULL +); + +CREATE TABLE IF NOT EXISTS document_vector_build_progress ( + generation_id INTEGER PRIMARY KEY REFERENCES document_vector_generations(id) ON DELETE CASCADE, + after_chunk_id INTEGER NOT NULL CHECK (after_chunk_id > 0), + updated_at DATETIME NOT NULL +); + +-- Publication rows are deliberately complete before a vector backend exists: +-- the durable token is the only identifier a backend receives. +CREATE TABLE IF NOT EXISTS document_vector_publications ( + generation_id INTEGER NOT NULL REFERENCES document_vector_generations(id) ON DELETE RESTRICT, + extraction_id TEXT NOT NULL, + extraction_profile_id TEXT NOT NULL, + canonical_blob_hash TEXT NOT NULL CHECK (length(canonical_blob_hash) = 64), + extraction_input_key TEXT NOT NULL, + chunk_id INTEGER NOT NULL, + chunk_key TEXT NOT NULL, + chunk_checksum TEXT NOT NULL, + source_sequence INTEGER NOT NULL, + token TEXT NOT NULL UNIQUE, + state TEXT NOT NULL CHECK (state IN ('pending', 'ready', 'failed')), + lease_owner TEXT, + lease_fence INTEGER NOT NULL DEFAULT 0, + lease_until DATETIME, + attempt_count INTEGER NOT NULL DEFAULT 0 CHECK (attempt_count >= 0), + next_retry_at DATETIME, + error_code TEXT, + backend_cleaned_at DATETIME, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (generation_id, extraction_id, chunk_id) +); +CREATE INDEX IF NOT EXISTS idx_document_vector_publications_cleanup + ON document_vector_publications(generation_id, backend_cleaned_at, token); + -- Foreign-key cascades can remove occurrences before asynchronous attachment -- reconciliation observes the deletion. Invalidate search cursors at the -- authoritative row mutation so every deletion path is covered. diff --git a/internal/store/schema_pg.sql b/internal/store/schema_pg.sql index 84dab7f7e..8e081f1ea 100644 --- a/internal/store/schema_pg.sql +++ b/internal/store/schema_pg.sql @@ -2192,6 +2192,10 @@ CREATE TABLE IF NOT EXISTS document_extractions ( units_processed INTEGER, returned_model TEXT, manifest_checksum TEXT, + normalization_version INTEGER, + document_family TEXT, + unit_kind TEXT, + normalized_truncated BOOLEAN NOT NULL DEFAULT FALSE, terminal_reason TEXT, source_sequence BIGINT NOT NULL DEFAULT 0, created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, @@ -2206,6 +2210,8 @@ CREATE INDEX IF NOT EXISTS idx_document_extractions_owner ON document_extractions(profile_id, canonical_blob_hash, extraction_input_key, state); CREATE INDEX IF NOT EXISTS idx_document_extractions_lease ON document_extractions(state, lease_until); +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_extractions_vector_identity + ON document_extractions(id, profile_id, canonical_blob_hash, extraction_input_key, source_sequence); CREATE TABLE IF NOT EXISTS document_extraction_claims ( profile_id TEXT NOT NULL REFERENCES document_extraction_profiles(id) ON DELETE CASCADE, @@ -2242,6 +2248,7 @@ CREATE TABLE IF NOT EXISTS document_units ( checksum TEXT NOT NULL, char_count INTEGER NOT NULL, truncated BOOLEAN NOT NULL DEFAULT FALSE, + heading_marks JSONB NOT NULL DEFAULT '[]'::jsonb, PRIMARY KEY (extraction_id, unit_index), CHECK (unit_index >= 0), CHECK (char_count >= 0) @@ -2269,6 +2276,8 @@ CREATE TABLE IF NOT EXISTS document_chunks ( CHECK (first_unit_index >= 0 AND last_unit_index >= first_unit_index), CHECK (synthetic_prefix_len >= 0 AND char_count >= 0) ); +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_chunks_vector_identity + ON document_chunks(id, extraction_id, chunk_key, checksum); CREATE INDEX IF NOT EXISTS idx_document_chunks_search_fts ON document_chunks USING GIN(search_fts); @@ -2320,6 +2329,77 @@ CREATE TABLE IF NOT EXISTS document_index_state ( INSERT INTO document_index_state(singleton, revision) VALUES (1, 0) ON CONFLICT (singleton) DO NOTHING; +CREATE TABLE IF NOT EXISTS document_vector_generations ( + id BIGINT GENERATED ALWAYS AS IDENTITY PRIMARY KEY, + fingerprint TEXT NOT NULL, + target_extraction_profile_id TEXT NOT NULL REFERENCES document_extraction_profiles(id) ON DELETE RESTRICT, + embedding_profile TEXT NOT NULL, + model TEXT NOT NULL, + dimension INTEGER NOT NULL CHECK (dimension > 0), + state TEXT NOT NULL CHECK (state IN ('building', 'active', 'retired')), + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + activated_at TIMESTAMPTZ, + retired_at TIMESTAMPTZ +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_vector_generations_building + ON document_vector_generations(state) WHERE state = 'building'; +CREATE UNIQUE INDEX IF NOT EXISTS idx_document_vector_generations_active + ON document_vector_generations(state) WHERE state = 'active'; +CREATE INDEX IF NOT EXISTS idx_document_vector_generations_live_fingerprint + ON document_vector_generations(fingerprint) WHERE state <> 'retired'; + +CREATE TABLE IF NOT EXISTS document_vector_consents ( + egress_fingerprint TEXT PRIMARY KEY, + purpose TEXT NOT NULL CHECK (purpose IN ('document_embedding', 'query_embedding')), + generation_fingerprint TEXT NOT NULL, + target_extraction_profile_id TEXT NOT NULL REFERENCES document_extraction_profiles(id) ON DELETE RESTRICT, + embedding_profile TEXT NOT NULL, + model TEXT NOT NULL, + dimension INTEGER NOT NULL CHECK (dimension > 0), + consented_at TIMESTAMPTZ NOT NULL +); + +CREATE TABLE IF NOT EXISTS document_vector_provider_usage ( + fingerprint TEXT PRIMARY KEY, + provider_calls BIGINT NOT NULL DEFAULT 0 CHECK (provider_calls >= 0), + provider_documents BIGINT NOT NULL DEFAULT 0 CHECK (provider_documents >= 0), + provider_chunks BIGINT NOT NULL DEFAULT 0 CHECK (provider_chunks >= 0), + provider_input_chars BIGINT NOT NULL DEFAULT 0 CHECK (provider_input_chars >= 0), + updated_at TIMESTAMPTZ NOT NULL +); + +CREATE TABLE IF NOT EXISTS document_vector_build_progress ( + generation_id BIGINT PRIMARY KEY REFERENCES document_vector_generations(id) ON DELETE CASCADE, + after_chunk_id BIGINT NOT NULL CHECK (after_chunk_id > 0), + updated_at TIMESTAMPTZ NOT NULL +); + +CREATE TABLE IF NOT EXISTS document_vector_publications ( + generation_id BIGINT NOT NULL REFERENCES document_vector_generations(id) ON DELETE RESTRICT, + extraction_id TEXT NOT NULL, + extraction_profile_id TEXT NOT NULL, + canonical_blob_hash TEXT NOT NULL CHECK (length(canonical_blob_hash) = 64), + extraction_input_key TEXT NOT NULL, + chunk_id BIGINT NOT NULL, + chunk_key TEXT NOT NULL, + chunk_checksum TEXT NOT NULL, + source_sequence BIGINT NOT NULL, + token TEXT NOT NULL UNIQUE, + state TEXT NOT NULL CHECK (state IN ('pending', 'ready', 'failed')), + lease_owner TEXT, + lease_fence BIGINT NOT NULL DEFAULT 0, + lease_until TIMESTAMPTZ, + attempt_count INTEGER NOT NULL DEFAULT 0 CHECK (attempt_count >= 0), + next_retry_at TIMESTAMPTZ, + error_code TEXT, + backend_cleaned_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (generation_id, extraction_id, chunk_id) +); +CREATE INDEX IF NOT EXISTS idx_document_vector_publications_cleanup + ON document_vector_publications(generation_id, backend_cleaned_at, token); + -- Foreign-key cascades can remove occurrences before asynchronous attachment -- reconciliation observes the deletion. Invalidate search cursors at the -- authoritative row mutation so every deletion path is covered. diff --git a/internal/store/sqlite_error_test.go b/internal/store/sqlite_error_test.go index a5f2ef5ef..61f7e45e3 100644 --- a/internal/store/sqlite_error_test.go +++ b/internal/store/sqlite_error_test.go @@ -17,7 +17,7 @@ func TestIsSQLiteError_ValueForm(t *testing.T) { } // Wrap the error - wrappedErr := fmt.Errorf("insert failed: %w", sqliteErr) + wrappedErr := errors.Join(sqliteErr) // sqlite3.Error.Error() returns the code description, e.g. "constraint failed" assert.True(t, isSQLiteError(wrappedErr, "constraint failed"), @@ -33,7 +33,6 @@ func TestIsSQLiteError_PointerForm(t *testing.T) { Code: sqlite3.ErrConstraint, ExtendedCode: sqlite3.ErrConstraintForeignKey, } - // sqlite3.Error.Error() returns the code description, e.g. "constraint failed" assert.True(t, isSQLiteError(sqliteErr, "constraint failed"), "isSQLiteError should match constraint error via pointer, got: %v", sqliteErr.Error()) diff --git a/internal/vector/backend.go b/internal/vector/backend.go index 4242185e4..2c881ad83 100644 --- a/internal/vector/backend.go +++ b/internal/vector/backend.go @@ -8,6 +8,12 @@ import ( // GenerationID identifies one index generation. type GenerationID int64 +// DocumentInput preserves one document boundary around its ordered embedding +// inputs. Providers may partially complete only a leading document prefix. +type DocumentInput struct { + Chunks []string +} + // SourceBasis identifies the source text used for chunk offsets. type SourceBasis uint8 diff --git a/internal/vector/config.go b/internal/vector/config.go index 85ef39a52..65ab83114 100644 --- a/internal/vector/config.go +++ b/internal/vector/config.go @@ -203,6 +203,39 @@ func (e EmbeddingsConfig) EffectiveAPIFormat() EmbeddingAPIFormat { return e.APIFormat } +// Validate checks embedding-provider settings without requiring the ordinary +// message-vector lane to be enabled. Document vectors use the same provider +// policy under their own enablement and consent controls. +func (e EmbeddingsConfig) Validate() error { + switch e.EffectiveAPIFormat() { + case APIFormatOpenAI, APIFormatVoyageContextual: + default: + return fmt.Errorf("vector.embeddings.api_format: unknown format %q (supported: %q, %q)", + e.APIFormat, APIFormatOpenAI, APIFormatVoyageContextual) + } + if e.EffectiveAPIFormat() == APIFormatVoyageContextual && e.Model != "voyage-context-4" { + return fmt.Errorf("vector.embeddings.model: api_format=%q requires %q, got %q", + APIFormatVoyageContextual, "voyage-context-4", e.Model) + } + if e.Endpoint == "" { + return errors.New("vector.embeddings.endpoint: required") + } + u, err := url.Parse(e.Endpoint) + if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { + return fmt.Errorf("vector.embeddings.endpoint: must be an http or https URL with a host (got %q)", e.Endpoint) + } + if e.Model == "" { + return fmt.Errorf("vector.embeddings.model: required (the index generation fingerprint is %q, which is ambiguous without a model name)", e.Fingerprint()) + } + if e.Dimension <= 0 { + return fmt.Errorf("vector.embeddings.dimension: must be positive, got %d", e.Dimension) + } + if e.BatchSize <= 0 { + return fmt.Errorf("vector.embeddings.batch_size: must be positive, got %d", e.BatchSize) + } + return nil +} + // PreprocessConfig controls message text preprocessing before embedding. // // Fields are pointers so the decoder can distinguish "unset" (nil, @@ -490,32 +523,8 @@ func (c *Config) Validate() error { if !c.Enabled { return nil } - switch c.Embeddings.EffectiveAPIFormat() { - case APIFormatOpenAI, APIFormatVoyageContextual: - default: - return fmt.Errorf("vector.embeddings.api_format: unknown format %q (supported: %q, %q)", - c.Embeddings.APIFormat, APIFormatOpenAI, APIFormatVoyageContextual) - } - if c.Embeddings.EffectiveAPIFormat() == APIFormatVoyageContextual && - c.Embeddings.Model != "voyage-context-4" { - return fmt.Errorf("vector.embeddings.model: api_format=%q requires %q, got %q", - APIFormatVoyageContextual, "voyage-context-4", c.Embeddings.Model) - } - if c.Embeddings.Endpoint == "" { - return errors.New("vector.embeddings.endpoint: required") - } - u, err := url.Parse(c.Embeddings.Endpoint) - if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { - return fmt.Errorf("vector.embeddings.endpoint: must be an http or https URL with a host (got %q)", c.Embeddings.Endpoint) - } - if c.Embeddings.Model == "" { - return fmt.Errorf("vector.embeddings.model: required (the index generation fingerprint is %q, which is ambiguous without a model name)", c.Embeddings.Fingerprint()) - } - if c.Embeddings.Dimension <= 0 { - return fmt.Errorf("vector.embeddings.dimension: must be positive, got %d", c.Embeddings.Dimension) - } - if c.Embeddings.BatchSize <= 0 { - return fmt.Errorf("vector.embeddings.batch_size: must be positive, got %d", c.Embeddings.BatchSize) + if err := c.Embeddings.Validate(); err != nil { + return err } if c.People.Enabled { if _, err := c.SemanticPersonEmbeddingProfile(); err != nil { diff --git a/internal/vector/document/backend.go b/internal/vector/document/backend.go new file mode 100644 index 000000000..e001f2a11 --- /dev/null +++ b/internal/vector/document/backend.go @@ -0,0 +1,97 @@ +package document + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "math" + "unicode/utf8" +) + +// ErrInvalidVector reports an invalid document-vector request or a vector +// that cannot participate in cosine similarity. +var ErrInvalidVector = errors.New("invalid document vector") + +// Embedding pairs an opaque publication token with its vector. +type Embedding struct { + Token string + Vector []float32 +} + +// Hit is one cosine-similarity result. Higher scores are better and Rank is +// one-based after deterministic score/token ordering. +type Hit struct { + Token string + Score float64 + Rank int +} + +// HitPage is one deterministic global vector page. Exhausted is authoritative. +type HitPage struct { + Hits []Hit + NextCursor string + Exhausted bool +} + +const pageCursorVersion = 1 + +type pageCursor struct { + Version int `json:"v"` + Distance float64 `json:"d"` + Token string `json:"t"` + Rank int `json:"r"` +} + +// EncodePageCursor records the last globally ordered vector hit. Backends use +// it to continue after that hit even when earlier obsolete rows are deleted. +func EncodePageCursor(distance float64, token string, rank int) (string, error) { + if math.IsNaN(distance) || math.IsInf(distance, 0) || token == "" || !utf8.ValidString(token) || rank < 1 { + return "", fmt.Errorf("%w: invalid document vector cursor", ErrInvalidVector) + } + payload, err := json.Marshal(pageCursor{ + Version: pageCursorVersion, Distance: distance, Token: token, Rank: rank, + }) + if err != nil { + return "", fmt.Errorf("encode document vector cursor: %w", err) + } + return base64.RawURLEncoding.EncodeToString(payload), nil +} + +// DecodePageCursor returns the last distance, token, and stable rank carried +// by an opaque backend cursor. An empty cursor starts at the first hit. +func DecodePageCursor(cursor string) (distance float64, token string, rank int, err error) { + if cursor == "" { + return 0, "", 0, nil + } + payload, err := base64.RawURLEncoding.DecodeString(cursor) + if err != nil { + return 0, "", 0, fmt.Errorf("%w: invalid document vector cursor", ErrInvalidVector) + } + var decoded pageCursor + if err := json.Unmarshal(payload, &decoded); err != nil || + decoded.Version != pageCursorVersion || math.IsNaN(decoded.Distance) || + math.IsInf(decoded.Distance, 0) || decoded.Token == "" || + !utf8.ValidString(decoded.Token) || decoded.Rank < 1 { + return 0, "", 0, fmt.Errorf("%w: invalid document vector cursor", ErrInvalidVector) + } + return decoded.Distance, decoded.Token, decoded.Rank, nil +} + +// PagedBackend supports scope-complete semantic retrieval without treating a +// global top-k cutoff as the final scoped candidate bound. +type PagedBackend interface { + SearchPage(ctx context.Context, generationID GenerationID, dimension int, query []float32, cursor string, limit int) (HitPage, error) +} + +// Backend stores unpublished vectors by opaque publication token. Publication +// authority remains in the main store; this interface deliberately knows +// nothing about archive or document identities encoded by callers. A canceled +// PutUnpublished must expose neither a partial batch nor a partially replaced +// token so claim-heartbeat loss can safely abort publication. +type Backend interface { + PutUnpublished(ctx context.Context, generationID GenerationID, dimension int, embeddings []Embedding) error + DeleteTokens(ctx context.Context, generationID GenerationID, tokens []string) error + Search(ctx context.Context, generationID GenerationID, dimension int, query []float32, k int) ([]Hit, error) +} diff --git a/internal/vector/document/fusion.go b/internal/vector/document/fusion.go new file mode 100644 index 000000000..f75a38980 --- /dev/null +++ b/internal/vector/document/fusion.go @@ -0,0 +1,128 @@ +package document + +import ( + "fmt" + "slices" + "sort" + + docembedding "go.kenn.io/docbank/document/embedding" + "go.kenn.io/msgvault/internal/store" +) + +type fusedSearchCandidate struct { + result store.DocumentSearchResult + lexical *store.DocumentSearchResult + semantic *store.DocumentSearchResult +} + +func fuseSearchResults( + lexical, semantic []store.DocumentSearchResult, limit int, +) ([]store.DocumentSearchResult, bool, error) { + byOccurrence := make(map[string]*fusedSearchCandidate, len(lexical)+len(semantic)) + for index := range lexical { + row := lexical[index] + candidate := byOccurrence[row.OccurrenceKey] + if candidate == nil { + candidate = &fusedSearchCandidate{} + byOccurrence[row.OccurrenceKey] = candidate + } + if candidate.lexical == nil || row.Rank < candidate.lexical.Rank { + rowCopy := row + candidate.lexical = &rowCopy + } + } + for index := range semantic { + row := semantic[index] + candidate := byOccurrence[row.OccurrenceKey] + if candidate == nil { + candidate = &fusedSearchCandidate{} + byOccurrence[row.OccurrenceKey] = candidate + } + if candidate.semantic == nil || semanticResultLess(row, *candidate.semantic) { + rowCopy := row + candidate.semantic = &rowCopy + } + } + lexicalCandidates := make([]docembedding.RankedCandidate, 0, len(byOccurrence)) + semanticCandidates := make([]docembedding.RankedCandidate, 0, len(byOccurrence)) + for _, candidate := range byOccurrence { + if candidate.lexical != nil { + lexicalCandidates = append(lexicalCandidates, docembedding.RankedCandidate{ + Key: candidate.lexical.OccurrenceKey, Rank: candidate.lexical.Rank, + }) + } + if candidate.semantic != nil { + semanticCandidates = append(semanticCandidates, docembedding.RankedCandidate{ + Key: candidate.semantic.OccurrenceKey, Rank: candidate.semantic.SemanticRank, + Score: candidate.semantic.SemanticScore, + }) + } + } + sort.Slice(lexicalCandidates, func(i, j int) bool { return lexicalCandidates[i].Rank < lexicalCandidates[j].Rank }) + sort.Slice(semanticCandidates, func(i, j int) bool { + if semanticCandidates[i].Rank != semanticCandidates[j].Rank { + return semanticCandidates[i].Rank < semanticCandidates[j].Rank + } + return semanticCandidates[i].Key < semanticCandidates[j].Key + }) + makeCandidateRanksStrict(lexicalCandidates) + makeCandidateRanksStrict(semanticCandidates) + fused, err := docembedding.FuseReciprocalRank(docembedding.FusionInput{ + Lexical: docembedding.ScopedCandidates{Candidates: lexicalCandidates}, + Semantic: docembedding.ScopedCandidates{Candidates: semanticCandidates}, + }, limit) + if err != nil { + return nil, false, fmt.Errorf("fuse Docbank ranked candidates: %w", err) + } + results := make([]store.DocumentSearchResult, 0, len(fused.Candidates)) + for _, shared := range fused.Candidates { + candidate := byOccurrence[shared.Key] + if candidate.lexical != nil { + candidate.result = *candidate.lexical + candidate.result.LexicalRank = candidate.lexical.Rank + candidate.result.MatchedSignals = slices.Clone(candidate.lexical.MatchedSignals) + if candidate.semantic != nil { + if candidate.result.PersonProvenance == nil { + candidate.result.PersonProvenance = candidate.semantic.PersonProvenance + } + candidate.result.SemanticRank = candidate.semantic.SemanticRank + candidate.result.SemanticScore = candidate.semantic.SemanticScore + candidate.result.VectorToken = candidate.semantic.VectorToken + candidate.result.VectorGenerationID = candidate.semantic.VectorGenerationID + candidate.result.VectorGenerationFingerprint = candidate.semantic.VectorGenerationFingerprint + candidate.result.VectorEmbeddingProfile = candidate.semantic.VectorEmbeddingProfile + candidate.result.VectorModel = candidate.semantic.VectorModel + candidate.result.VectorDimension = candidate.semantic.VectorDimension + candidate.result.MatchedSignals = append(candidate.result.MatchedSignals, "semantic") + } + } else { + candidate.result = *candidate.semantic + candidate.result.SemanticRank = candidate.semantic.SemanticRank + candidate.result.MatchedSignals = []string{"semantic"} + } + candidate.result.FusionScore = shared.Score + candidate.result.Rank = shared.Rank + results = append(results, candidate.result) + } + return results, fused.Truncated, nil +} + +func makeCandidateRanksStrict(candidates []docembedding.RankedCandidate) { + lastRank := 0 + for index := range candidates { + if candidates[index].Rank <= lastRank { + candidates[index].Rank = lastRank + 1 + } + lastRank = candidates[index].Rank + } +} + +func semanticResultLess(left, right store.DocumentSearchResult) bool { + if left.SemanticRank != right.SemanticRank { + return left.SemanticRank < right.SemanticRank + } + if left.SemanticScore != right.SemanticScore { + return left.SemanticScore > right.SemanticScore + } + return left.VectorToken < right.VectorToken +} diff --git a/internal/vector/document/fusion_test.go b/internal/vector/document/fusion_test.go new file mode 100644 index 000000000..c6f68a1fb --- /dev/null +++ b/internal/vector/document/fusion_test.go @@ -0,0 +1,127 @@ +package document + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/personscope" + "go.kenn.io/msgvault/internal/store" +) + +func TestFuseSearchResultsCollapsesDuplicateSemanticChunksPerOccurrence(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + semantic := []store.DocumentSearchResult{ + {OccurrenceKey: "occ-a", SemanticRank: 1, SemanticScore: .9, VectorToken: "token-a"}, + {OccurrenceKey: "occ-a", SemanticRank: 2, SemanticScore: .8, VectorToken: "token-duplicate"}, + {OccurrenceKey: "occ-b", SemanticRank: 3, SemanticScore: .7, VectorToken: "token-b"}, + } + + got, truncated, err := fuseSearchResults(nil, semantic, 10) + + requirements.NoError(err) + requirements.Len(got, 2) + assertions.False(truncated) + assertions.Equal("occ-a", got[0].OccurrenceKey) + assertions.Equal("token-a", got[0].VectorToken) + assertions.Equal(1, got[0].SemanticRank) + assertions.Equal([]string{"semantic"}, got[0].MatchedSignals) + assertions.Equal("occ-b", got[1].OccurrenceKey) + assertions.Equal(1, got[0].Rank) +} + +func TestFuseSearchResultsUsesOneRRFContributionPerSignal(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + provenance := &personscope.Provenance{ParticipantIDs: []int64{7}} + lexical := []store.DocumentSearchResult{ + {OccurrenceKey: "occ-b", Rank: 1, MatchedSignals: []string{"filename"}, PersonProvenance: provenance}, + } + semantic := []store.DocumentSearchResult{ + {OccurrenceKey: "occ-a", SemanticRank: 1, SemanticScore: .9, VectorToken: "token-a"}, + {OccurrenceKey: "occ-b", SemanticRank: 2, SemanticScore: .8, VectorToken: "token-b"}, + {OccurrenceKey: "occ-b", SemanticRank: 3, SemanticScore: .7, VectorToken: "token-duplicate"}, + } + + got, truncated, err := fuseSearchResults(lexical, semantic, 10) + + requirements.NoError(err) + requirements.Len(got, 2) + assertions.False(truncated) + assertions.Equal("occ-b", got[0].OccurrenceKey) + assertions.Equal([]string{"filename", "semantic"}, got[0].MatchedSignals) + assertions.Equal(1, got[0].LexicalRank) + assertions.Equal(2, got[0].SemanticRank) + assertions.Equal(provenance, got[0].PersonProvenance) + assertions.InDelta(1.0/61.0+1.0/62.0, got[0].FusionScore, 1e-12) + assertions.Equal("occ-a", got[1].OccurrenceKey) + assertions.InDelta(1.0/61.0, got[1].FusionScore, 1e-12) +} + +func TestFuseSearchResultsPreservesLaneRanksAndLexicalExcerpt(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + lexical := []store.DocumentSearchResult{ + { + OccurrenceKey: "occ-a", Rank: 4, MatchedSignals: []string{"content"}, + ChunkKey: "lexical-chunk", ChunkOrdinal: 2, HeadingPath: []string{"Lexical"}, + FirstUnitIndex: 4, LastUnitIndex: 5, + Excerpt: "matching lexical excerpt", HighlightStart: 9, HighlightEnd: 16, + }, + } + semantic := []store.DocumentSearchResult{ + { + OccurrenceKey: "occ-a", SemanticRank: 7, SemanticScore: .9, VectorToken: "token-a", + ChunkKey: "semantic-chunk", ChunkOrdinal: 8, HeadingPath: []string{"Semantic"}, + FirstUnitIndex: 12, LastUnitIndex: 13, Excerpt: "semantic excerpt", + }, + } + + got, truncated, err := fuseSearchResults(lexical, semantic, 10) + + requirements.NoError(err) + requirements.Len(got, 1) + assertions.False(truncated) + assertions.Equal(4, got[0].LexicalRank) + assertions.Equal(7, got[0].SemanticRank) + assertions.InDelta(1.0/64.0+1.0/67.0, got[0].FusionScore, 1e-12) + assertions.Equal("matching lexical excerpt", got[0].Excerpt) + assertions.Equal(9, got[0].HighlightStart) + assertions.Equal(16, got[0].HighlightEnd) + assertions.Equal("lexical-chunk", got[0].ChunkKey) + assertions.Equal(2, got[0].ChunkOrdinal) + assertions.Equal([]string{"Lexical"}, got[0].HeadingPath) + assertions.Equal(4, got[0].FirstUnitIndex) + assertions.Equal(5, got[0].LastUnitIndex) + assertions.Equal("token-a", got[0].VectorToken) +} + +func TestFuseSearchResultsBreaksSemanticTiesByOccurrenceIdentity(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + semantic := []store.DocumentSearchResult{ + {OccurrenceKey: "occ-z", SemanticRank: 1, SemanticScore: .5}, + {OccurrenceKey: "occ-a", SemanticRank: 1, SemanticScore: .5}, + } + + got, truncated, err := fuseSearchResults(nil, semantic, 10) + + require.NoError(err) + require.Len(got, 2) + assert.False(truncated) + assert.Equal([]string{"occ-a", "occ-z"}, []string{got[0].OccurrenceKey, got[1].OccurrenceKey}) +} + +func TestFuseSearchResultsReportsDisjointCandidateOverflow(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + lexical := []store.DocumentSearchResult{{OccurrenceKey: "occ-lexical", Rank: 1}} + semantic := []store.DocumentSearchResult{{OccurrenceKey: "occ-semantic", SemanticRank: 1}} + + got, truncated, err := fuseSearchResults(lexical, semantic, 1) + + require.NoError(err) + require.Len(got, 1) + assert.True(truncated) +} diff --git a/internal/vector/document/model.go b/internal/vector/document/model.go new file mode 100644 index 000000000..6f4b9457d --- /dev/null +++ b/internal/vector/document/model.go @@ -0,0 +1,165 @@ +// Package document defines the independent attachment-document vector corpus. +package document + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "strings" + + docembedding "go.kenn.io/docbank/document/embedding" + "go.kenn.io/msgvault/internal/vector" +) + +// GenerationID identifies one document-vector generation. +type GenerationID int64 + +// GenerationState describes a document-vector generation's lifecycle state. +type GenerationState string + +const ( + GenerationBuilding GenerationState = "building" + GenerationActive GenerationState = "active" + GenerationRetired GenerationState = "retired" +) + +// GenerationSpec is the immutable identity of one document-vector corpus +// generation. Its Fingerprint is derived from the extraction profile and +// embedding policy that determine the stored vectors. +type GenerationSpec struct { + Fingerprint string + ExtractionProfileID string + EmbeddingProfile string + Model string + Dimension int +} + +// SearchMode uses Docbank's shared retrieval modes. +type SearchMode = docembedding.SearchMode + +const ( + SearchModeAuto = docembedding.SearchModeAuto + SearchModeLexical = docembedding.SearchModeLexical + SearchModeSemantic = docembedding.SearchModeSemantic + SearchModeHybrid = docembedding.SearchModeHybrid +) + +// ParseSearchMode normalizes a document search mode. Omitted values normalize +// to auto; SearchService keeps auto lexical so query text leaves the process +// only on an explicit semantic or hybrid request. +func ParseSearchMode(value string) (SearchMode, error) { + mode := docembedding.SearchMode(strings.ToLower(strings.TrimSpace(value))) + if mode == "" { + return SearchModeAuto, nil + } + switch mode { + case SearchModeAuto, SearchModeLexical, SearchModeSemantic, SearchModeHybrid: + return mode, nil + default: + return "", fmt.Errorf("unsupported document search mode %q", value) + } +} + +// EmbeddingRecipe returns the shared raw preparation recipe used by this PR. +func EmbeddingRecipe(cfg vector.Config) (docembedding.Recipe, error) { + recipe, err := docembedding.NewRecipe(docembedding.RecipeConfig{ + Mode: docembedding.RepresentationRaw, MaxInputRunes: cfg.Embeddings.MaxInputChars, + }) + if err != nil { + return docembedding.Recipe{}, fmt.Errorf( + "construct document embedding recipe from vector.embeddings.max_input_chars=%d: %w", + cfg.Embeddings.MaxInputChars, err, + ) + } + return recipe, nil +} + +func vectorSpaceIdentity(cfg vector.Config) (docembedding.VectorSpaceIdentity, error) { + deploymentFingerprint, err := (docembedding.EgressIdentity{ + Purpose: docembedding.EgressDocumentEmbedding, Provider: string(cfg.Embeddings.EffectiveAPIFormat()), + Endpoint: cfg.Embeddings.Endpoint, Model: cfg.Embeddings.Model, ModelRevision: cfg.Embeddings.Model, + }).Fingerprint() + if err != nil { + return docembedding.VectorSpaceIdentity{}, fmt.Errorf("fingerprint document embedding deployment: %w", err) + } + return docembedding.VectorSpaceIdentity{ + Provider: string(cfg.Embeddings.EffectiveAPIFormat()), Model: cfg.Embeddings.Model, + ModelRevision: deploymentFingerprint, Dimension: cfg.Embeddings.Dimension, + Normalization: "provider-output-v1", + }, nil +} + +// Fingerprint returns the immutable, non-secret document-vector corpus +// identity. The canonical endpoint is represented only by a deployment hash so +// OpenAI-compatible services cannot reuse each other's vector spaces. Credential +// location, message preprocessing, and message scope remain excluded. +func Fingerprint(extractionProfileID string, cfg vector.Config) (string, error) { + recipe, err := EmbeddingRecipe(cfg) + if err != nil { + return "", err + } + spaceIdentity, err := vectorSpaceIdentity(cfg) + if err != nil { + return "", err + } + spaceFingerprint, err := spaceIdentity.Fingerprint() + if err != nil { + return "", fmt.Errorf("fingerprint document vector space: %w", err) + } + payload := struct { + Version int `json:"version"` + ExtractionProfileID string `json:"extraction_profile_id"` + RecipeFingerprint string `json:"recipe_fingerprint"` + VectorSpace string `json:"vector_space_fingerprint"` + }{ + Version: 1, ExtractionProfileID: extractionProfileID, + RecipeFingerprint: recipe.Fingerprint(), VectorSpace: spaceFingerprint, + } + encoded, err := json.Marshal(payload) + if err != nil { + return "", fmt.Errorf("marshal document vector fingerprint: %w", err) + } + hash := sha256.Sum256(encoded) + return hex.EncodeToString(hash[:]), nil +} + +// EgressFingerprint binds operator consent to both Msgvault's exact corpus +// policy and Docbank's canonical provider destination. +func EgressFingerprint(extractionProfileID string, cfg vector.Config) (string, error) { + return egressFingerprint(extractionProfileID, docembedding.EgressDocumentEmbedding, cfg) +} + +// QueryEgressFingerprint is purpose-separated from document-text consent. +func QueryEgressFingerprint(extractionProfileID string, cfg vector.Config) (string, error) { + return egressFingerprint(extractionProfileID, docembedding.EgressQueryEmbedding, cfg) +} + +func egressFingerprint(extractionProfileID string, purpose docembedding.EgressPurpose, cfg vector.Config) (string, error) { + destination, err := (docembedding.EgressIdentity{ + Purpose: purpose, Provider: string(cfg.Embeddings.EffectiveAPIFormat()), + Endpoint: cfg.Embeddings.Endpoint, Model: cfg.Embeddings.Model, + ModelRevision: cfg.Embeddings.Model, + }).Fingerprint() + if err != nil { + return "", fmt.Errorf("fingerprint document egress destination: %w", err) + } + generationFingerprint, err := Fingerprint(extractionProfileID, cfg) + if err != nil { + return "", err + } + payload := struct { + Version int `json:"version"` + GenerationFingerprint string `json:"generation_fingerprint"` + DestinationFingerprint string `json:"destination_fingerprint"` + }{ + Version: 1, GenerationFingerprint: generationFingerprint, + DestinationFingerprint: destination, + } + encoded, err := json.Marshal(payload) + if err != nil { + return "", fmt.Errorf("marshal document egress consent identity: %w", err) + } + hash := sha256.Sum256(encoded) + return hex.EncodeToString(hash[:]), nil +} diff --git a/internal/vector/document/model_test.go b/internal/vector/document/model_test.go new file mode 100644 index 000000000..7d427d66a --- /dev/null +++ b/internal/vector/document/model_test.go @@ -0,0 +1,121 @@ +package document + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/vector" +) + +func requireFingerprint(t *testing.T, extractionProfileID string, config vector.Config) string { + t.Helper() + fingerprint, err := Fingerprint(extractionProfileID, config) + require.NoError(t, err) + return fingerprint +} + +func TestFingerprintBindsDocumentExtractionAndEmbeddingPolicy(t *testing.T) { + config := vector.Config{Embeddings: vector.EmbeddingsConfig{ + Endpoint: "https://embeddings.example.test/v1", + APIFormat: vector.APIFormatOpenAI, Model: "embed-v1", Dimension: 768, MaxInputChars: 8192, + }} + + baseline := requireFingerprint(t, "extract-v1", config) + assert.Regexp(t, "^[0-9a-f]{64}$", baseline) + assert.Equal(t, baseline, requireFingerprint(t, "extract-v1", config)) + + tests := []struct { + name string + mutate func(*vector.Config) + }{ + {name: "api format", mutate: func(c *vector.Config) { c.Embeddings.APIFormat = vector.APIFormatVoyageContextual }}, + {name: "endpoint", mutate: func(c *vector.Config) { c.Embeddings.Endpoint = "https://other.example.test/v1" }}, + {name: "model", mutate: func(c *vector.Config) { c.Embeddings.Model = "embed-v2" }}, + {name: "dimension", mutate: func(c *vector.Config) { c.Embeddings.Dimension++ }}, + {name: "max input chars", mutate: func(c *vector.Config) { c.Embeddings.MaxInputChars++ }}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + changed := config + test.mutate(&changed) + assert.NotEqual(t, baseline, requireFingerprint(t, "extract-v1", changed)) + }) + } + assert.NotEqual(t, baseline, requireFingerprint(t, "extract-v2", config)) +} + +func TestFingerprintExcludesMessageCorpusAndCredentials(t *testing.T) { + config := vector.Config{Embeddings: vector.EmbeddingsConfig{ + Endpoint: "https://embeddings.example.test/v1", + APIFormat: vector.APIFormatOpenAI, Model: "embed-v1", Dimension: 768, MaxInputChars: 8192, + }} + baseline := requireFingerprint(t, "extract-v1", config) + + changed := config + changed.Embeddings.APIKeyEnv = "SYNTHETIC_EMBEDDING_KEY" + stripQuotes := false + changed.Preprocess.StripQuotes = &stripQuotes + changed.Embed.Scope.MessageTypes = []string{"email"} + changed.Embed.Scope.SourceIDs = []int64{7} + + assert.Equal(t, baseline, requireFingerprint(t, "extract-v1", changed)) +} + +func TestEgressFingerprintBindsCanonicalDestinationAndCorpus(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + config := vector.Config{Embeddings: vector.EmbeddingsConfig{ + Endpoint: "https://trusted.example.test/v1", + APIFormat: vector.APIFormatOpenAI, Model: "embed-v1", Dimension: 768, MaxInputChars: 8192, + }} + corpus := requireFingerprint(t, "extract-v1", config) + baseline, err := EgressFingerprint("extract-v1", config) + requirements.NoError(err) + + credentialChange := config + credentialChange.Embeddings.APIKeyEnv = "OTHER_SYNTHETIC_EMBEDDING_KEY" + credentialFingerprint, err := EgressFingerprint("extract-v1", credentialChange) + requirements.NoError(err) + assertions.Equal(baseline, credentialFingerprint) + + hosted := config + hosted.Embeddings.Endpoint = "https://hosted.example.test/v1" + hostedFingerprint, err := EgressFingerprint("extract-v1", hosted) + requirements.NoError(err) + assertions.NotEqual(baseline, hostedFingerprint) + assertions.NotEqual(corpus, requireFingerprint(t, "extract-v1", hosted)) + rotatedCorpusFingerprint, err := EgressFingerprint("extract-v2", config) + requirements.NoError(err) + assertions.NotEqual(baseline, rotatedCorpusFingerprint) + queryFingerprint, err := QueryEgressFingerprint("extract-v1", config) + requirements.NoError(err) + assertions.NotEqual(baseline, queryFingerprint) + + unsafeEndpoint := config + unsafeEndpoint.Embeddings.Endpoint = "https://user:secret@trusted.example.test/v1?api_key=secret" + _, err = EgressFingerprint("extract-v1", unsafeEndpoint) + requirements.Error(err) +} + +func TestParseSearchMode(t *testing.T) { + for _, test := range []struct { + input string + want SearchMode + }{ + {input: "", want: SearchModeAuto}, + {input: " AUTO ", want: SearchModeAuto}, + {input: "lexical", want: SearchModeLexical}, + {input: "SEMANTIC", want: SearchModeSemantic}, + {input: "hybrid", want: SearchModeHybrid}, + } { + t.Run(test.input, func(t *testing.T) { + got, err := ParseSearchMode(test.input) + require.NoError(t, err) + assert.Equal(t, test.want, got) + }) + } + + _, err := ParseSearchMode("unsupported") + require.ErrorContains(t, err, "unsupported") +} diff --git a/internal/vector/document/provider.go b/internal/vector/document/provider.go new file mode 100644 index 000000000..5441feff8 --- /dev/null +++ b/internal/vector/document/provider.go @@ -0,0 +1,13 @@ +package document + +import ( + "context" + + "go.kenn.io/msgvault/internal/vector" +) + +// Provider is the document-building subset of embed.SemanticClient. +// Query embedding belongs to retrieval, not generation publication. +type Provider interface { + EmbedDocuments(ctx context.Context, documents []vector.DocumentInput) ([][][]float32, error) +} diff --git a/internal/vector/document/reconcile.go b/internal/vector/document/reconcile.go new file mode 100644 index 000000000..cbfa23c3e --- /dev/null +++ b/internal/vector/document/reconcile.go @@ -0,0 +1,252 @@ +package document + +import ( + "context" + "errors" + "fmt" + "time" + + "go.kenn.io/msgvault/internal/store" +) + +const maxReconcileLimit = 1000 + +type WorkerRunner interface { + Run(ctx context.Context, generationID GenerationID, limit int) (RunResult, error) +} + +type ReconcileLedger interface { + GetDocumentVectorGeneration(ctx context.Context, id int64) (store.DocumentVectorGeneration, error) + GetDocumentVectorGenerationStatus(ctx context.Context, generationID int64, afterToken string, limit int) (store.DocumentVectorGenerationStatus, error) + ParkObsoleteDocumentVectorTokens(ctx context.Context, generationID int64, afterToken string, limit int, now time.Time) (store.DocumentVectorCleanupPage, error) + FinalizeObsoleteDocumentVectorToken(ctx context.Context, generationID int64, token string, now time.Time) (bool, error) + GetDocumentVectorCoverage(ctx context.Context, generationID int64) (store.DocumentVectorCoverage, error) + ActivateDocumentVectorGeneration(ctx context.Context, generationID int64, now time.Time) error + PurgeRetiredDocumentVectorGeneration(ctx context.Context, generationID int64) (bool, error) +} + +var ( + _ WorkerRunner = (*Worker)(nil) + _ ReconcileLedger = (*store.Store)(nil) +) + +type ReconcilerDeps struct { + Ledger ReconcileLedger + Worker WorkerRunner + Backend Backend + Now func() time.Time + + AfterCleanupGenerationID GenerationID + AfterCleanupToken string +} + +type ReconcileResult struct { + WorkerRan bool `json:"worker_ran"` + Worker RunResult `json:"worker"` + + CleanupListed int `json:"cleanup_listed"` + CleanupDeleted int `json:"cleanup_deleted"` + CleanupFinalized int `json:"cleanup_finalized"` + CleanupAfterGenerationID GenerationID `json:"cleanup_after_generation_id,omitempty"` + CleanupAfterToken string `json:"cleanup_after_token,omitempty"` + CleanupExhausted bool `json:"cleanup_exhausted"` + + Status GenerationStatus `json:"status"` + Coverage store.DocumentVectorCoverage `json:"coverage"` + Blocked bool `json:"blocked"` + Activated bool `json:"activated"` + Purged bool `json:"purged"` + Converged bool `json:"converged"` +} + +type Reconciler struct { + deps ReconcilerDeps + cleanupCursorGenerationID GenerationID + cleanupAfterToken string +} + +func NewReconciler(deps ReconcilerDeps) *Reconciler { + if deps.Now == nil { + deps.Now = func() time.Time { return time.Now().UTC() } + } + return &Reconciler{ + deps: deps, cleanupCursorGenerationID: deps.AfterCleanupGenerationID, + cleanupAfterToken: deps.AfterCleanupToken, + } +} + +func (r *Reconciler) Run(ctx context.Context, generationID GenerationID, limit int) (ReconcileResult, error) { + var result ReconcileResult + if err := r.validate(generationID, limit); err != nil { + return result, err + } + generation, err := r.deps.Ledger.GetDocumentVectorGeneration(ctx, int64(generationID)) + if err != nil { + return result, fmt.Errorf("read document vector generation for reconciliation: %w", err) + } + if generation.ID != int64(generationID) { + return result, store.ErrDocumentVectorInvalidGenerationState + } + if generation.State == store.DocumentVectorGenerationBuilding && r.deps.Worker == nil { + return result, errors.New("document vector reconciliation worker is required for a building generation") + } + r.bindCleanupCursor(generationID) + + var reconcileErr error + if generation.State == store.DocumentVectorGenerationBuilding { + result.WorkerRan = true + result.Worker, err = r.deps.Worker.Run(ctx, generationID, limit) + if ctxErr := ctx.Err(); ctxErr != nil { + return result, ctxErr + } + reconcileErr = errors.Join(reconcileErr, err) + } + + cleanupErr := r.cleanup(ctx, generationID, limit, &result) + if ctxErr := ctx.Err(); ctxErr != nil { + return result, errors.Join(reconcileErr, ctxErr) + } + reconcileErr = errors.Join(reconcileErr, cleanupErr) + + status, statusErr := r.deps.Ledger.GetDocumentVectorGenerationStatus(ctx, int64(generationID), "", limit) + if statusErr == nil { + result.Status = status + if status.CleanupPending == 0 { + r.resetCleanupCursor(&result) + } + } + reconcileErr = errors.Join(reconcileErr, statusErr) + if generation.State == store.DocumentVectorGenerationBuilding && statusErr == nil && status.Blocked { + result.Blocked = true + reconcileErr = errors.Join(reconcileErr, store.ErrDocumentVectorGenerationBlocked) + } + + switch generation.State { + case store.DocumentVectorGenerationBuilding, store.DocumentVectorGenerationActive: + coverage, coverageErr := r.deps.Ledger.GetDocumentVectorCoverage(ctx, int64(generationID)) + if coverageErr == nil { + result.Coverage = coverage + } + reconcileErr = errors.Join(reconcileErr, coverageErr) + if generation.State == store.DocumentVectorGenerationBuilding && coverageErr == nil && + coverage.Complete() && reconcileErr == nil { + activationErr := r.deps.Ledger.ActivateDocumentVectorGeneration(ctx, int64(generationID), r.now()) + if activationErr == nil { + result.Activated = true + if refreshed, refreshErr := r.deps.Ledger.GetDocumentVectorGenerationStatus(ctx, int64(generationID), "", limit); refreshErr == nil { + result.Status = refreshed + } else { + reconcileErr = refreshErr + } + } else { + reconcileErr = activationErr + } + } + result.Converged = (generation.State == store.DocumentVectorGenerationActive || result.Activated) && + result.Coverage.Complete() && result.Status.CleanupPending == 0 && reconcileErr == nil + case store.DocumentVectorGenerationRetired: + if result.Status.CleanupPending == 0 && reconcileErr == nil { + result.Purged, err = r.deps.Ledger.PurgeRetiredDocumentVectorGeneration(ctx, int64(generationID)) + reconcileErr = errors.Join(reconcileErr, err) + } + result.Converged = result.Purged && reconcileErr == nil + default: + reconcileErr = errors.Join(reconcileErr, store.ErrDocumentVectorInvalidGenerationState) + } + return result, reconcileErr +} + +func (r *Reconciler) cleanup( + ctx context.Context, generationID GenerationID, limit int, result *ReconcileResult, +) error { + page, err := r.deps.Ledger.ParkObsoleteDocumentVectorTokens( + ctx, int64(generationID), r.cleanupAfterToken, limit, r.now(), + ) + if err != nil { + return fmt.Errorf("park obsolete document vector tokens: %w", err) + } + result.CleanupListed = len(page.Tokens) + if len(page.Tokens) == 0 { + r.resetCleanupCursor(result) + return nil + } + opaque := make([]string, len(page.Tokens)) + for index := range page.Tokens { + opaque[index] = page.Tokens[index].Token + } + if err := r.deps.Backend.DeleteTokens(ctx, generationID, opaque); err != nil { + r.invalidateCleanupCursor(result) + return fmt.Errorf("delete obsolete document vector tokens: %w", err) + } + result.CleanupDeleted = len(opaque) + var finalizeErr error + for _, token := range opaque { + finalized, err := r.deps.Ledger.FinalizeObsoleteDocumentVectorToken( + ctx, int64(generationID), token, r.now(), + ) + if err != nil { + finalizeErr = errors.Join(finalizeErr, fmt.Errorf("finalize obsolete document vector token %q: %w", token, err)) + continue + } + if finalized { + result.CleanupFinalized++ + } + } + if finalizeErr != nil { + r.invalidateCleanupCursor(result) + return finalizeErr + } + if page.Exhausted { + r.resetCleanupCursor(result) + } else { + r.cleanupCursorGenerationID = GenerationID(page.AfterGenerationID) + r.cleanupAfterToken = page.AfterToken + result.CleanupAfterGenerationID = GenerationID(page.AfterGenerationID) + result.CleanupAfterToken = page.AfterToken + } + return nil +} + +func (r *Reconciler) validate(generationID GenerationID, limit int) error { + if generationID <= 0 { + return errors.New("document vector reconciliation generation must be positive") + } + if limit < 1 || limit > maxReconcileLimit { + return fmt.Errorf("document vector reconciliation limit must be between 1 and %d", maxReconcileLimit) + } + if r == nil || r.deps.Ledger == nil || r.deps.Backend == nil || r.deps.Now == nil { + return errors.New("document vector reconciliation dependencies are required") + } + if r.deps.AfterCleanupGenerationID < 0 || + (r.deps.AfterCleanupGenerationID == 0) != (r.deps.AfterCleanupToken == "") { + return errors.New("document vector reconciliation cleanup cursor is invalid") + } + return nil +} + +func (r *Reconciler) now() time.Time { + return r.deps.Now().UTC().Truncate(time.Millisecond) +} + +func (r *Reconciler) bindCleanupCursor(generationID GenerationID) { + if r.cleanupCursorGenerationID != 0 && r.cleanupCursorGenerationID != generationID { + r.cleanupAfterToken = "" + } + r.cleanupCursorGenerationID = generationID +} + +func (r *Reconciler) resetCleanupCursor(result *ReconcileResult) { + r.cleanupCursorGenerationID = 0 + r.cleanupAfterToken = "" + result.CleanupAfterGenerationID = 0 + result.CleanupAfterToken = "" + result.CleanupExhausted = true +} + +func (r *Reconciler) invalidateCleanupCursor(result *ReconcileResult) { + r.cleanupCursorGenerationID = 0 + r.cleanupAfterToken = "" + result.CleanupAfterGenerationID = 0 + result.CleanupAfterToken = "" + result.CleanupExhausted = false +} diff --git a/internal/vector/document/reconcile_test.go b/internal/vector/document/reconcile_test.go new file mode 100644 index 000000000..988de5579 --- /dev/null +++ b/internal/vector/document/reconcile_test.go @@ -0,0 +1,387 @@ +package document + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/store" +) + +func TestReconcilerRunActivatesCompleteBuildingGeneration(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + worker := &fakeDocumentVectorWorkerRunner{result: RunResult{Claimed: 1, Embedded: 1, Published: 1}} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.coverage = store.DocumentVectorCoverage{Required: 1, Ready: 1} + backend := &fakeDocumentVectorBackend{} + reconciler := NewReconciler(ReconcilerDeps{ + Ledger: ledger, Worker: worker, Backend: backend, + Now: func() time.Time { return workerNow }, + }) + + result, err := reconciler.Run(t.Context(), 1, 10) + requirements.NoError(err) + assertions.True(result.WorkerRan) + assertions.Equal(worker.result, result.Worker) + assertions.Equal(1, worker.calls) + assertions.Equal([]int{10}, worker.limits) + assertions.Equal(ledger.coverage, result.Coverage) + assertions.True(result.Activated) + assertions.True(result.Converged) + assertions.Equal(1, ledger.activateCalls) + assertions.Zero(ledger.purgeCalls) + assertions.Empty(backend.deletes) +} + +func TestReconcilerRunRejectsInvalidRequestAndMissingBuildingWorker(t *testing.T) { + requirements := require.New(t) + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + reconciler := NewReconciler(ReconcilerDeps{ + Ledger: ledger, Backend: &fakeDocumentVectorBackend{}, Now: func() time.Time { return workerNow }, + }) + + _, err := reconciler.Run(t.Context(), 0, 1) + requirements.ErrorContains(err, "generation") + _, err = reconciler.Run(t.Context(), 1, 0) + requirements.ErrorContains(err, "limit") + _, err = reconciler.Run(t.Context(), 1, 1001) + requirements.ErrorContains(err, "limit") + _, err = reconciler.Run(t.Context(), 1, 1) + requirements.ErrorContains(err, "worker") +} + +func TestReconcilerRunTimestampsCleanupAndActivationAfterWorker(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + worker := &fakeDocumentVectorWorkerRunner{} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.coverage = store.DocumentVectorCoverage{Required: 1, Ready: 1} + ledger.obsolete = []store.DocumentVectorCleanupToken{{GenerationID: 1, Token: "dddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddd"}} + ledger.status.CleanupPending = 1 + base := workerNow.Add(123456 * time.Nanosecond) + clockCalls := 0 + reconciler := NewReconciler(ReconcilerDeps{ + Ledger: ledger, Worker: worker, Backend: &fakeDocumentVectorBackend{}, + Now: func() time.Time { + at := base.Add(time.Duration(clockCalls) * time.Second) + clockCalls++ + return at + }, + }) + + _, err := reconciler.Run(t.Context(), 1, 10) + requirements.NoError(err) + assertions.Equal(base.UTC().Truncate(time.Millisecond), ledger.parkAt) + assertions.Equal(base.Add(time.Second).UTC().Truncate(time.Millisecond), ledger.finalizeAt) + assertions.Equal(base.Add(2*time.Second).UTC().Truncate(time.Millisecond), ledger.activateAt) +} + +func TestReconcilerRunContinuesSafeWorkAfterWorkerErrorWithoutActivation(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + workerErr := errors.New("provider unavailable") + worker := &fakeDocumentVectorWorkerRunner{err: workerErr} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.coverage = store.DocumentVectorCoverage{Required: 1, Ready: 1} + ledger.obsolete = []store.DocumentVectorCleanupToken{{GenerationID: 1, Token: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}} + ledger.status.CleanupPending = 1 + backend := &fakeDocumentVectorBackend{} + reconciler := NewReconciler(ReconcilerDeps{Ledger: ledger, Worker: worker, Backend: backend, Now: func() time.Time { return workerNow }}) + + result, err := reconciler.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, workerErr) + assertions.Equal(1, worker.calls) + assertions.Equal(1, result.CleanupDeleted) + assertions.Equal(1, result.CleanupFinalized) + assertions.Equal(1, ledger.statusCalls) + assertions.Equal(1, ledger.coverageCalls) + assertions.Zero(ledger.activateCalls) + assertions.False(result.Activated) + assertions.False(result.Converged) + assertions.Equal([][]string{{ledger.obsolete[0].Token}}, backend.deletes) +} + +func TestReconcilerRunReportsTerminalGenerationAsBlocked(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.status.Terminal = 1 + ledger.status.Blocked = true + ledger.coverage = store.DocumentVectorCoverage{Required: 2, Ready: 1} + reconciler := NewReconciler(ReconcilerDeps{ + Ledger: ledger, Worker: &fakeDocumentVectorWorkerRunner{}, Backend: &fakeDocumentVectorBackend{}, + Now: func() time.Time { return workerNow }, + }) + + result, err := reconciler.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, store.ErrDocumentVectorGenerationBlocked) + assertions.True(result.Blocked) + assertions.False(result.Activated) + assertions.False(result.Converged) + assertions.Zero(ledger.activateCalls) +} + +func TestReconcilerRunNeverWorksActiveGeneration(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationActive) + ledger.coverage = store.DocumentVectorCoverage{Required: 2, Ready: 2} + reconciler := NewReconciler(ReconcilerDeps{ + Ledger: ledger, Backend: &fakeDocumentVectorBackend{}, + Now: func() time.Time { return workerNow }, + }) + + result, err := reconciler.Run(t.Context(), 1, 10) + requirements.NoError(err) + assertions.False(result.WorkerRan) + assertions.True(result.Converged) + assertions.Zero(ledger.activateCalls) + assertions.Zero(ledger.purgeCalls) +} + +func TestReconcilerRunCleansRetiredGenerationWithoutWorkerDependency(t *testing.T) { + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationRetired) + ledger.purgeResult = true + reconciler := NewReconciler(ReconcilerDeps{ + Ledger: ledger, Backend: &fakeDocumentVectorBackend{}, Now: func() time.Time { return workerNow }, + }) + + result, err := reconciler.Run(t.Context(), 1, 10) + require.NoError(t, err) + assert.False(t, result.WorkerRan) + assert.True(t, result.Purged) +} + +func TestReconcilerRunReplaysRetiredDeleteAfterFinalizeCrashThenPurges(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + finalizeErr := errors.New("main store unavailable") + worker := &fakeDocumentVectorWorkerRunner{} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationRetired) + ledger.obsolete = []store.DocumentVectorCleanupToken{{GenerationID: 1, Token: "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"}} + ledger.status.CleanupPending = 1 + ledger.finalizeErr = finalizeErr + ledger.purgeResult = true + backend := &fakeDocumentVectorBackend{} + reconciler := NewReconciler(ReconcilerDeps{Ledger: ledger, Worker: worker, Backend: backend, Now: func() time.Time { return workerNow }}) + + first, err := reconciler.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, finalizeErr) + assertions.Zero(worker.calls) + assertions.Equal(1, first.CleanupDeleted) + assertions.Zero(first.CleanupFinalized) + assertions.Zero(first.CleanupAfterGenerationID) + assertions.Empty(first.CleanupAfterToken) + assertions.Zero(ledger.purgeCalls) + + ledger.finalizeErr = nil + second, err := reconciler.Run(t.Context(), 1, 10) + requirements.NoError(err) + assertions.Equal(1, second.CleanupDeleted) + assertions.Equal(1, second.CleanupFinalized) + assertions.True(second.Purged) + assertions.True(second.Converged) + assertions.Equal(1, ledger.purgeCalls) + assertions.Equal([][]string{{ledger.obsolete[0].Token}, {ledger.obsolete[0].Token}}, backend.deletes) +} + +func TestReconcilerRunReplaysParkedPageAfterBackendDeleteFailure(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + deleteErr := errors.New("backend unavailable") + worker := &fakeDocumentVectorWorkerRunner{} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.obsolete = []store.DocumentVectorCleanupToken{{GenerationID: 1, Token: "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc"}} + ledger.status.CleanupPending = 1 + backend := &fakeDocumentVectorBackend{deleteErr: deleteErr} + reconciler := NewReconciler(ReconcilerDeps{Ledger: ledger, Worker: worker, Backend: backend, Now: func() time.Time { return workerNow }}) + + first, err := reconciler.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, deleteErr) + assertions.Equal(1, first.CleanupListed) + assertions.Zero(first.CleanupFinalized) + assertions.Zero(first.CleanupAfterGenerationID) + assertions.Empty(first.CleanupAfterToken) + assertions.False(first.CleanupExhausted) + + backend.deleteErr = nil + second, err := reconciler.Run(t.Context(), 1, 10) + requirements.NoError(err) + assertions.Equal(1, second.CleanupDeleted) + assertions.Equal(1, second.CleanupFinalized) + assertions.Equal([][]string{{ledger.obsolete[0].Token}, {ledger.obsolete[0].Token}}, backend.deletes) +} + +func TestReconcilerRunDoesNotAdvancePastPartialFinalizeFailure(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + finalizeErr := errors.New("finalize interrupted") + worker := &fakeDocumentVectorWorkerRunner{} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.obsolete = []store.DocumentVectorCleanupToken{ + {GenerationID: 1, Token: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}, + {GenerationID: 1, Token: "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"}, + } + ledger.status.CleanupPending = 2 + ledger.finalizeErrFor = ledger.obsolete[1].Token + ledger.finalizeErr = finalizeErr + backend := &fakeDocumentVectorBackend{} + reconciler := NewReconciler(ReconcilerDeps{Ledger: ledger, Worker: worker, Backend: backend, Now: func() time.Time { return workerNow }}) + + first, err := reconciler.Run(t.Context(), 1, 2) + requirements.ErrorIs(err, finalizeErr) + assertions.Equal(1, first.CleanupFinalized) + assertions.Zero(first.CleanupAfterGenerationID) + assertions.Empty(first.CleanupAfterToken) + + ledger.finalizeErr = nil + second, err := reconciler.Run(t.Context(), 1, 2) + requirements.NoError(err) + assertions.Equal(1, second.CleanupListed) + assertions.Equal(ledger.obsolete[1].Token, ledger.finalized[len(ledger.finalized)-1]) + assertions.Equal([][]string{ + {ledger.obsolete[0].Token, ledger.obsolete[1].Token}, + {ledger.obsolete[1].Token}, + }, backend.deletes) +} + +func TestReconcilerRunBoundsCleanupAndReturnsRestorableCursor(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + worker := &fakeDocumentVectorWorkerRunner{} + ledger := newFakeDocumentVectorReconcileLedger(store.DocumentVectorGenerationBuilding) + ledger.coverage = store.DocumentVectorCoverage{Required: 1} + ledger.obsolete = []store.DocumentVectorCleanupToken{ + {GenerationID: 1, Token: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}, + {GenerationID: 1, Token: "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"}, + } + ledger.status.CleanupPending = 2 + backend := &fakeDocumentVectorBackend{} + reconciler := NewReconciler(ReconcilerDeps{Ledger: ledger, Worker: worker, Backend: backend, Now: func() time.Time { return workerNow }}) + + first, err := reconciler.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.Equal(1, first.CleanupListed) + assertions.Equal(GenerationID(1), first.CleanupAfterGenerationID) + assertions.Equal(ledger.obsolete[0].Token, first.CleanupAfterToken) + assertions.False(first.CleanupExhausted) + + second, err := reconciler.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.Equal(1, second.CleanupListed) + assertions.Zero(second.CleanupAfterGenerationID) + assertions.Empty(second.CleanupAfterToken) + assertions.True(second.CleanupExhausted) + assertions.Equal([][]string{{ledger.obsolete[0].Token}, {ledger.obsolete[1].Token}}, backend.deletes) +} + +type fakeDocumentVectorWorkerRunner struct { + result RunResult + err error + calls int + limits []int +} + +func (w *fakeDocumentVectorWorkerRunner) Run(_ context.Context, _ GenerationID, limit int) (RunResult, error) { + w.calls++ + w.limits = append(w.limits, limit) + return w.result, w.err +} + +type fakeDocumentVectorReconcileLedger struct { + generation store.DocumentVectorGeneration + status store.DocumentVectorGenerationStatus + coverage store.DocumentVectorCoverage + obsolete []store.DocumentVectorCleanupToken + finalized []string + finalizeErr error + finalizeErrFor string + activateErr error + purgeResult bool + purgeErr error + activateCalls int + purgeCalls int + statusCalls int + coverageCalls int + finalizedSet map[string]bool + parkAt time.Time + finalizeAt time.Time + activateAt time.Time +} + +func newFakeDocumentVectorReconcileLedger(state store.DocumentVectorGenerationState) *fakeDocumentVectorReconcileLedger { + return &fakeDocumentVectorReconcileLedger{ + generation: store.DocumentVectorGeneration{ID: 1, State: state}, + status: store.DocumentVectorGenerationStatus{GenerationID: 1, State: state, FailuresExhausted: true}, + finalizedSet: map[string]bool{}, + } +} + +func (l *fakeDocumentVectorReconcileLedger) GetDocumentVectorGeneration(context.Context, int64) (store.DocumentVectorGeneration, error) { + return l.generation, nil +} + +func (l *fakeDocumentVectorReconcileLedger) GetDocumentVectorGenerationStatus(context.Context, int64, string, int) (store.DocumentVectorGenerationStatus, error) { + l.statusCalls++ + return l.status, nil +} + +func (l *fakeDocumentVectorReconcileLedger) ParkObsoleteDocumentVectorTokens(_ context.Context, generationID int64, after string, limit int, now time.Time) (store.DocumentVectorCleanupPage, error) { + l.parkAt = now + page := store.DocumentVectorCleanupPage{} + var tokens []store.DocumentVectorCleanupToken + for _, token := range l.obsolete { + if token.Token <= after || l.finalizedSet[token.Token] { + continue + } + tokens = append(tokens, token) + if len(tokens) == limit { + break + } + } + page.Tokens = tokens + if len(tokens) < limit { + page.Exhausted = true + } else { + page.AfterGenerationID = generationID + page.AfterToken = tokens[len(tokens)-1].Token + } + return page, nil +} + +func (l *fakeDocumentVectorReconcileLedger) FinalizeObsoleteDocumentVectorToken(_ context.Context, _ int64, token string, now time.Time) (bool, error) { + l.finalizeAt = now + if l.finalizeErr != nil && (l.finalizeErrFor == "" || l.finalizeErrFor == token) { + return false, l.finalizeErr + } + l.finalized = append(l.finalized, token) + l.finalizedSet[token] = true + if l.status.CleanupPending > 0 { + l.status.CleanupPending-- + } + return true, nil +} + +func (l *fakeDocumentVectorReconcileLedger) GetDocumentVectorCoverage(context.Context, int64) (store.DocumentVectorCoverage, error) { + l.coverageCalls++ + return l.coverage, nil +} + +func (l *fakeDocumentVectorReconcileLedger) ActivateDocumentVectorGeneration(_ context.Context, _ int64, now time.Time) error { + l.activateCalls++ + l.activateAt = now + if l.activateErr == nil { + l.generation.State = store.DocumentVectorGenerationActive + l.status.State = store.DocumentVectorGenerationActive + } + return l.activateErr +} + +func (l *fakeDocumentVectorReconcileLedger) PurgeRetiredDocumentVectorGeneration(context.Context, int64) (bool, error) { + l.purgeCalls++ + return l.purgeResult, l.purgeErr +} diff --git a/internal/vector/document/search.go b/internal/vector/document/search.go new file mode 100644 index 000000000..c38918c94 --- /dev/null +++ b/internal/vector/document/search.go @@ -0,0 +1,512 @@ +package document + +import ( + "context" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "slices" + "strconv" + "strings" + "unicode/utf8" + + docembedding "go.kenn.io/docbank/document/embedding" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/vector" +) + +const ( + searchCursorVersion = 1 + maxSearchPageSize = 100 + maxSearchQueryBytes = 1024 + maxSearchQueryTerms = 20 +) + +var ErrSemanticSearchUnavailable = errors.New("semantic document search is unavailable") + +// QueryEmbedder embeds one search query using the provider's query role. +type QueryEmbedder interface { + EmbedQuery(ctx context.Context, text string) ([]float32, error) +} + +type SearchLedger interface { + SearchDocuments(ctx context.Context, request store.DocumentSearchRequest) (store.DocumentSearchResponse, error) + GetDocumentIndexRevision(ctx context.Context) (int64, error) + GetActiveDocumentVectorGeneration(ctx context.Context) (*store.DocumentVectorGeneration, error) + GetDocumentVectorTargetProfileID(ctx context.Context) (string, error) + ResolveDocumentVectorSearchOccurrences(ctx context.Context, generationID int64, hits []store.DocumentVectorSearchHit, request store.DocumentSearchRequest, limit int) ([]store.DocumentSearchResult, bool, error) +} + +var _ SearchLedger = (*store.Store)(nil) + +type SearchDeps struct { + Ledger SearchLedger + Embedder QueryEmbedder + Backend Backend + // ExpectedFingerprint binds daemon capability to the currently configured + // and consented policy. Empty preserves the standalone service contract. + ExpectedFingerprint string +} + +type SearchService struct{ deps SearchDeps } + +type searchCursor struct { + Version int `json:"version"` + RequestHash string `json:"request_hash"` + Revision int64 `json:"revision"` + EffectiveMode string `json:"effective_mode"` + GenerationID int64 `json:"generation_id"` + GenerationFingerprint string `json:"generation_fingerprint"` + CandidateLimit int `json:"candidate_limit"` + CandidateDigest string `json:"candidate_digest"` + Offset int `json:"offset"` +} + +func NewSearchService(deps SearchDeps) *SearchService { return &SearchService{deps: deps} } + +func (s *SearchService) Search(ctx context.Context, request store.DocumentSearchRequest) (store.DocumentSearchResponse, error) { + prepared, requestedMode, requestHash, err := normalizeSearchRequest(request) + if err != nil { + return store.DocumentSearchResponse{}, err + } + if s == nil || s.deps.Ledger == nil { + return store.DocumentSearchResponse{}, errors.New("document search ledger is required") + } + var generation *store.DocumentVectorGeneration + if requestedMode == SearchModeSemantic || requestedMode == SearchModeHybrid { + generation, err = s.deps.Ledger.GetActiveDocumentVectorGeneration(ctx) + if err != nil { + return store.DocumentSearchResponse{}, fmt.Errorf("read active document vector generation: %w", err) + } + if generation != nil { + target, targetErr := s.deps.Ledger.GetDocumentVectorTargetProfileID(ctx) + if targetErr != nil || generation.TargetExtractionProfileID != target { + return store.DocumentSearchResponse{}, ErrSemanticSearchUnavailable + } + } + } + effectiveMode, err := s.effectiveMode(requestedMode, generation) + if err != nil { + return store.DocumentSearchResponse{}, err + } + if effectiveMode == SearchModeLexical { + prepared.SearchMode = "" + response, searchErr := s.deps.Ledger.SearchDocuments(ctx, prepared) + if searchErr != nil { + return store.DocumentSearchResponse{}, searchErr + } + response.EffectiveMode = string(SearchModeLexical) + for index := range response.Results { + response.Results[index].LexicalRank = response.Results[index].Rank + } + return response, nil + } + revision, err := s.deps.Ledger.GetDocumentIndexRevision(ctx) + if err != nil { + return store.DocumentSearchResponse{}, fmt.Errorf("read document search revision: %w", err) + } + generationID, generationFingerprint := generation.ID, generation.Fingerprint + offset, expectedDigest, err := validateSearchCursor(prepared.Cursor, requestHash, revision, effectiveMode, generationID, generationFingerprint, prepared.CandidateLimit) + if err != nil { + return store.DocumentSearchResponse{}, err + } + + prepared.Cursor = "" + var candidates []store.DocumentSearchResult + truncated := false + switch effectiveMode { + case SearchModeSemantic, SearchModeHybrid: + var lexical []store.DocumentSearchResult + if effectiveMode == SearchModeHybrid { + var more bool + lexical, more, err = s.collectLexical(ctx, prepared) + if err != nil { + return store.DocumentSearchResponse{}, err + } + truncated = more + } + semantic, more, semanticErr := s.collectSemantic(ctx, prepared, *generation) + if semanticErr != nil { + return store.DocumentSearchResponse{}, semanticErr + } + truncated = truncated || more + var fusionMore bool + candidates, fusionMore, err = fuseSearchResults(lexical, semantic, prepared.CandidateLimit) + if err != nil { + return store.DocumentSearchResponse{}, fmt.Errorf("fuse document candidates: %w", err) + } + truncated = truncated || fusionMore + default: + return store.DocumentSearchResponse{}, fmt.Errorf("%w: unsupported effective mode", store.ErrDocumentSearchInvalidRequest) + } + digest, err := digestSearchCandidates(candidates) + if err != nil { + return store.DocumentSearchResponse{}, err + } + if expectedDigest != "" && expectedDigest != digest { + return store.DocumentSearchResponse{}, store.ErrDocumentSearchCursorStale + } + + response := store.DocumentSearchResponse{ + Revision: revision, EffectiveMode: string(effectiveMode), + VectorGenerationID: generationID, VectorGenerationFingerprint: generationFingerprint, + Truncated: truncated, + } + if offset >= len(candidates) { + return response, nil + } + end := min(offset+prepared.PageSize, len(candidates)) + response.Results = slices.Clone(candidates[offset:end]) + if end < len(candidates) { + response.NextCursor, err = encodeSearchCursor(searchCursor{ + Version: searchCursorVersion, RequestHash: requestHash, Revision: revision, + EffectiveMode: string(effectiveMode), GenerationID: generationID, + GenerationFingerprint: generationFingerprint, CandidateLimit: prepared.CandidateLimit, + CandidateDigest: digest, Offset: end, + }) + if err != nil { + return store.DocumentSearchResponse{}, err + } + } + return response, nil +} + +func (s *SearchService) effectiveMode(requested SearchMode, generation *store.DocumentVectorGeneration) (SearchMode, error) { + semanticReady := generation != nil && s.deps.Embedder != nil && s.deps.Backend != nil && + (s.deps.ExpectedFingerprint == "" || generation.Fingerprint == s.deps.ExpectedFingerprint) + switch requested { + case SearchModeLexical: + return SearchModeLexical, nil + case SearchModeAuto: + // Automatic searches stay local. Sending a query to the configured + // provider requires an explicit semantic or hybrid request. + return SearchModeLexical, nil + case SearchModeSemantic, SearchModeHybrid: + if !semanticReady { + return "", ErrSemanticSearchUnavailable + } + return requested, nil + default: + return "", fmt.Errorf("%w: unsupported search mode", store.ErrDocumentSearchInvalidRequest) + } +} + +func (s *SearchService) collectLexical(ctx context.Context, request store.DocumentSearchRequest) ([]store.DocumentSearchResult, bool, error) { + pageSize := min(request.CandidateLimit, maxSearchPageSize) + lexicalRequest := request + lexicalRequest.PageSize = pageSize + lexicalRequest.Cursor = "" + lexicalRequest.SearchMode = "" + lexicalRequest.CandidateLimit = request.CandidateLimit + results := make([]store.DocumentSearchResult, 0, request.CandidateLimit) + for len(results) < request.CandidateLimit { + response, err := s.deps.Ledger.SearchDocuments(ctx, lexicalRequest) + if err != nil { + return nil, false, err + } + remaining := request.CandidateLimit - len(results) + if len(response.Results) > remaining { + results = append(results, response.Results[:remaining]...) + return results, true, nil + } + results = append(results, response.Results...) + if response.NextCursor == "" { + return results, response.Truncated, nil + } + if len(results) == request.CandidateLimit { + return results, true, nil + } + lexicalRequest.Cursor = response.NextCursor + } + return results, true, nil +} + +func (s *SearchService) collectSemantic(ctx context.Context, request store.DocumentSearchRequest, generation store.DocumentVectorGeneration) ([]store.DocumentSearchResult, bool, error) { + query, err := s.deps.Embedder.EmbedQuery(ctx, request.Query) + if err != nil { + return nil, false, fmt.Errorf("embed document search query: %w", err) + } + if err := validateSearchVector(query, generation.Dimension); err != nil { + return nil, false, err + } + backend, ok := s.deps.Backend.(PagedBackend) + if !ok { + return nil, false, errors.New("document vector backend does not support scoped paging") + } + source := &semanticScopedSource{ + ledger: s.deps.Ledger, backend: backend, generation: generation, + query: query, request: request, seen: make(map[string]bool), + results: make(map[string]store.DocumentSearchResult), + } + collected, err := docembedding.CollectScopedCandidates( + ctx, source, request.CandidateLimit, semanticBackendPageSize, + ) + if err != nil { + return nil, false, fmt.Errorf("collect scoped document vectors: %w", err) + } + results := make([]store.DocumentSearchResult, len(collected.Candidates)) + for index, candidate := range collected.Candidates { + result := source.results[candidate.Key] + result.SemanticRank = candidate.Rank + result.Rank = candidate.Rank + results[index] = result + } + return results, collected.Truncated, nil +} + +const semanticBackendPageSize = 100 + +type semanticScopedSource struct { + ledger SearchLedger + backend PagedBackend + generation store.DocumentVectorGeneration + query []float32 + request store.DocumentSearchRequest + + backendCursor string + backendDone bool + pageSerial int + buffer []store.DocumentSearchResult + seen map[string]bool + results map[string]store.DocumentSearchResult +} + +func (s *semanticScopedSource) SearchPage(ctx context.Context, request docembedding.PageRequest) (docembedding.CandidatePage, error) { + expected := "" + if s.pageSerial > 0 { + expected = strconv.Itoa(s.pageSerial) + } + if request.Cursor != expected { + return docembedding.CandidatePage{}, errors.New("semantic scoped cursor is invalid") + } + for len(s.buffer) < request.Limit && !s.backendDone { + page, err := s.backend.SearchPage( + ctx, GenerationID(s.generation.ID), s.generation.Dimension, + s.query, s.backendCursor, semanticBackendPageSize, + ) + if err != nil { + return docembedding.CandidatePage{}, fmt.Errorf("search document vector backend page: %w", err) + } + storeHits := make([]store.DocumentVectorSearchHit, len(page.Hits)) + for index, hit := range page.Hits { + storeHits[index] = store.DocumentVectorSearchHit{Token: hit.Token, Score: hit.Score, Rank: hit.Rank} + } + resolved, expandedMore, err := s.ledger.ResolveDocumentVectorSearchOccurrences( + ctx, s.generation.ID, storeHits, s.request, docembedding.MaxCandidateLimit+1, + ) + if err != nil { + return docembedding.CandidatePage{}, err + } + for _, result := range resolved { + if s.seen[result.OccurrenceKey] { + continue + } + s.seen[result.OccurrenceKey] = true + s.results[result.OccurrenceKey] = result + s.buffer = append(s.buffer, result) + } + if expandedMore && len(s.buffer) < request.Limit { + return docembedding.CandidatePage{}, errors.New("semantic occurrence expansion exceeded its bounded page") + } + s.backendDone = page.Exhausted + s.backendCursor = page.NextCursor + } + + take := min(request.Limit, len(s.buffer)) + candidates := make([]docembedding.Candidate, take) + for index := range candidates { + result := s.buffer[index] + candidates[index] = docembedding.Candidate{Key: result.OccurrenceKey, Score: result.SemanticScore} + } + s.buffer = s.buffer[take:] + exhausted := s.backendDone && len(s.buffer) == 0 + page := docembedding.CandidatePage{Candidates: candidates, Exhausted: exhausted} + if !exhausted { + s.pageSerial++ + page.NextCursor = strconv.Itoa(s.pageSerial) + } + return page, nil +} + +func validateSearchVector(value []float32, dimension int) error { + if dimension <= 0 || len(value) != dimension { + return fmt.Errorf("%w: query vector length %d does not match dimension %d", vector.ErrDimensionMismatch, len(value), dimension) + } + norm := float64(0) + for _, component := range value { + if math.IsNaN(float64(component)) || math.IsInf(float64(component), 0) { + return fmt.Errorf("%w: query vector contains a nonfinite component", ErrInvalidVector) + } + norm += float64(component) * float64(component) + } + if norm == 0 { + return fmt.Errorf("%w: query vector has zero norm", ErrInvalidVector) + } + return nil +} + +func normalizeSearchRequest(request store.DocumentSearchRequest) (store.DocumentSearchRequest, SearchMode, string, error) { + mode, err := ParseSearchMode(request.SearchMode) + if err != nil { + return request, "", "", fmt.Errorf("%w: %w", store.ErrDocumentSearchInvalidRequest, err) + } + request.Query = strings.ToLower(strings.Join(strings.Fields(request.Query), " ")) + if request.Query == "" || len(request.Query) > maxSearchQueryBytes || !utf8.ValidString(request.Query) || len(strings.Fields(request.Query)) > maxSearchQueryTerms { + return request, "", "", fmt.Errorf("%w: requires a bounded UTF-8 query", store.ErrDocumentSearchInvalidRequest) + } + if request.PageSize == 0 { + request.PageSize = 20 + } + switch mode { + case SearchModeAuto, SearchModeLexical: + if request.CandidateLimit == 0 { + request.CandidateLimit = store.DefaultLexicalDocumentSearchCandidateLimit + } + if request.CandidateLimit < 1 || request.CandidateLimit > store.MaxLexicalDocumentSearchCandidateLimit { + return request, "", "", fmt.Errorf("%w: lexical candidate limit must be between 1 and %d", + store.ErrDocumentSearchInvalidRequest, store.MaxLexicalDocumentSearchCandidateLimit) + } + case SearchModeSemantic, SearchModeHybrid: + options, normalizeErr := docembedding.NormalizeSearchOptions(docembedding.SearchOptions{ + Mode: mode, CandidateLimit: request.CandidateLimit, + }) + if normalizeErr != nil { + return request, "", "", fmt.Errorf("%w: %w", store.ErrDocumentSearchInvalidRequest, normalizeErr) + } + request.CandidateLimit = options.CandidateLimit + } + if request.PageSize < 1 || request.PageSize > maxSearchPageSize || request.AttachmentID < 0 || request.MessageID < 0 { + return request, "", "", fmt.Errorf("%w: request has invalid bounds", store.ErrDocumentSearchInvalidRequest) + } + request.SourceIDs = slices.Clone(request.SourceIDs) + slices.Sort(request.SourceIDs) + request.SourceIDs = slices.Compact(request.SourceIDs) + for _, id := range request.SourceIDs { + if id <= 0 { + return request, "", "", fmt.Errorf("%w: source IDs must be positive", store.ErrDocumentSearchInvalidRequest) + } + } + request.MessageTypes = slices.Clone(request.MessageTypes) + for index := range request.MessageTypes { + request.MessageTypes[index] = strings.ToLower(strings.TrimSpace(request.MessageTypes[index])) + if request.MessageTypes[index] == "" { + return request, "", "", fmt.Errorf("%w: message types must be nonempty", store.ErrDocumentSearchInvalidRequest) + } + } + slices.Sort(request.MessageTypes) + request.MessageTypes = slices.Compact(request.MessageTypes) + request.SearchMode = string(mode) + hash, err := hashSearchRequest(request) + return request, mode, hash, err +} + +func hashSearchRequest(request store.DocumentSearchRequest) (string, error) { + request.Cursor = "" + encoded, err := json.Marshal(request) + if err != nil { + return "", fmt.Errorf("encode document semantic search request: %w", err) + } + digest := sha256.Sum256(encoded) + return hex.EncodeToString(digest[:]), nil +} + +func digestSearchCandidates(results []store.DocumentSearchResult) (string, error) { + type candidateIdentity struct { + OccurrenceKey string `json:"occurrence_key"` + ChunkKey string `json:"chunk_key"` + ExtractionID string `json:"extraction_id"` + VectorToken string `json:"vector_token"` + VectorGenerationFingerprint string `json:"vector_generation_fingerprint"` + VectorEmbeddingProfile string `json:"vector_embedding_profile"` + VectorModel string `json:"vector_model"` + AttachmentID int64 `json:"attachment_id"` + MessageID int64 `json:"message_id"` + VectorGenerationID int64 `json:"vector_generation_id"` + LexicalRank int `json:"lexical_rank"` + SemanticRank int `json:"semantic_rank"` + VectorDimension int `json:"vector_dimension"` + SemanticScore float64 `json:"semantic_score"` + FusionScore float64 `json:"fusion_score"` + MatchedSignals []string `json:"matched_signals"` + } + identities := make([]candidateIdentity, len(results)) + for index, result := range results { + identities[index] = candidateIdentity{ + OccurrenceKey: result.OccurrenceKey, ChunkKey: result.ChunkKey, ExtractionID: result.ExtractionID, + VectorToken: result.VectorToken, VectorGenerationFingerprint: result.VectorGenerationFingerprint, + VectorEmbeddingProfile: result.VectorEmbeddingProfile, VectorModel: result.VectorModel, + AttachmentID: result.AttachmentID, MessageID: result.MessageID, VectorGenerationID: result.VectorGenerationID, + LexicalRank: result.LexicalRank, SemanticRank: result.SemanticRank, VectorDimension: result.VectorDimension, + SemanticScore: result.SemanticScore, FusionScore: result.FusionScore, MatchedSignals: result.MatchedSignals, + } + } + encoded, err := json.Marshal(identities) + if err != nil { + return "", fmt.Errorf("encode document search candidate digest: %w", err) + } + digest := sha256.Sum256(encoded) + return hex.EncodeToString(digest[:]), nil +} + +func validateSearchCursor(value, requestHash string, revision int64, mode SearchMode, generationID int64, generationFingerprint string, candidateLimit int) (int, string, error) { + if value == "" { + return 0, "", nil + } + cursor, err := decodeSearchCursor(value) + if err != nil { + return 0, "", err + } + if cursor.RequestHash != requestHash || cursor.CandidateLimit != candidateLimit { + return 0, "", store.ErrDocumentSearchInvalidCursor + } + if cursor.Offset > cursor.CandidateLimit { + return 0, "", store.ErrDocumentSearchInvalidCursor + } + if cursor.Revision != revision || cursor.EffectiveMode != string(mode) || cursor.GenerationID != generationID || cursor.GenerationFingerprint != generationFingerprint { + return 0, "", store.ErrDocumentSearchCursorStale + } + return cursor.Offset, cursor.CandidateDigest, nil +} + +func encodeSearchCursor(cursor searchCursor) (string, error) { + encoded, err := json.Marshal(cursor) + if err != nil { + return "", fmt.Errorf("encode document semantic search cursor: %w", err) + } + return base64.RawURLEncoding.EncodeToString(encoded), nil +} + +func decodeSearchCursor(value string) (searchCursor, error) { + decoded, err := base64.RawURLEncoding.DecodeString(value) + if err != nil || len(decoded) > 4096 { + return searchCursor{}, store.ErrDocumentSearchInvalidCursor + } + var cursor searchCursor + decoder := json.NewDecoder(strings.NewReader(string(decoded))) + decoder.DisallowUnknownFields() + if err := decoder.Decode(&cursor); err != nil || cursor.Version != searchCursorVersion || !validSearchDigest(cursor.RequestHash) || !validSearchDigest(cursor.CandidateDigest) || cursor.Revision < 0 || cursor.Offset < 1 || cursor.Offset > store.MaxLexicalDocumentSearchCandidateLimit || cursor.CandidateLimit < 1 || cursor.CandidateLimit > store.MaxLexicalDocumentSearchCandidateLimit || (cursor.GenerationID == 0) != (cursor.GenerationFingerprint == "") || (cursor.GenerationFingerprint != "" && !validSearchDigest(cursor.GenerationFingerprint)) { + return searchCursor{}, store.ErrDocumentSearchInvalidCursor + } + mode, err := ParseSearchMode(cursor.EffectiveMode) + if err != nil || mode == SearchModeAuto || + (mode != SearchModeLexical && cursor.CandidateLimit > store.MaxDocumentSearchCandidateLimit) { + return searchCursor{}, store.ErrDocumentSearchInvalidCursor + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return searchCursor{}, store.ErrDocumentSearchInvalidCursor + } + return cursor, nil +} + +func validSearchDigest(value string) bool { + if len(value) != sha256.Size*2 { + return false + } + _, err := hex.DecodeString(value) + return err == nil && value == strings.ToLower(value) +} diff --git a/internal/vector/document/search_test.go b/internal/vector/document/search_test.go new file mode 100644 index 000000000..3eb5f6964 --- /dev/null +++ b/internal/vector/document/search_test.go @@ -0,0 +1,537 @@ +package document + +import ( + "context" + "errors" + "fmt" + "math" + "strconv" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + docbankdocument "go.kenn.io/docbank/document" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/testutil/storetest" + "go.kenn.io/msgvault/internal/vector" +) + +type countingSearchLedger struct { + *store.Store + + searchCalls int +} + +func (ledger *countingSearchLedger) SearchDocuments( + ctx context.Context, + request store.DocumentSearchRequest, +) (store.DocumentSearchResponse, error) { + ledger.searchCalls++ + return ledger.Store.SearchDocuments(ctx, request) +} + +func TestSearchServiceAppliesCandidateBoundsByMode(t *testing.T) { + fixture := storetest.New(t) + service := NewSearchService(SearchDeps{Ledger: fixture.Store}) + + for _, mode := range []SearchMode{SearchModeAuto, SearchModeLexical} { + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "synthetic query", SearchMode: string(mode), CandidateLimit: 1001, + }) + require.NoError(t, err) + } + for _, mode := range []SearchMode{SearchModeSemantic, SearchModeHybrid} { + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "synthetic query", SearchMode: string(mode), CandidateLimit: 1001, + }) + require.ErrorIs(t, err, store.ErrDocumentSearchInvalidRequest) + } +} + +func TestSearchServiceDefaultsCandidateLimitToSharedContract(t *testing.T) { + fixture := seedSemanticSearch(t, "default candidate evidence") + embedder := &searchQueryEmbedder{vector: []float32{1, 0, 0}} + backend := &searchBackend{hits: []Hit{{Token: fixture.claims[0].Token, Score: .9, Rank: 1}}} + service := NewSearchService(SearchDeps{ + Ledger: fixture.store.Store, Embedder: embedder, Backend: backend, + }) + + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "candidate", SearchMode: string(SearchModeSemantic), + }) + require.NoError(t, err) + require.Len(t, backend.searches, 1) + assert.Equal(t, store.DefaultDocumentSearchCandidateLimit, backend.searches[0].k) +} + +func TestSearchServiceAutoFallsBackToLexicalWhenSemanticCapabilityIsAbsent(t *testing.T) { + fixture := storetest.New(t) + service := NewSearchService(SearchDeps{Ledger: fixture.Store}) + + response, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "synthetic query", SearchMode: string(SearchModeAuto), + }) + + require.NoError(t, err) + assert.Equal(t, string(SearchModeLexical), response.EffectiveMode) + assert.Empty(t, response.Results) +} + +func TestSearchServiceExplicitSemanticDoesNotMasqueradeAsLexical(t *testing.T) { + fixture := storetest.New(t) + service := NewSearchService(SearchDeps{Ledger: fixture.Store}) + + for _, mode := range []SearchMode{SearchModeSemantic, SearchModeHybrid} { + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "synthetic query", SearchMode: string(mode), + }) + require.ErrorIs(t, err, ErrSemanticSearchUnavailable) + assert.NotErrorIs(t, err, store.ErrDocumentSearchUnavailable) + } +} + +func TestSearchServiceSemanticReturnsAuthoritativeOccurrenceProvenance(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + fixture := seedSemanticSearch(t, "alpha evidence", "beta evidence") + embedder := &searchQueryEmbedder{vector: []float32{1, 0, 0}} + backend := &searchBackend{hits: []Hit{{Token: fixture.claims[1].Token, Score: .75, Rank: 1}}} + service := NewSearchService(SearchDeps{Ledger: fixture.store.Store, Embedder: embedder, Backend: backend}) + + response, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: " FIND beta ", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + }) + + requirements.NoError(err) + requirements.Len(response.Results, 1) + result := response.Results[0] + assertions.Equal(string(SearchModeSemantic), response.EffectiveMode) + assertions.Equal(fixture.generation.ID, response.VectorGenerationID) + assertions.Equal(fixture.generation.Fingerprint, response.VectorGenerationFingerprint) + assertions.Equal(fixture.claims[1].Token, result.VectorToken) + assertions.Equal(fixture.claims[1].ChunkKey, result.ChunkKey) + assertions.Equal(fixture.claims[1].ChunkOrdinal, result.ChunkOrdinal) + assertions.Equal(fixture.claims[1].ExtractionID, result.ExtractionID) + assertions.Equal(fixture.claims[1].ExtractionProfileID, result.ProfileID) + assertions.Equal(fixture.generation.EmbeddingProfile, result.VectorEmbeddingProfile) + assertions.Equal(fixture.generation.Model, result.VectorModel) + assertions.Equal(fixture.generation.Dimension, result.VectorDimension) + assertions.Equal(1, result.SemanticRank) + assertions.Equal(1, result.Rank) + assertions.Equal([]string{"semantic"}, result.MatchedSignals) + assertions.Equal([]string{"find beta"}, embedder.queries) + requirements.Len(backend.searches, 1) + assertions.Equal(fixture.generation.ID, int64(backend.searches[0].generationID)) + assertions.Equal(3, backend.searches[0].dimension) + assertions.Equal([]float32{1, 0, 0}, backend.searches[0].query) + assertions.Equal(semanticBackendPageSize, backend.searches[0].k) +} + +func TestSearchServiceSemanticPagesPastOutOfScopeGlobalHits(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + fixture := seedSemanticSearch(t, "scoped evidence") + hits := make([]Hit, semanticBackendPageSize+1) + for index := range semanticBackendPageSize { + hits[index] = Hit{Token: fmt.Sprintf("%064x", index+1), Score: 1 - float64(index)/1000} + } + hits[semanticBackendPageSize] = Hit{Token: fixture.claims[0].Token, Score: .5} + backend := &searchBackend{hits: hits} + service := NewSearchService(SearchDeps{ + Ledger: fixture.store.Store, Embedder: &searchQueryEmbedder{vector: []float32{1, 0, 0}}, + Backend: backend, ExpectedFingerprint: fixture.generation.Fingerprint, + }) + + response, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "scoped evidence", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + }) + + require.NoError(err) + require.Len(response.Results, 1) + assert.Equal(fixture.claims[0].Token, response.Results[0].VectorToken) + assert.Len(backend.searches, 2) +} + +func TestSearchServiceAutoKeepsQueryLocalWhenSemanticCapabilityIsReady(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + fixture := seedSemanticSearch(t, "nebula evidence") + embedder := &searchQueryEmbedder{vector: []float32{1, 0, 0}} + backend := &searchBackend{hits: []Hit{{Token: fixture.claims[0].Token, Score: .9, Rank: 1}}} + service := NewSearchService(SearchDeps{Ledger: fixture.store.Store, Embedder: embedder, Backend: backend}) + + for _, mode := range []string{"", string(SearchModeAuto)} { + response, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "nebula", SearchMode: mode, CandidateLimit: 10, + }) + + requirements.NoError(err) + requirements.Len(response.Results, 1) + assertions.Equal(string(SearchModeLexical), response.EffectiveMode) + assertions.Equal([]string{"content"}, response.Results[0].MatchedSignals) + } + assertions.Empty(embedder.queries) + assertions.Empty(backend.searches) +} + +func TestSearchServiceExplicitLexicalPreservesStoreRankingWithoutProviderWork(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + fixture := seedSemanticSearch(t, "nebula evidence") + copyMessageID := fixture.store.CreateMessage("lexical-ranking-copy") + copyAttachmentID := addSemanticSearchAttachment( + t, fixture.store, copyMessageID, fixture.claims[0].CanonicalBlobHash, "nebula-copy.pdf", "provider:lexical-copy", + ) + _, eligible, err := fixture.store.Store.ReconcileDocumentOccurrence(t.Context(), copyAttachmentID, 2) + requirements.NoError(err) + requirements.True(eligible) + direct, err := fixture.store.Store.SearchDocuments(t.Context(), store.DocumentSearchRequest{Query: "nebula"}) + requirements.NoError(err) + embedder := &searchQueryEmbedder{vector: []float32{1, 0, 0}} + backend := &searchBackend{hits: []Hit{{Token: fixture.claims[0].Token, Score: .9, Rank: 1}}} + service := NewSearchService(SearchDeps{Ledger: fixture.store.Store, Embedder: embedder, Backend: backend}) + + response, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "nebula", SearchMode: string(SearchModeLexical), CandidateLimit: 10, + }) + + requirements.NoError(err) + requirements.Len(response.Results, len(direct.Results)) + for index := range direct.Results { + assertions.Equal(direct.Results[index].OccurrenceKey, response.Results[index].OccurrenceKey) + assertions.Equal(direct.Results[index].MatchedSignals, response.Results[index].MatchedSignals) + assertions.Equal(direct.Results[index].Rank, response.Results[index].Rank) + } + assertions.Empty(embedder.queries) + assertions.Empty(backend.searches) +} + +func TestSearchServiceLexicalUsesStoreCursorWithoutExhaustingCandidates(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + fixture := seedSemanticSearch(t, "paged lexical evidence") + for index := range 100 { + messageID := fixture.store.CreateMessage(fmt.Sprintf("lexical-page-%03d", index)) + attachmentID := addSemanticSearchAttachment( + t, fixture.store, messageID, fixture.claims[0].CanonicalBlobHash, + fmt.Sprintf("paged-%03d.pdf", index), fmt.Sprintf("provider:paged-%03d", index), + ) + _, eligible, err := fixture.store.Store.ReconcileDocumentOccurrence(t.Context(), attachmentID, int64(index+2)) + requirements.NoError(err) + requirements.True(eligible) + } + ledger := &countingSearchLedger{Store: fixture.store.Store} + service := NewSearchService(SearchDeps{Ledger: ledger}) + request := store.DocumentSearchRequest{ + Query: "paged", SearchMode: string(SearchModeLexical), CandidateLimit: 101, PageSize: 1, + } + + first, err := service.Search(t.Context(), request) + requirements.NoError(err) + requirements.Len(first.Results, 1) + requirements.NotEmpty(first.NextCursor) + assertions.Equal(1, ledger.searchCalls) + assertions.Equal(1, first.Results[0].Rank) + assertions.Equal(1, first.Results[0].LexicalRank) + + request.Cursor = first.NextCursor + second, err := service.Search(t.Context(), request) + requirements.NoError(err) + requirements.Len(second.Results, 1) + assertions.Equal(2, ledger.searchCalls) + assertions.Equal(2, second.Results[0].Rank) + assertions.Equal(2, second.Results[0].LexicalRank) + assertions.NotEqual(first.Results[0].OccurrenceKey, second.Results[0].OccurrenceKey) +} + +func TestSearchServiceExplicitSemanticModesDoNotFallbackAfterProviderAttempt(t *testing.T) { + fixture := seedSemanticSearch(t, "nebula evidence") + providerErr := errors.New("synthetic provider unavailable") + service := NewSearchService(SearchDeps{ + Ledger: fixture.store.Store, Embedder: &searchQueryEmbedder{err: providerErr}, Backend: &searchBackend{}, + }) + + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "nebula", SearchMode: string(SearchModeHybrid), CandidateLimit: 10, + }) + require.ErrorIs(t, err, providerErr) + + backendErr := errors.New("synthetic backend unavailable") + service = NewSearchService(SearchDeps{ + Ledger: fixture.store.Store, Embedder: &searchQueryEmbedder{vector: []float32{1, 0, 0}}, + Backend: &searchBackend{err: backendErr}, + }) + _, err = service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "nebula", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + }) + require.ErrorIs(t, err, backendErr) +} + +func TestSearchServiceRejectsInvalidQueryVectorBeforeBackendWork(t *testing.T) { + fixture := seedSemanticSearch(t, "vector evidence") + tests := []struct { + name string + vector []float32 + wantErr error + }{ + {name: "dimension", vector: []float32{1, 0}, wantErr: vector.ErrDimensionMismatch}, + {name: "zero norm", vector: []float32{0, 0, 0}, wantErr: ErrInvalidVector}, + {name: "nonfinite", vector: []float32{1, float32(math.NaN()), 0}, wantErr: ErrInvalidVector}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + backend := &searchBackend{} + service := NewSearchService(SearchDeps{ + Ledger: fixture.store.Store, Embedder: &searchQueryEmbedder{vector: test.vector}, Backend: backend, + }) + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "vector", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + }) + require.ErrorIs(t, err, test.wantErr) + assert.Empty(t, backend.searches) + }) + } +} + +func TestSearchServiceCursorBindsFixedCandidatesAndChecksRevisionBeforeEmbedding(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + fixture := seedSemanticSearch(t, "cursor evidence") + copyMessageID := fixture.store.CreateMessage("semantic-cursor-copy") + copyAttachmentID := addSemanticSearchAttachment( + t, fixture.store, copyMessageID, fixture.claims[0].CanonicalBlobHash, "cursor-copy.pdf", "provider:cursor-copy", + ) + _, eligible, err := fixture.store.Store.ReconcileDocumentOccurrence(t.Context(), copyAttachmentID, 2) + requirements.NoError(err) + requirements.True(eligible) + embedder := &searchQueryEmbedder{vector: []float32{1, 0, 0}} + backend := &searchBackend{hits: []Hit{{Token: fixture.claims[0].Token, Score: .9, Rank: 1}}} + service := NewSearchService(SearchDeps{Ledger: fixture.store.Store, Embedder: embedder, Backend: backend}) + request := store.DocumentSearchRequest{ + Query: "cursor", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, PageSize: 1, + } + + first, err := service.Search(t.Context(), request) + requirements.NoError(err) + requirements.Len(first.Results, 1) + requirements.NotEmpty(first.NextCursor) + providerCalls := len(embedder.queries) + _, err = service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "different", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + PageSize: 1, Cursor: first.NextCursor, + }) + requirements.ErrorIs(err, store.ErrDocumentSearchInvalidCursor) + assertions.Len(embedder.queries, providerCalls, "mismatched request must fail before provider work") + + request.Cursor = first.NextCursor + second, err := service.Search(t.Context(), request) + requirements.NoError(err) + requirements.Len(second.Results, 1) + assertions.NotEqual(first.Results[0].OccurrenceKey, second.Results[0].OccurrenceKey) + assertions.Equal(2, second.Results[0].Rank) + assertions.Empty(second.NextCursor) + + request.Cursor = first.NextCursor + backend.hits[0].Score = .7 + _, err = service.Search(t.Context(), request) + requirements.ErrorIs(err, store.ErrDocumentSearchCursorStale) + + backend.hits[0].Score = .9 + fresh, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "cursor", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, PageSize: 1, + }) + requirements.NoError(err) + providerCalls = len(embedder.queries) + thirdMessageID := fixture.store.CreateMessage("semantic-cursor-revision") + thirdAttachmentID := addSemanticSearchAttachment( + t, fixture.store, thirdMessageID, fixture.claims[0].CanonicalBlobHash, "third.pdf", "provider:third", + ) + _, eligible, err = fixture.store.Store.ReconcileDocumentOccurrence(t.Context(), thirdAttachmentID, 3) + requirements.NoError(err) + requirements.True(eligible) + _, err = service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "cursor", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + PageSize: 1, Cursor: fresh.NextCursor, + }) + requirements.ErrorIs(err, store.ErrDocumentSearchCursorStale) + assertions.Len(embedder.queries, providerCalls, "stale revision must fail before provider work") +} + +func TestSearchCursorRejectsOffsetBeyondCandidateSet(t *testing.T) { + cursor, err := encodeSearchCursor(searchCursor{ + Version: searchCursorVersion, RequestHash: strings.Repeat("a", 64), Revision: 1, + EffectiveMode: string(SearchModeLexical), CandidateLimit: 2, + CandidateDigest: strings.Repeat("b", 64), Offset: 3, + }) + require.NoError(t, err) + + _, _, err = validateSearchCursor(cursor, strings.Repeat("a", 64), 1, SearchModeLexical, 0, "", 2) + require.ErrorIs(t, err, store.ErrDocumentSearchInvalidCursor) +} + +type semanticSearchFixture struct { + store *storetest.Fixture + generation store.DocumentVectorGeneration + claims []store.DocumentVectorChunkClaim +} + +func seedSemanticSearch(t *testing.T, chunks ...string) semanticSearchFixture { + t.Helper() + requirements := require.New(t) + normalizationPolicy, err := docbankdocument.NewNormalizePolicy(1) + requirements.NoError(err) + fixture := storetest.New(t) + profileFingerprint := strings.Repeat("a", 64) + profile := store.DocumentExtractionProfile{ + ID: "profile-" + profileFingerprint, Fingerprint: profileFingerprint, + Provider: "mistral", Endpoint: "https://api.example.invalid/v1/ocr", Region: "test", + Model: "ocr-test", RetentionPosture: "standard", TrainingPosture: "opted-out", + AllowedMediaTypes: []string{"application/pdf"}, PolicyJSON: []byte(`{"policy":1}`), + } + _, err = fixture.Store.EnsureDocumentExtractionProfile(t.Context(), profile) + requirements.NoError(err) + requirements.NoError(fixture.Store.RecordDocumentProviderConsent(t.Context(), store.DocumentProviderConsent{ + ProfileID: profile.ID, ProfileFingerprint: profile.Fingerprint, + RetentionPosture: profile.RetentionPosture, TrainingPosture: profile.TrainingPosture, + })) + messageID := fixture.CreateMessage("semantic-search-document") + hash := strings.Repeat("b", 64) + attachmentID := addSemanticSearchAttachment(t, fixture, messageID, hash, "semantic.pdf", "mime:1") + occurrence, eligible, err := fixture.Store.ReconcileDocumentOccurrence(t.Context(), attachmentID, 1) + requirements.NoError(err) + requirements.True(eligible) + claim, err := fixture.Store.ClaimDocumentExtraction(t.Context(), store.DocumentExtractionClaimInput{ + ExtractionID: "semantic-extraction", ProfileID: profile.ID, CanonicalBlobHash: hash, + ExtractionInputKey: "original", OccurrenceAttachmentID: occurrence.AttachmentID, + OccurrenceMIMEType: occurrence.MIMEType, OccurrenceMessageType: "email", + LeaseOwner: "semantic-extractor", LeaseUntil: time.Now().UTC().Add(time.Minute), + LocalBytes: 128, SourceSequence: occurrence.SourceSequence, + }) + requirements.NoError(err) + unitText := strings.Join(chunks, " ") + publication := store.DocumentExtractionPublication{ + ExtractionID: claim.ExtractionID, ProfileID: claim.ProfileID, + CanonicalBlobHash: claim.CanonicalBlobHash, ExtractionInputKey: claim.ExtractionInputKey, + OccurrenceAttachmentID: claim.OccurrenceAttachmentID, OccurrenceMIMEType: claim.OccurrenceMIMEType, + OccurrenceMessageType: claim.OccurrenceMessageType, LeaseOwner: claim.LeaseOwner, LeaseFence: claim.LeaseFence, + ReturnedModel: profile.Model, UnitsProcessed: 1, RequestCount: 1, + ManifestChecksum: strings.Repeat("c", 64), NormalizationVersion: normalizationPolicy.Identity().Version, + DocumentFamily: "pdf", UnitKind: "page", + Units: []store.DocumentPublishedUnit{{ + Index: 0, Kind: "page", Text: unitText, Checksum: strings.Repeat("d", 64), CharCount: len([]rune(unitText)), + }}, + } + offset := 0 + for index, text := range chunks { + length := len([]rune(text)) + publication.Chunks = append(publication.Chunks, store.DocumentPublishedChunk{ + Key: fmt.Sprintf("chunk-%d", index), Ordinal: index, Text: text, + FirstUnitIndex: 0, LastUnitIndex: 0, Checksum: fmt.Sprintf("%064x", index+1), + CharCount: length, + Spans: []store.DocumentPublishedSpan{{UnitIndex: 0, CharStart: offset, CharEnd: offset + length}}, + }) + offset += length + 1 + } + requirements.NoError(fixture.Store.PublishDocumentExtraction(t.Context(), publication)) + generation, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), store.DocumentVectorGenerationSpec{ + Fingerprint: strings.Repeat("f", 64), TargetExtractionProfileID: profile.ID, + EmbeddingProfile: "vector.embeddings", Model: "embed-test", Dimension: 3, + }) + requirements.NoError(err) + now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) + claims := make([]store.DocumentVectorChunkClaim, 0, len(chunks)) + for range chunks { + vectorClaim, claimErr := fixture.Store.ClaimDocumentVectorChunk( + t.Context(), generation.ID, 0, 1000, "semantic-worker", now, time.Minute, + ) + requirements.NoError(claimErr) + requirements.NotNil(vectorClaim) + requirements.NoError(fixture.Store.CommitDocumentVectorPublication( + t.Context(), generation.ID, vectorClaim.Token, vectorClaim.LeaseOwner, + vectorClaim.LeaseFence, now.Add(time.Second), + )) + claims = append(claims, *vectorClaim) + } + requirements.NoError(fixture.Store.ActivateDocumentVectorGeneration(t.Context(), generation.ID, now.Add(2*time.Second))) + return semanticSearchFixture{store: fixture, generation: generation, claims: claims} +} + +func addSemanticSearchAttachment( + t *testing.T, fixture *storetest.Fixture, messageID int64, hash, filename, sourcePartKey string, +) int64 { + t.Helper() + require.NoError(t, fixture.Store.UpsertAttachmentRecord(t.Context(), messageID, store.AttachmentWrite{ + Filename: filename, MIMEType: "application/pdf", Size: 128, + StoragePath: hash[:2] + "/" + hash, ContentHash: hash, + Role: store.AttachmentRoleStandalone, RoleSource: store.AttachmentRoleSourceImporterSemantics, + SourcePartKey: sourcePartKey, + })) + var attachmentID int64 + require.NoError(t, fixture.Store.DB().QueryRow(fixture.Store.Rebind( + `SELECT id FROM attachments WHERE message_id = ?`), messageID).Scan(&attachmentID)) + return attachmentID +} + +type searchQueryEmbedder struct { + vector []float32 + err error + queries []string +} + +func (e *searchQueryEmbedder) EmbedQuery(_ context.Context, query string) ([]float32, error) { + e.queries = append(e.queries, query) + return e.vector, e.err +} + +type searchBackendCall struct { + generationID GenerationID + dimension int + query []float32 + k int +} + +type searchBackend struct { + hits []Hit + err error + searches []searchBackendCall +} + +func (*searchBackend) PutUnpublished(context.Context, GenerationID, int, []Embedding) error { + return nil +} +func (*searchBackend) DeleteTokens(context.Context, GenerationID, []string) error { return nil } +func (b *searchBackend) Search( + _ context.Context, generationID GenerationID, dimension int, query []float32, k int, +) ([]Hit, error) { + b.searches = append(b.searches, searchBackendCall{ + generationID: generationID, dimension: dimension, query: append([]float32(nil), query...), k: k, + }) + return append([]Hit(nil), b.hits...), b.err +} + +func (b *searchBackend) SearchPage( + _ context.Context, generationID GenerationID, dimension int, query []float32, cursor string, k int, +) (HitPage, error) { + offset := 0 + if cursor != "" { + var err error + offset, err = strconv.Atoi(cursor) + if err != nil { + return HitPage{}, fmt.Errorf("parse search backend cursor: %w", err) + } + } + b.searches = append(b.searches, searchBackendCall{ + generationID: generationID, dimension: dimension, query: append([]float32(nil), query...), k: k, + }) + end := min(offset+k, len(b.hits)) + page := HitPage{Hits: append([]Hit(nil), b.hits[offset:end]...), Exhausted: end == len(b.hits)} + for index := range page.Hits { + page.Hits[index].Rank = offset + index + 1 + } + if !page.Exhausted { + page.NextCursor = strconv.Itoa(end) + } + return page, b.err +} diff --git a/internal/vector/document/status.go b/internal/vector/document/status.go new file mode 100644 index 000000000..07a03743e --- /dev/null +++ b/internal/vector/document/status.go @@ -0,0 +1,45 @@ +package document + +import ( + "context" + "errors" + "time" + + "go.kenn.io/msgvault/internal/store" +) + +type FailureDiagnostic = store.DocumentVectorFailureDiagnostic +type GenerationStatus = store.DocumentVectorGenerationStatus +type FailureResetResult = store.DocumentVectorFailureResetResult + +// StatusLedger is the authoritative inspection and manual-recovery seam. +type StatusLedger interface { + GetDocumentVectorGenerationStatus(ctx context.Context, generationID int64, afterToken string, limit int) (store.DocumentVectorGenerationStatus, error) + ResetDocumentVectorFailures(ctx context.Context, generationID int64, afterToken string, limit int, now time.Time) (store.DocumentVectorFailureResetResult, error) +} + +var _ StatusLedger = (*store.Store)(nil) + +// StatusService exposes bounded, non-PII generation diagnostics and explicit +// failed-publication recovery without exposing store implementation details. +type StatusService struct{ ledger StatusLedger } + +func NewStatusService(ledger StatusLedger) *StatusService { return &StatusService{ledger: ledger} } + +func (s *StatusService) Inspect( + ctx context.Context, generationID GenerationID, afterToken string, limit int, +) (GenerationStatus, error) { + if s == nil || s.ledger == nil { + return GenerationStatus{}, errors.New("document vector status ledger is required") + } + return s.ledger.GetDocumentVectorGenerationStatus(ctx, int64(generationID), afterToken, limit) +} + +func (s *StatusService) ResetFailures( + ctx context.Context, generationID GenerationID, afterToken string, limit int, now time.Time, +) (FailureResetResult, error) { + if s == nil || s.ledger == nil { + return FailureResetResult{}, errors.New("document vector status ledger is required") + } + return s.ledger.ResetDocumentVectorFailures(ctx, int64(generationID), afterToken, limit, now) +} diff --git a/internal/vector/document/status_test.go b/internal/vector/document/status_test.go new file mode 100644 index 000000000..8d63920d9 --- /dev/null +++ b/internal/vector/document/status_test.go @@ -0,0 +1,47 @@ +package document + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/msgvault/internal/store" +) + +func TestStatusServiceInspectsAndResetsFailures(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + wantStatus := store.DocumentVectorGenerationStatus{ + GenerationID: 7, State: store.DocumentVectorGenerationBuilding, + Pending: 2, Retryable: 1, Terminal: 3, ReadyLive: 4, Obsolete: 5, CleanupPending: 2, + } + wantReset := store.DocumentVectorFailureResetResult{Scanned: 3, Reset: 2, Exhausted: true} + ledger := &fakeDocumentVectorStatusLedger{status: wantStatus, reset: wantReset} + service := NewStatusService(ledger) + + status, err := service.Inspect(t.Context(), 7, "", 25) + requirements.NoError(err) + assertions.Equal(wantStatus, status) + resetAt := time.Date(2026, time.August, 20, 23, 30, 0, 987654000, time.UTC) + reset, err := service.ResetFailures(t.Context(), 7, "", 25, resetAt) + requirements.NoError(err) + assertions.Equal(wantReset, reset) + assertions.Equal(resetAt, ledger.resetAt) +} + +type fakeDocumentVectorStatusLedger struct { + status store.DocumentVectorGenerationStatus + reset store.DocumentVectorFailureResetResult + resetAt time.Time +} + +func (l *fakeDocumentVectorStatusLedger) GetDocumentVectorGenerationStatus(context.Context, int64, string, int) (store.DocumentVectorGenerationStatus, error) { + return l.status, nil +} + +func (l *fakeDocumentVectorStatusLedger) ResetDocumentVectorFailures(_ context.Context, _ int64, _ string, _ int, now time.Time) (store.DocumentVectorFailureResetResult, error) { + l.resetAt = now + return l.reset, nil +} diff --git a/internal/vector/document/worker.go b/internal/vector/document/worker.go new file mode 100644 index 000000000..eca628e67 --- /dev/null +++ b/internal/vector/document/worker.go @@ -0,0 +1,650 @@ +package document + +import ( + "context" + "errors" + "fmt" + "math" + "strings" + "sync" + "time" + "unicode/utf8" + + docbankdocument "go.kenn.io/docbank/document" + docembedding "go.kenn.io/docbank/document/embedding" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/vector" +) + +const ( + maxWorkerRunLimit = 1000 + maxWorkerRetryDelay = 7 * 24 * time.Hour +) + +var ( + errInvalidProviderShape = vector.ErrInvalidProviderShape + errInvalidProviderVector = vector.ErrInvalidProviderVector + errInputPreparation = errors.New("document embedding input preparation failed") + errBackendPut = errors.New("document vector backend put failed") +) + +// Ledger is the authoritative publication surface used by Worker. +// *store.Store satisfies it. +type Ledger interface { + GetDocumentVectorGeneration(ctx context.Context, id int64) (store.DocumentVectorGeneration, error) + ListDocumentVectorChunkCandidates(ctx context.Context, generationID, afterChunkID int64, limit int) ([]store.DocumentVectorChunkCandidate, error) + ClaimDocumentVectorChunk(ctx context.Context, generationID, afterChunkID int64, scanLimit int, owner string, now time.Time, leaseDuration time.Duration) (*store.DocumentVectorChunkClaim, error) + RenewDocumentVectorChunkClaim(ctx context.Context, generationID int64, token, owner string, fence int64, now time.Time, leaseDuration time.Duration) (time.Time, error) + CommitDocumentVectorPublication(ctx context.Context, generationID int64, token, owner string, fence int64, now time.Time) error + FailDocumentVectorChunk(ctx context.Context, generationID int64, token, owner string, fence int64, now time.Time, nextRetryAt *time.Time, terminal bool, errorCode string) error +} + +var _ Ledger = (*store.Store)(nil) + +// WorkerDeps are the bounded collaborators and policy values for Worker. +type WorkerDeps struct { + Ledger Ledger + Provider Provider + Backend Backend + + Owner string + Dimension int + MaxInputChars int + LeaseDuration time.Duration + HeartbeatInterval time.Duration + RetryDelay time.Duration + MaxAttempts int + Recipe docembedding.Recipe + // ContextualDocuments keeps each durable extraction chunk in its own stable + // provider document. A worker batch must never redefine contextual scope. + ContextualDocuments bool + // AfterGenerationID and AfterChunkID restore the bounded scan cursor reported + // by a prior run of the same generation. Task 6b may persist this pair; Worker + // also carries it across sequential runs and resets it when generations change. + AfterGenerationID GenerationID + AfterChunkID int64 + Now func() time.Time + prepareInputs func(context.Context, Ledger, docembedding.Recipe, []*store.DocumentVectorChunkClaim) (map[string]string, error) +} + +// RunResult reports only locally observable accounting. Provider token usage +// is deliberately absent because EmbedDocuments does not report it. +type RunResult struct { + Claimed int `json:"claimed"` + Embedded int `json:"embedded"` + Published int `json:"published"` + Retry int `json:"retry"` + Terminal int `json:"terminal"` + SourceChanged int `json:"source_changed"` + + ProviderCalls int `json:"provider_calls"` + ProviderDocuments int `json:"provider_documents"` + ProviderChunks int `json:"provider_chunks"` + ProviderInputChars int `json:"provider_input_chars"` + + AfterGenerationID GenerationID `json:"after_generation_id,omitempty"` + AfterChunkID int64 `json:"after_chunk_id,omitempty"` + Exhausted bool `json:"exhausted"` +} + +// Worker publishes one bounded page of a building document-vector generation. +// A Worker is safe for sequential use. +type Worker struct { + deps WorkerDeps + cursorGenerationID GenerationID + afterChunkID int64 +} + +// workerClaimHeartbeat owns every live claim from provider dispatch through +// backend publication. releaseForTransition serializes the final renewal and +// ownership removal against periodic renewal before Commit or Fail begins. +type workerClaimHeartbeat struct { + worker *Worker + generationID GenerationID + ctx context.Context + cancel context.CancelFunc + stopCh chan struct{} + doneCh chan struct{} + stopOnce sync.Once + + mu sync.Mutex + active map[string]*store.DocumentVectorChunkClaim + runErr error +} + +func newWorkerClaimHeartbeat( + ctx context.Context, worker *Worker, generationID GenerationID, claims []*store.DocumentVectorChunkClaim, +) *workerClaimHeartbeat { + workCtx, cancel := context.WithCancel(ctx) + heartbeat := &workerClaimHeartbeat{ + worker: worker, generationID: generationID, ctx: workCtx, cancel: cancel, + stopCh: make(chan struct{}), doneCh: make(chan struct{}), + active: make(map[string]*store.DocumentVectorChunkClaim, len(claims)), + } + for _, claim := range claims { + heartbeat.active[claim.Token] = claim + } + go heartbeat.run(ctx) + return heartbeat +} + +func (h *workerClaimHeartbeat) context() context.Context { return h.ctx } + +func (h *workerClaimHeartbeat) err() error { + h.mu.Lock() + defer h.mu.Unlock() + return h.runErr +} + +func (h *workerClaimHeartbeat) releaseForTransition(claim *store.DocumentVectorChunkClaim) error { + h.mu.Lock() + defer h.mu.Unlock() + if h.runErr != nil { + return h.runErr + } + if err := h.ctx.Err(); err != nil { + return err + } + if _, ok := h.active[claim.Token]; !ok { + return fmt.Errorf("document vector claim %q is no longer heartbeat-owned", claim.Token) + } + if _, err := h.worker.deps.Ledger.RenewDocumentVectorChunkClaim( + h.ctx, int64(h.generationID), claim.Token, claim.LeaseOwner, claim.LeaseFence, + h.worker.deps.Now(), h.worker.deps.LeaseDuration, + ); err != nil { + h.failLocked(claim.Token, err) + return h.runErr + } + delete(h.active, claim.Token) + return nil +} + +func (h *workerClaimHeartbeat) run(parent context.Context) { + defer close(h.doneCh) + ticker := time.NewTicker(h.worker.deps.HeartbeatInterval) + defer ticker.Stop() + for { + select { + case <-parent.Done(): + h.mu.Lock() + if h.runErr == nil { + h.runErr = parent.Err() + } + h.cancel() + h.mu.Unlock() + return + case <-h.stopCh: + return + case <-ticker.C: + h.mu.Lock() + for _, claim := range h.active { + _, err := h.worker.deps.Ledger.RenewDocumentVectorChunkClaim( + h.ctx, int64(h.generationID), claim.Token, claim.LeaseOwner, + claim.LeaseFence, h.worker.deps.Now(), h.worker.deps.LeaseDuration, + ) + if err != nil { + h.failLocked(claim.Token, err) + h.mu.Unlock() + return + } + } + h.mu.Unlock() + } + } +} + +func (h *workerClaimHeartbeat) failLocked(token string, err error) { + if h.runErr == nil { + h.runErr = fmt.Errorf("renew document vector claim %q: %w", token, err) + } + h.cancel() +} + +func (h *workerClaimHeartbeat) stop() { + h.stopOnce.Do(func() { close(h.stopCh) }) + <-h.doneCh + h.cancel() +} + +func NewWorker(deps WorkerDeps) *Worker { + if deps.Now == nil { + deps.Now = func() time.Time { return time.Now().UTC() } + } + if deps.Recipe.Fingerprint() == "" && deps.MaxInputChars > 0 { + deps.Recipe, _ = docembedding.NewRecipe(docembedding.RecipeConfig{ + Mode: docembedding.RepresentationRaw, MaxInputRunes: deps.MaxInputChars, + }) + } + if deps.prepareInputs == nil { + deps.prepareInputs = prepareDocbankClaimInputs + } + return &Worker{ + deps: deps, + cursorGenerationID: deps.AfterGenerationID, + afterChunkID: deps.AfterChunkID, + } +} + +func (w *Worker) Run(ctx context.Context, generationID GenerationID, limit int) (RunResult, error) { + var result RunResult + if limit < 1 || limit > maxWorkerRunLimit { + return result, fmt.Errorf("document vector worker limit must be between 1 and %d", maxWorkerRunLimit) + } + if err := w.validate(generationID); err != nil { + return result, err + } + generation, err := w.deps.Ledger.GetDocumentVectorGeneration(ctx, int64(generationID)) + if err != nil { + return result, fmt.Errorf("read document vector generation: %w", err) + } + if generation.ID != int64(generationID) || generation.State != store.DocumentVectorGenerationBuilding { + return result, store.ErrDocumentVectorInvalidGenerationState + } + if generation.Dimension != w.deps.Dimension { + return result, fmt.Errorf("document vector worker dimension %d does not match generation dimension %d", w.deps.Dimension, generation.Dimension) + } + w.bindCursor(generationID) + result.AfterGenerationID = generationID + + claims, err := w.collectClaims(ctx, generationID, limit, &result) + if err != nil || len(claims) == 0 { + return result, err + } + heartbeat := newWorkerClaimHeartbeat(ctx, w, generationID, claims) + defer heartbeat.stop() + preparedTexts := make(map[string]string, len(claims)) + preparedClaims := make([]*store.DocumentVectorChunkClaim, 0, len(claims)) + var runErr error + for _, claimGroup := range groupWorkerClaimsByExtraction(claims) { + groupTexts, prepareErr := w.deps.prepareInputs( + heartbeat.context(), w.deps.Ledger, w.deps.Recipe, claimGroup, + ) + if prepareErr != nil { + cause := fmt.Errorf("%w: extraction %q: %w", errInputPreparation, claimGroup[0].ExtractionID, prepareErr) + failureErr := w.failClaims(ctx, generationID, claimGroup, cause, heartbeat, &result) + runErr = errors.Join(runErr, cause, failureErr) + continue + } + for _, claim := range claimGroup { + text, ok := groupTexts[claim.Token] + if !ok { + cause := fmt.Errorf("%w: extraction %q omitted token %q", errInputPreparation, claim.ExtractionID, claim.Token) + failureErr := w.failClaims(ctx, generationID, []*store.DocumentVectorChunkClaim{claim}, cause, heartbeat, &result) + runErr = errors.Join(runErr, cause, failureErr) + continue + } + preparedTexts[claim.Token] = text + preparedClaims = append(preparedClaims, claim) + } + } + if len(preparedClaims) == 0 { + return result, errors.Join(runErr, heartbeat.err()) + } + groups, inputs := groupWorkerClaims(preparedClaims, preparedTexts, w.deps.ContextualDocuments) + result.ProviderCalls = 1 + result.ProviderDocuments = len(inputs) + for _, input := range inputs { + result.ProviderChunks += len(input.Chunks) + for _, text := range input.Chunks { + result.ProviderInputChars += utf8.RuneCountInString(text) + } + } + + vectors, providerErr := w.deps.Provider.EmbedDocuments(heartbeat.context(), inputs) + if ctxErr := ctx.Err(); ctxErr != nil { + return result, ctxErr + } + if heartbeatErr := heartbeat.err(); heartbeatErr != nil { + return result, heartbeatErr + } + outcomes, responseErr := validateWorkerProviderDocuments(inputs, vectors, providerErr, w.deps.Dimension) + var embeddings []Embedding + var completedClaims []*store.DocumentVectorChunkClaim + for documentIndex, outcome := range outcomes { + if outcome.err != nil { + continue + } + for chunkIndex, vector := range outcome.vectors { + claim := groups[documentIndex][chunkIndex] + embeddings = append(embeddings, Embedding{Token: claim.Token, Vector: vector}) + completedClaims = append(completedClaims, claim) + } + } + result.Embedded = len(completedClaims) + if len(embeddings) > 0 { + if err := w.deps.Backend.PutUnpublished(heartbeat.context(), generationID, w.deps.Dimension, embeddings); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return result, ctxErr + } + if heartbeatErr := heartbeat.err(); heartbeatErr != nil { + return result, heartbeatErr + } + putErr := fmt.Errorf("%w: %w", errBackendPut, err) + failureErr := w.failClaims(ctx, generationID, completedClaims, putErr, heartbeat, &result) + runErr = errors.Join(runErr, fmt.Errorf("put unpublished document vectors: %w", err), failureErr) + } else { + for _, claim := range completedClaims { + if ctxErr := ctx.Err(); ctxErr != nil { + return result, errors.Join(runErr, ctxErr) + } + if err := heartbeat.releaseForTransition(claim); err != nil { + runErr = errors.Join(runErr, err) + break + } + err := w.deps.Ledger.CommitDocumentVectorPublication(heartbeat.context(), int64(generationID), claim.Token, claim.LeaseOwner, claim.LeaseFence, w.deps.Now()) + switch { + case err == nil: + result.Published++ + case errors.Is(err, store.ErrDocumentVectorSourceChanged): + result.SourceChanged++ + if deleteErr := w.deps.Backend.DeleteTokens(heartbeat.context(), generationID, []string{claim.Token}); deleteErr != nil { + runErr = errors.Join(runErr, fmt.Errorf("delete changed document vector token: %w", deleteErr)) + } + default: + runErr = errors.Join(runErr, err) + } + } + } + } + for documentIndex, outcome := range outcomes { + if outcome.err == nil { + continue + } + failureErr := w.failClaims(ctx, generationID, groups[documentIndex], outcome.err, heartbeat, &result) + runErr = errors.Join(runErr, outcome.err, failureErr) + } + runErr = errors.Join(runErr, responseErr) + runErr = errors.Join(runErr, heartbeat.err()) + return result, runErr +} + +func (w *Worker) validate(generationID GenerationID) error { + if generationID <= 0 { + return errors.New("document vector worker generation must be positive") + } + if w.deps.Ledger == nil || w.deps.Provider == nil || w.deps.Backend == nil { + return errors.New("document vector worker ledger, provider, and backend are required") + } + if strings.TrimSpace(w.deps.Owner) == "" || w.deps.Dimension <= 0 || w.deps.MaxInputChars <= 0 || + w.deps.LeaseDuration <= 0 || w.deps.HeartbeatInterval <= 0 || + w.deps.HeartbeatInterval >= w.deps.LeaseDuration || w.deps.RetryDelay <= 0 || + w.deps.RetryDelay > maxWorkerRetryDelay || + w.deps.MaxAttempts <= 0 || w.deps.AfterGenerationID < 0 || w.deps.AfterChunkID < 0 || + (w.deps.AfterGenerationID == 0) != (w.deps.AfterChunkID == 0) { + return errors.New("document vector worker policy is invalid") + } + if w.deps.Recipe.Fingerprint() == "" || w.deps.Recipe.Values().Mode != docembedding.RepresentationRaw { + return errors.New("document vector worker requires a valid raw Docbank embedding recipe") + } + return nil +} + +type normalizedDocumentLedger interface { + LoadNormalizedDocument(ctx context.Context, extractionID string) (docbankdocument.NormalizedDocument, error) +} + +func prepareDocbankClaimInputs( + ctx context.Context, ledger Ledger, recipe docembedding.Recipe, + claims []*store.DocumentVectorChunkClaim, +) (map[string]string, error) { + source, ok := ledger.(normalizedDocumentLedger) + if !ok { + return nil, errors.New("document vector ledger cannot load normalized documents") + } + inputsByToken := make(map[string]string, len(claims)) + plans := make(map[string]map[string]docembedding.EmbeddingInput) + for _, claim := range claims { + byChunk := plans[claim.ExtractionID] + if byChunk == nil { + normalized, err := source.LoadNormalizedDocument(ctx, claim.ExtractionID) + if err != nil { + return nil, fmt.Errorf("load normalized document for embedding: %w", err) + } + plan, err := docembedding.BuildEmbeddingPlan( + normalized, docembedding.DocumentContext{}, recipe, nil, + ) + if err != nil { + return nil, fmt.Errorf("build document embedding plan: %w", err) + } + byChunk = make(map[string]docembedding.EmbeddingInput, len(plan.Inputs)) + for _, input := range plan.Inputs { + if input.Kind != docembedding.RepresentationKindRaw || len(input.SourceRefs) != 1 { + return nil, errors.New("raw document embedding plan returned an invalid input") + } + byChunk[input.SourceRefs[0].ChunkKey] = input + } + plans[claim.ExtractionID] = byChunk + } + input, ok := byChunk[claim.ChunkKey] + if !ok || len(input.SourceRefs) != 1 || input.SourceRefs[0].ChunkChecksum != claim.ChunkChecksum { + return nil, fmt.Errorf("document embedding plan does not match claimed chunk %q", claim.ChunkKey) + } + inputsByToken[claim.Token] = input.Text + } + return inputsByToken, nil +} + +func (w *Worker) bindCursor(generationID GenerationID) { + if w.cursorGenerationID != 0 && w.cursorGenerationID != generationID { + w.afterChunkID = 0 + } + w.cursorGenerationID = generationID +} + +func (w *Worker) collectClaims(ctx context.Context, generationID GenerationID, limit int, result *RunResult) ([]*store.DocumentVectorChunkClaim, error) { + claims := make([]*store.DocumentVectorChunkClaim, 0, limit) + after := w.afterChunkID + result.AfterChunkID = after + candidates, err := w.deps.Ledger.ListDocumentVectorChunkCandidates(ctx, int64(generationID), after, maxWorkerRunLimit) + if err != nil { + return nil, fmt.Errorf("list document vector candidates: %w", err) + } + if len(candidates) == 0 { + result.Exhausted = true + w.resetCursor(result) + return claims, nil + } + processed := 0 + for _, candidate := range candidates { + processed++ + if candidate.ChunkID <= after { + continue + } + claim, err := w.deps.Ledger.ClaimDocumentVectorChunk(ctx, int64(generationID), after, 1, w.deps.Owner, w.deps.Now(), w.deps.LeaseDuration) + if err != nil { + return nil, fmt.Errorf("claim document vector chunk: %w", err) + } + after = candidate.ChunkID + if claim != nil { + claims = append(claims, claim) + result.Claimed++ + if claim.ChunkID > after { + after = claim.ChunkID + } + } + result.AfterChunkID = after + if len(claims) == limit { + break + } + } + if processed == len(candidates) && len(candidates) < maxWorkerRunLimit { + result.Exhausted = true + w.resetCursor(result) + } else { + w.afterChunkID = after + } + return claims, nil +} + +func (w *Worker) resetCursor(result *RunResult) { + w.cursorGenerationID = 0 + w.afterChunkID = 0 + result.AfterGenerationID = 0 + result.AfterChunkID = 0 +} + +func groupWorkerClaims( + claims []*store.DocumentVectorChunkClaim, preparedTexts map[string]string, contextualDocuments bool, +) ([][]*store.DocumentVectorChunkClaim, []vector.DocumentInput) { + if contextualDocuments { + groups := make([][]*store.DocumentVectorChunkClaim, len(claims)) + inputs := make([]vector.DocumentInput, len(claims)) + for index, claim := range claims { + groups[index] = []*store.DocumentVectorChunkClaim{claim} + inputs[index] = vector.DocumentInput{Chunks: []string{preparedTexts[claim.Token]}} + } + return groups, inputs + } + var groups [][]*store.DocumentVectorChunkClaim + var inputs []vector.DocumentInput + indices := make(map[string]int) + for _, claim := range claims { + index, ok := indices[claim.ExtractionID] + if !ok { + index = len(groups) + indices[claim.ExtractionID] = index + groups = append(groups, nil) + inputs = append(inputs, vector.DocumentInput{}) + } + groups[index] = append(groups[index], claim) + inputs[index].Chunks = append(inputs[index].Chunks, preparedTexts[claim.Token]) + } + return groups, inputs +} + +func groupWorkerClaimsByExtraction(claims []*store.DocumentVectorChunkClaim) [][]*store.DocumentVectorChunkClaim { + groups := make([][]*store.DocumentVectorChunkClaim, 0) + indices := make(map[string]int) + for _, claim := range claims { + index, ok := indices[claim.ExtractionID] + if !ok { + index = len(groups) + indices[claim.ExtractionID] = index + groups = append(groups, nil) + } + groups[index] = append(groups[index], claim) + } + return groups +} + +func (w *Worker) failClaims( + ctx context.Context, + generationID GenerationID, + claims []*store.DocumentVectorChunkClaim, + cause error, + heartbeat *workerClaimHeartbeat, + result *RunResult, +) error { + var failureErr error + for _, claim := range claims { + if ctxErr := ctx.Err(); ctxErr != nil { + return errors.Join(failureErr, ctxErr) + } + if err := heartbeat.releaseForTransition(claim); err != nil { + return errors.Join(failureErr, err) + } + now := w.deps.Now() + terminal, code := workerFailureDisposition(cause, claim.AttemptCount, w.deps.MaxAttempts) + var retryAt *time.Time + if !terminal { + deadline := now.Add(w.deps.RetryDelay) + retryAt = &deadline + } + err := w.deps.Ledger.FailDocumentVectorChunk( + heartbeat.context(), int64(generationID), claim.Token, claim.LeaseOwner, claim.LeaseFence, + now, retryAt, terminal, code, + ) + if err != nil { + failureErr = errors.Join(failureErr, err) + continue + } + if terminal { + result.Terminal++ + } else { + result.Retry++ + } + } + return failureErr +} + +func workerFailureDisposition(cause error, attemptCount, maxAttempts int) (bool, string) { + if attemptCount >= maxAttempts { + return true, "attempt_limit" + } + switch { + case errors.Is(cause, errInvalidProviderShape): + return true, "invalid_provider_shape" + case errors.Is(cause, errInvalidProviderVector): + return true, "invalid_provider_vector" + case errors.Is(cause, vector.ErrPermanent4xx): + return true, "provider_rejected" + case errors.Is(cause, errInputPreparation): + return false, "input_preparation" + case errors.Is(cause, errBackendPut): + return false, "backend_transient" + default: + return false, "provider_transient" + } +} + +type workerProviderDocumentOutcome struct { + vectors [][]float32 + err error +} + +func validateWorkerProviderDocuments( + inputs []vector.DocumentInput, vectors [][][]float32, providerErr error, dimension int, +) ([]workerProviderDocumentOutcome, error) { + outcomes := make([]workerProviderDocumentOutcome, len(inputs)) + var responseErr error + if len(vectors) > len(inputs) { + responseErr = fmt.Errorf("%w: document count got %d, expected at most %d", errInvalidProviderShape, len(vectors), len(inputs)) + } + for documentIndex := range inputs { + if documentIndex >= len(vectors) { + if providerErr != nil { + outcomes[documentIndex].err = providerErr + } else { + outcomes[documentIndex].err = fmt.Errorf( + "%w: document count got %d, expected %d", errInvalidProviderShape, len(vectors), len(inputs), + ) + } + continue + } + documentVectors := vectors[documentIndex] + if len(documentVectors) != len(inputs[documentIndex].Chunks) { + outcomes[documentIndex].err = fmt.Errorf( + "%w: document %d chunk count got %d, expected %d", + errInvalidProviderShape, documentIndex, len(documentVectors), len(inputs[documentIndex].Chunks), + ) + continue + } + for _, embedding := range documentVectors { + if vectorErr := validateWorkerProviderVector(embedding, dimension); vectorErr != nil { + outcomes[documentIndex].err = vectorErr + break + } + } + if outcomes[documentIndex].err == nil { + outcomes[documentIndex].vectors = documentVectors + } + } + return outcomes, errors.Join(responseErr, providerErr) +} + +func validateWorkerProviderVector(embedding []float32, dimension int) error { + if len(embedding) != dimension { + return fmt.Errorf("%w: dimension got %d, expected %d", errInvalidProviderVector, len(embedding), dimension) + } + var squaredNorm float64 + for _, value := range embedding { + if math.IsNaN(float64(value)) || math.IsInf(float64(value), 0) { + return fmt.Errorf("%w: nonfinite component", errInvalidProviderVector) + } + squaredNorm += float64(value) * float64(value) + } + if squaredNorm == 0 { + return fmt.Errorf("%w: zero norm", errInvalidProviderVector) + } + return nil +} diff --git a/internal/vector/document/worker_test.go b/internal/vector/document/worker_test.go new file mode 100644 index 000000000..c713e1aad --- /dev/null +++ b/internal/vector/document/worker_test.go @@ -0,0 +1,1149 @@ +package document + +import ( + "context" + "errors" + "fmt" + "math" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + "unicode/utf8" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + docbankdocument "go.kenn.io/docbank/document" + docembedding "go.kenn.io/docbank/document/embedding" + "go.kenn.io/msgvault/internal/store" + "go.kenn.io/msgvault/internal/vector/embed" +) + +var _ Provider = (embed.SemanticClient)(nil) + +func TestPrepareDocbankClaimInputsUsesValidatedRawPlan(t *testing.T) { + assert := assert.New(t) + require := require.New(t) + policy, err := docbankdocument.NewNormalizePolicy(10_000) + require.NoError(err) + normalized, err := docbankdocument.NormalizeDocument(docbankdocument.SourceDocument{ + Family: "pdf", UnitKind: "page", Units: []docbankdocument.SourceUnit{{ + Index: 0, Markdown: "# Findings\n\nImportant evidence", + }}, + }, policy) + require.NoError(err) + recipe, err := docembedding.NewRecipe(docembedding.RecipeConfig{ + Mode: docembedding.RepresentationRaw, MaxInputRunes: 1000, + }) + require.NoError(err) + claim := workerClaim("extract-a", 1, normalized.Chunks[0].Text, strings.Repeat("a", 64)) + claim.ChunkKey = normalized.Chunks[0].Key + claim.ChunkChecksum = normalized.Chunks[0].Checksum + ledger := &normalizedFakeDocumentVectorLedger{ + fakeDocumentVectorLedger: newFakeDocumentVectorLedger(claim), normalized: normalized, + } + + inputs, err := prepareDocbankClaimInputs(t.Context(), ledger, recipe, []*store.DocumentVectorChunkClaim{claim}) + + require.NoError(err) + assert.Contains(inputs[claim.Token], "Heading: Findings") + assert.Contains(inputs[claim.Token], "Source: page 1") + assert.Contains(inputs[claim.Token], "Content:\n# Findings\nImportant evidence") +} + +type normalizedFakeDocumentVectorLedger struct { + *fakeDocumentVectorLedger + + normalized docbankdocument.NormalizedDocument +} + +func (l *normalizedFakeDocumentVectorLedger) LoadNormalizedDocument(context.Context, string) (docbankdocument.NormalizedDocument, error) { + return l.normalized, nil +} + +func TestWorkerRunRejectsInvalidLimit(t *testing.T) { + worker := NewWorker(WorkerDeps{}) + + for _, limit := range []int{0, 1001} { + _, err := worker.Run(t.Context(), 1, limit) + require.ErrorContains(t, err, "limit") + } +} + +func TestWorkerRunRejectsInvalidGenerationAndDimension(t *testing.T) { + ledger := newFakeDocumentVectorLedger() + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{}, &fakeDocumentVectorBackend{}) + + _, err := worker.Run(t.Context(), 0, 1) + require.ErrorContains(t, err, "generation") + + ledger.generation.State = store.DocumentVectorGenerationActive + _, err = worker.Run(t.Context(), 1, 1) + require.ErrorIs(t, err, store.ErrDocumentVectorInvalidGenerationState) + + ledger.generation.State = store.DocumentVectorGenerationBuilding + ledger.generation.Dimension = 4 + _, err = worker.Run(t.Context(), 1, 1) + require.ErrorContains(t, err, "dimension") +} + +func TestWorkerRunRejectsUnboundedRetryAndIncompleteRestoredCursor(t *testing.T) { + ledger := newFakeDocumentVectorLedger() + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{}, &fakeDocumentVectorBackend{}) + worker.deps.RetryDelay = 7*24*time.Hour + time.Millisecond + + _, err := worker.Run(t.Context(), 1, 1) + require.ErrorContains(t, err, "policy") + + worker = newFakeWorker(ledger, &fakeDocumentVectorProvider{}, &fakeDocumentVectorBackend{}) + worker.deps.AfterChunkID = 12 + worker.afterChunkID = 12 + _, err = worker.Run(t.Context(), 1, 1) + require.ErrorContains(t, err, "policy") +} + +func TestWorkerRunPreservesExactTextAndExtractionBoundaries(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, " alpha\n", "token-a"), + workerClaim("extract-a", 2, "beta", "token-b"), + workerClaim("extract-b", 3, "γamma", "token-c"), + ) + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{ + {{1, 2, 3}, {4, 5, 6}}, + {{7, 8, 9}}, + }} + provider.call = func(_ context.Context, _ []embed.DocumentInput) ([][][]float32, error) { + assert.Equal(t, 3, ledger.claimCalls, "all provider inputs must be claimed first") + return provider.vectors, nil + } + backend := &fakeDocumentVectorBackend{} + ledger.beforeCommit = func(testedToken string) { + assert.NotEmpty(t, backend.puts, "backend put must precede fenced commit") + assert.Contains(t, embeddingTokens(backend.puts[0]), testedToken) + } + worker := newFakeWorker(ledger, provider, backend) + + result, err := worker.Run(t.Context(), 1, 4) + requirements.NoError(err) + requirements.Len(provider.calls, 1) + assertions.Equal([]embed.DocumentInput{ + {Chunks: []string{" alpha\n", "beta"}}, + {Chunks: []string{"γamma"}}, + }, provider.calls[0]) + assertions.Equal(3, result.Claimed) + assertions.Equal(3, result.Embedded) + assertions.Equal(3, result.Published) + assertions.Equal(1, result.ProviderCalls) + assertions.Equal(2, result.ProviderDocuments) + assertions.Equal(3, result.ProviderChunks) + assertions.Equal(utf8.RuneCountInString(" alpha\n")+utf8.RuneCountInString("beta")+utf8.RuneCountInString("γamma"), + result.ProviderInputChars, + ) + assertions.True(result.Exhausted) + assertions.Zero(result.AfterGenerationID) + assertions.Zero(result.AfterChunkID) + requirements.Len(backend.puts, 1) + assertions.Equal([]string{"token-a", "token-b", "token-c"}, embeddingTokens(backend.puts[0])) + assertions.Equal([]string{"token-a", "token-b", "token-c"}, ledger.committed) +} + +func TestWorkerRunCapsProviderInputOnRuneBoundary(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "αβγδε", "token-a"), + ) + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}}}} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + worker.deps.MaxInputChars = 3 + recipe, err := docembedding.NewRecipe(docembedding.RecipeConfig{ + Mode: docembedding.RepresentationRaw, MaxInputRunes: 3, + }) + requirements.NoError(err) + worker.deps.Recipe = recipe + + result, err := worker.Run(t.Context(), 1, 1) + requirements.NoError(err) + requirements.Len(provider.calls, 1) + assertions.Equal([]embed.DocumentInput{{Chunks: []string{"αβγ"}}}, provider.calls[0]) + assertions.Equal(3, result.ProviderInputChars) +} + +func TestWorkerRunNeverClaimsMoreThanLimit(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "one", "token-a"), + workerClaim("extract-a", 2, "two", "token-b"), + workerClaim("extract-a", 3, "three", "token-c"), + ) + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}, {4, 5, 6}}}} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + + result, err := worker.Run(t.Context(), 1, 2) + requirements.NoError(err) + assertions.Equal(2, result.Claimed) + assertions.False(result.Exhausted) + assertions.Equal(2, ledger.claimCalls) + requirements.Len(provider.calls, 1) + assertions.Equal([]embed.DocumentInput{{Chunks: []string{"one", "two"}}}, provider.calls[0]) +} + +func TestWorkerRunContextualBoundariesDoNotDependOnLimit(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + claims := func() []*store.DocumentVectorChunkClaim { + return []*store.DocumentVectorChunkClaim{ + workerClaim("extract-a", 1, "one", "token-a"), + workerClaim("extract-a", 2, "two", "token-b"), + workerClaim("extract-a", 3, "three", "token-c"), + } + } + wantInputs := []embed.DocumentInput{ + {Chunks: []string{"one"}}, {Chunks: []string{"two"}}, {Chunks: []string{"three"}}, + } + + boundedProvider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}}, {{4, 5, 6}}}} + boundedWorker := newFakeWorker(newFakeDocumentVectorLedger(claims()...), boundedProvider, &fakeDocumentVectorBackend{}) + boundedWorker.deps.ContextualDocuments = true + first, err := boundedWorker.Run(t.Context(), 1, 2) + requirements.NoError(err) + requirements.False(first.Exhausted) + boundedProvider.vectors = [][][]float32{{{7, 8, 9}}} + second, err := boundedWorker.Run(t.Context(), 1, 2) + requirements.NoError(err) + requirements.True(second.Exhausted) + + unboundedProvider := &fakeDocumentVectorProvider{vectors: [][][]float32{ + {{1, 2, 3}}, {{4, 5, 6}}, {{7, 8, 9}}, + }} + unboundedWorker := newFakeWorker(newFakeDocumentVectorLedger(claims()...), unboundedProvider, &fakeDocumentVectorBackend{}) + unboundedWorker.deps.ContextualDocuments = true + _, err = unboundedWorker.Run(t.Context(), 1, 3) + requirements.NoError(err) + + assertions.Equal(wantInputs, append(boundedProvider.calls[0], boundedProvider.calls[1]...)) + assertions.Equal(wantInputs, unboundedProvider.calls[0]) +} + +func TestWorkerRunBoundsScanningAndCarriesCursorWithoutStarvation(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + claims := make([]*store.DocumentVectorChunkClaim, 1500) + for index := range claims { + claims[index] = workerClaim( + "extract-a", int64(index+1), fmt.Sprintf("chunk-%d", index+1), fmt.Sprintf("token-%d", index+1), + ) + } + ledger := newFakeDocumentVectorLedger(claims...) + ledger.unclaimableThrough = 1499 + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}}}} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + + first, err := worker.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.Zero(first.Claimed) + assertions.Equal(int64(1000), first.AfterChunkID) + assertions.False(first.Exhausted) + assertions.Equal(1000, ledger.claimCalls) + assertions.Empty(provider.calls) + + second, err := worker.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.Equal(1, second.Claimed) + assertions.Equal(1, second.Published) + assertions.True(second.Exhausted) + assertions.Zero(second.AfterGenerationID) + assertions.Zero(second.AfterChunkID) + assertions.Equal(1500, ledger.claimCalls) + requirements.Len(provider.calls, 1) + assertions.Equal([]embed.DocumentInput{{Chunks: []string{"chunk-1500"}}}, provider.calls[0]) +} + +func TestWorkerRunExhaustedCursorRoundTripsAsReset(t *testing.T) { + tests := []struct { + name string + claims []*store.DocumentVectorChunkClaim + unclaimable int64 + wantClaimCalls int + }{ + {name: "empty corpus"}, + { + name: "nonempty tail", + claims: []*store.DocumentVectorChunkClaim{workerClaim("extract-a", 1, "leased", "token-a")}, + unclaimable: 1, + wantClaimCalls: 2, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(test.claims...) + ledger.unclaimableThrough = test.unclaimable + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{}, &fakeDocumentVectorBackend{}) + + first, err := worker.Run(t.Context(), 1, 1) + requirements.NoError(err) + requirements.True(first.Exhausted) + + resumeDeps := worker.deps + resumeDeps.AfterGenerationID = first.AfterGenerationID + resumeDeps.AfterChunkID = first.AfterChunkID + resumed := NewWorker(resumeDeps) + second, err := resumed.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.True(second.Exhausted) + assertions.Zero(first.AfterGenerationID) + assertions.Zero(first.AfterChunkID) + assertions.Zero(second.AfterGenerationID) + assertions.Zero(second.AfterChunkID) + assertions.Equal(test.wantClaimCalls, ledger.claimCalls) + }) + } +} + +func TestWorkerRunResetsInMemoryCursorWhenGenerationChanges(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + claims := make([]*store.DocumentVectorChunkClaim, 1000) + for index := range claims { + claims[index] = workerClaim("extract-a", int64(index+1), "old", fmt.Sprintf("old-%d", index+1)) + } + ledger := newFakeDocumentVectorLedger(claims...) + ledger.unclaimableThrough = 1000 + provider := &fakeDocumentVectorProvider{} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + + first, err := worker.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.Equal(GenerationID(1), first.AfterGenerationID) + assertions.Equal(int64(1000), first.AfterChunkID) + + ledger.generation.ID = 2 + ledger.unclaimableThrough = 0 + newClaim := workerClaim("extract-b", 1, "new generation", "new-token") + newClaim.GenerationID = 2 + nextClaim := workerClaim("extract-b", 2, "next", "next-token") + nextClaim.GenerationID = 2 + ledger.claims = []*store.DocumentVectorChunkClaim{newClaim, nextClaim} + provider.vectors = [][][]float32{{{1, 2, 3}}} + second, err := worker.Run(t.Context(), 2, 1) + requirements.NoError(err) + assertions.Equal(1, second.Published) + assertions.Equal(GenerationID(2), second.AfterGenerationID) + assertions.Equal(int64(1), second.AfterChunkID) + assertions.False(second.Exhausted) + requirements.Len(provider.calls, 1) + assertions.Equal([]embed.DocumentInput{{Chunks: []string{"new generation"}}}, provider.calls[0]) +} + +func TestWorkerRunPublishesCompletedProviderPrefixAndRetriesOnlySuffix(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + providerErr := errors.New("provider unavailable") + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-b", 2, "second", "token-b"), + ) + provider := &fakeDocumentVectorProvider{ + vectors: [][][]float32{{{1, 2, 3}}}, + err: providerErr, + } + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, providerErr) + assertions.Equal(1, result.Embedded) + assertions.Equal(1, result.Published) + assertions.Equal(1, result.Retry) + assertions.Zero(result.Terminal) + assertions.Equal([]string{"token-a"}, ledger.committed) + requirements.Len(ledger.failures, 1) + assertions.Equal("token-b", ledger.failures[0].token) + assertions.False(ledger.failures[0].terminal) + assertions.Equal("provider_transient", ledger.failures[0].errorCode) + requirements.NotNil(ledger.failures[0].retryAt) + assertions.Equal(workerNow.Add(time.Minute), *ledger.failures[0].retryAt) +} + +func TestWorkerRunMakesPermanentAndMalformedProviderFailuresTerminal(t *testing.T) { + tests := []struct { + name string + vectors [][][]float32 + err error + wantCode string + }{ + { + name: "permanent provider rejection", + err: fmt.Errorf("request rejected: %w", embed.ErrPermanent4xx), + wantCode: "provider_rejected", + }, + { + name: "document cardinality", + vectors: [][][]float32{{{1, 2, 3}}, {{4, 5, 6}}}, + wantCode: "invalid_provider_shape", + }, + { + name: "chunk cardinality", + vectors: [][][]float32{{{1, 2, 3}}}, + wantCode: "invalid_provider_shape", + }, + { + name: "vector dimension", + vectors: [][][]float32{{{1, 2}, {4, 5, 6}}}, + wantCode: "invalid_provider_vector", + }, + { + name: "nonfinite vector", + vectors: [][][]float32{{{1, float32(math.NaN()), 3}, {4, 5, 6}}}, + wantCode: "invalid_provider_vector", + }, + { + name: "zero norm vector", + vectors: [][][]float32{{{0, 0, 0}, {4, 5, 6}}}, + wantCode: "invalid_provider_vector", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-a", 2, "second", "token-b"), + ) + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{ + vectors: test.vectors, err: test.err, + }, &fakeDocumentVectorBackend{}) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.Error(err) + assertions.Zero(result.Published) + assertions.Equal(2, result.Terminal) + requirements.Len(ledger.failures, 2) + assertions.Equal(test.wantCode, ledger.failures[0].errorCode) + assertions.Equal(test.wantCode, ledger.failures[1].errorCode) + assertions.True(ledger.failures[0].terminal) + assertions.Nil(ledger.failures[0].retryAt) + }) + } +} + +func TestWorkerRunMakesMalformedHTTPProviderResponsesTerminal(t *testing.T) { + tests := []struct { + name string + response string + wantCode string + }{ + { + name: "response shape", + response: `{"data":[]}`, + wantCode: "invalid_provider_shape", + }, + { + name: "vector dimension", + response: `{"data":[{"embedding":[1,2],"index":0}]}`, + wantCode: "invalid_provider_vector", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, err := w.Write([]byte(test.response)) + assert.NoError(t, err) + })) + t.Cleanup(server.Close) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + client := embed.NewClient(embed.Config{ + Endpoint: server.URL, Model: "embed-test", Dimension: 3, MaxRetries: 1, + }) + worker := newFakeWorker(ledger, client, &fakeDocumentVectorBackend{}) + + result, err := worker.Run(t.Context(), 1, 1) + + require.Error(t, err) + assert.Equal(t, 1, result.Terminal) + assert.Zero(t, result.Retry) + require.Len(t, ledger.failures, 1) + assert.True(t, ledger.failures[0].terminal) + assert.Equal(t, test.wantCode, ledger.failures[0].errorCode) + }) + } +} + +func TestWorkerRunIsolatesMalformedProviderDocumentFromHealthySuffix(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-b", 2, "second", "token-b"), + workerClaim("extract-c", 3, "third", "token-c"), + ) + backend := &fakeDocumentVectorBackend{} + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{vectors: [][][]float32{ + {{1, 2, 3}}, {{4, 5}}, {{7, 8, 9}}, + }}, backend) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, errInvalidProviderVector) + assertions.Equal(2, result.Embedded) + assertions.Equal(2, result.Published) + assertions.Equal(1, result.Terminal) + assertions.Equal([]string{"token-a", "token-c"}, ledger.committed) + requirements.Len(ledger.failures, 1) + assertions.Equal("token-b", ledger.failures[0].token) + assertions.Equal("invalid_provider_vector", ledger.failures[0].errorCode) + requirements.Len(backend.puts, 1) + assertions.Equal([]string{"token-a", "token-c"}, embeddingTokens(backend.puts[0])) +} + +func TestWorkerRunRecordsPreparationFailureAndPublishesHealthyExtraction(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-b", 2, "second", "token-b"), + ) + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}}}} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + worker.deps.prepareInputs = func(_ context.Context, _ Ledger, _ docembedding.Recipe, claims []*store.DocumentVectorChunkClaim) (map[string]string, error) { + if claims[0].ExtractionID == "extract-b" { + return nil, errors.New("normalized document is corrupt") + } + return map[string]string{claims[0].Token: claims[0].Text}, nil + } + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, errInputPreparation) + assertions.Equal(1, result.Published) + assertions.Equal(1, result.Retry) + assertions.Zero(result.Terminal) + assertions.Equal([]string{"token-a"}, ledger.committed) + requirements.Len(ledger.failures, 1) + assertions.Equal("token-b", ledger.failures[0].token) + assertions.Equal("input_preparation", ledger.failures[0].errorCode) + requirements.Len(provider.calls, 1) + assertions.Equal([]embed.DocumentInput{{Chunks: []string{"first"}}}, provider.calls[0]) +} + +func TestWorkerRunPreparationFailureHonorsAttemptCeiling(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + claim := workerClaim("extract-a", 1, "first", "token-a") + claim.AttemptCount = 3 + ledger := newFakeDocumentVectorLedger(claim) + provider := &fakeDocumentVectorProvider{} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + worker.deps.prepareInputs = func(context.Context, Ledger, docembedding.Recipe, []*store.DocumentVectorChunkClaim) (map[string]string, error) { + return nil, errors.New("normalized document is corrupt") + } + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, errInputPreparation) + assertions.Zero(result.Published) + assertions.Zero(result.Retry) + assertions.Equal(1, result.Terminal) + requirements.Len(ledger.failures, 1) + assertions.Equal("attempt_limit", ledger.failures[0].errorCode) + assertions.Empty(provider.calls) +} + +func TestWorkerRunAttemptCeilingMakesTransientFailureTerminal(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + claim := workerClaim("extract-a", 1, "first", "token-a") + claim.AttemptCount = 3 + ledger := newFakeDocumentVectorLedger(claim) + providerErr := errors.New("provider unavailable") + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{err: providerErr}, &fakeDocumentVectorBackend{}) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, providerErr) + assertions.Zero(result.Retry) + assertions.Equal(1, result.Terminal) + requirements.Len(ledger.failures, 1) + assertions.Equal("attempt_limit", ledger.failures[0].errorCode) + assertions.True(ledger.failures[0].terminal) +} + +func TestWorkerRunDeletesOnlySourceChangedTokens(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-source"), + workerClaim("extract-b", 2, "second", "token-lost"), + ) + ledger.commitErr["token-source"] = store.ErrDocumentVectorSourceChanged + ledger.commitErr["token-lost"] = store.ErrDocumentVectorClaimLost + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{ + {{1, 2, 3}}, {{4, 5, 6}}, + }} + backend := &fakeDocumentVectorBackend{} + worker := newFakeWorker(ledger, provider, backend) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, store.ErrDocumentVectorClaimLost) + assertions.Zero(result.Published) + assertions.Equal(1, result.SourceChanged) + assertions.Equal([][]string{{"token-source"}}, backend.deletes) + assertions.NotContains(backend.deletes[0], "token-lost") +} + +func TestWorkerRunContextInterruptionLeavesClaimsLeasedForTakeover(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + started := make(chan struct{}) + provider := &fakeDocumentVectorProvider{call: func(ctx context.Context, _ []embed.DocumentInput) ([][][]float32, error) { + close(started) + <-ctx.Done() + return nil, ctx.Err() + }} + backend := &fakeDocumentVectorBackend{} + worker := newFakeWorker(ledger, provider, backend) + ctx, cancel := context.WithCancel(t.Context()) + done := make(chan error, 1) + go func() { + _, err := worker.Run(ctx, 1, 1) + done <- err + }() + <-started + cancel() + requirements.ErrorIs(<-done, context.Canceled) + assertions.Empty(ledger.failures) + assertions.Empty(ledger.committed) + assertions.Empty(backend.puts) + + takeover := workerClaim("extract-a", 1, "first", "token-a") + takeover.LeaseOwner = "worker-b" + takeover.LeaseFence = 2 + takeover.AttemptCount = 2 + ledger.claims = append(ledger.claims, takeover) + resumeProvider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}}}} + resumer := newFakeWorker(ledger, resumeProvider, backend) + resumer.deps.Owner = "worker-b" + result, err := resumer.Run(t.Context(), 1, 1) + requirements.NoError(err) + assertions.Equal(1, result.Published) + assertions.Equal([]string{"token-a"}, ledger.committed) +} + +func TestWorkerRunReplayDoesNotRepublishReadyWork(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + provider := &fakeDocumentVectorProvider{vectors: [][][]float32{{{1, 2, 3}}}} + backend := &fakeDocumentVectorBackend{} + worker := newFakeWorker(ledger, provider, backend) + + first, err := worker.Run(t.Context(), 1, 10) + requirements.NoError(err) + second, err := worker.Run(t.Context(), 1, 10) + requirements.NoError(err) + assertions.Equal(1, first.Published) + assertions.Zero(second.Claimed) + assertions.True(second.Exhausted) + assertions.Len(provider.calls, 1) + assertions.Len(backend.puts, 1) + assertions.Equal([]string{"token-a"}, ledger.committed) +} + +func TestWorkerRunRetriesBackendPutFailureWithoutCommitting(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + backendErr := errors.New("backend unavailable") + backend := &fakeDocumentVectorBackend{putErr: backendErr} + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{ + vectors: [][][]float32{{{1, 2, 3}}}, + }, backend) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, backendErr) + assertions.Equal(1, result.Embedded) + assertions.Zero(result.Published) + assertions.Equal(1, result.Retry) + assertions.Empty(ledger.committed) + requirements.Len(ledger.failures, 1) + assertions.Equal("backend_transient", ledger.failures[0].errorCode) +} + +func TestWorkerRunClassifiesProviderSuffixWhenPrefixBackendPutFails(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-b", 2, "second", "token-b"), + ) + providerErr := errors.New("provider unavailable") + backendErr := errors.New("backend unavailable") + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{ + vectors: [][][]float32{{{1, 2, 3}}}, err: providerErr, + }, &fakeDocumentVectorBackend{putErr: backendErr}) + + result, err := worker.Run(t.Context(), 1, 10) + requirements.ErrorIs(err, providerErr) + requirements.ErrorIs(err, backendErr) + assertions.Equal(2, result.Retry) + requirements.Len(ledger.failures, 2) + assertions.Equal("token-a", ledger.failures[0].token) + assertions.Equal("backend_transient", ledger.failures[0].errorCode) + assertions.Equal("token-b", ledger.failures[1].token) + assertions.Equal("provider_transient", ledger.failures[1].errorCode) +} + +func TestWorkerRunRenewsClaimsDuringLongProviderCall(t *testing.T) { + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + ledger.renewed = make(chan struct{}) + provider := &fakeDocumentVectorProvider{call: func(ctx context.Context, _ []embed.DocumentInput) ([][][]float32, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-ledger.renewed: + return [][][]float32{{{1, 2, 3}}}, nil + case <-time.After(100 * time.Millisecond): + return nil, errors.New("heartbeat did not renew claim") + } + }} + worker := newFakeWorker(ledger, provider, &fakeDocumentVectorBackend{}) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + + result, err := worker.Run(t.Context(), 1, 1) + require.NoError(t, err) + assert.Equal(t, 1, result.Published) + assert.Positive(t, ledger.renewCallCount()) +} + +func TestWorkerRunRenewsClaimsDuringLongBackendPut(t *testing.T) { + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + ledger.renewed = make(chan struct{}) + backend := &fakeDocumentVectorBackend{put: func(ctx context.Context, _ []Embedding) error { + select { + case <-ctx.Done(): + return ctx.Err() + case <-ledger.renewed: + return nil + case <-time.After(100 * time.Millisecond): + return errors.New("heartbeat did not renew during backend put") + } + }} + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{ + vectors: [][][]float32{{{1, 2, 3}}}, + }, backend) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + + result, err := worker.Run(t.Context(), 1, 1) + require.NoError(t, err) + assert.Equal(t, 1, result.Published) + assert.Positive(t, ledger.renewCallCount()) +} + +func TestWorkerRunCancelsBackendPutWhenHeartbeatLosesClaim(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + ledger.renewErr = store.ErrDocumentVectorClaimLost + backendCanceled := make(chan struct{}) + backend := &fakeDocumentVectorBackend{put: func(ctx context.Context, _ []Embedding) error { + select { + case <-ctx.Done(): + close(backendCanceled) + return ctx.Err() + case <-time.After(100 * time.Millisecond): + return errors.New("heartbeat did not cancel backend put") + } + }} + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{ + vectors: [][][]float32{{{1, 2, 3}}}, + }, backend) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + + result, err := worker.Run(t.Context(), 1, 1) + requirements.ErrorIs(err, store.ErrDocumentVectorClaimLost) + assertions.Zero(result.Published) + assertions.Zero(result.Retry) + assertions.Zero(result.Terminal) + assertions.Empty(ledger.committed) + assertions.Empty(ledger.failures) + assertions.Empty(backend.deletes) + select { + case <-backendCanceled: + default: + requirements.Fail("backend context was not canceled after renewal loss") + } +} + +func TestWorkerRunStopsRenewingClaimBeforeCommitAfterSynchronousRenewal(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + commitStarted := make(chan struct{}) + releaseCommit := make(chan struct{}) + ledger.beforeCommit = func(string) { + close(commitStarted) + <-releaseCommit + } + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{ + vectors: [][][]float32{{{1, 2, 3}}}, + }, &fakeDocumentVectorBackend{}) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + type outcome struct { + result RunResult + err error + } + done := make(chan outcome, 1) + go func() { + result, err := worker.Run(t.Context(), 1, 1) + done <- outcome{result: result, err: err} + }() + select { + case <-commitStarted: + case <-time.After(time.Second): + requirements.FailNow("worker did not reach commit") + } + before := ledger.renewTokenCallCount("token-a") + time.Sleep(30 * time.Millisecond) + after := ledger.renewTokenCallCount("token-a") + close(releaseCommit) + completed := <-done + time.Sleep(30 * time.Millisecond) + afterStop := ledger.renewTokenCallCount("token-a") + + requirements.NoError(completed.err) + assertions.Equal(1, completed.result.Published) + assertions.Equal(1, before, "claim is synchronously renewed before commit") + assertions.Equal(before, after, "heartbeat no longer owns a claim once commit begins") + assertions.Equal(after, afterStop, "worker shutdown leaves no heartbeat goroutine renewing claims") +} + +func TestWorkerRunCancelsCommitWhenAnotherClaimLosesHeartbeat(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-b", 2, "second", "token-b"), + ) + ledger.renewErrByToken["token-b"] = store.ErrDocumentVectorClaimLost + commitCanceled := make(chan struct{}) + ledger.commit = func(ctx context.Context, token string) error { + if token != "token-a" { + return nil + } + select { + case <-ctx.Done(): + close(commitCanceled) + return ctx.Err() + case <-time.After(100 * time.Millisecond): + return errors.New("heartbeat loss did not cancel commit") + } + } + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{vectors: [][][]float32{ + {{1, 2, 3}}, {{4, 5, 6}}, + }}, &fakeDocumentVectorBackend{}) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + + result, err := worker.Run(t.Context(), 1, 2) + requirements.ErrorIs(err, store.ErrDocumentVectorClaimLost) + assertions.Zero(result.Published) + assertions.Empty(ledger.committed) + assertions.Empty(ledger.failures) + select { + case <-commitCanceled: + default: + requirements.Fail("commit context was not canceled after another renewal loss") + } +} + +func TestWorkerRunCancelsFailureTransitionWhenAnotherClaimLosesHeartbeat(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger( + workerClaim("extract-a", 1, "first", "token-a"), + workerClaim("extract-b", 2, "second", "token-b"), + ) + ledger.renewErrByToken["token-b"] = store.ErrDocumentVectorClaimLost + failureCanceled := make(chan struct{}) + ledger.fail = func(ctx context.Context, token string) error { + if token != "token-a" { + return nil + } + select { + case <-ctx.Done(): + close(failureCanceled) + return ctx.Err() + case <-time.After(100 * time.Millisecond): + return errors.New("heartbeat loss did not cancel failure transition") + } + } + providerErr := errors.New("provider unavailable") + worker := newFakeWorker(ledger, &fakeDocumentVectorProvider{err: providerErr}, &fakeDocumentVectorBackend{}) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + + result, err := worker.Run(t.Context(), 1, 2) + requirements.ErrorIs(err, store.ErrDocumentVectorClaimLost) + requirements.ErrorIs(err, providerErr) + assertions.Zero(result.Retry) + assertions.Zero(result.Terminal) + assertions.Empty(ledger.committed) + assertions.Empty(ledger.failures) + select { + case <-failureCanceled: + default: + requirements.Fail("failure context was not canceled after another renewal loss") + } +} + +func TestWorkerRunCancelsPublicationWhenHeartbeatLosesClaim(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) + ledger.renewErr = store.ErrDocumentVectorClaimLost + providerCanceled := make(chan struct{}) + provider := &fakeDocumentVectorProvider{call: func(ctx context.Context, _ []embed.DocumentInput) ([][][]float32, error) { + select { + case <-ctx.Done(): + close(providerCanceled) + return nil, ctx.Err() + case <-time.After(100 * time.Millisecond): + return nil, errors.New("heartbeat did not cancel provider") + } + }} + backend := &fakeDocumentVectorBackend{} + worker := newFakeWorker(ledger, provider, backend) + worker.deps.LeaseDuration = 100 * time.Millisecond + worker.deps.HeartbeatInterval = 10 * time.Millisecond + + result, err := worker.Run(t.Context(), 1, 1) + requirements.ErrorIs(err, store.ErrDocumentVectorClaimLost) + assertions.Zero(result.Embedded) + assertions.Empty(backend.puts) + assertions.Empty(backend.deletes) + assertions.Empty(ledger.committed) + assertions.Empty(ledger.failures) + select { + case <-providerCanceled: + default: + requirements.Fail("provider context was not canceled after renewal loss") + } +} + +type fakeDocumentVectorLedger struct { + mu sync.Mutex + generation store.DocumentVectorGeneration + claims []*store.DocumentVectorChunkClaim + claimCalls int + committed []string + failures []fakeDocumentVectorFailure + commitErr map[string]error + commit func(context.Context, string) error + fail func(context.Context, string) error + beforeCommit func(string) + renewErr error + renewErrByToken map[string]error + renewCalls int + renewTokens map[string]int + renewed chan struct{} + renewOnce sync.Once + unclaimableThrough int64 +} + +type fakeDocumentVectorFailure struct { + token string + retryAt *time.Time + terminal bool + errorCode string +} + +func newFakeDocumentVectorLedger(claims ...*store.DocumentVectorChunkClaim) *fakeDocumentVectorLedger { + return &fakeDocumentVectorLedger{ + generation: store.DocumentVectorGeneration{ + ID: 1, State: store.DocumentVectorGenerationBuilding, + DocumentVectorGenerationSpec: store.DocumentVectorGenerationSpec{Dimension: 3}, + }, + claims: claims, commitErr: map[string]error{}, renewTokens: map[string]int{}, + renewErrByToken: map[string]error{}, + } +} + +func (l *fakeDocumentVectorLedger) GetDocumentVectorGeneration(context.Context, int64) (store.DocumentVectorGeneration, error) { + return l.generation, nil +} + +func (l *fakeDocumentVectorLedger) ListDocumentVectorChunkCandidates(_ context.Context, _ int64, after int64, limit int) ([]store.DocumentVectorChunkCandidate, error) { + candidates := make([]store.DocumentVectorChunkCandidate, 0, min(limit, len(l.claims))) + for _, claim := range l.claims { + if claim.ChunkID > after { + candidates = append(candidates, claim.DocumentVectorChunkCandidate) + if len(candidates) == limit { + break + } + } + } + return candidates, nil +} + +func (l *fakeDocumentVectorLedger) ClaimDocumentVectorChunk(_ context.Context, _ int64, after int64, _ int, _ string, _ time.Time, _ time.Duration) (*store.DocumentVectorChunkClaim, error) { + l.claimCalls++ + for index, claim := range l.claims { + if claim.ChunkID <= after { + continue + } + if claim.ChunkID <= l.unclaimableThrough { + return nil, nil //nolint:nilnil // No claim is a valid ledger result. + } + l.claims = append(l.claims[:index], l.claims[index+1:]...) + return claim, nil + } + return nil, nil //nolint:nilnil // No claim is a valid ledger result. +} + +func (l *fakeDocumentVectorLedger) RenewDocumentVectorChunkClaim(_ context.Context, _ int64, token, _ string, _ int64, _ time.Time, _ time.Duration) (time.Time, error) { + l.mu.Lock() + l.renewCalls++ + l.renewTokens[token]++ + err := l.renewErr + if tokenErr := l.renewErrByToken[token]; tokenErr != nil { + err = tokenErr + } + renewed := l.renewed + l.mu.Unlock() + if renewed != nil { + l.renewOnce.Do(func() { close(renewed) }) + } + return workerNow.Add(time.Minute), err +} + +func (l *fakeDocumentVectorLedger) renewTokenCallCount(token string) int { + l.mu.Lock() + defer l.mu.Unlock() + return l.renewTokens[token] +} + +func (l *fakeDocumentVectorLedger) renewCallCount() int { + l.mu.Lock() + defer l.mu.Unlock() + return l.renewCalls +} + +func (l *fakeDocumentVectorLedger) CommitDocumentVectorPublication(ctx context.Context, _ int64, token, _ string, _ int64, _ time.Time) error { + if l.beforeCommit != nil { + l.beforeCommit(token) + } + if l.commit != nil { + if err := l.commit(ctx, token); err != nil { + return err + } + } + if err := l.commitErr[token]; err != nil { + return err + } + l.committed = append(l.committed, token) + return nil +} + +func (l *fakeDocumentVectorLedger) FailDocumentVectorChunk(ctx context.Context, _ int64, token, _ string, _ int64, _ time.Time, retryAt *time.Time, terminal bool, errorCode string) error { + if l.fail != nil { + if err := l.fail(ctx, token); err != nil { + return err + } + } + l.failures = append(l.failures, fakeDocumentVectorFailure{ + token: token, retryAt: retryAt, terminal: terminal, errorCode: errorCode, + }) + return nil +} + +type fakeDocumentVectorProvider struct { + calls [][]embed.DocumentInput + vectors [][][]float32 + err error + call func(context.Context, []embed.DocumentInput) ([][][]float32, error) +} + +func (p *fakeDocumentVectorProvider) EmbedDocuments(ctx context.Context, inputs []embed.DocumentInput) ([][][]float32, error) { + cloned := make([]embed.DocumentInput, len(inputs)) + for index := range inputs { + cloned[index].Chunks = append([]string(nil), inputs[index].Chunks...) + } + p.calls = append(p.calls, cloned) + if p.call != nil { + return p.call(ctx, inputs) + } + return p.vectors, p.err +} + +type fakeDocumentVectorBackend struct { + puts [][]Embedding + deletes [][]string + putErr error + deleteErr error + put func(context.Context, []Embedding) error +} + +func (b *fakeDocumentVectorBackend) PutUnpublished(ctx context.Context, _ GenerationID, _ int, embeddings []Embedding) error { + b.puts = append(b.puts, append([]Embedding(nil), embeddings...)) + if b.put != nil { + return b.put(ctx, embeddings) + } + return b.putErr +} + +func (b *fakeDocumentVectorBackend) DeleteTokens(_ context.Context, _ GenerationID, tokens []string) error { + b.deletes = append(b.deletes, append([]string(nil), tokens...)) + return b.deleteErr +} + +func (*fakeDocumentVectorBackend) Search(context.Context, GenerationID, int, []float32, int) ([]Hit, error) { + return nil, errors.New("not used") +} + +var workerNow = time.Date(2026, 8, 20, 22, 0, 0, 0, time.UTC) + +func newFakeWorker(ledger *fakeDocumentVectorLedger, provider Provider, backend Backend) *Worker { + return NewWorker(WorkerDeps{ + Ledger: ledger, Provider: provider, Backend: backend, + Owner: "worker-a", Dimension: 3, MaxInputChars: 1000, LeaseDuration: time.Minute, + HeartbeatInterval: 10 * time.Second, RetryDelay: time.Minute, + MaxAttempts: 3, Now: func() time.Time { return workerNow }, + prepareInputs: func(_ context.Context, _ Ledger, recipe docembedding.Recipe, claims []*store.DocumentVectorChunkClaim) (map[string]string, error) { + inputs := make(map[string]string, len(claims)) + for _, claim := range claims { + runes := []rune(claim.Text) + if len(runes) > recipe.Values().MaxInputRunes { + runes = runes[:recipe.Values().MaxInputRunes] + } + inputs[claim.Token] = string(runes) + } + return inputs, nil + }, + }) +} + +func workerClaim(extractionID string, chunkID int64, text, token string) *store.DocumentVectorChunkClaim { + return &store.DocumentVectorChunkClaim{ + GenerationID: 1, ExtractionID: extractionID, ChunkID: chunkID, + ChunkKey: token + "-key", Text: text, + Token: token, LeaseOwner: "worker-a", LeaseFence: 1, + LeaseUntil: time.Date(2026, 8, 20, 22, 1, 0, 0, time.UTC), AttemptCount: 1, + } +} + +func embeddingTokens(embeddings []Embedding) []string { + tokens := make([]string, len(embeddings)) + for index := range embeddings { + tokens[index] = embeddings[index].Token + } + return tokens +} diff --git a/internal/vector/embed/client.go b/internal/vector/embed/client.go index 0472a2dff..ba63cc853 100644 --- a/internal/vector/embed/client.go +++ b/internal/vector/embed/client.go @@ -14,6 +14,8 @@ import ( "strconv" "strings" "time" + + "go.kenn.io/msgvault/internal/vector" ) // ErrPermanent4xx marks a non-retryable HTTP 4xx response from the @@ -21,7 +23,7 @@ import ( // it; the error message still carries the status code and a bounded // response body. 429 (rate-limited) and 5xx are NOT wrapped — they // flow through the retry loop as transient errors. -var ErrPermanent4xx = errors.New("embed: non-retryable 4xx response") +var ErrPermanent4xx = vector.ErrPermanent4xx // Config controls an embeddings Client. The zero value is not usable; callers // must set Endpoint, Model, and Dimension at a minimum. @@ -45,6 +47,9 @@ type Config struct { // BeforeRequest reauthorizes each concrete HTTP attempt. A returned error // is propagated without retrying. Nil leaves the client ungated. BeforeRequest BeforeRequestFunc + // RejectRedirects prevents provider responses from replaying embedding input + // to another URL. BeforeRequest clients always reject redirects as well. + RejectRedirects bool } // Client calls an OpenAI-compatible /v1/embeddings endpoint. @@ -61,7 +66,7 @@ func NewClient(cfg Config) *Client { if cfg.MaxRetries == 0 { cfg.MaxRetries = 3 } - return &Client{cfg: cfg, http: newHTTPClient(cfg.Timeout, cfg.BeforeRequest)} + return &Client{cfg: cfg, http: newHTTPClient(cfg.Timeout, cfg.BeforeRequest, cfg.RejectRedirects)} } // embeddingRequest is the JSON body sent to the server. @@ -149,7 +154,7 @@ func (c *Client) EmbedQuery(ctx context.Context, text string) ([]float32, error) return nil, err } if len(vecs) != 1 { - return nil, fmt.Errorf("embed query: expected exactly one vector, got %d", len(vecs)) + return nil, fmt.Errorf("%w: embed query expected exactly one vector, got %d", vector.ErrInvalidProviderShape, len(vecs)) } return vecs[0], nil } @@ -167,7 +172,7 @@ func (c *Client) EmbedDocuments(ctx context.Context, documents []DocumentInput) return nil, err } if len(vecs) != len(inputs) { - return nil, fmt.Errorf("embed documents: expected %d vectors, got %d", len(inputs), len(vecs)) + return nil, fmt.Errorf("%w: embed documents expected %d vectors, got %d", vector.ErrInvalidProviderShape, len(inputs), len(vecs)) } documentVecs := make([][][]float32, len(documents)) @@ -241,22 +246,22 @@ func (c *Client) doOnce(ctx context.Context, body []byte, want int) ([][]float32 return nil, &retryError{err: fmt.Errorf("decode response: %w", err)} } if len(r.Data) != want { - return nil, fmt.Errorf("embed: response count mismatch: got %d, expected %d", len(r.Data), want) + return nil, fmt.Errorf("%w: embed response count mismatch: got %d, expected %d", vector.ErrInvalidProviderShape, len(r.Data), want) } vecs := make([][]float32, want) for _, d := range r.Data { if d.Index < 0 || d.Index >= want { - return nil, fmt.Errorf("embed: invalid index %d (len=%d)", d.Index, want) + return nil, fmt.Errorf("%w: embed invalid index %d (len=%d)", vector.ErrInvalidProviderShape, d.Index, want) } if len(d.Embedding) != c.cfg.Dimension { - return nil, fmt.Errorf("embed: dimension mismatch: got %d, configured %d", - len(d.Embedding), c.cfg.Dimension) + return nil, fmt.Errorf("%w: embed dimension mismatch: got %d, configured %d", + vector.ErrInvalidProviderVector, len(d.Embedding), c.cfg.Dimension) } vecs[d.Index] = d.Embedding } for i, v := range vecs { if v == nil { - return nil, fmt.Errorf("embed: missing embedding at index %d", i) + return nil, fmt.Errorf("%w: embed missing embedding at index %d", vector.ErrInvalidProviderShape, i) } } return vecs, nil diff --git a/internal/vector/embed/document.go b/internal/vector/embed/document.go index e58acdfc3..a796d716f 100644 --- a/internal/vector/embed/document.go +++ b/internal/vector/embed/document.go @@ -58,9 +58,7 @@ type Document struct { } // DocumentInput is one document's ordered embedding inputs. -type DocumentInput struct { - Chunks []string -} +type DocumentInput = vector.DocumentInput // SemanticClient separates query embedding from document embedding while // preserving document boundaries. EmbedDocuments can return a completed diff --git a/internal/vector/embed/request_gate.go b/internal/vector/embed/request_gate.go index 84d9faf0e..0874d1da2 100644 --- a/internal/vector/embed/request_gate.go +++ b/internal/vector/embed/request_gate.go @@ -34,12 +34,14 @@ func (t beforeRequestTransport) RoundTrip(request *http.Request) (*http.Response return t.base.RoundTrip(request) } -func newHTTPClient(timeout time.Duration, before BeforeRequestFunc) *http.Client { +func newHTTPClient(timeout time.Duration, before BeforeRequestFunc, rejectRedirects bool) *http.Client { client := &http.Client{Timeout: timeout} if before != nil { client.Transport = beforeRequestTransport{ base: http.DefaultTransport, before: before, } + } + if before != nil || rejectRedirects { client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } diff --git a/internal/vector/embed/voyage.go b/internal/vector/embed/voyage.go index e57b268eb..c01259ae7 100644 --- a/internal/vector/embed/voyage.go +++ b/internal/vector/embed/voyage.go @@ -10,6 +10,8 @@ import ( "net/http" "strings" "time" + + "go.kenn.io/msgvault/internal/vector" ) // VoyageConfig controls the contextualized embeddings client. @@ -24,6 +26,9 @@ type VoyageConfig struct { // BeforeRequest reauthorizes each concrete HTTP attempt. A returned error // is propagated without retrying. Nil leaves the client ungated. BeforeRequest BeforeRequestFunc + // RejectRedirects prevents provider responses from replaying embedding input + // to another URL. BeforeRequest clients always reject redirects as well. + RejectRedirects bool } // VoyageClient calls Voyage's nested contextualized embeddings endpoint. @@ -50,7 +55,7 @@ func NewVoyageClient(cfg VoyageConfig) *VoyageClient { cfg.Limits.MaxUTF8Bytes = defaultVoyageRequestLimits.MaxUTF8Bytes } cfg.Limits = capVoyageRequestLimits(cfg.Limits) - return &VoyageClient{cfg: cfg, http: newHTTPClient(cfg.Timeout, cfg.BeforeRequest)} + return &VoyageClient{cfg: cfg, http: newHTTPClient(cfg.Timeout, cfg.BeforeRequest, cfg.RejectRedirects)} } type voyageRequest struct { @@ -227,35 +232,35 @@ func (c *VoyageClient) doVoyageOnce(ctx context.Context, body []byte, inputs [][ func (c *VoyageClient) decodeVoyageResponse(response voyageResponse, inputs [][]string) ([][][]float32, error) { if len(response.Data) != len(inputs) { - return nil, fmt.Errorf("embed: Voyage outer response count mismatch: got %d, expected %d", len(response.Data), len(inputs)) + return nil, fmt.Errorf("%w: Voyage outer response count mismatch: got %d, expected %d", vector.ErrInvalidProviderShape, len(response.Data), len(inputs)) } results := make([][][]float32, len(inputs)) outerSeen := make([]bool, len(inputs)) for _, outer := range response.Data { if outer.Index < 0 || outer.Index >= len(inputs) { - return nil, fmt.Errorf("embed: Voyage invalid outer index %d (len=%d)", outer.Index, len(inputs)) + return nil, fmt.Errorf("%w: Voyage invalid outer index %d (len=%d)", vector.ErrInvalidProviderShape, outer.Index, len(inputs)) } if outerSeen[outer.Index] { - return nil, fmt.Errorf("embed: Voyage duplicate outer index %d", outer.Index) + return nil, fmt.Errorf("%w: Voyage duplicate outer index %d", vector.ErrInvalidProviderShape, outer.Index) } outerSeen[outer.Index] = true expectedChunks := len(inputs[outer.Index]) if len(outer.Data) != expectedChunks { - return nil, fmt.Errorf("embed: Voyage inner response count mismatch at outer index %d: got %d, expected %d", outer.Index, len(outer.Data), expectedChunks) + return nil, fmt.Errorf("%w: Voyage inner response count mismatch at outer index %d: got %d, expected %d", vector.ErrInvalidProviderShape, outer.Index, len(outer.Data), expectedChunks) } vectors := make([][]float32, expectedChunks) innerSeen := make([]bool, expectedChunks) for _, inner := range outer.Data { if inner.Index < 0 || inner.Index >= expectedChunks { - return nil, fmt.Errorf("embed: Voyage invalid inner index %d at outer index %d (len=%d)", inner.Index, outer.Index, expectedChunks) + return nil, fmt.Errorf("%w: Voyage invalid inner index %d at outer index %d (len=%d)", vector.ErrInvalidProviderShape, inner.Index, outer.Index, expectedChunks) } if innerSeen[inner.Index] { - return nil, fmt.Errorf("embed: Voyage duplicate inner index %d at outer index %d", inner.Index, outer.Index) + return nil, fmt.Errorf("%w: Voyage duplicate inner index %d at outer index %d", vector.ErrInvalidProviderShape, inner.Index, outer.Index) } innerSeen[inner.Index] = true if len(inner.Embedding) != c.cfg.Dimension { - return nil, fmt.Errorf("embed: Voyage dimension mismatch at outer index %d, inner index %d: got %d, configured %d", outer.Index, inner.Index, len(inner.Embedding), c.cfg.Dimension) + return nil, fmt.Errorf("%w: Voyage dimension mismatch at outer index %d, inner index %d: got %d, configured %d", vector.ErrInvalidProviderVector, outer.Index, inner.Index, len(inner.Embedding), c.cfg.Dimension) } vectors[inner.Index] = inner.Embedding } @@ -263,7 +268,7 @@ func (c *VoyageClient) decodeVoyageResponse(response voyageResponse, inputs [][] } for i, seen := range outerSeen { if !seen { - return nil, fmt.Errorf("embed: Voyage missing outer index %d", i) + return nil, fmt.Errorf("%w: Voyage missing outer index %d", vector.ErrInvalidProviderShape, i) } } return results, nil diff --git a/internal/vector/errors.go b/internal/vector/errors.go index 91ba27db8..dc15d87ed 100644 --- a/internal/vector/errors.go +++ b/internal/vector/errors.go @@ -5,6 +5,16 @@ import "errors" // Sentinel errors used across the vector package. Callers should use // errors.Is to check for these. var ( + // ErrPermanent4xx marks a non-retryable HTTP 4xx response from an + // embeddings provider. Rate limits and server errors remain transient. + ErrPermanent4xx = errors.New("embed: non-retryable 4xx response") + // ErrInvalidProviderShape identifies a complete provider response whose + // document, chunk, or index layout cannot match the request. + ErrInvalidProviderShape = errors.New("invalid embedding provider response shape") + // ErrInvalidProviderVector identifies a provider vector that cannot belong + // to the configured vector space. + ErrInvalidProviderVector = errors.New("invalid embedding provider vector") + // ErrNotEnabled is returned when vector search is requested but // [vector] is not configured. ErrNotEnabled = errors.New("vector search not enabled") diff --git a/internal/vector/pgvector/document_vectors.go b/internal/vector/pgvector/document_vectors.go new file mode 100644 index 000000000..fb012bbc3 --- /dev/null +++ b/internal/vector/pgvector/document_vectors.go @@ -0,0 +1,274 @@ +//go:build pgvector + +package pgvector + +import ( + "context" + "database/sql" + "errors" + "fmt" + "math" + "strings" + "unicode/utf8" + + "go.kenn.io/msgvault/internal/vector" + "go.kenn.io/msgvault/internal/vector/document" +) + +const ( + pgDocumentVectorBatchLimit = 1000 + pgDocumentVectorTokenLimit = 1024 +) + +// DocumentBackend is the independent attachment-document vector store. It +// borrows its parent's pool and does not own or close it. +type DocumentBackend struct { + db *sql.DB +} + +var _ document.Backend = (*DocumentBackend)(nil) + +// DocumentBackend returns a non-owning document-vector view of this backend. +func (b *Backend) DocumentBackend() *DocumentBackend { + return &DocumentBackend{db: b.db} +} + +func (b *DocumentBackend) PutUnpublished(ctx context.Context, generationID document.GenerationID, dimension int, embeddings []document.Embedding) error { + if err := validatePGDocumentPut(generationID, dimension, embeddings); err != nil { + return err + } + if len(embeddings) == 0 { + return nil + } + if err := EnsureDocumentVectorIndex(ctx, b.db, dimension); err != nil { + return err + } + tx, err := b.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin document vector put: %w", err) + } + defer func() { _ = tx.Rollback() }() + + // Establish a lockable generation row before inspecting its dimension. + // For the first concurrent writers, PostgreSQL's unique-index conflict + // serializes the INSERTs: the loser waits for the winner, then locks and + // observes the committed authority row. The authority insert and all token + // writes share this transaction, so a later batch failure rolls everything + // back and lets the next writer establish the generation cleanly. + if _, err := tx.ExecContext(ctx, ` + INSERT INTO document_vector_backend_generations (generation_id, dimension) + VALUES ($1, $2) + ON CONFLICT (generation_id) DO NOTHING`, int64(generationID), dimension); err != nil { + return fmt.Errorf("establish document vector generation: %w", err) + } + var authoritativeDimension int + if err := tx.QueryRowContext(ctx, ` + SELECT dimension FROM document_vector_backend_generations + WHERE generation_id = $1 FOR UPDATE`, int64(generationID)).Scan(&authoritativeDimension); err != nil { + return fmt.Errorf("lock document vector generation: %w", err) + } + if authoritativeDimension != dimension { + return fmt.Errorf("%w: generation %d already uses dimension %d, got %d", + document.ErrInvalidVector, generationID, authoritativeDimension, dimension) + } + + stmt, err := tx.PrepareContext(ctx, ` + INSERT INTO document_vector_embeddings (token, generation_id, dimension, embedding) + VALUES ($1, $2, $3, $4::vector) + ON CONFLICT (token) DO UPDATE SET + dimension = excluded.dimension, + embedding = excluded.embedding + WHERE document_vector_embeddings.generation_id = excluded.generation_id + RETURNING generation_id`) + if err != nil { + return fmt.Errorf("prepare document vector put: %w", err) + } + defer func() { _ = stmt.Close() }() + for _, embedding := range embeddings { + var storedGeneration int64 + err := stmt.QueryRowContext(ctx, embedding.Token, int64(generationID), dimension, + vectorLiteral(embedding.Vector)).Scan(&storedGeneration) + if errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("%w: token %q belongs to another generation", + document.ErrInvalidVector, embedding.Token) + } + if err != nil { + return fmt.Errorf("put document vector %q: %w", embedding.Token, err) + } + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit document vector put: %w", err) + } + return nil +} + +func (b *DocumentBackend) DeleteTokens(ctx context.Context, generationID document.GenerationID, tokens []string) error { + if err := validatePGDocumentTokens(generationID, tokens); err != nil { + return err + } + if len(tokens) == 0 { + return nil + } + unique := make([]string, 0, len(tokens)) + seen := make(map[string]struct{}, len(tokens)) + for _, token := range tokens { + if _, ok := seen[token]; ok { + continue + } + seen[token] = struct{}{} + unique = append(unique, token) + } + tx, err := b.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin document vector delete: %w", err) + } + defer func() { _ = tx.Rollback() }() + if _, err := tx.ExecContext(ctx, ` + DELETE FROM document_vector_embeddings + WHERE generation_id = $1 AND token = ANY($2::text[])`, int64(generationID), textArray(unique)); err != nil { + return fmt.Errorf("delete document vectors: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit document vector delete: %w", err) + } + return nil +} + +func (b *DocumentBackend) Search(ctx context.Context, generationID document.GenerationID, dimension int, query []float32, k int) ([]document.Hit, error) { + page, err := b.SearchPage(ctx, generationID, dimension, query, "", k) + return page.Hits, err +} + +func (b *DocumentBackend) SearchPage(ctx context.Context, generationID document.GenerationID, dimension int, query []float32, cursor string, k int) (document.HitPage, error) { + if err := validatePGDocumentSearch(generationID, dimension, query, k); err != nil { + return document.HitPage{}, err + } + afterDistance, afterToken, afterRank, err := document.DecodePageCursor(cursor) + if err != nil { + return document.HitPage{}, err + } + // Scope is resolved after this backend page, so Exhausted must mean the + // complete generation was consumed. Materialize every generation distance + // before sorting so PostgreSQL cannot substitute the approximate HNSW order. + pagePredicate := "" + limitPlaceholder := "$3" + args := []any{vectorLiteral(query), int64(generationID), k + 1} + if cursor != "" { + pagePredicate = "WHERE distance > $3 OR (distance = $3 AND token > $4)" + limitPlaceholder = "$5" + args = []any{vectorLiteral(query), int64(generationID), afterDistance, afterToken, k + 1} + } + stmt := fmt.Sprintf(` + WITH exact AS MATERIALIZED ( + SELECT token, (embedding::vector(%[1]d)) <=> $1::vector AS distance + FROM document_vector_embeddings + WHERE generation_id = $2 AND dimension = %[1]d + ) + SELECT token, 1.0 - distance AS score, distance + FROM exact + %[2]s + ORDER BY distance ASC, token ASC + LIMIT %[3]s`, dimension, pagePredicate, limitPlaceholder) + rows, err := b.db.QueryContext(ctx, stmt, args...) + if err != nil { + return document.HitPage{}, fmt.Errorf("search document vectors: %w", err) + } + defer func() { _ = rows.Close() }() + hits := make([]document.Hit, 0, k+1) + distances := make([]float64, 0, k+1) + for rows.Next() { + var hit document.Hit + var distance float64 + if err := rows.Scan(&hit.Token, &hit.Score, &distance); err != nil { + return document.HitPage{}, fmt.Errorf("scan document vector hit: %w", err) + } + hit.Rank = afterRank + len(hits) + 1 + hits = append(hits, hit) + distances = append(distances, distance) + } + if err := rows.Err(); err != nil { + return document.HitPage{}, fmt.Errorf("iterate document vector hits: %w", err) + } + page := document.HitPage{Exhausted: len(hits) <= k} + if len(hits) > k { + page.Hits = hits[:k] + } else { + page.Hits = hits + } + if !page.Exhausted { + page.NextCursor, err = document.EncodePageCursor(distances[k-1], hits[k-1].Token, hits[k-1].Rank) + if err != nil { + return document.HitPage{}, err + } + } + return page, nil +} + +func validatePGDocumentPut(generationID document.GenerationID, dimension int, embeddings []document.Embedding) error { + if generationID <= 0 || dimension <= 0 || len(embeddings) > pgDocumentVectorBatchLimit { + return fmt.Errorf("%w: generation, dimension, or batch bound", document.ErrInvalidVector) + } + seen := make(map[string]struct{}, len(embeddings)) + for i, embedding := range embeddings { + if err := validatePGDocumentToken(embedding.Token); err != nil { + return fmt.Errorf("embedding %d: %w", i, err) + } + if _, ok := seen[embedding.Token]; ok { + return fmt.Errorf("%w: duplicate token %q", document.ErrInvalidVector, embedding.Token) + } + seen[embedding.Token] = struct{}{} + if len(embedding.Vector) != dimension { + return fmt.Errorf("%w: token %q has %d dimensions, want %d", + vector.ErrDimensionMismatch, embedding.Token, len(embedding.Vector), dimension) + } + if err := validatePGDocumentVector(embedding.Vector); err != nil { + return fmt.Errorf("token %q: %w", embedding.Token, err) + } + } + return nil +} + +func validatePGDocumentTokens(generationID document.GenerationID, tokens []string) error { + if generationID <= 0 || len(tokens) > pgDocumentVectorBatchLimit { + return fmt.Errorf("%w: generation or token batch bound", document.ErrInvalidVector) + } + for _, token := range tokens { + if err := validatePGDocumentToken(token); err != nil { + return err + } + } + return nil +} + +func validatePGDocumentSearch(generationID document.GenerationID, dimension int, query []float32, k int) error { + if generationID <= 0 || dimension <= 0 || k <= 0 || k > pgDocumentVectorBatchLimit { + return fmt.Errorf("%w: generation, dimension, or result bound", document.ErrInvalidVector) + } + if len(query) != dimension { + return fmt.Errorf("%w: query has %d dimensions, want %d", vector.ErrDimensionMismatch, len(query), dimension) + } + return validatePGDocumentVector(query) +} + +func validatePGDocumentToken(token string) error { + if token == "" || len(token) > pgDocumentVectorTokenLimit || + !utf8.ValidString(token) || strings.ContainsRune(token, 0) { + return fmt.Errorf("%w: token must contain 1..%d bytes", document.ErrInvalidVector, pgDocumentVectorTokenLimit) + } + return nil +} + +func validatePGDocumentVector(values []float32) error { + var norm float64 + for _, value := range values { + f := float64(value) + if math.IsNaN(f) || math.IsInf(f, 0) { + return fmt.Errorf("%w: non-finite component", document.ErrInvalidVector) + } + norm += f * f + } + if norm == 0 { + return fmt.Errorf("%w: zero-norm vector", document.ErrInvalidVector) + } + return nil +} diff --git a/internal/vector/pgvector/document_vectors_stub.go b/internal/vector/pgvector/document_vectors_stub.go new file mode 100644 index 000000000..21b69509b --- /dev/null +++ b/internal/vector/pgvector/document_vectors_stub.go @@ -0,0 +1,28 @@ +//go:build !pgvector + +package pgvector + +import ( + "context" + + "go.kenn.io/msgvault/internal/vector/document" +) + +// DocumentBackend is unavailable without pgvector support. +type DocumentBackend struct{} + +var _ document.Backend = (*DocumentBackend)(nil) + +func (b *Backend) DocumentBackend() *DocumentBackend { return &DocumentBackend{} } + +func (b *DocumentBackend) PutUnpublished(context.Context, document.GenerationID, int, []document.Embedding) error { + return ErrNotBuilt +} + +func (b *DocumentBackend) DeleteTokens(context.Context, document.GenerationID, []string) error { + return ErrNotBuilt +} + +func (b *DocumentBackend) Search(context.Context, document.GenerationID, int, []float32, int) ([]document.Hit, error) { + return nil, ErrNotBuilt +} diff --git a/internal/vector/pgvector/document_vectors_test.go b/internal/vector/pgvector/document_vectors_test.go new file mode 100644 index 000000000..e147ecca0 --- /dev/null +++ b/internal/vector/pgvector/document_vectors_test.go @@ -0,0 +1,327 @@ +//go:build pgvector + +package pgvector + +import ( + "context" + "fmt" + "math" + "strings" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/msgvault/internal/vector" + "go.kenn.io/msgvault/internal/vector/document" +) + +func TestDocumentBackendConcurrentFirstPutSerializesGenerationDimensionPostgres(t *testing.T) { + db := openPGTestDB(t) + parent, err := Open(t.Context(), Options{DB: db}) + require.NoError(t, err) + backend := parent.DocumentBackend() + require.NoError(t, EnsureDocumentVectorIndex(t.Context(), db, 3)) + require.NoError(t, EnsureDocumentVectorIndex(t.Context(), db, 4)) + + for round := range 12 { + generationID := document.GenerationID(100 + round) + start := make(chan struct{}) + type result struct { + dimension int + err error + } + results := make(chan result, 2) + var ready sync.WaitGroup + ready.Add(2) + put := func(dimension int, embedding document.Embedding) { + ready.Done() + <-start + results <- result{dimension: dimension, err: backend.PutUnpublished( + t.Context(), generationID, dimension, []document.Embedding{embedding})} + } + go put(3, document.Embedding{Token: "concurrent-3-" + string(rune('a'+round)), Vector: []float32{1, 0, 0}}) + go put(4, document.Embedding{Token: "concurrent-4-" + string(rune('a'+round)), Vector: []float32{1, 0, 0, 0}}) + ready.Wait() + close(start) + + first, second := <-results, <-results + outcomes := []result{first, second} + winnerDimension := 0 + for _, outcome := range outcomes { + if outcome.err == nil { + require.Zero(t, winnerDimension, "round %d admitted two dimensions", round) + winnerDimension = outcome.dimension + continue + } + require.ErrorIs(t, outcome.err, document.ErrInvalidVector, "round %d loser", round) + } + require.NotZero(t, winnerDimension, "round %d admitted no writer", round) + + var authorityDimension, rowCount, minDimension, maxDimension int + require.NoError(t, db.QueryRow(` + SELECT dimension FROM document_vector_backend_generations WHERE generation_id = $1`, + int64(generationID)).Scan(&authorityDimension)) + require.NoError(t, db.QueryRow(` + SELECT COUNT(*), MIN(dimension), MAX(dimension) + FROM document_vector_embeddings WHERE generation_id = $1`, + int64(generationID)).Scan(&rowCount, &minDimension, &maxDimension)) + assert.Equal(t, winnerDimension, authorityDimension) + assert.Equal(t, 1, rowCount) + assert.Equal(t, winnerDimension, minDimension) + assert.Equal(t, winnerDimension, maxDimension) + } +} + +func TestDocumentBackendGenerationAuthorityRollsBackWithBatchPostgres(t *testing.T) { + db := openPGTestDB(t) + parent, err := Open(t.Context(), Options{DB: db}) + require.NoError(t, err) + backend := parent.DocumentBackend() + require.NoError(t, backend.PutUnpublished(t.Context(), 801, 3, []document.Embedding{ + {Token: "authority-owner", Vector: []float32{1, 0, 0}}, + })) + + err = backend.PutUnpublished(t.Context(), 802, 3, []document.Embedding{ + {Token: "authority-partial", Vector: []float32{0, 1, 0}}, + {Token: "authority-owner", Vector: []float32{0, 0, 1}}, + }) + require.ErrorIs(t, err, document.ErrInvalidVector) + var generations, embeddings int + require.NoError(t, db.QueryRow(` + SELECT COUNT(*) FROM document_vector_backend_generations WHERE generation_id = 802`).Scan(&generations)) + require.NoError(t, db.QueryRow(` + SELECT COUNT(*) FROM document_vector_embeddings WHERE generation_id = 802`).Scan(&embeddings)) + assert.Zero(t, generations, "failed token batch rolls back generation authority") + assert.Zero(t, embeddings, "failed token batch rolls back embeddings") + + require.NoError(t, backend.PutUnpublished(t.Context(), 802, 4, []document.Embedding{ + {Token: "authority-after-rollback", Vector: []float32{1, 0, 0, 0}}, + }), "a later dimension may establish authority after the failed transaction") +} + +func TestDocumentBackendSearchPageReadsBeyondHNSWBudgetPostgres(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + db := openPGTestDB(t) + db.SetMaxOpenConns(1) + parent, err := Open(t.Context(), Options{DB: db}) + requirements.NoError(err) + backend := parent.DocumentBackend() + + const total = 1105 + embeddings := make([]document.Embedding, total) + for index := range embeddings { + embeddings[index] = document.Embedding{ + Token: fmt.Sprintf("%064x", index), Vector: []float32{1, float32(index) / total, 0}, + } + } + requirements.NoError(backend.PutUnpublished(t.Context(), 901, 3, embeddings[:1000])) + requirements.NoError(backend.PutUnpublished(t.Context(), 901, 3, embeddings[1000:])) + _, err = db.ExecContext(t.Context(), `SET enable_seqscan = off`) + requirements.NoError(err) + _, err = db.ExecContext(t.Context(), `SET hnsw.ef_search = 40`) + requirements.NoError(err) + + firstPage, err := backend.SearchPage(t.Context(), 901, 3, []float32{1, 0, 0}, "", 1000) + requirements.NoError(err) + requirements.False(firstPage.Exhausted) + page, err := backend.SearchPage(t.Context(), 901, 3, []float32{1, 0, 0}, firstPage.NextCursor, 100) + + requirements.NoError(err) + requirements.Len(page.Hits, 100) + assertions.False(page.Exhausted) + assertions.Equal(fmt.Sprintf("%064x", 1000), page.Hits[0].Token) + assertions.Equal(1001, page.Hits[0].Rank) +} + +func TestDocumentBackendContractPostgres(t *testing.T) { + db := openPGTestDB(t) + parent, err := Open(context.Background(), Options{DB: db, Dimension: 3}) + require.NoError(t, err) + backend := parent.DocumentBackend() + testPGDocumentBackendContract(t, backend, func(token string) int { + t.Helper() + var count int + require.NoError(t, db.QueryRow( + `SELECT COUNT(*) FROM document_vector_embeddings WHERE token = $1`, token).Scan(&count)) + return count + }) + + var definition string + require.NoError(t, db.QueryRow(`SELECT pg_get_indexdef(indexrelid) FROM pg_index WHERE indexrelid = $1::regclass`, + DocumentVectorIndexName(3)).Scan(&definition)) + assert.Contains(t, definition, "vector_cosine_ops") + assert.Contains(t, definition, "WHERE (dimension = 3)") +} + +func testPGDocumentBackendContract(t *testing.T, backend document.Backend, tokenCount func(string) int) { + t.Helper() + ctx := context.Background() + gen1 := document.GenerationID(11) + gen2 := document.GenerationID(12) + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, nil)) + require.NoError(t, backend.DeleteTokens(ctx, gen1, nil)) + empty, err := backend.Search(ctx, 99, 5, []float32{1, 0, 0, 0, 0}, 1) + require.NoError(t, err) + assert.Empty(t, empty) + embeddings := []document.Embedding{ + {Token: "token-d", Vector: []float32{-1, 0, 0}}, + {Token: "token-c", Vector: []float32{0, -1, 0}}, + {Token: "token-a", Vector: []float32{1, 0, 0}}, + {Token: "token-b", Vector: []float32{0, 1, 0}}, + } + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, embeddings)) + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, embeddings)) + canceledCtx, cancel := context.WithCancel(ctx) + cancel() + err = backend.PutUnpublished(canceledCtx, gen1, 3, []document.Embedding{ + {Token: "canceled-first", Vector: []float32{1, 0, 0}}, + {Token: "canceled-second", Vector: []float32{0, 1, 0}}, + }) + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, tokenCount("canceled-first")) + assert.Zero(t, tokenCount("canceled-second"), "a canceled put exposes no partial batch") + + hits, err := backend.Search(ctx, gen1, 3, []float32{1, 0, 0}, 4) + require.NoError(t, err) + require.Len(t, hits, 4) + assert.Equal(t, []string{"token-a", "token-b", "token-c", "token-d"}, pgHitTokens(hits)) + assert.Equal(t, []int{1, 2, 3, 4}, pgHitRanks(hits)) + assert.InDelta(t, 1, hits[0].Score, 1e-6) + assert.InDelta(t, 0, hits[1].Score, 1e-6) + assert.InDelta(t, 0, hits[2].Score, 1e-6) + assert.InDelta(t, -1, hits[3].Score, 1e-6) + paged, ok := backend.(document.PagedBackend) + require.True(t, ok) + firstPage, err := paged.SearchPage(ctx, gen1, 3, []float32{1, 0, 0}, "", 2) + require.NoError(t, err) + assert.False(t, firstPage.Exhausted) + assert.Equal(t, []string{"token-a", "token-b"}, pgHitTokens(firstPage.Hits)) + require.NoError(t, backend.DeleteTokens(ctx, gen1, []string{"token-a"})) + secondPage, err := paged.SearchPage(ctx, gen1, 3, []float32{1, 0, 0}, firstPage.NextCursor, 2) + require.NoError(t, err) + assert.True(t, secondPage.Exhausted) + assert.Equal(t, []string{"token-c", "token-d"}, pgHitTokens(secondPage.Hits)) + assert.Equal(t, []int{3, 4}, pgHitRanks(secondPage.Hits)) + + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, []document.Embedding{ + {Token: "token-a", Vector: []float32{0, 0, 1}}, + })) + hits, err = backend.Search(ctx, gen1, 3, []float32{0, 0, 1}, 1) + require.NoError(t, err) + require.Len(t, hits, 1) + assert.Equal(t, "token-a", hits[0].Token) + + require.NoError(t, backend.PutUnpublished(ctx, gen2, 4, []document.Embedding{ + {Token: "token-four", Vector: []float32{0, 0, 0, 1}}, + })) + hits, err = backend.Search(ctx, gen2, 4, []float32{0, 0, 0, 1}, 10) + require.NoError(t, err) + require.Len(t, hits, 1) + assert.Equal(t, "token-four", hits[0].Token) + hits, err = backend.Search(ctx, gen1, 3, []float32{0, 0, 1}, 10) + require.NoError(t, err) + assert.NotContains(t, pgHitTokens(hits), "token-four") + + err = backend.PutUnpublished(ctx, gen2, 4, []document.Embedding{ + {Token: "token-new", Vector: []float32{1, 0, 0, 0}}, + {Token: "token-a", Vector: []float32{0, 1, 0, 0}}, + }) + require.ErrorIs(t, err, document.ErrInvalidVector) + assert.Zero(t, tokenCount("token-new")) + + err = backend.PutUnpublished(ctx, gen1, 3, []document.Embedding{ + {Token: "token-partial", Vector: []float32{1, 0, 0}}, + {Token: "token-wrong", Vector: []float32{1, 0}}, + }) + require.ErrorIs(t, err, vector.ErrDimensionMismatch) + assert.Zero(t, tokenCount("token-partial")) + err = backend.PutUnpublished(ctx, gen1, 3, []document.Embedding{ + {Token: "token-partial-nonfinite", Vector: []float32{1, 0, 0}}, + {Token: "token-nonfinite", Vector: []float32{1, float32(math.NaN()), 0}}, + }) + require.ErrorIs(t, err, document.ErrInvalidVector) + assert.Zero(t, tokenCount("token-partial-nonfinite")) + + for name, batch := range map[string][]document.Embedding{ + "duplicate": { + {Token: "duplicate", Vector: []float32{1, 0, 0}}, + {Token: "duplicate", Vector: []float32{0, 1, 0}}, + }, + "empty token": {{Token: "", Vector: []float32{1, 0, 0}}}, + "long token": {{Token: strings.Repeat("x", 1025), Vector: []float32{1, 0, 0}}}, + "nul token": {{Token: "x\x00y", Vector: []float32{1, 0, 0}}}, + "nonfinite": {{Token: "nan", Vector: []float32{1, float32(math.NaN()), 0}}}, + "infinite": {{Token: "inf", Vector: []float32{1, float32(math.Inf(1)), 0}}}, + "zero norm": {{Token: "zero", Vector: []float32{0, 0, 0}}}, + } { + t.Run("put rejects "+name, func(t *testing.T) { + require.ErrorIs(t, backend.PutUnpublished(ctx, gen1, 3, batch), document.ErrInvalidVector) + }) + } + + tooMany := make([]document.Embedding, 1001) + for i := range tooMany { + tooMany[i] = document.Embedding{Token: "bounded-" + string(rune(i+1)), Vector: []float32{1, 0, 0}} + } + for name, call := range map[string]func() error{ + "put generation": func() error { return backend.PutUnpublished(ctx, 0, 3, embeddings[:1]) }, + "put dimension": func() error { return backend.PutUnpublished(ctx, gen1, 0, embeddings[:1]) }, + "put batch": func() error { return backend.PutUnpublished(ctx, gen1, 3, tooMany) }, + "delete gen": func() error { return backend.DeleteTokens(ctx, 0, []string{"token-a"}) }, + "delete batch": func() error { return backend.DeleteTokens(ctx, gen1, make([]string, 1001)) }, + "search gen": func() error { + _, err := backend.Search(ctx, 0, 3, []float32{1, 0, 0}, 1) + return err + }, + "search dimension": func() error { + _, err := backend.Search(ctx, gen1, 0, []float32{1, 0, 0}, 1) + return err + }, + "search k zero": func() error { + _, err := backend.Search(ctx, gen1, 3, []float32{1, 0, 0}, 0) + return err + }, + "search k large": func() error { + _, err := backend.Search(ctx, gen1, 3, []float32{1, 0, 0}, 1001) + return err + }, + "search nonfinite": func() error { + _, err := backend.Search(ctx, gen1, 3, []float32{1, float32(math.Inf(-1)), 0}, 1) + return err + }, + } { + t.Run("bounds reject "+name, func(t *testing.T) { + require.ErrorIs(t, call(), document.ErrInvalidVector) + }) + } + _, err = backend.Search(ctx, gen1, 3, []float32{1, 0}, 1) + require.ErrorIs(t, err, vector.ErrDimensionMismatch) + + require.NoError(t, backend.DeleteTokens(ctx, gen2, []string{"token-a"})) + assert.Equal(t, 1, tokenCount("token-a")) + require.NoError(t, backend.DeleteTokens(ctx, gen1, []string{"token-a"})) + require.NoError(t, backend.DeleteTokens(ctx, gen1, []string{"token-a"})) + assert.Zero(t, tokenCount("token-a")) + hits, err = backend.Search(ctx, gen1, 3, []float32{0, 0, 1}, 10) + require.NoError(t, err) + assert.NotContains(t, pgHitTokens(hits), "token-a") +} + +func pgHitTokens(hits []document.Hit) []string { + out := make([]string, len(hits)) + for i := range hits { + out[i] = hits[i].Token + } + return out +} + +func pgHitRanks(hits []document.Hit) []int { + out := make([]int, len(hits)) + for i := range hits { + out[i] = hits[i].Rank + } + return out +} diff --git a/internal/vector/pgvector/migrate.go b/internal/vector/pgvector/migrate.go index c8fdc6093..3bb4304af 100644 --- a/internal/vector/pgvector/migrate.go +++ b/internal/vector/pgvector/migrate.go @@ -128,6 +128,9 @@ func Migrate(ctx context.Context, db migrateExecer, defaultDim int, skipExtensio if err := EnsurePersonVectorIndex(ctx, db, defaultDim); err != nil { return err } + if err := EnsureDocumentVectorIndex(ctx, db, defaultDim); err != nil { + return err + } } return nil } @@ -183,3 +186,38 @@ func EnsureVectorIndex(ctx context.Context, db migrateExecer, dim int) error { func VectorIndexName(dim int) string { return fmt.Sprintf("idx_embeddings_hnsw_d%d", dim) } + +// EnsureDocumentVectorIndex creates the dedicated per-dimension cosine HNSW +// index for attachment-document vectors. It never indexes message embeddings. +func EnsureDocumentVectorIndex(ctx context.Context, db migrateExecer, dim int) error { + if dim <= 0 { + return fmt.Errorf("invalid document vector dimension %d", dim) + } + stmt := fmt.Sprintf( + `CREATE INDEX IF NOT EXISTS %s + ON document_vector_embeddings + USING hnsw ((embedding::vector(%d)) vector_cosine_ops) + WHERE dimension = %d`, + DocumentVectorIndexName(dim), dim, dim, + ) + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin document hnsw index tx for dim %d: %w", dim, err) + } + defer func() { _ = tx.Rollback() }() + if _, err := tx.ExecContext(ctx, "SET LOCAL statement_timeout = 0"); err != nil { + return fmt.Errorf("disable statement_timeout for document hnsw index: %w", err) + } + if _, err := tx.ExecContext(ctx, stmt); err != nil { + return fmt.Errorf("create document hnsw index for dim %d: %w", dim, err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit document hnsw index for dim %d: %w", dim, err) + } + return nil +} + +// DocumentVectorIndexName returns the dedicated dimension-specific HNSW name. +func DocumentVectorIndexName(dim int) string { + return fmt.Sprintf("idx_document_vector_embeddings_hnsw_d%d", dim) +} diff --git a/internal/vector/pgvector/schema.sql b/internal/vector/pgvector/schema.sql index 26545e370..70d55fe89 100644 --- a/internal/vector/pgvector/schema.sql +++ b/internal/vector/pgvector/schema.sql @@ -144,7 +144,6 @@ CREATE TABLE IF NOT EXISTS embedding_document_progress ( reconcile_cursor TEXT NOT NULL DEFAULT '', journal_cursor TEXT NOT NULL DEFAULT '' ); - -- Visual vectors share the PostgreSQL database with their authoritative -- publication rows but remain independently keyed by an opaque publication -- token. A prepared vector cannot be searched until visual_publications points @@ -159,3 +158,24 @@ CREATE INDEX IF NOT EXISTS idx_visual_vectors_hnsw_d1024 ON visual_vectors USING hnsw ((embedding::vector(1024)) vector_cosine_ops) WHERE dimension = 1024; + +-- Independent attachment-document vectors. Tokens are opaque globally unique +-- publication identities; generation IDs intentionally do not reference the +-- message-vector index_generations table. +CREATE TABLE IF NOT EXISTS document_vector_backend_generations ( + generation_id BIGINT PRIMARY KEY, + dimension INTEGER NOT NULL CHECK (dimension > 0), + UNIQUE (generation_id, dimension) +); + +CREATE TABLE IF NOT EXISTS document_vector_embeddings ( + token TEXT PRIMARY KEY, + generation_id BIGINT NOT NULL, + dimension INTEGER NOT NULL CHECK (dimension > 0), + embedding vector NOT NULL, + CONSTRAINT document_vector_embeddings_generation_dimension_fkey + FOREIGN KEY (generation_id, dimension) + REFERENCES document_vector_backend_generations(generation_id, dimension) +); +CREATE INDEX IF NOT EXISTS idx_document_vector_embeddings_generation + ON document_vector_embeddings(generation_id, dimension, token); diff --git a/internal/vector/sqlitevec/document_vectors.go b/internal/vector/sqlitevec/document_vectors.go new file mode 100644 index 000000000..74cdf11a0 --- /dev/null +++ b/internal/vector/sqlitevec/document_vectors.go @@ -0,0 +1,294 @@ +//go:build sqlite_vec + +package sqlitevec + +import ( + "context" + "database/sql" + "errors" + "fmt" + "math" + "strings" + "unicode/utf8" + + "go.kenn.io/msgvault/internal/vector" + "go.kenn.io/msgvault/internal/vector/document" +) + +const ( + documentVectorBatchLimit = 1000 + documentVectorTokenLimit = 1024 +) + +// DocumentBackend is the independent attachment-document vector store. It +// borrows its parent's connection and does not own or close it. +type DocumentBackend struct { + db *sql.DB +} + +var _ document.Backend = (*DocumentBackend)(nil) + +// DocumentBackend returns a non-owning document-vector view of this backend. +func (b *Backend) DocumentBackend() *DocumentBackend { + return &DocumentBackend{db: b.db} +} + +func (b *DocumentBackend) PutUnpublished(ctx context.Context, generationID document.GenerationID, dimension int, embeddings []document.Embedding) error { + if err := validateDocumentPut(generationID, dimension, embeddings); err != nil { + return err + } + if len(embeddings) == 0 { + return nil + } + if err := EnsureDocumentVectorTable(ctx, b.db, dimension); err != nil { + return err + } + + tx, err := b.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin document vector put: %w", err) + } + defer func() { _ = tx.Rollback() }() + + var existingDimension int + err = tx.QueryRowContext(ctx, + `SELECT dimension FROM document_vector_embeddings WHERE generation_id = ? LIMIT 1`, int64(generationID)).Scan(&existingDimension) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("lookup document generation dimension: %w", err) + } + if err == nil && existingDimension != dimension { + return fmt.Errorf("%w: generation %d already uses dimension %d, got %d", + document.ErrInvalidVector, generationID, existingDimension, dimension) + } + + vecTable := DocumentVectorTableName(dimension) + for _, embedding := range embeddings { + var rowID int64 + var existingGeneration int64 + err := tx.QueryRowContext(ctx, + `SELECT document_vector_id, generation_id FROM document_vector_embeddings WHERE token = ?`, embedding.Token). + Scan(&rowID, &existingGeneration) + switch { + case errors.Is(err, sql.ErrNoRows): + err = tx.QueryRowContext(ctx, ` + INSERT INTO document_vector_embeddings (token, generation_id, dimension) + VALUES (?, ?, ?) RETURNING document_vector_id`, + embedding.Token, int64(generationID), dimension).Scan(&rowID) + if err != nil { + return fmt.Errorf("insert document vector metadata: %w", err) + } + case err != nil: + return fmt.Errorf("lookup document vector token: %w", err) + case existingGeneration != int64(generationID): + return fmt.Errorf("%w: token %q belongs to generation %d", + document.ErrInvalidVector, embedding.Token, existingGeneration) + default: + if _, err := tx.ExecContext(ctx, + fmt.Sprintf(`DELETE FROM %s WHERE generation_id = ? AND document_vector_id = ?`, vecTable), + int64(generationID), rowID); err != nil { + return fmt.Errorf("delete replaced document vector: %w", err) + } + } + if _, err := tx.ExecContext(ctx, fmt.Sprintf(` + INSERT INTO %s (generation_id, document_vector_id, embedding) + VALUES (?, ?, ?)`, vecTable), int64(generationID), rowID, float32SliceBlob(embedding.Vector)); err != nil { + return fmt.Errorf("insert document vector %q: %w", embedding.Token, err) + } + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit document vector put: %w", err) + } + return nil +} + +func (b *DocumentBackend) DeleteTokens(ctx context.Context, generationID document.GenerationID, tokens []string) error { + if err := validateDocumentTokens(generationID, tokens); err != nil { + return err + } + if len(tokens) == 0 { + return nil + } + tx, err := b.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin document vector delete: %w", err) + } + defer func() { _ = tx.Rollback() }() + seen := make(map[string]struct{}, len(tokens)) + for _, token := range tokens { + if _, ok := seen[token]; ok { + continue + } + seen[token] = struct{}{} + var rowID int64 + var dimension int + err := tx.QueryRowContext(ctx, ` + SELECT document_vector_id, dimension FROM document_vector_embeddings + WHERE generation_id = ? AND token = ?`, int64(generationID), token).Scan(&rowID, &dimension) + if errors.Is(err, sql.ErrNoRows) { + continue + } + if err != nil { + return fmt.Errorf("lookup document vector for delete: %w", err) + } + if _, err := tx.ExecContext(ctx, fmt.Sprintf(` + DELETE FROM %s WHERE generation_id = ? AND document_vector_id = ?`, DocumentVectorTableName(dimension)), + int64(generationID), rowID); err != nil { + return fmt.Errorf("delete document vector %q: %w", token, err) + } + if _, err := tx.ExecContext(ctx, ` + DELETE FROM document_vector_embeddings + WHERE generation_id = ? AND token = ?`, int64(generationID), token); err != nil { + return fmt.Errorf("delete document vector metadata %q: %w", token, err) + } + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit document vector delete: %w", err) + } + return nil +} + +func (b *DocumentBackend) Search(ctx context.Context, generationID document.GenerationID, dimension int, query []float32, k int) ([]document.Hit, error) { + page, err := b.SearchPage(ctx, generationID, dimension, query, "", k) + return page.Hits, err +} + +func (b *DocumentBackend) SearchPage(ctx context.Context, generationID document.GenerationID, dimension int, query []float32, cursor string, k int) (document.HitPage, error) { + if err := validateDocumentSearch(generationID, dimension, query, k); err != nil { + return document.HitPage{}, err + } + afterDistance, afterToken, afterRank, err := document.DecodePageCursor(cursor) + if err != nil { + return document.HitPage{}, err + } + var exists int + err = b.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?`, DocumentVectorTableName(dimension)).Scan(&exists) + if err != nil { + return document.HitPage{}, fmt.Errorf("check document vector table: %w", err) + } + if exists == 0 { + return document.HitPage{Exhausted: true}, nil + } + pagePredicate := "" + args := []any{float32SliceBlob(query), int64(generationID), int64(generationID), dimension} + if cursor != "" { + pagePredicate = "WHERE distance > ? OR (distance = ? AND token > ?)" + args = append(args, afterDistance, afterDistance, afterToken) + } + args = append(args, k+1) + q := fmt.Sprintf(` + WITH exact AS ( + SELECT m.token, vec_distance_cosine(v.embedding, ?) AS distance + FROM %s v + JOIN document_vector_embeddings m + ON m.document_vector_id = v.document_vector_id + WHERE v.generation_id = ? AND m.generation_id = ? AND m.dimension = ? + ) + SELECT token, 1.0 - distance AS score, distance + FROM exact + %s + ORDER BY distance ASC, token ASC + LIMIT ?`, DocumentVectorTableName(dimension), pagePredicate) + rows, err := b.db.QueryContext(ctx, q, args...) + if err != nil { + return document.HitPage{}, fmt.Errorf("search document vectors: %w", err) + } + defer func() { _ = rows.Close() }() + hits := make([]document.Hit, 0, k+1) + distances := make([]float64, 0, k+1) + for rows.Next() { + var hit document.Hit + var distance float64 + if err := rows.Scan(&hit.Token, &hit.Score, &distance); err != nil { + return document.HitPage{}, fmt.Errorf("scan document vector hit: %w", err) + } + hit.Rank = afterRank + len(hits) + 1 + hits = append(hits, hit) + distances = append(distances, distance) + } + if err := rows.Err(); err != nil { + return document.HitPage{}, fmt.Errorf("iterate document vector hits: %w", err) + } + page := document.HitPage{Exhausted: len(hits) <= k} + if len(hits) > k { + page.Hits = hits[:k] + } else { + page.Hits = hits + } + if !page.Exhausted { + page.NextCursor, err = document.EncodePageCursor(distances[k-1], hits[k-1].Token, hits[k-1].Rank) + if err != nil { + return document.HitPage{}, err + } + } + return page, nil +} + +func validateDocumentPut(generationID document.GenerationID, dimension int, embeddings []document.Embedding) error { + if generationID <= 0 || dimension <= 0 || len(embeddings) > documentVectorBatchLimit { + return fmt.Errorf("%w: generation, dimension, or batch bound", document.ErrInvalidVector) + } + seen := make(map[string]struct{}, len(embeddings)) + for i, embedding := range embeddings { + if err := validateDocumentToken(embedding.Token); err != nil { + return fmt.Errorf("embedding %d: %w", i, err) + } + if _, ok := seen[embedding.Token]; ok { + return fmt.Errorf("%w: duplicate token %q", document.ErrInvalidVector, embedding.Token) + } + seen[embedding.Token] = struct{}{} + if len(embedding.Vector) != dimension { + return fmt.Errorf("%w: token %q has %d dimensions, want %d", + vector.ErrDimensionMismatch, embedding.Token, len(embedding.Vector), dimension) + } + if err := validateDocumentVector(embedding.Vector); err != nil { + return fmt.Errorf("token %q: %w", embedding.Token, err) + } + } + return nil +} + +func validateDocumentTokens(generationID document.GenerationID, tokens []string) error { + if generationID <= 0 || len(tokens) > documentVectorBatchLimit { + return fmt.Errorf("%w: generation or token batch bound", document.ErrInvalidVector) + } + for _, token := range tokens { + if err := validateDocumentToken(token); err != nil { + return err + } + } + return nil +} + +func validateDocumentSearch(generationID document.GenerationID, dimension int, query []float32, k int) error { + if generationID <= 0 || dimension <= 0 || k <= 0 || k > documentVectorBatchLimit { + return fmt.Errorf("%w: generation, dimension, or result bound", document.ErrInvalidVector) + } + if len(query) != dimension { + return fmt.Errorf("%w: query has %d dimensions, want %d", vector.ErrDimensionMismatch, len(query), dimension) + } + return validateDocumentVector(query) +} + +func validateDocumentToken(token string) error { + if token == "" || len(token) > documentVectorTokenLimit || + !utf8.ValidString(token) || strings.ContainsRune(token, 0) { + return fmt.Errorf("%w: token must contain 1..%d bytes", document.ErrInvalidVector, documentVectorTokenLimit) + } + return nil +} + +func validateDocumentVector(values []float32) error { + var norm float64 + for _, value := range values { + f := float64(value) + if math.IsNaN(f) || math.IsInf(f, 0) { + return fmt.Errorf("%w: non-finite component", document.ErrInvalidVector) + } + norm += f * f + } + if norm == 0 { + return fmt.Errorf("%w: zero-norm vector", document.ErrInvalidVector) + } + return nil +} diff --git a/internal/vector/sqlitevec/document_vectors_stub.go b/internal/vector/sqlitevec/document_vectors_stub.go new file mode 100644 index 000000000..171dc7c11 --- /dev/null +++ b/internal/vector/sqlitevec/document_vectors_stub.go @@ -0,0 +1,28 @@ +//go:build !sqlite_vec + +package sqlitevec + +import ( + "context" + + "go.kenn.io/msgvault/internal/vector/document" +) + +// DocumentBackend is unavailable without sqlite_vec support. +type DocumentBackend struct{} + +var _ document.Backend = (*DocumentBackend)(nil) + +func (b *Backend) DocumentBackend() *DocumentBackend { return &DocumentBackend{} } + +func (b *DocumentBackend) PutUnpublished(context.Context, document.GenerationID, int, []document.Embedding) error { + return ErrNotBuilt +} + +func (b *DocumentBackend) DeleteTokens(context.Context, document.GenerationID, []string) error { + return ErrNotBuilt +} + +func (b *DocumentBackend) Search(context.Context, document.GenerationID, int, []float32, int) ([]document.Hit, error) { + return nil, ErrNotBuilt +} diff --git a/internal/vector/sqlitevec/document_vectors_test.go b/internal/vector/sqlitevec/document_vectors_test.go new file mode 100644 index 000000000..613127f7f --- /dev/null +++ b/internal/vector/sqlitevec/document_vectors_test.go @@ -0,0 +1,268 @@ +//go:build sqlite_vec + +package sqlitevec + +import ( + "context" + "fmt" + "math" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/msgvault/internal/vector" + "go.kenn.io/msgvault/internal/vector/document" +) + +func openDocumentBackendForTest(t *testing.T) (*Backend, *DocumentBackend) { + t.Helper() + b, err := Open(context.Background(), Options{ + Path: filepath.Join(t.TempDir(), "vectors.db"), + Dimension: 3, + }) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, b.Close()) }) + return b, b.DocumentBackend() +} + +func TestDocumentBackendContractSQLite(t *testing.T) { + parent, backend := openDocumentBackendForTest(t) + testDocumentBackendContract(t, backend, func(token string) int { + t.Helper() + var count int + require.NoError(t, parent.db.QueryRow( + `SELECT COUNT(*) FROM document_vector_embeddings WHERE token = ?`, token).Scan(&count)) + return count + }) + require.NoError(t, backend.PutUnpublished(t.Context(), 21, 3, []document.Embedding{ + {Token: "stable-rowid", Vector: []float32{1, 0, 0}}, + })) + var before, after int64 + require.NoError(t, parent.db.QueryRow( + `SELECT document_vector_id FROM document_vector_embeddings WHERE token = ?`, "stable-rowid").Scan(&before)) + require.NoError(t, backend.PutUnpublished(t.Context(), 21, 3, []document.Embedding{ + {Token: "stable-rowid", Vector: []float32{0, 1, 0}}, + })) + require.NoError(t, parent.db.QueryRow( + `SELECT document_vector_id FROM document_vector_embeddings WHERE token = ?`, "stable-rowid").Scan(&after)) + assert.Equal(t, before, after, "replacement preserves the vec0-linked metadata rowid") +} + +func TestDocumentBackendSearchPagesEqualDistancesWithoutOverlapSQLite(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + _, backend := openDocumentBackendForTest(t) + + const total = 300 + embeddings := make([]document.Embedding, total) + for index := range embeddings { + embeddings[index] = document.Embedding{ + Token: fmt.Sprintf("%064x", index), Vector: []float32{1, 0, 0}, + } + } + requirements.NoError(backend.PutUnpublished(t.Context(), 31, 3, embeddings)) + + first, err := backend.SearchPage(t.Context(), 31, 3, []float32{1, 0, 0}, "", 100) + requirements.NoError(err) + requirements.False(first.Exhausted) + second, err := backend.SearchPage(t.Context(), 31, 3, []float32{1, 0, 0}, first.NextCursor, 100) + requirements.NoError(err) + requirements.False(second.Exhausted) + + got := append(hitTokens(first.Hits), hitTokens(second.Hits)...) + want := make([]string, 200) + for index := range want { + want[index] = fmt.Sprintf("%064x", index) + } + assertions.Equal(want, got) +} + +func testDocumentBackendContract(t *testing.T, backend document.Backend, tokenCount func(string) int) { + t.Helper() + ctx := context.Background() + gen1 := document.GenerationID(11) + gen2 := document.GenerationID(12) + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, nil)) + require.NoError(t, backend.DeleteTokens(ctx, gen1, nil)) + empty, err := backend.Search(ctx, 99, 5, []float32{1, 0, 0, 0, 0}, 1) + require.NoError(t, err) + assert.Empty(t, empty) + embeddings := []document.Embedding{ + {Token: "token-d", Vector: []float32{-1, 0, 0}}, + {Token: "token-c", Vector: []float32{0, -1, 0}}, + {Token: "token-a", Vector: []float32{1, 0, 0}}, + {Token: "token-b", Vector: []float32{0, 1, 0}}, + } + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, embeddings)) + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, embeddings), "same batch is idempotent") + canceledCtx, cancel := context.WithCancel(ctx) + cancel() + err = backend.PutUnpublished(canceledCtx, gen1, 3, []document.Embedding{ + {Token: "canceled-first", Vector: []float32{1, 0, 0}}, + {Token: "canceled-second", Vector: []float32{0, 1, 0}}, + }) + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, tokenCount("canceled-first")) + assert.Zero(t, tokenCount("canceled-second"), "a canceled put exposes no partial batch") + + hits, err := backend.Search(ctx, gen1, 3, []float32{1, 0, 0}, 4) + require.NoError(t, err) + require.Len(t, hits, 4) + assert.Equal(t, []string{"token-a", "token-b", "token-c", "token-d"}, hitTokens(hits)) + assert.Equal(t, []int{1, 2, 3, 4}, hitRanks(hits)) + assert.InDelta(t, 1, hits[0].Score, 1e-6) + assert.InDelta(t, 0, hits[1].Score, 1e-6) + assert.InDelta(t, 0, hits[2].Score, 1e-6) + assert.InDelta(t, -1, hits[3].Score, 1e-6) + paged, ok := backend.(document.PagedBackend) + require.True(t, ok) + firstPage, err := paged.SearchPage(ctx, gen1, 3, []float32{1, 0, 0}, "", 2) + require.NoError(t, err) + assert.False(t, firstPage.Exhausted) + assert.Equal(t, []string{"token-a", "token-b"}, hitTokens(firstPage.Hits)) + require.NoError(t, backend.DeleteTokens(ctx, gen1, []string{"token-a"})) + secondPage, err := paged.SearchPage(ctx, gen1, 3, []float32{1, 0, 0}, firstPage.NextCursor, 2) + require.NoError(t, err) + assert.True(t, secondPage.Exhausted) + assert.Equal(t, []string{"token-c", "token-d"}, hitTokens(secondPage.Hits)) + assert.Equal(t, []int{3, 4}, hitRanks(secondPage.Hits)) + + require.NoError(t, backend.PutUnpublished(ctx, gen1, 3, []document.Embedding{ + {Token: "token-a", Vector: []float32{0, 0, 1}}, + })) + hits, err = backend.Search(ctx, gen1, 3, []float32{0, 0, 1}, 1) + require.NoError(t, err) + require.Len(t, hits, 1) + assert.Equal(t, "token-a", hits[0].Token) + assert.InDelta(t, 1, hits[0].Score, 1e-6) + + require.NoError(t, backend.PutUnpublished(ctx, gen2, 4, []document.Embedding{ + {Token: "token-four", Vector: []float32{0, 0, 0, 1}}, + })) + hits, err = backend.Search(ctx, gen2, 4, []float32{0, 0, 0, 1}, 10) + require.NoError(t, err) + require.Len(t, hits, 1) + assert.Equal(t, "token-four", hits[0].Token) + hits, err = backend.Search(ctx, gen1, 3, []float32{0, 0, 1}, 10) + require.NoError(t, err) + assert.NotContains(t, hitTokens(hits), "token-four") + + err = backend.PutUnpublished(ctx, gen2, 4, []document.Embedding{ + {Token: "token-new", Vector: []float32{1, 0, 0, 0}}, + {Token: "token-a", Vector: []float32{0, 1, 0, 0}}, + }) + require.ErrorIs(t, err, document.ErrInvalidVector) + assert.Zero(t, tokenCount("token-new"), "generation collision rolls back the whole batch") + + err = backend.PutUnpublished(ctx, gen1, 3, []document.Embedding{ + {Token: "token-partial", Vector: []float32{1, 0, 0}}, + {Token: "token-wrong", Vector: []float32{1, 0}}, + }) + require.ErrorIs(t, err, vector.ErrDimensionMismatch) + assert.Zero(t, tokenCount("token-partial"), "dimension rejection happens before any write") + err = backend.PutUnpublished(ctx, gen1, 3, []document.Embedding{ + {Token: "token-partial-nonfinite", Vector: []float32{1, 0, 0}}, + {Token: "token-nonfinite", Vector: []float32{1, float32(math.NaN()), 0}}, + }) + require.ErrorIs(t, err, document.ErrInvalidVector) + assert.Zero(t, tokenCount("token-partial-nonfinite"), "nonfinite rejection happens before any write") + + for name, batch := range map[string][]document.Embedding{ + "duplicate": { + {Token: "duplicate", Vector: []float32{1, 0, 0}}, + {Token: "duplicate", Vector: []float32{0, 1, 0}}, + }, + "empty token": {{Token: "", Vector: []float32{1, 0, 0}}}, + "long token": {{Token: strings.Repeat("x", 1025), Vector: []float32{1, 0, 0}}}, + "nul token": {{Token: "x\x00y", Vector: []float32{1, 0, 0}}}, + "nonfinite": {{Token: "nan", Vector: []float32{1, float32(math.NaN()), 0}}}, + "infinite": {{Token: "inf", Vector: []float32{1, float32(math.Inf(1)), 0}}}, + "zero norm": {{Token: "zero", Vector: []float32{0, 0, 0}}}, + } { + t.Run("put rejects "+name, func(t *testing.T) { + err := backend.PutUnpublished(ctx, gen1, 3, batch) + require.ErrorIs(t, err, document.ErrInvalidVector) + }) + } + + tooMany := make([]document.Embedding, 1001) + for i := range tooMany { + tooMany[i] = document.Embedding{Token: "bounded-" + string(rune(i+1)), Vector: []float32{1, 0, 0}} + } + for name, call := range map[string]func() error{ + "put generation": func() error { return backend.PutUnpublished(ctx, 0, 3, embeddings[:1]) }, + "put dimension": func() error { return backend.PutUnpublished(ctx, gen1, 0, embeddings[:1]) }, + "put batch": func() error { return backend.PutUnpublished(ctx, gen1, 3, tooMany) }, + "delete gen": func() error { return backend.DeleteTokens(ctx, 0, []string{"token-a"}) }, + "delete batch": func() error { return backend.DeleteTokens(ctx, gen1, make([]string, 1001)) }, + "search gen": func() error { + _, err := backend.Search(ctx, 0, 3, []float32{1, 0, 0}, 1) + return err + }, + "search dimension": func() error { + _, err := backend.Search(ctx, gen1, 0, []float32{1, 0, 0}, 1) + return err + }, + "search k zero": func() error { + _, err := backend.Search(ctx, gen1, 3, []float32{1, 0, 0}, 0) + return err + }, + "search k large": func() error { + _, err := backend.Search(ctx, gen1, 3, []float32{1, 0, 0}, 1001) + return err + }, + "search nonfinite": func() error { + _, err := backend.Search(ctx, gen1, 3, []float32{1, float32(math.Inf(-1)), 0}, 1) + return err + }, + } { + t.Run("bounds reject "+name, func(t *testing.T) { + require.ErrorIs(t, call(), document.ErrInvalidVector) + }) + } + _, err = backend.Search(ctx, gen1, 3, []float32{1, 0}, 1) + require.ErrorIs(t, err, vector.ErrDimensionMismatch) + + require.NoError(t, backend.DeleteTokens(ctx, gen2, []string{"token-a"}), "wrong generation is an idempotent no-op") + assert.Equal(t, 1, tokenCount("token-a")) + require.NoError(t, backend.DeleteTokens(ctx, gen1, []string{"token-a"})) + require.NoError(t, backend.DeleteTokens(ctx, gen1, []string{"token-a"}), "repeat delete is idempotent") + assert.Zero(t, tokenCount("token-a")) + hits, err = backend.Search(ctx, gen1, 3, []float32{0, 0, 1}, 10) + require.NoError(t, err) + assert.NotContains(t, hitTokens(hits), "token-a") +} + +func hitTokens(hits []document.Hit) []string { + out := make([]string, len(hits)) + for i := range hits { + out[i] = hits[i].Token + } + return out +} + +func hitRanks(hits []document.Hit) []int { + out := make([]int, len(hits)) + for i := range hits { + out[i] = hits[i].Rank + } + return out +} + +func TestDocumentVectorMigrationSQLite(t *testing.T) { + parent, _ := openDocumentBackendForTest(t) + var tableSQL string + require.NoError(t, parent.db.QueryRow( + `SELECT sql FROM sqlite_master WHERE type = 'table' AND name = 'document_vector_embeddings'`).Scan(&tableSQL)) + assert.Contains(t, tableSQL, "token") + assert.NotContains(t, tableSQL, "REFERENCES index_generations") + + require.NoError(t, EnsureDocumentVectorTable(context.Background(), parent.db, 4)) + require.NoError(t, EnsureDocumentVectorTable(context.Background(), parent.db, 4), "lazy table migration is idempotent") + require.NoError(t, parent.db.QueryRow( + `SELECT sql FROM sqlite_master WHERE type = 'table' AND name = ?`, DocumentVectorTableName(4)).Scan(&tableSQL)) + assert.Contains(t, tableSQL, "distance_metric=cosine") +} diff --git a/internal/vector/sqlitevec/migrate.go b/internal/vector/sqlitevec/migrate.go index 7b9e6fa0a..5ca8976f3 100644 --- a/internal/vector/sqlitevec/migrate.go +++ b/internal/vector/sqlitevec/migrate.go @@ -554,3 +554,26 @@ func EnsureVectorTable(ctx context.Context, db *sql.DB, dim int) error { func VectorTableName(dim int) string { return fmt.Sprintf("vectors_vec_d%d", dim) } + +// EnsureDocumentVectorTable lazily creates the dedicated cosine vec0 table +// for one document-vector dimension. Its row IDs join only to +// document_vector_embeddings and never to message embeddings. +func EnsureDocumentVectorTable(ctx context.Context, db *sql.DB, dim int) error { + if dim <= 0 { + return fmt.Errorf("invalid document vector dimension %d", dim) + } + q := fmt.Sprintf(`CREATE VIRTUAL TABLE IF NOT EXISTS %s USING vec0( + generation_id INTEGER PARTITION KEY, + document_vector_id INTEGER PRIMARY KEY, + embedding FLOAT[%d] distance_metric=cosine + )`, DocumentVectorTableName(dim), dim) + if _, err := db.ExecContext(ctx, q); err != nil { + return fmt.Errorf("create %s: %w", DocumentVectorTableName(dim), err) + } + return nil +} + +// DocumentVectorTableName returns the dedicated dimension-specific vec0 name. +func DocumentVectorTableName(dim int) string { + return fmt.Sprintf("document_vectors_vec_d%d", dim) +} diff --git a/internal/vector/sqlitevec/schema.sql b/internal/vector/sqlitevec/schema.sql index 1283af066..7aa4a661b 100644 --- a/internal/vector/sqlitevec/schema.sql +++ b/internal/vector/sqlitevec/schema.sql @@ -140,7 +140,6 @@ CREATE TABLE IF NOT EXISTS embedding_document_progress ( reconcile_cursor TEXT NOT NULL DEFAULT '', journal_cursor TEXT NOT NULL DEFAULT '' ); - -- Visual vectors remain opaque in vectors.db. Ownership, source evidence, -- publication state, and model inputs stay in the authoritative archive. -- The fixed 1,024 dimension is the voyage-multimodal-3.5 contract. @@ -153,3 +152,15 @@ CREATE TABLE IF NOT EXISTS visual_vectors ( CREATE VIRTUAL TABLE IF NOT EXISTS visual_vectors_vec USING vec0( embedding float[1024] distance_metric=cosine ); + +-- Independent attachment-document vectors. Publication tokens are globally +-- unique opaque identities; generation IDs belong to the main document-vector +-- ledger and intentionally do not reference message index_generations. +CREATE TABLE IF NOT EXISTS document_vector_embeddings ( + document_vector_id INTEGER PRIMARY KEY AUTOINCREMENT, + token TEXT NOT NULL UNIQUE, + generation_id INTEGER NOT NULL, + dimension INTEGER NOT NULL CHECK (dimension > 0) +); +CREATE INDEX IF NOT EXISTS idx_document_vector_embeddings_generation + ON document_vector_embeddings(generation_id, dimension, token); diff --git a/nix/package.nix b/nix/package.nix index dd7e87db9..e49c7606b 100644 --- a/nix/package.nix +++ b/nix/package.nix @@ -16,7 +16,7 @@ buildGoModule { src = gitignoreSource ../.; - vendorHash = "sha256-mywedDZlx89nFjHuZmSqRxLnPMzVVNR4dqM2RwGWYbI="; + vendorHash = "sha256-0fb5ejfR9ctrGQCZjw209yHsRZsWyRXs4hVV5MogVkg="; proxyVendor = true; # Bun's copyfile backend can install incomplete packages when fetchBunDeps' diff --git a/pkg/client/generated/client.go b/pkg/client/generated/client.go index 2c6980d74..9abc8eaa9 100644 --- a/pkg/client/generated/client.go +++ b/pkg/client/generated/client.go @@ -339,6 +339,10 @@ type ClientInterface interface { GetDocumentIndexStatus(ctx context.Context, options *GetDocumentIndexStatusRequestOptions, reqEditors ...runtime.RequestEditorFn) (*GetDocumentIndexStatusResponse, error) GetDocumentIndexStatusWithResponse(ctx context.Context, options *GetDocumentIndexStatusRequestOptions, reqEditors ...runtime.RequestEditorFn) (*GetDocumentIndexStatusResp, error) + // GetDocumentVectorStatus Get document vector generation, consent, usage, and failure status + GetDocumentVectorStatus(ctx context.Context, options *GetDocumentVectorStatusRequestOptions, reqEditors ...runtime.RequestEditorFn) (*GetDocumentVectorStatusResponse, error) + GetDocumentVectorStatusWithResponse(ctx context.Context, options *GetDocumentVectorStatusRequestOptions, reqEditors ...runtime.RequestEditorFn) (*GetDocumentVectorStatusResp, error) + // SearchDomains Search analytical domains SearchDomains(ctx context.Context, options *SearchDomainsRequestOptions, reqEditors ...runtime.RequestEditorFn) (*SearchDomainsResponse, error) SearchDomainsWithResponse(ctx context.Context, options *SearchDomainsRequestOptions, reqEditors ...runtime.RequestEditorFn) (*SearchDomainsResp, error) @@ -5550,6 +5554,69 @@ func (c *Client) GetDocumentIndexStatus(ctx context.Context, options *GetDocumen return responseParser(ctx, resp) } +// GetDocumentVectorStatus Get document vector generation, consent, usage, and failure status +func (c *Client) GetDocumentVectorStatus(ctx context.Context, options *GetDocumentVectorStatusRequestOptions, reqEditors ...runtime.RequestEditorFn) (*GetDocumentVectorStatusResponse, error) { + var err error + reqParams := runtime.RequestOptionsParameters{ + RequestURL: c.apiClient.GetBaseURL() + "/api/v1/documents/vectors/status", + Method: "GET", + Options: options, + } + + req, err := c.apiClient.CreateRequest(ctx, reqParams, reqEditors...) + if err != nil { + return nil, fmt.Errorf("error creating request: %w", err) + } + + responseParser := func(ctx context.Context, resp *runtime.Response) (*GetDocumentVectorStatusResponse, error) { + bodyBytes := resp.Content + if resp.StatusCode != 200 { + target := new(GetDocumentVectorStatusErrorResponse) + // Handle empty error response body gracefully - skip unmarshal if no content + if len(bodyBytes) > 0 { + if err = json.Unmarshal(bodyBytes, target); err != nil { + return nil, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusErrorResponse", + Body: bodyBytes, + Err: err, + } + } + } + // Return error with (possibly empty) target + if errTarget, ok := any(*target).(error); ok { + return nil, runtime.NewClientAPIError(errTarget, runtime.WithStatusCode(resp.StatusCode)) + } + return nil, runtime.NewClientAPIError(fmt.Errorf("API error (status %d): %v", resp.StatusCode, *target), + runtime.WithStatusCode(resp.StatusCode)) + } + target := new(GetDocumentVectorStatusResponse) + // Handle empty response body gracefully + if len(bodyBytes) == 0 { + return target, nil + } + if err = json.Unmarshal(bodyBytes, target); err != nil { + return nil, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusResponse", + Body: bodyBytes, + Err: err, + } + } + return target, nil + } + + resp, err := c.apiClient.ExecuteRequest(ctx, req, "/api/v1/documents/vectors/status") + if err != nil { + return nil, fmt.Errorf("error executing request: %w", err) + } + return responseParser(ctx, resp) +} + // SearchDomains Search analytical domains func (c *Client) SearchDomains(ctx context.Context, options *SearchDomainsRequestOptions, reqEditors ...runtime.RequestEditorFn) (*SearchDomainsResponse, error) { var err error diff --git a/pkg/client/generated/client_options.go b/pkg/client/generated/client_options.go index 141c6b320..9345e29a3 100644 --- a/pkg/client/generated/client_options.go +++ b/pkg/client/generated/client_options.go @@ -2930,6 +2930,50 @@ func (o *GetDocumentIndexStatusRequestOptions) GetHeader() (map[string]string, e return nil, nil } +// GetDocumentVectorStatusRequestOptions is the options needed to make a request to GetDocumentVectorStatus. +type GetDocumentVectorStatusRequestOptions struct { + Query *GetDocumentVectorStatusQuery +} + +// Validate validates all the fields in the options. +// Use it if fields validation was not run. +func (o *GetDocumentVectorStatusRequestOptions) Validate() error { + var errors runtime.ValidationErrors + + if o.Query != nil { + if v, ok := any(o.Query).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Query", err) + } + } + } + if len(errors) == 0 { + return nil + } + + return errors +} + +// GetPathParams returns the path params as a map. +func (o *GetDocumentVectorStatusRequestOptions) GetPathParams() (map[string]any, error) { + return nil, nil +} + +// GetQuery returns the query params as a map. +func (o *GetDocumentVectorStatusRequestOptions) GetQuery() (map[string]any, error) { + return runtime.AsMap[any](o.Query) +} + +// GetBody returns the payload in any type that can be marshalled to JSON by the client. +func (o *GetDocumentVectorStatusRequestOptions) GetBody() any { + return nil +} + +// GetHeader returns the headers as a map. +func (o *GetDocumentVectorStatusRequestOptions) GetHeader() (map[string]string, error) { + return nil, nil +} + // SearchDomainsRequestOptions is the options needed to make a request to SearchDomains. type SearchDomainsRequestOptions struct { Body *SearchDomainsBody diff --git a/pkg/client/generated/client_with_response.go b/pkg/client/generated/client_with_response.go index 7c0931867..828c70de8 100644 --- a/pkg/client/generated/client_with_response.go +++ b/pkg/client/generated/client_with_response.go @@ -6447,6 +6447,117 @@ func (c *Client) GetDocumentIndexStatusWithResponse(ctx context.Context, options } } +// GetDocumentVectorStatus Get document vector generation, consent, usage, and failure status +func (c *Client) GetDocumentVectorStatusWithResponse(ctx context.Context, options *GetDocumentVectorStatusRequestOptions, reqEditors ...runtime.RequestEditorFn) (*GetDocumentVectorStatusResp, error) { + var err error + reqParams := runtime.RequestOptionsParameters{ + RequestURL: c.apiClient.GetBaseURL() + "/api/v1/documents/vectors/status", + Method: "GET", + Options: options, + } + + req, err := c.apiClient.CreateRequest(ctx, reqParams, reqEditors...) + if err != nil { + return nil, fmt.Errorf("error creating request: %w", err) + } + + resp, err := c.apiClient.ExecuteRequest(ctx, req, "/api/v1/documents/vectors/status") + if err != nil { + return nil, fmt.Errorf("error executing request: %w", err) + } + + out := &GetDocumentVectorStatusResp{ + HTTPResponse: resp.Raw, + Body: resp.Content, + StatusCode: resp.StatusCode, + } + + switch resp.StatusCode { + case 200: + out.JSON200 = new(GetDocumentVectorStatusResponse) + bodyBytes := resp.Content + if len(bodyBytes) > 0 { + if err := json.Unmarshal(bodyBytes, out.JSON200); err != nil { + return out, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusResponse", + Body: bodyBytes, + Err: err, + } + } + } + return out, nil + case 400: + out.JSON400 = new(GetDocumentVectorStatusErrorResponse) + bodyBytes := resp.Content + if len(bodyBytes) > 0 { + if err := json.Unmarshal(bodyBytes, out.JSON400); err != nil { + return out, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusErrorResponse", + Body: bodyBytes, + Err: err, + } + } + } + return out, runtime.NewClientAPIError(fmt.Errorf("API error (status %d)", resp.StatusCode), runtime.WithStatusCode(resp.StatusCode)) + case 403: + out.JSON403 = new(GetDocumentVectorStatusErrorResponseJSON) + bodyBytes := resp.Content + if len(bodyBytes) > 0 { + if err := json.Unmarshal(bodyBytes, out.JSON403); err != nil { + return out, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusErrorResponseJSON", + Body: bodyBytes, + Err: err, + } + } + } + return out, runtime.NewClientAPIError(fmt.Errorf("API error (status %d)", resp.StatusCode), runtime.WithStatusCode(resp.StatusCode)) + case 429: + out.JSON429 = new(GetDocumentVectorStatusErrorResponseJSON429) + bodyBytes := resp.Content + if len(bodyBytes) > 0 { + if err := json.Unmarshal(bodyBytes, out.JSON429); err != nil { + return out, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusErrorResponseJSON429", + Body: bodyBytes, + Err: err, + } + } + } + return out, runtime.NewClientAPIError(fmt.Errorf("API error (status %d)", resp.StatusCode), runtime.WithStatusCode(resp.StatusCode)) + case 503: + out.JSON503 = new(GetDocumentVectorStatusErrorResponseJSON503) + bodyBytes := resp.Content + if len(bodyBytes) > 0 { + if err := json.Unmarshal(bodyBytes, out.JSON503); err != nil { + return out, &runtime.ResponseDecodeError{ + StatusCode: resp.StatusCode, + ContentType: resp.Headers.Get("Content-Type"), + ContentLength: len(bodyBytes), + TargetType: "GetDocumentVectorStatusErrorResponseJSON503", + Body: bodyBytes, + Err: err, + } + } + } + return out, runtime.NewClientAPIError(fmt.Errorf("API error (status %d)", resp.StatusCode), runtime.WithStatusCode(resp.StatusCode)) + default: + return out, runtime.NewClientAPIError(fmt.Errorf("unexpected status code: %d", resp.StatusCode), runtime.WithStatusCode(resp.StatusCode)) + } +} + // SearchDomains Search analytical domains func (c *Client) SearchDomainsWithResponse(ctx context.Context, options *SearchDomainsRequestOptions, reqEditors ...runtime.RequestEditorFn) (*SearchDomainsResp, error) { var err error diff --git a/pkg/client/generated/queries.go b/pkg/client/generated/queries.go index 593d16f19..b83d5172f 100644 --- a/pkg/client/generated/queries.go +++ b/pkg/client/generated/queries.go @@ -365,6 +365,12 @@ type SearchDocumentsQuery struct { // Cursor Opaque cursor from the previous document search page Cursor *string `json:"cursor,omitempty"` + + // Mode Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider + Mode *string `json:"mode,omitempty"` + + // CandidateLimit Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000) + CandidateLimit *int64 `json:"candidate_limit,omitempty"` } func (s SearchDocumentsQuery) Validate() error { @@ -389,6 +395,17 @@ func (g GetDocumentIndexStatusQuery) Validate() error { return runtime.ConvertValidatorError(typesValidator.Struct(g)) } +type GetDocumentVectorStatusQuery struct { + // GenerationID Generation whose bounded failures to inspect + GenerationID *int64 `json:"generation_id,omitempty"` + + // AfterToken Stable failure cursor token + AfterToken *string `json:"after_token,omitempty"` + + // Limit Maximum failure diagnostics (default 20, max 1000) + Limit *int64 `json:"limit,omitempty"` +} + type ListIdentityMatchCandidatesQuery struct { // State Candidate state filter (candidate, accepted, rejected, conflict); repeat or comma-separate for multiple values State *string `json:"state,omitempty"` diff --git a/pkg/client/generated/responses.go b/pkg/client/generated/responses.go index aec821c74..a507855b5 100644 --- a/pkg/client/generated/responses.go +++ b/pkg/client/generated/responses.go @@ -593,6 +593,16 @@ type GetDocumentIndexStatusErrorResponseJSON429 = ErrorResponse type GetDocumentIndexStatusErrorResponseJSON503 = ErrorResponse +type GetDocumentVectorStatusResponse = DocumentVectorOperationsResponse + +type GetDocumentVectorStatusErrorResponse = ErrorResponse + +type GetDocumentVectorStatusErrorResponseJSON = ErrorResponse + +type GetDocumentVectorStatusErrorResponseJSON429 = ErrorResponse + +type GetDocumentVectorStatusErrorResponseJSON503 = ErrorResponse + type SearchDomainsResponse = DomainSearchHTTPResponse type SearchDomainsErrorResponse = ErrorResponse @@ -3322,6 +3332,17 @@ type GetDocumentIndexStatusResp struct { JSON503 *GetDocumentIndexStatusErrorResponseJSON503 } +type GetDocumentVectorStatusResp struct { + HTTPResponse *http.Response + Body []byte + StatusCode int + JSON200 *GetDocumentVectorStatusResponse + JSON400 *GetDocumentVectorStatusErrorResponse + JSON403 *GetDocumentVectorStatusErrorResponseJSON + JSON429 *GetDocumentVectorStatusErrorResponseJSON429 + JSON503 *GetDocumentVectorStatusErrorResponseJSON503 +} + type SearchDomainsResp struct { HTTPResponse *http.Response Body []byte diff --git a/pkg/client/generated/types.go b/pkg/client/generated/types.go index fd2f6e452..d4e76872d 100644 --- a/pkg/client/generated/types.go +++ b/pkg/client/generated/types.go @@ -2105,10 +2105,13 @@ func (d DocumentIndexStatusResponse) Validate() error { } type DocumentSearchResponse struct { - NextCursor *string `json:"next_cursor,omitempty"` - Results []DocumentSearchResult `json:"results,omitempty" validate:"required"` - Revision int64 `json:"revision"` - Truncated *bool `json:"truncated,omitempty"` + EffectiveMode *string `json:"effective_mode,omitempty"` + NextCursor *string `json:"next_cursor,omitempty"` + Results []DocumentSearchResult `json:"results,omitempty" validate:"required"` + Revision int64 `json:"revision"` + Truncated *bool `json:"truncated,omitempty"` + VectorGenerationFingerprint *string `json:"vector_generation_fingerprint,omitempty"` + VectorGenerationID *int64 `json:"vector_generation_id,omitempty"` } func (d DocumentSearchResponse) Validate() error { @@ -2127,35 +2130,45 @@ func (d DocumentSearchResponse) Validate() error { } type DocumentSearchResult struct { - AttachmentID int64 `json:"attachment_id"` - CanonicalBlobHash string `json:"canonical_blob_hash" validate:"required"` - ChunkKey string `json:"chunk_key" validate:"required"` - ChunkOrdinal int64 `json:"chunk_ordinal"` - ContainingTitle *string `json:"containing_title,omitempty"` - ConversationID int64 `json:"conversation_id"` - Excerpt string `json:"excerpt" validate:"required"` - ExtractionID string `json:"extraction_id" validate:"required"` - Filename *string `json:"filename,omitempty"` - FirstUnitIndex int64 `json:"first_unit_index"` - HeadingPath []string `json:"heading_path,omitempty"` - HighlightEnd int64 `json:"highlight_end"` - HighlightStart int64 `json:"highlight_start"` - LastUnitIndex int64 `json:"last_unit_index"` - MatchedSignals []string `json:"matched_signals,omitempty" validate:"required"` - MessageID int64 `json:"message_id"` - MimeType *string `json:"mime_type,omitempty"` - Model string `json:"model" validate:"required"` - OccurredAt *time.Time `json:"occurred_at,omitempty"` - OccurrenceKey string `json:"occurrence_key" validate:"required"` - OtherLiveCopies int64 `json:"other_live_copies"` - PersonProvenance *Provenance `json:"person_provenance,omitempty"` - ProfileID string `json:"profile_id" validate:"required"` - Provider string `json:"provider" validate:"required"` - Rank int64 `json:"rank"` - SourceID int64 `json:"source_id"` - SourceMessageID *string `json:"source_message_id,omitempty"` - SourcePartKey *string `json:"source_part_key,omitempty"` - Truncated bool `json:"truncated"` + AttachmentID int64 `json:"attachment_id"` + CanonicalBlobHash string `json:"canonical_blob_hash" validate:"required"` + ChunkKey string `json:"chunk_key" validate:"required"` + ChunkOrdinal int64 `json:"chunk_ordinal"` + ContainingTitle *string `json:"containing_title,omitempty"` + ConversationID int64 `json:"conversation_id"` + Excerpt string `json:"excerpt" validate:"required"` + ExtractionID string `json:"extraction_id" validate:"required"` + Filename *string `json:"filename,omitempty"` + FirstUnitIndex int64 `json:"first_unit_index"` + FusionScore *float64 `json:"fusion_score,omitempty"` + HeadingPath []string `json:"heading_path,omitempty"` + HighlightEnd int64 `json:"highlight_end"` + HighlightStart int64 `json:"highlight_start"` + LastUnitIndex int64 `json:"last_unit_index"` + LexicalRank *int64 `json:"lexical_rank,omitempty"` + MatchedSignals []string `json:"matched_signals,omitempty" validate:"required"` + MessageID int64 `json:"message_id"` + MimeType *string `json:"mime_type,omitempty"` + Model string `json:"model" validate:"required"` + OccurredAt *time.Time `json:"occurred_at,omitempty"` + OccurrenceKey string `json:"occurrence_key" validate:"required"` + OtherLiveCopies int64 `json:"other_live_copies"` + PersonProvenance *Provenance `json:"person_provenance,omitempty"` + ProfileID string `json:"profile_id" validate:"required"` + Provider string `json:"provider" validate:"required"` + Rank int64 `json:"rank"` + SemanticRank *int64 `json:"semantic_rank,omitempty"` + SemanticScore *float64 `json:"semantic_score,omitempty"` + SourceID int64 `json:"source_id"` + SourceMessageID *string `json:"source_message_id,omitempty"` + SourcePartKey *string `json:"source_part_key,omitempty"` + Truncated bool `json:"truncated"` + VectorDimension *int64 `json:"vector_dimension,omitempty"` + VectorEmbeddingProfile *string `json:"vector_embedding_profile,omitempty"` + VectorGenerationFingerprint *string `json:"vector_generation_fingerprint,omitempty"` + VectorGenerationID *int64 `json:"vector_generation_id,omitempty"` + VectorModel *string `json:"vector_model,omitempty"` + VectorToken *string `json:"vector_token,omitempty"` } func (d DocumentSearchResult) Validate() error { @@ -2200,6 +2213,215 @@ func (d DocumentSearchResult) Validate() error { return errors } +type DocumentVectorConsent struct { + ConsentedAt time.Time `json:"consented_at" validate:"required"` + Dimension int64 `json:"dimension"` + EgressFingerprint string `json:"egress_fingerprint" validate:"required"` + EmbeddingProfile string `json:"embedding_profile" validate:"required"` + Fingerprint string `json:"fingerprint" validate:"required"` + Model string `json:"model" validate:"required"` + Purpose string `json:"purpose" validate:"required"` + TargetExtractionProfileID string `json:"target_extraction_profile_id" validate:"required"` +} + +func (d DocumentVectorConsent) Validate() error { + return runtime.ConvertValidatorError(typesValidator.Struct(d)) +} + +type DocumentVectorCoverage struct { + Ready int64 `json:"ready"` + Required int64 `json:"required"` +} + +type DocumentVectorFailureDiagnostic struct { + AttemptCount int64 `json:"attempt_count"` + ErrorCode string `json:"error_code" validate:"required"` + NextRetryAt *time.Time `json:"next_retry_at,omitempty"` + Terminal bool `json:"terminal"` + Token string `json:"token" validate:"required"` +} + +func (d DocumentVectorFailureDiagnostic) Validate() error { + return runtime.ConvertValidatorError(typesValidator.Struct(d)) +} + +type DocumentVectorGeneration struct { + ActivatedAt *time.Time `json:"activated_at,omitempty"` + CreatedAt time.Time `json:"created_at" validate:"required"` + Dimension int64 `json:"dimension"` + EmbeddingProfile string `json:"embedding_profile" validate:"required"` + Fingerprint string `json:"fingerprint" validate:"required"` + ID int64 `json:"id"` + Model string `json:"model" validate:"required"` + RetiredAt *time.Time `json:"retired_at,omitempty"` + State string `json:"state" validate:"required"` + TargetExtractionProfileID string `json:"target_extraction_profile_id" validate:"required"` +} + +func (d DocumentVectorGeneration) Validate() error { + return runtime.ConvertValidatorError(typesValidator.Struct(d)) +} + +type DocumentVectorGenerationSpec struct { + Dimension int64 `json:"dimension"` + EmbeddingProfile string `json:"embedding_profile" validate:"required"` + Fingerprint string `json:"fingerprint" validate:"required"` + Model string `json:"model" validate:"required"` + TargetExtractionProfileID string `json:"target_extraction_profile_id" validate:"required"` +} + +func (d DocumentVectorGenerationSpec) Validate() error { + return runtime.ConvertValidatorError(typesValidator.Struct(d)) +} + +type DocumentVectorGenerationStatus struct { + Blocked bool `json:"blocked"` + CleanupPending int64 `json:"cleanup_pending"` + FailureAfterGenerationID *int64 `json:"failure_after_generation_id,omitempty"` + FailureAfterToken *string `json:"failure_after_token,omitempty"` + Failures []DocumentVectorFailureDiagnostic `json:"failures,omitempty" validate:"required"` + FailuresExhausted bool `json:"failures_exhausted"` + GenerationID int64 `json:"generation_id"` + Pending int64 `json:"pending"` + ReadyLive int64 `json:"ready_live"` + Retryable int64 `json:"retryable"` + StaleObsolete int64 `json:"stale_obsolete"` + State string `json:"state" validate:"required"` + Terminal int64 `json:"terminal"` +} + +func (d DocumentVectorGenerationStatus) Validate() error { + var errors runtime.ValidationErrors + for i, item := range d.Failures { + if v, ok := any(item).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append(fmt.Sprintf("Failures[%d]", i), err) + } + } + } + if err := typesValidator.Var(d.State, "required"); err != nil { + errors = errors.Append("State", err) + } + if len(errors) == 0 { + return nil + } + return errors +} + +type DocumentVectorOperationsResponse struct { + Configured bool `json:"configured"` + Enabled bool `json:"enabled"` + ScheduledRegistrationRequiresRestart *bool `json:"scheduled_registration_requires_restart,omitempty"` + Status *DocumentVectorOperationsStatus `json:"status,omitempty"` +} + +func (d DocumentVectorOperationsResponse) Validate() error { + var errors runtime.ValidationErrors + if d.Status != nil { + if v, ok := any(d.Status).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Status", err) + } + } + } + if len(errors) == 0 { + return nil + } + return errors +} + +type DocumentVectorOperationsStatus struct { + Active *DocumentVectorGeneration `json:"active,omitempty"` + Building *DocumentVectorGeneration `json:"building,omitempty"` + ConfiguredDocumentEgressFingerprint string `json:"configured_document_egress_fingerprint" validate:"required"` + ConfiguredQueryEgressFingerprint string `json:"configured_query_egress_fingerprint" validate:"required"` + ConfiguredSpec DocumentVectorGenerationSpec `json:"configured_spec"` + Coverage *DocumentVectorCoverage `json:"coverage,omitempty"` + DocumentConsent *DocumentVectorConsent `json:"document_consent,omitempty"` + QueryConsent *DocumentVectorConsent `json:"query_consent,omitempty"` + Selected *DocumentVectorGenerationStatus `json:"selected,omitempty"` + Usage DocumentVectorProviderUsage `json:"usage"` +} + +func (d DocumentVectorOperationsStatus) Validate() error { + var errors runtime.ValidationErrors + if d.Active != nil { + if v, ok := any(d.Active).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Active", err) + } + } + } + if d.Building != nil { + if v, ok := any(d.Building).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Building", err) + } + } + } + if err := typesValidator.Var(d.ConfiguredDocumentEgressFingerprint, "required"); err != nil { + errors = errors.Append("ConfiguredDocumentEgressFingerprint", err) + } + if err := typesValidator.Var(d.ConfiguredQueryEgressFingerprint, "required"); err != nil { + errors = errors.Append("ConfiguredQueryEgressFingerprint", err) + } + if v, ok := any(d.ConfiguredSpec).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("ConfiguredSpec", err) + } + } + if d.Coverage != nil { + if v, ok := any(d.Coverage).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Coverage", err) + } + } + } + if d.DocumentConsent != nil { + if v, ok := any(d.DocumentConsent).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("DocumentConsent", err) + } + } + } + if d.QueryConsent != nil { + if v, ok := any(d.QueryConsent).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("QueryConsent", err) + } + } + } + if d.Selected != nil { + if v, ok := any(d.Selected).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Selected", err) + } + } + } + if v, ok := any(d.Usage).(runtime.Validator); ok { + if err := v.Validate(); err != nil { + errors = errors.Append("Usage", err) + } + } + if len(errors) == 0 { + return nil + } + return errors +} + +type DocumentVectorProviderUsage struct { + Fingerprint string `json:"fingerprint" validate:"required"` + ProviderCalls int64 `json:"provider_calls"` + ProviderChunks int64 `json:"provider_chunks"` + ProviderDocuments int64 `json:"provider_documents"` + ProviderInputChars int64 `json:"provider_input_chars"` + UpdatedAt *time.Time `json:"updated_at,omitempty"` +} + +func (d DocumentVectorProviderUsage) Validate() error { + return runtime.ConvertValidatorError(typesValidator.Struct(d)) +} + type DomainContextSummaryHTTPResponse struct { CacheRevision string `json:"cache_revision" validate:"required"` CandidateSnapshotID *string `json:"candidate_snapshot_id,omitempty"` diff --git a/pkg/client/openapi.yaml b/pkg/client/openapi.yaml index 67c23a549..a2be356c9 100644 --- a/pkg/client/openapi.yaml +++ b/pkg/client/openapi.yaml @@ -2491,6 +2491,8 @@ components: type: object DocumentSearchResponse: properties: + effective_mode: + type: string next_cursor: type: string results: @@ -2503,6 +2505,11 @@ components: type: integer truncated: type: boolean + vector_generation_fingerprint: + type: string + vector_generation_id: + format: int64 + type: integer required: - results - revision @@ -2533,6 +2540,9 @@ components: first_unit_index: format: int64 type: integer + fusion_score: + format: double + type: number heading_path: items: type: string @@ -2547,6 +2557,9 @@ components: last_unit_index: format: int64 type: integer + lexical_rank: + format: int64 + type: integer matched_signals: items: type: string @@ -2576,6 +2589,12 @@ components: rank: format: int64 type: integer + semantic_rank: + format: int64 + type: integer + semantic_score: + format: double + type: number source_id: format: int64 type: integer @@ -2585,6 +2604,20 @@ components: type: string truncated: type: boolean + vector_dimension: + format: int64 + type: integer + vector_embedding_profile: + type: string + vector_generation_fingerprint: + type: string + vector_generation_id: + format: int64 + type: integer + vector_model: + type: string + vector_token: + type: string required: - attachment_id - message_id @@ -2608,6 +2641,245 @@ components: - truncated - rank type: object + DocumentVectorConsent: + properties: + consented_at: + format: date-time + type: string + dimension: + format: int64 + type: integer + egress_fingerprint: + type: string + embedding_profile: + type: string + fingerprint: + type: string + model: + type: string + purpose: + type: string + target_extraction_profile_id: + type: string + required: + - consented_at + - egress_fingerprint + - purpose + - fingerprint + - target_extraction_profile_id + - embedding_profile + - model + - dimension + type: object + DocumentVectorCoverage: + properties: + ready: + format: int64 + type: integer + required: + format: int64 + type: integer + required: + - required + - ready + type: object + DocumentVectorFailureDiagnostic: + properties: + attempt_count: + format: int64 + type: integer + error_code: + type: string + next_retry_at: + format: date-time + type: string + terminal: + type: boolean + token: + type: string + required: + - token + - attempt_count + - terminal + - error_code + type: object + DocumentVectorGeneration: + properties: + activated_at: + format: date-time + type: string + created_at: + format: date-time + type: string + dimension: + format: int64 + type: integer + embedding_profile: + type: string + fingerprint: + type: string + id: + format: int64 + type: integer + model: + type: string + retired_at: + format: date-time + type: string + state: + type: string + target_extraction_profile_id: + type: string + required: + - id + - state + - created_at + - fingerprint + - target_extraction_profile_id + - embedding_profile + - model + - dimension + type: object + DocumentVectorGenerationSpec: + properties: + dimension: + format: int64 + type: integer + embedding_profile: + type: string + fingerprint: + type: string + model: + type: string + target_extraction_profile_id: + type: string + required: + - fingerprint + - target_extraction_profile_id + - embedding_profile + - model + - dimension + type: object + DocumentVectorGenerationStatus: + properties: + blocked: + type: boolean + cleanup_pending: + format: int64 + type: integer + failure_after_generation_id: + format: int64 + type: integer + failure_after_token: + type: string + failures: + items: + $ref: "#/components/schemas/DocumentVectorFailureDiagnostic" + nullable: true + type: array + failures_exhausted: + type: boolean + generation_id: + format: int64 + type: integer + pending: + format: int64 + type: integer + ready_live: + format: int64 + type: integer + retryable: + format: int64 + type: integer + stale_obsolete: + format: int64 + type: integer + state: + type: string + terminal: + format: int64 + type: integer + required: + - generation_id + - state + - blocked + - pending + - retryable + - terminal + - ready_live + - stale_obsolete + - cleanup_pending + - failures + - failures_exhausted + type: object + DocumentVectorOperationsResponse: + properties: + configured: + type: boolean + enabled: + type: boolean + scheduled_registration_requires_restart: + type: boolean + status: + $ref: "#/components/schemas/DocumentVectorOperationsStatus" + required: + - enabled + - configured + type: object + DocumentVectorOperationsStatus: + properties: + active: + $ref: "#/components/schemas/DocumentVectorGeneration" + building: + $ref: "#/components/schemas/DocumentVectorGeneration" + configured_document_egress_fingerprint: + type: string + configured_query_egress_fingerprint: + type: string + configured_spec: + $ref: "#/components/schemas/DocumentVectorGenerationSpec" + coverage: + $ref: "#/components/schemas/DocumentVectorCoverage" + document_consent: + $ref: "#/components/schemas/DocumentVectorConsent" + query_consent: + $ref: "#/components/schemas/DocumentVectorConsent" + selected: + $ref: "#/components/schemas/DocumentVectorGenerationStatus" + usage: + $ref: "#/components/schemas/DocumentVectorProviderUsage" + required: + - configured_spec + - configured_document_egress_fingerprint + - configured_query_egress_fingerprint + - usage + type: object + DocumentVectorProviderUsage: + properties: + fingerprint: + type: string + provider_calls: + format: int64 + type: integer + provider_chunks: + format: int64 + type: integer + provider_documents: + format: int64 + type: integer + provider_input_chars: + format: int64 + type: integer + updated_at: + format: date-time + type: string + required: + - fingerprint + - provider_calls + - provider_documents + - provider_chunks + - provider_input_chars + type: object DomainContextSummaryHTTPResponse: properties: cache_revision: @@ -13433,6 +13705,17 @@ paths: name: cursor schema: type: string + - description: "Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider" + in: query + name: mode + schema: + type: string + - description: "Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000)" + in: query + name: candidate_limit + schema: + format: int64 + type: integer responses: "200": content: @@ -13560,6 +13843,69 @@ paths: summary: Get extracted document index status tags: - API + /api/v1/documents/vectors/status: + get: + operationId: getDocumentVectorStatus + parameters: + - description: Generation whose bounded failures to inspect + in: query + name: generation_id + schema: + format: int64 + type: integer + - description: Stable failure cursor token + in: query + name: after_token + schema: + type: string + - description: Maximum failure diagnostics (default 20, max 1000) + in: query + name: limit + schema: + format: int64 + type: integer + responses: + "200": + content: + application/json: + schema: + $ref: "#/components/schemas/DocumentVectorOperationsResponse" + description: OK + "400": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + "403": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + "429": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + "503": + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + default: + content: + application/json: + schema: + $ref: "#/components/schemas/ErrorResponse" + description: Error + security: + - apiKey: [] + summary: Get document vector generation, consent, usage, and failure status + tags: + - API /api/v1/domains/search: post: operationId: searchDomains diff --git a/web/src/lib/api/generated/schema.d.ts b/web/src/lib/api/generated/schema.d.ts index 9aaf9574b..4dd8fe86e 100644 --- a/web/src/lib/api/generated/schema.d.ts +++ b/web/src/lib/api/generated/schema.d.ts @@ -1079,6 +1079,23 @@ export interface paths { patch?: never; trace?: never; }; + "/api/v1/documents/vectors/status": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get document vector generation, consent, usage, and failure status */ + get: operations["getDocumentVectorStatus"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/api/v1/domains/search": { parameters: { query?: never; @@ -4216,11 +4233,15 @@ export interface components { [key: string]: unknown; }; DocumentSearchResponse: { + effective_mode?: string; next_cursor?: string; results: components["schemas"]["DocumentSearchResult"][] | null; /** Format: int64 */ revision: number; truncated?: boolean; + vector_generation_fingerprint?: string; + /** Format: int64 */ + vector_generation_id?: number; } & { [key: string]: unknown; }; @@ -4239,6 +4260,8 @@ export interface components { filename?: string; /** Format: int64 */ first_unit_index: number; + /** Format: double */ + fusion_score?: number; heading_path?: string[] | null; /** Format: int64 */ highlight_end: number; @@ -4246,6 +4269,8 @@ export interface components { highlight_start: number; /** Format: int64 */ last_unit_index: number; + /** Format: int64 */ + lexical_rank?: number; matched_signals: string[] | null; /** Format: int64 */ message_id: number; @@ -4262,10 +4287,146 @@ export interface components { /** Format: int64 */ rank: number; /** Format: int64 */ + semantic_rank?: number; + /** Format: double */ + semantic_score?: number; + /** Format: int64 */ source_id: number; source_message_id?: string; source_part_key?: string; truncated: boolean; + /** Format: int64 */ + vector_dimension?: number; + vector_embedding_profile?: string; + vector_generation_fingerprint?: string; + /** Format: int64 */ + vector_generation_id?: number; + vector_model?: string; + vector_token?: string; + } & { + [key: string]: unknown; + }; + DocumentVectorConsent: { + /** Format: date-time */ + consented_at: string; + /** Format: int64 */ + dimension: number; + egress_fingerprint: string; + embedding_profile: string; + fingerprint: string; + model: string; + purpose: string; + target_extraction_profile_id: string; + } & { + [key: string]: unknown; + }; + DocumentVectorCoverage: { + /** Format: int64 */ + ready: number; + /** Format: int64 */ + required: number; + } & { + [key: string]: unknown; + }; + DocumentVectorFailureDiagnostic: { + /** Format: int64 */ + attempt_count: number; + error_code: string; + /** Format: date-time */ + next_retry_at?: string; + terminal: boolean; + token: string; + } & { + [key: string]: unknown; + }; + DocumentVectorGeneration: { + /** Format: date-time */ + activated_at?: string; + /** Format: date-time */ + created_at: string; + /** Format: int64 */ + dimension: number; + embedding_profile: string; + fingerprint: string; + /** Format: int64 */ + id: number; + model: string; + /** Format: date-time */ + retired_at?: string; + state: string; + target_extraction_profile_id: string; + } & { + [key: string]: unknown; + }; + DocumentVectorGenerationSpec: { + /** Format: int64 */ + dimension: number; + embedding_profile: string; + fingerprint: string; + model: string; + target_extraction_profile_id: string; + } & { + [key: string]: unknown; + }; + DocumentVectorGenerationStatus: { + blocked: boolean; + /** Format: int64 */ + cleanup_pending: number; + /** Format: int64 */ + failure_after_generation_id?: number; + failure_after_token?: string; + failures: components["schemas"]["DocumentVectorFailureDiagnostic"][] | null; + failures_exhausted: boolean; + /** Format: int64 */ + generation_id: number; + /** Format: int64 */ + pending: number; + /** Format: int64 */ + ready_live: number; + /** Format: int64 */ + retryable: number; + /** Format: int64 */ + stale_obsolete: number; + state: string; + /** Format: int64 */ + terminal: number; + } & { + [key: string]: unknown; + }; + DocumentVectorOperationsResponse: { + configured: boolean; + enabled: boolean; + scheduled_registration_requires_restart?: boolean; + status?: components["schemas"]["DocumentVectorOperationsStatus"]; + } & { + [key: string]: unknown; + }; + DocumentVectorOperationsStatus: { + active?: components["schemas"]["DocumentVectorGeneration"]; + building?: components["schemas"]["DocumentVectorGeneration"]; + configured_document_egress_fingerprint: string; + configured_query_egress_fingerprint: string; + configured_spec: components["schemas"]["DocumentVectorGenerationSpec"]; + coverage?: components["schemas"]["DocumentVectorCoverage"]; + document_consent?: components["schemas"]["DocumentVectorConsent"]; + query_consent?: components["schemas"]["DocumentVectorConsent"]; + selected?: components["schemas"]["DocumentVectorGenerationStatus"]; + usage: components["schemas"]["DocumentVectorProviderUsage"]; + } & { + [key: string]: unknown; + }; + DocumentVectorProviderUsage: { + fingerprint: string; + /** Format: int64 */ + provider_calls: number; + /** Format: int64 */ + provider_chunks: number; + /** Format: int64 */ + provider_documents: number; + /** Format: int64 */ + provider_input_chars: number; + /** Format: date-time */ + updated_at?: string; } & { [key: string]: unknown; }; @@ -11500,6 +11661,10 @@ export interface operations { limit?: number; /** @description Opaque cursor from the previous document search page */ cursor?: string; + /** @description Search mode: lexical (default and auto); semantic/hybrid send the query to the embedding provider */ + mode?: string; + /** @description Maximum candidates (default/max: lexical 10000; semantic/hybrid 100/1000) */ + candidate_limit?: number; }; header?: never; path?: never; @@ -11655,6 +11820,78 @@ export interface operations { }; }; }; + getDocumentVectorStatus: { + parameters: { + query?: { + /** @description Generation whose bounded failures to inspect */ + generation_id?: number; + /** @description Stable failure cursor token */ + after_token?: string; + /** @description Maximum failure diagnostics (default 20, max 1000) */ + limit?: number; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description OK */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["DocumentVectorOperationsResponse"]; + }; + }; + /** @description Error */ + 400: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ErrorResponse"]; + }; + }; + /** @description Error */ + 403: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ErrorResponse"]; + }; + }; + /** @description Error */ + 429: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ErrorResponse"]; + }; + }; + /** @description Error */ + 503: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ErrorResponse"]; + }; + }; + /** @description Error */ + default: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ErrorResponse"]; + }; + }; + }; + }; searchDomains: { parameters: { query?: never; From 3079e0c03fce9a5bf53d1dcfb5ceb2f7760f3d31 Mon Sep 17 00:00:00 2001 From: Wes McKinney Date: Mon, 24 Aug 2026 09:18:40 -0500 Subject: [PATCH 2/4] fix(documents): serialize SQLite vector operations Cron and post-sync convergence can run at the same time in one SQLite daemon. The archive ownership lock only excludes other processes, so both passes could mutate the same generation concurrently. Serialize the existing SQLite operation boundary in process. PostgreSQL retains its schema-wide advisory lock. Generated with Codex Co-authored-by: Codex --- internal/store/document_vector_lock.go | 8 +++-- internal/store/document_vectors_test.go | 41 +++++++++++++++++++++++++ internal/store/store.go | 3 +- internal/vector/document/worker_test.go | 16 +++++----- 4 files changed, 57 insertions(+), 11 deletions(-) diff --git a/internal/store/document_vector_lock.go b/internal/store/document_vector_lock.go index 3654465a9..774b4a8e6 100644 --- a/internal/store/document_vector_lock.go +++ b/internal/store/document_vector_lock.go @@ -9,14 +9,16 @@ import ( const documentVectorOperationLockSQL = `hashtextextended( current_database() || ':' || current_schema() || ':msgvault.document_vectors', 0)` -// WithDocumentVectorOperationLock serializes the document-vector writer across -// every process using the same PostgreSQL archive schema. SQLite writers are -// already serialized by the archive ownership lock. +// WithDocumentVectorOperationLock serializes document-vector writers. SQLite +// uses a process-local lock because the archive ownership lock only excludes +// other processes. PostgreSQL uses a schema-scoped advisory lock. func (s *Store) WithDocumentVectorOperationLock(ctx context.Context, operation func() error) (retErr error) { if operation == nil { return errors.New("document vector operation is required") } if !s.IsPostgreSQL() { + s.documentVectorOperationMu.Lock() + defer s.documentVectorOperationMu.Unlock() return operation() } conn, err := s.db.Conn(ctx) diff --git a/internal/store/document_vectors_test.go b/internal/store/document_vectors_test.go index 5677303fd..156762a1c 100644 --- a/internal/store/document_vectors_test.go +++ b/internal/store/document_vectors_test.go @@ -69,6 +69,47 @@ func TestDocumentVectorOperationLockSerializesPostgresWriters(t *testing.T) { requirements.NoError(<-secondDone) } +func TestDocumentVectorOperationLockSerializesSQLiteWriters(t *testing.T) { + requirements := require.New(t) + if store.IsPostgresURL(os.Getenv("MSGVAULT_TEST_DB")) { + t.Skip("SQLite-only process-local writer lock") + } + f := storetest.New(t) + firstEntered := make(chan struct{}) + releaseFirst := make(chan struct{}) + firstDone := make(chan error, 1) + go func() { + firstDone <- f.Store.WithDocumentVectorOperationLock(t.Context(), func() error { + close(firstEntered) + <-releaseFirst + return nil + }) + }() + <-firstEntered + + secondEntered := make(chan struct{}) + secondDone := make(chan error, 1) + go func() { + secondDone <- f.Store.WithDocumentVectorOperationLock(t.Context(), func() error { + close(secondEntered) + return nil + }) + }() + select { + case <-secondEntered: + requirements.FailNow("second document-vector writer entered while the first held the lock") + case <-time.After(100 * time.Millisecond): + } + close(releaseFirst) + requirements.NoError(<-firstDone) + select { + case <-secondEntered: + case <-time.After(2 * time.Second): + requirements.FailNow("second document-vector writer did not acquire the released lock") + } + requirements.NoError(<-secondDone) +} + func runDocumentVectorGenerationLifecycleContract(t *testing.T) { t.Helper() t.Run("consent usage progress and same-policy rebuild", testDocumentVectorOperationsState) diff --git a/internal/store/store.go b/internal/store/store.go index 61d863199..5a254691e 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -51,7 +51,8 @@ type Store struct { fts5Available bool // Whether FTS5 is available for full-text search closeCleanup func() - sqliteOptimizeMu sync.Mutex + sqliteOptimizeMu sync.Mutex + documentVectorOperationMu sync.Mutex // Test-only seams into migration, backfill, and transaction paths, nil in // production and settable only from export_test.go. They belong to the diff --git a/internal/vector/document/worker_test.go b/internal/vector/document/worker_test.go index c713e1aad..5297ac7c3 100644 --- a/internal/vector/document/worker_test.go +++ b/internal/vector/document/worker_test.go @@ -453,10 +453,12 @@ func TestWorkerRunMakesMalformedHTTPProviderResponsesTerminal(t *testing.T) { } for _, test := range tests { t.Run(test.name, func(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") _, err := w.Write([]byte(test.response)) - assert.NoError(t, err) + assertions.NoError(err) })) t.Cleanup(server.Close) ledger := newFakeDocumentVectorLedger(workerClaim("extract-a", 1, "first", "token-a")) @@ -467,12 +469,12 @@ func TestWorkerRunMakesMalformedHTTPProviderResponsesTerminal(t *testing.T) { result, err := worker.Run(t.Context(), 1, 1) - require.Error(t, err) - assert.Equal(t, 1, result.Terminal) - assert.Zero(t, result.Retry) - require.Len(t, ledger.failures, 1) - assert.True(t, ledger.failures[0].terminal) - assert.Equal(t, test.wantCode, ledger.failures[0].errorCode) + requirements.Error(err) + assertions.Equal(1, result.Terminal) + assertions.Zero(result.Retry) + requirements.Len(ledger.failures, 1) + assertions.True(ledger.failures[0].terminal) + assertions.Equal(test.wantCode, ledger.failures[0].errorCode) }) } } From 6e34a3df1c1a61b839e70205c757e262f97c89e7 Mon Sep 17 00:00:00 2001 From: Wes McKinney Date: Mon, 24 Aug 2026 15:25:02 -0500 Subject: [PATCH 3/4] fix(documents): reject searches across authority changes Semantic candidate collection spans provider and backend calls. An activation or reindex during that window can apply new authority to old hits and silently return a partial response. Revalidate the active generation, target, fingerprint, and document revision after collection. Return the existing stale-cursor signal so clients can retry deliberately without automatic duplicate query egress. Generated with Codex Co-authored-by: Codex --- internal/vector/document/search.go | 33 +++++++++++++ internal/vector/document/search_test.go | 63 +++++++++++++++++++++++++ 2 files changed, 96 insertions(+) diff --git a/internal/vector/document/search.go b/internal/vector/document/search.go index c38918c94..ab2715a46 100644 --- a/internal/vector/document/search.go +++ b/internal/vector/document/search.go @@ -144,6 +144,9 @@ func (s *SearchService) Search(ctx context.Context, request store.DocumentSearch default: return store.DocumentSearchResponse{}, fmt.Errorf("%w: unsupported effective mode", store.ErrDocumentSearchInvalidRequest) } + if err := s.validateSearchAuthorityUnchanged(ctx, *generation, revision); err != nil { + return store.DocumentSearchResponse{}, err + } digest, err := digestSearchCandidates(candidates) if err != nil { return store.DocumentSearchResponse{}, err @@ -176,6 +179,36 @@ func (s *SearchService) Search(ctx context.Context, request store.DocumentSearch return response, nil } +func (s *SearchService) validateSearchAuthorityUnchanged( + ctx context.Context, + expected store.DocumentVectorGeneration, + expectedRevision int64, +) error { + active, err := s.deps.Ledger.GetActiveDocumentVectorGeneration(ctx) + if err != nil { + return fmt.Errorf("revalidate active document vector generation: %w", err) + } + if active == nil || active.ID != expected.ID || active.Fingerprint != expected.Fingerprint || + active.TargetExtractionProfileID != expected.TargetExtractionProfileID { + return store.ErrDocumentSearchCursorStale + } + target, err := s.deps.Ledger.GetDocumentVectorTargetProfileID(ctx) + if errors.Is(err, store.ErrDocumentVectorInvalidGenerationState) { + return store.ErrDocumentSearchCursorStale + } + if err != nil { + return fmt.Errorf("revalidate document vector target profile: %w", err) + } + revision, err := s.deps.Ledger.GetDocumentIndexRevision(ctx) + if err != nil { + return fmt.Errorf("revalidate document search revision: %w", err) + } + if target != expected.TargetExtractionProfileID || revision != expectedRevision { + return store.ErrDocumentSearchCursorStale + } + return nil +} + func (s *SearchService) effectiveMode(requested SearchMode, generation *store.DocumentVectorGeneration) (SearchMode, error) { semanticReady := generation != nil && s.deps.Embedder != nil && s.deps.Backend != nil && (s.deps.ExpectedFingerprint == "" || generation.Fingerprint == s.deps.ExpectedFingerprint) diff --git a/internal/vector/document/search_test.go b/internal/vector/document/search_test.go index 3eb5f6964..2ce611c6d 100644 --- a/internal/vector/document/search_test.go +++ b/internal/vector/document/search_test.go @@ -24,6 +24,35 @@ type countingSearchLedger struct { searchCalls int } +type changingSearchAuthorityLedger struct { + SearchLedger + + generationAfterCollection *store.DocumentVectorGeneration + revisionAfterCollection bool + generationReads int + revisionReads int +} + +func (ledger *changingSearchAuthorityLedger) GetActiveDocumentVectorGeneration( + ctx context.Context, +) (*store.DocumentVectorGeneration, error) { + ledger.generationReads++ + if ledger.generationReads > 1 && ledger.generationAfterCollection != nil { + generation := *ledger.generationAfterCollection + return &generation, nil + } + return ledger.SearchLedger.GetActiveDocumentVectorGeneration(ctx) +} + +func (ledger *changingSearchAuthorityLedger) GetDocumentIndexRevision(ctx context.Context) (int64, error) { + ledger.revisionReads++ + revision, err := ledger.SearchLedger.GetDocumentIndexRevision(ctx) + if err == nil && ledger.revisionReads > 1 && ledger.revisionAfterCollection { + revision++ + } + return revision, err +} + func (ledger *countingSearchLedger) SearchDocuments( ctx context.Context, request store.DocumentSearchRequest, @@ -129,6 +158,40 @@ func TestSearchServiceSemanticReturnsAuthoritativeOccurrenceProvenance(t *testin assertions.Equal(semanticBackendPageSize, backend.searches[0].k) } +func TestSearchServiceRejectsAuthorityChangesDuringCandidateCollection(t *testing.T) { + fixture := seedSemanticSearch(t, "authority evidence") + changedGeneration := fixture.generation + changedGeneration.ID++ + changedGeneration.Fingerprint = strings.Repeat("e", 64) + for _, test := range []struct { + name string + generationAfter *store.DocumentVectorGeneration + revisionAfter bool + }{ + {name: "generation activation", generationAfter: &changedGeneration}, + {name: "document reindex", revisionAfter: true}, + } { + t.Run(test.name, func(t *testing.T) { + assertions := assert.New(t) + requirements := require.New(t) + ledger := &changingSearchAuthorityLedger{ + SearchLedger: fixture.store.Store, generationAfterCollection: test.generationAfter, + revisionAfterCollection: test.revisionAfter, + } + embedder := &searchQueryEmbedder{vector: []float32{1, 0, 0}} + backend := &searchBackend{hits: []Hit{{Token: fixture.claims[0].Token, Score: .9, Rank: 1}}} + service := NewSearchService(SearchDeps{Ledger: ledger, Embedder: embedder, Backend: backend}) + + _, err := service.Search(t.Context(), store.DocumentSearchRequest{ + Query: "authority", SearchMode: string(SearchModeSemantic), CandidateLimit: 10, + }) + + requirements.ErrorIs(err, store.ErrDocumentSearchCursorStale) + assertions.Equal([]string{"authority"}, embedder.queries) + }) + } +} + func TestSearchServiceSemanticPagesPastOutOfScopeGlobalHits(t *testing.T) { assert := assert.New(t) require := require.New(t) From d014222e5424f559e58a958c70c1ceef78bce091 Mon Sep 17 00:00:00 2001 From: Wes McKinney Date: Tue, 25 Aug 2026 08:25:57 -0500 Subject: [PATCH 4/4] fix(documents): keep vector maintenance bounded Retired generations must remain removable after an operator disables or changes the embedding provider. Open only the existing vector backend for retired cleanup so provider setup cannot strand obsolete storage. Normalized evidence is reloaded for each bounded worker pass. Fetch all chunk spans in one query so large documents do not multiply database round trips. The PostgreSQL test lane creates many schemas concurrently. Give its stock service a larger test-only lock budget so the expanded schema does not exhaust the default lock table. Generated with Codex Co-authored-by: Codex --- .github/workflows/ci.yml | 1 + ...documents_vector_cleanup_sqlitevec_test.go | 63 +++++++++++++++++++ cmd/msgvault/cmd/documents_vector_runtime.go | 38 +++++++++++ internal/store/document_normalized.go | 47 ++++++++------ internal/store/document_publication_test.go | 25 +++++--- internal/vector/pgvector/document_vectors.go | 10 +++ .../vector/pgvector/document_vectors_stub.go | 3 + 7 files changed, 159 insertions(+), 28 deletions(-) create mode 100644 cmd/msgvault/cmd/documents_vector_cleanup_sqlitevec_test.go diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 08e3e6c08..3e30d455e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -506,6 +506,7 @@ jobs: POSTGRES_USER: msgvault_test POSTGRES_PASSWORD: msgvault_test POSTGRES_DB: msgvault_test + POSTGRES_INITDB_ARGS: "-c max_locks_per_transaction=256" ports: - 5432:5432 options: >- diff --git a/cmd/msgvault/cmd/documents_vector_cleanup_sqlitevec_test.go b/cmd/msgvault/cmd/documents_vector_cleanup_sqlitevec_test.go new file mode 100644 index 000000000..9e7dfeb6e --- /dev/null +++ b/cmd/msgvault/cmd/documents_vector_cleanup_sqlitevec_test.go @@ -0,0 +1,63 @@ +//go:build sqlite_vec + +package cmd + +import ( + "path/filepath" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + vectordocument "go.kenn.io/msgvault/internal/vector/document" + "go.kenn.io/msgvault/internal/vector/sqlitevec" +) + +func TestRunConfiguredDocumentVectorGenerationCleansRetiredWhenEmbeddingsDisabled(t *testing.T) { + fixture, spec := documentVectorCommandFixture(t) + vectorPath := filepath.Join(t.TempDir(), "vectors.db") + cfg.Vector.DBPath = vectorPath + + generation, _, err := fixture.Store.EnsureDocumentVectorGeneration(t.Context(), spec) + require.NoError(t, err) + token := strings.Repeat("9", 64) + backend, err := sqlitevec.Open(t.Context(), sqlitevec.Options{ + Path: vectorPath, Dimension: spec.Dimension, + }) + require.NoError(t, err) + require.NoError(t, backend.DocumentBackend().PutUnpublished(t.Context(), vectordocument.GenerationID(generation.ID), spec.Dimension, []vectordocument.Embedding{{ + Token: token, Vector: []float32{1, 0, 0}, + }})) + require.NoError(t, backend.Close()) + + _, err = fixture.Store.DB().Exec(fixture.Store.Rebind(` + INSERT INTO document_vector_publications + (generation_id, extraction_id, extraction_profile_id, canonical_blob_hash, + extraction_input_key, chunk_id, chunk_key, chunk_checksum, source_sequence, token, state) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'ready')`), generation.ID, + "disabled-cleanup-extraction", spec.TargetExtractionProfileID, strings.Repeat("a", 64), + "original", 1, "disabled-cleanup-chunk", "disabled-cleanup-checksum", 1, token) + require.NoError(t, err) + retired, err := fixture.Store.RetireDocumentVectorGeneration(t.Context(), generation.ID, time.Now()) + require.NoError(t, err) + require.True(t, retired) + cfg.Vector.Enabled = false + + result, err := runConfiguredDocumentVectorGeneration(t.Context(), fixture.Store, generation.ID, 1) + + require.NoError(t, err) + assert.True(t, result.Purged) + assert.True(t, result.Converged) + _, err = fixture.Store.GetDocumentVectorGeneration(t.Context(), generation.ID) + require.ErrorContains(t, err, "not found") + + backend, err = sqlitevec.Open(t.Context(), sqlitevec.Options{Path: vectorPath}) + require.NoError(t, err) + t.Cleanup(func() { _ = backend.Close() }) + var remaining int + require.NoError(t, backend.DB().QueryRow( + `SELECT COUNT(*) FROM document_vector_embeddings WHERE token = ?`, token, + ).Scan(&remaining)) + assert.Zero(t, remaining) +} diff --git a/cmd/msgvault/cmd/documents_vector_runtime.go b/cmd/msgvault/cmd/documents_vector_runtime.go index 460039868..d473bb749 100644 --- a/cmd/msgvault/cmd/documents_vector_runtime.go +++ b/cmd/msgvault/cmd/documents_vector_runtime.go @@ -5,12 +5,16 @@ package cmd import ( "context" "errors" + "fmt" + "path/filepath" "time" "github.com/google/uuid" "go.kenn.io/msgvault/internal/store" "go.kenn.io/msgvault/internal/vector" vectordocument "go.kenn.io/msgvault/internal/vector/document" + "go.kenn.io/msgvault/internal/vector/pgvector" + "go.kenn.io/msgvault/internal/vector/sqlitevec" ) func nextDocumentVectorWorkerOwner() string { @@ -47,6 +51,21 @@ func runConfiguredDocumentVectorGeneration(ctx context.Context, st *store.Store, if limit < 1 || limit > 1000 { return vectordocument.ReconcileResult{}, errors.New("document vector operation limit must be between 1 and 1000") } + generation, err := st.GetDocumentVectorGeneration(ctx, generationID) + if err != nil { + return vectordocument.ReconcileResult{}, err + } + if generation.State == store.DocumentVectorGenerationRetired { + backend, closeBackend, err := openDocumentVectorCleanupBackend(ctx, st, cfg.DatabaseDSN()) + if err != nil { + return vectordocument.ReconcileResult{}, err + } + defer func() { _ = closeBackend() }() + reconciler := vectordocument.NewReconciler(vectordocument.ReconcilerDeps{ + Ledger: st, Backend: backend, Now: func() time.Time { return time.Now().UTC() }, + }) + return reconciler.Run(ctx, vectordocument.GenerationID(generationID), limit) + } vf, err := setupVectorFeatures(ctx, st, cfg.DatabaseDSN(), false) if err != nil { return vectordocument.ReconcileResult{}, err @@ -58,6 +77,25 @@ func runConfiguredDocumentVectorGeneration(ctx context.Context, st *store.Store, return runDocumentVectorWithFeatures(ctx, st, vf, generationID, limit) } +func openDocumentVectorCleanupBackend(ctx context.Context, st *store.Store, mainPath string) (vectordocument.Backend, func() error, error) { + if store.IsPostgresURL(mainPath) { + backend, err := pgvector.DocumentBackendForDB(st.DB()) + if err != nil { + return nil, nil, fmt.Errorf("open pgvector document cleanup backend: %w", err) + } + return backend, func() error { return nil }, nil + } + vectorPath := cfg.Vector.DBPath + if vectorPath == "" { + vectorPath = filepath.Join(cfg.Data.DataDir, "vectors.db") + } + backend, err := sqlitevec.Open(ctx, sqlitevec.Options{Path: vectorPath}) + if err != nil { + return nil, nil, fmt.Errorf("open document cleanup backend: %w", err) + } + return backend.DocumentBackend(), backend.Close, nil +} + func runDocumentVectorWithFeatures(ctx context.Context, st *store.Store, vf *vectorFeatures, generationID int64, limit int) (vectordocument.ReconcileResult, error) { limit = min(limit, max(1, vf.Cfg.Embeddings.BatchSize)) generation, err := st.GetDocumentVectorGeneration(ctx, generationID) diff --git a/internal/store/document_normalized.go b/internal/store/document_normalized.go index ecbc5e2a5..a5a200b3e 100644 --- a/internal/store/document_normalized.go +++ b/internal/store/document_normalized.go @@ -130,30 +130,37 @@ func (s *Store) LoadNormalizedDocument(ctx context.Context, extractionID string) if err := chunkRows.Close(); err != nil { return document.NormalizedDocument{}, fmt.Errorf("close normalized document chunks: %w", err) } + chunkIndexes := make(map[string]int, len(normalized.Chunks)) for index := range normalized.Chunks { - chunk := &normalized.Chunks[index] - spanRows, err := s.db.QueryContext(ctx, s.Rebind(` - SELECT unit_index, start_char, end_char - FROM document_chunk_spans - WHERE extraction_id = ? AND chunk_key = ? ORDER BY span_ordinal`), extractionID, chunk.Key) - if err != nil { - return document.NormalizedDocument{}, fmt.Errorf("read normalized document chunk spans: %w", err) - } - for spanRows.Next() { - var span document.ChunkSpan - if err := spanRows.Scan(&span.UnitIndex, &span.CharStart, &span.CharEnd); err != nil { - _ = spanRows.Close() - return document.NormalizedDocument{}, fmt.Errorf("scan normalized document chunk span: %w", err) - } - chunk.Spans = append(chunk.Spans, span) - } - if err := spanRows.Err(); err != nil { + chunkIndexes[normalized.Chunks[index].Key] = index + } + spanRows, err := s.db.QueryContext(ctx, s.Rebind(` + SELECT chunk_key, unit_index, start_char, end_char + FROM document_chunk_spans + WHERE extraction_id = ? ORDER BY chunk_key, span_ordinal`), extractionID) + if err != nil { + return document.NormalizedDocument{}, fmt.Errorf("read normalized document chunk spans: %w", err) + } + for spanRows.Next() { + var chunkKey string + var span document.ChunkSpan + if err := spanRows.Scan(&chunkKey, &span.UnitIndex, &span.CharStart, &span.CharEnd); err != nil { _ = spanRows.Close() - return document.NormalizedDocument{}, fmt.Errorf("iterate normalized document chunk spans: %w", err) + return document.NormalizedDocument{}, fmt.Errorf("scan normalized document chunk span: %w", err) } - if err := spanRows.Close(); err != nil { - return document.NormalizedDocument{}, fmt.Errorf("close normalized document chunk spans: %w", err) + index, ok := chunkIndexes[chunkKey] + if !ok { + _ = spanRows.Close() + return document.NormalizedDocument{}, fmt.Errorf("normalized document chunk span references unknown chunk %q", chunkKey) } + normalized.Chunks[index].Spans = append(normalized.Chunks[index].Spans, span) + } + if err := spanRows.Err(); err != nil { + _ = spanRows.Close() + return document.NormalizedDocument{}, fmt.Errorf("iterate normalized document chunk spans: %w", err) + } + if err := spanRows.Close(); err != nil { + return document.NormalizedDocument{}, fmt.Errorf("close normalized document chunk spans: %w", err) } if err := document.ValidateNormalizedDocument(normalized); err != nil { return document.NormalizedDocument{}, fmt.Errorf("validate stored normalized document: %w", err) diff --git a/internal/store/document_publication_test.go b/internal/store/document_publication_test.go index 0c2b0c920..5ded92c36 100644 --- a/internal/store/document_publication_test.go +++ b/internal/store/document_publication_test.go @@ -28,10 +28,11 @@ func TestDocumentExtractionPublicationRoundTripsNormalizedV3Identity(t *testing. require.NoError(err) normalized, err := docbankdocument.NormalizeDocument(docbankdocument.SourceDocument{ Family: "pdf", UnitKind: "page", Units: []docbankdocument.SourceUnit{{ - Index: 0, Markdown: "# Evidence\n\nStored identity", + Index: 0, Markdown: "# Evidence\n\n" + strings.Repeat("stored identity evidence ", 300), }}, }, policy) require.NoError(err) + require.Greater(len(normalized.Chunks), 1) publication := publicationFor(t, claim, normalized.Chunks[0].Text, normalized.Chunks[0].Checksum) publication.ManifestChecksum = normalized.Checksum publication.NormalizationVersion = normalized.PolicyVersion @@ -46,20 +47,28 @@ func TestDocumentExtractionPublicationRoundTripsNormalizedV3Identity(t *testing. CharCount: normalized.Units[0].CharCount, Truncated: normalized.Units[0].Truncated, HeadingMarks: normalized.Units[0].HeadingMarks, } - publication.Chunks[0].Key = normalized.Chunks[0].Key - publication.Chunks[0].HeadingPath = normalized.Chunks[0].HeadingPath - publication.Chunks[0].Truncated = normalized.Chunks[0].Truncated - publication.Chunks[0].Spans[0] = store.DocumentPublishedSpan{ - UnitIndex: normalized.Chunks[0].Spans[0].UnitIndex, - CharStart: normalized.Chunks[0].Spans[0].CharStart, - CharEnd: normalized.Chunks[0].Spans[0].CharEnd, + publication.Chunks = make([]store.DocumentPublishedChunk, len(normalized.Chunks)) + for index, chunk := range normalized.Chunks { + publication.Chunks[index] = store.DocumentPublishedChunk{ + Key: chunk.Key, Ordinal: chunk.Ordinal, Text: chunk.Text, + HeadingPath: chunk.HeadingPath, Checksum: chunk.Checksum, + CharCount: chunk.CharCount, Truncated: chunk.Truncated, + Spans: []store.DocumentPublishedSpan{{ + UnitIndex: chunk.Spans[0].UnitIndex, + CharStart: chunk.Spans[0].CharStart, + CharEnd: chunk.Spans[0].CharEnd, + }}, + } } require.NoError(f.Store.PublishDocumentExtraction(t.Context(), publication)) + logs := captureAttachmentQueryLogs(t) loaded, err := f.Store.LoadNormalizedDocument(t.Context(), claim.ExtractionID) require.NoError(err) assert.Equal(normalized, loaded) + assert.Equal(1, strings.Count(logs.String(), "FROM document_chunk_spans"), + "normalized document loading must fetch every chunk span in one query") } func TestDocumentExtractionPublicationKeepsOldHeadUntilAtomicSwitch(t *testing.T) { diff --git a/internal/vector/pgvector/document_vectors.go b/internal/vector/pgvector/document_vectors.go index fb012bbc3..e443362f1 100644 --- a/internal/vector/pgvector/document_vectors.go +++ b/internal/vector/pgvector/document_vectors.go @@ -33,6 +33,16 @@ func (b *Backend) DocumentBackend() *DocumentBackend { return &DocumentBackend{db: b.db} } +// DocumentBackendForDB returns a non-owning document-vector backend for an +// already initialized PostgreSQL database. It is used by cleanup operations +// that must not initialize an embedding provider or rerun backend migrations. +func DocumentBackendForDB(db *sql.DB) (*DocumentBackend, error) { + if db == nil { + return nil, errors.New("pgvector document backend database is required") + } + return &DocumentBackend{db: db}, nil +} + func (b *DocumentBackend) PutUnpublished(ctx context.Context, generationID document.GenerationID, dimension int, embeddings []document.Embedding) error { if err := validatePGDocumentPut(generationID, dimension, embeddings); err != nil { return err diff --git a/internal/vector/pgvector/document_vectors_stub.go b/internal/vector/pgvector/document_vectors_stub.go index 21b69509b..06d6ba3dc 100644 --- a/internal/vector/pgvector/document_vectors_stub.go +++ b/internal/vector/pgvector/document_vectors_stub.go @@ -4,6 +4,7 @@ package pgvector import ( "context" + "database/sql" "go.kenn.io/msgvault/internal/vector/document" ) @@ -15,6 +16,8 @@ var _ document.Backend = (*DocumentBackend)(nil) func (b *Backend) DocumentBackend() *DocumentBackend { return &DocumentBackend{} } +func DocumentBackendForDB(*sql.DB) (*DocumentBackend, error) { return nil, ErrNotBuilt } + func (b *DocumentBackend) PutUnpublished(context.Context, document.GenerationID, int, []document.Embedding) error { return ErrNotBuilt }