Share the VLM URL builder and centralize the embedder empty guard

This commit is contained in:
Yiorgis Gozadinos 2026-06-29 14:37:13 +03:00
parent 654cb2b94c
commit afd6e3ae00
No known key found for this signature in database
8 changed files with 59 additions and 122 deletions

View file

@ -7,6 +7,23 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from docling_core.types.doc.document import DoclingDocument
from haiku.rag.config import AppConfig
from haiku.rag.config.models import ModelConfig
def vlm_api_url(config: "AppConfig", model: "ModelConfig") -> str:
"""Construct the VLM chat-completions URL for a picture-description model."""
if model.base_url:
return f"{model.base_url.rstrip('/')}/v1/chat/completions"
if model.provider == "ollama":
return f"{config.providers.ollama.base_url.rstrip('/')}/v1/chat/completions"
if model.provider == "openai":
return "https://api.openai.com/v1/chat/completions"
raise ValueError(f"Unsupported VLM provider: {model.provider}")
class DocumentConverter(ABC):
"""Abstract base class for document converters.

View file

@ -5,7 +5,7 @@ from pathlib import Path
from typing import TYPE_CHECKING, ClassVar
from haiku.rag.config import AppConfig
from haiku.rag.converters.base import DocumentConverter
from haiku.rag.converters.base import DocumentConverter, vlm_api_url
from haiku.rag.converters.text_utils import TextFileHandler
if TYPE_CHECKING:
@ -13,7 +13,7 @@ if TYPE_CHECKING:
from docling.document_converter import FormatOption
from docling_core.types.doc.document import DoclingDocument
from haiku.rag.config.models import ConversionOptions, ModelConfig
from haiku.rag.config.models import ConversionOptions
class DoclingLocalConverter(DocumentConverter):
@ -62,21 +62,6 @@ class DoclingLocalConverter(DocumentConverter):
"""Return list of file extensions supported by this converter."""
return self.docling_extensions + TextFileHandler.text_extensions
def _get_vlm_api_url(self, model: "ModelConfig") -> str:
"""Construct VLM API URL from model config."""
if model.base_url:
base = model.base_url.rstrip("/")
return f"{base}/v1/chat/completions"
if model.provider == "ollama":
base = self.config.providers.ollama.base_url.rstrip("/")
return f"{base}/v1/chat/completions"
if model.provider == "openai":
return "https://api.openai.com/v1/chat/completions"
raise ValueError(f"Unsupported VLM provider: {model.provider}")
def _get_ocr_options(self, opts: "ConversionOptions"):
"""Get OCR options based on configuration."""
from docling.datamodel.pipeline_options import (
@ -145,7 +130,7 @@ class DoclingLocalConverter(DocumentConverter):
pipeline_options.enable_remote_services = True
pipeline_options.picture_description_options = PictureDescriptionApiOptions(
url=AnyUrl(self._get_vlm_api_url(pic_desc.model)),
url=AnyUrl(vlm_api_url(self.config, pic_desc.model)),
params=dict(
model=pic_desc.model.name,
max_completion_tokens=pic_desc.max_tokens,

View file

@ -6,15 +6,13 @@ from pathlib import Path
from typing import TYPE_CHECKING, ClassVar
from haiku.rag.config import AppConfig
from haiku.rag.converters.base import DocumentConverter
from haiku.rag.converters.base import DocumentConverter, vlm_api_url
from haiku.rag.converters.text_utils import TextFileHandler
from haiku.rag.providers.docling_serve import DoclingServeClient
if TYPE_CHECKING:
from docling_core.types.doc.document import DoclingDocument
from haiku.rag.config.models import ModelConfig
class DoclingServeConverter(DocumentConverter):
"""Converter that uses docling-serve for document conversion.
@ -69,21 +67,6 @@ class DoclingServeConverter(DocumentConverter):
"""Return list of file extensions supported by this converter."""
return self.docling_serve_extensions + TextFileHandler.text_extensions
def _get_vlm_api_url(self, model: "ModelConfig") -> str:
"""Construct VLM API URL from model config."""
if model.base_url:
base = model.base_url.rstrip("/")
return f"{base}/v1/chat/completions"
if model.provider == "ollama":
base = self.config.providers.ollama.base_url.rstrip("/")
return f"{base}/v1/chat/completions"
if model.provider == "openai":
return "https://api.openai.com/v1/chat/completions"
raise ValueError(f"Unsupported VLM provider: {model.provider}")
def _build_conversion_data(self) -> dict[str, str | list[str]]:
"""Build form data for conversion request.
@ -124,7 +107,7 @@ class DoclingServeConverter(DocumentConverter):
if runs_vlm:
prompt = self.config.prompts.picture_description
picture_description_api = {
"url": self._get_vlm_api_url(pic_desc.model),
"url": vlm_api_url(self.config, pic_desc.model),
"params": {
"model": pic_desc.model.name,
"max_completion_tokens": pic_desc.max_tokens,

View file

@ -54,6 +54,9 @@ class EmbedderWrapper:
"""Embed documents/chunks for indexing."""
if not texts:
return []
return await self._embed_documents(texts)
async def _embed_documents(self, texts: list[str]) -> list[list[float]]:
assert self._embedder is not None
result = await self._embedder.embed_documents(texts)
return [list(e) for e in result.embeddings]

View file

@ -43,9 +43,7 @@ class CohereMultimodalEmbedder(EmbedderWrapper):
rows = await self._embed_texts([text], "search_query")
return rows[0]
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
async def _embed_documents(self, texts: list[str]) -> list[list[float]]:
return await self._embed_texts(texts, "search_document")
async def embed_image(self, image: "bytes | PILImage.Image") -> list[float]:

View file

@ -86,9 +86,7 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
)
return rows[0]
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
async def _embed_documents(self, texts: list[str]) -> list[list[float]]:
return await self._post(
{
"model": self._model_name,

View file

@ -37,9 +37,7 @@ class VoyageMultimodalEmbedder(EmbedderWrapper):
)
return list(result.embeddings[0])
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
async def _embed_documents(self, texts: list[str]) -> list[list[float]]:
result = await self._client.multimodal_embed(
inputs=[[text] for text in texts],
model=self._model_name,

View file

@ -9,12 +9,43 @@ import pytest
from docling_core.types.doc.document import DoclingDocument
from haiku.rag.config import AppConfig
from haiku.rag.config.models import ModelConfig
from haiku.rag.converters import get_converter
from haiku.rag.converters.base import vlm_api_url
from haiku.rag.converters.docling_local import DoclingLocalConverter
from haiku.rag.converters.docling_serve import DoclingServeConverter
from haiku.rag.converters.text_utils import TextFileHandler
class TestVlmApiUrl:
"""URL construction for picture-description VLM models (shared by both converters)."""
def test_ollama_uses_provider_base_url(self):
url = vlm_api_url(
AppConfig(), ModelConfig(provider="ollama", name="ministral-3")
)
assert url == "http://localhost:11434/v1/chat/completions"
def test_custom_base_url_takes_precedence(self):
url = vlm_api_url(
AppConfig(),
ModelConfig(
provider="openai", name="gpt-4-vision", base_url="http://my-vllm:8000"
),
)
assert url == "http://my-vllm:8000/v1/chat/completions"
def test_openai_uses_public_endpoint(self):
url = vlm_api_url(
AppConfig(), ModelConfig(provider="openai", name="gpt-4-vision")
)
assert url == "https://api.openai.com/v1/chat/completions"
def test_unsupported_provider_raises(self):
with pytest.raises(ValueError, match="Unsupported VLM provider"):
vlm_api_url(AppConfig(), ModelConfig(provider="unsupported", name="test"))
@pytest.fixture(scope="module")
def vcr_cassette_dir():
return str(Path(__file__).parent / "cassettes" / "test_converters")
@ -763,44 +794,6 @@ class TestDoclingLocalConverter:
"Pages should have image data when generate_page_images=True"
)
def test_get_vlm_api_url_with_ollama(self, config):
"""Test VLM API URL construction for Ollama provider."""
converter = DoclingLocalConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(provider="ollama", name="ministral-3")
url = converter._get_vlm_api_url(model)
assert url == "http://localhost:11434/v1/chat/completions"
def test_get_vlm_api_url_with_custom_base_url(self, config):
"""Test VLM API URL construction with custom base_url."""
converter = DoclingLocalConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(
provider="openai", name="gpt-4-vision", base_url="http://my-vllm:8000"
)
url = converter._get_vlm_api_url(model)
assert url == "http://my-vllm:8000/v1/chat/completions"
def test_get_vlm_api_url_with_openai(self, config):
"""Test VLM API URL construction for OpenAI provider."""
converter = DoclingLocalConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(provider="openai", name="gpt-4-vision")
url = converter._get_vlm_api_url(model)
assert url == "https://api.openai.com/v1/chat/completions"
def test_get_vlm_api_url_unsupported_provider(self, config):
"""Test VLM API URL construction raises error for unsupported provider."""
converter = DoclingLocalConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(provider="unsupported", name="test")
with pytest.raises(ValueError, match="Unsupported VLM provider"):
converter._get_vlm_api_url(model)
def test_ocr_engine_config_applied(self, config):
"""Test that ocr_engine config is stored correctly."""
config.processing.conversion_options.ocr_engine = "rapidocr"
@ -1266,44 +1259,6 @@ class TestDoclingServeConverterPictureDescription:
config.providers.docling_serve.api_key = ""
return config
def test_get_vlm_api_url_with_ollama(self, config):
"""Test VLM API URL construction for Ollama provider."""
converter = DoclingServeConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(provider="ollama", name="ministral-3")
url = converter._get_vlm_api_url(model)
assert url == "http://localhost:11434/v1/chat/completions"
def test_get_vlm_api_url_with_custom_base_url(self, config):
"""Test VLM API URL construction with custom base_url."""
converter = DoclingServeConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(
provider="openai", name="gpt-4-vision", base_url="http://my-vllm:8000"
)
url = converter._get_vlm_api_url(model)
assert url == "http://my-vllm:8000/v1/chat/completions"
def test_get_vlm_api_url_with_openai(self, config):
"""Test VLM API URL construction for OpenAI provider."""
converter = DoclingServeConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(provider="openai", name="gpt-4-vision")
url = converter._get_vlm_api_url(model)
assert url == "https://api.openai.com/v1/chat/completions"
def test_get_vlm_api_url_unsupported_provider(self, config):
"""Test VLM API URL construction raises error for unsupported provider."""
converter = DoclingServeConverter(config)
from haiku.rag.config.models import ModelConfig
model = ModelConfig(provider="unsupported", name="test")
with pytest.raises(ValueError, match="Unsupported VLM provider"):
converter._get_vlm_api_url(model)
@pytest.mark.asyncio
async def test_picture_description_options_passed_to_api(self, config):
"""Picture-description options reach the docling-serve API when the