diff --git a/src/haiku/rag/logging.py b/src/haiku/rag/logging.py index 3627a5e9..34aeaf5a 100644 --- a/src/haiku/rag/logging.py +++ b/src/haiku/rag/logging.py @@ -6,6 +6,7 @@ from rich.logging import RichHandler logging.basicConfig(level=logging.DEBUG) logging.getLogger("httpx").setLevel(logging.WARNING) logging.getLogger("httpcore").setLevel(logging.WARNING) +logging.getLogger("docling").setLevel(logging.WARNING) def get_logger() -> logging.Logger: diff --git a/src/haiku/rag/qa/agent.py b/src/haiku/rag/qa/agent.py index b76112bc..711106ba 100644 --- a/src/haiku/rag/qa/agent.py +++ b/src/haiku/rag/qa/agent.py @@ -1,6 +1,6 @@ from pydantic import BaseModel, Field from pydantic_ai import Agent, RunContext -from pydantic_ai.models.openai import OpenAIModel +from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.ollama import OllamaProvider from haiku.rag.client import HaikuRAG @@ -61,7 +61,7 @@ class QuestionAnswerAgent: def _get_model(self, provider: str, model: str): """Get the appropriate model object for the provider.""" if provider == "ollama": - return OpenAIModel( + return OpenAIChatModel( model_name=model, provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"), ) diff --git a/tests/generate_benchmark_db.py b/tests/generate_benchmark_db.py index 7318a27b..47baabd9 100644 --- a/tests/generate_benchmark_db.py +++ b/tests/generate_benchmark_db.py @@ -6,6 +6,7 @@ from llm_judge import LLMJudge from rich.console import Console from rich.progress import Progress +from haiku.rag import logging # noqa from haiku.rag.client import HaikuRAG from haiku.rag.qa import get_qa_agent