Cover get_model provider branches and visualize_chunk edge paths
Parametrize the get_model tests whose only assertion was the returned model type, extending the table to the reasoning-off, groq-parsed and Bedrock o-series/qwen/unmapped branches. Fold format_citations_no_title into the superset test that already covered its scenario. Add tests for cosine_similarity on parallel vectors, multi-citation Rich rendering, picture rendering with missing and undecodable bytes, a missing docling distribution, and the visualize_chunk short-circuits for absent documents, absent rasters and refless chunks.
This commit is contained in:
parent
1e8e5e9f6f
commit
7c76c3399e
3 changed files with 441 additions and 74 deletions
|
|
@ -2279,3 +2279,242 @@ def test_check_source_accessible_file_uri(tmp_path):
|
||||||
|
|
||||||
assert check_source_accessible(existing.as_uri()) is True
|
assert check_source_accessible(existing.as_uri()) is True
|
||||||
assert check_source_accessible((tmp_path / "gone.txt").as_uri()) is False
|
assert check_source_accessible((tmp_path / "gone.txt").as_uri()) is False
|
||||||
|
|
||||||
|
|
||||||
|
def _bbox_doc(*, with_page_image: bool, pages: tuple[int, ...] = (1,)):
|
||||||
|
"""DoclingDocument with one paragraph per page, each carrying a bbox.
|
||||||
|
|
||||||
|
``with_page_image=False`` produces pages with no raster, so bounding boxes
|
||||||
|
resolve but there is nothing to draw them on.
|
||||||
|
"""
|
||||||
|
from docling_core.types.doc.base import BoundingBox, Size
|
||||||
|
from docling_core.types.doc.document import ImageRef, ProvenanceItem
|
||||||
|
from PIL import Image as PilImageModule
|
||||||
|
|
||||||
|
doc = DoclingDocument(name="bbox-doc")
|
||||||
|
size = Size(width=612.0, height=792.0)
|
||||||
|
for page_no in pages:
|
||||||
|
image = (
|
||||||
|
ImageRef.from_pil(
|
||||||
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
||||||
|
)
|
||||||
|
if with_page_image
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
doc.add_page(page_no=page_no, size=size, image=image)
|
||||||
|
doc.add_text(
|
||||||
|
label=DocItemLabel.PARAGRAPH,
|
||||||
|
text=f"Content on page {page_no}.",
|
||||||
|
prov=ProvenanceItem(
|
||||||
|
page_no=page_no,
|
||||||
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
||||||
|
charspan=(0, 20),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return doc
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_returns_empty_without_page_rasters(temp_db_path):
|
||||||
|
"""Boxes resolve, but a document ingested without page images has nothing
|
||||||
|
to render them onto."""
|
||||||
|
docling_doc = _bbox_doc(with_page_image=False)
|
||||||
|
chunks = [
|
||||||
|
Chunk(
|
||||||
|
content="Content on page 1.",
|
||||||
|
metadata={
|
||||||
|
"doc_item_refs": ["#/texts/0"],
|
||||||
|
"page_numbers": [1],
|
||||||
|
"labels": ["paragraph"],
|
||||||
|
},
|
||||||
|
order=0,
|
||||||
|
embedding=[0.1] * 2560,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-raster")
|
||||||
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
|
||||||
|
assert await client.visualize_chunk(stored[0]) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_skips_pages_without_a_raster(temp_db_path):
|
||||||
|
"""A document where only some pages carry a raster renders just those."""
|
||||||
|
from docling_core.types.doc.base import BoundingBox, Size
|
||||||
|
from docling_core.types.doc.document import ImageRef, ProvenanceItem
|
||||||
|
from PIL import Image as PilImageModule
|
||||||
|
|
||||||
|
docling_doc = DoclingDocument(name="mixed-rasters")
|
||||||
|
size = Size(width=612.0, height=792.0)
|
||||||
|
docling_doc.add_page(
|
||||||
|
page_no=1,
|
||||||
|
size=size,
|
||||||
|
image=ImageRef.from_pil(
|
||||||
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
||||||
|
),
|
||||||
|
)
|
||||||
|
docling_doc.add_page(page_no=2, size=size, image=None)
|
||||||
|
for page_no in (1, 2):
|
||||||
|
docling_doc.add_text(
|
||||||
|
label=DocItemLabel.PARAGRAPH,
|
||||||
|
text=f"Content on page {page_no}.",
|
||||||
|
prov=ProvenanceItem(
|
||||||
|
page_no=page_no,
|
||||||
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
||||||
|
charspan=(0, 20),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
chunks = [
|
||||||
|
Chunk(
|
||||||
|
content="Content on page 1.\nContent on page 2.",
|
||||||
|
metadata={
|
||||||
|
"doc_item_refs": ["#/texts/0", "#/texts/1"],
|
||||||
|
"page_numbers": [1, 2],
|
||||||
|
"labels": ["paragraph", "paragraph"],
|
||||||
|
},
|
||||||
|
order=0,
|
||||||
|
embedding=[0.1] * 2560,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.import_document(
|
||||||
|
docling_doc, chunks, uri="test://mixed-rasters"
|
||||||
|
)
|
||||||
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
|
||||||
|
images = await client.visualize_chunk(stored[0])
|
||||||
|
|
||||||
|
assert len(images) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_returns_empty_when_pages_row_missing(temp_db_path):
|
||||||
|
docling_doc = _bbox_doc(with_page_image=True)
|
||||||
|
chunks = [
|
||||||
|
Chunk(
|
||||||
|
content="Content on page 1.",
|
||||||
|
metadata={
|
||||||
|
"doc_item_refs": ["#/texts/0"],
|
||||||
|
"page_numbers": [1],
|
||||||
|
"labels": ["paragraph"],
|
||||||
|
},
|
||||||
|
order=0,
|
||||||
|
embedding=[0.1] * 2560,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-row")
|
||||||
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
|
||||||
|
async def no_pages_row(document_id):
|
||||||
|
return None
|
||||||
|
|
||||||
|
client.document_repository.get_pages_data = no_pages_row # type: ignore[method-assign]
|
||||||
|
|
||||||
|
assert await client.visualize_chunk(stored[0]) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_skips_box_on_unstored_page(temp_db_path):
|
||||||
|
"""A bounding box referencing a page the document never registered is
|
||||||
|
skipped rather than raising."""
|
||||||
|
from docling_core.types.doc.base import BoundingBox, Size
|
||||||
|
from docling_core.types.doc.document import ImageRef, ProvenanceItem
|
||||||
|
from PIL import Image as PilImageModule
|
||||||
|
|
||||||
|
docling_doc = DoclingDocument(name="orphan-page-box")
|
||||||
|
docling_doc.add_page(
|
||||||
|
page_no=1,
|
||||||
|
size=Size(width=612.0, height=792.0),
|
||||||
|
image=ImageRef.from_pil(
|
||||||
|
PilImageModule.new("RGB", (612, 792), color="white"), dpi=72
|
||||||
|
),
|
||||||
|
)
|
||||||
|
docling_doc.add_text(
|
||||||
|
label=DocItemLabel.PARAGRAPH,
|
||||||
|
text="Content attributed to a page with no raster.",
|
||||||
|
prov=ProvenanceItem(
|
||||||
|
page_no=3,
|
||||||
|
bbox=BoundingBox(l=50, t=700, r=550, b=650),
|
||||||
|
charspan=(0, 20),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
chunks = [
|
||||||
|
Chunk(
|
||||||
|
content="Content attributed to a page with no raster.",
|
||||||
|
metadata={
|
||||||
|
"doc_item_refs": ["#/texts/0"],
|
||||||
|
"page_numbers": [3],
|
||||||
|
"labels": ["paragraph"],
|
||||||
|
},
|
||||||
|
order=0,
|
||||||
|
embedding=[0.1] * 2560,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.import_document(
|
||||||
|
docling_doc, chunks, uri="test://orphan-page"
|
||||||
|
)
|
||||||
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
|
||||||
|
assert await client.visualize_chunk(stored[0]) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_without_refs_falls_back_to_chunk_metadata(temp_db_path):
|
||||||
|
"""A chunk carrying no doc_item_refs has nothing to expand from."""
|
||||||
|
docling_doc = _bbox_doc(with_page_image=True)
|
||||||
|
chunks = [
|
||||||
|
Chunk(
|
||||||
|
content="Content on page 1.",
|
||||||
|
metadata={"page_numbers": [1], "labels": ["paragraph"]},
|
||||||
|
order=0,
|
||||||
|
embedding=[0.1] * 2560,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.import_document(docling_doc, chunks, uri="test://no-refs")
|
||||||
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
|
||||||
|
assert await client.visualize_chunk(stored[0]) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_falls_back_when_expansion_drops_refs(temp_db_path):
|
||||||
|
"""If expansion returns results carrying no refs, the original search
|
||||||
|
results' refs are used instead."""
|
||||||
|
from haiku.rag.client import search as search_module
|
||||||
|
|
||||||
|
docling_doc = _bbox_doc(with_page_image=True)
|
||||||
|
chunks = [
|
||||||
|
Chunk(
|
||||||
|
content="Content on page 1.",
|
||||||
|
metadata={
|
||||||
|
"doc_item_refs": ["#/texts/0"],
|
||||||
|
"page_numbers": [1],
|
||||||
|
"labels": ["paragraph"],
|
||||||
|
},
|
||||||
|
order=0,
|
||||||
|
embedding=[0.1] * 2560,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.import_document(docling_doc, chunks, uri="test://drops-refs")
|
||||||
|
stored = await client.chunk_repository.get_by_document_id(doc.id)
|
||||||
|
|
||||||
|
async def expansion_without_refs(_client, results):
|
||||||
|
return [r.model_copy(update={"doc_item_refs": []}) for r in results]
|
||||||
|
|
||||||
|
with patch.object(search_module, "expand_context", expansion_without_refs):
|
||||||
|
images = await client.visualize_chunk(stored[0])
|
||||||
|
|
||||||
|
assert len(images) == 1
|
||||||
|
|
|
||||||
|
|
@ -480,20 +480,26 @@ def _picture_only_result(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_dedup_keeps_higher_scoring_picture_chunk():
|
@pytest.mark.parametrize(
|
||||||
|
"first_score,second_score",
|
||||||
|
# Whichever duplicate scores higher wins, regardless of arrival order.
|
||||||
|
[(0.7, 0.9), (0.9, 0.7)],
|
||||||
|
ids=["later_wins", "earlier_wins"],
|
||||||
|
)
|
||||||
|
def test_dedup_keeps_higher_scoring_picture_chunk(first_score, second_score):
|
||||||
"""Two results referencing the same single picture self_ref collapse
|
"""Two results referencing the same single picture self_ref collapse
|
||||||
to the one with the higher score."""
|
to the one with the higher score."""
|
||||||
from haiku.rag.client.search import _dedup_picture_chunks
|
from haiku.rag.client.search import _dedup_picture_chunks
|
||||||
|
|
||||||
text_chunk = _picture_only_result("#/pictures/0", score=0.7)
|
text_chunk = _picture_only_result("#/pictures/0", score=first_score)
|
||||||
pic_chunk = _picture_only_result("#/pictures/0", score=0.9)
|
pic_chunk = _picture_only_result("#/pictures/0", score=second_score)
|
||||||
other = _picture_only_result("#/pictures/1", score=0.6)
|
other = _picture_only_result("#/pictures/1", score=0.6)
|
||||||
|
|
||||||
deduped = _dedup_picture_chunks([text_chunk, pic_chunk, other])
|
deduped = _dedup_picture_chunks([text_chunk, pic_chunk, other])
|
||||||
|
|
||||||
assert len(deduped) == 2
|
assert len(deduped) == 2
|
||||||
chosen = next(r for r in deduped if r.doc_item_refs == ["#/pictures/0"])
|
chosen = next(r for r in deduped if r.doc_item_refs == ["#/pictures/0"])
|
||||||
assert chosen.score == 0.9
|
assert chosen.score == max(first_score, second_score)
|
||||||
assert any(r.doc_item_refs == ["#/pictures/1"] for r in deduped)
|
assert any(r.doc_item_refs == ["#/pictures/1"] for r in deduped)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -525,3 +531,56 @@ def test_dedup_does_not_collapse_across_documents():
|
||||||
|
|
||||||
deduped = _dedup_picture_chunks([a, b])
|
deduped = _dedup_picture_chunks([a, b])
|
||||||
assert len(deduped) == 2
|
assert len(deduped) == 2
|
||||||
|
|
||||||
|
|
||||||
|
# visualize_chunk short-circuits
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_expand_context_passes_through_results_without_document(temp_db_path):
|
||||||
|
"""A result with no document_id can't be expanded; it is returned as-is."""
|
||||||
|
from haiku.rag.client.search import expand_context
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
|
orphan = SearchResult(content="loose text", score=0.5, chunk_id="c1")
|
||||||
|
assert await expand_context(rag, [orphan]) == [orphan]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_returns_empty_for_no_chunks(temp_db_path):
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
|
assert await rag.visualize_chunk([]) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_returns_empty_without_document_id(temp_db_path):
|
||||||
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
|
assert await rag.visualize_chunk(Chunk(content="x", metadata={})) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_returns_empty_when_document_missing(temp_db_path):
|
||||||
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
|
chunk = Chunk(content="x", document_id="does-not-exist", metadata={})
|
||||||
|
assert await rag.visualize_chunk(chunk) == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_visualize_chunk_returns_empty_when_docling_blob_absent(temp_db_path):
|
||||||
|
"""A markdown-ingested document has no docling structure to resolve boxes in."""
|
||||||
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
from haiku.rag.store.models.document import Document
|
||||||
|
from haiku.rag.store.repositories.document import DocumentRepository
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||||
|
doc = await DocumentRepository(rag.store).create(
|
||||||
|
Document(content="plain body", uri="test://plain")
|
||||||
|
)
|
||||||
|
assert doc.id is not None
|
||||||
|
chunk = Chunk(content="plain body", document_id=doc.id, metadata={})
|
||||||
|
|
||||||
|
assert await rag.visualize_chunk(chunk) == []
|
||||||
|
|
|
||||||
|
|
@ -139,27 +139,38 @@ Emoji test: 🚀 ✅ 📝"""
|
||||||
assert "🚀" in result_markdown
|
assert "🚀" in result_markdown
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_ollama():
|
@pytest.mark.parametrize(
|
||||||
"""Test get_model returns OpenAIChatModel for Ollama."""
|
"kwargs",
|
||||||
model_config = ModelConfig(provider="ollama", name="llama3")
|
[
|
||||||
result = get_model(model_config)
|
{"provider": "ollama", "name": "llama3"},
|
||||||
assert isinstance(result, OpenAIChatModel)
|
{"provider": "ollama", "name": "gpt-oss", "enable_thinking": False},
|
||||||
|
{"provider": "ollama", "name": "gpt-oss", "enable_thinking": True},
|
||||||
|
{"provider": "ollama", "name": "llama3", "temperature": 0.5, "max_tokens": 100},
|
||||||
def test_get_model_ollama_without_thinking():
|
{"provider": "openai", "name": "gpt-4o"},
|
||||||
"""Test get_model configures thinking for gpt-oss on Ollama."""
|
{"provider": "openai", "name": "o1", "enable_thinking": True},
|
||||||
model_config = ModelConfig(provider="ollama", name="gpt-oss", enable_thinking=False)
|
{"provider": "openai", "name": "o1", "enable_thinking": False},
|
||||||
result = get_model(model_config)
|
{
|
||||||
assert isinstance(result, OpenAIChatModel)
|
"provider": "openai",
|
||||||
|
"name": "gpt-4o",
|
||||||
|
"enable_thinking": False,
|
||||||
def test_get_model_ollama_with_settings():
|
"temperature": 0.7,
|
||||||
"""Test get_model applies temperature and max_tokens for Ollama."""
|
"max_tokens": 500,
|
||||||
model_config = ModelConfig(
|
},
|
||||||
provider="ollama", name="llama3", temperature=0.5, max_tokens=100
|
],
|
||||||
)
|
ids=[
|
||||||
result = get_model(model_config)
|
"ollama",
|
||||||
assert isinstance(result, OpenAIChatModel)
|
"ollama_thinking_off",
|
||||||
|
"ollama_thinking_on",
|
||||||
|
"ollama_with_settings",
|
||||||
|
"openai",
|
||||||
|
"openai_reasoning_thinking_on",
|
||||||
|
"openai_reasoning_thinking_off",
|
||||||
|
"openai_all_settings",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_get_model_returns_openai_chat_model(kwargs):
|
||||||
|
"""Every ollama and openai configuration resolves to an OpenAIChatModel."""
|
||||||
|
assert isinstance(get_model(ModelConfig(**kwargs)), OpenAIChatModel)
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_ollama_appends_v1_to_per_model_base_url():
|
def test_get_model_ollama_appends_v1_to_per_model_base_url():
|
||||||
|
|
@ -183,20 +194,6 @@ def test_get_model_ollama_does_not_double_append_v1():
|
||||||
assert not url.endswith("/v1/v1")
|
assert not url.endswith("/v1/v1")
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_openai():
|
|
||||||
"""Test get_model returns OpenAIChatModel for OpenAI."""
|
|
||||||
model_config = ModelConfig(provider="openai", name="gpt-4o")
|
|
||||||
result = get_model(model_config)
|
|
||||||
assert isinstance(result, OpenAIChatModel)
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_openai_with_thinking():
|
|
||||||
"""Test get_model configures thinking for OpenAI reasoning models."""
|
|
||||||
model_config = ModelConfig(provider="openai", name="o1", enable_thinking=True)
|
|
||||||
result = get_model(model_config)
|
|
||||||
assert isinstance(result, OpenAIChatModel)
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_openai_non_reasoning_model_ignores_thinking():
|
def test_get_model_openai_non_reasoning_model_ignores_thinking():
|
||||||
"""Test that non-reasoning OpenAI models don't get reasoning_effort setting."""
|
"""Test that non-reasoning OpenAI models don't get reasoning_effort setting."""
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
|
|
@ -299,14 +296,15 @@ def test_get_model_anthropic():
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(not HAS_ANTHROPIC, reason="Anthropic not installed")
|
@pytest.mark.skipif(not HAS_ANTHROPIC, reason="Anthropic not installed")
|
||||||
def test_get_model_anthropic_with_thinking():
|
@pytest.mark.parametrize("enable_thinking", [True, False])
|
||||||
|
def test_get_model_anthropic_with_thinking(enable_thinking):
|
||||||
"""Test get_model configures thinking for Anthropic."""
|
"""Test get_model configures thinking for Anthropic."""
|
||||||
from pydantic_ai.models.anthropic import AnthropicModel
|
from pydantic_ai.models.anthropic import AnthropicModel
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="anthropic",
|
provider="anthropic",
|
||||||
name="claude-3-5-sonnet-20241022",
|
name="claude-3-5-sonnet-20241022",
|
||||||
enable_thinking=True,
|
enable_thinking=enable_thinking,
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, AnthropicModel)
|
assert isinstance(result, AnthropicModel)
|
||||||
|
|
@ -345,12 +343,15 @@ def test_get_model_groq():
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(not HAS_GROQ, reason="Groq not installed")
|
@pytest.mark.skipif(not HAS_GROQ, reason="Groq not installed")
|
||||||
def test_get_model_groq_with_thinking():
|
@pytest.mark.parametrize("enable_thinking", [True, False])
|
||||||
|
def test_get_model_groq_with_thinking(enable_thinking):
|
||||||
"""Test get_model configures thinking format for Groq."""
|
"""Test get_model configures thinking format for Groq."""
|
||||||
from pydantic_ai.models.groq import GroqModel
|
from pydantic_ai.models.groq import GroqModel
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="groq", name="llama-3.3-70b-versatile", enable_thinking=False
|
provider="groq",
|
||||||
|
name="llama-3.3-70b-versatile",
|
||||||
|
enable_thinking=enable_thinking,
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, GroqModel)
|
assert isinstance(result, GroqModel)
|
||||||
|
|
@ -369,14 +370,26 @@ def test_get_model_bedrock():
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(not HAS_BEDROCK, reason="Bedrock not installed")
|
@pytest.mark.skipif(not HAS_BEDROCK, reason="Bedrock not installed")
|
||||||
def test_get_model_bedrock_with_thinking():
|
@pytest.mark.parametrize("enable_thinking", [True, False])
|
||||||
"""Test get_model configures thinking for Bedrock Claude models."""
|
@pytest.mark.parametrize(
|
||||||
|
"name",
|
||||||
|
[
|
||||||
|
"anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||||
|
"openai.o3-mini-v1:0",
|
||||||
|
"qwen.qwen3-32b-v1:0",
|
||||||
|
# A family with no reasoning mapping leaves the request fields untouched.
|
||||||
|
"meta.llama3-70b-instruct-v1:0",
|
||||||
|
],
|
||||||
|
ids=["claude", "o_series", "qwen", "unmapped_family"],
|
||||||
|
)
|
||||||
|
def test_get_model_bedrock_with_thinking(name, enable_thinking):
|
||||||
|
"""Each Bedrock model family maps thinking onto its own request field."""
|
||||||
from pydantic_ai.models.bedrock import BedrockConverseModel
|
from pydantic_ai.models.bedrock import BedrockConverseModel
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="bedrock",
|
provider="bedrock",
|
||||||
name="anthropic.claude-3-5-sonnet-20241022-v2:0",
|
name=name,
|
||||||
enable_thinking=True,
|
enable_thinking=enable_thinking,
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, BedrockConverseModel)
|
assert isinstance(result, BedrockConverseModel)
|
||||||
|
|
@ -390,19 +403,6 @@ def test_get_model_unknown_provider():
|
||||||
assert result == "mistral:mistral-large-latest"
|
assert result == "mistral:mistral-large-latest"
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_with_all_settings():
|
|
||||||
"""Test get_model applies all settings together."""
|
|
||||||
model_config = ModelConfig(
|
|
||||||
provider="openai",
|
|
||||||
name="gpt-4o",
|
|
||||||
enable_thinking=False,
|
|
||||||
temperature=0.7,
|
|
||||||
max_tokens=500,
|
|
||||||
)
|
|
||||||
result = get_model(model_config)
|
|
||||||
assert isinstance(result, OpenAIChatModel)
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_package_versions():
|
def test_get_package_versions():
|
||||||
"""Test get_package_versions returns expected keys."""
|
"""Test get_package_versions returns expected keys."""
|
||||||
from haiku.rag.utils import get_package_versions
|
from haiku.rag.utils import get_package_versions
|
||||||
|
|
@ -526,20 +526,7 @@ def test_format_citations_multiple_pages():
|
||||||
result = format_citations([citation])
|
result = format_citations([citation])
|
||||||
assert "[1] test://doc" in result
|
assert "[1] test://doc" in result
|
||||||
assert "pp. 1-3" in result
|
assert "pp. 1-3" in result
|
||||||
|
# No title: the URI stands in, and the document id never leaks.
|
||||||
|
|
||||||
def test_format_citations_no_title():
|
|
||||||
from haiku.rag.store.models.citation import Citation
|
|
||||||
from haiku.rag.utils import format_citations
|
|
||||||
|
|
||||||
citation = Citation(
|
|
||||||
document_id="doc1",
|
|
||||||
chunk_id="chunk1",
|
|
||||||
document_uri="test://doc",
|
|
||||||
content="Content",
|
|
||||||
)
|
|
||||||
result = format_citations([citation])
|
|
||||||
assert "[1] test://doc" in result
|
|
||||||
assert "doc1" not in result
|
assert "doc1" not in result
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -758,3 +745,85 @@ def test_parse_model_option():
|
||||||
for bad in ["just-a-name", ":model", "provider:"]:
|
for bad in ["just-a-name", ":model", "provider:"]:
|
||||||
with pytest.raises(ValueError, match="Invalid model format"):
|
with pytest.raises(ValueError, match="Invalid model format"):
|
||||||
parse_model_option(bad)
|
parse_model_option(bad)
|
||||||
|
|
||||||
|
|
||||||
|
def test_cosine_similarity_identical_vectors():
|
||||||
|
from haiku.rag.utils import cosine_similarity
|
||||||
|
|
||||||
|
assert cosine_similarity([1.0, 0.0], [1.0, 0.0]) == pytest.approx(1.0)
|
||||||
|
assert cosine_similarity([1.0, 0.0], [0.0, 1.0]) == pytest.approx(0.0)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_format_citations_rich_separates_multiple_citations():
|
||||||
|
from haiku.rag.store.models.citation import Citation
|
||||||
|
from haiku.rag.utils import format_citations_rich
|
||||||
|
|
||||||
|
citations = [
|
||||||
|
Citation(
|
||||||
|
document_id=f"doc{i}",
|
||||||
|
chunk_id=f"chunk{i}",
|
||||||
|
document_uri=f"test://doc{i}",
|
||||||
|
document_title=f"Doc {i}",
|
||||||
|
content=f"Body {i}",
|
||||||
|
)
|
||||||
|
for i in (1, 2)
|
||||||
|
]
|
||||||
|
|
||||||
|
output = _render_rich(await format_citations_rich(citations))
|
||||||
|
|
||||||
|
assert "[1] Doc 1 (test://doc1)" in output
|
||||||
|
assert "[2] Doc 2 (test://doc2)" in output
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"stored,renders",
|
||||||
|
[
|
||||||
|
(None, False),
|
||||||
|
(b"not a real image", False),
|
||||||
|
("png", True),
|
||||||
|
],
|
||||||
|
ids=["no_bytes", "undecodable_bytes", "valid_png"],
|
||||||
|
)
|
||||||
|
async def test_render_picture_handles_stored_bytes(stored, renders):
|
||||||
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
|
from haiku.rag.utils import _render_picture
|
||||||
|
|
||||||
|
if stored == "png":
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
|
from PIL import Image as PILImage
|
||||||
|
|
||||||
|
buf = BytesIO()
|
||||||
|
PILImage.new("RGB", (4, 4), "red").save(buf, format="PNG")
|
||||||
|
stored = buf.getvalue()
|
||||||
|
|
||||||
|
client = AsyncMock()
|
||||||
|
client.document_item_repository.get_picture_bytes = AsyncMock(return_value=stored)
|
||||||
|
|
||||||
|
result = await _render_picture(client, "doc1", "#/pictures/0")
|
||||||
|
|
||||||
|
assert (result is not None) is renders
|
||||||
|
|
||||||
|
|
||||||
|
async def test_render_picture_without_client_returns_none():
|
||||||
|
from haiku.rag.utils import _render_picture
|
||||||
|
|
||||||
|
assert await _render_picture(None, "doc1", "#/pictures/0") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_package_versions_reports_missing_docling(monkeypatch):
|
||||||
|
from importlib import metadata as importlib_metadata
|
||||||
|
|
||||||
|
from haiku.rag.utils import get_package_versions
|
||||||
|
|
||||||
|
real_version = importlib_metadata.version
|
||||||
|
|
||||||
|
def fake_version(name):
|
||||||
|
if name == "docling":
|
||||||
|
raise importlib_metadata.PackageNotFoundError(name)
|
||||||
|
return real_version(name)
|
||||||
|
|
||||||
|
monkeypatch.setattr(importlib_metadata, "version", fake_version)
|
||||||
|
|
||||||
|
assert get_package_versions()["docling"] == "not installed"
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue