feat(embedding): support boundary and corpus model roles

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
2026-08-10 14:12:45 +03:30
parent 0ff9d8dd21
commit 736391b137

View File

@@ -1,68 +1,219 @@
"""OpenAI embedding service using text-embedding-3-small. """Embedding service with Cloud (OpenAI) and Local (Ollama) Providers.
All strategies share the same embedding model (fixed, not configurable) ADR-0024: two roles — Boundary (semantic cuts) and Corpus (storage + query).
to ensure fair comparison. Batch support up to 2048 texts per call. Callers snapshot both at operation start so a mid-flight Admin switch cannot mix models.
Legacy Active Embedding Model maps to Corpus (and migrates into both defaults).
""" """
from __future__ import annotations from __future__ import annotations
import logging import logging
from typing import Literal
from openai import OpenAI from src.chunking.embedding_models import (
EmbeddingModelSpec,
from src.core.config import settings Provider,
from src.core.dependencies import get_openai_client apply_task_prefixes,
default_model_id,
get_model,
)
from src.core.dependencies import get_ollama_client, get_openai_client
from src.core.exceptions import EmbeddingError from src.core.exceptions import EmbeddingError
from src.storage import sqlite as db
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# OpenAI batch limit for text-embedding-3-small _OPENAI_BATCH_SIZE = 2048
_BATCH_SIZE = 2048 _OLLAMA_BATCH_SIZE = 64
_CORPUS_SETTING = "corpus_embedding_model_id"
_BOUNDARY_SETTING = "boundary_embedding_model_id"
_LEGACY_ACTIVE_SETTING = "active_embedding_model_id"
Purpose = Literal["document", "query"]
def embed_texts(texts: list[str]) -> list[list[float]]: def _ensure_role_defaults_migrated() -> None:
"""Embed a list of texts and return their vectors. """One-shot: legacy Active → Corpus + Boundary when role keys unset."""
corpus = db.get_app_setting(_CORPUS_SETTING)
boundary = db.get_app_setting(_BOUNDARY_SETTING)
if corpus and boundary:
return
legacy = db.get_app_setting(_LEGACY_ACTIVE_SETTING) or default_model_id()
if not corpus:
db.set_app_setting(_CORPUS_SETTING, legacy)
if not boundary:
db.set_app_setting(_BOUNDARY_SETTING, legacy)
For the contextual_structure strategy, these are the enriched texts
(not raw content) — this is by design (ADR 0011).
Args: def _resolve_role(setting_key: str) -> EmbeddingModelSpec:
texts: List of strings to embed. _ensure_role_defaults_migrated()
stored = db.get_app_setting(setting_key)
model_id = stored or default_model_id()
try:
return get_model(model_id)
except KeyError:
logger.warning(
"Stored Embedding Model '%s' (%s) not in registry; falling back to %s",
model_id,
setting_key,
default_model_id(),
)
return get_model(default_model_id())
Returns:
List of embedding vectors (same order as input).
Raises: def get_corpus_embedding_model() -> EmbeddingModelSpec:
EmbeddingError: If the OpenAI API call fails. """Default Corpus Embedding Model (storage + query)."""
return _resolve_role(_CORPUS_SETTING)
def get_boundary_embedding_model() -> EmbeddingModelSpec:
"""Default Boundary Embedding Model (semantic cuts)."""
return _resolve_role(_BOUNDARY_SETTING)
def set_corpus_embedding_model(model_id: str) -> EmbeddingModelSpec:
"""Persist Default Corpus. Raises KeyError if unknown."""
_ensure_role_defaults_migrated()
model = get_model(model_id)
db.set_app_setting(_CORPUS_SETTING, model.id)
# Keep legacy key in sync for older readers
db.set_app_setting(_LEGACY_ACTIVE_SETTING, model.id)
logger.info("Corpus Embedding Model set to %s (%s)", model.id, model.provider.value)
return model
def set_boundary_embedding_model(model_id: str) -> EmbeddingModelSpec:
"""Persist Default Boundary. Raises KeyError if unknown."""
_ensure_role_defaults_migrated()
model = get_model(model_id)
db.set_app_setting(_BOUNDARY_SETTING, model.id)
logger.info("Boundary Embedding Model set to %s (%s)", model.id, model.provider.value)
return model
def resolve_corpus_model(model_id: str | None = None) -> EmbeddingModelSpec:
"""Snapshot Corpus for an operation (explicit id or Admin default)."""
if model_id:
return get_model(model_id)
return get_corpus_embedding_model()
def resolve_boundary_model(model_id: str | None = None) -> EmbeddingModelSpec:
"""Snapshot Boundary for an operation (explicit id or Admin default)."""
if model_id:
return get_model(model_id)
return get_boundary_embedding_model()
# ── Legacy aliases (Corpus) ───────────────────────────────────────
def get_active_embedding_model() -> EmbeddingModelSpec:
"""Deprecated: alias for get_corpus_embedding_model (ADR-0024)."""
return get_corpus_embedding_model()
def set_active_embedding_model(model_id: str) -> EmbeddingModelSpec:
"""Deprecated: sets Corpus default (and legacy active key)."""
return set_corpus_embedding_model(model_id)
def snapshot_active_model() -> EmbeddingModelSpec:
"""Deprecated: snapshot Corpus Embedding Model."""
return get_corpus_embedding_model()
def embed_texts(
texts: list[str],
*,
model: EmbeddingModelSpec | None = None,
purpose: Purpose = "document",
) -> list[list[float]]:
"""Embed texts with the given (or Corpus) Embedding Model.
For contextual_retrieval, pass enriched texts — ADR 0011.
""" """
if not texts: if not texts:
return [] return []
model = model or get_corpus_embedding_model()
prepared = apply_task_prefixes(texts, model=model, purpose=purpose)
try:
if model.provider == Provider.CLOUD:
return _embed_openai(prepared, model)
if model.provider == Provider.LOCAL:
return _embed_ollama(prepared, model)
raise EmbeddingError(f"Unsupported Provider: {model.provider}")
except EmbeddingError:
raise
except Exception as exc:
raise EmbeddingError(f"Embedding failed ({model.id}): {exc}") from exc
def embed_single(
text: str,
*,
model: EmbeddingModelSpec | None = None,
purpose: Purpose = "query",
) -> list[float]:
"""Embed a single text (default purpose=query for retrieval)."""
results = embed_texts([text], model=model, purpose=purpose)
return results[0]
def _embed_openai(texts: list[str], model: EmbeddingModelSpec) -> list[list[float]]:
client = get_openai_client() client = get_openai_client()
all_embeddings: list[list[float]] = [] all_embeddings: list[list[float]] = []
try: for start in range(0, len(texts), _OPENAI_BATCH_SIZE):
for start in range(0, len(texts), _BATCH_SIZE): batch = texts[start : start + _OPENAI_BATCH_SIZE]
batch = texts[start:start + _BATCH_SIZE]
response = client.embeddings.create( response = client.embeddings.create(
model=settings.embedding_model, model=model.model_name,
input=batch, input=batch,
) )
# Sort by index to guarantee order matches input
sorted_data = sorted(response.data, key=lambda x: x.index) sorted_data = sorted(response.data, key=lambda x: x.index)
all_embeddings.extend([item.embedding for item in sorted_data]) vectors = [item.embedding for item in sorted_data]
_validate_dimensions(vectors, model)
all_embeddings.extend(vectors)
logger.debug( logger.debug(
"Embedded batch %d-%d (%d texts)", "OpenAI embedded batch %d-%d with %s",
start, start + len(batch), len(batch), start,
start + len(batch),
model.model_name,
) )
return all_embeddings return all_embeddings
except Exception as exc:
raise EmbeddingError(f"Embedding failed: {exc}") from exc
def embed_single(text: str) -> list[float]: def _embed_ollama(texts: list[str], model: EmbeddingModelSpec) -> list[list[float]]:
"""Embed a single text (convenience wrapper).""" client = get_ollama_client()
results = embed_texts([text]) all_embeddings: list[list[float]] = []
return results[0]
for start in range(0, len(texts), _OLLAMA_BATCH_SIZE):
batch = texts[start : start + _OLLAMA_BATCH_SIZE]
response = client.embeddings.create(
model=model.model_name,
input=batch,
)
sorted_data = sorted(response.data, key=lambda x: x.index)
vectors = [item.embedding for item in sorted_data]
_validate_dimensions(vectors, model)
all_embeddings.extend(vectors)
logger.debug(
"Ollama embedded batch %d-%d with %s",
start,
start + len(batch),
model.model_name,
)
return all_embeddings
def _validate_dimensions(vectors: list[list[float]], model: EmbeddingModelSpec) -> None:
for i, vec in enumerate(vectors):
if len(vec) != model.dimension:
raise EmbeddingError(
f"Embedding dimension mismatch for {model.id}: "
f"expected {model.dimension}, got {len(vec)} (index {i})"
)