nanobot: 2026-09-10 12:33:37
This commit is contained in:
127
skills/wiki/scripts/wiki_embed.py
Normal file
127
skills/wiki/scripts/wiki_embed.py
Normal file
@@ -0,0 +1,127 @@
|
||||
#!/usr/bin/env python3
|
||||
# /// script
|
||||
# requires-python = ">=3.11"
|
||||
# dependencies = ["requests", "pyyaml", "markdown-it-py", "sqlite-vec==0.1.6"]
|
||||
# ///
|
||||
"""Embedding client and the identity guard over the vector space.
|
||||
|
||||
Two invariants live here:
|
||||
|
||||
* Vectors are stored L2-normalized, so cosine distance equals the dot product.
|
||||
* The query prefix must be bit-identical at index and at query time. That is the real
|
||||
reason it is persisted in `meta` rather than only read from the config — mixing
|
||||
vectors produced under two different contracts degrades results silently, which is
|
||||
the most expensive kind of bug.
|
||||
|
||||
An unreachable Ollama is not an error here: the caller stores chunks with
|
||||
`embedded_at IS NULL` and the query side says out loud that it is FTS-only.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
|
||||
import requests
|
||||
from wiki_chunker import CHUNKER_VERSION
|
||||
from wiki_config import EmbeddingConfig
|
||||
from wiki_db import EMBEDDING_DIMS
|
||||
|
||||
SCHEMA_VERSION = "1"
|
||||
SQLITE_VEC_VERSION = "0.1.6"
|
||||
REQUEST_TIMEOUT_SECONDS = 60
|
||||
|
||||
# Only these keys make two vectors comparable; the rest of `meta` is informational.
|
||||
GUARDED_META_KEYS = (
|
||||
"embedding_model",
|
||||
"embedding_dims",
|
||||
"normalized",
|
||||
"query_prefix",
|
||||
"chunker_version",
|
||||
)
|
||||
|
||||
|
||||
class EmbeddingUnavailable(RuntimeError):
|
||||
"""Ollama could not be reached or refused the request."""
|
||||
|
||||
|
||||
class IndexIdentityMismatch(RuntimeError):
|
||||
"""The index was built under a different embedding contract — a reindex is needed."""
|
||||
|
||||
|
||||
def expected_meta(config: EmbeddingConfig) -> dict[str, str]:
|
||||
return {
|
||||
"embedding_model": config.model,
|
||||
"embedding_dims": str(config.dims),
|
||||
"normalized": "l2",
|
||||
"query_prefix": config.query_prefix,
|
||||
"chunker_version": CHUNKER_VERSION,
|
||||
"schema_version": SCHEMA_VERSION,
|
||||
"sqlite_vec_version": SQLITE_VEC_VERSION,
|
||||
}
|
||||
|
||||
|
||||
def meta_mismatches(stored: dict[str, str], config: EmbeddingConfig) -> list[str]:
|
||||
"""Guarded meta keys that disagree with the config. Empty on a fresh (unwritten) index."""
|
||||
if not stored:
|
||||
return []
|
||||
expected = expected_meta(config)
|
||||
return [key for key in GUARDED_META_KEYS if stored.get(key) != expected[key]]
|
||||
|
||||
|
||||
def require_matching_index(stored: dict[str, str], config: EmbeddingConfig) -> None:
|
||||
"""Refuse to query an index built under a different contract."""
|
||||
if config.dims != EMBEDDING_DIMS:
|
||||
raise IndexIdentityMismatch(f"reindex needed: config dims {config.dims} != schema dims {EMBEDDING_DIMS}")
|
||||
mismatches = meta_mismatches(stored, config)
|
||||
if mismatches:
|
||||
detail = ", ".join(
|
||||
f"{key}: index={stored.get(key)!r} config={expected_meta(config)[key]!r}" for key in mismatches
|
||||
)
|
||||
raise IndexIdentityMismatch(f"reindex needed: {detail}")
|
||||
|
||||
|
||||
def l2_normalize(vector: Sequence[float]) -> list[float]:
|
||||
norm = math.sqrt(sum(value * value for value in vector))
|
||||
if norm == 0.0:
|
||||
return list(vector)
|
||||
return [value / norm for value in vector]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OllamaEmbedder:
|
||||
config: EmbeddingConfig
|
||||
|
||||
def embed_documents(self, texts: Sequence[str]) -> list[list[float]]:
|
||||
"""Documents are embedded without the instruct prefix (the model is asymmetric)."""
|
||||
return self._embed(list(texts))
|
||||
|
||||
def embed_query(self, text: str) -> list[float]:
|
||||
return self._embed([self.config.query_prefix + text])[0]
|
||||
|
||||
def probe(self) -> None:
|
||||
"""Raise EmbeddingUnavailable unless the endpoint answers an embed request."""
|
||||
self._embed(["ping"])
|
||||
|
||||
def _embed(self, inputs: list[str]) -> list[list[float]]:
|
||||
if not inputs:
|
||||
return []
|
||||
payload = {
|
||||
"model": self.config.model,
|
||||
"input": inputs,
|
||||
"keep_alive": self.config.keep_alive,
|
||||
}
|
||||
try:
|
||||
response = requests.post(f"{self.config.endpoint}/api/embed", json=payload, timeout=REQUEST_TIMEOUT_SECONDS)
|
||||
response.raise_for_status()
|
||||
embeddings = response.json()["embeddings"]
|
||||
except (requests.RequestException, KeyError, ValueError) as exc:
|
||||
raise EmbeddingUnavailable(f"{self.config.endpoint}: {exc}") from exc
|
||||
|
||||
if len(embeddings) != len(inputs):
|
||||
raise EmbeddingUnavailable(f"asked for {len(inputs)} vectors, got {len(embeddings)}")
|
||||
for vector in embeddings:
|
||||
if len(vector) != self.config.dims:
|
||||
raise EmbeddingUnavailable(f"model returned {len(vector)} dims, config says {self.config.dims}")
|
||||
return [l2_normalize(vector) for vector in embeddings]
|
||||
Reference in New Issue
Block a user