feat(providers): add Mistral embedding provider, route registry through dynaconf
Adds a hosted Mistral embedding option (mistral-embed, 1024-dim) alongside the existing Bedrock / OpenAI / Ollama / Simple providers. Implementation mirrors OpenAIProvider: lazy dimension detection with a known-models lookup, chunked batch requests, defensive index sort, and a 429-aware retry decorator. In the same change, ProviderRegistry switches from os.getenv to the dynaconf-backed Settings dataclass so all five providers share a single configuration path. config.py gains the previously-uncovered Bedrock keys, the new Mistral keys, the missing OPENAI_GENERATION_MODEL / OLLAMA_GENERATION_MODEL, and SIMPLE_EMBEDDING_DIMENSION. Auto-detection priority: Bedrock → OpenAI → Mistral → Ollama → Simple. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.7
parent
61cadf7935
commit
3268a13d11
@@ -3,6 +3,7 @@
|
||||
from .anthropic import AnthropicProvider
|
||||
from .base import Provider
|
||||
from .bedrock import BedrockProvider
|
||||
from .mistral import MistralProvider
|
||||
from .ollama import OllamaProvider
|
||||
from .openai import OpenAIProvider
|
||||
from .registry import get_provider, reset_provider
|
||||
@@ -13,6 +14,7 @@ __all__ = [
|
||||
"OllamaProvider",
|
||||
"OpenAIProvider",
|
||||
"AnthropicProvider",
|
||||
"MistralProvider",
|
||||
"SimpleProvider",
|
||||
"BedrockProvider",
|
||||
"get_provider",
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Mistral provider for embeddings.
|
||||
|
||||
Currently supports embeddings only (``mistral-embed``, 1024-dim). Generation
|
||||
can be added later if needed; see ADR-015.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from functools import wraps
|
||||
|
||||
import anyio
|
||||
from mistralai.client import Mistral
|
||||
from mistralai.client.errors.sdkerror import SDKError
|
||||
|
||||
from .base import Provider
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_RETRIES = 5
|
||||
INITIAL_RETRY_DELAY = 2.0
|
||||
MAX_RETRY_DELAY = 60.0
|
||||
|
||||
|
||||
def retry_on_rate_limit(func):
|
||||
"""Retry on Mistral 429 (rate limit) responses with exponential backoff."""
|
||||
|
||||
@wraps(func)
|
||||
async def wrapper(*args, **kwargs):
|
||||
retry_delay = INITIAL_RETRY_DELAY
|
||||
last_error: Exception | None = None
|
||||
|
||||
for attempt in range(1, MAX_RETRIES + 1):
|
||||
try:
|
||||
return await func(*args, **kwargs)
|
||||
except SDKError as e:
|
||||
# SDKError carries a status_code attribute populated from the
|
||||
# raw response. Only 429 is retryable here.
|
||||
status = getattr(e, "status_code", None)
|
||||
if status != 429:
|
||||
raise
|
||||
last_error = e
|
||||
if attempt < MAX_RETRIES:
|
||||
logger.warning(
|
||||
"Mistral rate limit hit (attempt %d/%d), retrying in %.1fs...",
|
||||
attempt,
|
||||
MAX_RETRIES,
|
||||
retry_delay,
|
||||
)
|
||||
await anyio.sleep(retry_delay)
|
||||
retry_delay = min(retry_delay * 2, MAX_RETRY_DELAY)
|
||||
|
||||
logger.error("Mistral rate limit exceeded after %d attempts", MAX_RETRIES)
|
||||
raise last_error # type: ignore[misc]
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
|
||||
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_on_rate_limit
|
||||
async def embed(self, text: str) -> list[float]:
|
||||
"""Generate an embedding for a single text."""
|
||||
if not self.supports_embeddings:
|
||||
raise NotImplementedError(
|
||||
"Embedding not supported - no embedding_model configured"
|
||||
)
|
||||
|
||||
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``."""
|
||||
if not self.supports_embeddings:
|
||||
raise NotImplementedError(
|
||||
"Embedding not supported - no embedding_model configured"
|
||||
)
|
||||
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
all_embeddings: list[list[float]] = []
|
||||
for i in range(0, len(texts), BATCH_SIZE):
|
||||
batch = texts[i : i + BATCH_SIZE]
|
||||
batch_embeddings = await self._embed_batch_request(batch)
|
||||
all_embeddings.extend(batch_embeddings)
|
||||
|
||||
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
|
||||
|
||||
@retry_on_rate_limit
|
||||
async def _embed_batch_request(self, batch: list[str]) -> list[list[float]]:
|
||||
"""Single batch request with rate-limit retry."""
|
||||
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"
|
||||
)
|
||||
return result
|
||||
|
||||
def get_dimension(self) -> int:
|
||||
if not self.supports_embeddings:
|
||||
raise NotImplementedError(
|
||||
"Embedding not supported - no embedding_model configured"
|
||||
)
|
||||
|
||||
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 Mistral SDK manages its own httpx client lifecycle; close it
|
||||
# via the SDK's context-manager hook if present, otherwise no-op.
|
||||
close = getattr(self.client, "__aexit__", None)
|
||||
if close is not None:
|
||||
try:
|
||||
await close(None, None, None)
|
||||
except Exception: # pragma: no cover - best-effort cleanup
|
||||
logger.debug("Mistral client close raised; ignoring", exc_info=True)
|
||||
@@ -1,10 +1,11 @@
|
||||
"""Provider registry and factory for auto-detection and instantiation."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from ..config import get_settings
|
||||
from .base import Provider
|
||||
from .bedrock import BedrockProvider
|
||||
from .mistral import MistralProvider
|
||||
from .ollama import OllamaProvider
|
||||
from .openai import OpenAIProvider
|
||||
from .simple import SimpleProvider
|
||||
@@ -16,117 +17,112 @@ class ProviderRegistry:
|
||||
"""
|
||||
Registry for provider auto-detection and instantiation.
|
||||
|
||||
Checks environment variables in priority order and creates appropriate provider:
|
||||
1. Bedrock (AWS_REGION + BEDROCK_*_MODEL)
|
||||
2. OpenAI (OPENAI_API_KEY)
|
||||
3. Ollama (OLLAMA_BASE_URL)
|
||||
4. Simple (fallback for testing/development)
|
||||
Reads configuration via dynaconf-backed Settings (see ``config.py``).
|
||||
Checks provider settings in priority order and creates the appropriate
|
||||
provider:
|
||||
|
||||
1. Bedrock (``AWS_REGION`` or ``BEDROCK_*_MODEL``)
|
||||
2. OpenAI (``OPENAI_API_KEY``)
|
||||
3. Mistral (``MISTRAL_API_KEY``)
|
||||
4. Ollama (``OLLAMA_BASE_URL``)
|
||||
5. Simple (fallback for testing/development)
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create_provider() -> Provider:
|
||||
"""
|
||||
Auto-detect and create provider based on environment variables.
|
||||
Auto-detect and create provider based on configured settings.
|
||||
|
||||
Settings are sourced via :func:`nextcloud_mcp_server.config.get_settings`,
|
||||
which reads from settings files and environment variables (env vars
|
||||
always win, see ADR-024/025).
|
||||
|
||||
Priority order:
|
||||
1. Bedrock - if AWS_REGION or BEDROCK_EMBEDDING_MODEL is set
|
||||
2. OpenAI - if OPENAI_API_KEY is set
|
||||
3. Ollama - if OLLAMA_BASE_URL is set
|
||||
4. Simple - fallback for testing/development
|
||||
|
||||
1. Bedrock - if ``aws_region`` or ``bedrock_embedding_model`` is set
|
||||
2. OpenAI - if ``openai_api_key`` is set
|
||||
3. Mistral - if ``mistral_api_key`` is set
|
||||
4. Ollama - if ``ollama_base_url`` is set
|
||||
5. Simple - fallback for testing/development
|
||||
|
||||
Returns:
|
||||
Provider instance
|
||||
|
||||
Environment Variables:
|
||||
Bedrock:
|
||||
- AWS_REGION: AWS region (e.g., "us-east-1")
|
||||
- AWS_ACCESS_KEY_ID: AWS access key (optional, uses credential chain)
|
||||
- AWS_SECRET_ACCESS_KEY: AWS secret key (optional)
|
||||
- BEDROCK_EMBEDDING_MODEL: Model ID for embeddings (e.g., "amazon.titan-embed-text-v2:0")
|
||||
- BEDROCK_GENERATION_MODEL: Model ID for text generation (e.g., "anthropic.claude-3-sonnet-20240229-v1:0")
|
||||
|
||||
OpenAI:
|
||||
- OPENAI_API_KEY: OpenAI API key (or GITHUB_TOKEN for GitHub Models)
|
||||
- OPENAI_BASE_URL: Base URL override (e.g., "https://models.github.ai/inference")
|
||||
- OPENAI_EMBEDDING_MODEL: Model for embeddings (default: "text-embedding-3-small")
|
||||
- OPENAI_GENERATION_MODEL: Model for text generation (e.g., "gpt-4o-mini")
|
||||
|
||||
Ollama:
|
||||
- OLLAMA_BASE_URL: Ollama API base URL (e.g., "http://localhost:11434")
|
||||
- OLLAMA_EMBEDDING_MODEL: Model for embeddings (default: "nomic-embed-text")
|
||||
- OLLAMA_GENERATION_MODEL: Model for text generation (e.g., "llama3.2:1b")
|
||||
- OLLAMA_VERIFY_SSL: Verify SSL certificates (default: "true")
|
||||
|
||||
Simple (no configuration needed, fallback):
|
||||
- SIMPLE_EMBEDDING_DIMENSION: Embedding dimension (default: 384)
|
||||
"""
|
||||
# 1. Check for Bedrock
|
||||
aws_region = os.getenv("AWS_REGION")
|
||||
bedrock_embedding_model = os.getenv("BEDROCK_EMBEDDING_MODEL")
|
||||
bedrock_generation_model = os.getenv("BEDROCK_GENERATION_MODEL")
|
||||
settings = get_settings()
|
||||
|
||||
if aws_region or bedrock_embedding_model or bedrock_generation_model:
|
||||
# 1. Bedrock
|
||||
if (
|
||||
settings.aws_region
|
||||
or settings.bedrock_embedding_model
|
||||
or settings.bedrock_generation_model
|
||||
):
|
||||
logger.info(
|
||||
f"Using Bedrock provider: region={aws_region}, "
|
||||
f"embedding_model={bedrock_embedding_model}, "
|
||||
f"generation_model={bedrock_generation_model}"
|
||||
"Using Bedrock provider: region=%s, embedding_model=%s, "
|
||||
"generation_model=%s",
|
||||
settings.aws_region,
|
||||
settings.bedrock_embedding_model,
|
||||
settings.bedrock_generation_model,
|
||||
)
|
||||
return BedrockProvider(
|
||||
region_name=aws_region,
|
||||
embedding_model=bedrock_embedding_model,
|
||||
generation_model=bedrock_generation_model,
|
||||
aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID"),
|
||||
aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY"),
|
||||
region_name=settings.aws_region,
|
||||
embedding_model=settings.bedrock_embedding_model,
|
||||
generation_model=settings.bedrock_generation_model,
|
||||
aws_access_key_id=settings.aws_access_key_id,
|
||||
aws_secret_access_key=settings.aws_secret_access_key,
|
||||
)
|
||||
|
||||
# 2. Check for OpenAI
|
||||
openai_api_key = os.getenv("OPENAI_API_KEY")
|
||||
if openai_api_key:
|
||||
base_url = os.getenv("OPENAI_BASE_URL")
|
||||
embedding_model = os.getenv(
|
||||
"OPENAI_EMBEDDING_MODEL", "text-embedding-3-small"
|
||||
)
|
||||
generation_model = os.getenv("OPENAI_GENERATION_MODEL")
|
||||
|
||||
# 2. OpenAI
|
||||
if settings.openai_api_key:
|
||||
logger.info(
|
||||
f"Using OpenAI provider: base_url={base_url or 'default'}, "
|
||||
f"embedding_model={embedding_model}, "
|
||||
f"generation_model={generation_model}"
|
||||
"Using OpenAI provider: base_url=%s, embedding_model=%s, "
|
||||
"generation_model=%s",
|
||||
settings.openai_base_url or "default",
|
||||
settings.openai_embedding_model,
|
||||
settings.openai_generation_model,
|
||||
)
|
||||
return OpenAIProvider(
|
||||
api_key=openai_api_key,
|
||||
base_url=base_url,
|
||||
embedding_model=embedding_model,
|
||||
generation_model=generation_model,
|
||||
api_key=settings.openai_api_key,
|
||||
base_url=settings.openai_base_url,
|
||||
embedding_model=settings.openai_embedding_model,
|
||||
generation_model=settings.openai_generation_model,
|
||||
)
|
||||
|
||||
# 3. Check for Ollama (local LLM)
|
||||
ollama_url = os.getenv("OLLAMA_BASE_URL")
|
||||
if ollama_url:
|
||||
embedding_model = os.getenv("OLLAMA_EMBEDDING_MODEL", "nomic-embed-text")
|
||||
generation_model = os.getenv("OLLAMA_GENERATION_MODEL")
|
||||
verify_ssl = os.getenv("OLLAMA_VERIFY_SSL", "true").lower() == "true"
|
||||
|
||||
# 3. Mistral
|
||||
if settings.mistral_api_key:
|
||||
logger.info(
|
||||
f"Using Ollama provider: {ollama_url}, "
|
||||
f"embedding_model={embedding_model}, "
|
||||
f"generation_model={generation_model}"
|
||||
"Using Mistral provider: base_url=%s, embedding_model=%s",
|
||||
settings.mistral_base_url or "default",
|
||||
settings.mistral_embedding_model,
|
||||
)
|
||||
return MistralProvider(
|
||||
api_key=settings.mistral_api_key,
|
||||
base_url=settings.mistral_base_url,
|
||||
embedding_model=settings.mistral_embedding_model,
|
||||
)
|
||||
|
||||
# 4. Ollama
|
||||
if settings.ollama_base_url:
|
||||
logger.info(
|
||||
"Using Ollama provider: %s, embedding_model=%s, generation_model=%s",
|
||||
settings.ollama_base_url,
|
||||
settings.ollama_embedding_model,
|
||||
settings.ollama_generation_model,
|
||||
)
|
||||
return OllamaProvider(
|
||||
base_url=ollama_url,
|
||||
embedding_model=embedding_model,
|
||||
generation_model=generation_model,
|
||||
verify_ssl=verify_ssl,
|
||||
base_url=settings.ollama_base_url,
|
||||
embedding_model=settings.ollama_embedding_model,
|
||||
generation_model=settings.ollama_generation_model,
|
||||
verify_ssl=settings.ollama_verify_ssl,
|
||||
)
|
||||
|
||||
# 4. Fallback to Simple provider for development/testing
|
||||
dimension = int(os.getenv("SIMPLE_EMBEDDING_DIMENSION", "384"))
|
||||
# 5. Simple (fallback)
|
||||
logger.warning(
|
||||
"No provider configured (AWS_REGION, OPENAI_API_KEY, OLLAMA_BASE_URL not set). "
|
||||
"No provider configured (AWS_REGION, OPENAI_API_KEY, "
|
||||
"MISTRAL_API_KEY, OLLAMA_BASE_URL not set). "
|
||||
"Using SimpleProvider for testing/development. "
|
||||
"For production, configure Bedrock, OpenAI, or Ollama."
|
||||
"For production, configure Bedrock, OpenAI, Mistral, or Ollama."
|
||||
)
|
||||
return SimpleProvider(dimension=dimension)
|
||||
return SimpleProvider(dimension=settings.simple_embedding_dimension)
|
||||
|
||||
|
||||
# Singleton instance
|
||||
|
||||
Reference in New Issue
Block a user