diff --git a/src/backend/chat/agents/__init__.py b/src/backend/chat/agents/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/backend/chat/agents/base.py b/src/backend/chat/agents/base.py new file mode 100644 index 0000000..82fa196 --- /dev/null +++ b/src/backend/chat/agents/base.py @@ -0,0 +1,38 @@ +"""Base module for PydanticAI agents.""" + +from django.conf import settings +from django.core.exceptions import ImproperlyConfigured + +from pydantic_ai import Agent +from pydantic_ai.models.openai import OpenAIModel +from pydantic_ai.providers.openai import OpenAIProvider + +from chat.tools import get_pydantic_tools_by_name + + +def _get_pydantic_agent(model_hrid, mcp_servers=None, **kwargs) -> Agent: + """Get the PydanticAI Agent instance with the configured settings.""" + try: + _model = settings.LLM_CONFIGURATIONS[model_hrid] + except KeyError as exc: + raise ImproperlyConfigured(f"LLM model configuration '{model_hrid}' not found.") from exc + + _model_instance = OpenAIModel( + model_name=_model.model_name, + provider=OpenAIProvider( + base_url=_model.provider.base_url, + api_key=_model.provider.api_key, + ) + if _model.provider + else None, + ) + _system_prompt = _model.system_prompt + _tools = [get_pydantic_tools_by_name(tool_name) for tool_name in _model.tools] + + return Agent( + model=_model_instance, + system_prompt=_system_prompt, + mcp_servers=mcp_servers or [], + tools=_tools, + **kwargs, + ) diff --git a/src/backend/chat/agents/conversation.py b/src/backend/chat/agents/conversation.py new file mode 100644 index 0000000..feef8e9 --- /dev/null +++ b/src/backend/chat/agents/conversation.py @@ -0,0 +1,41 @@ +"""Build the main conversation agent.""" + +import logging + +from django.conf import settings +from django.utils import formats, timezone + +from pydantic_ai import Agent + +from core.enums import get_language_name + +from .base import _get_pydantic_agent + +logger = logging.getLogger(__name__) + + +def build_conversation_agent( + mcp_servers, model_hrid=None, language=None, instrument=False +) -> Agent[None, str]: + """Create a Pydantic AI Agent instance with the configured settings.""" + model_hrid = model_hrid or settings.LLM_DEFAULT_MODEL_HRID + + agent = _get_pydantic_agent(model_hrid, mcp_servers, instrument=instrument) + + @agent.system_prompt + def add_the_date() -> str: + """ + Dynamic system prompt function to add the current date. + + Warning: this will always use the date in the server timezone, + not the user's timezone... + """ + _formatted_date = formats.date_format(timezone.now(), "l d/m/Y", use_l10n=False) + return f"Today is {_formatted_date}." + + @agent.system_prompt + def enforce_response_language() -> str: + """Dynamic system prompt function to set the expected language to use.""" + return f"Answer in {get_language_name(language).lower()}." if language else "" + + return agent diff --git a/src/backend/chat/agents/routing.py b/src/backend/chat/agents/routing.py new file mode 100644 index 0000000..2b0527d --- /dev/null +++ b/src/backend/chat/agents/routing.py @@ -0,0 +1,54 @@ +"""Build the routing agent.""" + +import logging + +from django.conf import settings +from django.core.exceptions import ImproperlyConfigured + +from pydantic import BaseModel +from pydantic_ai import Agent, NativeOutput + +from .base import _get_pydantic_agent + +logger = logging.getLogger(__name__) + + +class UserIntent(BaseModel): + """Model to represent the detected user intent.""" + + web_search: bool = False + attachment_summary: bool = False + + +def build_routing_agent(model_hrid=None, instrument=False) -> Agent[None, str] | None: + """ + Create a Pydantic AI routing Agent instance with the configured settings. + + This agent is used to detect the user intent from the user prompt. + + Args: + model_hrid (str | None): The HRID of the routing model to use. + If None, the default routing model from settings will be used. + Returns: + Agent | None: The Pydantic AI Agent instance or None if not configured. + Raises: + ImproperlyConfigured: If the routing model configuration is invalid. + """ + model_hrid = model_hrid or settings.LLM_ROUTING_MODEL_HRID + + try: + agent = _get_pydantic_agent( + model_hrid, + output_type=NativeOutput([UserIntent]), + instrument=instrument, + ) + except ImproperlyConfigured: + logger.info("AI routing model does not exist -> disabled") + return None + + # Simple detection of configuration not set + if not agent.model.model_name: + logger.info("AI routing model configuration not set -> disabled") + return None + + return agent diff --git a/src/backend/chat/clients/pydantic_ai.py b/src/backend/chat/clients/pydantic_ai.py index 0aff09c..39ad5e2 100644 --- a/src/backend/chat/clients/pydantic_ai.py +++ b/src/backend/chat/clients/pydantic_ai.py @@ -16,15 +16,12 @@ from typing import Dict, List, Optional, Tuple from django.conf import settings from django.core.cache import cache -from django.core.exceptions import ImproperlyConfigured -from django.utils import formats, timezone from django.utils.module_loading import import_string from django.utils.translation import gettext_lazy as _ from asgiref.sync import sync_to_async from langfuse import get_client -from pydantic import BaseModel -from pydantic_ai import Agent, NativeOutput +from pydantic_ai import Agent from pydantic_ai.messages import ( BinaryContent, FunctionToolCallEvent, @@ -45,13 +42,12 @@ from pydantic_ai.messages import ( ToolReturnPart, UserPromptPart, ) -from pydantic_ai.models.openai import OpenAIModel -from pydantic_ai.providers.openai import OpenAIProvider -from core.enums import get_language_name from core.feature_flags.helpers import is_feature_enabled from chat.agent_rag.document_search.albert_api import AlbertRagDocumentSearch +from chat.agents.conversation import build_conversation_agent +from chat.agents.routing import UserIntent, build_routing_agent from chat.ai_sdk_types import ( LanguageModelV1Source, SourceUIPart, @@ -64,109 +60,12 @@ from chat.clients.pydantic_ui_message_converter import ( ui_message_to_user_content, ) from chat.mcp_servers import get_mcp_servers -from chat.tools import get_pydantic_tools_by_name from chat.vercel_ai_sdk.core import events_v4, events_v5 from chat.vercel_ai_sdk.encoder import EventEncoder logger = logging.getLogger(__name__) -def _get_pydantic_agent(model_hrid, mcp_servers=None, **kwargs) -> Agent: - """Get the PydanticAI Agent instance with the configured settings.""" - try: - _model = settings.LLM_CONFIGURATIONS[model_hrid] - except KeyError as exc: - raise ImproperlyConfigured(f"LLM model configuration '{model_hrid}' not found.") from exc - - _model_instance = OpenAIModel( - model_name=_model.model_name, - provider=OpenAIProvider( - base_url=_model.provider.base_url, - api_key=_model.provider.api_key, - ) - if _model.provider - else None, - ) - _system_prompt = _model.system_prompt - _tools = [get_pydantic_tools_by_name(tool_name) for tool_name in _model.tools] - - return Agent( - model=_model_instance, - system_prompt=_system_prompt, - mcp_servers=mcp_servers or [], - tools=_tools, - **kwargs, - ) - - -def _build_pydantic_agent( - mcp_servers, model_hrid=None, language=None, instrument=False -) -> Agent[None, str]: - """Create a Pydantic AI Agent instance with the configured settings.""" - model_hrid = model_hrid or settings.LLM_DEFAULT_MODEL_HRID - - agent = _get_pydantic_agent(model_hrid, mcp_servers, instrument=instrument) - - @agent.system_prompt - def add_the_date() -> str: - """ - Dynamic system prompt function to add the current date. - - Warning: this will always use the date in the server timezone, - not the user's timezone... - """ - _formatted_date = formats.date_format(timezone.now(), "l d/m/Y", use_l10n=False) - return f"Today is {_formatted_date}." - - @agent.system_prompt - def enforce_response_language() -> str: - """Dynamic system prompt function to set the expected language to use.""" - return f"Answer in {get_language_name(language).lower()}." if language else "" - - return agent - - -def _build_routing_agent(model_hrid=None, instrument=False) -> Agent[None, str] | None: - """ - Create a Pydantic AI routing Agent instance with the configured settings. - - This agent is used to detect the user intent from the user prompt. - - Args: - model_hrid (str | None): The HRID of the routing model to use. - If None, the default routing model from settings will be used. - Returns: - Agent | None: The Pydantic AI Agent instance or None if not configured. - Raises: - ImproperlyConfigured: If the routing model configuration is invalid. - """ - model_hrid = model_hrid or settings.LLM_ROUTING_MODEL_HRID - - try: - agent = _get_pydantic_agent( - model_hrid, - output_type=NativeOutput([UserIntent]), - instrument=instrument, - ) - except ImproperlyConfigured: - logger.info("AI routing model does not exist -> disabled") - return None - - # Simple detection of configuration not set - if not agent.model.model_name: - logger.info("AI routing model configuration not set -> disabled") - return None - - return agent - - -class UserIntent(BaseModel): - """Model to represent the detected user intent.""" - - web_search: bool = False - attachment_summary: bool = False - - class AIAgentService: # pylint: disable=too-many-instance-attributes """Service class for AI-related operations (Pydantic-AI edition).""" @@ -291,7 +190,7 @@ class AIAgentService: # pylint: disable=too-many-instance-attributes ) return UserIntent() - agent = _build_routing_agent(instrument=self._store_analytics) + agent = build_routing_agent(instrument=self._store_analytics) if not agent: return UserIntent() @@ -560,12 +459,15 @@ class AIAgentService: # pylint: disable=too-many-instance-attributes # MCP servers (if any) can be initialized here mcp_servers = [await stack.enter_async_context(mcp) for mcp in get_mcp_servers()] - async with _build_pydantic_agent( + # Build the agent + conversation_agent = build_conversation_agent( mcp_servers, model_hrid=self.model_hrid, language=self.language, instrument=self._store_analytics, - ).iter( + ) + + async with conversation_agent.iter( [user_prompt] + input_images, message_history=history, ) as run: diff --git a/src/backend/chat/tests/agents/__init__.py b/src/backend/chat/tests/agents/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/backend/chat/tests/clients/pydantic_ai/test__build_pydantic_agent.py b/src/backend/chat/tests/agents/test_build_conversation_agent.py similarity index 87% rename from src/backend/chat/tests/clients/pydantic_ai/test__build_pydantic_agent.py rename to src/backend/chat/tests/agents/test_build_conversation_agent.py index d71281f..f263457 100644 --- a/src/backend/chat/tests/clients/pydantic_ai/test__build_pydantic_agent.py +++ b/src/backend/chat/tests/agents/test_build_conversation_agent.py @@ -1,4 +1,4 @@ -"""Test cases for the _build_pydantic_agent function in the chat.clients.pydantic_ai module.""" +"""Test cases for the build_conversation_agent function in the chat.clients.pydantic_ai module.""" # pylint:disable=protected-access @@ -7,7 +7,7 @@ from freezegun import freeze_time from pydantic_ai import Agent from pydantic_ai.models.openai import OpenAIModel -from chat.clients.pydantic_ai import _build_pydantic_agent +from chat.agents.conversation import build_conversation_agent @pytest.fixture(autouse=True) @@ -27,7 +27,7 @@ def base_settings(settings): def test_build_pydantic_agent_success_no_tools(): """Test successful agent creation without tools.""" - agent = _build_pydantic_agent([]) + agent = build_conversation_agent([]) assert isinstance(agent, Agent) assert agent._system_prompts == ("You are a helpful assistant",) @@ -43,7 +43,7 @@ def test_build_pydantic_agent_with_tools(settings): """Test successful agent creation with tools.""" settings.AI_AGENT_TOOLS = ["get_current_weather"] - agent = _build_pydantic_agent([]) + agent = build_conversation_agent([]) assert isinstance(agent, Agent) assert agent._system_prompts == ("You are a helpful assistant",) @@ -61,7 +61,7 @@ def test_add_dynamic_system_prompt(): Ensure add_the_date and enforce_response_language system prompt are registered and returns proper values. """ - agent = _build_pydantic_agent([]) + agent = build_conversation_agent([]) assert len(agent._system_prompt_functions) == 2 @@ -71,5 +71,5 @@ def test_add_dynamic_system_prompt(): assert agent._system_prompt_functions[1].function.__name__ == "enforce_response_language" assert agent._system_prompt_functions[1].function() == "" - agent = _build_pydantic_agent([], language="fr-fr") + agent = build_conversation_agent([], language="fr-fr") assert agent._system_prompt_functions[1].function() == "Answer in french." diff --git a/src/backend/chat/tools/__init__.py b/src/backend/chat/tools/__init__.py index ccffb65..5f3441d 100644 --- a/src/backend/chat/tools/__init__.py +++ b/src/backend/chat/tools/__init__.py @@ -7,7 +7,7 @@ from .web_search_tavily import tavily_web_search def get_pydantic_tools_by_name(name: str) -> Tool: - """Get a Pydantic AI agent by its name.""" + """Get a tool by its name.""" tool_dict = { "get_current_weather": Tool(get_current_weather, takes_ctx=False), "tavily_web_search": Tool(tavily_web_search, takes_ctx=False),