diff --git a/litellm/integrations/advisor_interception/__init__.py b/litellm/integrations/advisor_interception/__init__.py new file mode 100644 index 00000000000..43b23d18e00 --- /dev/null +++ b/litellm/integrations/advisor_interception/__init__.py @@ -0,0 +1,22 @@ +""" +Advisor interception module. + +Provides completion/chat-completions advisor orchestration for providers that +do not natively support Anthropic's ``advisor_20260301`` tool type. +""" + +from litellm.integrations.advisor_interception.handler import ( + AdvisorInterceptionLogger, +) +from litellm.integrations.advisor_interception.tools import ( + get_litellm_advisor_tool, + get_litellm_advisor_tool_openai, + is_advisor_tool, +) + +__all__ = [ + "AdvisorInterceptionLogger", + "get_litellm_advisor_tool", + "get_litellm_advisor_tool_openai", + "is_advisor_tool", +] diff --git a/litellm/integrations/advisor_interception/handler.py b/litellm/integrations/advisor_interception/handler.py new file mode 100644 index 00000000000..bcf9f08549d --- /dev/null +++ b/litellm/integrations/advisor_interception/handler.py @@ -0,0 +1,626 @@ +""" +Advisor interception for chat completions. + +This mirrors the websearch interception pattern: +- Convert advisor tools to provider-compatible tool schema pre-request. +- Detect advisor tool calls in chat completion responses. +- Execute advisor sub-calls and continue the agentic loop server-side. +""" + +import json +from typing import Any, Dict, List, Optional, Tuple, Union + +import litellm +from litellm._logging import verbose_logger +from litellm.constants import ADVISOR_MAX_USES, ADVISOR_NATIVE_PROVIDERS +from litellm.integrations.advisor_interception.tools import ( + LITELLM_ADVISOR_TOOL_NAME, + get_litellm_advisor_tool, + get_litellm_advisor_tool_openai, + is_advisor_tool, + is_advisor_tool_chat_completion, +) +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.utils import CallTypes, LlmProviders + + +class AdvisorInterceptionLogger(CustomLogger): + """ + Intercept advisor tool calls in chat completions and orchestrate sub-calls. + """ + + def __init__( + self, + enabled_providers: Optional[List[Union[LlmProviders, str]]] = None, + default_advisor_model: str = "claude-opus-4-6", + ): + super().__init__() + if enabled_providers is None: + self.enabled_providers = None + else: + self.enabled_providers = [ + p.value if isinstance(p, LlmProviders) else p for p in enabled_providers + ] + self.default_advisor_model = default_advisor_model + self._advisor_config_by_call_id: Dict[str, Dict[str, Any]] = {} + self._skip_post_hook_call_ids: set[str] = set() + + async def async_pre_request_hook( + self, model: str, messages: List[Dict], kwargs: Dict + ) -> Optional[Dict]: + """ + Convert advisor tools into provider-compatible form before request. + """ + custom_llm_provider = kwargs.get("litellm_params", {}).get( + "custom_llm_provider", "" + ) + return self._convert_tools_for_provider( + kwargs=kwargs, custom_llm_provider=custom_llm_provider + ) + + async def async_pre_call_deployment_hook( + self, kwargs: Dict[str, Any], call_type: Optional[Any] + ) -> Optional[dict]: + """ + Pre-call hook used by completion/chat-completions paths. + """ + if kwargs.pop("_advisor_interception_skip_post_hook", False): + call_id = kwargs.get("litellm_call_id") + if isinstance(call_id, str): + self._skip_post_hook_call_ids.add(call_id) + + custom_llm_provider = kwargs.get("custom_llm_provider", "") or kwargs.get( + "litellm_params", {} + ).get("custom_llm_provider", "") + if not custom_llm_provider: + try: + _, custom_llm_provider, _, _ = litellm.get_llm_provider( + model=kwargs.get("model", "") + ) + except Exception: + custom_llm_provider = "" + return self._convert_tools_for_provider( + kwargs=kwargs, custom_llm_provider=custom_llm_provider + ) + + async def async_post_call_success_deployment_hook( + self, + request_data: Dict[str, Any], + response: Any, + call_type: Optional[CallTypes], + ) -> Optional[Any]: + """ + Fallback advisor interception for providers that do not call + async_should_run_chat_completion_agentic_loop internally. + """ + if call_type not in {CallTypes.completion, CallTypes.acompletion}: + return None + + call_id = request_data.get("litellm_call_id") + if isinstance(call_id, str) and call_id in self._skip_post_hook_call_ids: + self._skip_post_hook_call_ids.remove(call_id) + return None + + model = request_data.get("model") + messages = request_data.get("messages") + if not isinstance(model, str) or not isinstance(messages, list): + return None + + custom_llm_provider = request_data.get("custom_llm_provider", "") or request_data.get( + "litellm_params", {} + ).get("custom_llm_provider", "") + if not custom_llm_provider: + try: + _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model) + except Exception: + custom_llm_provider = "" + + tools = request_data.get("tools") + stream = bool(request_data.get("stream", False)) + + should_run, tools_dict = await self.async_should_run_chat_completion_agentic_loop( + response=response, + model=model, + messages=messages, + tools=tools if isinstance(tools, list) else None, + stream=stream, + custom_llm_provider=custom_llm_provider, + kwargs=request_data, + ) + if not should_run: + return None + + optional_params = self._build_optional_params_from_request_data(request_data) + return await self.async_run_chat_completion_agentic_loop( + tools=tools_dict, + model=model, + messages=messages, + response=response, + optional_params=optional_params, + logging_obj=request_data.get("litellm_logging_obj"), + stream=stream, + kwargs=request_data, + ) + + async def async_should_run_chat_completion_agentic_loop( + self, + response: Any, + model: str, + messages: List[Dict], + tools: Optional[List[Dict]], + stream: bool, + custom_llm_provider: str, + kwargs: Dict, + ) -> Tuple[bool, Dict]: + """ + Determine whether advisor agentic loop should run for chat completions. + """ + if ( + self.enabled_providers is not None + and custom_llm_provider not in self.enabled_providers + ): + return False, {} + if custom_llm_provider in ADVISOR_NATIVE_PROVIDERS: + return False, {} + + call_id = kwargs.get("litellm_call_id") + has_advisor_config = isinstance(call_id, str) and ( + call_id in self._advisor_config_by_call_id + ) + has_advisor_tool = has_advisor_config or ( + bool(tools) and any(is_advisor_tool_chat_completion(t) for t in tools) + ) + if not has_advisor_tool: + return False, {} + + advisor_calls, raw_tool_calls = self._extract_advisor_tool_calls(response) + if not advisor_calls: + return False, {} + + # If there are mixed tool calls, do not hijack the request. + if len(advisor_calls) != len(raw_tool_calls): + verbose_logger.debug( + "AdvisorInterception: Mixed tool calls detected, skipping advisor interception" + ) + return False, {} + + advisor_config = {} + if isinstance(call_id, str): + advisor_config = self._advisor_config_by_call_id.get(call_id, {}) + return True, { + "advisor_calls": advisor_calls, + "raw_tool_calls": raw_tool_calls, + "advisor_config": advisor_config, + "response_format": "openai", + } + + async def async_run_chat_completion_agentic_loop( + self, + tools: Dict, + model: str, + messages: List[Dict], + response: Any, + optional_params: Dict, + logging_obj: Any, + stream: bool, + kwargs: Dict, + ) -> Any: + """ + Execute advisor sub-calls and continue chat-completion loop. + """ + advisor_config = tools.get("advisor_config", {}) or {} + max_uses = int(advisor_config.get("max_uses", ADVISOR_MAX_USES)) + advisor_model = advisor_config.get("advisor_model") or self.default_advisor_model + advisor_api_key = advisor_config.get("api_key") + advisor_api_base = advisor_config.get("api_base") + call_id = kwargs.get("litellm_call_id") + + current_messages: List[Dict] = list(messages) + current_response = response + advisor_uses = 0 + + try: + while True: + advisor_calls, raw_tool_calls = self._extract_advisor_tool_calls( + current_response + ) + if not advisor_calls: + return current_response + + assistant_content = self._extract_message_content(current_response) + assistant_message: Dict[str, Any] = { + "role": "assistant", + "tool_calls": raw_tool_calls, + } + if assistant_content is not None: + assistant_message["content"] = assistant_content + + tool_messages: List[Dict[str, Any]] = [] + for advisor_call in advisor_calls: + advisor_uses += 1 + if advisor_uses > max_uses: + raise ValueError( + "Advisor orchestration loop exceeded max_uses={}. " + "Increase max_uses in advisor tool config.".format(max_uses) + ) + + question = advisor_call.get( + "question", "Please provide guidance on the current task." + ) + advisor_messages = self._build_advisor_context( + messages=current_messages, + assistant_content=assistant_content, + question=question, + ) + advisor_response = await litellm.acompletion( + model=advisor_model, + messages=advisor_messages, + tools=None, + max_tokens=optional_params.get("max_tokens", 1024), + stream=False, + api_key=advisor_api_key, + api_base=advisor_api_base, + _advisor_interception_skip_post_hook=True, + ) + advisor_text = self._extract_text_content(advisor_response) + tool_messages.append( + { + "role": "tool", + "tool_call_id": advisor_call["id"], + "content": advisor_text, + } + ) + + current_messages = current_messages + [assistant_message] + tool_messages + + optional_params_clean = { + k: v + for k, v in optional_params.items() + if k + not in { + "tools", + "extra_body", + "model_alias_map", + "stream_response", + "custom_prompt_dict", + } + } + kwargs_for_followup = self._prepare_followup_kwargs(kwargs) + current_response = await litellm.acompletion( + model=self._get_full_model_name(model=model, kwargs=kwargs), + messages=current_messages, + tools=optional_params.get("tools"), + _advisor_interception_skip_post_hook=True, + **optional_params_clean, + **kwargs_for_followup, + ) + finally: + if isinstance(call_id, str): + self._advisor_config_by_call_id.pop(call_id, None) + + def _convert_tools_for_provider( + self, kwargs: Dict[str, Any], custom_llm_provider: str + ) -> Optional[Dict[str, Any]]: + if ( + self.enabled_providers is not None + and custom_llm_provider not in self.enabled_providers + ): + return None + + tools = kwargs.get("tools") + if not tools: + return None + if not any(is_advisor_tool(t) for t in tools): + return None + + advisor_cfg = self._extract_advisor_config(tools) + advisor_model = advisor_cfg.get("advisor_model") or self.default_advisor_model + if not advisor_model: + raise ValueError( + "Advisor tool requires a 'model'. Either pass native advisor tool " + "with `model`, or set default_advisor_model on AdvisorInterceptionLogger." + ) + max_uses = advisor_cfg.get("max_uses") + if max_uses is None: + max_uses = ADVISOR_MAX_USES + api_key = advisor_cfg.get("api_key") + api_base = advisor_cfg.get("api_base") + + converted_tools: List[Dict] = [] + if custom_llm_provider in ADVISOR_NATIVE_PROVIDERS: + for tool in tools: + if is_advisor_tool(tool): + converted_tools.append( + get_litellm_advisor_tool( + model=advisor_model, + max_uses=int(max_uses), + api_key=api_key, + api_base=api_base, + ) + ) + else: + converted_tools.append(tool) + kwargs["tools"] = converted_tools + return kwargs + + for tool in tools: + if is_advisor_tool(tool): + converted_tools.append(get_litellm_advisor_tool_openai()) + else: + converted_tools.append(tool) + kwargs["tools"] = converted_tools + call_id = kwargs.get("litellm_call_id") + if isinstance(call_id, str): + self._advisor_config_by_call_id[call_id] = { + "advisor_model": advisor_model, + "max_uses": int(max_uses), + "api_key": api_key, + "api_base": api_base, + } + if kwargs.get("stream"): + kwargs["stream"] = False + kwargs["_advisor_interception_converted_stream"] = True + return kwargs + + def _extract_advisor_config(self, tools: List[Dict]) -> Dict[str, Any]: + advisor_model = None + max_uses = None + api_key = None + api_base = None + + for tool in tools: + if not is_advisor_tool(tool): + continue + + if tool.get("type") == "advisor_20260301": + advisor_model = tool.get("model", advisor_model) + max_uses = tool.get("max_uses", max_uses) + api_key = tool.get("api_key", api_key) + api_base = tool.get("api_base", api_base) + + return { + "advisor_model": advisor_model, + "max_uses": max_uses, + "api_key": api_key, + "api_base": api_base, + } + + def _extract_advisor_tool_calls(self, response: Any) -> Tuple[List[Dict], List[Dict]]: + message = self._extract_first_choice_message(response) + if not message: + return [], [] + + tool_calls = message.get("tool_calls", []) + function_call = message.get("function_call") + advisor_calls: List[Dict] = [] + raw_tool_calls: List[Dict] = [] + + for tool_call in tool_calls: + function = tool_call.get("function", {}) + function_name = function.get("name") + if function_name not in {"advisor", LITELLM_ADVISOR_TOOL_NAME}: + continue + + arguments = function.get("arguments", {}) + if isinstance(arguments, str): + try: + parsed_args = json.loads(arguments) + except json.JSONDecodeError: + parsed_args = {} + elif isinstance(arguments, dict): + parsed_args = arguments + else: + parsed_args = {} + + advisor_calls.append( + { + "id": tool_call.get("id"), + "question": parsed_args.get("question"), + } + ) + raw_tool_calls.append(tool_call) + + # Some providers (e.g. Gemini in certain modes) can return legacy function_call. + if not advisor_calls and function_call is not None: + if isinstance(function_call, dict): + function_name = function_call.get("name") + function_arguments = function_call.get("arguments", {}) + else: + function_name = getattr(function_call, "name", None) + function_arguments = getattr(function_call, "arguments", {}) + if function_name in {"advisor", LITELLM_ADVISOR_TOOL_NAME}: + if isinstance(function_arguments, str): + try: + parsed_args = json.loads(function_arguments) + except json.JSONDecodeError: + parsed_args = {} + elif isinstance(function_arguments, dict): + parsed_args = function_arguments + else: + parsed_args = {} + + synthetic_id = "call_litellm_advisor_0" + advisor_calls.append( + { + "id": synthetic_id, + "question": parsed_args.get("question"), + } + ) + raw_tool_calls.append( + { + "id": synthetic_id, + "type": "function", + "function": { + "name": function_name, + "arguments": json.dumps(parsed_args), + }, + } + ) + return advisor_calls, raw_tool_calls + + @staticmethod + def _extract_first_choice_message(response: Any) -> Optional[Dict]: + if isinstance(response, dict): + choices = response.get("choices", []) + else: + choices = getattr(response, "choices", None) or [] + if not choices: + return None + + first_choice = choices[0] + if isinstance(first_choice, dict): + message = first_choice.get("message") + else: + message = getattr(first_choice, "message", None) + + if message is None: + return None + if isinstance(message, dict): + return message + + # pydantic object -> dict + tool_calls = getattr(message, "tool_calls", None) or [] + normalized_tool_calls = [] + for tc in tool_calls: + if isinstance(tc, dict): + normalized_tool_calls.append(tc) + else: + function = getattr(tc, "function", None) + normalized_tool_calls.append( + { + "id": getattr(tc, "id", None), + "type": getattr(tc, "type", None), + "function": { + "name": getattr(function, "name", None) if function else None, + "arguments": getattr(function, "arguments", None) + if function + else None, + }, + } + ) + return { + "role": getattr(message, "role", "assistant"), + "content": getattr(message, "content", None), + "tool_calls": normalized_tool_calls, + "function_call": getattr(message, "function_call", None), + } + + @staticmethod + def _extract_message_content(response: Any) -> Optional[str]: + message = AdvisorInterceptionLogger._extract_first_choice_message(response) + if not message: + return None + return message.get("content") + + @staticmethod + def _extract_text_content(response: Any) -> str: + message = AdvisorInterceptionLogger._extract_first_choice_message(response) + if not message: + return "" + content = message.get("content") + if isinstance(content, str): + return content + if isinstance(content, list): + parts: List[str] = [] + for block in content: + if isinstance(block, dict): + text = block.get("text") + if isinstance(text, str): + parts.append(text) + return "\n".join(parts).strip() + return "" + + @staticmethod + def _build_advisor_context( + messages: List[Dict], + assistant_content: Optional[str], + question: str, + ) -> List[Dict]: + advisor_messages = list(messages) + if assistant_content: + advisor_messages.append({"role": "assistant", "content": assistant_content}) + advisor_messages.append({"role": "user", "content": question}) + return advisor_messages + + @staticmethod + def _prepare_followup_kwargs(kwargs: Dict) -> Dict: + internal_params = { + "_advisor_interception", + "_advisor_interception_converted_stream", + "acompletion", + "litellm_logging_obj", + "custom_llm_provider", + "model_alias_map", + "stream_response", + "custom_prompt_dict", + "model", + "messages", + "tools", + "max_tokens", + "max_completion_tokens", + "temperature", + "top_p", + "n", + "stop", + "presence_penalty", + "frequency_penalty", + "logit_bias", + "user", + "response_format", + "seed", + "tool_choice", + "parallel_tool_calls", + "reasoning_effort", + "verbosity", + "extra_headers", + "api_version", + "metadata", + "web_search_options", + "safety_identifier", + "service_tier", + "stream", + } + return { + k: v + for k, v in kwargs.items() + if not k.startswith("_advisor_interception") and k not in internal_params + } + + @staticmethod + def _get_full_model_name(model: str, kwargs: Dict) -> str: + full_model_name = model + custom_llm_provider = kwargs.get("custom_llm_provider") + if custom_llm_provider and "/" not in model: + full_model_name = "{}/{}".format(custom_llm_provider, model) + return full_model_name + + @staticmethod + def _build_optional_params_from_request_data( + request_data: Dict[str, Any] + ) -> Dict[str, Any]: + allowed_keys = { + "max_tokens", + "max_completion_tokens", + "temperature", + "top_p", + "n", + "stop", + "presence_penalty", + "frequency_penalty", + "logit_bias", + "user", + "response_format", + "seed", + "tools", + "tool_choice", + "parallel_tool_calls", + "reasoning_effort", + "verbosity", + "extra_headers", + "api_version", + "metadata", + "web_search_options", + "safety_identifier", + "service_tier", + } + return {k: v for k, v in request_data.items() if k in allowed_keys} diff --git a/litellm/integrations/advisor_interception/tools.py b/litellm/integrations/advisor_interception/tools.py new file mode 100644 index 00000000000..aff1c50c2a7 --- /dev/null +++ b/litellm/integrations/advisor_interception/tools.py @@ -0,0 +1,102 @@ +""" +LiteLLM advisor tool definitions and detection helpers. +""" + +from typing import Any, Dict, Optional + +from litellm.constants import ADVISOR_TOOL_DESCRIPTION +from litellm.types.llms.anthropic import ANTHROPIC_ADVISOR_TOOL_TYPE + +LITELLM_ADVISOR_TOOL_NAME = "litellm_advisor" +_ADVISOR_TOOL_NAMES = {LITELLM_ADVISOR_TOOL_NAME, "advisor"} + + +def get_litellm_advisor_tool( + model: str, + max_uses: Optional[int] = None, + api_key: Optional[str] = None, + api_base: Optional[str] = None, +) -> Dict[str, Any]: + """ + Get the canonical advisor tool definition in Anthropic native format. + """ + tool: Dict[str, Any] = { + "type": ANTHROPIC_ADVISOR_TOOL_TYPE, + "name": "advisor", + "model": model, + } + if max_uses is not None: + tool["max_uses"] = max_uses + if api_key is not None: + tool["api_key"] = api_key + if api_base is not None: + tool["api_base"] = api_base + return tool + + +def get_litellm_advisor_tool_openai() -> Dict[str, Any]: + """ + Get the canonical advisor tool definition in OpenAI function format. + """ + return { + "type": "function", + "function": { + "name": LITELLM_ADVISOR_TOOL_NAME, + "description": ADVISOR_TOOL_DESCRIPTION, + "parameters": { + "type": "object", + "properties": { + "question": { + "type": "string", + "description": "The question or challenge you want guidance on.", + } + }, + "required": ["question"], + }, + }, + } + + +def is_advisor_tool_chat_completion(tool: Any) -> bool: + """ + Strict chat-completions advisor tool check. + """ + if isinstance(tool, dict): + tool_type = tool.get("type") + function_def = tool.get("function", {}) or {} + function_declarations = tool.get("function_declarations", []) + else: + tool_type = getattr(tool, "type", None) + function_def = getattr(tool, "function", None) or {} + function_declarations = getattr(tool, "function_declarations", None) or [] + + if tool_type == "function": + if not isinstance(function_def, dict): + function_def = { + "name": getattr(function_def, "name", None), + } + function_name = function_def.get("name") + return function_name in _ADVISOR_TOOL_NAMES + # Gemini tool schema: {"function_declarations": [{"name": "..."}]} + if isinstance(function_declarations, list): + for declaration in function_declarations: + if isinstance(declaration, dict): + name = declaration.get("name") + else: + name = getattr(declaration, "name", None) + if name in _ADVISOR_TOOL_NAMES: + return True + return False + + +def is_advisor_tool(tool: Any) -> bool: + """ + Check whether a tool is an advisor tool in any supported format. + """ + tool_type = tool.get("type") if isinstance(tool, dict) else getattr(tool, "type", None) + if tool_type == ANTHROPIC_ADVISOR_TOOL_TYPE: + return True + if is_advisor_tool_chat_completion(tool): + return True + tool_name = tool.get("name") if isinstance(tool, dict) else getattr(tool, "name", None) + return tool_name in _ADVISOR_TOOL_NAMES diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 9ecae363ed7..dc5e7f31382 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -281,6 +281,13 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 ) ) imported_list.append(websearch_interception_obj) + elif isinstance(callback, str) and callback == "advisor_interception": + from litellm.integrations.advisor_interception.handler import ( + AdvisorInterceptionLogger, + ) + + advisor_interception_obj = AdvisorInterceptionLogger() + imported_list.append(advisor_interception_obj) elif isinstance(callback, str) and callback == "datadog_cost_management": from litellm.integrations.datadog.datadog_cost_management import ( DatadogCostManagementLogger, diff --git a/tests/test_litellm/integrations/advisor_interception/test_advisor_interception_handler.py b/tests/test_litellm/integrations/advisor_interception/test_advisor_interception_handler.py new file mode 100644 index 00000000000..631356c617e --- /dev/null +++ b/tests/test_litellm/integrations/advisor_interception/test_advisor_interception_handler.py @@ -0,0 +1,150 @@ +import pytest + +from litellm.integrations.advisor_interception.handler import AdvisorInterceptionLogger +from litellm.integrations.advisor_interception.tools import ( + LITELLM_ADVISOR_TOOL_NAME, + get_litellm_advisor_tool_openai, +) +from litellm.types.utils import ( + ChatCompletionMessageToolCall, + Choices, + Function, + Message, + ModelResponse, +) + + +@pytest.mark.asyncio +async def test_pre_request_hook_non_native_converts_advisor_tool(): + logger = AdvisorInterceptionLogger(enabled_providers=["openai"]) + kwargs = { + "tools": [ + { + "type": "advisor_20260301", + "name": "advisor", + "model": "claude-opus-4-6", + "max_uses": 2, + } + ], + "stream": True, + "litellm_call_id": "call-1", + "litellm_params": {"custom_llm_provider": "openai"}, + } + + result = await logger.async_pre_request_hook( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Help"}], + kwargs=kwargs, + ) + assert result is not None + tool = result["tools"][0] + assert tool["type"] == "function" + assert tool["function"]["name"] == LITELLM_ADVISOR_TOOL_NAME + assert logger._advisor_config_by_call_id["call-1"]["advisor_model"] == "claude-opus-4-6" + assert logger._advisor_config_by_call_id["call-1"]["max_uses"] == 2 + assert result["stream"] is False + assert result["_advisor_interception_converted_stream"] is True + + +@pytest.mark.asyncio +async def test_pre_request_hook_anthropic_converts_standard_tool_to_native(): + logger = AdvisorInterceptionLogger(default_advisor_model="claude-opus-4-6") + kwargs = { + "tools": [get_litellm_advisor_tool_openai()], + "litellm_params": {"custom_llm_provider": "anthropic"}, + } + + result = await logger.async_pre_request_hook( + model="claude-sonnet-4-6", + messages=[{"role": "user", "content": "Help"}], + kwargs=kwargs, + ) + assert result is not None + tool = result["tools"][0] + assert tool["type"] == "advisor_20260301" + assert tool["name"] == "advisor" + assert tool["model"] == "claude-opus-4-6" + + +@pytest.mark.asyncio +async def test_should_run_chat_completion_agentic_loop_detects_advisor_tool_call(): + logger = AdvisorInterceptionLogger(enabled_providers=["openai"]) + mock_response = ModelResponse( + id="test", + choices=[ + Choices( + finish_reason="tool_calls", + index=0, + message=Message( + role="assistant", + content=None, + tool_calls=[ + ChatCompletionMessageToolCall( + id="call_123", + type="function", + function=Function( + name=LITELLM_ADVISOR_TOOL_NAME, + arguments='{"question":"What should I do?"}', + ), + ) + ], + ), + ) + ], + model="gpt-4o-mini", + object="chat.completion", + created=123, + ) + + should_run, tools_dict = await logger.async_should_run_chat_completion_agentic_loop( + response=mock_response, + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Help"}], + tools=[get_litellm_advisor_tool_openai()], + stream=False, + custom_llm_provider="openai", + kwargs={ + "litellm_call_id": "call-2", + }, + ) + assert should_run is True + assert len(tools_dict["advisor_calls"]) == 1 + assert tools_dict["advisor_calls"][0]["question"] == "What should I do?" + + +@pytest.mark.asyncio +async def test_should_run_chat_completion_agentic_loop_detects_legacy_function_call(): + logger = AdvisorInterceptionLogger(enabled_providers=["gemini"]) + mock_response = ModelResponse( + id="test", + choices=[ + Choices( + finish_reason="function_call", + index=0, + message=Message( + role="assistant", + content=None, + function_call={ + "name": LITELLM_ADVISOR_TOOL_NAME, + "arguments": '{"question":"Can you confirm advisor path?"}', + }, + ), + ) + ], + model="gemini/gemini-2.5-flash", + object="chat.completion", + created=123, + ) + + should_run, tools_dict = await logger.async_should_run_chat_completion_agentic_loop( + response=mock_response, + model="gemini/gemini-2.5-flash", + messages=[{"role": "user", "content": "Help"}], + tools=[get_litellm_advisor_tool_openai()], + stream=False, + custom_llm_provider="gemini", + kwargs={"litellm_call_id": "call-3"}, + ) + assert should_run is True + assert len(tools_dict["advisor_calls"]) == 1 + assert tools_dict["advisor_calls"][0]["question"] == "Can you confirm advisor path?"