Add per-endpoint api_key to model and embedding config

This commit is contained in:
Yiorgis Gozadinos 2026-08-23 23:45:47 +03:00
parent e9f6fea598
commit 2310b7a8b3
No known key found for this signature in database
16 changed files with 440 additions and 56 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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.

View file

@ -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,

View file

@ -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,

View file

@ -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))

View file

@ -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":

View file

@ -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

View file

@ -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()

View file

@ -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

View file

@ -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):

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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")
)