embeddings_queries now records the embedding request's token count (the unit upstream providers bill on) instead of an operation count, and fires on the indexing path too. Previously only semantic search recorded it (value=1), so a re-indexing run produced no embeddings_queries events at all — only pages_chunks. - Provider layer: additive embed_with_usage / embed_batch_with_usage surface the per-request token count (Mistral/OpenAI usage.total_tokens, Bedrock Titan inputTextTokenCount, Ollama prompt_eval_count); a char-based estimate is the fallback (Simple, and any provider/response without a token field). Gateway and EmbeddingService forward through. The count travels as a return value / a per-request SearchAlgorithm attribute — never on the singleton — so concurrent indexing + search can't mis-attribute bills. - Indexing (vector/processor.py): records embeddings_queries (value=batch tokens) alongside the existing pages_chunks event. - Search (server/semantic.py): value is now the query embedding's token count, relayed from BM25HybridSearchAlgorithm via query_token_count. The astrolabe_embeddings_queries Stripe meter (sum aggregation) now sums tokens with no CP/Terraform change. The meter "queries"->tokens naming/unit clarification (homelab-terraform #254) + CP rollup/portal copy is a follow-up. Deck #67. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
243 lines
8.8 KiB
Python
243 lines
8.8 KiB
Python
"""Mistral provider for embeddings.
|
|
|
|
Currently supports embeddings only (``mistral-embed``, 1024-dim). Generation
|
|
can be added later if needed; see ADR-015.
|
|
"""
|
|
|
|
import logging
|
|
|
|
# mistralai 2.x ships no top-level __init__.py, so `from mistralai import …`
|
|
# raises ImportError. The canonical public paths are `mistralai.client` (which
|
|
# re-exports the SDK class via `client/__init__.py`) and `mistralai.client.errors`
|
|
# (which lazy-loads SDKError). There is no `mistralai.models` subpackage either.
|
|
from mistralai.client import Mistral
|
|
from mistralai.client.errors import SDKError
|
|
|
|
from ._retry import retry_on_rate_limit
|
|
from .base import Provider
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Well-known Mistral embedding model dimensions
|
|
MISTRAL_EMBEDDING_DIMENSIONS: dict[str, int] = {
|
|
"mistral-embed": 1024,
|
|
}
|
|
|
|
# Conservative chunk size for batch embeddings. Mistral allows large batches,
|
|
# but we keep this in line with sibling providers (OpenAI=100, Ollama=32).
|
|
BATCH_SIZE = 64
|
|
|
|
_NO_EMBEDDING_MODEL_MSG = "Embedding not supported - no embedding_model configured"
|
|
|
|
|
|
def _is_rate_limit(exc: BaseException) -> bool:
|
|
"""True only for HTTP 429 SDKErrors."""
|
|
return getattr(exc, "status_code", None) == 429
|
|
|
|
|
|
_retry_429 = retry_on_rate_limit(
|
|
SDKError, is_rate_limit=_is_rate_limit, provider_name="Mistral"
|
|
)
|
|
|
|
|
|
class MistralProvider(Provider):
|
|
"""
|
|
Mistral provider — embeddings only.
|
|
|
|
Uses the official ``mistralai`` SDK. Lazy dimension detection mirrors the
|
|
OpenAI provider: known models populate the cached dimension at construction
|
|
time; unknown models get their dimension detected on the first ``embed()``
|
|
call.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
api_key: str,
|
|
embedding_model: str | None = "mistral-embed",
|
|
base_url: str | None = None,
|
|
):
|
|
"""
|
|
Initialize the Mistral provider.
|
|
|
|
Args:
|
|
api_key: Mistral API key.
|
|
embedding_model: Embedding model ID (default: ``mistral-embed``).
|
|
Pass ``None`` to disable embeddings (the provider will then
|
|
support no capabilities, which is mostly useful for tests).
|
|
base_url: Optional base URL override (e.g. proxies, on-prem).
|
|
"""
|
|
self.embedding_model = embedding_model
|
|
self._dimension: int | None = None
|
|
|
|
self.client = Mistral(api_key=api_key, server_url=base_url)
|
|
|
|
if embedding_model and embedding_model in MISTRAL_EMBEDDING_DIMENSIONS:
|
|
self._dimension = MISTRAL_EMBEDDING_DIMENSIONS[embedding_model]
|
|
|
|
logger.info(
|
|
"Initialized Mistral provider: base_url=%s, embedding_model=%s, "
|
|
"dimension=%s",
|
|
base_url or "default",
|
|
embedding_model,
|
|
self._dimension,
|
|
)
|
|
|
|
@property
|
|
def supports_embeddings(self) -> bool:
|
|
return self.embedding_model is not None
|
|
|
|
@property
|
|
def supports_generation(self) -> bool:
|
|
return False
|
|
|
|
@_retry_429
|
|
async def embed(self, text: str) -> list[float]:
|
|
"""Generate an embedding for a single text."""
|
|
if not self.supports_embeddings:
|
|
raise NotImplementedError(_NO_EMBEDDING_MODEL_MSG)
|
|
|
|
assert self.embedding_model is not None
|
|
response = await self.client.embeddings.create_async(
|
|
model=self.embedding_model,
|
|
inputs=[text],
|
|
)
|
|
|
|
if not response.data or response.data[0].embedding is None:
|
|
raise RuntimeError(
|
|
f"Mistral embeddings API returned no embedding for model "
|
|
f"{self.embedding_model}"
|
|
)
|
|
|
|
embedding = response.data[0].embedding
|
|
|
|
if self._dimension is None:
|
|
self._dimension = len(embedding)
|
|
logger.info(
|
|
"Detected embedding dimension: %d for model %s",
|
|
self._dimension,
|
|
self.embedding_model,
|
|
)
|
|
|
|
return embedding
|
|
|
|
async def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
|
"""Generate embeddings for multiple texts, chunking by ``BATCH_SIZE``."""
|
|
embeddings, _ = await self.embed_batch_with_usage(texts)
|
|
return embeddings
|
|
|
|
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
|
|
"""Embed one text, reporting the Mistral request's token count."""
|
|
embeddings, tokens = await self.embed_batch_with_usage([text])
|
|
if not embeddings:
|
|
raise RuntimeError(
|
|
f"Mistral embeddings API returned no embedding for model "
|
|
f"{self.embedding_model}"
|
|
)
|
|
return embeddings[0], tokens
|
|
|
|
async def embed_batch_with_usage(
|
|
self, texts: list[str]
|
|
) -> tuple[list[list[float]], int]:
|
|
"""Embed multiple texts, summing the Mistral-reported token usage.
|
|
|
|
Returns ``(embeddings, total_tokens)`` where ``total_tokens`` is the
|
|
sum of ``response.usage.total_tokens`` across the ``BATCH_SIZE`` sub-
|
|
requests (the unit Mistral bills on). Used by the usage-metering hooks
|
|
to record ``embeddings_queries`` by tokens (Deck #67).
|
|
"""
|
|
if not self.supports_embeddings:
|
|
raise NotImplementedError(_NO_EMBEDDING_MODEL_MSG)
|
|
|
|
if not texts:
|
|
return [], 0
|
|
|
|
all_embeddings: list[list[float]] = []
|
|
total_tokens = 0
|
|
for i in range(0, len(texts), BATCH_SIZE):
|
|
batch = texts[i : i + BATCH_SIZE]
|
|
batch_embeddings, batch_tokens = await self._embed_batch_request(batch)
|
|
all_embeddings.extend(batch_embeddings)
|
|
total_tokens += batch_tokens
|
|
|
|
if self._dimension is None and batch_embeddings:
|
|
self._dimension = len(batch_embeddings[0])
|
|
logger.info(
|
|
"Detected embedding dimension: %d for model %s",
|
|
self._dimension,
|
|
self.embedding_model,
|
|
)
|
|
|
|
return all_embeddings, total_tokens
|
|
|
|
@_retry_429
|
|
async def _embed_batch_request(
|
|
self, batch: list[str]
|
|
) -> tuple[list[list[float]], int]:
|
|
"""Single batch request with rate-limit retry.
|
|
|
|
Returns ``(embeddings, token_count)``; ``token_count`` comes from the
|
|
response's ``usage.total_tokens`` and falls back to a char-based
|
|
estimate if the API omits usage.
|
|
"""
|
|
assert self.embedding_model is not None
|
|
response = await self.client.embeddings.create_async(
|
|
model=self.embedding_model,
|
|
inputs=batch,
|
|
)
|
|
|
|
# Defensive: response.data items have Optional fields. Sort by index
|
|
# (default 0 if missing) and reject None embeddings explicitly.
|
|
sorted_data = sorted(response.data or [], key=lambda x: x.index or 0)
|
|
result: list[list[float]] = []
|
|
for item in sorted_data:
|
|
if item.embedding is None:
|
|
raise RuntimeError(
|
|
f"Mistral embeddings API returned a null embedding for "
|
|
f"model {self.embedding_model}"
|
|
)
|
|
result.append(item.embedding)
|
|
|
|
if len(result) != len(batch):
|
|
raise RuntimeError(
|
|
f"Mistral embeddings API returned {len(result)} embeddings "
|
|
f"for {len(batch)} inputs"
|
|
)
|
|
|
|
usage = getattr(response, "usage", None)
|
|
total_tokens = getattr(usage, "total_tokens", None) if usage else None
|
|
# Guard on numeric type (not just ``is not None``): a real response
|
|
# gives an int, but test doubles / partial responses can surface a
|
|
# non-numeric attribute — fall back to the estimate there.
|
|
tokens = (
|
|
int(total_tokens)
|
|
if isinstance(total_tokens, (int, float))
|
|
else self._estimate_tokens(batch)
|
|
)
|
|
return result, tokens
|
|
|
|
def get_dimension(self) -> int:
|
|
if not self.supports_embeddings:
|
|
raise NotImplementedError(_NO_EMBEDDING_MODEL_MSG)
|
|
|
|
if self._dimension is None:
|
|
raise RuntimeError(
|
|
f"Embedding dimension not detected yet for model "
|
|
f"{self.embedding_model}. Call embed() first or use a known "
|
|
"model."
|
|
)
|
|
return self._dimension
|
|
|
|
async def generate(self, prompt: str, max_tokens: int = 500) -> str:
|
|
raise NotImplementedError(
|
|
"MistralProvider does not support generation. "
|
|
"Use OpenAI, Anthropic, or Bedrock for text generation."
|
|
)
|
|
|
|
async def close(self) -> None:
|
|
# The mistralai 2.x client (Speakeasy-generated) does not expose a
|
|
# public close()/aclose() — only the async-context-manager protocol
|
|
# (__aenter__/__aexit__). Calling __aexit__ directly is internal API
|
|
# and brittle across SDK patch versions; the underlying httpx client
|
|
# is closed during garbage collection, so we leave this as a no-op.
|
|
return None
|