Additional fixes to remove type ignores

This commit is contained in:
Yiorgis Gozadinos 2026-01-19 16:00:18 +02:00
parent 5e55c0df54
commit f05159859c
No known key found for this signature in database
4 changed files with 40 additions and 18 deletions

View file

@ -177,6 +177,10 @@ def get_model(
return OpenAIChatModel(model_name=model, settings=openai_settings)
elif provider == "anthropic":
from anthropic.types.beta import (
BetaThinkingConfigDisabledParam,
BetaThinkingConfigEnabledParam,
)
from pydantic_ai.models.anthropic import AnthropicModel, AnthropicModelSettings
anthropic_settings: Any = None
@ -184,12 +188,19 @@ def get_model(
# Apply thinking control
if model_config.enable_thinking is not None:
if model_config.enable_thinking:
thinking_config: BetaThinkingConfigEnabledParam = {
"type": "enabled",
"budget_tokens": 4096,
}
anthropic_settings = AnthropicModelSettings(
anthropic_thinking={"type": "enabled", "budget_tokens": 4096} # ty: ignore[invalid-argument-type]
anthropic_thinking=thinking_config
)
else:
thinking_disabled: BetaThinkingConfigDisabledParam = {
"type": "disabled"
}
anthropic_settings = AnthropicModelSettings(
anthropic_thinking={"type": "disabled"} # ty: ignore[invalid-argument-type]
anthropic_thinking=thinking_disabled
)
anthropic_settings = apply_common_settings(

View file

@ -81,9 +81,10 @@ def test_chat_agent_has_dynamic_system_prompt():
agent = create_chat_agent(Config)
# The agent should have at least one system prompt function registered
# (the add_background_context function)
assert len(agent._system_prompt_functions) >= 1
system_prompt_functions = getattr(agent, "_system_prompt_functions")
assert len(system_prompt_functions) >= 1
# Verify it's the add_background_context function
func_names = [r.function.__name__ for r in agent._system_prompt_functions] # ty: ignore[unresolved-attribute]
func_names = [r.function.__name__ for r in system_prompt_functions]
assert "add_background_context" in func_names

View file

@ -2,7 +2,7 @@ import logging
import os
import tempfile
from pathlib import Path
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, cast
# Prevent tests from loading user's local haiku.rag.yaml by setting env var
# to a test config file BEFORE any haiku.rag imports.
@ -26,7 +26,7 @@ from datasets import Dataset, load_dataset, load_from_disk # noqa: E402
if TYPE_CHECKING:
from vcr import VCR
pydantic_ai.models.ALLOW_MODEL_REQUESTS = False # ty: ignore[invalid-assignment]
setattr(pydantic_ai.models, "ALLOW_MODEL_REQUESTS", False)
logging.getLogger("vcr.cassette").setLevel(logging.WARNING)
@ -35,8 +35,7 @@ def qa_corpus() -> Dataset:
ds_path = Path(__file__).parent / "data" / "dataset"
ds_path.mkdir(parents=True, exist_ok=True)
try:
ds: Dataset = load_from_disk(ds_path) # ty: ignore[invalid-assignment]
return ds
return cast(Dataset, load_from_disk(ds_path))
except FileNotFoundError:
ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"]
corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories")

View file

@ -1,9 +1,20 @@
from typing import TypedDict
import pytest
from datasets import Dataset
from haiku.rag.client import HaikuRAG, RebuildMode
class ChunkData(TypedDict):
id: str
document_id: str
content: str
content_fts: str
metadata: str
order: int
@pytest.mark.vcr()
async def test_rebuild_full(qa_corpus: Dataset, temp_db_path):
"""Test full rebuild: converts, chunks, and embeds all documents."""
@ -132,15 +143,15 @@ async def test_rebuild_embed_only_with_changed_vector_dim(
chunks_before = await client.chunk_repository.get_by_document_id(doc.id)
assert len(chunks_before) > 0
chunk_data = [
{
"id": c.id,
"document_id": c.document_id,
"content": c.content,
"content_fts": c.content,
"metadata": json.dumps(c.metadata),
"order": c.order,
}
chunk_data: list[ChunkData] = [
ChunkData(
id=c.id or "",
document_id=c.document_id or "",
content=c.content,
content_fts=c.content,
metadata=json.dumps(c.metadata),
order=c.order,
)
for c in chunks_before
]
@ -167,7 +178,7 @@ async def test_rebuild_embed_only_with_changed_vector_dim(
content=c["content"],
content_fts=c["content_fts"],
metadata=c["metadata"],
order=c["order"], # ty: ignore[invalid-argument-type]
order=c["order"],
vector=[0.1] * 4096,
)
for c in chunk_data