diff --git a/litellm/integrations/advisor_interception/handler.py b/litellm/integrations/advisor_interception/handler.py index 1fee42eb8b3..14b420dc0b9 100644 --- a/litellm/integrations/advisor_interception/handler.py +++ b/litellm/integrations/advisor_interception/handler.py @@ -21,6 +21,7 @@ from litellm.integrations.advisor_interception.tools import ( is_advisor_tool_chat_completion, ) from litellm.integrations.custom_logger import CustomLogger +from litellm.types.integrations.advisor_interception import AdvisorInterceptionConfig from litellm.types.utils import CallTypes, LlmProviders @@ -32,7 +33,7 @@ class AdvisorInterceptionLogger(CustomLogger): def __init__( self, enabled_providers: Optional[List[Union[LlmProviders, str]]] = None, - default_advisor_model: str = "claude-opus-4-6", + default_advisor_model: Optional[str] = None, ): super().__init__() if enabled_providers is None: @@ -44,12 +45,69 @@ class AdvisorInterceptionLogger(CustomLogger): 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() + self._converted_stream_call_ids: set[str] = set() + + @classmethod + def from_config_yaml( + cls, config: AdvisorInterceptionConfig + ) -> "AdvisorInterceptionLogger": + """ + Initialize AdvisorInterceptionLogger from proxy config.yaml parameters. + + Args: + config: Configuration dictionary from litellm_settings.advisor_interception_params + + Example: + From proxy_config.yaml: + litellm_settings: + advisor_interception_params: + default_advisor_model: "advisor-model" + enabled_providers: ["openai", "vertex_ai"] + """ + enabled_providers_str = config.get("enabled_providers", None) + default_advisor_model = config.get("default_advisor_model", None) + + enabled_providers: Optional[List[Union[LlmProviders, str]]] = None + if enabled_providers_str is not None: + enabled_providers = [] + for provider in enabled_providers_str: + try: + provider_enum = LlmProviders(provider) + enabled_providers.append(provider_enum) + except ValueError: + enabled_providers.append(provider) + + return cls( + enabled_providers=enabled_providers, + default_advisor_model=default_advisor_model, + ) + + @staticmethod + def initialize_from_proxy_config( + litellm_settings: Dict[str, Any], + callback_specific_params: Dict[str, Any], + ) -> "AdvisorInterceptionLogger": + """ + Static method to initialize AdvisorInterceptionLogger from proxy config. + + Used in callback_utils.py to simplify initialization logic. + """ + advisor_params: AdvisorInterceptionConfig = {} + if "advisor_interception_params" in litellm_settings: + advisor_params = litellm_settings["advisor_interception_params"] + elif "advisor_interception" in callback_specific_params: + advisor_params = callback_specific_params["advisor_interception"] + + return AdvisorInterceptionLogger.from_config_yaml(advisor_params) 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. + + Skips conversion for anthropic_messages call type because the Messages + API path has its own AdvisorOrchestrationHandler interceptor. """ custom_llm_provider = kwargs.get("litellm_params", {}).get( "custom_llm_provider", "" @@ -63,7 +121,14 @@ class AdvisorInterceptionLogger(CustomLogger): ) -> Optional[dict]: """ Pre-call hook used by completion/chat-completions paths. + + Skips conversion for anthropic_messages call type because the Messages + API path has its own AdvisorOrchestrationHandler interceptor that + expects the raw advisor_20260301 tool definition. """ + if call_type == CallTypes.anthropic_messages: + return None + if kwargs.pop("_advisor_interception_skip_post_hook", False): call_id = kwargs.get("litellm_call_id") if isinstance(call_id, str): @@ -97,13 +162,23 @@ class AdvisorInterceptionLogger(CustomLogger): return None call_id = request_data.get("litellm_call_id") + converted_stream = ( + isinstance(call_id, str) and call_id in self._converted_stream_call_ids + ) + if converted_stream: + self._converted_stream_call_ids.discard(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) + if converted_stream: + return self._wrap_as_streaming_if_needed(response) return None model = request_data.get("model") messages = request_data.get("messages") if not isinstance(model, str) or not isinstance(messages, list): + if converted_stream: + return self._wrap_as_streaming_if_needed(response) return None custom_llm_provider = request_data.get("custom_llm_provider", "") or request_data.get( @@ -130,10 +205,12 @@ class AdvisorInterceptionLogger(CustomLogger): if not should_run: if isinstance(call_id, str): self._advisor_config_by_call_id.pop(call_id, None) + if converted_stream: + return self._wrap_as_streaming_if_needed(response) return None optional_params = self._build_optional_params_from_request_data(request_data) - return await self.async_run_chat_completion_agentic_loop( + result = await self.async_run_chat_completion_agentic_loop( tools=tools_dict, model=model, messages=messages, @@ -143,6 +220,9 @@ class AdvisorInterceptionLogger(CustomLogger): stream=stream, kwargs=request_data, ) + if converted_stream: + return self._wrap_as_streaming_if_needed(result) + return result async def async_should_run_chat_completion_agentic_loop( self, @@ -162,8 +242,17 @@ class AdvisorInterceptionLogger(CustomLogger): and custom_llm_provider not in self.enabled_providers ): return False, {} + + # Only skip the orchestration loop for native providers when the advisor + # model is actually supported natively (Anthropic Claude Opus 4.6). + # For Anthropic executors with a non-native advisor model, fall through + # to the orchestration loop below. if custom_llm_provider in ADVISOR_NATIVE_PROVIDERS: - return False, {} + call_id_check = kwargs.get("litellm_call_id") + advisor_cfg = self._advisor_config_by_call_id.get(call_id_check, {}) if isinstance(call_id_check, str) else {} + advisor_model_check = advisor_cfg.get("advisor_model") or self.default_advisor_model or "" + if self._is_native_anthropic_advisor_model(advisor_model_check): + return False, {} call_id = kwargs.get("litellm_call_id") has_advisor_config = isinstance(call_id, str) and ( @@ -219,14 +308,24 @@ class AdvisorInterceptionLogger(CustomLogger): 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 + if not advisor_model: + raise ValueError( + "No advisor model configured. Either:\n" + " 1. Set 'default_advisor_model' in advisor_interception_params in your proxy config YAML, or\n" + " 2. Pass 'model' in the native advisor_20260301 tool definition.\n" + "The advisor model should be a model_name from your model_list for correct credential resolution." + ) advisor_api_key = advisor_config.get("api_key") advisor_api_base = advisor_config.get("api_base") call_id = kwargs.get("litellm_call_id") + llm_router = self._get_llm_router() + current_messages: List[Dict] = list(messages) current_response = response advisor_uses = 0 total_response_cost = self._safe_get_response_cost(current_response) + advisor_interactions: List[Dict[str, str]] = [] try: while True: @@ -237,6 +336,9 @@ class AdvisorInterceptionLogger(CustomLogger): self._set_response_cost_if_possible( response=current_response, response_cost=total_response_cost ) + self._inject_advisor_results_into_response( + current_response, advisor_interactions + ) return current_response if len(advisor_calls) != len(raw_tool_calls): verbose_logger.debug( @@ -269,18 +371,20 @@ class AdvisorInterceptionLogger(CustomLogger): assistant_content=assistant_content, question=question, ) - advisor_response = await litellm.acompletion( - model=advisor_model, + advisor_response = await self._call_advisor_model( + llm_router=llm_router, + advisor_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, ) total_response_cost += self._safe_get_response_cost(advisor_response) advisor_text = self._extract_text_content(advisor_response) + advisor_interactions.append({ + "tool_use_id": advisor_call["id"], + "advisor_text": advisor_text, + }) tool_messages.append( { "role": "tool", @@ -297,6 +401,7 @@ class AdvisorInterceptionLogger(CustomLogger): if k not in { "tools", + "tool_choice", # never force tool use on follow-up turns "extra_body", "model_alias_map", "stream_response", @@ -304,19 +409,165 @@ class AdvisorInterceptionLogger(CustomLogger): } } kwargs_for_followup = self._prepare_followup_kwargs(kwargs) - current_response = await litellm.acompletion( - model=self._get_full_model_name(model=model, kwargs=kwargs), + executor_model = self._get_full_model_name(model=model, kwargs=kwargs) + current_response = await self._call_executor_model( + llm_router=llm_router, + model=executor_model, messages=current_messages, tools=optional_params.get("tools"), - _advisor_interception_skip_post_hook=True, - **optional_params_clean, - **kwargs_for_followup, + optional_params_clean=optional_params_clean, + kwargs_for_followup=kwargs_for_followup, ) total_response_cost += self._safe_get_response_cost(current_response) finally: if isinstance(call_id, str): self._advisor_config_by_call_id.pop(call_id, None) + @staticmethod + def _wrap_as_streaming_if_needed(response: Any) -> Any: + """ + Wrap a ModelResponse in a MockResponseIterator so the proxy can + async-iterate it when the original request was stream=True but the + advisor hook converted it to stream=False for the agentic loop. + """ + from litellm.types.utils import ModelResponse as _ModelResponse + + if isinstance(response, _ModelResponse): + from litellm.llms.base_llm.base_model_iterator import ( + MockResponseIterator, + ) + + return MockResponseIterator(response) + return response + + @staticmethod + def _get_llm_router() -> Optional[Any]: + """Import the proxy router at runtime. Returns None in SDK-only usage.""" + try: + from litellm.proxy.proxy_server import llm_router + except ImportError: + verbose_logger.debug( + "AdvisorInterception: Could not import llm_router from proxy_server, " + "falling back to direct litellm.acompletion()" + ) + llm_router = None + return llm_router + + @staticmethod + def _is_native_anthropic_advisor_model(advisor_model: str) -> bool: + """ + Return True only when the advisor model resolves to Anthropic Claude Opus 4.6, + which is the only model Anthropic supports as a native advisor. + + Handles bare model names, litellm provider-prefixed names + (e.g. ``anthropic/claude-opus-4-6``) and proxy model aliases. + """ + # Resolve proxy alias → underlying litellm model string first. + try: + llm_router = AdvisorInterceptionLogger._get_llm_router() + if llm_router is not None: + for deployment in llm_router.model_list or []: + if deployment.get("model_name") == advisor_model: + advisor_model = ( + deployment.get("litellm_params", {}).get("model") + or advisor_model + ) + break + except Exception: + pass + + normalized = advisor_model.lower().replace("_", "-") + # Must be an Anthropic model and specifically opus-4-6. + is_anthropic = normalized.startswith("anthropic/") or ( + "/" not in normalized and "claude" in normalized + ) + is_opus_46 = "claude-opus-4-6" in normalized or "claude-opus-4.6" in normalized + return is_anthropic and is_opus_46 + + @staticmethod + async def _call_advisor_model( + llm_router: Optional[Any], + advisor_model: str, + messages: List[Dict], + max_tokens: int, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> Any: + """ + Call the advisor model, routing through the proxy router when available + so that deployed credentials and load-balancing are used. + """ + if llm_router is not None: + try: + return await llm_router.acompletion( + model=advisor_model, + messages=messages, + tools=None, + max_tokens=max_tokens, + stream=False, + _advisor_interception_skip_post_hook=True, + ) + except Exception: + verbose_logger.debug( + "AdvisorInterception: Router call for advisor model '%s' failed, " + "falling back to direct litellm.acompletion()", + advisor_model, + ) + + kwargs: Dict[str, Any] = {} + if api_key is not None: + kwargs["api_key"] = api_key + if api_base is not None: + kwargs["api_base"] = api_base + return await litellm.acompletion( + model=advisor_model, + messages=messages, + tools=None, + max_tokens=max_tokens, + stream=False, + _advisor_interception_skip_post_hook=True, + **kwargs, + ) + + @staticmethod + async def _call_executor_model( + llm_router: Optional[Any], + model: str, + messages: List[Dict], + tools: Optional[List[Dict]], + optional_params_clean: Dict[str, Any], + kwargs_for_followup: Dict[str, Any], + ) -> Any: + """ + Call the executor model for the follow-up turn, routing through the + proxy router when available. + """ + if llm_router is not None: + try: + return await llm_router.acompletion( + model=model, + messages=messages, + tools=tools, + _advisor_interception_skip_post_hook=True, + **optional_params_clean, + **kwargs_for_followup, + ) + except Exception: + verbose_logger.debug( + "AdvisorInterception: Router call for executor model '%s' failed, " + "falling back to direct litellm.acompletion()", + model, + ) + + return await litellm.acompletion( + model=model, + messages=messages, + tools=tools, + _advisor_interception_skip_post_hook=True, + **optional_params_clean, + **kwargs_for_followup, + ) + def _convert_tools_for_provider( self, kwargs: Dict[str, Any], custom_llm_provider: str ) -> Optional[Dict[str, Any]]: @@ -336,8 +587,10 @@ class AdvisorInterceptionLogger(CustomLogger): 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." + "No advisor model configured. Either:\n" + " 1. Set 'default_advisor_model' in advisor_interception_params in your proxy config YAML, or\n" + " 2. Pass 'model' in the native advisor_20260301 tool definition.\n" + "The advisor model should be a model_name from your model_list for correct credential resolution." ) max_uses = advisor_cfg.get("max_uses") if max_uses is None: @@ -346,7 +599,11 @@ class AdvisorInterceptionLogger(CustomLogger): api_base = advisor_cfg.get("api_base") converted_tools: List[Dict] = [] - if custom_llm_provider in ADVISOR_NATIVE_PROVIDERS: + use_native = ( + custom_llm_provider in ADVISOR_NATIVE_PROVIDERS + and self._is_native_anthropic_advisor_model(advisor_model) + ) + if use_native: for tool in tools: if is_advisor_tool(tool): converted_tools.append( @@ -378,7 +635,12 @@ class AdvisorInterceptionLogger(CustomLogger): } if kwargs.get("stream"): kwargs["stream"] = False - kwargs["_advisor_interception_converted_stream"] = True + call_id_for_stream = kwargs.get("litellm_call_id") + if isinstance(call_id_for_stream, str): + self._converted_stream_call_ids.add(call_id_for_stream) + litellm_params = kwargs.get("litellm_params") + if isinstance(litellm_params, dict): + litellm_params["_advisor_interception_converted_stream"] = True return kwargs def _extract_advisor_config(self, tools: List[Dict]) -> Dict[str, Any]: @@ -477,6 +739,66 @@ class AdvisorInterceptionLogger(CustomLogger): ) return advisor_calls, raw_tool_calls + @staticmethod + def _inject_advisor_results_into_response( + response: Any, advisor_interactions: List[Dict[str, str]] + ) -> None: + """ + Add ``advisor_tool_result`` blocks to ``provider_specific_fields`` + of the final chat-completion response message. + + This gives callers the same advisor visibility as the Anthropic + native ``/v1/messages`` path. + """ + if not advisor_interactions: + return + + advisor_results: List[Dict] = [] + for interaction in advisor_interactions: + tool_use_id = interaction["tool_use_id"] + advisor_text = interaction["advisor_text"] + advisor_results.append({ + "type": "server_tool_use", + "id": tool_use_id, + "name": "advisor", + }) + advisor_results.append({ + "type": "advisor_tool_result", + "tool_use_id": tool_use_id, + "content": { + "type": "advisor_result", + "text": advisor_text, + }, + }) + + message = AdvisorInterceptionLogger._extract_first_choice_message_obj(response) + if message is None: + return + + existing_psf = getattr(message, "provider_specific_fields", None) or {} + existing_psf["advisor_tool_results"] = advisor_results + try: + message.provider_specific_fields = existing_psf + except Exception: + try: + setattr(message, "provider_specific_fields", existing_psf) + except Exception: + pass + + @staticmethod + def _extract_first_choice_message_obj(response: Any) -> Any: + """Return the raw message object (not dict-normalised) from the first choice.""" + 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): + return first_choice.get("message") + return getattr(first_choice, "message", None) + @staticmethod def _extract_first_choice_message(response: Any) -> Optional[Dict]: if isinstance(response, dict): @@ -605,7 +927,6 @@ class AdvisorInterceptionLogger(CustomLogger): def _prepare_followup_kwargs(kwargs: Dict) -> Dict: internal_params = { "_advisor_interception", - "_advisor_interception_converted_stream", "acompletion", "litellm_logging_obj", "custom_llm_provider",