auto-append /v1 to per-model Ollama base_url, fix flaky tests
This commit is contained in:
parent
65d9c74224
commit
aa3e9406cf
8 changed files with 161 additions and 20 deletions
|
|
@ -167,7 +167,9 @@ def get_embedder(config: AppConfig = Config) -> EmbedderWrapper:
|
||||||
|
|
||||||
if provider == "ollama":
|
if provider == "ollama":
|
||||||
# Use model-level base_url if set, otherwise fall back to providers config
|
# Use model-level base_url if set, otherwise fall back to providers config
|
||||||
base_url = embedding_model.base_url or f"{config.providers.ollama.base_url}/v1"
|
base_url = embedding_model.base_url or config.providers.ollama.base_url
|
||||||
|
if not base_url.rstrip("/").endswith("/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),
|
||||||
|
|
|
||||||
|
|
@ -158,8 +158,11 @@ def get_model(
|
||||||
model_settings, OpenAIChatModelSettings, model_config
|
model_settings, OpenAIChatModelSettings, model_config
|
||||||
)
|
)
|
||||||
|
|
||||||
# Use model-level base_url if set, otherwise fall back to providers config
|
# Ollama's OpenAI-compatible API lives under /v1. Append it if the
|
||||||
base_url = model_config.base_url or f"{app_config.providers.ollama.base_url}/v1"
|
# configured base_url doesn't already include it.
|
||||||
|
base_url = model_config.base_url or app_config.providers.ollama.base_url
|
||||||
|
if not base_url.rstrip("/").endswith("/v1"):
|
||||||
|
base_url = base_url.rstrip("/") + "/v1"
|
||||||
|
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
|
|
|
||||||
File diff suppressed because one or more lines are too long
|
|
@ -5,6 +5,7 @@ from haiku.rag.client.documents import (
|
||||||
_store_document_with_chunks,
|
_store_document_with_chunks,
|
||||||
_update_document_with_chunks,
|
_update_document_with_chunks,
|
||||||
)
|
)
|
||||||
|
from haiku.rag.config import AppConfig
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.store.engine import Store
|
||||||
from haiku.rag.store.models.document_item import (
|
from haiku.rag.store.models.document_item import (
|
||||||
DocumentItem,
|
DocumentItem,
|
||||||
|
|
@ -646,9 +647,11 @@ class TestPictureDataPreservedThroughRoundTrip:
|
||||||
|
|
||||||
docling_doc = _docling_doc_with_picture()
|
docling_doc = _docling_doc_with_picture()
|
||||||
|
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
config = AppConfig()
|
||||||
# Preservation only kicks in under modes that retain picture bytes.
|
# Preservation only kicks in under modes that retain picture bytes.
|
||||||
rag._config.processing.pictures = "image"
|
config.processing.pictures = "image"
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, config=config, create=True) as rag:
|
||||||
document = Document(content="Hello world", uri="test://doc")
|
document = Document(content="Hello world", uri="test://doc")
|
||||||
document.set_docling(docling_doc)
|
document.set_docling(docling_doc)
|
||||||
created = await _store_document_with_chunks(rag, document, [], docling_doc)
|
created = await _store_document_with_chunks(rag, document, [], docling_doc)
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,5 @@
|
||||||
|
import logging
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
|
|
@ -7,6 +9,16 @@ from haiku.rag.config.loader import (
|
||||||
generate_default_config,
|
generate_default_config,
|
||||||
load_yaml_config,
|
load_yaml_config,
|
||||||
)
|
)
|
||||||
|
from haiku.rag.config.loader import logger as loader_logger
|
||||||
|
|
||||||
|
|
||||||
|
class _ListHandler(logging.Handler):
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__(level=logging.WARNING)
|
||||||
|
self.records: list[logging.LogRecord] = []
|
||||||
|
|
||||||
|
def emit(self, record: logging.LogRecord) -> None:
|
||||||
|
self.records.append(record)
|
||||||
|
|
||||||
|
|
||||||
def test_load_yaml_config(tmp_path):
|
def test_load_yaml_config(tmp_path):
|
||||||
|
|
@ -237,7 +249,7 @@ def _write(tmp_path, body: str):
|
||||||
return p
|
return p
|
||||||
|
|
||||||
|
|
||||||
def test_load_yaml_legacy_picture_description_maps_to_description(tmp_path, caplog):
|
def test_load_yaml_legacy_picture_description_maps_to_description(tmp_path):
|
||||||
"""`picture_description.enabled=true` (with or without the image flag)
|
"""`picture_description.enabled=true` (with or without the image flag)
|
||||||
maps to `processing.pictures: description`."""
|
maps to `processing.pictures: description`."""
|
||||||
config_file = _write(
|
config_file = _write(
|
||||||
|
|
@ -251,15 +263,21 @@ processing:
|
||||||
timeout: 120
|
timeout: 120
|
||||||
""",
|
""",
|
||||||
)
|
)
|
||||||
with caplog.at_level("WARNING", logger="haiku.rag.config.loader"):
|
handler = _ListHandler()
|
||||||
|
loader_logger.addHandler(handler)
|
||||||
|
try:
|
||||||
data = load_yaml_config(config_file)
|
data = load_yaml_config(config_file)
|
||||||
|
finally:
|
||||||
|
loader_logger.removeHandler(handler)
|
||||||
cfg = AppConfig.model_validate(data)
|
cfg = AppConfig.model_validate(data)
|
||||||
assert cfg.processing.pictures == "description"
|
assert cfg.processing.pictures == "description"
|
||||||
assert cfg.processing.conversion_options.picture_description.timeout == 120
|
assert cfg.processing.conversion_options.picture_description.timeout == 120
|
||||||
assert any("picture_description.enabled=true" in m.message for m in caplog.records)
|
assert any(
|
||||||
|
"picture_description.enabled=true" in r.getMessage() for r in handler.records
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_load_yaml_legacy_generate_picture_images_maps_to_image(tmp_path, caplog):
|
def test_load_yaml_legacy_generate_picture_images_maps_to_image(tmp_path):
|
||||||
"""`generate_picture_images=true` alone maps to `pictures: image`."""
|
"""`generate_picture_images=true` alone maps to `pictures: image`."""
|
||||||
config_file = _write(
|
config_file = _write(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
|
|
@ -269,14 +287,20 @@ processing:
|
||||||
generate_picture_images: true
|
generate_picture_images: true
|
||||||
""",
|
""",
|
||||||
)
|
)
|
||||||
with caplog.at_level("WARNING", logger="haiku.rag.config.loader"):
|
handler = _ListHandler()
|
||||||
|
loader_logger.addHandler(handler)
|
||||||
|
try:
|
||||||
data = load_yaml_config(config_file)
|
data = load_yaml_config(config_file)
|
||||||
|
finally:
|
||||||
|
loader_logger.removeHandler(handler)
|
||||||
cfg = AppConfig.model_validate(data)
|
cfg = AppConfig.model_validate(data)
|
||||||
assert cfg.processing.pictures == "image"
|
assert cfg.processing.pictures == "image"
|
||||||
assert any("generate_picture_images=true" in m.message for m in caplog.records)
|
assert any(
|
||||||
|
"generate_picture_images=true" in r.getMessage() for r in handler.records
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_load_yaml_no_legacy_fields_keeps_default_none(tmp_path, caplog):
|
def test_load_yaml_no_legacy_fields_keeps_default_none(tmp_path):
|
||||||
"""Empty processing block leaves the default `none` mode untouched and
|
"""Empty processing block leaves the default `none` mode untouched and
|
||||||
does not warn."""
|
does not warn."""
|
||||||
config_file = _write(
|
config_file = _write(
|
||||||
|
|
@ -286,11 +310,15 @@ processing:
|
||||||
chunk_size: 256
|
chunk_size: 256
|
||||||
""",
|
""",
|
||||||
)
|
)
|
||||||
with caplog.at_level("WARNING", logger="haiku.rag.config.loader"):
|
handler = _ListHandler()
|
||||||
|
loader_logger.addHandler(handler)
|
||||||
|
try:
|
||||||
data = load_yaml_config(config_file)
|
data = load_yaml_config(config_file)
|
||||||
|
finally:
|
||||||
|
loader_logger.removeHandler(handler)
|
||||||
cfg = AppConfig.model_validate(data)
|
cfg = AppConfig.model_validate(data)
|
||||||
assert cfg.processing.pictures == "none"
|
assert cfg.processing.pictures == "none"
|
||||||
assert not caplog.records
|
assert not handler.records
|
||||||
|
|
||||||
|
|
||||||
def test_load_yaml_explicit_pictures_wins_over_legacy(tmp_path):
|
def test_load_yaml_explicit_pictures_wins_over_legacy(tmp_path):
|
||||||
|
|
|
||||||
|
|
@ -75,3 +75,39 @@ def test_unsupported_provider_raises():
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="Unsupported embedding provider"):
|
with pytest.raises(ValueError, match="Unsupported embedding provider"):
|
||||||
get_embedder(custom_config)
|
get_embedder(custom_config)
|
||||||
|
|
||||||
|
|
||||||
|
def test_ollama_embedder_appends_v1_when_missing():
|
||||||
|
"""Per-model base_url without /v1 should get it appended for Ollama."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="ollama",
|
||||||
|
name="qwen3-embedding:4b",
|
||||||
|
vector_dim=2560,
|
||||||
|
base_url="http://my-ollama:11434",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
embedder = get_embedder(config)
|
||||||
|
pa_model = embedder._embedder._model # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
||||||
|
assert str(pa_model.base_url).rstrip("/").endswith("/v1") # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
||||||
|
|
||||||
|
|
||||||
|
def test_ollama_embedder_does_not_double_append_v1():
|
||||||
|
"""If the user already includes /v1 we leave it alone."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="ollama",
|
||||||
|
name="qwen3-embedding:4b",
|
||||||
|
vector_dim=2560,
|
||||||
|
base_url="http://my-ollama:11434/v1",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
embedder = get_embedder(config)
|
||||||
|
pa_model = embedder._embedder._model # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
||||||
|
url = str(pa_model.base_url).rstrip("/") # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
||||||
|
assert url.endswith("/v1")
|
||||||
|
assert not url.endswith("/v1/v1")
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,7 @@ from pydantic_ai.usage import RunUsage
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.client.search import _populate_image_data
|
from haiku.rag.client.search import _populate_image_data
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import AppConfig, Config
|
||||||
from haiku.rag.context import expand_with_items
|
from haiku.rag.context import expand_with_items
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.store.models.document_item import DocumentItem
|
from haiku.rag.store.models.document_item import DocumentItem
|
||||||
|
|
@ -167,6 +167,7 @@ async def test_expand_context_preserves_picture_refs_with_empty_text(temp_db_pat
|
||||||
assert "picture" in out.labels
|
assert "picture" in out.labels
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.vcr()
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_rechunk_preserves_picture_data(temp_db_path):
|
async def test_rechunk_preserves_picture_data(temp_db_path):
|
||||||
"""``rebuild --rechunk`` keeps ``picture_data`` for every picture row."""
|
"""``rebuild --rechunk`` keeps ``picture_data`` for every picture row."""
|
||||||
|
|
@ -177,8 +178,10 @@ async def test_rechunk_preserves_picture_data(temp_db_path):
|
||||||
|
|
||||||
docling_doc = _docling_doc_with_picture()
|
docling_doc = _docling_doc_with_picture()
|
||||||
|
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
config = AppConfig()
|
||||||
rag._config.processing.pictures = "image"
|
config.processing.pictures = "image"
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, config=config, create=True) as rag:
|
||||||
document = Document(content="x", uri="test://doc")
|
document = Document(content="x", uri="test://doc")
|
||||||
document.set_docling(docling_doc)
|
document.set_docling(docling_doc)
|
||||||
created = await _store_document_with_chunks(rag, document, [], docling_doc)
|
created = await _store_document_with_chunks(rag, document, [], docling_doc)
|
||||||
|
|
@ -193,6 +196,7 @@ async def test_rechunk_preserves_picture_data(temp_db_path):
|
||||||
assert after.get("#/pictures/0") == before.get("#/pictures/0")
|
assert after.get("#/pictures/0") == before.get("#/pictures/0")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.vcr()
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_update_clears_picture_data_when_mode_none(temp_db_path):
|
async def test_update_clears_picture_data_when_mode_none(temp_db_path):
|
||||||
"""Switching to ``pictures="none"`` and re-running update_document
|
"""Switching to ``pictures="none"`` and re-running update_document
|
||||||
|
|
@ -207,9 +211,11 @@ async def test_update_clears_picture_data_when_mode_none(temp_db_path):
|
||||||
|
|
||||||
docling_doc = _docling_doc_with_picture()
|
docling_doc = _docling_doc_with_picture()
|
||||||
|
|
||||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
config = AppConfig()
|
||||||
|
config.processing.pictures = "image"
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, config=config, create=True) as rag:
|
||||||
# Ingest under "image" so picture bytes land in document_items.
|
# Ingest under "image" so picture bytes land in document_items.
|
||||||
rag._config.processing.pictures = "image"
|
|
||||||
document = Document(content="x", uri="test://doc")
|
document = Document(content="x", uri="test://doc")
|
||||||
document.set_docling(docling_doc)
|
document.set_docling(docling_doc)
|
||||||
created = await _store_document_with_chunks(rag, document, [], docling_doc)
|
created = await _store_document_with_chunks(rag, document, [], docling_doc)
|
||||||
|
|
|
||||||
|
|
@ -162,6 +162,27 @@ def test_get_model_ollama_with_settings():
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_model_ollama_appends_v1_to_per_model_base_url():
|
||||||
|
"""Per-model base_url without /v1 should get it appended."""
|
||||||
|
model_config = ModelConfig(
|
||||||
|
provider="ollama", name="qwen3.6", base_url="http://my-ollama:11434"
|
||||||
|
)
|
||||||
|
result = get_model(model_config)
|
||||||
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
assert str(result.client.base_url).rstrip("/").endswith("/v1")
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_model_ollama_does_not_double_append_v1():
|
||||||
|
"""If the per-model base_url already ends with /v1, leave it alone."""
|
||||||
|
model_config = ModelConfig(
|
||||||
|
provider="ollama", name="qwen3.6", base_url="http://my-ollama:11434/v1"
|
||||||
|
)
|
||||||
|
result = get_model(model_config)
|
||||||
|
url = str(result.client.base_url).rstrip("/")
|
||||||
|
assert url.endswith("/v1")
|
||||||
|
assert not url.endswith("/v1/v1")
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_openai():
|
def test_get_model_openai():
|
||||||
"""Test get_model returns OpenAIChatModel for OpenAI."""
|
"""Test get_model returns OpenAIChatModel for OpenAI."""
|
||||||
model_config = ModelConfig(provider="openai", name="gpt-4o")
|
model_config = ModelConfig(provider="openai", name="gpt-4o")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue