Add per-endpoint api_key to model and embedding config
This commit is contained in:
parent
e9f6fea598
commit
2310b7a8b3
16 changed files with 440 additions and 56 deletions
|
|
@ -1,6 +1,15 @@
|
||||||
# Changelog
|
# Changelog
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- `api_key` on model and embedding-model config, overriding the provider's environment variable. Honored on the `openai` and `ollama` providers, `vllm` embedders and rerankers, the picture-description VLM endpoint, and `doctor`'s endpoint probes; other providers raise.
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
- `doctor`'s docling-serve probe sends `X-Api-Key`, so an instance requiring a key is reported reachable rather than unreachable.
|
||||||
|
- The picture-description request to the public OpenAI endpoint sends `OPENAI_API_KEY`; it carried no authorization header.
|
||||||
|
|
||||||
## [0.77.0] - 2026-08-21
|
## [0.77.0] - 2026-08-21
|
||||||
|
|
||||||
## [0.77.0] - 2026-08-21
|
## [0.77.0] - 2026-08-21
|
||||||
|
|
|
||||||
|
|
@ -29,8 +29,32 @@ qa:
|
||||||
- **max_tokens**: Maximum tokens in response. Default: unset (provider default), except title generation (100).
|
- **max_tokens**: Maximum tokens in response. Default: unset (provider default), except title generation (100).
|
||||||
- **enable_thinking**: Control reasoning behavior (see below)
|
- **enable_thinking**: Control reasoning behavior (see below)
|
||||||
- **base_url**: Custom endpoint for OpenAI-compatible servers (vLLM, LM Studio, etc.)
|
- **base_url**: Custom endpoint for OpenAI-compatible servers (vLLM, LM Studio, etc.)
|
||||||
|
- **api_key**: Key for this endpoint, overriding the provider's environment variable (see [Per-endpoint API keys](#per-endpoint-api-keys))
|
||||||
- **extra_body**: Raw dict forwarded to the model SDK (see [Raw Provider Pass-through](#raw-provider-pass-through))
|
- **extra_body**: Raw dict forwarded to the model SDK (see [Raw Provider Pass-through](#raw-provider-pass-through))
|
||||||
|
|
||||||
|
### Per-endpoint API keys
|
||||||
|
|
||||||
|
The `openai` provider reads `OPENAI_API_KEY`, so several `openai`-compatible endpoints in one config would otherwise share a single key. Set `api_key` per model to give each its own, and keep the secret in the environment with [variable expansion](index.md#environment-variables):
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
qa:
|
||||||
|
model:
|
||||||
|
provider: openai
|
||||||
|
name: some-model
|
||||||
|
base_url: https://vendor-a.example/v1
|
||||||
|
api_key: ${VENDOR_A_KEY}
|
||||||
|
|
||||||
|
embeddings:
|
||||||
|
model:
|
||||||
|
provider: openai
|
||||||
|
name: some-embedding-model
|
||||||
|
vector_dim: 1024
|
||||||
|
base_url: https://vendor-b.example/v1
|
||||||
|
api_key: ${VENDOR_B_KEY}
|
||||||
|
```
|
||||||
|
|
||||||
|
`api_key` is honored on the `openai` and `ollama` providers, on `vllm` embedders and rerankers, and on the picture-description VLM endpoint (which otherwise falls back to `OPENAI_API_KEY` only for the public OpenAI endpoint, never for a custom `base_url`). Other providers (`anthropic`, `cohere`, `voyageai`, …) reach their vendor SDK by name and read their own environment variable; setting `api_key` there raises rather than being dropped silently.
|
||||||
|
|
||||||
### Thinking Control
|
### Thinking Control
|
||||||
|
|
||||||
The `enable_thinking` setting controls whether models use explicit reasoning steps before answering.
|
The `enable_thinking` setting controls whether models use explicit reasoning steps before answering.
|
||||||
|
|
@ -107,7 +131,7 @@ Same mechanism, opposite direction. Without `extra_body` the Gemma-4 chat templa
|
||||||
|
|
||||||
## Embedding Providers
|
## Embedding Providers
|
||||||
|
|
||||||
Embedding models require three settings: `provider`, `name`, and `vector_dim`. Optionally, use `base_url` for OpenAI-compatible servers.
|
Embedding models require three settings: `provider`, `name`, and `vector_dim`. Optionally, use `base_url` for OpenAI-compatible servers and [`api_key`](#per-endpoint-api-keys) for the key that endpoint expects.
|
||||||
|
|
||||||
### Batch Size
|
### Batch Size
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -23,6 +23,11 @@ class ModelConfig(ConfigModel):
|
||||||
provider: Model provider (ollama, openai, anthropic, etc.)
|
provider: Model provider (ollama, openai, anthropic, etc.)
|
||||||
name: Model name/identifier
|
name: Model name/identifier
|
||||||
base_url: Optional base URL for OpenAI-compatible servers (vLLM, LM Studio, etc.)
|
base_url: Optional base URL for OpenAI-compatible servers (vLLM, LM Studio, etc.)
|
||||||
|
api_key: Key sent to the endpoint, overriding the provider's own
|
||||||
|
environment variable. Lets several openai-compatible endpoints each
|
||||||
|
carry their own key; typically written as `${VENDOR_KEY}`. Honored
|
||||||
|
on the openai and ollama providers, and on the picture-description
|
||||||
|
VLM endpoint.
|
||||||
enable_thinking: Control reasoning behavior (true/false/None for default)
|
enable_thinking: Control reasoning behavior (true/false/None for default)
|
||||||
temperature: Sampling temperature (0.0 to 1.0+)
|
temperature: Sampling temperature (0.0 to 1.0+)
|
||||||
max_tokens: Maximum tokens to generate
|
max_tokens: Maximum tokens to generate
|
||||||
|
|
@ -37,6 +42,7 @@ class ModelConfig(ConfigModel):
|
||||||
provider: str = "ollama"
|
provider: str = "ollama"
|
||||||
name: str = "gpt-oss"
|
name: str = "gpt-oss"
|
||||||
base_url: str | None = None
|
base_url: str | None = None
|
||||||
|
api_key: str | None = None
|
||||||
|
|
||||||
enable_thinking: bool | None = None
|
enable_thinking: bool | None = None
|
||||||
temperature: float | None = None
|
temperature: float | None = None
|
||||||
|
|
@ -53,6 +59,9 @@ class EmbeddingModelConfig(ConfigModel):
|
||||||
name: Model name/identifier
|
name: Model name/identifier
|
||||||
vector_dim: Vector dimensions produced by the model
|
vector_dim: Vector dimensions produced by the model
|
||||||
base_url: Optional base URL for OpenAI-compatible servers (vLLM, LM Studio, etc.)
|
base_url: Optional base URL for OpenAI-compatible servers (vLLM, LM Studio, etc.)
|
||||||
|
api_key: Key sent to the endpoint, overriding the provider's own
|
||||||
|
environment variable. Honored on the openai, ollama and vllm
|
||||||
|
providers.
|
||||||
multimodal: Whether the model embeds images into the same vector space as
|
multimodal: Whether the model embeds images into the same vector space as
|
||||||
text. Supported on the vllm, voyageai, and cohere providers; other
|
text. Supported on the vllm, voyageai, and cohere providers; other
|
||||||
providers raise when this is set.
|
providers raise when this is set.
|
||||||
|
|
@ -62,6 +71,7 @@ class EmbeddingModelConfig(ConfigModel):
|
||||||
name: str = "qwen3-embedding:4b"
|
name: str = "qwen3-embedding:4b"
|
||||||
vector_dim: int = Field(default=2560, gt=0)
|
vector_dim: int = Field(default=2560, gt=0)
|
||||||
base_url: str | None = None
|
base_url: str | None = None
|
||||||
|
api_key: str | None = None
|
||||||
multimodal: bool = False
|
multimodal: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
"""Base class for document converters."""
|
"""Base class for document converters."""
|
||||||
|
|
||||||
|
import os
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
@ -25,6 +26,22 @@ def vlm_api_url(config: "AppConfig", model: "ModelConfig") -> str:
|
||||||
raise ValueError(f"Unsupported VLM provider: {model.provider}")
|
raise ValueError(f"Unsupported VLM provider: {model.provider}")
|
||||||
|
|
||||||
|
|
||||||
|
def vlm_api_headers(model: "ModelConfig") -> dict[str, str]:
|
||||||
|
"""Auth headers for the picture-description VLM endpoint. Docling posts to
|
||||||
|
it directly, so the key travels as a header rather than through an SDK.
|
||||||
|
|
||||||
|
The public OpenAI endpoint falls back to ``OPENAI_API_KEY``. A custom
|
||||||
|
``base_url`` never does: that key belongs to api.openai.com, not to
|
||||||
|
whatever self-hosted server the model points at.
|
||||||
|
"""
|
||||||
|
key = model.api_key
|
||||||
|
if not key and model.provider == "openai" and not model.base_url:
|
||||||
|
key = os.environ.get("OPENAI_API_KEY")
|
||||||
|
if key:
|
||||||
|
return {"Authorization": f"Bearer {key}"}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
class DocumentConverter(ABC):
|
class DocumentConverter(ABC):
|
||||||
"""Abstract base class for document converters.
|
"""Abstract base class for document converters.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -9,7 +9,11 @@ from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, ClassVar
|
from typing import TYPE_CHECKING, ClassVar
|
||||||
|
|
||||||
from haiku.rag.config import AppConfig
|
from haiku.rag.config import AppConfig
|
||||||
from haiku.rag.converters.base import DocumentConverter, vlm_api_url
|
from haiku.rag.converters.base import (
|
||||||
|
DocumentConverter,
|
||||||
|
vlm_api_headers,
|
||||||
|
vlm_api_url,
|
||||||
|
)
|
||||||
from haiku.rag.converters.text_utils import TextFileHandler, docling_safe_name
|
from haiku.rag.converters.text_utils import TextFileHandler, docling_safe_name
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|
@ -148,6 +152,7 @@ class DoclingLocalConverter(DocumentConverter):
|
||||||
pipeline_options.enable_remote_services = True
|
pipeline_options.enable_remote_services = True
|
||||||
pipeline_options.picture_description_options = PictureDescriptionApiOptions(
|
pipeline_options.picture_description_options = PictureDescriptionApiOptions(
|
||||||
url=AnyUrl(vlm_api_url(self.config, pic_desc.model)),
|
url=AnyUrl(vlm_api_url(self.config, pic_desc.model)),
|
||||||
|
headers=vlm_api_headers(pic_desc.model),
|
||||||
params=dict(
|
params=dict(
|
||||||
model=pic_desc.model.name,
|
model=pic_desc.model.name,
|
||||||
max_completion_tokens=pic_desc.max_tokens,
|
max_completion_tokens=pic_desc.max_tokens,
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,11 @@ from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, ClassVar
|
from typing import TYPE_CHECKING, ClassVar
|
||||||
|
|
||||||
from haiku.rag.config import AppConfig
|
from haiku.rag.config import AppConfig
|
||||||
from haiku.rag.converters.base import DocumentConverter, vlm_api_url
|
from haiku.rag.converters.base import (
|
||||||
|
DocumentConverter,
|
||||||
|
vlm_api_headers,
|
||||||
|
vlm_api_url,
|
||||||
|
)
|
||||||
from haiku.rag.converters.text_utils import TextFileHandler, docling_safe_name
|
from haiku.rag.converters.text_utils import TextFileHandler, docling_safe_name
|
||||||
from haiku.rag.providers.docling_serve import DoclingServeClient
|
from haiku.rag.providers.docling_serve import DoclingServeClient
|
||||||
|
|
||||||
|
|
@ -105,6 +109,7 @@ class DoclingServeConverter(DocumentConverter):
|
||||||
prompt = self.config.prompts.picture_description
|
prompt = self.config.prompts.picture_description
|
||||||
picture_description_api = {
|
picture_description_api = {
|
||||||
"url": vlm_api_url(self.config, pic_desc.model),
|
"url": vlm_api_url(self.config, pic_desc.model),
|
||||||
|
"headers": vlm_api_headers(pic_desc.model),
|
||||||
"params": {
|
"params": {
|
||||||
"model": pic_desc.model.name,
|
"model": pic_desc.model.name,
|
||||||
"max_completion_tokens": pic_desc.max_tokens,
|
"max_completion_tokens": pic_desc.max_tokens,
|
||||||
|
|
|
||||||
|
|
@ -10,7 +10,11 @@ import yaml
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from haiku.rag.config import AppConfig
|
from haiku.rag.config import AppConfig
|
||||||
from haiku.rag.config.models import DuplicateDetectionConfig
|
from haiku.rag.config.models import (
|
||||||
|
DuplicateDetectionConfig,
|
||||||
|
EmbeddingModelConfig,
|
||||||
|
ModelConfig,
|
||||||
|
)
|
||||||
from haiku.rag.store.engine import Store, connect_lancedb
|
from haiku.rag.store.engine import Store, connect_lancedb
|
||||||
from haiku.rag.store.info import get_database_stats
|
from haiku.rag.store.info import get_database_stats
|
||||||
from haiku.rag.store.repositories.settings import SettingsRepository
|
from haiku.rag.store.repositories.settings import SettingsRepository
|
||||||
|
|
@ -83,42 +87,37 @@ def _sample(ids: list[str]) -> list[str]:
|
||||||
return [*ids[:_SAMPLE_LIMIT], f"... (+{extra} more)"]
|
return [*ids[:_SAMPLE_LIMIT], f"... (+{extra} more)"]
|
||||||
|
|
||||||
|
|
||||||
def _active_models(config: AppConfig) -> list[tuple[str, str, str | None]]:
|
def _active_models(config: AppConfig) -> list[ModelConfig | EmbeddingModelConfig]:
|
||||||
"""(provider, name, base_url) for every model role the config activates.
|
"""Every model role the config activates.
|
||||||
|
|
||||||
Picture-description and title models are only included when their feature
|
Picture-description and title models are only included when their feature
|
||||||
is enabled (``processing.pictures == "description"`` / ``auto_title``), so
|
is enabled (``processing.pictures == "description"`` / ``auto_title``), so
|
||||||
doctor checks exactly the providers the next ingest will use.
|
doctor checks exactly the providers the next ingest will use.
|
||||||
"""
|
"""
|
||||||
models = [
|
models: list[ModelConfig | EmbeddingModelConfig] = [config.embeddings.model]
|
||||||
(
|
|
||||||
config.embeddings.model.provider,
|
|
||||||
config.embeddings.model.name,
|
|
||||||
config.embeddings.model.base_url,
|
|
||||||
)
|
|
||||||
]
|
|
||||||
for model in (config.reranking.model, config.qa.model, config.analysis.model):
|
for model in (config.reranking.model, config.qa.model, config.analysis.model):
|
||||||
if model is not None:
|
if model is not None:
|
||||||
models.append((model.provider, model.name, model.base_url))
|
models.append(model)
|
||||||
|
|
||||||
proc = config.processing
|
proc = config.processing
|
||||||
if proc.pictures == "description":
|
if proc.pictures == "description":
|
||||||
pd = proc.conversion_options.picture_description.model
|
models.append(proc.conversion_options.picture_description.model)
|
||||||
models.append((pd.provider, pd.name, pd.base_url))
|
|
||||||
if proc.auto_title:
|
if proc.auto_title:
|
||||||
tm = proc.title_model
|
models.append(proc.title_model)
|
||||||
models.append((tm.provider, tm.name, tm.base_url))
|
|
||||||
return models
|
return models
|
||||||
|
|
||||||
|
|
||||||
def _check_api_keys(config: AppConfig, environ: dict[str, str]) -> CheckResult:
|
def _check_api_keys(config: AppConfig, environ: dict[str, str]) -> CheckResult:
|
||||||
# A custom base_url points at a self-hosted OpenAI-compatible endpoint that
|
# A custom base_url points at a self-hosted OpenAI-compatible endpoint that
|
||||||
# uses a placeholder key, so the SaaS key is only required when a provider
|
# uses a placeholder key, so the SaaS key is only required when a provider
|
||||||
# is used without one. Reachability of custom endpoints is the probe's job.
|
# is used without one, and without a key in the config. Reachability of
|
||||||
|
# custom endpoints is the probe's job.
|
||||||
need_key = {
|
need_key = {
|
||||||
provider
|
model.provider
|
||||||
for provider, _name, base_url in _active_models(config)
|
for model in _active_models(config)
|
||||||
if not base_url and provider in _PROVIDER_ENV_VARS
|
if not model.base_url
|
||||||
|
and not model.api_key
|
||||||
|
and model.provider in _PROVIDER_ENV_VARS
|
||||||
}
|
}
|
||||||
missing = [
|
missing = [
|
||||||
f"{provider} ({_PROVIDER_ENV_VARS[provider]})"
|
f"{provider} ({_PROVIDER_ENV_VARS[provider]})"
|
||||||
|
|
@ -896,31 +895,43 @@ def _provider_targets(
|
||||||
local: set[str] = set()
|
local: set[str] = set()
|
||||||
ollama_base = config.providers.ollama.base_url
|
ollama_base = config.providers.ollama.base_url
|
||||||
|
|
||||||
def add_model(provider: str, name: str, base_url: str | None) -> None:
|
def add_model(model: ModelConfig | EmbeddingModelConfig) -> None:
|
||||||
resolved = _resolve_endpoint(provider, base_url, ollama_base)
|
resolved = _resolve_endpoint(model.provider, model.base_url, ollama_base)
|
||||||
if resolved is None:
|
if resolved is None:
|
||||||
return
|
return
|
||||||
if resolved == "local":
|
if resolved == "local":
|
||||||
local.add(provider)
|
local.add(model.provider)
|
||||||
return
|
return
|
||||||
probe_url, kind, display = resolved
|
probe_url, kind, display = resolved
|
||||||
entry = targets.setdefault(
|
entry = targets.setdefault(
|
||||||
probe_url, {"kind": kind, "display": display, "models": set()}
|
probe_url,
|
||||||
|
{"kind": kind, "display": display, "models": set(), "headers": {}},
|
||||||
)
|
)
|
||||||
if name:
|
# A secured endpoint answers the probe only with its key. Models sharing
|
||||||
entry["models"].add(name)
|
# a probe URL share the endpoint, so the first key configured for it wins.
|
||||||
|
if model.api_key and not entry["headers"]:
|
||||||
|
entry["headers"] = {"Authorization": f"Bearer {model.api_key}"}
|
||||||
|
if model.name:
|
||||||
|
entry["models"].add(model.name)
|
||||||
|
|
||||||
proc = config.processing
|
proc = config.processing
|
||||||
if proc.converter == "docling-serve" or proc.chunker == "docling-serve":
|
if proc.converter == "docling-serve" or proc.chunker == "docling-serve":
|
||||||
|
docling_key = config.providers.docling_serve.api_key
|
||||||
|
headers = {"X-Api-Key": docling_key} if docling_key else {}
|
||||||
for url in config.providers.docling_serve.base_urls:
|
for url in config.providers.docling_serve.base_urls:
|
||||||
base = url.rstrip("/")
|
base = url.rstrip("/")
|
||||||
targets.setdefault(
|
targets.setdefault(
|
||||||
f"{base}/health",
|
f"{base}/health",
|
||||||
{"kind": "docling-serve", "display": base, "models": set()},
|
{
|
||||||
|
"kind": "docling-serve",
|
||||||
|
"display": base,
|
||||||
|
"models": set(),
|
||||||
|
"headers": headers,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
for provider, name, base_url in _active_models(config):
|
for model in _active_models(config):
|
||||||
add_model(provider, name, base_url)
|
add_model(model)
|
||||||
|
|
||||||
return targets, local
|
return targets, local
|
||||||
|
|
||||||
|
|
@ -934,10 +945,10 @@ def _model_present(expected: str, available: set[str]) -> bool:
|
||||||
|
|
||||||
|
|
||||||
async def _probe_endpoint(
|
async def _probe_endpoint(
|
||||||
client: httpx.AsyncClient, url: str
|
client: httpx.AsyncClient, url: str, headers: dict[str, str]
|
||||||
) -> tuple[bool, str | None, dict | None]:
|
) -> tuple[bool, str | None, dict | None]:
|
||||||
try:
|
try:
|
||||||
response = await client.get(url)
|
response = await client.get(url, headers=headers)
|
||||||
except httpx.HTTPError as exc:
|
except httpx.HTTPError as exc:
|
||||||
return False, str(exc), None
|
return False, str(exc), None
|
||||||
if not response.is_success:
|
if not response.is_success:
|
||||||
|
|
@ -996,7 +1007,10 @@ async def run_provider_checks(
|
||||||
on_progress("Probing provider endpoints")
|
on_progress("Probing provider endpoints")
|
||||||
async with httpx.AsyncClient(timeout=_PROBE_TIMEOUT_S) as client:
|
async with httpx.AsyncClient(timeout=_PROBE_TIMEOUT_S) as client:
|
||||||
probes = await asyncio.gather(
|
probes = await asyncio.gather(
|
||||||
*(_probe_endpoint(client, url) for url in targets)
|
*(
|
||||||
|
_probe_endpoint(client, url, targets[url]["headers"])
|
||||||
|
for url in targets
|
||||||
|
)
|
||||||
)
|
)
|
||||||
for url, (reachable, error, payload) in zip(targets, probes):
|
for url, (reachable, error, payload) in zip(targets, probes):
|
||||||
results.append(_endpoint_result(targets[url], reachable, error, payload))
|
results.append(_endpoint_result(targets[url], reachable, error, payload))
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,7 @@ from pydantic_ai.providers.ollama import OllamaProvider
|
||||||
from pydantic_ai.providers.openai import OpenAIProvider
|
from pydantic_ai.providers.openai import OpenAIProvider
|
||||||
|
|
||||||
from haiku.rag.config import AppConfig, get_config
|
from haiku.rag.config import AppConfig, get_config
|
||||||
|
from haiku.rag.utils import check_api_key_supported
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from PIL import Image as PILImage
|
from PIL import Image as PILImage
|
||||||
|
|
@ -198,6 +199,7 @@ def get_embedder(config: AppConfig | None = None) -> EmbedderWrapper:
|
||||||
provider = embedding_model.provider
|
provider = embedding_model.provider
|
||||||
model_name = embedding_model.name
|
model_name = embedding_model.name
|
||||||
vector_dim = embedding_model.vector_dim
|
vector_dim = embedding_model.vector_dim
|
||||||
|
check_api_key_supported(embedding_model, {"openai", "ollama", "vllm"})
|
||||||
|
|
||||||
if embedding_model.multimodal:
|
if embedding_model.multimodal:
|
||||||
return _get_multimodal_embedder(embedding_model)
|
return _get_multimodal_embedder(embedding_model)
|
||||||
|
|
@ -209,15 +211,18 @@ def get_embedder(config: AppConfig | None = None) -> EmbedderWrapper:
|
||||||
base_url = base_url.rstrip("/") + "/v1"
|
base_url = base_url.rstrip("/") + "/v1"
|
||||||
model = OpenAIEmbeddingModel(
|
model = OpenAIEmbeddingModel(
|
||||||
model_name,
|
model_name,
|
||||||
provider=OllamaProvider(base_url=base_url),
|
provider=OllamaProvider(base_url=base_url, api_key=embedding_model.api_key),
|
||||||
)
|
)
|
||||||
return EmbedderWrapper(Embedder(model), vector_dim)
|
return EmbedderWrapper(Embedder(model), vector_dim)
|
||||||
|
|
||||||
if provider == "openai":
|
if provider == "openai":
|
||||||
if embedding_model.base_url:
|
if embedding_model.base_url or embedding_model.api_key:
|
||||||
model = OpenAIEmbeddingModel(
|
model = OpenAIEmbeddingModel(
|
||||||
model_name,
|
model_name,
|
||||||
provider=OpenAIProvider(base_url=embedding_model.base_url),
|
provider=OpenAIProvider(
|
||||||
|
base_url=embedding_model.base_url,
|
||||||
|
api_key=embedding_model.api_key,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
return EmbedderWrapper(Embedder(model), vector_dim)
|
return EmbedderWrapper(Embedder(model), vector_dim)
|
||||||
return EmbedderWrapper(Embedder(f"openai:{model_name}"), vector_dim)
|
return EmbedderWrapper(Embedder(f"openai:{model_name}"), vector_dim)
|
||||||
|
|
@ -238,7 +243,11 @@ def get_embedder(config: AppConfig | None = None) -> EmbedderWrapper:
|
||||||
|
|
||||||
base_url = _vllm_base_url(embedding_model.base_url)
|
base_url = _vllm_base_url(embedding_model.base_url)
|
||||||
return VLLMMultimodalEmbedder(
|
return VLLMMultimodalEmbedder(
|
||||||
model_name, vector_dim, base_url=base_url, supports_images=False
|
model_name,
|
||||||
|
vector_dim,
|
||||||
|
base_url=base_url,
|
||||||
|
api_key=embedding_model.api_key,
|
||||||
|
supports_images=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
raise ValueError(f"Unsupported embedding provider: {provider}")
|
raise ValueError(f"Unsupported embedding provider: {provider}")
|
||||||
|
|
@ -268,7 +277,11 @@ def _get_multimodal_embedder(
|
||||||
|
|
||||||
base_url = _vllm_base_url(embedding_model.base_url)
|
base_url = _vllm_base_url(embedding_model.base_url)
|
||||||
return VLLMMultimodalEmbedder(
|
return VLLMMultimodalEmbedder(
|
||||||
model_name, vector_dim, base_url=base_url, supports_images=True
|
model_name,
|
||||||
|
vector_dim,
|
||||||
|
base_url=base_url,
|
||||||
|
api_key=embedding_model.api_key,
|
||||||
|
supports_images=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
if provider == "voyageai":
|
if provider == "voyageai":
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
from haiku.rag.config import AppConfig, get_config
|
from haiku.rag.config import AppConfig, get_config
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
|
from haiku.rag.utils import check_api_key_supported
|
||||||
|
|
||||||
|
|
||||||
def get_reranker(config: AppConfig | None = None) -> RerankerBase | None:
|
def get_reranker(config: AppConfig | None = None) -> RerankerBase | None:
|
||||||
|
|
@ -14,6 +15,8 @@ def get_reranker(config: AppConfig | None = None) -> RerankerBase | None:
|
||||||
if model is None:
|
if model is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
check_api_key_supported(model, {"vllm"})
|
||||||
|
|
||||||
if config.reranking.multimodal and model.provider != "vllm":
|
if config.reranking.multimodal and model.provider != "vllm":
|
||||||
raise ValueError("reranking.multimodal is only supported on the vllm provider")
|
raise ValueError("reranking.multimodal is only supported on the vllm provider")
|
||||||
|
|
||||||
|
|
@ -27,7 +30,7 @@ def get_reranker(config: AppConfig | None = None) -> RerankerBase | None:
|
||||||
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
|
|
||||||
return VLLMReranker(model.name, model.base_url)
|
return VLLMReranker(model.name, model.base_url, api_key=model.api_key)
|
||||||
|
|
||||||
if model.provider == "zeroentropy":
|
if model.provider == "zeroentropy":
|
||||||
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
||||||
|
|
|
||||||
|
|
@ -24,9 +24,15 @@ def _document(chunk: Chunk) -> str | dict:
|
||||||
|
|
||||||
|
|
||||||
class VLLMReranker(RerankerBase):
|
class VLLMReranker(RerankerBase):
|
||||||
def __init__(self, model: str, base_url: str):
|
def __init__(self, model: str, base_url: str, api_key: str | None = None):
|
||||||
self._model = model
|
self._model = model
|
||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
|
self._headers = {
|
||||||
|
"accept": "application/json",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
}
|
||||||
|
if api_key:
|
||||||
|
self._headers["Authorization"] = f"Bearer {api_key}"
|
||||||
# One client reused across rerank calls (connection kept alive).
|
# One client reused across rerank calls (connection kept alive).
|
||||||
# Multimodal document batches can take far longer than httpx's 5s
|
# Multimodal document batches can take far longer than httpx's 5s
|
||||||
# default timeout to score.
|
# default timeout to score.
|
||||||
|
|
@ -43,10 +49,7 @@ class VLLMReranker(RerankerBase):
|
||||||
response = await self._client.post(
|
response = await self._client.post(
|
||||||
f"{self._base_url}/v1/rerank",
|
f"{self._base_url}/v1/rerank",
|
||||||
json={"model": self._model, "query": query, "documents": documents},
|
json={"model": self._model, "query": query, "documents": documents},
|
||||||
headers={
|
headers=self._headers,
|
||||||
"accept": "application/json",
|
|
||||||
"Content-Type": "application/json",
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,7 @@ if TYPE_CHECKING:
|
||||||
from rich.console import RenderableType
|
from rich.console import RenderableType
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config.models import AppConfig, ModelConfig
|
from haiku.rag.config.models import AppConfig, EmbeddingModelConfig, ModelConfig
|
||||||
from haiku.rag.store.models.citation import Citation
|
from haiku.rag.store.models.citation import Citation
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -28,6 +28,22 @@ def parse_model_option(value: str) -> "ModelConfig":
|
||||||
return ModelConfig(provider=parts[0], name=parts[1])
|
return ModelConfig(provider=parts[0], name=parts[1])
|
||||||
|
|
||||||
|
|
||||||
|
def check_api_key_supported(
|
||||||
|
model_config: "ModelConfig | EmbeddingModelConfig", supported: set[str]
|
||||||
|
) -> None:
|
||||||
|
"""Reject a configured api_key on a provider whose client we never build.
|
||||||
|
|
||||||
|
Those providers reach their vendor SDK by name and read their own
|
||||||
|
environment variable, so a key in the config would be dropped silently.
|
||||||
|
"""
|
||||||
|
if model_config.api_key and model_config.provider not in supported:
|
||||||
|
raise ValueError(
|
||||||
|
f"api_key is not supported on the '{model_config.provider}' provider "
|
||||||
|
f"(supported: {', '.join(sorted(supported))}). Set that provider's "
|
||||||
|
"own API key environment variable instead."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||||
"""Compute cosine similarity between two vectors."""
|
"""Compute cosine similarity between two vectors."""
|
||||||
dot_product = sum(a * b for a, b in zip(vec1, vec2))
|
dot_product = sum(a * b for a, b in zip(vec1, vec2))
|
||||||
|
|
@ -135,6 +151,7 @@ def get_model(
|
||||||
|
|
||||||
provider = model_config.provider
|
provider = model_config.provider
|
||||||
model = model_config.name
|
model = model_config.name
|
||||||
|
check_api_key_supported(model_config, {"openai", "ollama"})
|
||||||
|
|
||||||
if provider == "ollama":
|
if provider == "ollama":
|
||||||
model_settings = None
|
model_settings = None
|
||||||
|
|
@ -158,7 +175,7 @@ def get_model(
|
||||||
|
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=base_url),
|
provider=OllamaProvider(base_url=base_url, api_key=model_config.api_key),
|
||||||
settings=model_settings,
|
settings=model_settings,
|
||||||
profile=_OPENAI_COMPAT_PROFILE,
|
profile=_OPENAI_COMPAT_PROFILE,
|
||||||
)
|
)
|
||||||
|
|
@ -188,12 +205,22 @@ def get_model(
|
||||||
if model_config.base_url:
|
if model_config.base_url:
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OpenAIProvider(base_url=model_config.base_url),
|
provider=OpenAIProvider(
|
||||||
|
base_url=model_config.base_url, api_key=model_config.api_key
|
||||||
|
),
|
||||||
settings=openai_settings,
|
settings=openai_settings,
|
||||||
profile=_OPENAI_COMPAT_PROFILE,
|
profile=_OPENAI_COMPAT_PROFILE,
|
||||||
)
|
)
|
||||||
|
|
||||||
return OpenAIChatModel(model_name=model, settings=openai_settings)
|
return OpenAIChatModel(
|
||||||
|
model_name=model,
|
||||||
|
provider=(
|
||||||
|
OpenAIProvider(api_key=model_config.api_key)
|
||||||
|
if model_config.api_key
|
||||||
|
else "openai"
|
||||||
|
),
|
||||||
|
settings=openai_settings,
|
||||||
|
)
|
||||||
|
|
||||||
elif provider == "anthropic":
|
elif provider == "anthropic":
|
||||||
from anthropic.types.beta import BetaThinkingConfigDisabledParam
|
from anthropic.types.beta import BetaThinkingConfigDisabledParam
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,7 @@ from docling_core.types.doc.document import DoclingDocument
|
||||||
from haiku.rag.config import AppConfig
|
from haiku.rag.config import AppConfig
|
||||||
from haiku.rag.config.models import ModelConfig
|
from haiku.rag.config.models import ModelConfig
|
||||||
from haiku.rag.converters import docling_local, get_converter
|
from haiku.rag.converters import docling_local, get_converter
|
||||||
from haiku.rag.converters.base import vlm_api_url
|
from haiku.rag.converters.base import vlm_api_headers, vlm_api_url
|
||||||
from haiku.rag.converters.docling_local import DoclingLocalConverter
|
from haiku.rag.converters.docling_local import DoclingLocalConverter
|
||||||
from haiku.rag.converters.docling_serve import DoclingServeConverter
|
from haiku.rag.converters.docling_serve import DoclingServeConverter
|
||||||
from haiku.rag.converters.text_utils import TextFileHandler, docling_safe_name
|
from haiku.rag.converters.text_utils import TextFileHandler, docling_safe_name
|
||||||
|
|
@ -50,6 +50,41 @@ class TestVlmApiUrl:
|
||||||
vlm_api_url(AppConfig(), ModelConfig(provider="unsupported", name="test"))
|
vlm_api_url(AppConfig(), ModelConfig(provider="unsupported", name="test"))
|
||||||
|
|
||||||
|
|
||||||
|
class TestVlmApiHeaders:
|
||||||
|
"""Auth headers for picture-description VLM models."""
|
||||||
|
|
||||||
|
def test_no_headers_without_api_key(self):
|
||||||
|
assert vlm_api_headers(ModelConfig(provider="ollama", name="ministral-3")) == {}
|
||||||
|
|
||||||
|
def test_public_openai_falls_back_to_environment(self, monkeypatch):
|
||||||
|
"""doctor accepts OPENAI_API_KEY for a keyless openai model, so the
|
||||||
|
request must use it."""
|
||||||
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-env")
|
||||||
|
headers = vlm_api_headers(ModelConfig(provider="openai", name="gpt-4-vision"))
|
||||||
|
assert headers == {"Authorization": "Bearer sk-env"}
|
||||||
|
|
||||||
|
def test_custom_base_url_never_gets_the_openai_environment_key(self, monkeypatch):
|
||||||
|
"""A self-hosted endpoint must not receive the public OpenAI key."""
|
||||||
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-env")
|
||||||
|
headers = vlm_api_headers(
|
||||||
|
ModelConfig(
|
||||||
|
provider="openai", name="qwen-vl", base_url="http://my-vllm:8000"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert headers == {}
|
||||||
|
|
||||||
|
def test_api_key_becomes_bearer_header(self):
|
||||||
|
headers = vlm_api_headers(
|
||||||
|
ModelConfig(
|
||||||
|
provider="openai",
|
||||||
|
name="gpt-4-vision",
|
||||||
|
base_url="http://my-vllm:8000",
|
||||||
|
api_key="sk-vlm",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert headers == {"Authorization": "Bearer sk-vlm"}
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="module")
|
@pytest.fixture(scope="module")
|
||||||
def vcr_cassette_dir():
|
def vcr_cassette_dir():
|
||||||
return str(Path(__file__).parent / "cassettes" / "test_converters")
|
return str(Path(__file__).parent / "cassettes" / "test_converters")
|
||||||
|
|
@ -970,6 +1005,21 @@ class TestDoclingLocalConverter:
|
||||||
pic_desc = converter.config.processing.conversion_options.picture_description
|
pic_desc = converter.config.processing.conversion_options.picture_description
|
||||||
assert pic_desc.timeout == 120
|
assert pic_desc.timeout == 120
|
||||||
|
|
||||||
|
def test_picture_description_api_key_becomes_auth_header(self, config):
|
||||||
|
"""The VLM endpoint is reached over plain HTTP, so its key travels as
|
||||||
|
an Authorization header on the picture-description options."""
|
||||||
|
config.processing.pictures = "description"
|
||||||
|
config.processing.conversion_options.picture_description.model.api_key = (
|
||||||
|
"sk-vlm"
|
||||||
|
)
|
||||||
|
converter = DoclingLocalConverter(config)
|
||||||
|
|
||||||
|
opts = converter._build_pipeline_options()
|
||||||
|
|
||||||
|
assert opts.picture_description_options.headers == {
|
||||||
|
"Authorization": "Bearer sk-vlm"
|
||||||
|
}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.vcr()
|
@pytest.mark.vcr()
|
||||||
async def test_picture_description_end_to_end(
|
async def test_picture_description_end_to_end(
|
||||||
|
|
@ -1551,6 +1601,9 @@ class TestDoclingServeConverterPictureDescription:
|
||||||
)
|
)
|
||||||
config.processing.conversion_options.picture_description.timeout = 120
|
config.processing.conversion_options.picture_description.timeout = 120
|
||||||
config.processing.conversion_options.picture_description.max_tokens = 300
|
config.processing.conversion_options.picture_description.max_tokens = 300
|
||||||
|
config.processing.conversion_options.picture_description.model.api_key = (
|
||||||
|
"sk-vlm"
|
||||||
|
)
|
||||||
config.prompts.picture_description = "Test prompt for picture description"
|
config.prompts.picture_description = "Test prompt for picture description"
|
||||||
converter = DoclingServeConverter(config)
|
converter = DoclingServeConverter(config)
|
||||||
|
|
||||||
|
|
@ -1583,6 +1636,7 @@ class TestDoclingServeConverterPictureDescription:
|
||||||
assert api_config["params"]["max_completion_tokens"] == 300
|
assert api_config["params"]["max_completion_tokens"] == 300
|
||||||
assert api_config["prompt"] == "Test prompt for picture description"
|
assert api_config["prompt"] == "Test prompt for picture description"
|
||||||
assert api_config["timeout"] == 120
|
assert api_config["timeout"] == 120
|
||||||
|
assert api_config["headers"] == {"Authorization": "Bearer sk-vlm"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_picture_description_disabled_by_default(self, config):
|
async def test_picture_description_disabled_by_default(self, config):
|
||||||
|
|
|
||||||
|
|
@ -147,7 +147,7 @@ def _stub_provider_probe(monkeypatch):
|
||||||
so database-integrity tests don't depend on a live Ollama. Provider tests
|
so database-integrity tests don't depend on a live Ollama. Provider tests
|
||||||
re-patch this with their own behavior."""
|
re-patch this with their own behavior."""
|
||||||
|
|
||||||
async def probe(_client, _url):
|
async def probe(_client, _url, _headers):
|
||||||
return (
|
return (
|
||||||
True,
|
True,
|
||||||
None,
|
None,
|
||||||
|
|
@ -693,14 +693,27 @@ def test_api_key_required_for_openai_without_base_url():
|
||||||
assert any("OPENAI_API_KEY" in d for d in result.details)
|
assert any("OPENAI_API_KEY" in d for d in result.details)
|
||||||
|
|
||||||
|
|
||||||
|
def test_api_key_not_required_when_config_supplies_it():
|
||||||
|
"""A key in the config is the point of the field; doctor must not demand
|
||||||
|
the environment variable as well."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="openai", name="x", vector_dim=4, api_key="sk-inline"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert _check_api_keys(config, {}).severity is Severity.OK
|
||||||
|
|
||||||
|
|
||||||
def test_active_models_includes_picture_description_when_enabled():
|
def test_active_models_includes_picture_description_when_enabled():
|
||||||
config = AppConfig(processing=ProcessingConfig(pictures="description"))
|
config = AppConfig(processing=ProcessingConfig(pictures="description"))
|
||||||
names = [name for _p, name, _b in _active_models(config)]
|
names = [model.name for model in _active_models(config)]
|
||||||
assert "ministral-3" in names
|
assert "ministral-3" in names
|
||||||
|
|
||||||
|
|
||||||
def test_active_models_excludes_picture_description_by_default():
|
def test_active_models_excludes_picture_description_by_default():
|
||||||
names = [name for _p, name, _b in _active_models(AppConfig())]
|
names = [model.name for model in _active_models(AppConfig())]
|
||||||
assert "ministral-3" not in names
|
assert "ministral-3" not in names
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -804,6 +817,37 @@ def test_provider_targets_includes_docling_serve():
|
||||||
assert targets["http://docling:5001/health"]["kind"] == "docling-serve"
|
assert targets["http://docling:5001/health"]["kind"] == "docling-serve"
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_targets_carries_model_api_key():
|
||||||
|
"""A secured endpoint answers the probe only with its key."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="openai",
|
||||||
|
name="x",
|
||||||
|
vector_dim=4,
|
||||||
|
base_url="http://vllm:8000/v1",
|
||||||
|
api_key="sk-probe",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
targets, _ = _provider_targets(config)
|
||||||
|
entry = targets["http://vllm:8000/v1/models"]
|
||||||
|
assert entry["headers"] == {"Authorization": "Bearer sk-probe"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_targets_carries_docling_serve_api_key():
|
||||||
|
config = AppConfig(
|
||||||
|
processing=ProcessingConfig(converter="docling-serve"),
|
||||||
|
providers=ProvidersConfig(
|
||||||
|
docling_serve=DoclingServeConfig(
|
||||||
|
base_url="http://docling:5001", api_key="ds-key"
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
targets, _ = _provider_targets(config)
|
||||||
|
assert targets["http://docling:5001/health"]["headers"] == {"X-Api-Key": "ds-key"}
|
||||||
|
|
||||||
|
|
||||||
def test_provider_targets_collects_local_providers():
|
def test_provider_targets_collects_local_providers():
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
|
|
@ -817,7 +861,7 @@ def test_provider_targets_collects_local_providers():
|
||||||
|
|
||||||
|
|
||||||
def _fake_probe(result):
|
def _fake_probe(result):
|
||||||
async def probe(_client, _url):
|
async def probe(_client, _url, _headers):
|
||||||
return result
|
return result
|
||||||
|
|
||||||
return probe
|
return probe
|
||||||
|
|
@ -901,12 +945,12 @@ async def test_run_doctor_includes_provider_results(temp_db_path, monkeypatch):
|
||||||
assert not report.failed
|
assert not report.failed
|
||||||
|
|
||||||
|
|
||||||
async def _probe_with_handler(handler):
|
async def _probe_with_handler(handler, headers: dict[str, str] | None = None):
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
transport = httpx.MockTransport(handler)
|
transport = httpx.MockTransport(handler)
|
||||||
async with httpx.AsyncClient(transport=transport) as client:
|
async with httpx.AsyncClient(transport=transport) as client:
|
||||||
return await _probe_endpoint(client, "http://x")
|
return await _probe_endpoint(client, "http://x", headers or {})
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -940,6 +984,22 @@ async def test_probe_endpoint_http_error_status():
|
||||||
assert error is not None and "503" in error
|
assert error is not None and "503" in error
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_probe_endpoint_sends_headers():
|
||||||
|
"""A secured endpoint needs its key on the probe request too."""
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
seen: dict[str, str] = {}
|
||||||
|
|
||||||
|
def handler(request):
|
||||||
|
seen.update(request.headers)
|
||||||
|
return httpx.Response(200, json={})
|
||||||
|
|
||||||
|
await _probe_with_handler(handler, {"Authorization": "Bearer sk-probe"})
|
||||||
|
|
||||||
|
assert seen["authorization"] == "Bearer sk-probe"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_probe_endpoint_connection_error():
|
async def test_probe_endpoint_connection_error():
|
||||||
import httpx
|
import httpx
|
||||||
|
|
|
||||||
|
|
@ -195,3 +195,74 @@ def test_voyageai_to_pil_rejects_unsupported_type():
|
||||||
|
|
||||||
with pytest.raises(TypeError, match="Unsupported image type"):
|
with pytest.raises(TypeError, match="Unsupported image type"):
|
||||||
_to_pil("not an image") # ty: ignore[invalid-argument-type]
|
_to_pil("not an image") # ty: ignore[invalid-argument-type]
|
||||||
|
|
||||||
|
|
||||||
|
def test_openai_embedder_api_key_from_config():
|
||||||
|
"""A config-supplied api_key reaches the embedding client."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="openai",
|
||||||
|
name="some-local-model",
|
||||||
|
vector_dim=768,
|
||||||
|
base_url="http://localhost:8000/v1",
|
||||||
|
api_key="sk-vendor-a",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
embedder = get_embedder(config)
|
||||||
|
|
||||||
|
assert embedder._embedder.model._client.api_key == "sk-vendor-a" # ty: ignore[unresolved-attribute]
|
||||||
|
|
||||||
|
|
||||||
|
def test_ollama_embedder_api_key_from_config():
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="ollama",
|
||||||
|
name="qwen3-embedding:4b",
|
||||||
|
vector_dim=512,
|
||||||
|
api_key="sk-proxy",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
embedder = get_embedder(config)
|
||||||
|
|
||||||
|
assert embedder._embedder.model._client.api_key == "sk-proxy" # ty: ignore[unresolved-attribute]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("multimodal", [False, True])
|
||||||
|
def test_vllm_embedder_api_key_from_config(multimodal):
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="vllm",
|
||||||
|
name="Qwen/Qwen3-VL-Embedding-8B",
|
||||||
|
vector_dim=1024,
|
||||||
|
base_url="http://vllm:8000/v1",
|
||||||
|
api_key="sk-vllm",
|
||||||
|
multimodal=multimodal,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
embedder = get_embedder(config)
|
||||||
|
|
||||||
|
assert embedder._headers()["Authorization"] == "Bearer sk-vllm" # ty: ignore[unresolved-attribute]
|
||||||
|
|
||||||
|
|
||||||
|
def test_embedder_api_key_rejected_on_unplumbed_provider():
|
||||||
|
"""Providers whose client we never build read their own vendor variable;
|
||||||
|
an api_key there would be silently dropped."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="voyageai", name="voyage-3", vector_dim=1024, api_key="sk-x"
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="api_key is not supported"):
|
||||||
|
get_embedder(config)
|
||||||
|
|
|
||||||
|
|
@ -134,6 +134,33 @@ class TestGetReranker:
|
||||||
with pytest.raises(ValueError, match="multimodal"):
|
with pytest.raises(ValueError, match="multimodal"):
|
||||||
get_reranker(config)
|
get_reranker(config)
|
||||||
|
|
||||||
|
def test_vllm_provider_api_key_from_config(self):
|
||||||
|
pytest.importorskip("haiku.rag.reranking.vllm")
|
||||||
|
|
||||||
|
config = AppConfig(
|
||||||
|
reranking=RerankingConfig(
|
||||||
|
model=ModelConfig(
|
||||||
|
provider="vllm",
|
||||||
|
name="BAAI/bge-reranker-v2-m3",
|
||||||
|
base_url="http://vllm:8000",
|
||||||
|
api_key="sk-rerank",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
reranker = get_reranker(config)
|
||||||
|
assert reranker._headers["Authorization"] == "Bearer sk-rerank" # ty: ignore[unresolved-attribute]
|
||||||
|
|
||||||
|
def test_api_key_rejected_on_unplumbed_provider(self):
|
||||||
|
"""Providers whose client we never build read their own vendor
|
||||||
|
variable; an api_key there would be silently dropped."""
|
||||||
|
config = AppConfig(
|
||||||
|
reranking=RerankingConfig(
|
||||||
|
model=ModelConfig(provider="cohere", name="rerank-v3.5", api_key="sk-x")
|
||||||
|
)
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="api_key is not supported"):
|
||||||
|
get_reranker(config)
|
||||||
|
|
||||||
def test_multimodal_vllm_provider_builds_reranker(self):
|
def test_multimodal_vllm_provider_builds_reranker(self):
|
||||||
pytest.importorskip("haiku.rag.reranking.vllm")
|
pytest.importorskip("haiku.rag.reranking.vllm")
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
|
|
|
||||||
|
|
@ -970,3 +970,45 @@ def test_get_package_versions_reports_missing_docling(monkeypatch):
|
||||||
monkeypatch.setattr(importlib_metadata, "version", fake_version)
|
monkeypatch.setattr(importlib_metadata, "version", fake_version)
|
||||||
|
|
||||||
assert get_package_versions()["docling"] == "not installed"
|
assert get_package_versions()["docling"] == "not installed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_model_openai_api_key_from_config():
|
||||||
|
"""A config-supplied api_key reaches the client, so several
|
||||||
|
openai-compatible endpoints can each carry their own key."""
|
||||||
|
result = get_model(
|
||||||
|
ModelConfig(
|
||||||
|
provider="openai",
|
||||||
|
name="qwen3.6",
|
||||||
|
base_url="http://vllm:8000/v1",
|
||||||
|
api_key="sk-vendor-a",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert result.client.api_key == "sk-vendor-a"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_model_openai_api_key_without_base_url_overrides_env():
|
||||||
|
result = get_model(
|
||||||
|
ModelConfig(provider="openai", name="gpt-4o", api_key="sk-vendor-b")
|
||||||
|
)
|
||||||
|
assert result.client.api_key == "sk-vendor-b"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_model_ollama_api_key_from_config():
|
||||||
|
result = get_model(
|
||||||
|
ModelConfig(
|
||||||
|
provider="ollama",
|
||||||
|
name="gpt-oss",
|
||||||
|
base_url="http://remote-ollama:11434/v1",
|
||||||
|
api_key="sk-proxy",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert result.client.api_key == "sk-proxy"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_model_api_key_rejected_on_unplumbed_provider():
|
||||||
|
"""Providers whose client we never build read their own vendor variable;
|
||||||
|
an api_key there would be silently dropped."""
|
||||||
|
with pytest.raises(ValueError, match="api_key is not supported"):
|
||||||
|
get_model(
|
||||||
|
ModelConfig(provider="anthropic", name="claude-sonnet-4-5", api_key="sk-x")
|
||||||
|
)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue