diff --git a/haiku_rag_slim/haiku/rag/utils.py b/haiku_rag_slim/haiku/rag/utils.py index f51770f7..b038dde0 100644 --- a/haiku_rag_slim/haiku/rag/utils.py +++ b/haiku_rag_slim/haiku/rag/utils.py @@ -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( diff --git a/tests/agents/chat/test_chat_agent.py b/tests/agents/chat/test_chat_agent.py index 0dba3637..392ced22 100644 --- a/tests/agents/chat/test_chat_agent.py +++ b/tests/agents/chat/test_chat_agent.py @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py index c92956b4..5362a6cf 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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") diff --git a/tests/test_rebuild.py b/tests/test_rebuild.py index 6dfa0b90..43e56304 100644 --- a/tests/test_rebuild.py +++ b/tests/test_rebuild.py @@ -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