From ec268b0d18c525076cb970bdd13d65896586f49d Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 23 Jun 2026 07:29:31 -0700 Subject: [PATCH 01/29] refactor(completion): extract provider dispatch into typed helpers so basedpyright can analyze it (#30813) --- litellm/main.py | 6636 +++++++++++-------- litellm/types/completion.py | 63 +- ruff-strict-budget.json | 2 +- tests/test_litellm/types/test_completion.py | 61 +- 4 files changed, 4045 insertions(+), 2717 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 63c5798e70a..1d75766c7e6 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -118,6 +118,10 @@ from litellm.llms.vertex_ai.common_utils import ( ) from litellm.realtime_api.main import _realtime_health_check from litellm.secret_managers.main import get_secret_bool, get_secret_str +from litellm.types.completion import ( + _CompletionDispatchContext, + _CompletionDispatchResult, +) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -1084,6 +1088,3831 @@ def _build_custom_pricing_entry( return entry +def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + _azure_detection_model = ctx._azure_detection_model + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + api_version = ctx.api_version + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + extra_headers = ctx.extra_headers + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + max_retries = ctx.max_retries + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + dynamic_params = False + if client is not None and ( + isinstance(client, openai.AzureOpenAI) + or isinstance(client, openai.AsyncAzureOpenAI) + ): + dynamic_params = _check_dynamic_azure_params( + azure_client_params={"api_version": api_version}, + azure_client=client, + ) + + api_type = get_secret("AZURE_API_TYPE") or "azure" + + api_base = api_base or litellm.api_base or get_secret("AZURE_API_BASE") + + api_version = ( + api_version + or litellm.api_version + or get_secret_str("AZURE_API_VERSION") + or litellm.AZURE_DEFAULT_API_VERSION + ) + + api_key = ( + api_key + or litellm.api_key + or litellm.azure_key + or get_secret_str("AZURE_OPENAI_API_KEY") + or get_secret_str("AZURE_API_KEY") + ) + + azure_ad_token = optional_params.get("extra_body", {}).pop( + "azure_ad_token", None + ) or get_secret_str("AZURE_AD_TOKEN") + + azure_ad_token_provider = litellm_params.get("azure_ad_token_provider", None) + + headers = headers or litellm.headers + + if extra_headers is not None: + optional_params["extra_headers"] = extra_headers + if max_retries is not None: + optional_params["max_retries"] = max_retries + + if litellm.AzureOpenAIO1Config().is_o_series_model(model=_azure_detection_model): + ## LOAD CONFIG - if set + config = litellm.AzureOpenAIO1Config.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + + response = azure_o1_chat_completions.completion( + model=model, + messages=messages, + headers=headers, + api_key=api_key, + api_base=api_base, + api_version=api_version, + dynamic_params=dynamic_params, + azure_ad_token=azure_ad_token, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + logging_obj=logging, + acompletion=acompletion, + timeout=timeout, # type: ignore + client=client, # pass AsyncAzureOpenAI, AzureOpenAI client + custom_llm_provider=custom_llm_provider, + ) + else: + ## LOAD CONFIG - if set + config = litellm.AzureOpenAIConfig.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + + ## COMPLETION CALL + response = azure_chat_completions.completion( + model=model, + messages=messages, + headers=headers, + api_key=api_key, + api_base=api_base, + api_version=api_version, + api_type=api_type, + dynamic_params=dynamic_params, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + logging_obj=logging, + acompletion=acompletion, + timeout=timeout, # type: ignore + client=client, # pass AsyncAzureOpenAI, AzureOpenAI client + ) + + if optional_params.get("stream", False): + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + additional_args={ + "headers": headers, + "api_version": api_version, + "api_base": api_base, + }, + ) + + return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + api_version = ctx.api_version + client = ctx.client + extra_headers = ctx.extra_headers + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + api_type = get_secret_str("AZURE_API_TYPE") or "azure" + + api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") + + if api_base is None: + raise ValueError( + "api_base is required for Azure OpenAI LLM provider. Either set it dynamically or set the AZURE_API_BASE environment variable." + ) + + api_version = ( + api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") + ) + + api_key = ( + api_key + or litellm.api_key + or litellm.azure_key + or get_secret_str("AZURE_OPENAI_API_KEY") + or get_secret_str("AZURE_API_KEY") + ) + + azure_ad_token = optional_params.get("extra_body", {}).pop( + "azure_ad_token", None + ) or get_secret_str("AZURE_AD_TOKEN") + + azure_ad_token_provider = litellm_params.get("azure_ad_token_provider", None) + + headers = headers or litellm.headers + + if extra_headers is not None: + optional_params["extra_headers"] = extra_headers + + ## LOAD CONFIG - if set + config = litellm.AzureOpenAIConfig.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + + ## COMPLETION CALL + response = azure_text_completions.completion( + model=model, + messages=messages, + headers=headers, + api_key=api_key, + api_base=api_base, + api_version=cast(str, api_version), + api_type=api_type, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + logging_obj=logging, + acompletion=acompletion, + timeout=timeout, + client=client, # pass AsyncAzureOpenAI, AzureOpenAI client + ) + + if optional_params.get("stream", False) or acompletion is True: + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + additional_args={ + "headers": headers, + "api_version": api_version, + "api_base": api_base, + }, + ) + + return response + + +def _complete_deepseek(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_azure_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + extra_headers = ctx.extra_headers + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo + + azure_ai_route = AzureFoundryModelInfo.get_azure_ai_route(model) + + # Check if this is an agents route - model format: azure_ai/agents/ + if azure_ai_route == "agents": + from litellm.llms.azure_ai.agents import AzureAIAgentsConfig + + api_base = AzureFoundryModelInfo.get_api_base(api_base) + if api_base is None: + raise ValueError( + "Azure AI Agents requests require an api_base. " + "Set `api_base` or the AZURE_AI_API_BASE env var." + ) + api_key = AzureFoundryModelInfo.get_api_key(api_key) + + response = AzureAIAgentsConfig.completion( + model=model, + messages=messages, + api_base=api_base, + api_key=api_key, + model_response=model_response, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, + acompletion=acompletion, + stream=stream, + headers=headers or litellm.headers, + ) + + # Check if this is a Claude model - route to Azure Anthropic handler + elif "claude" in model.lower(): + # Use Azure Anthropic handler for Claude models + api_base = AzureFoundryModelInfo.get_api_base(api_base) + if api_base is None: + raise ValueError( + "Azure Anthropic requests require an api_base. " + "Set `api_base` or the AZURE_AI_API_BASE env var." + ) + api_key = AzureFoundryModelInfo.get_api_key(api_key) + + # Ensure the URL ends with /v1/messages for Anthropic + if api_base: + api_base = api_base.rstrip("/") + if not api_base.endswith("/v1/messages"): + if "/anthropic" in api_base: + parts = api_base.split("/anthropic", 1) + api_base = parts[0] + "/anthropic" + else: + api_base = api_base + "/anthropic" + api_base = api_base + "/v1/messages" + + response = azure_anthropic_chat_completions.completion( + model=model, + messages=messages, + api_base=api_base, + acompletion=acompletion, + custom_prompt_dict=litellm.custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + headers=headers, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + ) + if optional_params.get("stream", False) or acompletion is True: + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + ) + response = response + else: + # Non-Claude models use standard Azure AI flow + api_base = AzureFoundryModelInfo.get_api_base(api_base) + # set API KEY + api_key = AzureFoundryModelInfo.get_api_key(api_key) + + headers = headers or litellm.headers + + if extra_headers is not None: + optional_params["extra_headers"] = extra_headers + + ## FOR COHERE + if "command-r" in model: # make sure tool call in messages are str + messages = stringify_json_tool_call_content(messages=messages) + + ## COMPLETION CALL + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, # type: ignore + client=client, # pass AsyncOpenAI, OpenAI client + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + if optional_params.get("stream", False): + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + additional_args={"headers": headers}, + ) + + return response + + +def _complete_text_completion_openai( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + text_completion = ctx.text_completion + timeout = ctx.timeout + + openai.api_type = "openai" + + api_base = ( + api_base + or litellm.api_base + or get_secret("OPENAI_BASE_URL") + or get_secret("OPENAI_API_BASE") + or "https://api.openai.com/v1" + ) + + openai.api_version = None + # set API KEY + + api_key = ( + api_key or litellm.api_key or litellm.openai_key or get_secret("OPENAI_API_KEY") + ) + + headers = headers or litellm.headers + + ## LOAD CONFIG - if set + config = litellm.OpenAITextCompletionConfig.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > openai_text_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + if litellm.organization: + openai.organization = litellm.organization + + ## COMPLETION CALL + _response = openai_text_completions.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + print_verbose=print_verbose, + api_key=api_key, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + acompletion=acompletion, + client=client, # pass AsyncOpenAI, OpenAI client + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + timeout=timeout, # type: ignore + ) + + if ( + optional_params.get("stream", False) is False + and acompletion is False + and text_completion is False + ): + # convert to chat completion response + _response = ( + litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object( + response_object=_response, model_response_object=model_response + ) + ) + + if optional_params.get("stream", False) or acompletion is True: + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=_response, + additional_args={"headers": headers}, + ) + return _response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_fireworks_ai( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_heroku(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_ragflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_xai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_groq(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there + or litellm.api_base + or get_secret("GROQ_API_BASE") + or "https://api.groq.com/openai/v1" + ) + + # set API KEY + api_key = ( + api_key + or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there + or litellm.groq_key + or get_secret("GROQ_API_KEY") + ) + + headers = headers or litellm.headers + + ## LOAD CONFIG - if set + config = litellm.GroqChatConfig.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + + +def _complete_bedrock_mantle( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = api_base or litellm.api_base or get_secret("BEDROCK_MANTLE_API_BASE") + api_key = api_key or litellm.api_key or get_secret("BEDROCK_MANTLE_API_KEY") + headers = headers or litellm.headers + config = litellm.BedrockMantleChatConfig.get_config() + for k, v in config.items(): + if k not in optional_params: + optional_params[k] = v + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + ) + + +def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + ( + api_base, + api_key, + headers, + ) = litellm.A2AConfig.resolve_agent_config_from_registry( + model=model, + api_base=api_base, + api_key=api_key, + headers=headers, + optional_params=optional_params, + ) + + # Fall back to environment variables and defaults + api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE") + + if api_base is None: + raise Exception( + "api_base is required for A2A provider. " + "Either provide api_base parameter, set A2A_API_BASE environment variable, " + "or register the agent in the proxy with model='a2a/'." + ) + + headers = headers or litellm.headers + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + provider_config=provider_config, + ) + + +def _complete_gigachat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.api_key + or litellm.gigachat_key + or get_secret("GIGACHAT_API_KEY") + or get_secret("GIGACHAT_CREDENTIALS") + ) + + headers = headers or litellm.headers or {} + + ## COMPLETION CALL + try: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_sap(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + headers = headers or litellm.headers + ## LOAD CONFIG - if set + config = litellm.GenAIHubOrchestrationConfig.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + + return sap_gen_ai_hub_chat_completions.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + shared_session=shared_session, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + api_key=api_key, + api_base=api_base, + stream=stream, + ) + + +def _complete_aiohttp_openai( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + extra_headers = ctx.extra_headers + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there + or litellm.api_base + or get_secret("OPENAI_BASE_URL") + or get_secret("OPENAI_API_BASE") + or "https://api.openai.com/v1" + ) + # set API KEY + api_key = ( + api_key + or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there + or litellm.openai_key + or get_secret("OPENAI_API_KEY") + ) + + headers = headers or litellm.headers + + if extra_headers is not None: + optional_params["extra_headers"] = extra_headers + return base_llm_aiohttp_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + ) + + +def _complete_cometapi(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.cometapi_key + or get_secret_str("COMETAPI_KEY") + or litellm.api_key + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret_str("COMETAPI_API_BASE") + or "https://api.cometapi.com/v1" + ) + + ## COMPLETION CALL + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + + ## LOGGING + logging.post_call(input=messages, api_key=api_key, original_response=response) + + return response + + +def _complete_minimax(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = api_key or get_secret_str("MINIMAX_API_KEY") or litellm.api_key + + api_base = ( + api_base + or litellm.api_base + or get_secret_str("MINIMAX_API_BASE") + or "https://api.minimax.io/v1" + ) + + response = base_llm_http_handler.completion( + model=model, + messages=messages, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + model_response=model_response, + encoding=_get_encoding(), + logging_obj=logging, + optional_params=optional_params, + timeout=timeout, + litellm_params=litellm_params, + shared_session=shared_session, + acompletion=acompletion, + stream=stream, + api_key=api_key, + headers=headers, + client=client, + provider_config=provider_config, + ) + logging.post_call(input=messages, api_key=api_key, original_response=response) + + return response + + +def _complete_hosted_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = api_base or litellm.api_base or get_secret_str("HOSTED_VLLM_API_BASE") + + response = base_llm_http_handler.completion( + model=model, + messages=messages, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + model_response=model_response, + encoding=_get_encoding(), + logging_obj=logging, + optional_params=optional_params, + timeout=timeout, + litellm_params=litellm_params, + shared_session=shared_session, + acompletion=acompletion, + stream=stream, + api_key=api_key, + headers=headers, + client=client, + provider_config=provider_config, + ) + logging.post_call(input=messages, api_key=api_key, original_response=response) + + return response + + +def _complete_custom_openai( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + custom_prompt_dict = ctx.custom_prompt_dict + extra_headers = ctx.extra_headers + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + metadata = ctx.metadata + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + organization = ctx.organization + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there + or litellm.api_base + or get_secret("OPENAI_BASE_URL") + or get_secret("OPENAI_API_BASE") + or "https://api.openai.com/v1" + ) + organization = ( + organization + or litellm.organization + or get_secret("OPENAI_ORGANIZATION") + or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105 + ) + openai.organization = organization + # set API KEY + api_key = ( + api_key + or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there + or litellm.openai_key + or get_secret("OPENAI_API_KEY") + ) + + headers = headers or litellm.headers + + # Add GitHub Copilot headers (same as /responses endpoint does) + if custom_llm_provider == "github_copilot": + from litellm.llms.github_copilot.authenticator import Authenticator + from litellm.llms.github_copilot.common_utils import ( + get_copilot_default_headers, + ) + + copilot_auth = Authenticator() + copilot_api_key = copilot_auth.get_api_key() + copilot_headers = get_copilot_default_headers(copilot_api_key) + if extra_headers: + copilot_headers.update(extra_headers) + extra_headers = copilot_headers + + if extra_headers is not None: + optional_params["extra_headers"] = extra_headers + + if ( + litellm.enable_preview_features and metadata is not None + ): # [PREVIEW] allow metadata to be passed to OPENAI + openai_metadata = get_requester_metadata(metadata) + if openai_metadata is not None: + optional_params["metadata"] = openai_metadata + + ## LOAD CONFIG - if set + config = litellm.OpenAIConfig.get_config() + for k, v in config.items(): + if ( + k not in optional_params + ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in + optional_params[k] = v + + ## COMPLETION CALL + use_base_llm_http_handler = get_secret_bool( + "EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER" + ) + + try: + if use_base_llm_http_handler: + response = base_llm_http_handler.completion( + model=model, + messages=messages, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + model_response=model_response, + encoding=_get_encoding(), + logging_obj=logging, + optional_params=optional_params, + timeout=timeout, + litellm_params=litellm_params, + shared_session=shared_session, + acompletion=acompletion, + stream=stream, + api_key=api_key, + headers=headers, + client=client, + provider_config=provider_config, + ) + else: + response = openai_chat_completions.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + print_verbose=print_verbose, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + timeout=timeout, # type: ignore + custom_prompt_dict=custom_prompt_dict, + client=client, # pass AsyncOpenAI, OpenAI client + organization=organization, + custom_llm_provider=custom_llm_provider, + shared_session=shared_session, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + if optional_params.get("stream", False): + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + additional_args={"headers": headers}, + ) + + return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_mistral(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = api_key or litellm.api_key or get_secret("MISTRAL_API_KEY") + api_base = ( + api_base + or litellm.api_base + or get_secret("MISTRAL_API_BASE") + or "https://api.mistral.ai/v1" + ) + + return base_llm_http_handler.completion( + model=model, + messages=messages, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + model_response=model_response, + encoding=_get_encoding(), + logging_obj=logging, + optional_params=optional_params, + timeout=timeout, + litellm_params=litellm_params, + shared_session=shared_session, + acompletion=acompletion, + stream=stream, + api_key=api_key, + headers=headers, + client=client, + provider_config=provider_config, + ) + + +def _complete_replicate(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + replicate_key = ( + api_key + or litellm.replicate_key + or litellm.api_key + or get_secret("REPLICATE_API_KEY") + or get_secret("REPLICATE_API_TOKEN") + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret("REPLICATE_API_BASE") + or "https://api.replicate.com/v1" + ) + + custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict + + model_response = replicate_chat_completion( # type: ignore + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), # for calculating input/output tokens + api_key=replicate_key, + logging_obj=logging, + custom_prompt_dict=custom_prompt_dict, + acompletion=acompletion, + headers=headers, + ) + + if optional_params.get("stream", False) is True: + ## LOGGING + logging.post_call( + input=messages, + api_key=replicate_key, + original_response=model_response, + ) + + return model_response + + +def _complete_anthropic_text( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.anthropic_key + or litellm.api_key + or os.environ.get("ANTHROPIC_API_KEY") + ) + custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict + api_base = cast( + Optional[str], + api_base + or litellm.api_base + or get_secret("ANTHROPIC_API_BASE") + or get_secret("ANTHROPIC_BASE_URL") + or "https://api.anthropic.com/v1/complete", + ) + + # Check if we should disable automatic URL suffix appending + disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX") + if ( + api_base is not None + and not disable_url_suffix + and not api_base.endswith("/v1/complete") + ): + api_base += "/v1/complete" + elif disable_url_suffix: + verbose_logger.debug( + "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/complete suffix" + ) + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="anthropic_text", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + ) + + +def _complete_anthropic(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.anthropic_key + or litellm.api_key + or os.environ.get("ANTHROPIC_API_KEY") + ) + custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict + # call /messages + # default route for all anthropic models + api_base = cast( + Optional[str], + api_base + or litellm.api_base + or get_secret("ANTHROPIC_API_BASE") + or get_secret("ANTHROPIC_BASE_URL") + or "https://api.anthropic.com/v1/messages", + ) + + # Check if we should disable automatic URL suffix appending + disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX") + if ( + api_base is not None + and not disable_url_suffix + and not api_base.endswith("/v1/messages") + ): + api_base += "/v1/messages" + elif disable_url_suffix: + verbose_logger.debug( + "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/messages suffix" + ) + + response = anthropic_chat_completions.completion( + model=model, + messages=messages, + api_base=api_base, + acompletion=acompletion, + custom_prompt_dict=litellm.custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), # for calculating input/output tokens + api_key=api_key, + logging_obj=logging, + headers=headers, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + ) + if optional_params.get("stream", False) or acompletion is True: + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + ) + return response + + +def _complete_nlp_cloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + nlp_cloud_key = ( + api_key + or litellm.nlp_cloud_key + or get_secret("NLP_CLOUD_API_KEY") + or litellm.api_key + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret("NLP_CLOUD_API_BASE") + or "https://api.nlpcloud.io/v1/gpu/" + ) + + response = nlp_cloud_chat_completion( + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + api_key=nlp_cloud_key, + logging_obj=logging, + ) + + if "stream" in optional_params and optional_params["stream"] is True: + # don't try to access stream object, + response = CustomStreamWrapper( + response, + model, + custom_llm_provider="nlp_cloud", + logging_obj=logging, + ) + + if optional_params.get("stream", False) or acompletion is True: + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + ) + + return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_aleph_alpha(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base = ctx.api_base + api_key = ctx.api_key + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + aleph_alpha_key = ( + api_key + or litellm.aleph_alpha_key + or get_secret("ALEPH_ALPHA_API_KEY") + or get_secret("ALEPHALPHA_API_KEY") + or litellm.api_key + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret("ALEPH_ALPHA_API_BASE") + or "https://api.aleph-alpha.com/complete" + ) + + model_response = aleph_alpha.completion( + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + default_max_tokens_to_sample=litellm.max_tokens, + api_key=aleph_alpha_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + ) + + if "stream" in optional_params and optional_params["stream"] is True: + # don't try to access stream object, + return CustomStreamWrapper( + model_response, + model, + custom_llm_provider="aleph_alpha", + logging_obj=logging, + ) + return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_cohere_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + extra_headers = ctx.extra_headers + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + cohere_key = ( + api_key + or litellm.cohere_key + or get_secret_str("COHERE_API_KEY") + or get_secret_str("CO_API_KEY") + or litellm.api_key + ) + + cohere_route = CohereModelInfo.get_cohere_route(model) + verbose_logger.debug(f"Cohere route: {cohere_route}") + # Set API base based on route + if cohere_route == "v2": + api_base = ( + api_base + or litellm.api_base + or get_secret_str("COHERE_API_BASE") + or "https://api.cohere.com/v2/chat" + ) + # Remove v2/ prefix from model name for the actual API call + if "v2/" in model: + model = model.replace("v2/", "") + else: + api_base = ( + api_base + or litellm.api_base + or get_secret_str("COHERE_API_BASE") + or "https://api.cohere.ai/v1/chat" + ) + + headers = headers or litellm.headers or {} + if headers is None: + headers = {} + + if extra_headers is not None: + headers.update(extra_headers) + + verbose_logger.debug(f"Model: {model}, API Base: {api_base}") + verbose_logger.debug(f"Provider Config: {provider_config}") + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="cohere_chat", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=cohere_key, + provider_config=provider_config, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + ) + + +def _complete_maritalk(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base = ctx.api_base + api_key = ctx.api_key + custom_prompt_dict = ctx.custom_prompt_dict + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + maritalk_key = ( + api_key + or litellm.maritalk_key + or get_secret("MARITALK_API_KEY") + or litellm.api_key + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret("MARITALK_API_BASE") + or "https://chat.maritaca.ai/api" + ) + + return openai_like_chat_completion.completion( + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + api_key=maritalk_key, + logging_obj=logging, + custom_llm_provider="maritalk", + custom_prompt_dict=custom_prompt_dict, + ) + + +def _complete_amazon_nova(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base = ctx.api_base + api_key = ctx.api_key + custom_llm_provider = ctx.custom_llm_provider + custom_prompt_dict = ctx.custom_prompt_dict + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.amazon_nova_api_key + or get_secret_str("AMAZON_NOVA_API_KEY") + or litellm.api_key + ) + api_base = ( + api_base + or litellm.api_base + or get_secret_str("AMAZON_NOVA_API_BASE") + or "https://api.nova.amazon.com/v1" + ) + return openai_like_chat_completion.completion( + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + custom_prompt_dict=custom_prompt_dict, + ) + + +def _complete_huggingface(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + huggingface_key = ( + api_key + or litellm.huggingface_key + or os.environ.get("HF_TOKEN") + or os.environ.get("HUGGINGFACE_API_KEY") + or litellm.api_key + ) + hf_headers = headers or litellm.headers + return base_llm_http_handler.completion( + model=model, + messages=messages, + headers=hf_headers, + model_response=model_response, + api_key=huggingface_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + ) + + +def _complete_oci(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + return base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + ) + + +def _complete_compactifai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + stream = ctx.stream + timeout = ctx.timeout + + api_key = api_key or get_secret_str("COMPACTIFAI_API_KEY") or litellm.api_key + + api_base = api_base or "https://api.compactif.ai/v1" + + ## COMPLETION CALL + return base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + + +def _complete_oobabooga(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base = ctx.api_base + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + model_response = oobabooga.completion( + model=model, + messages=messages, + model_response=model_response, + api_base=api_base, # type: ignore + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=None, + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + ) + if "stream" in optional_params and optional_params["stream"] is True: + # don't try to access stream object, + return CustomStreamWrapper( + model_response, + model, + custom_llm_provider="oobabooga", + logging_obj=logging, + ) + return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_databricks(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base # for databricks we check in get_llm_provider and pass in the api base from there + or litellm.api_base + or os.getenv("DATABRICKS_API_BASE") + ) + + # set API KEY + api_key = ( + api_key + or litellm.api_key # for databricks we check in get_llm_provider and pass in the api key from there + or litellm.databricks_key + or get_secret("DATABRICKS_API_KEY") + ) + + headers = headers or litellm.headers + + ## COMPLETION CALL + try: + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider="databricks", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + if optional_params.get("stream", False): + ## LOGGING + logging.post_call( + input=messages, + api_key=api_key, + original_response=response, + additional_args={"headers": headers}, + ) + + return response + + +def _complete_datarobot(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + stream = ctx.stream + timeout = ctx.timeout + + return base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + + +def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base + or litellm.api_base + or get_secret_str("OPENROUTER_API_BASE") + or "https://openrouter.ai/api/v1" + ) + + api_key = ( + api_key + or litellm.api_key + or litellm.openrouter_key + or get_secret_str("OPENROUTER_API_KEY") + or get_secret_str("OR_API_KEY") + ) + + openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai" + openrouter_app_name = get_secret("OR_APP_NAME") or "liteLLM" + + openrouter_headers = { + "HTTP-Referer": openrouter_site_url, + "X-Title": openrouter_app_name, + } + + _headers = headers or litellm.headers + if _headers: + openrouter_headers.update(_headers) + + headers = openrouter_headers + + ## Load Config + config = litellm.OpenrouterConfig.get_config() + for k, v in config.items(): + if k == "extra_body": + # we use openai 'extra_body' to pass openrouter specific params - transforms, route, models + if "extra_body" in optional_params: + optional_params[k].update(v) + else: + optional_params[k] = v + elif k not in optional_params: + optional_params[k] = v + + ## COMPLETION CALL + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="openrouter", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + ## LOGGING + logging.post_call( + input=messages, api_key=openai.api_key, original_response=response + ) + + return response + + +def _complete_vercel_ai_gateway( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base + or litellm.api_base + or get_secret_str("VERCEL_AI_GATEWAY_API_BASE") + or "https://ai-gateway.vercel.sh/v1" + ) + + api_key = api_key or litellm.api_key or get_secret("VERCEL_AI_GATEWAY_API_KEY") + + vercel_site_url = get_secret("VERCEL_SITE_URL") or "https://litellm.ai" + vercel_app_name = get_secret("VERCEL_APP_NAME") or "liteLLM" + + vercel_headers = { + "http-referer": vercel_site_url, + "x-title": vercel_app_name, + } + + _headers = headers or litellm.headers + if _headers: + vercel_headers.update(_headers) + + headers = vercel_headers + + ## Load Config + config = litellm.VercelAIGatewayConfig.get_config() + for k, v in config.items(): + if k == "extra_body": + # we use openai 'extra_body' to pass vercel specific params - providerOptions + if "extra_body" in optional_params: + optional_params[k].update(v) + else: + optional_params[k] = v + elif k not in optional_params: + optional_params[k] = v + + ## COMPLETION CALL + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="vercel_ai_gateway", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + ## LOGGING + logging.post_call( + input=messages, api_key=openai.api_key, original_response=response + ) + + return response + + +def _complete_vertex_ai_beta( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + vertex_ai_project = ( + optional_params.pop("vertex_project", None) + or optional_params.pop("vertex_ai_project", None) + or litellm.vertex_project + or get_secret("VERTEXAI_PROJECT") + ) + vertex_ai_location = ( + optional_params.pop("vertex_location", None) + or optional_params.pop("vertex_ai_location", None) + or litellm.vertex_location + or get_secret("VERTEXAI_LOCATION") + ) + vertex_credentials = ( + optional_params.pop("vertex_credentials", None) + or optional_params.pop("vertex_ai_credentials", None) + or get_secret("VERTEXAI_CREDENTIALS") + ) + + gemini_api_key = ( + api_key + or get_api_key_from_env() + or get_secret("PALM_API_KEY") # older palm api key should also work + or litellm.api_key + ) + + api_base = api_base or litellm.api_base or get_secret("GEMINI_API_BASE") + new_params = safe_deep_copy(optional_params or {}) + return vertex_chat_completion.completion( # type: ignore + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=_get_encoding(), + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + gemini_api_key=gemini_api_key, + logging_obj=logging, + acompletion=acompletion, + timeout=timeout, + custom_llm_provider=custom_llm_provider, # type: ignore + client=client, + api_base=api_base, + extra_headers=headers, + ) + + +def _complete_vertex_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + vertex_ai_project = ( + optional_params.pop("vertex_project", None) + or optional_params.pop("vertex_ai_project", None) + or litellm.vertex_project + or get_secret("VERTEXAI_PROJECT") + ) + vertex_ai_location = ( + optional_params.pop("vertex_location", None) + or optional_params.pop("vertex_ai_location", None) + or litellm.vertex_location + or get_secret("VERTEXAI_LOCATION") + ) + vertex_credentials = ( + optional_params.pop("vertex_credentials", None) + or optional_params.pop("vertex_ai_credentials", None) + or get_secret("VERTEXAI_CREDENTIALS") + ) + + api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE") + + new_params = safe_deep_copy(optional_params or {}) + model_route = get_vertex_ai_model_route(model=model, litellm_params=litellm_params) + + if model_route == VertexAIModelRoute.PARTNER_MODELS: + model_response = vertex_partner_models_chat_completion.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=_get_encoding(), + api_base=api_base, + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + logging_obj=logging, + acompletion=acompletion, + headers=headers, + custom_prompt_dict=custom_prompt_dict, + timeout=timeout, + client=client, + ) + elif model_route == VertexAIModelRoute.GEMINI: + model_response = vertex_chat_completion.completion( # type: ignore + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=_get_encoding(), + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + gemini_api_key=None, + logging_obj=logging, + acompletion=acompletion, + timeout=timeout, + custom_llm_provider=custom_llm_provider, # type: ignore + client=client, + api_base=api_base, + extra_headers=headers, + ) + elif model_route == VertexAIModelRoute.GEMMA: + # Vertex Gemma Models with custom prediction endpoint + model_response = vertex_gemma_chat_completion.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=_get_encoding(), + api_base=api_base, + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + logging_obj=logging, + acompletion=acompletion, + headers=headers, + custom_prompt_dict=custom_prompt_dict, + timeout=timeout, + client=client, + ) + elif model_route == VertexAIModelRoute.MODEL_GARDEN: + # Vertex Model Garden - OpenAI compatible models + model_response = vertex_model_garden_chat_completion.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=_get_encoding(), + api_base=api_base, + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + logging_obj=logging, + acompletion=acompletion, + headers=headers, + custom_prompt_dict=custom_prompt_dict, + timeout=timeout, + client=client, + ) + elif model_route == VertexAIModelRoute.AGENT_ENGINE: + # Vertex AI Agent Engine (Reasoning Engines) + from litellm.llms.vertex_ai.agent_engine.transformation import ( + VertexAgentEngineConfig, + ) + + vertex_agent_engine_config = VertexAgentEngineConfig() + + # Update litellm_params with vertex credentials + litellm_params["vertex_project"] = vertex_ai_project + litellm_params["vertex_location"] = vertex_ai_location + litellm_params["vertex_credentials"] = vertex_credentials + + model_response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + model_response=model_response, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + encoding=_get_encoding(), + api_key=None, + api_base=api_base, + logging_obj=logging, + acompletion=acompletion, + timeout=timeout, + client=client, + custom_llm_provider="vertex_ai", + provider_config=vertex_agent_engine_config, + headers=headers or {}, + ) + else: # VertexAIModelRoute.NON_GEMINI + model_response = vertex_ai_non_gemini.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + logging_obj=logging, + acompletion=acompletion, + ) + + if ( + "stream" in optional_params + and optional_params["stream"] is True + and acompletion is False + ): + return CustomStreamWrapper( + model_response, + model, + custom_llm_provider="vertex_ai", + logging_obj=logging, + ) + return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_predibase(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + custom_prompt_dict = ctx.custom_prompt_dict + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + tenant_id = ( + optional_params.pop("tenant_id", None) + or optional_params.pop("predibase_tenant_id", None) + or litellm.predibase_tenant_id + or get_secret("PREDIBASE_TENANT_ID") + ) + + if tenant_id is None: + raise ValueError( + "Missing Predibase Tenant ID - Required for making the request. Set dynamically (e.g. `completion(..tenant_id=)`) or in env - `PREDIBASE_TENANT_ID`." + ) + + api_base = ( + api_base + or optional_params.pop("api_base", None) + or optional_params.pop("base_url", None) + or litellm.api_base + or get_secret("PREDIBASE_API_BASE") + ) + + api_key = ( + api_key + or litellm.api_key + or litellm.predibase_key + or get_secret("PREDIBASE_API_KEY") + ) + + _model_response = predibase_chat_completions.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + acompletion=acompletion, + api_base=api_base, + custom_prompt_dict=custom_prompt_dict, + api_key=api_key, + tenant_id=tenant_id, + timeout=timeout, + ) + + if ( + "stream" in optional_params + and optional_params["stream"] is True + and acompletion is False + ): + return _model_response + return _model_response + + +def _complete_text_completion_codestral( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + custom_prompt_dict = ctx.custom_prompt_dict + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + api_base + or optional_params.pop("api_base", None) + or optional_params.pop("base_url", None) + or litellm.api_base + or "https://codestral.mistral.ai/v1/fim/completions" + ) + + api_key = api_key or litellm.api_key or get_secret("CODESTRAL_API_KEY") + + text_completion_model_response = litellm.TextCompletionResponse(stream=stream) + + _model_response = codestral_text_completions.completion( # type: ignore + model=model, + messages=messages, + model_response=text_completion_model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + acompletion=acompletion, + api_base=api_base, + custom_prompt_dict=custom_prompt_dict, + api_key=api_key, + timeout=timeout, + ) + + if ( + "stream" in optional_params + and optional_params["stream"] is True + and acompletion is False + ): + return _model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + return _model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_text_completion_inception( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + text_completion = ctx.text_completion + timeout = ctx.timeout + + passed_api_base = ( + api_base + or optional_params.pop("api_base", None) + or optional_params.pop("base_url", None) + ) + api_base = ( + passed_api_base + or get_secret_str("INCEPTION_API_BASE") + or "https://api.inceptionlabs.ai/v1" + ) + # FIM is served at `/v1/fim/completions`; the OpenAI client appends + # `/completions`, so point it at the `/v1/fim` base. + api_base = api_base.rstrip("/") + if not api_base.endswith("/fim"): + api_base += "/fim" + + # Don't forward the server-managed Inception key to a caller-supplied + # api_base; only resolve it for the default/server base, or when the + # caller passes their own key. + if passed_api_base is None or api_key: + api_key = ( + api_key or litellm.inception_key or get_secret_str("INCEPTION_API_KEY") + ) + + _response = openai_text_completions.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + api_key=api_key, # type: ignore[arg-type] + custom_llm_provider="text-completion-inception", + api_base=api_base, + acompletion=acompletion, + client=client, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + timeout=timeout, # type: ignore + ) + + if ( + optional_params.get("stream", False) is False + and acompletion is False + and text_completion is False + ): + _response = ( + litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object( + response_object=_response, model_response_object=model_response + ) + ) + + if optional_params.get("stream", False) or acompletion is True: + logging.post_call( + input=messages, + api_key=api_key, + original_response=_response, + additional_args={"headers": headers}, + ) + return _response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_sagemaker_chat( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + + +def _complete_sagemaker(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + custom_prompt_dict = ctx.custom_prompt_dict + hf_model_name = ctx.hf_model_name + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + return sagemaker_llm.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + custom_prompt_dict=custom_prompt_dict, + hf_model_name=hf_model_name, + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + acompletion=acompletion, + ) + + +def _complete_bedrock(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + provider_config = ctx.provider_config + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict + + if "aws_bedrock_client" in optional_params: + verbose_logger.warning( + "'aws_bedrock_client' is a deprecated param. Please move to another auth method - https://docs.litellm.ai/docs/providers/bedrock#boto3---authentication." + ) + # Extract credentials for legacy boto3 client and pass thru to httpx + aws_bedrock_client = optional_params.pop("aws_bedrock_client") + creds = aws_bedrock_client._get_credentials().get_frozen_credentials() + + if creds.access_key: + optional_params["aws_access_key_id"] = creds.access_key + if creds.secret_key: + optional_params["aws_secret_access_key"] = creds.secret_key + if creds.token: + optional_params["aws_session_token"] = creds.token + if ( + "aws_region_name" not in optional_params + or optional_params["aws_region_name"] is None + ): + optional_params["aws_region_name"] = aws_bedrock_client.meta.region_name + + bedrock_route = BedrockModelInfo.get_bedrock_route(model) + if bedrock_route == "claude_platform": + provider_config = ProviderConfigManager.get_provider_chat_config( + model=model, + provider=LlmProviders.BEDROCK, + ) + model = BedrockModelInfo.get_claude_platform_model(model) + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="bedrock", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + provider_config=provider_config, + ) + elif bedrock_route == "converse": + model = model.replace("converse/", "") + response = bedrock_converse_chat_completion.completion( + model=model, + messages=messages, + custom_prompt_dict=custom_prompt_dict, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + extra_headers=headers, # Use merged headers instead of original extra_headers + timeout=timeout, + acompletion=acompletion, + client=client, + api_base=api_base, + api_key=api_key, + ) + elif bedrock_route == "converse_like": + model = model.replace("converse_like/", "") + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider="bedrock", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + else: + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider="bedrock", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + ) + + return response + + +def _complete_watsonx(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + timeout = ctx.timeout + + return watsonx_chat_completion.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + print_verbose=print_verbose, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + timeout=timeout, # type: ignore + custom_prompt_dict=custom_prompt_dict, + client=client, # pass AsyncOpenAI, OpenAI client + encoding=_get_encoding(), + custom_llm_provider="watsonx", + ) + + +def _complete_watsonx_text( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or optional_params.pop("apikey", None) + or get_secret_str("WATSONX_APIKEY") + or get_secret_str("WATSONX_API_KEY") + or get_secret_str("WX_API_KEY") + ) + + api_base = ( + api_base + or optional_params.pop( + "url", + optional_params.pop("api_base", optional_params.pop("base_url", None)), + ) + or get_secret_str("WATSONX_API_BASE") + or get_secret_str("WATSONX_URL") + or get_secret_str("WX_URL") + or get_secret_str("WML_URL") + ) + + wx_credentials = optional_params.pop( + "wx_credentials", + optional_params.pop( + "watsonx_credentials", None + ), # follow {provider}_credentials, same as vertex ai + ) + + token: Optional[str] = None + if wx_credentials is not None: + api_base = wx_credentials.get("url", api_base) + api_key = wx_credentials.get("apikey", wx_credentials.get("api_key", api_key)) + token = wx_credentials.get( + "token", + wx_credentials.get( + "watsonx_token", None + ), # follow format of {provider}_token, same as azure - e.g. 'azure_ad_token=..' + ) + + if token is not None: + optional_params["token"] = token + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="watsonx_text", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + + +def _complete_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + custom_prompt_dict = ctx.custom_prompt_dict + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + + custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict + model_response = vllm_handler.completion( + model=model, + messages=messages, + custom_prompt_dict=custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + ) + + if "stream" in optional_params and optional_params["stream"] is True: ## [BETA] + # don't try to access stream object, + return CustomStreamWrapper( + model_response, + model, + custom_llm_provider="vllm", + logging_obj=logging, + ) + + ## RESPONSE OBJECT + return model_response + + +def _complete_ollama(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + litellm.api_base + or api_base + or get_secret("OLLAMA_API_BASE") + or "http://localhost:11434" + ) + if api_key is not None and "Authorization" not in headers: + headers["Authorization"] = f"Bearer {api_key}" + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="ollama", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + + +def _complete_ollama_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = ( + litellm.api_base + or api_base + or get_secret("OLLAMA_API_BASE") + or "http://localhost:11434" + ) + + api_key = ( + api_key + or litellm.ollama_key + or os.environ.get("OLLAMA_API_KEY") + or litellm.api_key + ) + if api_key is not None and "Authorization" not in headers: + headers["Authorization"] = f"Bearer {api_key}" + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="ollama_chat", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + client=client, + ) + + +def _complete_triton(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = litellm.api_base or api_base + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + ) + + +def _complete_cloudflare(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.cloudflare_api_key + or litellm.api_key + or get_secret("CLOUDFLARE_API_KEY") + ) + account_id = get_secret("CLOUDFLARE_ACCOUNT_ID") + api_base = ( + api_base + or litellm.api_base + or get_secret("CLOUDFLARE_API_BASE") + or f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/" + ) + + custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="cloudflare", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements + ) + + +def _complete_petals(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base = ctx.api_base + client = ctx.client + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + + api_base = api_base or litellm.api_base + + stream = optional_params.pop("stream", False) + model_response = petals_handler.completion( + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=_get_encoding(), + logging_obj=logging, + client=client, + ) + if stream is True: ## [BETA] + # Fake streaming for petals + resp_string = model_response["choices"][0]["message"]["content"] + return CustomStreamWrapper( + resp_string, + model, + custom_llm_provider="petals", + logging_obj=logging, + ) + return model_response + + +def _complete_snowflake(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + try: + client = ( + HTTPHandler(timeout=timeout) if stream is False else None + ) # Keep this here, otherwise, the httpx.client closes and streaming is impossible + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + ) + + except Exception as e: + ## LOGGING - log the original exception returned + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args={"headers": headers}, + ) + raise e + + return response + + +def _complete_gradient_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + api_base = litellm.api_base or api_base + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider="gradient_ai", + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + ) + + +def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.bytez_key + or get_secret_str("BYTEZ_API_KEY") + or litellm.api_key + ) + + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=bytez_transformation, + ) + + pass + + return response + + +def _complete_lemonade(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.lemonade_key + or get_secret_str("LEMONADE_API_KEY") + or litellm.api_key + ) + + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=lemonade_transformation, + ) + + pass + + return response + + +def _complete_ovhcloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + api_key = ( + api_key + or litellm.ovhcloud_key + or get_secret_str("OVHCLOUD_API_KEY") + or litellm.api_key + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret_str("OVHCLOUD_API_BASE") + or "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1" + ) + + response = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + timeout=timeout, # type: ignore + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=ovhcloud_transformation, + ) + + pass + + return response + + +def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base = ctx.api_base + headers = ctx.headers + kwargs = ctx.kwargs + max_tokens = ctx.max_tokens + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + temperature = ctx.temperature + top_p = ctx.top_p + + url = litellm.api_base or api_base or "" + if url is None or url == "": + raise ValueError( + "api_base not set. Set api_base or litellm.api_base for custom endpoints" + ) + + """ + assume input to custom LLM api bases follow this format: + resp = litellm.module_level_client.post( + api_base, + json={ + 'model': 'meta-llama/Llama-2-13b-hf', # model name + 'params': { + 'prompt': ["The capital of France is P"], + 'max_tokens': 32, + 'temperature': 0.7, + 'top_p': 1.0, + 'top_k': 40, + } + } + ) + + """ + prompt = " ".join([message["content"] for message in messages]) # type: ignore + resp = litellm.module_level_client.post( + url, + headers=headers, + json={ + "model": model, + "params": { + "prompt": [prompt], + "max_tokens": max_tokens, + "temperature": temperature, + "top_p": top_p, + "top_k": kwargs.get("top_k"), + }, + **kwargs.get("extra_body", {}), + }, + ) + response_json = resp.json() + """ + assume all responses from custom api_bases of this format: + { + 'data': [ + { + 'prompt': 'The capital of France is P', + 'output': ['The capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France'], + 'params': {'temperature': 0.7, 'top_k': 40, 'top_p': 1}}], + 'message': 'ok' + } + ] + } + """ + string_response = response_json["data"][0]["output"][0] + ## RESPONSE OBJECT + model_response.choices[0].message.content = string_response # type: ignore + model_response.created = int(time.time()) + model_response.model = model + return model_response + + +def _complete_custom_providers( + ctx: _CompletionDispatchContext, +) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + custom_prompt_dict = ctx.custom_prompt_dict + headers = ctx.headers + litellm_params = ctx.litellm_params + logger_fn = ctx.logger_fn + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + stream = ctx.stream + timeout = ctx.timeout + + custom_handler: Optional[CustomLLM] = None + for item in litellm.custom_provider_map: + if item["provider"] == custom_llm_provider: + custom_handler = item["custom_handler"] + + if custom_handler is None: + raise LiteLLMUnknownProvider( + model=model, custom_llm_provider=custom_llm_provider + ) + + ## ROUTE LLM CALL ## + handler_fn = custom_chat_llm_router( + async_fn=acompletion, stream=stream, custom_llm=custom_handler + ) + + headers = headers or litellm.headers or {} + + ## CALL FUNCTION + response = handler_fn( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + print_verbose=print_verbose, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + timeout=timeout, # type: ignore + custom_prompt_dict=custom_prompt_dict, + client=client, # pass AsyncOpenAI, OpenAI client + encoding=_get_encoding(), + ) + if stream is True: + return CustomStreamWrapper( + completion_stream=response, + model=model, + custom_llm_provider=custom_llm_provider, + logging_obj=logging, + ) + + return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract + + +def _complete_langgraph(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + from litellm.llms.langgraph.chat.transformation import LangGraphConfig + + ( + api_base, + api_key, + ) = LangGraphConfig()._get_openai_compatible_provider_info( + api_base=api_base or litellm.api_base, + api_key=api_key or litellm.api_key, + ) + + headers = headers or litellm.headers + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + ) + + +def _complete_langflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion = ctx.acompletion + api_base = ctx.api_base + api_key = ctx.api_key + client = ctx.client + custom_llm_provider = ctx.custom_llm_provider + headers = ctx.headers + litellm_params = ctx.litellm_params + logging = ctx.logging + messages = ctx.messages + model = ctx.model + model_response = ctx.model_response + optional_params = ctx.optional_params + shared_session = ctx.shared_session + stream = ctx.stream + timeout = ctx.timeout + + from litellm.llms.langflow.chat.transformation import LangFlowConfig + + ( + api_base, + api_key, + ) = LangFlowConfig()._get_openai_compatible_provider_info( + api_base=api_base or litellm.api_base, + api_key=api_key or litellm.api_key, + ) + + headers = headers or litellm.headers + + return base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + ) + + @tracer.wrap() @client def completion( # type: ignore @@ -1746,375 +5575,52 @@ def completion( # type: ignore optional_params ) + _dispatch_ctx = _CompletionDispatchContext( + _azure_detection_model=_azure_detection_model, + acompletion=acompletion, + api_base=api_base, + api_key=api_key, + api_version=api_version, + client=client, + custom_llm_provider=custom_llm_provider, + custom_prompt_dict=custom_prompt_dict, + extra_headers=extra_headers, + headers=headers, + hf_model_name=hf_model_name, + kwargs=kwargs, + litellm_params=litellm_params, + logger_fn=logger_fn, + logging=logging, + max_retries=max_retries, + max_tokens=max_tokens, + messages=messages, + metadata=metadata, + model=model, + model_response=model_response, + optional_params=optional_params, + organization=organization, + provider_config=provider_config, + shared_session=shared_session, + stream=stream, + temperature=temperature, + text_completion=text_completion, + timeout=timeout, + top_p=top_p, + ) if custom_llm_provider == "azure": # azure configs ## check dynamic params ## - dynamic_params = False - if client is not None and ( - isinstance(client, openai.AzureOpenAI) - or isinstance(client, openai.AsyncAzureOpenAI) - ): - dynamic_params = _check_dynamic_azure_params( - azure_client_params={"api_version": api_version}, - azure_client=client, - ) - - api_type = get_secret("AZURE_API_TYPE") or "azure" - - api_base = api_base or litellm.api_base or get_secret("AZURE_API_BASE") - - api_version = ( - api_version - or litellm.api_version - or get_secret_str("AZURE_API_VERSION") - or litellm.AZURE_DEFAULT_API_VERSION - ) - - api_key = ( - api_key - or litellm.api_key - or litellm.azure_key - or get_secret_str("AZURE_OPENAI_API_KEY") - or get_secret_str("AZURE_API_KEY") - ) - - azure_ad_token = optional_params.get("extra_body", {}).pop( - "azure_ad_token", None - ) or get_secret_str("AZURE_AD_TOKEN") - - azure_ad_token_provider = litellm_params.get( - "azure_ad_token_provider", None - ) - - headers = headers or litellm.headers - - if extra_headers is not None: - optional_params["extra_headers"] = extra_headers - if max_retries is not None: - optional_params["max_retries"] = max_retries - - if litellm.AzureOpenAIO1Config().is_o_series_model( - model=_azure_detection_model - ): - ## LOAD CONFIG - if set - config = litellm.AzureOpenAIO1Config.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - response = azure_o1_chat_completions.completion( - model=model, - messages=messages, - headers=headers, - api_key=api_key, - api_base=api_base, - api_version=api_version, - dynamic_params=dynamic_params, - azure_ad_token=azure_ad_token, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - logging_obj=logging, - acompletion=acompletion, - timeout=timeout, # type: ignore - client=client, # pass AsyncAzureOpenAI, AzureOpenAI client - custom_llm_provider=custom_llm_provider, - ) - else: - ## LOAD CONFIG - if set - config = litellm.AzureOpenAIConfig.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - ## COMPLETION CALL - response = azure_chat_completions.completion( - model=model, - messages=messages, - headers=headers, - api_key=api_key, - api_base=api_base, - api_version=api_version, - api_type=api_type, - dynamic_params=dynamic_params, - azure_ad_token=azure_ad_token, - azure_ad_token_provider=azure_ad_token_provider, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - logging_obj=logging, - acompletion=acompletion, - timeout=timeout, # type: ignore - client=client, # pass AsyncAzureOpenAI, AzureOpenAI client - ) - - if optional_params.get("stream", False): - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - additional_args={ - "headers": headers, - "api_version": api_version, - "api_base": api_base, - }, - ) + response = _complete_azure(_dispatch_ctx) elif custom_llm_provider == "azure_text": # azure configs - api_type = get_secret_str("AZURE_API_TYPE") or "azure" - - api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") - - if api_base is None: - raise ValueError( - "api_base is required for Azure OpenAI LLM provider. Either set it dynamically or set the AZURE_API_BASE environment variable." - ) - - api_version = ( - api_version - or litellm.api_version - or get_secret_str("AZURE_API_VERSION") - ) - - api_key = ( - api_key - or litellm.api_key - or litellm.azure_key - or get_secret_str("AZURE_OPENAI_API_KEY") - or get_secret_str("AZURE_API_KEY") - ) - - azure_ad_token = optional_params.get("extra_body", {}).pop( - "azure_ad_token", None - ) or get_secret_str("AZURE_AD_TOKEN") - - azure_ad_token_provider = litellm_params.get( - "azure_ad_token_provider", None - ) - - headers = headers or litellm.headers - - if extra_headers is not None: - optional_params["extra_headers"] = extra_headers - - ## LOAD CONFIG - if set - config = litellm.AzureOpenAIConfig.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - ## COMPLETION CALL - response = azure_text_completions.completion( - model=model, - messages=messages, - headers=headers, - api_key=api_key, - api_base=api_base, - api_version=cast(str, api_version), - api_type=api_type, - azure_ad_token=azure_ad_token, - azure_ad_token_provider=azure_ad_token_provider, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - logging_obj=logging, - acompletion=acompletion, - timeout=timeout, - client=client, # pass AsyncAzureOpenAI, AzureOpenAI client - ) - - if optional_params.get("stream", False) or acompletion is True: - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - additional_args={ - "headers": headers, - "api_version": api_version, - "api_base": api_base, - }, - ) + response = _complete_azure_text(_dispatch_ctx) elif custom_llm_provider == "deepseek": ## COMPLETION CALL - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_deepseek(_dispatch_ctx) elif custom_llm_provider == "azure_ai": - from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo - - azure_ai_route = AzureFoundryModelInfo.get_azure_ai_route(model) - - # Check if this is an agents route - model format: azure_ai/agents/ - if azure_ai_route == "agents": - from litellm.llms.azure_ai.agents import AzureAIAgentsConfig - - api_base = AzureFoundryModelInfo.get_api_base(api_base) - if api_base is None: - raise ValueError( - "Azure AI Agents requests require an api_base. " - "Set `api_base` or the AZURE_AI_API_BASE env var." - ) - api_key = AzureFoundryModelInfo.get_api_key(api_key) - - response = AzureAIAgentsConfig.completion( - model=model, - messages=messages, - api_base=api_base, - api_key=api_key, - model_response=model_response, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, - acompletion=acompletion, - stream=stream, - headers=headers or litellm.headers, - ) - - # Check if this is a Claude model - route to Azure Anthropic handler - elif "claude" in model.lower(): - # Use Azure Anthropic handler for Claude models - api_base = AzureFoundryModelInfo.get_api_base(api_base) - if api_base is None: - raise ValueError( - "Azure Anthropic requests require an api_base. " - "Set `api_base` or the AZURE_AI_API_BASE env var." - ) - api_key = AzureFoundryModelInfo.get_api_key(api_key) - - # Ensure the URL ends with /v1/messages for Anthropic - if api_base: - api_base = api_base.rstrip("/") - if not api_base.endswith("/v1/messages"): - if "/anthropic" in api_base: - parts = api_base.split("/anthropic", 1) - api_base = parts[0] + "/anthropic" - else: - api_base = api_base + "/anthropic" - api_base = api_base + "/v1/messages" - - response = azure_anthropic_chat_completions.completion( - model=model, - messages=messages, - api_base=api_base, - acompletion=acompletion, - custom_prompt_dict=litellm.custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - headers=headers, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - ) - if optional_params.get("stream", False) or acompletion is True: - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - ) - response = response - else: - # Non-Claude models use standard Azure AI flow - api_base = AzureFoundryModelInfo.get_api_base(api_base) - # set API KEY - api_key = AzureFoundryModelInfo.get_api_key(api_key) - - headers = headers or litellm.headers - - if extra_headers is not None: - optional_params["extra_headers"] = extra_headers - - ## FOR COHERE - if "command-r" in model: # make sure tool call in messages are str - messages = stringify_json_tool_call_content(messages=messages) - - ## COMPLETION CALL - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, # type: ignore - client=client, # pass AsyncOpenAI, OpenAI client - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e - - if optional_params.get("stream", False): - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - additional_args={"headers": headers}, - ) + response = _complete_azure_ai(_dispatch_ctx) elif ( custom_llm_provider == "text-completion-openai" or "ft:babbage-002" in model @@ -2123,535 +5629,42 @@ def completion( # type: ignore in litellm.openai_text_completion_compatible_providers and kwargs.get("text_completion") is True ): - openai.api_type = "openai" - - api_base = ( - api_base - or litellm.api_base - or get_secret("OPENAI_BASE_URL") - or get_secret("OPENAI_API_BASE") - or "https://api.openai.com/v1" - ) - - openai.api_version = None - # set API KEY - - api_key = ( - api_key - or litellm.api_key - or litellm.openai_key - or get_secret("OPENAI_API_KEY") - ) - - headers = headers or litellm.headers - - ## LOAD CONFIG - if set - config = litellm.OpenAITextCompletionConfig.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > openai_text_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - if litellm.organization: - openai.organization = litellm.organization - - if ( - len(messages) > 0 - and "content" in messages[0] - and isinstance(messages[0]["content"], list) - ): - # text-davinci-003 can accept a string or array, if it's an array, assume the array is set in messages[0]['content'] - # https://platform.openai.com/docs/api-reference/completions/create - prompt = messages[0]["content"] - else: - prompt = " ".join([message["content"] for message in messages]) # type: ignore - - ## COMPLETION CALL - _response = openai_text_completions.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - print_verbose=print_verbose, - api_key=api_key, - custom_llm_provider=custom_llm_provider, - api_base=api_base, - acompletion=acompletion, - client=client, # pass AsyncOpenAI, OpenAI client - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - timeout=timeout, # type: ignore - ) - - if ( - optional_params.get("stream", False) is False - and acompletion is False - and text_completion is False - ): - # convert to chat completion response - _response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object( - response_object=_response, model_response_object=model_response - ) - - if optional_params.get("stream", False) or acompletion is True: - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=_response, - additional_args={"headers": headers}, - ) - response = _response + response = _complete_text_completion_openai(_dispatch_ctx) elif custom_llm_provider == "fireworks_ai": ## COMPLETION CALL - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_fireworks_ai(_dispatch_ctx) elif custom_llm_provider == "heroku": - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_heroku(_dispatch_ctx) elif custom_llm_provider == "ragflow": ## COMPLETION CALL - RAGFlow uses HTTP handler to support custom URL paths - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_ragflow(_dispatch_ctx) elif custom_llm_provider == "xai": ## COMPLETION CALL - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_xai(_dispatch_ctx) elif custom_llm_provider == "groq": - api_base = ( - api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there - or litellm.api_base - or get_secret("GROQ_API_BASE") - or "https://api.groq.com/openai/v1" - ) - - # set API KEY - api_key = ( - api_key - or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there - or litellm.groq_key - or get_secret("GROQ_API_KEY") - ) - - headers = headers or litellm.headers - - ## LOAD CONFIG - if set - config = litellm.GroqChatConfig.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) + response = _complete_groq(_dispatch_ctx) elif custom_llm_provider == "bedrock_mantle": - api_base = ( - api_base or litellm.api_base or get_secret("BEDROCK_MANTLE_API_BASE") - ) - api_key = api_key or litellm.api_key or get_secret("BEDROCK_MANTLE_API_KEY") - headers = headers or litellm.headers - config = litellm.BedrockMantleChatConfig.get_config() - for k, v in config.items(): - if k not in optional_params: - optional_params[k] = v - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - client=client, - ) + response = _complete_bedrock_mantle(_dispatch_ctx) elif custom_llm_provider == "a2a": # A2A (Agent-to-Agent) Protocol # Resolve agent configuration from registry if model format is "a2a/" - ( - api_base, - api_key, - headers, - ) = litellm.A2AConfig.resolve_agent_config_from_registry( - model=model, - api_base=api_base, - api_key=api_key, - headers=headers, - optional_params=optional_params, - ) - - # Fall back to environment variables and defaults - api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE") - - if api_base is None: - raise Exception( - "api_base is required for A2A provider. " - "Either provide api_base parameter, set A2A_API_BASE environment variable, " - "or register the agent in the proxy with model='a2a/'." - ) - - headers = headers or litellm.headers - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - client=client, - provider_config=provider_config, - ) + response = _complete_a2a(_dispatch_ctx) elif custom_llm_provider == "gigachat": # GigaChat - Sber AI's LLM (Russia) - api_key = ( - api_key - or litellm.api_key - or litellm.gigachat_key - or get_secret("GIGACHAT_API_KEY") - or get_secret("GIGACHAT_CREDENTIALS") - ) - - headers = headers or litellm.headers or {} - - ## COMPLETION CALL - try: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_gigachat(_dispatch_ctx) elif custom_llm_provider == "sap": - headers = headers or litellm.headers - ## LOAD CONFIG - if set - config = litellm.GenAIHubOrchestrationConfig.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - response = sap_gen_ai_hub_chat_completions.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - shared_session=shared_session, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - api_key=api_key, - api_base=api_base, - stream=stream, - ) + response = _complete_sap(_dispatch_ctx) elif custom_llm_provider == "aiohttp_openai": # NEW aiohttp provider for 10-100x higher RPS - api_base = ( - api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there - or litellm.api_base - or get_secret("OPENAI_BASE_URL") - or get_secret("OPENAI_API_BASE") - or "https://api.openai.com/v1" - ) - # set API KEY - api_key = ( - api_key - or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there - or litellm.openai_key - or get_secret("OPENAI_API_KEY") - ) - - headers = headers or litellm.headers - - if extra_headers is not None: - optional_params["extra_headers"] = extra_headers - response = base_llm_aiohttp_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - ) + response = _complete_aiohttp_openai(_dispatch_ctx) elif custom_llm_provider == "cometapi": - api_key = ( - api_key - or litellm.cometapi_key - or get_secret_str("COMETAPI_KEY") - or litellm.api_key - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret_str("COMETAPI_API_BASE") - or "https://api.cometapi.com/v1" - ) - - ## COMPLETION CALL - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - - ## LOGGING - logging.post_call( - input=messages, api_key=api_key, original_response=response - ) + response = _complete_cometapi(_dispatch_ctx) elif custom_llm_provider == "minimax": - api_key = api_key or get_secret_str("MINIMAX_API_KEY") or litellm.api_key - - api_base = ( - api_base - or litellm.api_base - or get_secret_str("MINIMAX_API_BASE") - or "https://api.minimax.io/v1" - ) - - response = base_llm_http_handler.completion( - model=model, - messages=messages, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - model_response=model_response, - encoding=_get_encoding(), - logging_obj=logging, - optional_params=optional_params, - timeout=timeout, - litellm_params=litellm_params, - shared_session=shared_session, - acompletion=acompletion, - stream=stream, - api_key=api_key, - headers=headers, - client=client, - provider_config=provider_config, - ) - logging.post_call( - input=messages, api_key=api_key, original_response=response - ) + response = _complete_minimax(_dispatch_ctx) elif custom_llm_provider == "hosted_vllm": - api_base = ( - api_base or litellm.api_base or get_secret_str("HOSTED_VLLM_API_BASE") - ) - - response = base_llm_http_handler.completion( - model=model, - messages=messages, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - model_response=model_response, - encoding=_get_encoding(), - logging_obj=logging, - optional_params=optional_params, - timeout=timeout, - litellm_params=litellm_params, - shared_session=shared_session, - acompletion=acompletion, - stream=stream, - api_key=api_key, - headers=headers, - client=client, - provider_config=provider_config, - ) - logging.post_call( - input=messages, api_key=api_key, original_response=response - ) + response = _complete_hosted_vllm(_dispatch_ctx) elif ( model in litellm.open_ai_chat_completion_models or custom_llm_provider == "custom_openai" @@ -2676,205 +5689,17 @@ def completion( # type: ignore ): # allow user to make an openai call with a custom base # note: if a user sets a custom base - we should ensure this works # allow for the setting of dynamic and stateful api-bases - api_base = ( - api_base # for deepinfra/perplexity/anyscale/groq/friendliai we check in get_llm_provider and pass in the api base from there - or litellm.api_base - or get_secret("OPENAI_BASE_URL") - or get_secret("OPENAI_API_BASE") - or "https://api.openai.com/v1" - ) - organization = ( - organization - or litellm.organization - or get_secret("OPENAI_ORGANIZATION") - or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105 - ) - openai.organization = organization - # set API KEY - api_key = ( - api_key - or litellm.api_key # for deepinfra/perplexity/anyscale/friendliai we check in get_llm_provider and pass in the api key from there - or litellm.openai_key - or get_secret("OPENAI_API_KEY") - ) - - headers = headers or litellm.headers - - # Add GitHub Copilot headers (same as /responses endpoint does) - if custom_llm_provider == "github_copilot": - from litellm.llms.github_copilot.authenticator import Authenticator - from litellm.llms.github_copilot.common_utils import ( - get_copilot_default_headers, - ) - - copilot_auth = Authenticator() - copilot_api_key = copilot_auth.get_api_key() - copilot_headers = get_copilot_default_headers(copilot_api_key) - if extra_headers: - copilot_headers.update(extra_headers) - extra_headers = copilot_headers - - if extra_headers is not None: - optional_params["extra_headers"] = extra_headers - - if ( - litellm.enable_preview_features and metadata is not None - ): # [PREVIEW] allow metadata to be passed to OPENAI - openai_metadata = get_requester_metadata(metadata) - if openai_metadata is not None: - optional_params["metadata"] = openai_metadata - - ## LOAD CONFIG - if set - config = litellm.OpenAIConfig.get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - ## COMPLETION CALL - use_base_llm_http_handler = get_secret_bool( - "EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER" - ) - - try: - if use_base_llm_http_handler: - response = base_llm_http_handler.completion( - model=model, - messages=messages, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - model_response=model_response, - encoding=_get_encoding(), - logging_obj=logging, - optional_params=optional_params, - timeout=timeout, - litellm_params=litellm_params, - shared_session=shared_session, - acompletion=acompletion, - stream=stream, - api_key=api_key, - headers=headers, - client=client, - provider_config=provider_config, - ) - else: - response = openai_chat_completions.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - print_verbose=print_verbose, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - timeout=timeout, # type: ignore - custom_prompt_dict=custom_prompt_dict, - client=client, # pass AsyncOpenAI, OpenAI client - organization=organization, - custom_llm_provider=custom_llm_provider, - shared_session=shared_session, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e - - if optional_params.get("stream", False): - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - additional_args={"headers": headers}, - ) + response = _complete_custom_openai(_dispatch_ctx) elif custom_llm_provider == "mistral": - api_key = api_key or litellm.api_key or get_secret("MISTRAL_API_KEY") - api_base = ( - api_base - or litellm.api_base - or get_secret("MISTRAL_API_BASE") - or "https://api.mistral.ai/v1" - ) - - response = base_llm_http_handler.completion( - model=model, - messages=messages, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - model_response=model_response, - encoding=_get_encoding(), - logging_obj=logging, - optional_params=optional_params, - timeout=timeout, - litellm_params=litellm_params, - shared_session=shared_session, - acompletion=acompletion, - stream=stream, - api_key=api_key, - headers=headers, - client=client, - provider_config=provider_config, - ) + response = _complete_mistral(_dispatch_ctx) elif ( "replicate" in model or custom_llm_provider == "replicate" or model in litellm.replicate_models ): # Setting the relevant API KEY for replicate, replicate defaults to using os.environ.get("REPLICATE_API_TOKEN") - replicate_key = ( - api_key - or litellm.replicate_key - or litellm.api_key - or get_secret("REPLICATE_API_KEY") - or get_secret("REPLICATE_API_TOKEN") - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret("REPLICATE_API_BASE") - or "https://api.replicate.com/v1" - ) - - custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - - model_response = replicate_chat_completion( # type: ignore - model=model, - messages=messages, - api_base=api_base, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), # for calculating input/output tokens - api_key=replicate_key, - logging_obj=logging, - custom_prompt_dict=custom_prompt_dict, - acompletion=acompletion, - headers=headers, - ) - - if optional_params.get("stream", False) is True: - ## LOGGING - logging.post_call( - input=messages, - api_key=replicate_key, - original_response=model_response, - ) - - response = model_response + response = _complete_replicate(_dispatch_ctx) elif ( "clarifai" in model or custom_llm_provider == "clarifai" @@ -2882,614 +5707,36 @@ def completion( # type: ignore ): pass # Deprecated - handled in the openai compatible provider section above elif custom_llm_provider == "anthropic_text": - api_key = ( - api_key - or litellm.anthropic_key - or litellm.api_key - or os.environ.get("ANTHROPIC_API_KEY") - ) - custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - api_base = ( - api_base - or litellm.api_base - or get_secret("ANTHROPIC_API_BASE") - or get_secret("ANTHROPIC_BASE_URL") - or "https://api.anthropic.com/v1/complete" - ) - - # Check if we should disable automatic URL suffix appending - disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX") - if ( - api_base is not None - and not disable_url_suffix - and not api_base.endswith("/v1/complete") - ): - api_base += "/v1/complete" - elif disable_url_suffix: - verbose_logger.debug( - "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/complete suffix" - ) - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="anthropic_text", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - ) + response = _complete_anthropic_text(_dispatch_ctx) elif custom_llm_provider == "anthropic": - api_key = ( - api_key - or litellm.anthropic_key - or litellm.api_key - or os.environ.get("ANTHROPIC_API_KEY") - ) - custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - # call /messages - # default route for all anthropic models - api_base = ( - api_base - or litellm.api_base - or get_secret("ANTHROPIC_API_BASE") - or get_secret("ANTHROPIC_BASE_URL") - or "https://api.anthropic.com/v1/messages" - ) - - # Check if we should disable automatic URL suffix appending - disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX") - if ( - api_base is not None - and not disable_url_suffix - and not api_base.endswith("/v1/messages") - ): - api_base += "/v1/messages" - elif disable_url_suffix: - verbose_logger.debug( - "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/messages suffix" - ) - - response = anthropic_chat_completions.completion( - model=model, - messages=messages, - api_base=api_base, - acompletion=acompletion, - custom_prompt_dict=litellm.custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), # for calculating input/output tokens - api_key=api_key, - logging_obj=logging, - headers=headers, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - ) - if optional_params.get("stream", False) or acompletion is True: - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - ) - response = response + response = _complete_anthropic(_dispatch_ctx) elif custom_llm_provider == "nlp_cloud": - nlp_cloud_key = ( - api_key - or litellm.nlp_cloud_key - or get_secret("NLP_CLOUD_API_KEY") - or litellm.api_key - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret("NLP_CLOUD_API_BASE") - or "https://api.nlpcloud.io/v1/gpu/" - ) - - response = nlp_cloud_chat_completion( - model=model, - messages=messages, - api_base=api_base, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - api_key=nlp_cloud_key, - logging_obj=logging, - ) - - if "stream" in optional_params and optional_params["stream"] is True: - # don't try to access stream object, - response = CustomStreamWrapper( - response, - model, - custom_llm_provider="nlp_cloud", - logging_obj=logging, - ) - - if optional_params.get("stream", False) or acompletion is True: - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - ) - - response = response + response = _complete_nlp_cloud(_dispatch_ctx) elif custom_llm_provider == "aleph_alpha": - aleph_alpha_key = ( - api_key - or litellm.aleph_alpha_key - or get_secret("ALEPH_ALPHA_API_KEY") - or get_secret("ALEPHALPHA_API_KEY") - or litellm.api_key - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret("ALEPH_ALPHA_API_BASE") - or "https://api.aleph-alpha.com/complete" - ) - - model_response = aleph_alpha.completion( - model=model, - messages=messages, - api_base=api_base, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - default_max_tokens_to_sample=litellm.max_tokens, - api_key=aleph_alpha_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - ) - - if "stream" in optional_params and optional_params["stream"] is True: - # don't try to access stream object, - response = CustomStreamWrapper( - model_response, - model, - custom_llm_provider="aleph_alpha", - logging_obj=logging, - ) - return response - response = model_response + response = _complete_aleph_alpha(_dispatch_ctx) elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere": - cohere_key = ( - api_key - or litellm.cohere_key - or get_secret_str("COHERE_API_KEY") - or get_secret_str("CO_API_KEY") - or litellm.api_key - ) - - cohere_route = CohereModelInfo.get_cohere_route(model) - verbose_logger.debug(f"Cohere route: {cohere_route}") - # Set API base based on route - if cohere_route == "v2": - api_base = ( - api_base - or litellm.api_base - or get_secret_str("COHERE_API_BASE") - or "https://api.cohere.com/v2/chat" - ) - # Remove v2/ prefix from model name for the actual API call - if "v2/" in model: - model = model.replace("v2/", "") - else: - api_base = ( - api_base - or litellm.api_base - or get_secret_str("COHERE_API_BASE") - or "https://api.cohere.ai/v1/chat" - ) - - headers = headers or litellm.headers or {} - if headers is None: - headers = {} - - if extra_headers is not None: - headers.update(extra_headers) - - verbose_logger.debug(f"Model: {model}, API Base: {api_base}") - verbose_logger.debug(f"Provider Config: {provider_config}") - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="cohere_chat", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=cohere_key, - provider_config=provider_config, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - ) + response = _complete_cohere_chat(_dispatch_ctx) elif custom_llm_provider == "maritalk": - maritalk_key = ( - api_key - or litellm.maritalk_key - or get_secret("MARITALK_API_KEY") - or litellm.api_key - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret("MARITALK_API_BASE") - or "https://chat.maritaca.ai/api" - ) - - model_response = openai_like_chat_completion.completion( - model=model, - messages=messages, - api_base=api_base, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - api_key=maritalk_key, - logging_obj=logging, - custom_llm_provider="maritalk", - custom_prompt_dict=custom_prompt_dict, - ) - - response = model_response + response = _complete_maritalk(_dispatch_ctx) elif custom_llm_provider == "amazon_nova": - api_key = ( - api_key - or litellm.amazon_nova_api_key - or get_secret_str("AMAZON_NOVA_API_KEY") - or litellm.api_key - ) - api_base = ( - api_base - or litellm.api_base - or get_secret_str("AMAZON_NOVA_API_BASE") - or "https://api.nova.amazon.com/v1" - ) - response = openai_like_chat_completion.completion( - model=model, - messages=messages, - api_base=api_base, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - timeout=timeout, - custom_llm_provider=custom_llm_provider, - custom_prompt_dict=custom_prompt_dict, - ) + response = _complete_amazon_nova(_dispatch_ctx) elif custom_llm_provider == "huggingface": - huggingface_key = ( - api_key - or litellm.huggingface_key - or os.environ.get("HF_TOKEN") - or os.environ.get("HUGGINGFACE_API_KEY") - or litellm.api_key - ) - hf_headers = headers or litellm.headers - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=hf_headers, - model_response=model_response, - api_key=huggingface_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - ) + response = _complete_huggingface(_dispatch_ctx) elif custom_llm_provider == "oci": - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - ) + response = _complete_oci(_dispatch_ctx) elif custom_llm_provider == "compactifai": - api_key = ( - api_key or get_secret_str("COMPACTIFAI_API_KEY") or litellm.api_key - ) - - api_base = api_base or "https://api.compactif.ai/v1" - - ## COMPLETION CALL - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) + response = _complete_compactifai(_dispatch_ctx) elif custom_llm_provider == "oobabooga": - custom_llm_provider = "oobabooga" - model_response = oobabooga.completion( - model=model, - messages=messages, - model_response=model_response, - api_base=api_base, # type: ignore - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - api_key=None, - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - ) - if "stream" in optional_params and optional_params["stream"] is True: - # don't try to access stream object, - response = CustomStreamWrapper( - model_response, - model, - custom_llm_provider="oobabooga", - logging_obj=logging, - ) - return response - response = model_response + response = _complete_oobabooga(_dispatch_ctx) elif custom_llm_provider == "databricks": - api_base = ( - api_base # for databricks we check in get_llm_provider and pass in the api base from there - or litellm.api_base - or os.getenv("DATABRICKS_API_BASE") - ) - - # set API KEY - api_key = ( - api_key - or litellm.api_key # for databricks we check in get_llm_provider and pass in the api key from there - or litellm.databricks_key - or get_secret("DATABRICKS_API_KEY") - ) - - headers = headers or litellm.headers - - ## COMPLETION CALL - try: - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - custom_llm_provider="databricks", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e - - if optional_params.get("stream", False): - ## LOGGING - logging.post_call( - input=messages, - api_key=api_key, - original_response=response, - additional_args={"headers": headers}, - ) + response = _complete_databricks(_dispatch_ctx) elif custom_llm_provider == "datarobot": - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) + response = _complete_datarobot(_dispatch_ctx) elif custom_llm_provider == "openrouter": - api_base = ( - api_base - or litellm.api_base - or get_secret_str("OPENROUTER_API_BASE") - or "https://openrouter.ai/api/v1" - ) - - api_key = ( - api_key - or litellm.api_key - or litellm.openrouter_key - or get_secret_str("OPENROUTER_API_KEY") - or get_secret_str("OR_API_KEY") - ) - - openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai" - openrouter_app_name = get_secret("OR_APP_NAME") or "liteLLM" - - openrouter_headers = { - "HTTP-Referer": openrouter_site_url, - "X-Title": openrouter_app_name, - } - - _headers = headers or litellm.headers - if _headers: - openrouter_headers.update(_headers) - - headers = openrouter_headers - - ## Load Config - config = litellm.OpenrouterConfig.get_config() - for k, v in config.items(): - if k == "extra_body": - # we use openai 'extra_body' to pass openrouter specific params - transforms, route, models - if "extra_body" in optional_params: - optional_params[k].update(v) - else: - optional_params[k] = v - elif k not in optional_params: - optional_params[k] = v - - data = {"model": model, "messages": messages, **optional_params} - - ## COMPLETION CALL - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="openrouter", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) - ## LOGGING - logging.post_call( - input=messages, api_key=openai.api_key, original_response=response - ) + response = _complete_openrouter(_dispatch_ctx) elif custom_llm_provider == "vercel_ai_gateway": - api_base = ( - api_base - or litellm.api_base - or get_secret_str("VERCEL_AI_GATEWAY_API_BASE") - or "https://ai-gateway.vercel.sh/v1" - ) - - api_key = ( - api_key or litellm.api_key or get_secret("VERCEL_AI_GATEWAY_API_KEY") - ) - - vercel_site_url = get_secret("VERCEL_SITE_URL") or "https://litellm.ai" - vercel_app_name = get_secret("VERCEL_APP_NAME") or "liteLLM" - - vercel_headers = { - "http-referer": vercel_site_url, - "x-title": vercel_app_name, - } - - _headers = headers or litellm.headers - if _headers: - vercel_headers.update(_headers) - - headers = vercel_headers - - ## Load Config - config = litellm.VercelAIGatewayConfig.get_config() - for k, v in config.items(): - if k == "extra_body": - # we use openai 'extra_body' to pass vercel specific params - providerOptions - if "extra_body" in optional_params: - optional_params[k].update(v) - else: - optional_params[k] = v - elif k not in optional_params: - optional_params[k] = v - - data = {"model": model, "messages": messages, **optional_params} - - ## COMPLETION CALL - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="vercel_ai_gateway", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) - ## LOGGING - logging.post_call( - input=messages, api_key=openai.api_key, original_response=response - ) + response = _complete_vercel_ai_gateway(_dispatch_ctx) elif ( custom_llm_provider == "together_ai" or ("togethercomputer" in model) @@ -3504,1114 +5751,75 @@ def completion( # type: ignore "Palm was decommisioned on October 2024. Please use the `gemini/` route for Gemini Google AI Studio Models. Announcement: https://ai.google.dev/palm_docs/palm?hl=en" ) elif custom_llm_provider == "vertex_ai_beta" or custom_llm_provider == "gemini": - vertex_ai_project = ( - optional_params.pop("vertex_project", None) - or optional_params.pop("vertex_ai_project", None) - or litellm.vertex_project - or get_secret("VERTEXAI_PROJECT") - ) - vertex_ai_location = ( - optional_params.pop("vertex_location", None) - or optional_params.pop("vertex_ai_location", None) - or litellm.vertex_location - or get_secret("VERTEXAI_LOCATION") - ) - vertex_credentials = ( - optional_params.pop("vertex_credentials", None) - or optional_params.pop("vertex_ai_credentials", None) - or get_secret("VERTEXAI_CREDENTIALS") - ) - - gemini_api_key = ( - api_key - or get_api_key_from_env() - or get_secret("PALM_API_KEY") # older palm api key should also work - or litellm.api_key - ) - - api_base = api_base or litellm.api_base or get_secret("GEMINI_API_BASE") - new_params = safe_deep_copy(optional_params or {}) - response = vertex_chat_completion.completion( # type: ignore - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=new_params, - litellm_params=litellm_params, # type: ignore - logger_fn=logger_fn, - encoding=_get_encoding(), - vertex_location=vertex_ai_location, - vertex_project=vertex_ai_project, - vertex_credentials=vertex_credentials, - gemini_api_key=gemini_api_key, - logging_obj=logging, - acompletion=acompletion, - timeout=timeout, - custom_llm_provider=custom_llm_provider, # type: ignore - client=client, - api_base=api_base, - extra_headers=headers, - ) + response = _complete_vertex_ai_beta(_dispatch_ctx) elif custom_llm_provider == "vertex_ai": - vertex_ai_project = ( - optional_params.pop("vertex_project", None) - or optional_params.pop("vertex_ai_project", None) - or litellm.vertex_project - or get_secret("VERTEXAI_PROJECT") - ) - vertex_ai_location = ( - optional_params.pop("vertex_location", None) - or optional_params.pop("vertex_ai_location", None) - or litellm.vertex_location - or get_secret("VERTEXAI_LOCATION") - ) - vertex_credentials = ( - optional_params.pop("vertex_credentials", None) - or optional_params.pop("vertex_ai_credentials", None) - or get_secret("VERTEXAI_CREDENTIALS") - ) - - api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE") - - new_params = safe_deep_copy(optional_params or {}) - model_route = get_vertex_ai_model_route( - model=model, litellm_params=litellm_params - ) - - if model_route == VertexAIModelRoute.PARTNER_MODELS: - model_response = vertex_partner_models_chat_completion.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=new_params, - litellm_params=litellm_params, # type: ignore - logger_fn=logger_fn, - encoding=_get_encoding(), - api_base=api_base, - vertex_location=vertex_ai_location, - vertex_project=vertex_ai_project, - vertex_credentials=vertex_credentials, - logging_obj=logging, - acompletion=acompletion, - headers=headers, - custom_prompt_dict=custom_prompt_dict, - timeout=timeout, - client=client, - ) - elif model_route == VertexAIModelRoute.GEMINI: - model_response = vertex_chat_completion.completion( # type: ignore - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=new_params, - litellm_params=litellm_params, # type: ignore - logger_fn=logger_fn, - encoding=_get_encoding(), - vertex_location=vertex_ai_location, - vertex_project=vertex_ai_project, - vertex_credentials=vertex_credentials, - gemini_api_key=None, - logging_obj=logging, - acompletion=acompletion, - timeout=timeout, - custom_llm_provider=custom_llm_provider, # type: ignore - client=client, - api_base=api_base, - extra_headers=headers, - ) - elif model_route == VertexAIModelRoute.GEMMA: - # Vertex Gemma Models with custom prediction endpoint - model_response = vertex_gemma_chat_completion.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=new_params, - litellm_params=litellm_params, # type: ignore - logger_fn=logger_fn, - encoding=_get_encoding(), - api_base=api_base, - vertex_location=vertex_ai_location, - vertex_project=vertex_ai_project, - vertex_credentials=vertex_credentials, - logging_obj=logging, - acompletion=acompletion, - headers=headers, - custom_prompt_dict=custom_prompt_dict, - timeout=timeout, - client=client, - ) - elif model_route == VertexAIModelRoute.MODEL_GARDEN: - # Vertex Model Garden - OpenAI compatible models - model_response = vertex_model_garden_chat_completion.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=new_params, - litellm_params=litellm_params, # type: ignore - logger_fn=logger_fn, - encoding=_get_encoding(), - api_base=api_base, - vertex_location=vertex_ai_location, - vertex_project=vertex_ai_project, - vertex_credentials=vertex_credentials, - logging_obj=logging, - acompletion=acompletion, - headers=headers, - custom_prompt_dict=custom_prompt_dict, - timeout=timeout, - client=client, - ) - elif model_route == VertexAIModelRoute.AGENT_ENGINE: - # Vertex AI Agent Engine (Reasoning Engines) - from litellm.llms.vertex_ai.agent_engine.transformation import ( - VertexAgentEngineConfig, - ) - - vertex_agent_engine_config = VertexAgentEngineConfig() - - # Update litellm_params with vertex credentials - litellm_params["vertex_project"] = vertex_ai_project - litellm_params["vertex_location"] = vertex_ai_location - litellm_params["vertex_credentials"] = vertex_credentials - - model_response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - model_response=model_response, - optional_params=new_params, - litellm_params=litellm_params, # type: ignore - encoding=_get_encoding(), - api_key=None, - api_base=api_base, - logging_obj=logging, - acompletion=acompletion, - timeout=timeout, - client=client, - custom_llm_provider="vertex_ai", - provider_config=vertex_agent_engine_config, - headers=headers or {}, - ) - else: # VertexAIModelRoute.NON_GEMINI - model_response = vertex_ai_non_gemini.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=new_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - vertex_location=vertex_ai_location, - vertex_project=vertex_ai_project, - vertex_credentials=vertex_credentials, - logging_obj=logging, - acompletion=acompletion, - ) - - if ( - "stream" in optional_params - and optional_params["stream"] is True - and acompletion is False - ): - response = CustomStreamWrapper( - model_response, - model, - custom_llm_provider="vertex_ai", - logging_obj=logging, - ) - return response - response = model_response + response = _complete_vertex_ai(_dispatch_ctx) elif custom_llm_provider == "predibase": - tenant_id = ( - optional_params.pop("tenant_id", None) - or optional_params.pop("predibase_tenant_id", None) - or litellm.predibase_tenant_id - or get_secret("PREDIBASE_TENANT_ID") - ) - - if tenant_id is None: - raise ValueError( - "Missing Predibase Tenant ID - Required for making the request. Set dynamically (e.g. `completion(..tenant_id=)`) or in env - `PREDIBASE_TENANT_ID`." - ) - - api_base = ( - api_base - or optional_params.pop("api_base", None) - or optional_params.pop("base_url", None) - or litellm.api_base - or get_secret("PREDIBASE_API_BASE") - ) - - api_key = ( - api_key - or litellm.api_key - or litellm.predibase_key - or get_secret("PREDIBASE_API_KEY") - ) - - _model_response = predibase_chat_completions.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - acompletion=acompletion, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - api_key=api_key, - tenant_id=tenant_id, - timeout=timeout, - ) - - if ( - "stream" in optional_params - and optional_params["stream"] is True - and acompletion is False - ): - return _model_response - response = _model_response + response = _complete_predibase(_dispatch_ctx) elif custom_llm_provider == "text-completion-codestral": - api_base = ( - api_base - or optional_params.pop("api_base", None) - or optional_params.pop("base_url", None) - or litellm.api_base - or "https://codestral.mistral.ai/v1/fim/completions" - ) - - api_key = api_key or litellm.api_key or get_secret("CODESTRAL_API_KEY") - - text_completion_model_response = litellm.TextCompletionResponse( - stream=stream - ) - - _model_response = codestral_text_completions.completion( # type: ignore - model=model, - messages=messages, - model_response=text_completion_model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - acompletion=acompletion, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - api_key=api_key, - timeout=timeout, - ) - - if ( - "stream" in optional_params - and optional_params["stream"] is True - and acompletion is False - ): - return _model_response - response = _model_response + response = _complete_text_completion_codestral(_dispatch_ctx) elif custom_llm_provider == "text-completion-inception": - passed_api_base = ( - api_base - or optional_params.pop("api_base", None) - or optional_params.pop("base_url", None) - ) - api_base = ( - passed_api_base - or get_secret_str("INCEPTION_API_BASE") - or "https://api.inceptionlabs.ai/v1" - ) - # FIM is served at `/v1/fim/completions`; the OpenAI client appends - # `/completions`, so point it at the `/v1/fim` base. - api_base = api_base.rstrip("/") - if not api_base.endswith("/fim"): - api_base += "/fim" - - # Don't forward the server-managed Inception key to a caller-supplied - # api_base; only resolve it for the default/server base, or when the - # caller passes their own key. - if passed_api_base is None or api_key: - api_key = ( - api_key - or litellm.inception_key - or get_secret_str("INCEPTION_API_KEY") - ) - - _response = openai_text_completions.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - api_key=api_key, # type: ignore[arg-type] - custom_llm_provider="text-completion-inception", - api_base=api_base, - acompletion=acompletion, - client=client, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - timeout=timeout, # type: ignore - ) - - if ( - optional_params.get("stream", False) is False - and acompletion is False - and text_completion is False - ): - _response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object( - response_object=_response, model_response_object=model_response - ) - - if optional_params.get("stream", False) or acompletion is True: - logging.post_call( - input=messages, - api_key=api_key, - original_response=_response, - additional_args={"headers": headers}, - ) - response = _response + response = _complete_text_completion_inception(_dispatch_ctx) elif custom_llm_provider in ("sagemaker_chat", "sagemaker_nova"): # boto3 reads keys from .env # sagemaker_chat: HF Messages API endpoints # sagemaker_nova: Nova models on SageMaker (OpenAI-compatible) - model_response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) - - ## RESPONSE OBJECT - response = model_response + response = _complete_sagemaker_chat(_dispatch_ctx) elif custom_llm_provider == "sagemaker": # boto3 reads keys from .env - model_response = sagemaker_llm.completion( - model=model, - messages=messages, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - custom_prompt_dict=custom_prompt_dict, - hf_model_name=hf_model_name, - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - acompletion=acompletion, - ) - - ## RESPONSE OBJECT - response = model_response + response = _complete_sagemaker(_dispatch_ctx) elif custom_llm_provider == "bedrock": # boto3 reads keys from .env - custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - - if "aws_bedrock_client" in optional_params: - verbose_logger.warning( - "'aws_bedrock_client' is a deprecated param. Please move to another auth method - https://docs.litellm.ai/docs/providers/bedrock#boto3---authentication." - ) - # Extract credentials for legacy boto3 client and pass thru to httpx - aws_bedrock_client = optional_params.pop("aws_bedrock_client") - creds = aws_bedrock_client._get_credentials().get_frozen_credentials() - - if creds.access_key: - optional_params["aws_access_key_id"] = creds.access_key - if creds.secret_key: - optional_params["aws_secret_access_key"] = creds.secret_key - if creds.token: - optional_params["aws_session_token"] = creds.token - if ( - "aws_region_name" not in optional_params - or optional_params["aws_region_name"] is None - ): - optional_params["aws_region_name"] = ( - aws_bedrock_client.meta.region_name - ) - - bedrock_route = BedrockModelInfo.get_bedrock_route(model) - if bedrock_route == "claude_platform": - provider_config = ProviderConfigManager.get_provider_chat_config( - model=model, - provider=LlmProviders.BEDROCK, - ) - model = BedrockModelInfo.get_claude_platform_model(model) - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="bedrock", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - client=client, - provider_config=provider_config, - ) - return response - elif bedrock_route == "converse": - model = model.replace("converse/", "") - response = bedrock_converse_chat_completion.completion( - model=model, - messages=messages, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, # type: ignore - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - extra_headers=headers, # Use merged headers instead of original extra_headers - timeout=timeout, - acompletion=acompletion, - client=client, - api_base=api_base, - api_key=api_key, - ) - elif bedrock_route == "converse_like": - model = model.replace("converse_like/", "") - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - custom_llm_provider="bedrock", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) - else: - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - custom_llm_provider="bedrock", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - client=client, - ) + response = _complete_bedrock(_dispatch_ctx) elif custom_llm_provider == "watsonx": - response = watsonx_chat_completion.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - print_verbose=print_verbose, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - timeout=timeout, # type: ignore - custom_prompt_dict=custom_prompt_dict, - client=client, # pass AsyncOpenAI, OpenAI client - encoding=_get_encoding(), - custom_llm_provider="watsonx", - ) + response = _complete_watsonx(_dispatch_ctx) elif custom_llm_provider == "watsonx_text": - api_key = ( - api_key - or optional_params.pop("apikey", None) - or get_secret_str("WATSONX_APIKEY") - or get_secret_str("WATSONX_API_KEY") - or get_secret_str("WX_API_KEY") - ) - - api_base = ( - api_base - or optional_params.pop( - "url", - optional_params.pop( - "api_base", optional_params.pop("base_url", None) - ), - ) - or get_secret_str("WATSONX_API_BASE") - or get_secret_str("WATSONX_URL") - or get_secret_str("WX_URL") - or get_secret_str("WML_URL") - ) - - wx_credentials = optional_params.pop( - "wx_credentials", - optional_params.pop( - "watsonx_credentials", None - ), # follow {provider}_credentials, same as vertex ai - ) - - token: Optional[str] = None - if wx_credentials is not None: - api_base = wx_credentials.get("url", api_base) - api_key = wx_credentials.get( - "apikey", wx_credentials.get("api_key", api_key) - ) - token = wx_credentials.get( - "token", - wx_credentials.get( - "watsonx_token", None - ), # follow format of {provider}_token, same as azure - e.g. 'azure_ad_token=..' - ) - - if token is not None: - optional_params["token"] = token - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="watsonx_text", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) + response = _complete_watsonx_text(_dispatch_ctx) elif custom_llm_provider == "vllm": - custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - model_response = vllm_handler.completion( - model=model, - messages=messages, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - ) - - if ( - "stream" in optional_params and optional_params["stream"] is True - ): ## [BETA] - # don't try to access stream object, - response = CustomStreamWrapper( - model_response, - model, - custom_llm_provider="vllm", - logging_obj=logging, - ) - return response - - ## RESPONSE OBJECT - response = model_response + response = _complete_vllm(_dispatch_ctx) elif custom_llm_provider == "ollama": - api_base = ( - litellm.api_base - or api_base - or get_secret("OLLAMA_API_BASE") - or "http://localhost:11434" - ) - if api_key is not None and "Authorization" not in headers: - headers["Authorization"] = f"Bearer {api_key}" - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="ollama", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) + response = _complete_ollama(_dispatch_ctx) elif custom_llm_provider == "ollama_chat": - api_base = ( - litellm.api_base - or api_base - or get_secret("OLLAMA_API_BASE") - or "http://localhost:11434" - ) - - api_key = ( - api_key - or litellm.ollama_key - or os.environ.get("OLLAMA_API_KEY") - or litellm.api_key - ) - if api_key is not None and "Authorization" not in headers: - headers["Authorization"] = f"Bearer {api_key}" - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="ollama_chat", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - client=client, - ) + response = _complete_ollama_chat(_dispatch_ctx) elif custom_llm_provider == "triton": - api_base = litellm.api_base or api_base - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - ) + response = _complete_triton(_dispatch_ctx) elif custom_llm_provider == "cloudflare": - api_key = ( - api_key - or litellm.cloudflare_api_key - or litellm.api_key - or get_secret("CLOUDFLARE_API_KEY") - ) - account_id = get_secret("CLOUDFLARE_ACCOUNT_ID") - api_base = ( - api_base - or litellm.api_base - or get_secret("CLOUDFLARE_API_BASE") - or f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/" - ) - - custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="cloudflare", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - ) + response = _complete_cloudflare(_dispatch_ctx) elif custom_llm_provider == "petals" or model in litellm.petals_models: - api_base = api_base or litellm.api_base - - custom_llm_provider = "petals" - stream = optional_params.pop("stream", False) - model_response = petals_handler.completion( - model=model, - messages=messages, - api_base=api_base, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=_get_encoding(), - logging_obj=logging, - client=client, - ) - if stream is True: ## [BETA] - # Fake streaming for petals - resp_string = model_response["choices"][0]["message"]["content"] - response = CustomStreamWrapper( - resp_string, - model, - custom_llm_provider="petals", - logging_obj=logging, - ) - return response - response = model_response + response = _complete_petals(_dispatch_ctx) elif custom_llm_provider == "snowflake" or model in litellm.snowflake_models: - try: - client = ( - HTTPHandler(timeout=timeout) if stream is False else None - ) # Keep this here, otherwise, the httpx.client closes and streaming is impossible - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - ) - - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e + response = _complete_snowflake(_dispatch_ctx) elif custom_llm_provider == "gradient_ai": - api_base = litellm.api_base or api_base - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider="gradient_ai", - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - ) + response = _complete_gradient_ai(_dispatch_ctx) elif custom_llm_provider == "bytez": - api_key = ( - api_key - or litellm.bytez_key - or get_secret_str("BYTEZ_API_KEY") - or litellm.api_key - ) - - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=bytez_transformation, - ) - - pass + response = _complete_bytez(_dispatch_ctx) elif custom_llm_provider == "lemonade": - api_key = ( - api_key - or litellm.lemonade_key - or get_secret_str("LEMONADE_API_KEY") - or litellm.api_key - ) - - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=lemonade_transformation, - ) - - pass + response = _complete_lemonade(_dispatch_ctx) elif custom_llm_provider == "ovhcloud" or model in litellm.ovhcloud_models: - api_key = ( - api_key - or litellm.ovhcloud_key - or get_secret_str("OVHCLOUD_API_KEY") - or litellm.api_key - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret_str("OVHCLOUD_API_BASE") - or "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1" - ) - - response = base_llm_http_handler.completion( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - timeout=timeout, # type: ignore - client=client, - custom_llm_provider=custom_llm_provider, - encoding=_get_encoding(), - stream=stream, - provider_config=ovhcloud_transformation, - ) - - pass + response = _complete_ovhcloud(_dispatch_ctx) elif custom_llm_provider == "custom": - url = litellm.api_base or api_base or "" - if url is None or url == "": - raise ValueError( - "api_base not set. Set api_base or litellm.api_base for custom endpoints" - ) - - """ - assume input to custom LLM api bases follow this format: - resp = litellm.module_level_client.post( - api_base, - json={ - 'model': 'meta-llama/Llama-2-13b-hf', # model name - 'params': { - 'prompt': ["The capital of France is P"], - 'max_tokens': 32, - 'temperature': 0.7, - 'top_p': 1.0, - 'top_k': 40, - } - } - ) - - """ - prompt = " ".join([message["content"] for message in messages]) # type: ignore - resp = litellm.module_level_client.post( - url, - headers=headers, - json={ - "model": model, - "params": { - "prompt": [prompt], - "max_tokens": max_tokens, - "temperature": temperature, - "top_p": top_p, - "top_k": kwargs.get("top_k"), - }, - **kwargs.get("extra_body", {}), - }, - ) - response_json = resp.json() - """ - assume all responses from custom api_bases of this format: - { - 'data': [ - { - 'prompt': 'The capital of France is P', - 'output': ['The capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France is PARIS.\nThe capital of France'], - 'params': {'temperature': 0.7, 'top_k': 40, 'top_p': 1}}], - 'message': 'ok' - } - ] - } - """ - string_response = response_json["data"][0]["output"][0] - ## RESPONSE OBJECT - model_response.choices[0].message.content = string_response # type: ignore - model_response.created = int(time.time()) - model_response.model = model - response = model_response + response = _complete_custom(_dispatch_ctx) elif ( custom_llm_provider in litellm._custom_providers ): # Assume custom LLM provider # Get the Custom Handler - custom_handler: Optional[CustomLLM] = None - for item in litellm.custom_provider_map: - if item["provider"] == custom_llm_provider: - custom_handler = item["custom_handler"] - - if custom_handler is None: - raise LiteLLMUnknownProvider( - model=model, custom_llm_provider=custom_llm_provider - ) - - ## ROUTE LLM CALL ## - handler_fn = custom_chat_llm_router( - async_fn=acompletion, stream=stream, custom_llm=custom_handler - ) - - headers = headers or litellm.headers or {} - - ## CALL FUNCTION - response = handler_fn( - model=model, - messages=messages, - headers=headers, - model_response=model_response, - print_verbose=print_verbose, - api_key=api_key, - api_base=api_base, - acompletion=acompletion, - logging_obj=logging, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - timeout=timeout, # type: ignore - custom_prompt_dict=custom_prompt_dict, - client=client, # pass AsyncOpenAI, OpenAI client - encoding=_get_encoding(), - ) - if stream is True: - return CustomStreamWrapper( - completion_stream=response, - model=model, - custom_llm_provider=custom_llm_provider, - logging_obj=logging, - ) + response = _complete_custom_providers(_dispatch_ctx) elif custom_llm_provider == "langgraph": # LangGraph - Agent Runtime Provider - from litellm.llms.langgraph.chat.transformation import LangGraphConfig - - ( - api_base, - api_key, - ) = LangGraphConfig()._get_openai_compatible_provider_info( - api_base=api_base or litellm.api_base, - api_key=api_key or litellm.api_key, - ) - - headers = headers or litellm.headers - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - client=client, - ) + response = _complete_langgraph(_dispatch_ctx) elif custom_llm_provider == "langflow": # LangFlow - Visual AI Agent Platform - from litellm.llms.langflow.chat.transformation import LangFlowConfig - - ( - api_base, - api_key, - ) = LangFlowConfig()._get_openai_compatible_provider_info( - api_base=api_base or litellm.api_base, - api_key=api_key or litellm.api_key, - ) - - headers = headers or litellm.headers - - response = base_llm_http_handler.completion( - model=model, - stream=stream, - messages=messages, - acompletion=acompletion, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - litellm_params=litellm_params, - shared_session=shared_session, - custom_llm_provider=custom_llm_provider, - timeout=timeout, - headers=headers, - encoding=_get_encoding(), - api_key=api_key, - logging_obj=logging, - client=client, - ) + response = _complete_langflow(_dispatch_ctx) else: raise LiteLLMUnknownProvider( diff --git a/litellm/types/completion.py b/litellm/types/completion.py index cb263914be8..a91f6234fad 100644 --- a/litellm/types/completion.py +++ b/litellm/types/completion.py @@ -1,8 +1,28 @@ -from typing import Iterable, List, Optional, Union +from __future__ import annotations + +from dataclasses import dataclass +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Coroutine, + Iterable, + List, + Optional, + Union, +) from pydantic import BaseModel, ConfigDict from typing_extensions import Literal, Required, TypedDict +if TYPE_CHECKING: + import httpx + from aiohttp import ClientSession + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm import BaseConfig + from litellm.utils import CustomStreamWrapper, ModelResponse + class ChatCompletionSystemMessageParam(TypedDict, total=False): content: Required[str] @@ -191,3 +211,44 @@ class CompletionRequest(BaseModel): model_list: Optional[List[str]] = None model_config = ConfigDict(protected_namespaces=(), extra="allow") + + +@dataclass(frozen=True, slots=True) +class _CompletionDispatchContext: + _azure_detection_model: str + acompletion: bool + api_base: Optional[str] + api_key: Optional[str] + api_version: Optional[str] + client: Any + custom_llm_provider: str + custom_prompt_dict: dict + extra_headers: Optional[dict] + headers: dict + hf_model_name: Optional[str] + kwargs: dict + litellm_params: dict + logger_fn: Optional[Callable] + logging: LiteLLMLoggingObj + max_retries: Optional[int] + max_tokens: Optional[int] + messages: list + metadata: Optional[dict] + model: str + model_response: ModelResponse + optional_params: dict + organization: Optional[str] + provider_config: Optional[BaseConfig] + shared_session: Optional[ClientSession] + stream: Optional[bool] + temperature: Optional[float] + text_completion: bool + timeout: Optional[Union[float, str, httpx.Timeout]] + top_p: Optional[float] + + +_CompletionDispatchResult = Union[ + Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]], + "ModelResponse", + "CustomStreamWrapper", +] diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 62ebdb559fc..ae46f020de1 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -300,7 +300,7 @@ "slack": 3 }, "RET504": { - "baseline": 709, + "baseline": 702, "slack": 20 }, "RUF010": { diff --git a/tests/test_litellm/types/test_completion.py b/tests/test_litellm/types/test_completion.py index f24b00df3fc..cd51913c5dd 100644 --- a/tests/test_litellm/types/test_completion.py +++ b/tests/test_litellm/types/test_completion.py @@ -8,9 +8,16 @@ Usage: pytest tests/test_litellm/types/test_completion.py -v """ +import dataclasses from typing import List -from litellm.types.completion import CompletionRequest, ChatCompletionMessageParam +import pytest + +from litellm.types.completion import ( + ChatCompletionMessageParam, + CompletionRequest, + _CompletionDispatchContext, +) def test_completion_request_messages_type_validation(): @@ -146,3 +153,55 @@ def test_completion_request_with_all_params(): assert request.presence_penalty == 0.0 assert request.stream is False assert request.n == 1 + + +def _build_dispatch_context() -> _CompletionDispatchContext: + return _CompletionDispatchContext( + _azure_detection_model="gpt-4o", + acompletion=False, + api_base=None, + api_key=None, + api_version=None, + client=None, + custom_llm_provider="openai", + custom_prompt_dict={}, + extra_headers=None, + headers={}, + hf_model_name=None, + kwargs={}, + litellm_params={}, + logger_fn=None, + logging=None, # type: ignore[arg-type] + max_retries=None, + max_tokens=None, + messages=[], + metadata=None, + model="gpt-4o", + model_response=None, # type: ignore[arg-type] + optional_params={}, + organization=None, + provider_config=None, + shared_session=None, + stream=None, + temperature=None, + text_completion=False, + timeout=None, + top_p=None, + ) + + +def test_dispatch_context_is_frozen(): + """A helper must not be able to re-route the call by rebinding a dispatch + input mid-flight; this pins the frozen invariant the dispatch shape relies on.""" + ctx = _build_dispatch_context() + with pytest.raises(dataclasses.FrozenInstanceError): + ctx.model = "claude-haiku-4-5" # type: ignore[misc] + with pytest.raises(dataclasses.FrozenInstanceError): + ctx.custom_llm_provider = "anthropic" # type: ignore[misc] + + +def test_dispatch_context_uses_slots(): + """slots=True keeps the per-call context lightweight (no per-instance __dict__).""" + ctx = _build_dispatch_context() + assert not hasattr(ctx, "__dict__") + assert hasattr(type(ctx), "__slots__") From 80c5a848719f27a972194bb04084be127fc53e42 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 23 Jun 2026 20:01:44 +0530 Subject: [PATCH 02/29] chore: litellm oss staging (#30968) * fix: correct amazon.titan-embed-text-v2 input price to $0.02/1M tokens (#29693) * fix: correct amazon.titan-embed-text-v2 input price to $0.02/1M tokens * test: scope local cost map env var with monkeypatch to avoid test pollution * fix(sensitive_data_masker): fully mask secrets at or below the reveal threshold (#30764) * fix(sensitive_data_masker): fully mask secrets at or below the reveal threshold _mask_value did partial reveal by showing the first visible_prefix and last visible_suffix characters, but for a value whose length was at or below visible_prefix + visible_suffix (8 by default) it returned the value verbatim. A value of exactly 8 chars fell through the length guard and computed masked_length == 0, reconstructing the original string with no mask characters; anything shorter hit the early return. Either way short credentials were emitted in plaintext. mask_dict routes real secrets through this path, so an 8-char-or-shorter redis password, api key, or token could be written to logs and the UI unmasked. The sibling helper mask_sensitive_keys already guards this case; _mask_value now does the same by fully masking any value at or below the threshold. * fix(sensitive_data_masker): add mask_short_values opt-out for truncation callers Fully masking short values is the right default for secret masking, but CooldownCache reuses the masker purely to truncate exception messages to the first 50 characters, and it relies on short messages being returned readable. Masking those blanked out short exception text and broke its tests. Add a mask_short_values flag (default True, secure) and have CooldownCache pass False so it keeps the truncation behavior, while every secret-masking caller still gets short values fully masked. * fix(mcp_debug): opt out of short-value masking to keep diagnostic token preview MCPDebug uses the masker to preview auth tokens in debug headers and documents that values of 10 chars or fewer are shown unchanged so token types stay distinguishable. Pass mask_short_values=False so that diagnostic behavior is preserved while secret maskers keep masking short values. * fix(mcp_debug): mask short auth values in debug headers instead of echoing them Earlier this masker opted out of short-value masking to keep a token preview, but that echoes short authorization and token values verbatim in debug response headers, which is the same leak this change is meant to close. Auth material should never be emitted in full, so mask short values here too; the first/last character preview still applies to longer tokens. Only CooldownCache keeps the opt-out, since it truncates exception text rather than masking secrets. * test(mcp_debug): assert masked short value preserves length * refactor(fireworks_ai): remove deprecated audio transcriptions endpoint (#30917) Fireworks AI deprecated audio inference on 2026-06-10 (https://docs.fireworks.ai/updates/changelog#audio-inference-and-image-generation-deprecation). Live API testing confirms the endpoint is already non-functional: a valid Fireworks API key receives HTTP 401 "Unauthorized" from api.fireworks.ai/inference/v1/audio/transcriptions for every request, regardless of payload. The audio-prod.api.fireworks.ai host referenced in the test suite returns 401 for every path; the entire host is decommissioned. Remove the dead FireworksAIAudioTranscriptionConfig class and every reference to it across the codebase: - Delete litellm/llms/fireworks_ai/audio_transcription/ directory (17-line config class that inherited from OpenAIWhisperAudioTranscriptionConfig) - Remove the Fireworks branch from ProviderConfigManager.get_provider_audio_transcription_config() in litellm/utils.py; update the stale comment in get_optional_params_transcription that referenced fireworks ai - Remove the FireworksAIAudioTranscriptionConfig entries from LLM_CONFIG_NAMES and _LLM_CONFIGS_IMPORT_MAP in litellm/_lazy_imports_registry.py - Remove the TYPE_CHECKING re-export in litellm/__init__.py - Remove the transcription branch in the fireworks_ai case of get_supported_openai_params() in litellm/litellm_core_utils/get_supported_openai_params.py - Remove the whisper-v3 and whisper-v3-turbo entries from model_prices_and_context_window.json and litellm/model_prices_and_context_window_backup.json (both had mode: audio_transcription and zero-cost pricing) - Remove the TestFireworksAIAudioTranscription test class and its imports from tests/llm_translation/test_fireworks_ai_translation.py No other provider is affected. The openai_compatible_providers list, FireworksAIMixin, and the OpenAI Whisper transcription handler all stay because they are shared with other Fireworks endpoints and other providers. The provider_endpoints_support.json registry already had audio_transcriptions set to false for fireworks_ai. * feat: add darkbloom provider (#30876) * feat: add darkbloom provider * fix: document darkbloom provider endpoints * fix: address darkbloom review feedback * fix: update darkbloom tool metadata * fix: fail fast for non-Postgres database URLs (#30883) * fix(proxy): fail fast on non-PostgreSQL DATABASE_URL instead of hanging on startup LiteLLM's Prisma datasource is pinned to provider = 'postgresql', so a sqlite:// or mysql:// DATABASE_URL can never connect. Today that surfaces as an opaque startup stall where the port never binds, and a separate 'DB not connected' 500 on /key/generate when no DATABASE_URL is set at all leaves operators guessing what to configure. Validate the DATABASE_URL / DIRECT_URL scheme in run_server before any Prisma call and exit with an actionable message naming the unsupported scheme. Also reword CommonProxyErrors.db_not_connected_error to tell the operator to set DATABASE_URL to a postgresql:// connection string. Add regression tests covering postgres acceptance and sqlite/mysql/mssql rejection. * fix: resolve CI failures and proxy DB URL typing issue * fix(proxy): fail fast on non-PostgreSQL DATABASE_URLs with clear startup errors instead of hanging * Validate DIRECT_URL alongside DATABASE_URL startup guards * fix(bedrock): surface modeled HTTP status for mid-stream error events so 5xx is retryable (#24608) (#30946) * fix(bedrock): surface modeled HTTP status for mid-stream error events (#24608) * test(bedrock): mid-stream server errors trigger streaming fallback (#24608) * style(bedrock): black-format stream-error helper (#24608) * fix(mcp): re-land native tool preservation with typed annotations (#30645) * fix(mcp): preserve native tools in semantic filter hook with typed annotations * fix(mcp): tighten _is_mcp_tool Chat Completions shape check * fix(sambanova): return embeddings supported params instead of dropping them (#30937) * fix(router): send fallback metadata when streaming (#30914) When a streaming request triggers a fallback, there was previously no way to know it happened. This commit addresses this in a few ways: 1. The response now correctly populates the fallback headers (`x-litellm-attempted-fallbacks`) so callers know a fallback happened. 2. The correct model ID is passed in the streaming chunks. 3. A streaming chunk with the fallback error can be optionally sent back to the client (opt-in) by passing `include_fallback_errors: true` in the request. The format of the fallback errors while streaming is intentionally OpenAI compatible to not break existing libraries that parse these events. It was tested with Vercel's AI SDK (ai-sdk.dev). It is also opt-in, so it is not delieved unexpectedly to callers by default. * fix(mistral): drop output-only reasoning fields from input messages (#30884) LiteLLM attaches reasoning_content and thinking_blocks to assistant responses. Replaying those assistant turns verbatim forwarded the fields back to Mistral, whose input schema forbids unknown keys, so the whole request failed with a 422 extra_forbidden and reasoning models became unusable across multiple turns. Strip both fields from assistant messages before the request is built, in a spot that runs ahead of the image/file branch so it applies on every path. Fixes #30835 Co-authored-by: Cursor * fix(perplexity): bill search queries at the per-request price, not 1/1000 of it (#30652) * fix(perplexity): bill search queries at the per-request price, not 1/1000 The fallback cost calculator divided search_context_cost_per_query by 1000, but that field stores the per-request price in USD: sonar is {low: 0.005, medium: 0.008, high: 0.012}, matching Perplexity's published $5/$8/$12 per 1,000 requests expressed per request. The gemini cost calculator reads the same field per request with no division (its docstring calls it "the per-request cost"). The division understated search cost by 1000x on every Perplexity call that falls back to manual calculation (i.e. when the API does not return a pre-computed usage.cost). Use the value directly. Update the tests that had encoded the /1000 factor in their expectations, and drop an unused import flagged by ruff in the touched test file. * test(perplexity): update integration test search-cost expectations to per-request The integration tests still encoded the old /1000 search-cost factor, so they failed once the fallback calculator was corrected to bill search_context_cost_per_query per request. Update the four expected-cost computations (and the high-volume dollar-value comments) to match. * test(perplexity): drop unused mock imports flagged by ruff * fix: include model_access_groups when expanding all-team-models in get_team_models (#30622) * fix(fireworks_ai): return None for transcription in get_supported_openai_params Fireworks AI deprecated audio inference on 2026-06-10; the endpoint is decommissioned. Without an explicit transcription branch, requests with request_type='transcription' fell through to the else and returned FireworksAIConfig chat-completion params. Return None instead to signal the provider does not support transcription. * fix(proxy): gate include_fallback_errors behind expose_fallback_errors_to_caller setting Without an operator gate, any authenticated caller could set include_fallback_errors=True, trigger a fallback, and read raw upstream exception messages from the x-litellm-fallback-errors header and the litellm-fallback-metadata SSE event. Strip include_fallback_errors from request data in common_processing_pre_call_logic when expose_fallback_errors_to_caller is not set, so the router never builds the error list. Also gate _should_include_fallback_errors on the same setting as a secondary check for the streaming SSE injection path. * test(proxy): opt in to expose_fallback_errors_to_caller in streaming SSE test The operator gate added in e7ff3e1 means include_fallback_errors is only honoured when general_settings.expose_fallback_errors_to_caller is True. Set that flag via monkeypatch in the test that exercises the emit path. * test(prompt_templates): make test_convert_url hermetic instead of hitting picsum.photos test_convert_url called convert_url_to_base64 against a live picsum.photos URL and asserted nothing, so it added no real signal and broke CI whenever the host was unreachable (it was returning 522 and blocking this branch). Replace the live call with a mocked HTTP client and assert the produced base64 data URL, so the conversion path is exercised deterministically with no network dependency. This suite runs under VCR, which is why a transport level mock (respx) does not reliably intercept; mocking the client object itself is robust regardless. * fix(interactions): drop role from Interaction response to match Google spec Google removed the output-only role field from the Interaction schema (it now lives only on Turn), so the live OpenAPI compliance canary started failing with 'role' not in spec. Reconcile our generated types by removing role from Interaction, CreateModelInteractionParams, CreateAgentInteractionParams and from the LiteLLM InteractionsAPIResponse/InteractionsAPIStreamingResponse, stop stamping role=model in the responses-to-interactions transformation, and update the compliance and integration tests accordingly. Turn.role is kept since the spec still defines it. * fix: align all-team-models sentinel access * fix(router): forward include_fallback_errors through multi-hop fallbacks run_async_fallback received include_fallback_errors as an explicit named parameter, so it was bound out of **kwargs and never reached the nested async_function_with_fallbacks call. Multi-hop fallback chains (a fallback group that itself fails over) therefore stopped collecting fallback errors beyond the first hop when a caller opted in. Re-inject the flag into kwargs before the nested call so inner hops keep accumulating errors, which add_fallback_headers_to_response already merges across levels. --------- Co-authored-by: Srivatsa Kamballa Co-authored-by: Ahmad Shahzad <107808273+shzdehmd@users.noreply.github.com> Co-authored-by: Jeremy Chapeau <113923302+jychp@users.noreply.github.com> Co-authored-by: KRISH SONI <67964054+krishvsoni@users.noreply.github.com> Co-authored-by: Kent <72616338+kingdoooo@users.noreply.github.com> Co-authored-by: Ayush Shekhar <106994833+ayushh0110@users.noreply.github.com> Co-authored-by: dav nguyxn Co-authored-by: Tal Marian Co-authored-by: Hemant K <51333870+hemant1026@users.noreply.github.com> Co-authored-by: Cursor Co-authored-by: Yash Raj Pandey <55940078+devYRPauli@users.noreply.github.com> Co-authored-by: Zang Peiyu <166481866+factnn@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/__init__.py | 8 +- litellm/_lazy_imports_registry.py | 5 - litellm/constants.py | 1 + .../transformation.py | 3 - .../get_supported_openai_params.py | 8 +- .../sensitive_data_masker.py | 12 +- litellm/llms/bedrock/chat/invoke_handler.py | 19 +- litellm/llms/bedrock/common_utils.py | 65 ++- .../audio_transcription/transformation.py | 17 - litellm/llms/mistral/chat/transformation.py | 21 + litellm/llms/openai_like/providers.json | 8 + litellm/llms/perplexity/cost_calculator.py | 9 +- ...odel_prices_and_context_window_backup.json | 54 ++- .../provider_endpoints_support_backup.json | 17 + litellm/proxy/_types.py | 4 +- litellm/proxy/auth/auth_checks.py | 26 + litellm/proxy/auth/model_checks.py | 19 +- litellm/proxy/common_request_processing.py | 2 + litellm/proxy/db/db_url_settings.py | 93 +++- .../proxy/hooks/mcp_semantic_filter/hook.py | 173 +++++-- litellm/proxy/proxy_cli.py | 19 + litellm/proxy/proxy_server.py | 146 +++++- litellm/router.py | 87 +++- .../add_retry_fallback_headers.py | 118 ++++- litellm/router_utils/cooldown_cache.py | 1 + .../router_utils/fallback_event_handlers.py | 9 + litellm/types/interactions/generated.py | 11 - litellm/types/utils.py | 1 + litellm/utils.py | 4 +- model_prices_and_context_window.json | 54 ++- provider_endpoints_support.json | 17 + .../test_bedrock_completion.py | 26 +- .../test_bedrock_embedding_pricing.py | 34 ++ .../test_fireworks_ai_translation.py | 83 +--- tests/llm_translation/test_prompt_factory.py | 29 +- .../test_google_interactions_integration.py | 1 - .../interactions/test_openapi_compliance.py | 3 +- .../test_get_supported_openai_params.py | 14 + .../test_sensitive_data_masker.py | 43 ++ .../test_streaming_handler.py | 108 +++++ .../test_mistral_chat_transformation.py | 90 ++++ .../llms/openai_like/test_json_providers.py | 94 ++++ .../test_perplexity_cost_calculator.py | 18 +- .../perplexity/test_perplexity_integration.py | 17 +- .../mcp_server/test_mcp_debug.py | 10 +- .../mcp_server/test_semantic_tool_filter.py | 452 +++++++++++++++++- .../proxy/auth/test_auth_checks.py | 37 ++ .../proxy/auth/test_model_checks.py | 74 +++ .../proxy/db/test_db_url_settings.py | 91 +++- .../proxy_server/test_streaming_helpers.py | 351 +++++++++++++- tests/test_litellm/proxy/test_proxy_cli.py | 51 ++ .../test_add_retry_fallback_headers.py | 142 ++++++ .../test_fallback_event_handlers.py | 139 ++++++ ...test_router_streaming_fallback_metadata.py | 187 ++++++++ 54 files changed, 2752 insertions(+), 373 deletions(-) delete mode 100644 litellm/llms/fireworks_ai/audio_transcription/transformation.py create mode 100644 tests/llm_translation/test_bedrock_embedding_pricing.py create mode 100644 tests/test_litellm/router_utils/test_add_retry_fallback_headers.py create mode 100644 tests/test_litellm/router_utils/test_fallback_event_handlers.py create mode 100644 tests/test_litellm/test_router_streaming_fallback_metadata.py diff --git a/litellm/__init__.py b/litellm/__init__.py index d21234d2a81..b1ad63d72b0 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -673,6 +673,7 @@ elevenlabs_models: Set = set() dashscope_models: Set = set() moonshot_models: Set = set() publicai_models: Set = set() +darkbloom_models: Set = set() v0_models: Set = set() morph_models: Set = set() lambda_ai_models: Set = set() @@ -927,6 +928,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None): moonshot_models.add(key) elif value.get("litellm_provider") == "publicai": publicai_models.add(key) + elif value.get("litellm_provider") == "darkbloom": + darkbloom_models.add(key) elif value.get("litellm_provider") == "v0": v0_models.add(key) elif value.get("litellm_provider") == "morph": @@ -1075,6 +1078,7 @@ model_list = list( | dashscope_models | moonshot_models | publicai_models + | darkbloom_models | v0_models | morph_models | lambda_ai_models @@ -1179,6 +1183,7 @@ models_by_provider: dict = { "modelscope": modelscope_models, "moonshot": moonshot_models, "publicai": publicai_models, + "darkbloom": darkbloom_models, "v0": v0_models, "morph": morph_models, "lambda_ai": lambda_ai_models, @@ -1922,9 +1927,6 @@ if TYPE_CHECKING: from .llms.fireworks_ai.completion.transformation import ( FireworksAITextCompletionConfig as FireworksAITextCompletionConfig, ) - from .llms.fireworks_ai.audio_transcription.transformation import ( - FireworksAIAudioTranscriptionConfig as FireworksAIAudioTranscriptionConfig, - ) from .llms.fireworks_ai.embed.fireworks_ai_transformation import ( FireworksAIEmbeddingConfig as FireworksAIEmbeddingConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index e653b40fd04..4f131354d2e 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -260,7 +260,6 @@ LLM_CONFIG_NAMES = ( "SambaNovaEmbeddingConfig", "FireworksAIConfig", "FireworksAITextCompletionConfig", - "FireworksAIAudioTranscriptionConfig", "FireworksAIEmbeddingConfig", "FriendliaiChatConfig", "JinaAIEmbeddingConfig", @@ -1027,10 +1026,6 @@ _LLM_CONFIGS_IMPORT_MAP = { ".llms.fireworks_ai.completion.transformation", "FireworksAITextCompletionConfig", ), - "FireworksAIAudioTranscriptionConfig": ( - ".llms.fireworks_ai.audio_transcription.transformation", - "FireworksAIAudioTranscriptionConfig", - ), "FireworksAIEmbeddingConfig": ( ".llms.fireworks_ai.embed.fireworks_ai_transformation", "FireworksAIEmbeddingConfig", diff --git a/litellm/constants.py b/litellm/constants.py index c0e265c0e4a..083e9a1241b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -867,6 +867,7 @@ openai_compatible_providers: List = [ "docker_model_runner", "ragflow", "pinstripes", # Pinstripes - JSON-configured provider + "darkbloom", ] openai_text_completion_compatible_providers: List = ( [ # providers that support `/v1/completions` diff --git a/litellm/interactions/litellm_responses_transformation/transformation.py b/litellm/interactions/litellm_responses_transformation/transformation.py index 173d4ca8764..0ff1a97cd0b 100644 --- a/litellm/interactions/litellm_responses_transformation/transformation.py +++ b/litellm/interactions/litellm_responses_transformation/transformation.py @@ -300,9 +300,6 @@ class LiteLLMResponsesInteractionsConfig: "total_output_tokens": getattr(usage, "output_tokens", 0), } - # Add role - interactions_response_dict["role"] = "model" - # Add updated (same as created for now) interactions_response_dict["updated"] = created diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index e87042b9101..c22d3b99705 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -86,9 +86,7 @@ def get_supported_openai_params( model=model ) elif request_type == "transcription": - return litellm.FireworksAIAudioTranscriptionConfig().get_supported_openai_params( - model=model - ) + return None else: return litellm.FireworksAIConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "nvidia_nim": @@ -191,7 +189,9 @@ def get_supported_openai_params( ) elif custom_llm_provider == "sambanova": if request_type == "embeddings": - litellm.SambaNovaEmbeddingConfig().get_supported_openai_params(model=model) + return litellm.SambaNovaEmbeddingConfig().get_supported_openai_params( + model=model + ) else: return litellm.SambanovaConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "nebius": diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 4928dd08386..b14e12de7cd 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -12,6 +12,7 @@ class SensitiveDataMasker: visible_prefix: int = 4, visible_suffix: int = 4, mask_char: str = "*", + mask_short_values: bool = True, ): self.sensitive_patterns = sensitive_patterns or { "password", @@ -38,12 +39,17 @@ class SensitiveDataMasker: self.visible_prefix = visible_prefix self.visible_suffix = visible_suffix self.mask_char = mask_char + self.mask_short_values = mask_short_values def _mask_value(self, value: str) -> str: - if not value or len(str(value)) < (self.visible_prefix + self.visible_suffix): - return value - value_str = str(value) + if not value_str: + return value + if len(value_str) <= (self.visible_prefix + self.visible_suffix): + return ( + self.mask_char * len(value_str) if self.mask_short_values else value_str + ) + masked_length = len(value_str) - (self.visible_prefix + self.visible_suffix) # Handle the case where visible_suffix is 0 to avoid showing the entire string diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 75b560b4d6d..9fca7bc61af 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -70,6 +70,7 @@ from ..base_aws_llm import BaseAWSLLM from ..common_utils import ( BedrockError, ModelResponseIterator, + build_bedrock_stream_error, get_bedrock_response_stream_shape, get_bedrock_tool_name, ) @@ -1841,23 +1842,7 @@ class AWSEventStreamDecoder: parsed_response = self.parser.parse(response_dict, response_stream_shape) if response_dict["status_code"] != 200: - decoded_body = response_dict["body"].decode() - if isinstance(decoded_body, dict): - error_message = decoded_body.get("message") - elif isinstance(decoded_body, str): - error_message = decoded_body - else: - error_message = "" - exception_status = response_dict["headers"].get(":exception-type") - error_message = exception_status + " " + error_message - raise BedrockError( - status_code=response_dict["status_code"], - message=( - json.dumps(error_message) - if isinstance(error_message, dict) - else error_message - ), - ) + raise build_bedrock_stream_error(response_dict, response_stream_shape) if "chunk" in parsed_response: chunk = parsed_response.get("chunk") if not chunk: diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index bdc5da321c6..9f58e5c0f1c 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -7,9 +7,21 @@ Common utilities used across bedrock chat/embedding/image generation import functools import json import os -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Mapping, + Optional, + TypedDict, + Union, +) if TYPE_CHECKING: + from botocore.model import Shape + from litellm.types.llms.bedrock import BedrockCreateBatchRequest import httpx @@ -1132,6 +1144,39 @@ def get_bedrock_response_stream_shape(): return _load_bedrock_response_stream_shape() +class BedrockEventStreamResponseDict(TypedDict): + status_code: int + headers: Mapping[str, str] + body: bytes + + +def build_bedrock_stream_error( + response_dict: BedrockEventStreamResponseDict, + response_stream_shape: Shape | None, +) -> BedrockError: + """Build a BedrockError for a non-200 event-stream error event. + + botocore hard-codes HTTP 400 on every mid-stream error event, so the modeled + ResponseStream member's httpStatusCode is the real status. Resolve it from the + shape and fall back to the raw status when the type is not modeled. + """ + exception_type = response_dict["headers"].get(":exception-type") + decoded_body = response_dict["body"].decode() + message = f"{exception_type} {decoded_body}" if exception_type else decoded_body + + status_code = response_dict["status_code"] + if exception_type is not None and response_stream_shape is not None: + member = response_stream_shape.members.get(exception_type) + if member is not None: + modeled_status = ( + (member.metadata or {}).get("error", {}).get("httpStatusCode") + ) + if modeled_status is not None: + status_code = int(modeled_status) + + return BedrockError(status_code=status_code, message=message) + + class BedrockEventStreamDecoderBase: """ Base class for event stream decoding for Bedrock @@ -1156,23 +1201,7 @@ class BedrockEventStreamDecoderBase: parsed_response = self.parser.parse(response_dict, response_stream_shape) if response_dict["status_code"] != 200: - decoded_body = response_dict["body"].decode() - if isinstance(decoded_body, dict): - error_message = decoded_body.get("message") - elif isinstance(decoded_body, str): - error_message = decoded_body - else: - error_message = "" - exception_status = response_dict["headers"].get(":exception-type") - error_message = exception_status + " " + error_message - raise BedrockError( - status_code=response_dict["status_code"], - message=( - json.dumps(error_message) - if isinstance(error_message, dict) - else error_message - ), - ) + raise build_bedrock_stream_error(response_dict, response_stream_shape) if "chunk" in parsed_response: chunk = parsed_response.get("chunk") if not chunk: diff --git a/litellm/llms/fireworks_ai/audio_transcription/transformation.py b/litellm/llms/fireworks_ai/audio_transcription/transformation.py deleted file mode 100644 index 00bb5f26797..00000000000 --- a/litellm/llms/fireworks_ai/audio_transcription/transformation.py +++ /dev/null @@ -1,17 +0,0 @@ -from typing import List - -from litellm.types.llms.openai import OpenAIAudioTranscriptionOptionalParams - -from ...openai.transcriptions.whisper_transformation import ( - OpenAIWhisperAudioTranscriptionConfig, -) -from ..common_utils import FireworksAIMixin - - -class FireworksAIAudioTranscriptionConfig( - FireworksAIMixin, OpenAIWhisperAudioTranscriptionConfig -): - def get_supported_openai_params( - self, model: str - ) -> List[OpenAIAudioTranscriptionOptionalParams]: - return ["language", "prompt", "response_format", "timestamp_granularities"] diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index f1ad3708236..8d0cf993814 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -247,6 +247,8 @@ class MistralConfig(OpenAIGPTConfig): The above statement is not valid now. Need to plan to remove all the #1,2,3 Mistral API supports content as a list. """ + messages = [self._strip_output_only_fields(m) for m in messages] + ## 1. If 'image_url' or 'file' in content, then transform with base class and mistral-specific handling for m in messages: _content_block = m.get("content") @@ -409,6 +411,25 @@ class MistralConfig(OpenAIGPTConfig): return cleaned_tools + @classmethod + def _strip_output_only_fields(cls, message: AllMessageValues) -> AllMessageValues: + """ + ``reasoning_content`` and ``thinking_blocks`` are output-only fields that + LiteLLM attaches to assistant responses. Mistral's input schema forbids + unknown fields, so replaying them verbatim in a follow-up turn triggers a + 422 ``extra_forbidden``. Drop them before the request is sent. + """ + if message["role"] != "assistant": + return message + return cast( + AllMessageValues, + { + k: v + for k, v in message.items() + if k not in ("reasoning_content", "thinking_blocks") + }, + ) + @classmethod def _handle_name_in_message(cls, message: AllMessageValues) -> AllMessageValues: """ diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index 24943563937..d87346fea70 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -115,6 +115,14 @@ "max_completion_tokens": "max_tokens" } }, + "darkbloom": { + "base_url": "https://api.darkbloom.dev/v1", + "api_key_env": "DARKBLOOM_API_KEY", + "api_base_env": "DARKBLOOM_API_BASE", + "param_mappings": { + "max_completion_tokens": "max_tokens" + } + }, "neosantara": { "base_url": "https://api.neosantara.xyz/v1", "api_key_env": "NEOSANTARA_API_KEY", diff --git a/litellm/llms/perplexity/cost_calculator.py b/litellm/llms/perplexity/cost_calculator.py index bf055f91aa0..ec7ec397ea6 100644 --- a/litellm/llms/perplexity/cost_calculator.py +++ b/litellm/llms/perplexity/cost_calculator.py @@ -98,10 +98,11 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: if num_search_queries > 0 and search_cost_value is not None: # Handle both dict and float formats if isinstance(search_cost_value, dict): - # Use the "low" size as default - tests expect 0.005 / 1000 - search_cost_per_query = ( - _safe_float_cast(search_cost_value.get("search_context_size_low", 0)) - / 1000 + # search_context_cost_per_query stores the per-request price in USD + # (e.g. sonar low = $0.005/request). Use it directly, matching the + # gemini cost calculator which reads the same field per request. + search_cost_per_query = _safe_float_cast( + search_cost_value.get("search_context_size_low", 0) ) else: search_cost_per_query = _safe_float_cast(search_cost_value) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 5c962cf8440..4f022e1f882 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -571,7 +571,7 @@ "output_vector_size": 1536 }, "amazon.titan-embed-text-v2:0": { - "input_cost_per_token": 2e-07, + "input_cost_per_token": 2e-08, "litellm_provider": "bedrock", "max_input_tokens": 8192, "max_tokens": 8192, @@ -39908,24 +39908,6 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, - "fireworks_ai/accounts/fireworks/models/whisper-v3": { - "max_tokens": 4096, - "max_input_tokens": 4096, - "max_output_tokens": 4096, - "input_cost_per_token": 0.0, - "output_cost_per_token": 0.0, - "litellm_provider": "fireworks_ai", - "mode": "audio_transcription" - }, - "fireworks_ai/accounts/fireworks/models/whisper-v3-turbo": { - "max_tokens": 4096, - "max_input_tokens": 4096, - "max_output_tokens": 4096, - "input_cost_per_token": 0.0, - "output_cost_per_token": 0.0, - "litellm_provider": "fireworks_ai", - "mode": "audio_transcription" - }, "fireworks_ai/accounts/fireworks/models/yi-34b": { "max_tokens": 4096, "max_input_tokens": 4096, @@ -43061,6 +43043,40 @@ "supports_tool_choice": true, "supports_vision": false }, + "darkbloom/gemma-4-26b": { + "input_cost_per_token": 3e-08, + "litellm_provider": "darkbloom", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.65e-07, + "source": "https://www.darkbloom.dev/", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "darkbloom/gpt-oss-20b": { + "input_cost_per_token": 1.45e-08, + "litellm_provider": "darkbloom", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 7e-08, + "source": "https://www.darkbloom.dev/", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "deepseek/deepseek-v4-pro": { "cache_creation_input_token_cost": 0.0, "cache_read_input_token_cost": 3.625e-09, diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index db6183edaa0..dd7712aabca 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -1835,6 +1835,23 @@ "interactions": true } }, + "darkbloom": { + "display_name": "Darkbloom (`darkbloom`)", + "url": "https://docs.litellm.ai/docs/providers/darkbloom", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false + } + }, "predibase": { "display_name": "Predibase (`predibase`)", "url": "https://docs.litellm.ai/docs/providers/predibase", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ac90302fdaa..5bba842c7eb 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3358,7 +3358,9 @@ class ProxyException(Exception): class CommonProxyErrors(str, enum.Enum): db_not_connected_error = ( - "DB not connected. See https://docs.litellm.ai/docs/proxy/virtual_keys" + "DB not connected. This endpoint needs a database; set DATABASE_URL to a " + "PostgreSQL connection string (postgresql://...) to enable it. " + "See https://docs.litellm.ai/docs/proxy/virtual_keys" ) no_llm_router = "No models configured on proxy" not_allowed_access = "Admin-only endpoint. Not allowed to access this." diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 52788ed9238..88db2a2b7ea 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2954,6 +2954,26 @@ async def _get_agent_ids_from_access_groups( ) +def _resolve_all_team_model_sentinel_for_auth_check( + models: List[str], + llm_router: Optional[Router], + team_id: Optional[str], +) -> List[str]: + if ( + SpecialModelNames.all_team_models.value not in models + or team_id is None + or llm_router is None + ): + return models + proxy_models = llm_router.get_model_names() + non_sentinel_models = [ + model for model in models if model != SpecialModelNames.all_team_models.value + ] + if not proxy_models: + return non_sentinel_models or models + return list(dict.fromkeys(non_sentinel_models + proxy_models)) + + def _check_model_access_helper( model: str, llm_router: Optional[Router], @@ -2971,6 +2991,12 @@ def _check_model_access_helper( model_name=model, team_id=team_id ) + models = _resolve_all_team_model_sentinel_for_auth_check( + models=models, + llm_router=llm_router, + team_id=team_id, + ) + if ( len(access_groups) > 0 and llm_router is not None ): # check if token contains any model access groups diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index b89db51c6f1..aa53954da8f 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -122,9 +122,16 @@ def get_key_models( SpecialModelNames.all_team_models.value in all_models and user_api_key_dict.team_id is not None ): - all_models = list( - user_api_key_dict.team_models - ) # copy to avoid mutating cached objects + all_models = list(user_api_key_dict.team_models) + if SpecialModelNames.all_team_models.value in all_models: + all_models = [ + model + for model in all_models + if model != SpecialModelNames.all_team_models.value + ] + all_models.extend(proxy_model_list) + if include_model_access_groups: + all_models.extend(model_access_groups.keys()) if SpecialModelNames.all_proxy_models.value in all_models: all_models = list(proxy_model_list) # copy to avoid mutating caller's list if include_model_access_groups: @@ -160,6 +167,12 @@ def get_team_models( all_models_set.update(team_models) if SpecialModelNames.all_team_models.value in all_models_set: all_models_set.update(team_models) + # GH#30619: expand all-team-models sentinel + # to the actual proxy model list + all_models_set.discard(SpecialModelNames.all_team_models.value) + all_models_set.update(proxy_model_list) + if include_model_access_groups: + all_models_set.update(model_access_groups.keys()) if SpecialModelNames.all_proxy_models.value in all_models_set: all_models_set.update(proxy_model_list) if include_model_access_groups: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 8ef931e8d25..8dec08460b4 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1037,6 +1037,8 @@ class ProxyBaseLLMRequestProcessing: version=version, proxy_config=proxy_config, ) + if not general_settings.get("expose_fallback_errors_to_caller"): + self.data.pop("include_fallback_errors", None) if route_type in {"aresponses", "_aresponses_websocket"}: await _authorize_response_file_search_vector_stores( data=self.data, diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py index 58478db5e2e..ae2307658dd 100644 --- a/litellm/proxy/db/db_url_settings.py +++ b/litellm/proxy/db/db_url_settings.py @@ -32,7 +32,7 @@ password when their ``*_READ_REPLICA`` counterpart is unset. import os import urllib.parse -from typing import Optional, cast +from typing import Final, cast from pydantic import AliasChoices, Field from pydantic_settings import BaseSettings, SettingsConfigDict @@ -44,6 +44,41 @@ from litellm.proxy.auth import rds_iam_token _IAM_ENV_KEY = "IAM_TOKEN_DB_AUTH" _DEFAULT_PG_PORT = "5432" +# schema.prisma pins `provider = "postgresql"`, so these are the only schemes +# Prisma can actually connect with. +SUPPORTED_DB_SCHEMES: Final[frozenset[str]] = frozenset({"postgresql", "postgres"}) +_MISSING_SCHEME = "" + + +def unsupported_db_scheme(database_url: str) -> str | None: + """Return the connection URL scheme when it is not PostgreSQL, else None. + + A `sqlite://` / `mysql://` URL can never connect against the + postgresql-only datasource, but the resulting Prisma failure is opaque and + version-dependent (a confusing migration error, or a startup that never + binds). Callers use this to reject the URL up front with an actionable + error instead. + + A schemeless value (e.g. a malformed DSN like ``user:pass@host/db``) yields + the ``_MISSING_SCHEME`` placeholder rather than the raw URL, so callers that + log the return value never echo embedded credentials. + """ + scheme = urllib.parse.urlsplit(database_url).scheme.lower() + if scheme in SUPPORTED_DB_SCHEMES: + return None + return scheme or _MISSING_SCHEME + + +def unsupported_db_scheme_message(env_var: str, scheme: str) -> str: + """Operator-facing message naming the offending env var and scheme.""" + return ( + f"{env_var} uses unsupported scheme '{scheme}'. LiteLLM's database " + "features (virtual keys, store_model_in_db, spend tracking) require " + "PostgreSQL; use a 'postgresql://' connection string. SQLite and other " + "engines are not supported. " + "See https://docs.litellm.ai/docs/proxy/virtual_keys" + ) + class DatabaseURLSettings(BaseSettings): """Discrete ``DATABASE_*`` env vars, loaded once at process start. @@ -58,46 +93,47 @@ class DatabaseURLSettings(BaseSettings): iam_token_db_auth: bool = Field(default=False, validation_alias=_IAM_ENV_KEY) # Writer - database_url: Optional[str] = Field(default=None, validation_alias="DATABASE_URL") - database_host: Optional[str] = Field(default=None, validation_alias="DATABASE_HOST") + database_url: str | None = Field(default=None, validation_alias="DATABASE_URL") + direct_url: str | None = Field(default=None, validation_alias="DIRECT_URL") + database_host: str | None = Field(default=None, validation_alias="DATABASE_HOST") database_port: str = Field( default=_DEFAULT_PG_PORT, validation_alias="DATABASE_PORT" ) - database_user: Optional[str] = Field( + database_user: str | None = Field( default=None, validation_alias=AliasChoices("DATABASE_USER", "DATABASE_USERNAME"), ) - database_name: Optional[str] = Field(default=None, validation_alias="DATABASE_NAME") - database_schema: Optional[str] = Field( + database_name: str | None = Field(default=None, validation_alias="DATABASE_NAME") + database_schema: str | None = Field( default=None, validation_alias="DATABASE_SCHEMA" ) - database_password: Optional[str] = Field( + database_password: str | None = Field( default=None, validation_alias="DATABASE_PASSWORD" ) # Read replica - database_url_read_replica: Optional[str] = Field( + database_url_read_replica: str | None = Field( default=None, validation_alias="DATABASE_URL_READ_REPLICA" ) - database_host_read_replica: Optional[str] = Field( + database_host_read_replica: str | None = Field( default=None, validation_alias="DATABASE_HOST_READ_REPLICA" ) - database_port_read_replica: Optional[str] = Field( + database_port_read_replica: str | None = Field( default=None, validation_alias="DATABASE_PORT_READ_REPLICA" ) - database_user_read_replica: Optional[str] = Field( + database_user_read_replica: str | None = Field( default=None, validation_alias=AliasChoices( "DATABASE_USER_READ_REPLICA", "DATABASE_USERNAME_READ_REPLICA" ), ) - database_name_read_replica: Optional[str] = Field( + database_name_read_replica: str | None = Field( default=None, validation_alias="DATABASE_NAME_READ_REPLICA" ) - database_schema_read_replica: Optional[str] = Field( + database_schema_read_replica: str | None = Field( default=None, validation_alias="DATABASE_SCHEMA_READ_REPLICA" ) - database_password_read_replica: Optional[str] = Field( + database_password_read_replica: str | None = Field( default=None, validation_alias="DATABASE_PASSWORD_READ_REPLICA" ) @@ -106,7 +142,7 @@ class DatabaseURLSettings(BaseSettings): """Load the settings from ``os.environ`` (read at call time).""" return cls() - def build_writer_url(self) -> Optional[str]: + def build_writer_url(self) -> str | None: """Return the writer URL to set, or ``None`` to leave it as-is. Raises ``RuntimeError`` (naming the offending vars) when IAM auth is @@ -156,7 +192,7 @@ class DatabaseURLSettings(BaseSettings): ) return None - def build_reader_url(self) -> Optional[str]: + def build_reader_url(self) -> str | None: """Return the read-replica URL to set, or ``None`` to leave it as-is. Opt-in via ``DATABASE_HOST_READ_REPLICA``; never clobbers a @@ -217,11 +253,11 @@ class DatabaseURLSettings(BaseSettings): def _password_url( *, user: str, - password: Optional[str], + password: str | None, host: str, port: str, name: str, - schema: Optional[str], + schema: str | None, ) -> str: """Percent-encode credentials into a ``postgresql://`` URL. @@ -239,6 +275,26 @@ class DatabaseURLSettings(BaseSettings): url += f"?schema={schema}" return url + def _raise_for_unsupported_scheme(self) -> None: + """Reject an operator-pinned non-PostgreSQL writer / direct / reader URL. + + The componentized entrypoints (gateway / backend / migrations) call + ``apply_to_env`` and then hand the URL straight to Prisma, bypassing + the CLI's own guard. A pinned URL flows through untouched, so validate + the same three vars the CLI guard checks (DATABASE_URL, DIRECT_URL, and + the read replica) rather than letting Prisma stall on an unusable scheme. + """ + for env_var, url in ( + ("DATABASE_URL", self.database_url), + ("DIRECT_URL", self.direct_url), + ("DATABASE_URL_READ_REPLICA", self.database_url_read_replica), + ): + if not url: + continue + bad_scheme = unsupported_db_scheme(url) + if bad_scheme is not None: + raise RuntimeError(unsupported_db_scheme_message(env_var, bad_scheme)) + def apply_to_env(self) -> bool: """Write the assembled URL(s) into ``os.environ``. @@ -246,6 +302,7 @@ class DatabaseURLSettings(BaseSettings): password auth that assembled a fresh URL). False means there was nothing to do — an operator-pinned URL, or no discrete fields. """ + self._raise_for_unsupported_scheme() wrote_writer = False writer_url = self.build_writer_url() if writer_url is not None: diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 6343faaa965..9888baf897e 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -123,11 +123,70 @@ class SemanticToolFilterHook(CustomLogger): return openai_tools_as_dicts + def _is_mcp_tool(self, tool: object) -> bool: + """ + Check whether *tool* is registered in the MCP semantic router. + + Classification strategy (shape-first, lookup-second): + 1. Chat Completions format dicts are always native. + 2. Responses API function tools are always native. + 3. Everything else is looked up by name in the MCP registry. + """ + if ( + isinstance(tool, dict) + and tool.get("type") == "function" + and isinstance(tool.get("function"), dict) + ): + return False + if ( + isinstance(tool, dict) + and tool.get("type") == "function" + and isinstance(tool.get("name"), str) + ): + return False + name, _ = self.filter._extract_tool_info(tool) + return bool(name) and name in self.filter._tool_map + def _get_metadata_variable_name(self, data: dict) -> str: if "litellm_metadata" in data: return "litellm_metadata" return "metadata" + def _emit_filter_metadata( + self, + data: dict, + mcp_tools: list[object], + filtered_mcp_tools: list[object], + native_tools: list[object], + filtered_tools: list[object], + ) -> None: + """ + Emit response-header metadata when MCP tools were filtered. + + Stats report MCP-only counts so downstream consumers see accurate + semantic filter metrics. Skips metadata entirely for purely-native + requests to avoid spurious headers. + """ + if mcp_tools: + filter_stats = f"{len(mcp_tools)}->{len(filtered_mcp_tools)}" + tool_names_csv = self._get_tool_names_csv(filtered_mcp_tools) + + _metadata_variable_name = self._get_metadata_variable_name(data) + metadata = data.setdefault(_metadata_variable_name, {}) + metadata["litellm_semantic_filter_stats"] = filter_stats + metadata["litellm_semantic_filter_tools"] = tool_names_csv + + verbose_proxy_logger.info( + f"Semantic tool filter: {filter_stats} MCP tools " + f"({len(native_tools)} native preserved, " + f"{len(filtered_tools)} total)" + ) + else: + verbose_proxy_logger.info( + f"Semantic tool filter: all {len(native_tools)} tools " + f"are native, no MCP filtering applied" + ) + async def async_pre_call_hook( self, user_api_key_dict: "UserAPIKeyAuth", @@ -140,53 +199,55 @@ class SemanticToolFilterHook(CustomLogger): This hook is called before the LLM request is made. It filters the tools list to only include semantically relevant tools. - - Args: - user_api_key_dict: User authentication - cache: Cache instance - data: Request data containing messages and tools - call_type: Type of call (completion, acompletion, etc.) - - Returns: - Modified data dict with filtered tools, or None if no changes """ - # Only filter endpoints that support tools if call_type not in ("completion", "acompletion", "aresponses"): verbose_proxy_logger.debug( f"Skipping semantic filter for call_type={call_type}" ) return None - # Check if tools are present tools = data.get("tools") if not tools: verbose_proxy_logger.debug("No tools in request, skipping semantic filter") return None - original_tool_count = len(tools) - - # Check for MCP references (server_url="litellm_proxy") and expand them + # Expanded MCP tools are in OpenAI nested format which + # filter_tools/_extract_tool_info cannot name-match, so we skip + # semantic filtering and return early. if self._should_expand_mcp_tools(tools): verbose_proxy_logger.debug( "Detected litellm_proxy MCP references, expanding before semantic filtering" ) try: + native_tools_before_expand = [ + t + for t in tools + if not (isinstance(t, dict) and t.get("type") == "mcp") + ] + expanded_tools = await self._expand_mcp_tools(tools, user_api_key_dict) if not expanded_tools: + if native_tools_before_expand: + data["tools"] = native_tools_before_expand + verbose_proxy_logger.warning( + "No MCP tools expanded, preserving " + f"{len(native_tools_before_expand)} native tools" + ) + return data verbose_proxy_logger.warning( "No tools expanded from MCP references" ) return None + data["tools"] = native_tools_before_expand + expanded_tools verbose_proxy_logger.info( - f"Expanded {len(tools)} MCP reference(s) to {len(expanded_tools)} tools" + f"Expanded MCP references to {len(expanded_tools)} tools " + f"({len(native_tools_before_expand)} native preserved), " + f"skipping semantic filter (OpenAI nested format)" ) - - # Update tools for filtering - tools = expanded_tools - original_tool_count = len(tools) + return data except Exception as e: verbose_proxy_logger.error( @@ -194,7 +255,6 @@ class SemanticToolFilterHook(CustomLogger): ) return None - # Check if messages are present (try both "messages" and "input" for responses API) messages = data.get("messages", []) if not messages: messages = data.get("input", []) @@ -204,13 +264,11 @@ class SemanticToolFilterHook(CustomLogger): ) return None - # Check if filter is enabled if not self.filter.enabled: verbose_proxy_logger.debug("Semantic filter disabled, skipping") return None try: - # Extract user query from messages user_query = self.filter.extract_user_query(messages) if not user_query: verbose_proxy_logger.debug( @@ -218,33 +276,60 @@ class SemanticToolFilterHook(CustomLogger): ) return None + native_tools: list[object] = [] + mcp_tools: list[object] = [] + mcp_indices: set[int] = set() + for i, t in enumerate(tools): + if self._is_mcp_tool(t): + mcp_tools.append(t) + mcp_indices.add(i) + else: + native_tools.append(t) + verbose_proxy_logger.debug( - f"Applying semantic filter to {len(tools)} tools " - f"with query: '{user_query[:50]}...'" + f"Applying semantic filter: {len(mcp_tools)} MCP tools, " + f"{len(native_tools)} native tools, " + f"query: '{user_query[:50]}...'" ) - # Filter tools semantically - filtered_tools = await self.filter.filter_tools( - query=user_query, - available_tools=tools, # type: ignore - ) + if mcp_tools: + filtered_mcp_tools = await self.filter.filter_tools( + query=user_query, + available_tools=mcp_tools, # type: ignore + ) + else: + filtered_mcp_tools = [] + + filtered_mcp_names: set[str] = set() + for t in filtered_mcp_tools: + name, _ = self.filter._extract_tool_info(t) + if name: + filtered_mcp_names.add(name) + + filtered_tools: list[object] = [] + for i, t in enumerate(tools): + if i in mcp_indices: + name, _ = self.filter._extract_tool_info(t) + if name in filtered_mcp_names: + filtered_tools.append(t) + else: + filtered_tools.append(t) - # Always update tools and emit header (even if count unchanged) data["tools"] = filtered_tools - # Store filter stats and tool names for response header - filter_stats = f"{original_tool_count}->{len(filtered_tools)}" - tool_names_csv = self._get_tool_names_csv(filtered_tools) - - _metadata_variable_name = self._get_metadata_variable_name(data) - data[_metadata_variable_name][ - "litellm_semantic_filter_stats" - ] = filter_stats - data[_metadata_variable_name][ - "litellm_semantic_filter_tools" - ] = tool_names_csv - - verbose_proxy_logger.info(f"Semantic tool filter: {filter_stats} tools") + try: + self._emit_filter_metadata( + data=data, + mcp_tools=mcp_tools, + filtered_mcp_tools=filtered_mcp_tools, + native_tools=native_tools, + filtered_tools=filtered_tools, + ) + except Exception as e: + verbose_proxy_logger.warning( + f"Failed to emit semantic filter metadata: {e}", + exc_info=True, + ) return data @@ -266,7 +351,7 @@ class SemanticToolFilterHook(CustomLogger): from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH _metadata_variable_name = self._get_metadata_variable_name(data) - metadata = data[_metadata_variable_name] + metadata = data.get(_metadata_variable_name, {}) filter_stats = metadata.get("litellm_semantic_filter_stats") if not filter_stats: diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 9c4d7b1bb5d..d0281885482 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1195,6 +1195,25 @@ def run_server( os.getenv("DATABASE_URL", None) is not None or os.getenv("DIRECT_URL", None) is not None ): + from litellm.proxy.db.db_url_settings import ( + unsupported_db_scheme, + unsupported_db_scheme_message, + ) + + for _db_env in ("DATABASE_URL", "DIRECT_URL"): + _candidate_url = os.getenv(_db_env) + if _candidate_url is None: + continue + _bad_scheme = unsupported_db_scheme(_candidate_url) + if _bad_scheme is not None: + print( + f"\033[1;31mLiteLLM Proxy: " + f"{unsupported_db_scheme_message(_db_env, _bad_scheme)}" + "\033[0m", + file=sys.stderr, + flush=True, + ) + sys.exit(1) try: from litellm.secret_managers.main import get_secret diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c5676fbb32f..9238cf91301 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -106,6 +106,10 @@ from litellm.proxy.common_utils.callback_utils import ( process_callback, ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.router_utils.add_retry_fallback_headers import ( + get_fallback_errors_from_headers, + get_hidden_params_dict, +) from litellm.types.utils import ( ModelResponse, ModelResponseStream, @@ -7085,57 +7089,122 @@ def _get_client_requested_model_for_streaming(request_data: dict) -> str: return requested_model if isinstance(requested_model, str) else "" +def _is_positive_int_like(value: Any) -> bool: + try: + return int(value) > 0 + except (TypeError, ValueError): + return False + + +def _should_include_fallback_errors(request_data: dict[str, object]) -> bool: + if not general_settings.get("expose_fallback_errors_to_caller"): + return False + return request_data.get("include_fallback_errors") is True + + +def _get_streaming_fallback_metadata( + response_obj: object, +) -> tuple[bool, str | None, list[dict[str, object]]]: + additional_headers = get_hidden_params_dict(response_obj).get("additional_headers") + if not isinstance(additional_headers, dict): + return False, None, [] + + if not _is_positive_int_like( + additional_headers.get("x-litellm-attempted-fallbacks") + ): + return False, None, [] + + fallback_model = additional_headers.get("x-litellm-model-group") + fallback_errors = get_fallback_errors_from_headers(additional_headers) + if isinstance(fallback_model, str) and fallback_model: + return True, fallback_model, fallback_errors + return True, None, fallback_errors + + +def _format_fallback_metadata_sse_event( + *, + fallback_model: str | None, + fallback_errors: list[dict[str, object]], +) -> str: + import time + + payload = { + "id": "litellm-fallback-metadata", + "object": "chat.completion.chunk", + "created": int(time.time()), + "model": fallback_model or "", + "choices": [], + "litellm_fallback": { + "fallback_model": fallback_model, + "errors": fallback_errors, + }, + } + return f"data: {json.dumps(payload)}\n\n" + + def _restamp_streaming_chunk_model( *, chunk: Any, requested_model_from_client: str, request_data: dict, model_mismatch_logged: bool, -) -> Tuple[Any, bool]: + fallback_was_attempted: bool = False, + fallback_model_from_metadata: str | None = None, +) -> tuple[Any, bool]: + target_model = ( + fallback_model_from_metadata + if fallback_was_attempted + else requested_model_from_client + ) # Always return the client-requested model name (not provider-prefixed internal identifiers) # on streaming chunks. + # On fallback, use the public OpenAI-compatible model name. This keeps + # provider-prefixed internal identifiers from leaking into the public API. # # Note: This warning is intentionally verbose. A mismatch is a useful signal that an # internal provider/deployment identifier is leaking into the public API, and helps # maintainers/operators catch regressions while preserving OpenAI-compatible output. - if not requested_model_from_client or not isinstance(chunk, (BaseModel, dict)): + if not target_model or not isinstance(chunk, (BaseModel, dict)): return chunk, model_mismatch_logged # For Azure Model Router, preserve the actual model used in each chunk - if _is_azure_model_router_request(requested_model_from_client): + if not fallback_was_attempted and _is_azure_model_router_request( + requested_model_from_client + ): return chunk, model_mismatch_logged # For fastest_response batch completions, preserve the winning model's name # instead of stamping the comma-separated list the client sent. - if request_data.get("fastest_response", False): + if not fallback_was_attempted and request_data.get("fastest_response", False): return chunk, model_mismatch_logged downstream_model = ( chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None) ) - if downstream_model == requested_model_from_client: + if downstream_model == target_model: return chunk, model_mismatch_logged - if not model_mismatch_logged and downstream_model != requested_model_from_client: + if not model_mismatch_logged and downstream_model != target_model: verbose_proxy_logger.debug( - "litellm_call_id=%s: streaming chunk model mismatch - requested=%r downstream=%r. Overriding model to requested.", + "litellm_call_id=%s: streaming chunk model mismatch - target=%r downstream=%r fallback_was_attempted=%s. Overriding chunk model to target.", request_data.get("litellm_call_id"), - requested_model_from_client, + target_model, downstream_model, + fallback_was_attempted, ) model_mismatch_logged = True if isinstance(chunk, dict): - chunk["model"] = requested_model_from_client + chunk["model"] = target_model return chunk, model_mismatch_logged try: - setattr(chunk, "model", requested_model_from_client) + chunk.model = target_model except Exception as e: verbose_proxy_logger.error( "litellm_call_id=%s: failed to override chunk.model=%r on chunk_type=%s. error=%s", request_data.get("litellm_call_id"), - requested_model_from_client, + target_model, type(chunk), str(e), exc_info=True, @@ -7294,7 +7363,14 @@ async def async_data_generator( requested_model_from_client = _get_client_requested_model_for_streaming( request_data=request_data ) + ( + fallback_was_attempted, + fallback_model_from_metadata, + fallback_errors, + ) = _get_streaming_fallback_metadata(response) model_mismatch_logged = False + fallback_metadata_event_sent = False + include_fallback_errors = _should_include_fallback_errors(request_data) # Use a running string instead of list + join to avoid O(n^2) overhead. # Previously "".join(str_so_far_parts) was called every chunk, re-joining # the entire accumulated response. String += is O(n) amortized total. @@ -7332,13 +7408,37 @@ async def async_data_generator( str_so_far=_str_so_far, ) + # Mid-stream fallbacks surface metadata on individual chunks rather than + # the response wrapper. Keep scanning chunks until a fallback model is + # resolved, then latch it for the rest of the stream. + if fallback_model_from_metadata is None: + ( + chunk_fallback_was_attempted, + chunk_fallback_model, + chunk_fallback_errors, + ) = _get_streaming_fallback_metadata(chunk) + if chunk_fallback_was_attempted: + fallback_was_attempted = True + fallback_model_from_metadata = chunk_fallback_model + fallback_errors = fallback_errors or chunk_fallback_errors + + pending_fallback_event = ( + include_fallback_errors + and fallback_was_attempted + and fallback_errors + and not fallback_metadata_event_sent + ) + chunk, model_mismatch_logged = _restamp_streaming_chunk_model( chunk=chunk, requested_model_from_client=requested_model_from_client, request_data=request_data, model_mismatch_logged=model_mismatch_logged, + fallback_was_attempted=fallback_was_attempted, + fallback_model_from_metadata=fallback_model_from_metadata, ) + raw_passthrough = False if isinstance(chunk, BaseModel): chunk = _serialize_streaming_chunk(chunk) elif isinstance(chunk, bytes): @@ -7354,14 +7454,14 @@ async def async_data_generator( raise ValueError( "Raw SSE stream exceeded maximum buffered size without a frame delimiter" ) - continue - if chunk.startswith(("data:", "event:", ":")): + raw_passthrough = True + elif chunk.startswith(("data:", "event:", ":")): yield ( chunk if chunk.endswith(_SSE_FRAME_DELIMITERS) else chunk + "\n\n" ) - continue + raw_passthrough = True elif isinstance(chunk, str) and is_raw_sse_stream: raw_sse_buffer += chunk while True: @@ -7373,15 +7473,23 @@ async def async_data_generator( raise ValueError( "Raw SSE stream exceeded maximum buffered size without a frame delimiter" ) - continue + raw_passthrough = True elif isinstance(chunk, str) and chunk.startswith("data: "): error_message = chunk break - try: - yield _format_streaming_sse_chunk(chunk=chunk) - except Exception as e: - yield f"data: {str(e)}\n\n" + if not raw_passthrough: + try: + yield _format_streaming_sse_chunk(chunk=chunk) + except Exception as e: + yield f"data: {str(e)}\n\n" + + if pending_fallback_event: + yield _format_fallback_metadata_sse_event( + fallback_model=fallback_model_from_metadata, + fallback_errors=fallback_errors, + ) + fallback_metadata_event_sent = True stream_completed = True if not needs_iterator_wrap: diff --git a/litellm/router.py b/litellm/router.py index e54eadfb872..acce1c58a8b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -40,7 +40,6 @@ import anyio import httpx import openai from openai import AsyncOpenAI -from pydantic import BaseModel from typing_extensions import overload import litellm @@ -81,8 +80,10 @@ from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2 from litellm.router_strategy.simple_shuffle import simple_shuffle from litellm.router_strategy.tag_based_routing import get_deployments_for_tag from litellm.router_utils.add_retry_fallback_headers import ( + _HiddenParamsHost, add_fallback_headers_to_response, add_retry_headers_to_response, + get_hidden_params_dict, ) from litellm.router_utils.batch_utils import ( _get_router_metadata_variable_name, @@ -2165,6 +2166,36 @@ class Router: ) setattr(fallback_item, "usage", combined_usage) + @staticmethod + def _prepare_fallback_hidden_params( + fallback_response: object, + ) -> tuple[dict[str, object], dict[str, object]]: + fallback_hidden_params = get_hidden_params_dict(fallback_response) + fallback_headers = fallback_hidden_params.get("additional_headers") + if not isinstance(fallback_headers, dict): + return fallback_hidden_params, {} + return fallback_hidden_params, cast("dict[str, object]", fallback_headers) + + @staticmethod + def _apply_fallback_hidden_params_to_item( + fallback_item: object, + prepared_fallback_hidden_params: tuple[dict[str, object], dict[str, object]], + ) -> None: + if fallback_item is None or not hasattr(fallback_item, "_hidden_params"): + return + + fallback_hidden_params, fallback_headers = prepared_fallback_hidden_params + item_hidden_params = get_hidden_params_dict(fallback_item) + item_headers = item_hidden_params.get("additional_headers") + if not isinstance(item_headers, dict): + item_headers = {} + + cast(_HiddenParamsHost, fallback_item)._hidden_params = { + **item_hidden_params, + **fallback_hidden_params, + "additional_headers": {**item_headers, **fallback_headers}, + } + async def _acompletion_streaming_iterator( self, model_response: CustomStreamWrapper, @@ -2257,12 +2288,22 @@ class Router: model_group=model_group, args=(), kwargs=initial_kwargs, + include_fallback_errors=initial_kwargs.get( + "include_fallback_errors", False + ) + is True, ) ) # If fallback returns a streaming response, iterate over it if hasattr(fallback_response, "__aiter__"): + prepared_fallback_hidden_params = ( + Router._prepare_fallback_hidden_params(fallback_response) + ) async for fallback_item in fallback_response: # type: ignore + Router._apply_fallback_hidden_params_to_item( + fallback_item, prepared_fallback_hidden_params + ) if ( fallback_item and isinstance(fallback_item, ModelResponseStream) @@ -2686,11 +2727,21 @@ class Router: model_group=model_group, args=(), kwargs=initial_kwargs, + include_fallback_errors=initial_kwargs.get( + "include_fallback_errors", False + ) + is True, ) ) if hasattr(fallback_response, "__aiter__"): + prepared_fallback_hidden_params = ( + Router._prepare_fallback_hidden_params(fallback_response) + ) async for fallback_item in fallback_response: # type: ignore + Router._apply_fallback_hidden_params_to_item( + fallback_item, prepared_fallback_hidden_params + ) if partial_usage is not None: Router._combine_responses_fallback_usage( fallback_item, partial_usage @@ -2815,7 +2866,13 @@ class Router: ) if hasattr(fallback_response, "__iter__"): + prepared_fallback_hidden_params = ( + Router._prepare_fallback_hidden_params(fallback_response) + ) for fallback_item in fallback_response: + Router._apply_fallback_hidden_params_to_item( + fallback_item, prepared_fallback_hidden_params + ) if ( fallback_item and isinstance(fallback_item, ModelResponseStream) @@ -2972,6 +3029,7 @@ class Router: **kwargs, } input_kwargs.pop("silent_model", None) + input_kwargs.pop("include_fallback_errors", None) _response = litellm.acompletion(**input_kwargs) @@ -6478,6 +6536,7 @@ class Router: model_group: Optional[str], args: tuple, kwargs: dict, + include_fallback_errors: bool = False, ): """ Common utilities for async_function_with_fallbacks @@ -6501,6 +6560,8 @@ class Router: input_kwargs["max_fallbacks"] = self.max_fallbacks if "fallback_depth" not in input_kwargs: input_kwargs["fallback_depth"] = 0 + if include_fallback_errors: + input_kwargs["include_fallback_errors"] = True # ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list # Skip for error types that have their own dedicated fallback handlers @@ -6759,6 +6820,7 @@ class Router: If it fails after num_retries, fall back to another model group """ model_group: Optional[str] = kwargs.get("model") + include_fallback_errors = kwargs.get("include_fallback_errors", False) is True disable_fallbacks: Optional[bool] = kwargs.pop("disable_fallbacks", False) fallbacks: Optional[List] = kwargs.get("fallbacks", self.fallbacks) context_window_fallbacks: Optional[List] = kwargs.get( @@ -6802,6 +6864,7 @@ class Router: model_group, args, kwargs, + include_fallback_errors=include_fallback_errors, ) def _handle_mock_testing_fallbacks( @@ -9725,17 +9788,19 @@ class Router: # - if healthy_deployments > 1, return model group rate limit headers # - else return the model's rate limit headers """ - if ( - isinstance(response, BaseModel) - and hasattr(response, "_hidden_params") - and isinstance(response._hidden_params, dict) # type: ignore - ): - response._hidden_params.setdefault("additional_headers", {}) # type: ignore - response._hidden_params["additional_headers"][ # type: ignore - "x-litellm-model-group" - ] = model_group + if response is not None and hasattr(response, "_hidden_params"): + hidden_params = getattr(response, "_hidden_params", {}) or {} + if hasattr(hidden_params, "model_dump"): + hidden_params = hidden_params.model_dump() + if not isinstance(hidden_params, dict): + return response + response._hidden_params = hidden_params - additional_headers = response._hidden_params["additional_headers"] # type: ignore + additional_headers = hidden_params.get("additional_headers") + if not isinstance(additional_headers, dict): + additional_headers = {} + hidden_params["additional_headers"] = additional_headers + additional_headers["x-litellm-model-group"] = model_group # Lift QualityRouter routing decision into response headers for # transparency. The decision is stashed in request_kwargs.metadata diff --git a/litellm/router_utils/add_retry_fallback_headers.py b/litellm/router_utils/add_retry_fallback_headers.py index 6b921a0db8a..0b927714ca9 100644 --- a/litellm/router_utils/add_retry_fallback_headers.py +++ b/litellm/router_utils/add_retry_fallback_headers.py @@ -1,44 +1,99 @@ -from typing import Any, Optional, Union +import json +from typing import Protocol, TypedDict, cast from pydantic import BaseModel -from litellm.types.utils import HiddenParams + +class FallbackErrorInfo(TypedDict): + message: str + type: str + param: str | None + code: str | None -def _add_headers_to_response(response: Any, headers: dict) -> Any: +class _HiddenParamsHost(Protocol): + _hidden_params: dict[str, object] + + +def get_hidden_params_dict(response: object) -> dict[str, object]: + hidden_params: object = cast(object, getattr(response, "_hidden_params", None)) + if isinstance(hidden_params, BaseModel): + return cast("dict[str, object]", hidden_params.model_dump()) + if isinstance(hidden_params, dict): + return cast("dict[str, object]", hidden_params) + return {} + + +def _ensure_additional_headers_dict( + hidden_params: dict[str, object], +) -> dict[str, object]: + additional_headers = hidden_params.get("additional_headers") + if isinstance(additional_headers, dict): + return cast("dict[str, object]", additional_headers) + return {} + + +def get_fallback_error_info(error: Exception) -> FallbackErrorInfo: + message = cast(object, getattr(error, "message", str(error))) + error_type = cast(object, getattr(error, "type", error.__class__.__name__)) + param = cast(object, getattr(error, "param", None)) + code = cast(object, getattr(error, "status_code", getattr(error, "code", None))) + return FallbackErrorInfo( + message=str(message), + type=str(error_type), + param=str(param) if param is not None else None, + code=str(code) if code is not None else None, + ) + + +def _coerce_error_dicts(items: list[object]) -> list[dict[str, object]]: + return [cast("dict[str, object]", item) for item in items if isinstance(item, dict)] + + +def get_fallback_errors_from_headers( + additional_headers: dict[str, object], +) -> list[dict[str, object]]: + existing_errors = additional_headers.get("x-litellm-fallback-errors") + if isinstance(existing_errors, list): + return _coerce_error_dicts(cast("list[object]", existing_errors)) + if isinstance(existing_errors, str): + try: + parsed_errors: object = cast(object, json.loads(existing_errors)) + except json.JSONDecodeError: + return [] + if isinstance(parsed_errors, list): + return _coerce_error_dicts(cast("list[object]", parsed_errors)) + return [] + + +def _add_headers_to_response(response: object, headers: dict[str, object]) -> object: """ Helper function to add headers to a response's hidden params """ - if response is None or not isinstance(response, BaseModel): + if response is None: return response - hidden_params: Optional[Union[dict, HiddenParams]] = getattr( - response, "_hidden_params", {} - ) + if not isinstance(response, BaseModel) and not hasattr(response, "_hidden_params"): + return response - if hidden_params is None: - hidden_params_dict = {} - elif isinstance(hidden_params, HiddenParams): - hidden_params_dict = hidden_params.model_dump() - else: - hidden_params_dict = hidden_params + hidden_params = get_hidden_params_dict(response) + additional_headers = _ensure_additional_headers_dict(hidden_params) + additional_headers.update(headers) + hidden_params["additional_headers"] = additional_headers - hidden_params_dict.setdefault("additional_headers", {}) - hidden_params_dict["additional_headers"].update(headers) - - setattr(response, "_hidden_params", hidden_params_dict) + cast(_HiddenParamsHost, response)._hidden_params = hidden_params return response def add_retry_headers_to_response( - response: Any, + response: object, attempted_retries: int, - max_retries: Optional[int] = None, -) -> Any: + max_retries: int | None = None, +) -> object: """ Add retry headers to the request """ - retry_headers = { + retry_headers: dict[str, object] = { "x-litellm-attempted-retries": attempted_retries, } if max_retries is not None: @@ -48,9 +103,10 @@ def add_retry_headers_to_response( def add_fallback_headers_to_response( - response: Any, + response: object, attempted_fallbacks: int, -) -> Any: + fallback_errors: list[FallbackErrorInfo] | None = None, +) -> object: """ Add fallback headers to the response @@ -64,7 +120,19 @@ def add_fallback_headers_to_response( Note: It's intentional that we don't add max_fallbacks in response headers Want to avoid bloat in the response headers for performance. """ - fallback_headers = { + fallback_headers: dict[str, object] = { "x-litellm-attempted-fallbacks": attempted_fallbacks, } - return _add_headers_to_response(response, fallback_headers) + response = _add_headers_to_response(response, fallback_headers) + if fallback_errors is None or response is None: + return response + + hidden_params = get_hidden_params_dict(response) + additional_headers = _ensure_additional_headers_dict(hidden_params) + merged_errors = get_fallback_errors_from_headers(additional_headers) + [ + cast("dict[str, object]", error) for error in fallback_errors + ] + additional_headers["x-litellm-fallback-errors"] = json.dumps(merged_errors) + hidden_params["additional_headers"] = additional_headers + cast(_HiddenParamsHost, response)._hidden_params = hidden_params + return response diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index b210ea44596..dcfa44381c1 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -38,6 +38,7 @@ class CooldownCache: visible_prefix=50, # Show first 50 characters visible_suffix=0, # Show last 0 characters mask_char="*", # Use * for masking + mask_short_values=False, # Truncate long messages only; keep short ones readable ) def _common_add_cooldown_logic( diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index eb756e3cf8b..f0edc7fc9db 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -6,6 +6,7 @@ from litellm._logging import verbose_router_logger from litellm.integrations.custom_logger import CustomLogger from litellm.router_utils.add_retry_fallback_headers import ( add_fallback_headers_to_response, + get_fallback_error_info, ) from litellm.types.router import LiteLLMParamsTypedDict @@ -90,6 +91,7 @@ async def run_async_fallback( original_exception: Exception, max_fallbacks: int, fallback_depth: int, + include_fallback_errors: bool = False, **kwargs, ) -> Any: """ @@ -118,6 +120,7 @@ async def run_async_fallback( raise original_exception error_from_fallbacks = original_exception + fallback_errors = (get_fallback_error_info(original_exception),) for mg in fallback_model_group: if mg == original_model_group: @@ -136,6 +139,8 @@ async def run_async_fallback( fallback_depth = fallback_depth + 1 kwargs["fallback_depth"] = fallback_depth kwargs["max_fallbacks"] = max_fallbacks + if include_fallback_errors: + kwargs["include_fallback_errors"] = include_fallback_errors response = await litellm_router.async_function_with_fallbacks( *args, **kwargs ) @@ -143,6 +148,9 @@ async def run_async_fallback( response = add_fallback_headers_to_response( response=response, attempted_fallbacks=fallback_depth, + fallback_errors=( + list(fallback_errors) if include_fallback_errors else None + ), ) # callback for successfull_fallback_event(): await log_success_fallback_event( @@ -153,6 +161,7 @@ async def run_async_fallback( return response except Exception as e: error_from_fallbacks = e + fallback_errors = fallback_errors + (get_fallback_error_info(e),) await log_failure_fallback_event( original_model_group=original_model_group, kwargs=kwargs, diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index b38cd8f58b9..a07642073af 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -954,9 +954,6 @@ class Interaction(BaseModel): None, description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).", ) - role: Optional[str] = Field( - None, description="Output only. The role of the interaction." - ) outputs: Optional[List[Content]] = Field( None, description="Output only. Responses from the model." ) @@ -1031,9 +1028,6 @@ class CreateModelInteractionParams(BaseModel): None, description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).", ) - role: Optional[str] = Field( - None, description="Output only. The role of the interaction." - ) outputs: Optional[List[Content]] = Field( None, description="Output only. Responses from the model." ) @@ -1101,9 +1095,6 @@ class CreateAgentInteractionParams(BaseModel): None, description="Output only. The time at which the response was last updated in ISO 8601 format\n(YYYY-MM-DDThh:mm:ssZ).", ) - role: Optional[str] = Field( - None, description="Output only. The role of the interaction." - ) outputs: Optional[List[Content]] = Field( None, description="Output only. Responses from the model." ) @@ -1323,7 +1314,6 @@ class InteractionsAPIResponse(BaseLiteLLMOpenAIResponseObject): status: Optional[str] = None created: Optional[str] = None updated: Optional[str] = None - role: Optional[str] = None # Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026. outputs: Optional[List[Dict[str, Any]]] = None # New schema field (Api-Revision: 2026-05-20). @@ -1356,7 +1346,6 @@ class InteractionsAPIStreamingResponse(BaseLiteLLMOpenAIResponseObject): status: Optional[str] = None created: Optional[str] = None updated: Optional[str] = None - role: Optional[str] = None # Legacy schema field (Api-Revision: 2026-05-07). Remove after June 8, 2026. outputs: Optional[List[Dict[str, Any]]] = None # New schema field (Api-Revision: 2026-05-20). diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 16e693c7d7c..76854cd28f9 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3463,6 +3463,7 @@ class LlmProviders(str, Enum): TENSORMESH = "tensormesh" LIBERTAI = "libertai" PINSTRIPES = "pinstripes" + DARKBLOOM = "darkbloom" LITELLM_AGENT = "litellm_agent" CURSOR = "cursor" BEDROCK_MANTLE = "bedrock_mantle" diff --git a/litellm/utils.py b/litellm/utils.py index 29f703104da..a842e9e058d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3191,7 +3191,7 @@ def get_optional_params_transcription( model=model, drop_params=drop_params if drop_params is not None else False, ) - elif provider_config is not None: # handles fireworks ai, and any future providers + elif provider_config is not None: # custom audio transcription config supported_params = provider_config.get_supported_openai_params(model=model) _check_valid_arg(supported_params=supported_params) optional_params = provider_config.map_openai_params( @@ -8915,8 +8915,6 @@ class ProviderConfigManager: ) return AzureSpeechAudioTranscriptionConfig() - if litellm.LlmProviders.FIREWORKS_AI == provider: - return litellm.FireworksAIAudioTranscriptionConfig() elif litellm.LlmProviders.DEEPGRAM == provider: return litellm.DeepgramAudioTranscriptionConfig() elif litellm.LlmProviders.ELEVENLABS == provider: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3c50dde9277..56baa5c573f 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -571,7 +571,7 @@ "output_vector_size": 1536 }, "amazon.titan-embed-text-v2:0": { - "input_cost_per_token": 2e-07, + "input_cost_per_token": 2e-08, "litellm_provider": "bedrock", "max_input_tokens": 8192, "max_tokens": 8192, @@ -39946,24 +39946,6 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, - "fireworks_ai/accounts/fireworks/models/whisper-v3": { - "max_tokens": 4096, - "max_input_tokens": 4096, - "max_output_tokens": 4096, - "input_cost_per_token": 0.0, - "output_cost_per_token": 0.0, - "litellm_provider": "fireworks_ai", - "mode": "audio_transcription" - }, - "fireworks_ai/accounts/fireworks/models/whisper-v3-turbo": { - "max_tokens": 4096, - "max_input_tokens": 4096, - "max_output_tokens": 4096, - "input_cost_per_token": 0.0, - "output_cost_per_token": 0.0, - "litellm_provider": "fireworks_ai", - "mode": "audio_transcription" - }, "fireworks_ai/accounts/fireworks/models/yi-34b": { "max_tokens": 4096, "max_input_tokens": 4096, @@ -43543,5 +43525,39 @@ "supports_assistant_prefill": true, "supports_reasoning": false, "source": "https://pinstripes.io/pricing" + }, + "darkbloom/gemma-4-26b": { + "input_cost_per_token": 3e-08, + "litellm_provider": "darkbloom", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.65e-07, + "source": "https://www.darkbloom.dev/", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "darkbloom/gpt-oss-20b": { + "input_cost_per_token": 1.45e-08, + "litellm_provider": "darkbloom", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 7e-08, + "source": "https://www.darkbloom.dev/", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_system_messages": true, + "supports_tool_choice": true } } diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 7386ced3e6d..9991ff9e01e 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2008,6 +2008,23 @@ "interactions": true } }, + "darkbloom": { + "display_name": "Darkbloom (`darkbloom`)", + "url": "https://docs.litellm.ai/docs/providers/darkbloom", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false + } + }, "predibase": { "display_name": "Predibase (`predibase`)", "url": "https://docs.litellm.ai/docs/providers/predibase", diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index fa22ff6b392..f4c307e9c8a 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -2502,19 +2502,34 @@ async def test_bedrock_image_url_sync_client(): mock_post.assert_called_once() -def test_bedrock_error_handling_streaming(): +@pytest.mark.parametrize( + "exception_type, expected_status_code", + [ + ("internalServerException", 500), + ("serviceUnavailableException", 503), + ("modelTimeoutException", 408), + ("modelStreamErrorException", 424), + ("validationException", 400), + ], +) +def test_bedrock_error_handling_streaming(exception_type, expected_status_code): + """Bedrock event-stream error events arrive with botocore's hard-coded + status_code=400; the decoder must surface the modeled HTTP status instead + (e.g. internalServerException -> 500). For 5xx this is what makes the error + retryable downstream; for all types it replaces the misleading 400 with the + true code. Regression for #24608.""" from litellm.llms.bedrock.chat.invoke_handler import ( AWSEventStreamDecoder, BedrockError, ) - from unittest.mock import patch, Mock + from unittest.mock import Mock event = Mock() event.to_response_dict = Mock( return_value={ "status_code": 400, "headers": { - ":exception-type": "serviceUnavailableException", + ":exception-type": exception_type, ":content-type": "application/json", ":message-type": "exception", }, @@ -2525,11 +2540,10 @@ def test_bedrock_error_handling_streaming(): decoder = AWSEventStreamDecoder( model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0" ) - with pytest.raises(Exception) as e: + with pytest.raises(BedrockError) as e: decoder._parse_message_from_event(event) - assert isinstance(e.value, BedrockError) assert "Bedrock is unable to process your request." in e.value.message - assert e.value.status_code == 400 + assert e.value.status_code == expected_status_code @pytest.mark.parametrize( diff --git a/tests/llm_translation/test_bedrock_embedding_pricing.py b/tests/llm_translation/test_bedrock_embedding_pricing.py new file mode 100644 index 00000000000..099d73fed87 --- /dev/null +++ b/tests/llm_translation/test_bedrock_embedding_pricing.py @@ -0,0 +1,34 @@ +""" +Tests for AWS Bedrock embedding model pricing in the model cost map. + +Regression test for the Amazon Titan Text Embeddings V2 commercial price, +which was previously set 10x too high (2e-07 instead of 2e-08). +AWS lists Titan Text Embeddings V2 at $0.02 per 1M input tokens +(= $0.00002 per 1K tokens = 2e-08 per token). +""" + +import importlib + + +class TestBedrockEmbeddingPricing: + """Test suite for Bedrock embedding model pricing in the cost map.""" + + def test_titan_embed_v2_commercial_input_cost(self, monkeypatch): + """Titan Text Embeddings V2 should be priced at $0.02 / 1M tokens (2e-08).""" + # Scope the local-cost-map flag to this test only, so it does not leak + # into sibling tests. monkeypatch restores the environment on teardown. + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + + import litellm.litellm_core_utils.get_model_cost_map + import litellm + + # Reload so the cost map is re-read from the local file with the flag set. + importlib.reload(litellm.litellm_core_utils.get_model_cost_map) + importlib.reload(litellm) + + model = litellm.model_cost["amazon.titan-embed-text-v2:0"] + + assert model["input_cost_per_token"] == 2e-08 + assert model["output_cost_per_token"] == 0.0 + assert model["litellm_provider"] == "bedrock" + assert model["mode"] == "embedding" diff --git a/tests/llm_translation/test_fireworks_ai_translation.py b/tests/llm_translation/test_fireworks_ai_translation.py index 204f4d9e31b..27059581e4d 100644 --- a/tests/llm_translation/test_fireworks_ai_translation.py +++ b/tests/llm_translation/test_fireworks_ai_translation.py @@ -7,9 +7,10 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import litellm -from litellm import transcription +from litellm.litellm_core_utils.get_supported_openai_params import ( + get_supported_openai_params, +) from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig -from base_audio_transcription_unit_tests import BaseLLMAudioTranscriptionTest fireworks = FireworksAIConfig() @@ -69,74 +70,16 @@ def test_map_response_format(): assert result == {"response_format": response_format} -_AUDIO_FILE_PATH = os.path.join( - os.path.dirname(os.path.realpath(__file__)), "gettysburg.wav" -) - - -class TestFireworksAIAudioTranscription(BaseLLMAudioTranscriptionTest): - def get_base_audio_transcription_call_args(self) -> dict: - return { - "model": "fireworks_ai/whisper-v3", - "api_base": "https://audio-prod.api.fireworks.ai/v1", - } - - def get_custom_llm_provider(self) -> litellm.LlmProviders: - return litellm.LlmProviders.FIREWORKS_AI - - def test_audio_transcription(self): - from unittest.mock import MagicMock - - from openai.types.audio import Transcription - - audio_file = open(_AUDIO_FILE_PATH, "rb") - mock_client = MagicMock() - mock_client.audio.transcriptions.create.return_value = Transcription( - text="four score and seven years ago" - ) - - transcript = transcription( - **self.get_base_audio_transcription_call_args(), - file=audio_file, - api_key="fw-test-key", - client=mock_client, - ) - - assert transcript.text == "four score and seven years ago" - sent = mock_client.audio.transcriptions.create.call_args.kwargs - assert sent["model"] == "whisper-v3" - assert sent["file"] is audio_file - - @pytest.mark.asyncio - async def test_audio_transcription_async(self): - from unittest.mock import AsyncMock, MagicMock - - from openai.types.audio import Transcription - - audio_file = open(_AUDIO_FILE_PATH, "rb") - raw_response = MagicMock() - raw_response.headers = {} - raw_response.parse.return_value = Transcription( - text="four score and seven years ago" - ) - mock_client = MagicMock() - mock_client.audio.transcriptions.with_raw_response.create = AsyncMock( - return_value=raw_response - ) - - transcript = await litellm.atranscription( - **self.get_base_audio_transcription_call_args(), - file=audio_file, - api_key="fw-test-key", - client=mock_client, - ) - - assert transcript.text == "four score and seven years ago" - sent = ( - mock_client.audio.transcriptions.with_raw_response.create.call_args.kwargs - ) - assert sent["model"] == "whisper-v3" - assert sent["file"] is audio_file +def test_get_supported_openai_params_transcription_returns_none(): + # Fireworks AI deprecated audio transcription on 2026-06-10; the endpoint + # is decommissioned. Returning None (not chat-completion params) signals + # to callers that transcription is unsupported for this provider. + result = get_supported_openai_params( + model="fireworks_ai/accounts/fireworks/models/whisper-v3", + custom_llm_provider="fireworks_ai", + request_type="transcription", + ) + assert result is None @pytest.mark.parametrize( diff --git a/tests/llm_translation/test_prompt_factory.py b/tests/llm_translation/test_prompt_factory.py index 36e47e3c2f4..05a58a135d2 100644 --- a/tests/llm_translation/test_prompt_factory.py +++ b/tests/llm_translation/test_prompt_factory.py @@ -605,8 +605,33 @@ def test_no_messages_yields_user_text(): assert contents == expected_output -def test_convert_url(): - convert_url_to_base64("https://picsum.photos/id/237/200/300") +def test_convert_url(monkeypatch): + import base64 + from unittest.mock import MagicMock + + import httpx + + from litellm.litellm_core_utils.prompt_templates.image_handling import ( + in_memory_cache, + ) + + url = "https://picsum.photos/id/237/200/300" + image_bytes = b"\x89PNG\r\n\x1a\nfake-png-bytes" + + mock_client = MagicMock() + mock_client.get.return_value = httpx.Response( + 200, content=image_bytes, headers={"Content-Type": "image/png"} + ) + + monkeypatch.setattr(litellm, "user_url_validation", False, raising=False) + monkeypatch.setattr(litellm, "module_level_client", mock_client, raising=False) + in_memory_cache.flush_cache() + + result = convert_url_to_base64(url) + + expected = "data:image/png;base64," + base64.b64encode(image_bytes).decode("utf-8") + assert result == expected + mock_client.get.assert_called_once() def test_azure_tool_call_invoke_helper(): diff --git a/tests/test_litellm/interactions/test_google_interactions_integration.py b/tests/test_litellm/interactions/test_google_interactions_integration.py index cfff26d51ef..9c651cc94f5 100644 --- a/tests/test_litellm/interactions/test_google_interactions_integration.py +++ b/tests/test_litellm/interactions/test_google_interactions_integration.py @@ -299,7 +299,6 @@ class TestGoogleInteractionsResponseStructure: assert hasattr(response, "outputs") assert hasattr(response, "usage") assert hasattr(response, "model") or hasattr(response, "agent") - assert hasattr(response, "role") assert hasattr(response, "created") assert hasattr(response, "updated") diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 44c0b6e5b02..209e99895db 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -162,7 +162,8 @@ class TestResponseCompliance: # Keep this aligned with the live spec. schema = spec_dict["components"]["schemas"]["Interaction"] - # Output fields (readOnly). + # Output fields (readOnly). `role` was removed from the `Interaction` + # schema by Google; it now lives only on `Turn`. output_fields = [ "id", "status", diff --git a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py b/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py index 3c280c6ba92..84900e3f2ed 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py +++ b/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py @@ -132,3 +132,17 @@ def test_azure_base_model_detection_preserved(): assert params is not None assert "reasoning_effort" in params assert "tools" in params + + +def test_sambanova_embeddings_request_returns_list_not_none(): + """The sambanova embeddings branch resolved the config but dropped the result, + so embedding requests got ``None`` instead of the supported-params list while the + chat branch returned correctly. A list (the sambanova embeddings config exposes no + extra params, hence ``[]``) must reach the caller.""" + embedding_params = get_supported_openai_params( + model="E5-Mistral-7B-Instruct", + custom_llm_provider="sambanova", + request_type="embeddings", + ) + + assert embedding_params == [] diff --git a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py index 6808c4821c1..7239636fd48 100644 --- a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py +++ b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py @@ -126,6 +126,49 @@ def test_lists_with_sensitive_keys_are_masked(): assert masked["tags"] == ["prod", "test"] +def test_short_secrets_are_fully_masked(): + """ + Regression test: secrets at or below the reveal threshold (visible_prefix + + visible_suffix, 8 by default) were returned verbatim instead of masked. + An exactly-8-char value hit masked_length == 0 and round-tripped unchanged; + anything shorter hit the early return. Both leaked short credentials (e.g. an + 8-char redis password) in plaintext through mask_dict. + """ + masker = SensitiveDataMasker() + + # Boundary: exactly 8 chars previously returned verbatim. + assert masker._mask_value("abcd1234") == "********" + # Below threshold previously hit the early return and leaked verbatim. + assert masker._mask_value("sk-12") == "*****" + # Values above the threshold must still partially reveal, not over-mask. + assert masker._mask_value("abcd12345") == "abcd*2345" + + masked = masker.mask_dict({"redis_password": "pass1234", "api_key": "sk-7a"}) + assert masked["redis_password"] == "********" + assert masked["api_key"] == "*****" + + +def test_mask_short_values_false_keeps_short_values_readable(): + """ + mask_short_values=False opts out of full masking so short values are returned + as-is. This preserves the truncation use (e.g. CooldownCache shows the first 50 + chars of an exception and only masks longer tails), while longer values are still + partially masked. + """ + masker = SensitiveDataMasker( + visible_prefix=50, visible_suffix=0, mask_short_values=False + ) + + short = "Test exception for structure validation" + assert masker._mask_value(short) == short + + long_value = "x" * 60 + masked = masker._mask_value(long_value) + assert masked.startswith("x" * 50) + assert masked.endswith("*" * 10) + assert len(masked) == 60 + + def test_cost_per_token_fields_not_masked(): """ Regression test: cost fields like input_cost_per_token contain "token" in their name diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 734e49161c4..81af0ad3e6f 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -878,6 +878,114 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging): assert "invalid maxOutputTokens" in str(excinfo.value) +def _bedrock_error_event(exception_type: str): + """A mocked botocore event-stream error event: status_code is botocore's + hard-coded 400, with the real type in the :exception-type header.""" + event = Mock() + event.to_response_dict = Mock( + return_value={ + "status_code": 400, + "headers": { + ":exception-type": exception_type, + ":content-type": "application/json", + ":message-type": "exception", + }, + "body": b'{"message":"Bedrock had an internal error."}', + } + ) + return event + + +@pytest.mark.asyncio +async def test_bedrock_midstream_internal_server_error_wraps_for_fallback( + logging_obj: Logging, +): + """End-to-end regression for https://github.com/BerriAI/litellm/issues/24608: + a Bedrock mid-stream internalServerException event (botocore stamps it 400) + must flow through the real decoder, gain its modeled 500 status, and wrap + into MidStreamFallbackError so the Router can run streaming fallback. + + Calls the real AWSEventStreamDecoder, so reverting the decoder status fix + makes the decoder raise BedrockError(400) and the gate raises BadRequestError + directly -> this test fails without the fix.""" + from litellm.exceptions import MidStreamFallbackError + from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder + + decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + + async def _bedrock_stream(): + decoder._parse_message_from_event( + _bedrock_error_event("internalServerException") + ) + yield # unreachable; the line above raises + + async def _make_call(**kwargs): + return _bedrock_stream() + + response = CustomStreamWrapper( + completion_stream=None, + model="anthropic.claude-3-sonnet-20240229-v1:0", + logging_obj=logging_obj, + custom_llm_provider="bedrock", + make_call=_make_call, + ) + + with pytest.raises(MidStreamFallbackError): + await response.__anext__() + + +@pytest.mark.asyncio +async def test_bedrock_5xx_wraps_for_midstream_fallback(logging_obj: Logging): + """Gate contract: a Bedrock 5xx (here 503 serviceUnavailableException) wraps + into MidStreamFallbackError so the Router can run streaming fallback.""" + from litellm.exceptions import MidStreamFallbackError + from litellm.llms.bedrock.chat.invoke_handler import BedrockError + + async def _raise_503(**kwargs): + raise BedrockError( + status_code=503, + message="serviceUnavailableException Bedrock is unavailable.", + ) + + response = CustomStreamWrapper( + completion_stream=None, + model="anthropic.claude-3-sonnet-20240229-v1:0", + logging_obj=logging_obj, + custom_llm_provider="bedrock", + make_call=_raise_503, + ) + + with pytest.raises(MidStreamFallbackError): + await response.__anext__() + + +@pytest.mark.asyncio +async def test_bedrock_validation_error_raises_directly(logging_obj: Logging): + """Gate contract: a Bedrock validationException (400) is a client error and + must surface directly, never wrapped into MidStreamFallbackError.""" + from litellm.exceptions import MidStreamFallbackError + from litellm.llms.bedrock.chat.invoke_handler import BedrockError + + async def _raise_400(**kwargs): + raise BedrockError( + status_code=400, + message="validationException malformed input.", + ) + + response = CustomStreamWrapper( + completion_stream=None, + model="anthropic.claude-3-sonnet-20240229-v1:0", + logging_obj=logging_obj, + custom_llm_provider="bedrock", + make_call=_raise_400, + ) + + with pytest.raises(Exception) as excinfo: + await response.__anext__() + assert not isinstance(excinfo.value, MidStreamFallbackError) + assert getattr(excinfo.value, "status_code", None) == 400 + + @pytest.mark.asyncio async def test_async_streaming_read_timeout_triggers_midstream_fallback( logging_obj: Logging, diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index 3ee53bb46cd..7a3f372582f 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -719,3 +719,93 @@ class TestMistralFileHandling: # Check that file_ids are modified to match Mistral's expected format assert result[0]["content"][1]["file_id"] == "file-12345" # type: ignore assert result[0]["content"][2]["file_id"] == "file-67890" # type: ignore + + +class TestMistralStripsOutputOnlyFields: + """Mistral rejects unknown input fields with a 422 ``extra_forbidden``. + + LiteLLM attaches ``reasoning_content`` / ``thinking_blocks`` to assistant + responses, so replaying an assistant turn verbatim must not forward them. + Regression for https://github.com/BerriAI/litellm/issues/30835. + """ + + def test_assistant_reasoning_content_is_dropped(self): + messages = cast( + List[AllMessageValues], + [ + {"role": "user", "content": "Question?"}, + { + "role": "assistant", + "content": "Follow-up", + "reasoning_content": "Some internal reasoning text.", + "thinking_blocks": [ + {"type": "thinking", "thinking": "step", "signature": "mistral"} + ], + }, + ], + ) + + result = cast( + List[AllMessageValues], + MistralConfig()._transform_messages( + messages=messages, model="mistral-medium-3-5" + ), + ) + + assistant_message = result[-1] + assert "reasoning_content" not in assistant_message + assert "thinking_blocks" not in assistant_message + assert assistant_message["content"] == "Follow-up" + assert assistant_message["role"] == "assistant" + + def test_non_assistant_messages_are_untouched(self): + messages = cast( + List[AllMessageValues], + [{"role": "user", "content": "Question?", "reasoning_content": "noise"}], + ) + + result = cast( + List[AllMessageValues], + MistralConfig()._transform_messages( + messages=messages, model="mistral-medium-3-5" + ), + ) + + assert result[0].get("reasoning_content") == "noise" + + def test_reasoning_content_dropped_when_image_present(self): + """The image branch returns early, so stripping must run before it.""" + messages = cast( + List[AllMessageValues], + [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this"}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/cat.png"}, + }, + ], + }, + { + "role": "assistant", + "content": "A cat.", + "reasoning_content": "leaked reasoning", + }, + ], + ) + + with patch.object( + MistralConfig, + "_transform_messages_sync", + side_effect=lambda transformed, model: transformed, + ): + result = cast( + List[AllMessageValues], + MistralConfig()._transform_messages( + messages=messages, model="mistral-medium-3-5", is_async=False + ), + ) + + assert "reasoning_content" not in result[-1] diff --git a/tests/test_litellm/llms/openai_like/test_json_providers.py b/tests/test_litellm/llms/openai_like/test_json_providers.py index 39a4964f5f4..c8743e1809d 100644 --- a/tests/test_litellm/llms/openai_like/test_json_providers.py +++ b/tests/test_litellm/llms/openai_like/test_json_providers.py @@ -2,6 +2,7 @@ Tests for JSON-based provider configuration system. """ +import json import os import sys from unittest.mock import MagicMock, patch @@ -244,6 +245,99 @@ class TestPinstripes: assert result["temperature"] == 0.7 +class TestDarkbloom: + def test_darkbloom_json_config_exists(self): + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + darkbloom = JSONProviderRegistry.get("darkbloom") + assert darkbloom is not None + assert darkbloom.base_url == "https://api.darkbloom.dev/v1" + assert darkbloom.api_key_env == "DARKBLOOM_API_KEY" + assert darkbloom.api_base_env == "DARKBLOOM_API_BASE" + assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_darkbloom_provider_resolution(self): + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="darkbloom/gemma-4-26b", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "gemma-4-26b" + assert provider == "darkbloom" + assert api_key is None + assert api_base == "https://api.darkbloom.dev/v1" + + def test_darkbloom_dynamic_config(self): + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("darkbloom") + config_class = create_config_class(provider) + config = config_class() + + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://api.darkbloom.dev/v1" + + api_base, api_key = config._get_openai_compatible_provider_info( + "https://custom.darkbloom.dev/v1", "test-key" + ) + assert api_base == "https://custom.darkbloom.dev/v1" + assert api_key == "test-key" + + def test_darkbloom_complete_url_appends_endpoint(self): + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("darkbloom") + config_class = create_config_class(provider) + config = config_class() + + url = config.get_complete_url( + api_base="https://api.darkbloom.dev/v1", + api_key="test-key", + model="darkbloom/gemma-4-26b", + optional_params={}, + litellm_params={}, + stream=True, + ) + + assert url == "https://api.darkbloom.dev/v1/chat/completions" + + def test_darkbloom_provider_config_manager(self): + from litellm import LlmProviders + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_chat_config( + model="gemma-4-26b", provider=LlmProviders.DARKBLOOM + ) + + assert config is not None + assert config.custom_llm_provider == "darkbloom" + + def test_darkbloom_model_cost_map(self): + with open( + os.path.join(workspace_path, "model_prices_and_context_window.json") + ) as f: + model_cost = json.load(f) + + expected_models = { + "darkbloom/gemma-4-26b": (3e-08, 1.65e-07), + "darkbloom/gpt-oss-20b": (1.45e-08, 7e-08), + } + for model, (input_cost, output_cost) in expected_models.items(): + assert model in model_cost + assert model_cost[model]["litellm_provider"] == "darkbloom" + assert model_cost[model]["max_output_tokens"] == 32768 + assert model_cost[model]["supports_function_calling"] is True + assert model_cost[model]["supports_tool_choice"] is True + assert model_cost[model]["input_cost_per_token"] == input_cost + assert model_cost[model]["output_cost_per_token"] == output_cost + + class TestPublicAIIntegration: """Integration tests for PublicAI provider""" diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py index e2d1ab72c5e..46c1e457d7c 100644 --- a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py +++ b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py @@ -9,7 +9,7 @@ import json import math import os import sys -from unittest.mock import Mock, patch +from unittest.mock import patch import pytest @@ -120,10 +120,10 @@ class TestPerplexityCostCalculator: # Expected costs: # Input: 100 tokens * $2e-6 = $0.0002 # Output: 50 tokens * $8e-6 = $0.0004 - # Search: 3 queries * ($0.005 / 1000) = $0.000015 - # Total completion cost: $0.000415 + # Search: 3 queries * $0.005 per request = $0.015 + # Total completion cost: $0.0154 expected_prompt_cost = 100 * 2e-6 - expected_completion_cost = (50 * 8e-6) + (3 / 1000 * 0.005) + expected_completion_cost = (50 * 8e-6) + (3 * 0.005) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6) @@ -195,10 +195,10 @@ class TestPerplexityCostCalculator: # Total prompt cost = $0.00026 # Output (text): (50 - 15) tokens * $8e-6 = $0.00028 # Reasoning: 15 tokens * $3e-6 = $0.000045 - # Search: 2 queries * ($0.005 / 1000) = $0.00001 - # Total completion cost = $0.000335 + # Search: 2 queries * $0.005 per request = $0.01 + # Total completion cost = $0.010325 expected_prompt_cost = (100 * 2e-6) + (30 * 2e-6) - expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 / 1000 * 0.005) + expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 * 0.005) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6) @@ -311,7 +311,7 @@ class TestPerplexityCostCalculator: # Calculate expected total cost (reasoning is a subset of completion_tokens) expected_prompt_cost = (100 * 2e-6) + (15 * 2e-6) # Input + citation expected_completion_cost = ( - ((50 - 10) * 8e-6) + (10 * 3e-6) + (1 / 1000 * 0.005) + ((50 - 10) * 8e-6) + (10 * 3e-6) + (1 * 0.005) ) # Output (text) + reasoning + search expected_total = expected_prompt_cost + expected_completion_cost @@ -361,7 +361,7 @@ class TestPerplexityCostCalculator: expected_completion_cost = ( ((50 - reasoning_tokens) * 8e-6) + (reasoning_tokens * 3e-6) - + (search_queries / 1000 * 0.005) + + (search_queries * 0.005) ) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py index e59fbc9f272..8691e6a1ee5 100644 --- a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py +++ b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py @@ -9,7 +9,6 @@ import json import math import os import sys -from unittest.mock import Mock, patch import pytest @@ -106,8 +105,8 @@ class TestPerplexityIntegration: expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6) expected_completion_cost = ( - ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005) - ) + ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 * 0.005) + ) # Output (text) + reasoning + search expected_total = expected_prompt_cost + expected_completion_cost assert math.isclose(total_cost, expected_total, rel_tol=1e-6) @@ -152,8 +151,8 @@ class TestPerplexityIntegration: expected_prompt_cost = (200 * 2e-6) + (40 * 2e-6) expected_completion_cost = ( - ((100 - 25) * 8e-6) + (25 * 3e-6) + (3 / 1000 * 0.005) - ) + ((100 - 25) * 8e-6) + (25 * 3e-6) + (3 * 0.005) + ) # Output (text) + reasoning + search assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6) @@ -262,9 +261,9 @@ class TestPerplexityIntegration: expected_prompt_cost = (50000 * 2e-6) + (5000 * 2e-6) expected_completion_cost = ( - ((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 / 1000 * 0.005) - ) - expected_total = expected_prompt_cost + expected_completion_cost + ((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 * 0.005) + ) # $0.65 + expected_total = expected_prompt_cost + expected_completion_cost # $0.76 assert math.isclose(total_cost, expected_total, rel_tol=1e-6) assert total_cost > 0.25 @@ -326,7 +325,7 @@ class TestPerplexityIntegration: # Should calculate costs correctly expected_prompt_cost = (100 * 2e-6) + (10 * 2e-6) - expected_completion_cost = (50 * 8e-6) + (1 / 1000 * 0.005) + expected_completion_cost = (50 * 8e-6) + (1 * 0.005) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py index 468bd946ae9..d299239f68e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py @@ -41,9 +41,13 @@ class TestMask: def test_empty_returns_none_label(self): assert MCPDebug._mask("") == "(none)" - def test_short_value_unchanged(self): - # visible_prefix=6 + visible_suffix=4 = 10, so <= 10 chars unchanged - assert MCPDebug._mask("sk-1234") == "sk-1234" + def test_short_value_masked(self): + # Short auth values must not be echoed verbatim in debug headers, even though + # visible_prefix + visible_suffix would otherwise reveal the whole value. + masked = MCPDebug._mask("sk-1234") + assert "sk-1234" not in masked + assert set(masked) == {"*"} + assert len(masked) == len("sk-1234") def test_long_value_masked(self): result = MCPDebug._mask("Bearer eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py index 2558df8533b..cebc265a148 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py @@ -452,6 +452,430 @@ async def test_semantic_filter_hook_skips_no_tools(): print("✅ Hook correctly skips requests without tools") +@pytest.mark.asyncio +async def test_semantic_filter_hook_preserves_native_tools(): + """ + Regression test: mixed MCP + native tools. + + Given: 5 MCP tools (registered in _tool_map) + 2 native OpenAI-format + function tools (not in _tool_map) + When: The hook filters tools + Then: The native tools must survive unconditionally, and only MCP + tools go through the semantic filter. + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.types.utils import Embedding, EmbeddingResponse + + mock_router = Mock() + + def mock_embedding_sync(*args, **kwargs): + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync() + + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=2, + similarity_threshold=0.3, + enabled=True, + ) + + # --- MCP tools (registered in the semantic router) --- + mcp_tools = [ + MCPTool( + name=f"mcp_tool_{i}", + description=f"MCP tool {i}", + inputSchema={"type": "object"}, + ) + for i in range(5) + ] + filter_instance._build_router(mcp_tools) + + # --- Native OpenAI-format function tools (NOT in _tool_map) --- + native_tools = [ + { + "type": "function", + "function": { + "name": "get_current_weather", + "description": "Get the current weather", + "parameters": {"type": "object", "properties": {}}, + }, + }, + { + "type": "function", + "function": { + "name": "search_web", + "description": "Search the web", + "parameters": {"type": "object", "properties": {}}, + }, + }, + ] + + # Combine: MCP tools + native tools + all_tools = list(mcp_tools) + native_tools + + hook = SemanticToolFilterHook(filter_instance) + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "What is the weather?"}], + "tools": all_tools, + "metadata": {}, + } + + result = await hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=data, + call_type="completion", + ) + + assert result is not None, "Hook should return modified data" + filtered = result["tools"] + + # Native tools must survive + native_in_result = [ + t for t in filtered if isinstance(t, dict) and t.get("type") == "function" + ] + assert ( + len(native_in_result) == 2 + ), f"Both native tools must survive, got {len(native_in_result)}" + + # MCP tools should be filtered (top_k=2) + mcp_in_result = [t for t in filtered if not isinstance(t, dict)] + assert ( + len(mcp_in_result) <= 2 + ), f"MCP tools should be filtered to top_k=2, got {len(mcp_in_result)}" + + # Total should be native + filtered MCP + assert len(filtered) <= 4, f"Expected at most 4 tools, got {len(filtered)}" + + # Filter stats should be emitted (MCP tools were present) + assert "litellm_semantic_filter_stats" in result["metadata"] + + # Stats should report MCP-only counts, not inflated with native tools + stats = result["metadata"]["litellm_semantic_filter_stats"] + mcp_before, mcp_after = stats.split("->") + assert ( + int(mcp_before) == 5 + ), f"Stats 'from' should be MCP count (5), got {mcp_before}" + assert int(mcp_after) == len( + mcp_in_result + ), f"Stats 'to' should match filtered MCP count, got {mcp_after}" + + print( + f"✅ Hook preserves native tools: {len(all_tools)} -> {len(filtered)} " + f"({len(native_in_result)} native + {len(mcp_in_result)} MCP), " + f"stats={stats}" + ) + + +@pytest.mark.asyncio +async def test_semantic_filter_hook_all_native_tools(): + """ + Regression test: all-native request. + + Given: Only native OpenAI-format function tools (none registered in + the MCP semantic router) + When: The hook processes the request + Then: All tools pass through, and NO spurious semantic filter response + headers are emitted (no litellm_semantic_filter_stats in metadata). + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + + mock_router = Mock() + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=3, + similarity_threshold=0.3, + enabled=True, + ) + + # Build router with some MCP tools (so tool_router is not None) + mcp_tools = [ + MCPTool( + name="some_mcp_tool", + description="An MCP tool", + inputSchema={"type": "object"}, + ) + ] + + from litellm.types.utils import Embedding, EmbeddingResponse + + def mock_embedding_sync(*args, **kwargs): + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync() + + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + + filter_instance._build_router(mcp_tools) + + # --- Only native tools in the request --- + native_tools = [ + { + "type": "function", + "function": { + "name": f"native_func_{i}", + "description": f"Native function {i}", + "parameters": {"type": "object", "properties": {}}, + }, + } + for i in range(3) + ] + + hook = SemanticToolFilterHook(filter_instance) + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "tools": native_tools, + "metadata": {}, + } + + result = await hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=data, + call_type="completion", + ) + + assert result is not None, "Hook should return modified data" + filtered = result["tools"] + + # All native tools must pass through + assert ( + len(filtered) == 3 + ), f"All 3 native tools must pass through, got {len(filtered)}" + + # No spurious semantic filter stats (P2 fix) + assert ( + "litellm_semantic_filter_stats" not in result["metadata"] + ), "Should NOT emit semantic filter stats for all-native-tool requests" + + print( + f"✅ Hook passes through all {len(filtered)} native tools, " + f"no spurious filter headers emitted" + ) + + +@pytest.mark.asyncio +async def test_semantic_filter_hook_responses_api_name_collision(): + """ + Regression test: Responses API native tool with MCP-matching name. + + Given: A Responses-API native tool whose top-level ``name`` collides + with an MCP canonical name in ``_tool_map`` + When: The hook classifies tools + Then: The native tool must NOT be sent to the semantic filter, even + though its name matches an MCP canonical. + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.types.utils import Embedding, EmbeddingResponse + + mock_router = Mock() + + def mock_embedding_sync(*args, **kwargs): + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync() + + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=2, + similarity_threshold=0.3, + enabled=True, + ) + + # Register an MCP tool with name "github-search" + mcp_tools = [ + MCPTool( + name="github-search", + description="Search GitHub repos", + inputSchema={"type": "object"}, + ) + ] + filter_instance._build_router(mcp_tools) + + # Responses API native tool with SAME name as MCP canonical + responses_api_tool = { + "type": "function", + "name": "github-search", + "description": "Caller-owned search tool", + "parameters": {"type": "object"}, + } + + hook = SemanticToolFilterHook(filter_instance) + + # Verify classification: should be native, not MCP + assert not hook._is_mcp_tool(responses_api_tool), ( + "Responses API tool with type=function + top-level name " + "should be classified as native, not MCP" + ) + + # Full hook test: all-native request should preserve tools + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Search GitHub"}], + "tools": [responses_api_tool], + "metadata": {}, + } + + result = await hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=data, + call_type="completion", + ) + + # All tools are native → hook returns data with all tools preserved + filtered = (result or data)["tools"] + assert len(filtered) == 1, f"Native tool must survive, got {len(filtered)}" + assert filtered[0]["name"] == "github-search" + + print("✅ Responses API tool with MCP-matching name correctly classified as native") + + +@pytest.mark.asyncio +async def test_semantic_filter_hook_preserves_tool_order(): + """ + Regression test: tool ordering preservation. + + Given: An interleaved request [mcp_tool_A, native_tool, mcp_tool_B] + When: The hook filters tools (all MCP tools survive) + Then: The output order must match the original request order, + NOT native-first. + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.types.utils import Embedding, EmbeddingResponse + + mock_router = Mock() + + def mock_embedding_sync(*args, **kwargs): + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync() + + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=5, + similarity_threshold=0.3, + enabled=True, + ) + + # Register MCP tools + mcp_tool_a = MCPTool( + name="github-search", + description="Search GitHub", + inputSchema={"type": "object"}, + ) + mcp_tool_b = MCPTool( + name="github-issue", + description="Create GitHub issue", + inputSchema={"type": "object"}, + ) + filter_instance._build_router([mcp_tool_a, mcp_tool_b]) + + # Mock filter_tools to return both MCP tools (deterministic) + filter_instance.filter_tools = AsyncMock( # type: ignore[method-assign] + return_value=[mcp_tool_a, mcp_tool_b] + ) + + # Native tool (interleaved between MCP tools) + native_tool = { + "type": "function", + "function": { + "name": "weather_lookup", + "description": "Look up weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + + # Original order: [mcp_A, native, mcp_B] + original_tools = [mcp_tool_a, native_tool, mcp_tool_b] + + hook = SemanticToolFilterHook(filter_instance) + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Search GitHub and check weather"}], + "tools": original_tools, + "metadata": {}, + } + + result = await hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=data, + call_type="completion", + ) + + assert result is not None, "Hook should return modified data" + filtered = result["tools"] + + # All tools should survive + assert len(filtered) == 3, f"Expected 3 tools, got {len(filtered)}" + + # Order must be preserved: [mcp_A, native, mcp_B] + assert filtered[0] is mcp_tool_a, "First tool should be mcp_tool_a" + assert filtered[1] is native_tool, "Second tool should be native_tool" + assert filtered[2] is mcp_tool_b, "Third tool should be mcp_tool_b" + + print( + "✅ Tool ordering preserved: [mcp_A, native, mcp_B] maintained after filtering" + ) + + class TestGetToolsByNames: """ Regression coverage for SemanticMCPToolFilter._get_tools_by_names @@ -489,9 +913,7 @@ class TestGetToolsByNames: {"name": "send_email", "description": "send mail"}, ] - matched = filter_instance._get_tools_by_names( - ["send_email"], available_tools - ) + matched = filter_instance._get_tools_by_names(["send_email"], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "send_email" @@ -503,9 +925,7 @@ class TestGetToolsByNames: client_name = "litellm_" + canonical available_tools = [{"name": client_name, "description": "scrape"}] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 # Must return the incoming tool unchanged so the client-facing @@ -516,13 +936,9 @@ class TestGetToolsByNames: """Some clients use dash as alias separator; accept that too.""" filter_instance = self._make_filter() canonical = "weather_svc-get_weather" - available_tools = [ - {"name": "mcp-" + canonical, "description": "weather"} - ] + available_tools = [{"name": "mcp-" + canonical, "description": "weather"}] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "mcp-" + canonical @@ -552,9 +968,7 @@ class TestGetToolsByNames: {"name": "litellm_" + canonical, "description": "wrapped"}, ] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == canonical @@ -567,9 +981,7 @@ class TestGetToolsByNames: separator-anchored suffixes of ``litellm_api-fs-read_file``. """ filter_instance = self._make_filter() - available_tools = [ - {"name": "litellm_api-fs-read_file", "description": "read"} - ] + available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}] matched = filter_instance._get_tools_by_names( ["fs-read_file", "api-fs-read_file"], available_tools @@ -590,9 +1002,7 @@ class TestGetToolsByNames: {"name": "my_" + canonical, "description": "plain search"}, ] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "my_" + canonical diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 702bae77339..c8dc0ea5ed6 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -351,6 +351,43 @@ async def test_can_key_call_model_all_team_models_no_team_id_is_denied(): assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied +@pytest.mark.asyncio +async def test_can_team_access_model_all_team_models_expands_router_models(): + from litellm import Router + from litellm.proxy._types import SpecialModelNames + from litellm.proxy.auth.auth_checks import can_team_access_model + + team_object = LiteLLM_TeamTable( + team_id="team-123", + models=[SpecialModelNames.all_team_models.value], + ) + router = Router( + model_list=[ + { + "model_name": "allowed-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}, + } + ] + ) + + assert ( + await can_team_access_model( + model="allowed-model", + team_object=team_object, + llm_router=router, + ) + is True + ) + with pytest.raises(ProxyException) as exc_info: + await can_team_access_model( + model="blocked-model", + team_object=team_object, + llm_router=router, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + @pytest.mark.asyncio async def test_get_key_object_should_reconnect_once_on_db_connection_error(): mock_prisma_client = MagicMock() diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index 02b1f698132..261485e8965 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -543,3 +543,77 @@ async def test_get_available_models_for_user_expands_query_team_wildcard( ) assert "openai/gpt-4o-mini" in result + + +def test_get_key_models_all_team_models_recursive_team(): + """GH#30619: when key and team both have all-team-models, + the sentinel should expand to proxy_model_list.""" + from litellm.proxy.auth.model_checks import get_key_models + from litellm.proxy._types import SpecialModelNames + + user_api_key_dict = type( + "obj", (object,), + { + "models": [SpecialModelNames.all_team_models.value], + "team_id": "team-1", + "team_models": [SpecialModelNames.all_team_models.value], + }, + )() + proxy_model_list = ["model-a", "model-b"] + result = get_key_models(user_api_key_dict, proxy_model_list, {}) + assert SpecialModelNames.all_team_models.value not in result + assert set(result) == {"model-a", "model-b"} + + +def test_get_key_models_all_team_models_keeps_mixed_team_entries(): + from litellm.proxy.auth.model_checks import get_key_models + from litellm.proxy._types import SpecialModelNames + + user_api_key_dict = type( + "obj", + (object,), + { + "models": [SpecialModelNames.all_team_models.value], + "team_id": "team-1", + "team_models": [ + SpecialModelNames.all_team_models.value, + "restricted-model", + ], + }, + )() + result = get_key_models(user_api_key_dict, ["model-a", "model-b"], {}) + assert SpecialModelNames.all_team_models.value not in result + assert set(result) == {"model-a", "model-b", "restricted-model"} + + +def test_get_team_models_all_team_models_expands(): + """GH#30619: all-team-models in team_models should expand.""" + from litellm.proxy.auth.model_checks import get_team_models + from litellm.proxy._types import SpecialModelNames + + result = get_team_models( + [SpecialModelNames.all_team_models.value], + ["model-a", "model-b"], + {}, + ) + assert SpecialModelNames.all_team_models.value not in result + assert set(result) == {"model-a", "model-b"} + + +def test_get_team_models_all_team_models_expands_with_access_groups(): + """GH#30619: all-team-models with include_model_access_groups + should include access group keys.""" + from litellm.proxy.auth.model_checks import get_team_models + from litellm.proxy._types import SpecialModelNames + + result = get_team_models( + [SpecialModelNames.all_team_models.value], + ["model-a", "model-b"], + {"group-1": ["g1-model"], "group-2": ["g2-model"]}, + include_model_access_groups=True, + ) + assert SpecialModelNames.all_team_models.value not in result + assert "model-a" in result + assert "model-b" in result + assert "group-1" in result + assert "group-2" in result diff --git a/tests/test_litellm/proxy/db/test_db_url_settings.py b/tests/test_litellm/proxy/db/test_db_url_settings.py index b2212068a5b..573bd5ae584 100644 --- a/tests/test_litellm/proxy/db/test_db_url_settings.py +++ b/tests/test_litellm/proxy/db/test_db_url_settings.py @@ -16,7 +16,11 @@ from unittest.mock import patch import pytest -from litellm.proxy.db.db_url_settings import DatabaseURLSettings +from litellm.proxy.db.db_url_settings import ( + DatabaseURLSettings, + unsupported_db_scheme, + unsupported_db_scheme_message, +) def _apply() -> bool: @@ -27,6 +31,7 @@ def _apply() -> bool: _MANAGED_DB_ENV_VARS = ( "IAM_TOKEN_DB_AUTH", "DATABASE_URL", + "DIRECT_URL", "DATABASE_URL_READ_REPLICA", "DATABASE_HOST", "DATABASE_PORT", @@ -287,3 +292,87 @@ def test_password_reader_uses_own_credentials(monkeypatch): os.environ["DATABASE_URL_READ_REPLICA"] == "postgresql://litellm_ro:ro_pw@reader.example.com:5432/litellm_db" ) + + +@pytest.mark.parametrize( + "url", + [ + "postgresql://u:p@host:5432/db", + "postgres://u:p@host:5432/db", + "POSTGRESQL://u:p@host:5432/db", + "postgresql://host/db?schema=public", + ], +) +def test_unsupported_db_scheme_accepts_postgres(url): + assert unsupported_db_scheme(url) is None + + +@pytest.mark.parametrize( + "url,scheme", + [ + ("sqlite:///data/litellm.db", "sqlite"), + ("sqlite:///./local.db", "sqlite"), + ("mysql://u:p@host:3306/db", "mysql"), + ("mssql://host/db", "mssql"), + ], +) +def test_unsupported_db_scheme_rejects_non_postgres(url, scheme): + assert unsupported_db_scheme(url) == scheme + + +def test_unsupported_db_scheme_does_not_echo_schemeless_credentials(): + """A malformed schemeless DSN must not leak its embedded credentials + through the return value (which callers log).""" + leaky = "litellm:s3cr3t_password@db.internal:5432/litellm" + + result = unsupported_db_scheme(leaky) + + assert result is not None + assert "s3cr3t_password" not in result + assert "db.internal" not in result + + +def test_apply_to_env_rejects_pinned_sqlite_writer(monkeypatch): + """Componentized entrypoints pin DATABASE_URL and call apply_to_env; a + sqlite writer must raise here rather than reach Prisma.""" + monkeypatch.setenv("DATABASE_URL", "sqlite:///data/litellm.db") + + with pytest.raises(RuntimeError, match="sqlite"): + _apply() + + # The bad URL must not have been propagated as a usable connection string. + assert os.environ["DATABASE_URL"] == "sqlite:///data/litellm.db" + + +def test_apply_to_env_rejects_pinned_sqlite_direct_url(monkeypatch): + """DIRECT_URL reaches Prisma the same way DATABASE_URL does; a non-postgres + direct URL must be rejected in apply_to_env, matching the CLI startup guard.""" + monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db") + monkeypatch.setenv("DIRECT_URL", "sqlite:///data/litellm.db") + + with pytest.raises(RuntimeError, match="DIRECT_URL.*sqlite"): + _apply() + + +def test_apply_to_env_rejects_pinned_non_postgres_reader(monkeypatch): + monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db") + monkeypatch.setenv( + "DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db" + ) + + with pytest.raises(RuntimeError, match="DATABASE_URL_READ_REPLICA.*mysql"): + _apply() + + +def test_apply_to_env_accepts_pinned_postgres(monkeypatch): + monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@host:5432/db") + + # Operator-pinned URL: nothing reassembled, no error. + assert _apply() is False + + +def test_unsupported_db_scheme_message_names_var_and_scheme(): + msg = unsupported_db_scheme_message("DIRECT_URL", "sqlite") + assert "DIRECT_URL" in msg + assert "sqlite" in msg + assert "postgresql://" in msg diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index 33de1ede917..699606b5277 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -16,8 +16,6 @@ Pins covered: from __future__ import annotations import json -from typing import Any, AsyncIterator -from unittest.mock import AsyncMock, MagicMock import pytest @@ -26,8 +24,11 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import ( _apply_streaming_chunk_hooks, _fast_serialize_simple_model_response_stream, + _format_fallback_metadata_sse_event, _format_streaming_sse_chunk, _get_client_requested_model_for_streaming, + _get_streaming_fallback_metadata, + _is_positive_int_like, _restamp_streaming_chunk_model, _serialize_streaming_chunk, async_assistants_data_generator, @@ -71,6 +72,15 @@ async def _async_iter_raises(exc: Exception): raise exc +class _FakeStream: + def __init__(self, chunks, hidden_params=None): + self._chunks = chunks + self._hidden_params = hidden_params or {} + + def __aiter__(self): + return _async_iter(self._chunks) + + # --------------------------------------------------------------------------- # data_generator # --------------------------------------------------------------------------- @@ -274,6 +284,34 @@ def test_restamp_streaming_chunk_model_overrides_model_on_dict(): assert logged is True +def test_restamp_streaming_chunk_model_uses_fallback_model_from_metadata(): + chunk = _simple_chunk(model="openai/internal-fallback") + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="primary-model", + request_data={"litellm_call_id": "id-1"}, + model_mismatch_logged=False, + fallback_was_attempted=True, + fallback_model_from_metadata="fallback-model", + ) + assert new_chunk.model == "fallback-model" + assert logged is True + + +def test_restamp_streaming_chunk_model_preserves_fallback_model_without_group(): + chunk = _simple_chunk(model="openai/internal-fallback") + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="primary-model", + request_data={}, + model_mismatch_logged=False, + fallback_was_attempted=True, + fallback_model_from_metadata=None, + ) + assert new_chunk.model == "openai/internal-fallback" + assert logged is False + + def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged(): """For a non-BaseModel, non-dict chunk the helper returns it as-is along with the original ``model_mismatch_logged`` flag.""" @@ -288,6 +326,147 @@ def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged(): assert logged is False +def test_is_positive_int_like_invalid_and_edge_values(): + assert _is_positive_int_like(None) is False + assert _is_positive_int_like("not-a-number") is False + assert _is_positive_int_like(0) is False + assert _is_positive_int_like(-1) is False + assert _is_positive_int_like("1") is True + assert _is_positive_int_like(2) is True + + +def test_get_streaming_fallback_metadata_reads_headers(): + fallback_errors = [ + { + "message": "litellm.RateLimitError: upstream limited request", + "type": "RateLimitError", + "param": None, + "code": "429", + } + ] + stream = _FakeStream( + [], + hidden_params={ + "additional_headers": { + "x-litellm-attempted-fallbacks": "1", + "x-litellm-model-group": "fallback-model", + "x-litellm-fallback-errors": json.dumps(fallback_errors), + } + }, + ) + assert _get_streaming_fallback_metadata(stream) == ( + True, + "fallback-model", + fallback_errors, + ) + + +def test_get_streaming_fallback_metadata_no_additional_headers(): + stream = _FakeStream([], hidden_params={}) + assert _get_streaming_fallback_metadata(stream) == (False, None, []) + + +def test_get_streaming_fallback_metadata_zero_fallback_count(): + stream = _FakeStream( + [], + hidden_params={ + "additional_headers": {"x-litellm-attempted-fallbacks": 0} + }, + ) + assert _get_streaming_fallback_metadata(stream) == (False, None, []) + + +def test_get_streaming_fallback_metadata_no_model_group_returns_none_model(): + stream = _FakeStream( + [], + hidden_params={ + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + } + }, + ) + was_attempted, fallback_model, errors = _get_streaming_fallback_metadata(stream) + assert was_attempted is True + assert fallback_model is None + assert errors == [] + + +def test_restamp_streaming_chunk_model_azure_router_preserves_model(): + chunk = _simple_chunk(model="azure_ai/internal-deployment") + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="azure_ai/model-router", + request_data={}, + model_mismatch_logged=False, + ) + assert new_chunk.model == "azure_ai/internal-deployment" + assert logged is False + + +def test_restamp_streaming_chunk_model_fastest_response_preserves_model(): + chunk = _simple_chunk(model="winning-model") + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="gpt-4,claude-3", + request_data={"fastest_response": True}, + model_mismatch_logged=False, + ) + assert new_chunk.model == "winning-model" + assert logged is False + + +def test_restamp_streaming_chunk_model_setattr_exception_logs_and_returns(): + from pydantic import ConfigDict + + class FrozenChunk(_simple_chunk().__class__): + model_config = ConfigDict(frozen=True) + + chunk = FrozenChunk( + id="chatcmpl-test", + choices=[], + created=0, + model="openai/internal-x", + object="chat.completion.chunk", + ) + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="gpt-4", + request_data={"litellm_call_id": "test-id"}, + model_mismatch_logged=False, + ) + assert new_chunk.model == "openai/internal-x" + assert logged is True + + +def test_format_fallback_metadata_sse_event(): + fallback_errors = [ + { + "message": "litellm.RateLimitError: upstream limited request", + "type": "RateLimitError", + "param": None, + "code": "429", + } + ] + + event = _format_fallback_metadata_sse_event( + fallback_model="fallback-model", + fallback_errors=fallback_errors, + ) + + assert isinstance(event, str) + assert event.startswith("data: ") + payload = json.loads(event.removeprefix("data: ").removesuffix("\n\n")) + assert payload["choices"] == [] + assert payload["litellm_fallback"] == { + "fallback_model": "fallback-model", + "errors": fallback_errors, + } + assert payload["id"] == "litellm-fallback-metadata" + assert payload["object"] == "chat.completion.chunk" + assert payload["model"] == "fallback-model" + assert isinstance(payload["created"], int) + + # --------------------------------------------------------------------------- # _fast_serialize_simple_model_response_stream # --------------------------------------------------------------------------- @@ -473,7 +652,7 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch): # First chunk is bytes (fast path) wrapped via _format_streaming_sse_chunk. first = out[0] assert isinstance(first, bytes) - payload = json.loads(first.removeprefix(b"data: ").rstrip(b"\n\n")) + payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n")) assert normalize(payload) == { "id": "", "object": "chat.completion.chunk", @@ -488,6 +667,172 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch): } +@pytest.mark.asyncio +async def test_async_data_generator_uses_response_fallback_metadata(monkeypatch): + _patch_logging_flags(monkeypatch) + + response = _FakeStream( + [_simple_chunk(model="openai/internal-fallback", content="hello")], + hidden_params={ + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + } + }, + ) + out = [] + async for line in async_data_generator( + response=response, + user_api_key_dict=_user_auth(), + request_data={"model": "primary-model", "include_fallback_errors": True}, + ): + out.append(line) + + first = out[0] + assert isinstance(first, bytes) + payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n")) + assert payload["model"] == "fallback-model" + + +@pytest.mark.asyncio +async def test_async_data_generator_uses_chunk_fallback_metadata(monkeypatch): + _patch_logging_flags(monkeypatch) + + chunk = _simple_chunk(model="openai/internal-fallback", content="hello") + chunk._hidden_params = { + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + } + } + out = [] + async for line in async_data_generator( + response=_async_iter([chunk]), + user_api_key_dict=_user_auth(), + request_data={"model": "primary-model"}, + ): + out.append(line) + + first = out[0] + assert isinstance(first, bytes) + payload = json.loads(first.removeprefix(b"data: ").removesuffix(b"\n\n")) + assert payload["model"] == "fallback-model" + + +@pytest.mark.asyncio +async def test_async_data_generator_switches_model_mid_stream_on_fallback(monkeypatch): + """Pre-fallback chunks keep the client-requested model; once a chunk carries + fallback metadata the model latches to the fallback group for the rest of the + stream. This pins the client-visible mid-stream model change.""" + _patch_logging_flags(monkeypatch) + + primary_chunk = _simple_chunk(model="openai/internal-primary", content="hi") + fallback_chunk = _simple_chunk(model="openai/internal-fallback", content="there") + fallback_chunk._hidden_params = { + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + } + } + out = [] + async for line in async_data_generator( + response=_async_iter([primary_chunk, fallback_chunk]), + user_api_key_dict=_user_auth(), + request_data={"model": "primary-model"}, + ): + out.append(line) + + first_payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n")) + second_payload = json.loads(out[1].removeprefix(b"data: ").removesuffix(b"\n\n")) + assert first_payload["model"] == "primary-model" + assert second_payload["model"] == "fallback-model" + + +@pytest.mark.asyncio +async def test_async_data_generator_emits_fallback_error_metadata_event(monkeypatch): + _patch_logging_flags(monkeypatch) + monkeypatch.setitem(ps.general_settings, "expose_fallback_errors_to_caller", True) + + fallback_errors = [ + { + "message": "litellm.RateLimitError: upstream limited request", + "type": "RateLimitError", + "param": None, + "code": "429", + } + ] + response = _FakeStream( + [_simple_chunk(model="openai/internal-fallback", content="hello")], + hidden_params={ + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + "x-litellm-fallback-errors": json.dumps(fallback_errors), + } + }, + ) + out = [] + async for line in async_data_generator( + response=response, + user_api_key_dict=_user_auth(), + request_data={"model": "primary-model", "include_fallback_errors": True}, + ): + out.append(line) + + assert isinstance(out[0], bytes) + chunk_payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n")) + assert chunk_payload["model"] == "fallback-model" + assert isinstance(out[1], str) + assert out[1].startswith("data: ") + metadata_payload = json.loads(out[1].removeprefix("data: ").removesuffix("\n\n")) + assert metadata_payload["choices"] == [] + assert metadata_payload["litellm_fallback"] == { + "fallback_model": "fallback-model", + "errors": fallback_errors, + } + assert metadata_payload["id"] == "litellm-fallback-metadata" + assert metadata_payload["object"] == "chat.completion.chunk" + assert metadata_payload["model"] == "fallback-model" + assert isinstance(metadata_payload["created"], int) + + +@pytest.mark.asyncio +async def test_async_data_generator_skips_fallback_error_event_without_opt_in( + monkeypatch, +): + _patch_logging_flags(monkeypatch) + + fallback_errors = [ + { + "message": "litellm.RateLimitError: upstream limited request", + "type": "RateLimitError", + "param": None, + "code": "429", + } + ] + response = _FakeStream( + [_simple_chunk(model="openai/internal-fallback", content="hello")], + hidden_params={ + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + "x-litellm-fallback-errors": json.dumps(fallback_errors), + } + }, + ) + out = [] + async for line in async_data_generator( + response=response, + user_api_key_dict=_user_auth(), + request_data={"model": "primary-model"}, + ): + out.append(line) + + assert isinstance(out[0], bytes) + payload = json.loads(out[0].removeprefix(b"data: ").removesuffix(b"\n\n")) + assert payload["model"] == "fallback-model" + + @pytest.mark.asyncio async def test_async_data_generator_mid_stream_exception_yields_error_payload( monkeypatch, diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 56627c5be88..88dbec4020f 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -1708,6 +1708,57 @@ class TestRunServerDbSetup: use_migrate=True, use_v2_resolver=False ) + @patch("subprocess.run") + @patch("atexit.register") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") + @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") + def test_startup_exits_on_non_postgres_database_url( + self, + mock_should_update_schema, + mock_check_schema_diff, + mock_setup_database, + mock_atexit_register, + mock_subprocess_run, + ): + """A sqlite DATABASE_URL must exit immediately, before any prisma call, + instead of stalling on a migration against the postgresql-only schema.""" + from litellm.proxy.proxy_cli import run_server + + mock_subprocess_run.return_value = MagicMock(returncode=0) + mock_should_update_schema.return_value = True + + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + + clean_env = { + k: v + for k, v in os.environ.items() + if k not in ("DATABASE_URL", "DIRECT_URL") + } + clean_env["DATABASE_URL"] = "sqlite:///data/litellm.db" + + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + { + "proxy_server": mock_proxy_module, + "litellm.proxy.proxy_server": mock_proxy_module, + }, + ), + ): + with pytest.raises(SystemExit) as exc_info: + run_server.main( + ["--local", "--skip_server_startup"], standalone_mode=False + ) + assert exc_info.value.code == 1 + mock_setup_database.assert_not_called() + # --- Module-level helpers for worker startup hook tests --- diff --git a/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py b/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py new file mode 100644 index 00000000000..2aee0f0a4ef --- /dev/null +++ b/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py @@ -0,0 +1,142 @@ +import json + +from pydantic import BaseModel + +from litellm.router_utils.add_retry_fallback_headers import ( + add_fallback_headers_to_response, + add_retry_headers_to_response, + get_fallback_errors_from_headers, + get_hidden_params_dict, +) + + +class StreamingWrapper: + def __init__(self): + self._hidden_params = {"additional_headers": {"x-existing": "keep"}} + + +def test_add_fallback_headers_to_streaming_wrapper(): + response = StreamingWrapper() + + result = add_fallback_headers_to_response( + response=response, + attempted_fallbacks=1, + ) + + assert result is response + assert response._hidden_params["additional_headers"] == { + "x-existing": "keep", + "x-litellm-attempted-fallbacks": 1, + } + + +def test_add_fallback_headers_serializes_fallback_errors(): + response = StreamingWrapper() + fallback_errors = [ + { + "message": "litellm.RateLimitError: upstream limited request", + "type": "RateLimitError", + "param": None, + "code": "429", + } + ] + + result = add_fallback_headers_to_response( + response=response, + attempted_fallbacks=1, + fallback_errors=fallback_errors, + ) + + assert result is response + assert response._hidden_params["additional_headers"][ + "x-litellm-attempted-fallbacks" + ] == 1 + assert ( + json.loads( + response._hidden_params["additional_headers"]["x-litellm-fallback-errors"] + ) + == fallback_errors + ) + + +def test_add_retry_headers_to_streaming_wrapper(): + response = StreamingWrapper() + + result = add_retry_headers_to_response( + response=response, + attempted_retries=2, + max_retries=3, + ) + + assert result is response + assert response._hidden_params["additional_headers"] == { + "x-existing": "keep", + "x-litellm-attempted-retries": 2, + "x-litellm-max-retries": 3, + } + + +def test_get_hidden_params_dict_with_pydantic_model_hidden_params(): + class InnerHiddenParams(BaseModel): + additional_headers: dict = {} + + class Response: + def __init__(self): + self._hidden_params = InnerHiddenParams( + additional_headers={"x-custom": "value"} + ) + + result = get_hidden_params_dict(Response()) + assert result == {"additional_headers": {"x-custom": "value"}} + + +def test_get_hidden_params_dict_with_no_hidden_params(): + class PlainResponse: + pass + + assert get_hidden_params_dict(PlainResponse()) == {} + + +def test_add_fallback_headers_when_no_existing_additional_headers(): + class NoHeadersWrapper: + def __init__(self): + self._hidden_params = {} + + response = NoHeadersWrapper() + result = add_fallback_headers_to_response(response=response, attempted_fallbacks=2) + + assert result is response + assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 2 + + +def test_add_fallback_headers_returns_none_when_response_is_none(): + result = add_fallback_headers_to_response(response=None, attempted_fallbacks=1) + assert result is None + + +def test_add_fallback_headers_returns_unchanged_when_response_has_no_hidden_params(): + class PlainObject: + pass + + obj = PlainObject() + result = add_fallback_headers_to_response(response=obj, attempted_fallbacks=1) + assert result is obj + assert not hasattr(obj, "_hidden_params") + + +def test_get_fallback_errors_from_headers_existing_list_passthrough(): + errors = [{"message": "err", "type": "T", "param": None, "code": "400"}] + result = get_fallback_errors_from_headers({"x-litellm-fallback-errors": errors}) + assert result == errors + + +def test_get_fallback_errors_from_headers_invalid_json_returns_empty(): + result = get_fallback_errors_from_headers( + {"x-litellm-fallback-errors": "not-valid-json-{"} + ) + assert result == [] + + +def test_get_fallback_errors_from_headers_missing_key_returns_empty(): + result = get_fallback_errors_from_headers({}) + assert result == [] diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py new file mode 100644 index 00000000000..ca647bdce55 --- /dev/null +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -0,0 +1,139 @@ +import json + +import pytest + +from litellm.router_utils.fallback_event_handlers import run_async_fallback + + +class StreamingWrapper: + def __init__(self): + self._hidden_params = {"additional_headers": {}} + + +class FakeRouter: + def log_retry(self, kwargs, e): + return kwargs + + async def async_function_with_fallbacks(self, *args, **kwargs): + return StreamingWrapper() + + +class AlwaysFailRouter: + def log_retry(self, kwargs, e): + return kwargs + + async def async_function_with_fallbacks(self, *args, **kwargs): + raise RuntimeError("fallback model also failed") + + +@pytest.mark.asyncio +async def test_run_async_fallback_adds_errors_when_opted_in(): + response = await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + include_fallback_errors=True, + ) + + additional_headers = response._hidden_params["additional_headers"] + assert additional_headers["x-litellm-attempted-fallbacks"] == 1 + assert json.loads(additional_headers["x-litellm-fallback-errors"]) == [ + { + "message": "upstream limited request", + "type": "RuntimeError", + "param": None, + "code": None, + } + ] + + +@pytest.mark.asyncio +async def test_run_async_fallback_omits_errors_without_opt_in(): + response = await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + ) + + additional_headers = response._hidden_params["additional_headers"] + assert additional_headers["x-litellm-attempted-fallbacks"] == 1 + assert "x-litellm-fallback-errors" not in additional_headers + + +@pytest.mark.asyncio +async def test_run_async_fallback_raises_when_all_fallbacks_fail(): + with pytest.raises(RuntimeError, match="fallback model also failed"): + await run_async_fallback( + litellm_router=AlwaysFailRouter(), + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("original request failed"), + max_fallbacks=3, + fallback_depth=0, + include_fallback_errors=True, + ) + + +class RecordingRouter: + def __init__(self): + self.received_kwargs = None + + def log_retry(self, kwargs, e): + return kwargs + + async def async_function_with_fallbacks(self, *args, **kwargs): + self.received_kwargs = kwargs + return StreamingWrapper() + + +@pytest.mark.asyncio +async def test_run_async_fallback_forwards_include_fallback_errors_to_nested_call(): + """A nested fallback (multi-hop) must keep collecting errors, so the opt-in + flag has to reach the nested async_function_with_fallbacks call.""" + router = RecordingRouter() + await run_async_fallback( + litellm_router=router, + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + include_fallback_errors=True, + ) + + assert router.received_kwargs.get("include_fallback_errors") is True + + +@pytest.mark.asyncio +async def test_run_async_fallback_does_not_forward_flag_without_opt_in(): + router = RecordingRouter() + await run_async_fallback( + litellm_router=router, + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + ) + + assert "include_fallback_errors" not in router.received_kwargs + + +@pytest.mark.asyncio +async def test_run_async_fallback_skips_original_model_group(): + response = await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=["primary-model", "fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("original failed"), + max_fallbacks=3, + fallback_depth=0, + ) + + assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1 diff --git a/tests/test_litellm/test_router_streaming_fallback_metadata.py b/tests/test_litellm/test_router_streaming_fallback_metadata.py new file mode 100644 index 00000000000..6ed70dc7cfe --- /dev/null +++ b/tests/test_litellm/test_router_streaming_fallback_metadata.py @@ -0,0 +1,187 @@ +import json +from unittest.mock import MagicMock + +import pytest + +import litellm +from litellm.proxy.proxy_server import _should_include_fallback_errors +from litellm.router import Router +from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict + + +def test_apply_fallback_hidden_params_copies_from_fallback_response(): + fallback_errors = [ + { + "message": "litellm.RateLimitError: upstream limited request", + "type": "RateLimitError", + "param": None, + "code": "429", + } + ] + chunk = litellm.ModelResponseStream( + id="test", + model="openai/internal-fallback", + choices=[], + ) + chunk._hidden_params = { + "additional_headers": {"x-existing-chunk-header": "keep"}, + "model_id": "chunk-model-id", + } + fallback_response = MagicMock() + fallback_response._hidden_params = { + "additional_headers": { + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + "x-litellm-fallback-errors": json.dumps(fallback_errors), + }, + "api_base": "https://fallback.example", + } + + Router._apply_fallback_hidden_params_to_item( + fallback_item=chunk, + prepared_fallback_hidden_params=Router._prepare_fallback_hidden_params( + fallback_response + ), + ) + + assert chunk._hidden_params["api_base"] == "https://fallback.example" + assert chunk._hidden_params["model_id"] == "chunk-model-id" + assert chunk._hidden_params["additional_headers"] == { + "x-existing-chunk-header": "keep", + "x-litellm-attempted-fallbacks": 1, + "x-litellm-model-group": "fallback-model", + "x-litellm-fallback-errors": json.dumps(fallback_errors), + } + + +def _two_group_fallback_router() -> Router: + return litellm.Router( + model_list=[ + { + "model_name": "primary-model", + "litellm_params": {"model": "openai/gpt-fake", "api_key": "sk-fake"}, + }, + { + "model_name": "fallback-model", + "litellm_params": {"model": "openai/gpt-fake-2", "api_key": "sk-fake"}, + }, + ], + fallbacks=[{"primary-model": ["fallback-model"]}], + ) + + +def _additional_headers(response: object) -> dict: + return get_hidden_params_dict(response).get("additional_headers", {}) + + +@pytest.mark.asyncio +async def test_include_fallback_errors_propagates_through_router(): + router = _two_group_fallback_router() + + response = await router.acompletion( + model="primary-model", + messages=[{"role": "user", "content": "Hello"}], + mock_testing_fallbacks=True, + mock_response="fallback success", + include_fallback_errors=True, + ) + + headers = _additional_headers(response) + assert headers["x-litellm-attempted-fallbacks"] == 1 + errors = json.loads(headers["x-litellm-fallback-errors"]) + assert isinstance(errors, list) and len(errors) >= 1 + assert set(errors[0].keys()) == {"message", "type", "param", "code"} + + +@pytest.mark.asyncio +async def test_router_omits_fallback_errors_without_opt_in(): + router = _two_group_fallback_router() + + response = await router.acompletion( + model="primary-model", + messages=[{"role": "user", "content": "Hello"}], + mock_testing_fallbacks=True, + mock_response="fallback success", + ) + + headers = _additional_headers(response) + assert headers["x-litellm-attempted-fallbacks"] == 1 + assert "x-litellm-fallback-errors" not in headers + + +def test_prepare_fallback_hidden_params_no_additional_headers(): + class FakeResponse: + _hidden_params = {"api_base": "http://example.com"} + + hidden_params, headers = Router._prepare_fallback_hidden_params(FakeResponse()) + assert hidden_params == {"api_base": "http://example.com"} + assert headers == {} + + +def test_apply_fallback_hidden_params_to_item_none_item(): + Router._apply_fallback_hidden_params_to_item( + None, ({"api_base": "http://fallback.example"}, {"x-custom": "value"}) + ) + + +def test_apply_fallback_hidden_params_to_item_no_existing_additional_headers(): + class FakeChunk: + _hidden_params = {"model_id": "test-id"} + + chunk = FakeChunk() + Router._apply_fallback_hidden_params_to_item( + chunk, + ( + {"api_base": "http://fallback.example"}, + {"x-litellm-attempted-fallbacks": 1}, + ), + ) + + assert chunk._hidden_params["api_base"] == "http://fallback.example" + assert chunk._hidden_params["model_id"] == "test-id" + assert chunk._hidden_params["additional_headers"] == { + "x-litellm-attempted-fallbacks": 1 + } + + +@pytest.mark.asyncio +async def test_set_response_headers_adds_model_group_to_streaming_wrapper(): + class StreamingWrapper: + def __init__(self): + self._hidden_params = {"additional_headers": {"x-existing": "keep"}} + + router = litellm.Router(model_list=[]) + response = StreamingWrapper() + + result = await router.set_response_headers( + response=response, + model_group="fallback-model", + ) + + assert result is response + assert response._hidden_params["additional_headers"] == { + "x-existing": "keep", + "x-litellm-model-group": "fallback-model", + } + + +def test_should_include_fallback_errors_gated_by_operator_setting(): + request_data: dict = {"include_fallback_errors": True} + + import litellm.proxy.proxy_server as ps + + original = ps.general_settings.copy() if isinstance(ps.general_settings, dict) else {} + try: + ps.general_settings = {} + assert _should_include_fallback_errors(request_data) is False + + ps.general_settings = {"expose_fallback_errors_to_caller": False} + assert _should_include_fallback_errors(request_data) is False + + ps.general_settings = {"expose_fallback_errors_to_caller": True} + assert _should_include_fallback_errors(request_data) is True + + ps.general_settings = {"expose_fallback_errors_to_caller": True} + assert _should_include_fallback_errors({}) is False + finally: + ps.general_settings = original From 23808c2a094d36833956b8ab28963ba9ee8d8af0 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 23 Jun 2026 08:40:17 -0700 Subject: [PATCH 03/29] fix(typing): bring reportReturnType back under the basedpyright budget (#31103) --- litellm/main.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 1d75766c7e6..4fade2ac4b0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5044,9 +5044,7 @@ def completion( # type: ignore if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( tools=tools_for_mcp ): - # Return coroutine - acompletion will await it - # completion() can return a coroutine when MCP tools are present, which acompletion() awaits - return acompletion_with_mcp( # type: ignore[return-value] + return acompletion_with_mcp( # pyright: ignore[reportReturnType] # MCP path returns a coroutine that acompletion() awaits; completion()'s sync return type omits it model=model, messages=messages, functions=functions, @@ -5218,12 +5216,16 @@ def completion( # type: ignore logging: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, litellm_logging_obj) fallbacks = fallbacks or litellm.model_fallbacks if fallbacks is not None: - return completion_with_fallbacks(**args) + return completion_with_fallbacks( # pyright: ignore[reportReturnType] # fallback runner is untyped; resolves to ModelResponse|CustomStreamWrapper at runtime + **args + ) if model_list is not None: deployments = [ m["litellm_params"] for m in model_list if m["model_name"] == model ] - return litellm.batch_completion_models(deployments=deployments, **args) + return litellm.batch_completion_models( # pyright: ignore[reportReturnType] # batch path returns a list of responses, outside completion()'s single-response return type + deployments=deployments, **args + ) if litellm.model_alias_map and model in litellm.model_alias_map: model = litellm.model_alias_map[ model @@ -5545,7 +5547,7 @@ def completion( # type: ignore else: optional_params["reasoning_effort"] = {"summary": rs_val} - return responses_api_bridge.completion( + return responses_api_bridge.completion( # pyright: ignore[reportReturnType] # bridge returns a coroutine on the acompletion path; awaited by the async caller model=model, messages=messages, headers=headers, From e73cbfb0268283f017441c86842658b69a88c2c9 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 23 Jun 2026 21:10:52 +0530 Subject: [PATCH 04/29] fix(realtime): post-tool-call function_response id omission (#30446) --- litellm/llms/custom_httpx/llm_http_handler.py | 4 +- .../llms/gemini/realtime/transformation.py | 11 +-- .../llms/vertex_ai/realtime/transformation.py | 3 + .../test_vertex_ai_realtime_transformation.py | 71 +++++++++++++++++++ 4 files changed, 82 insertions(+), 7 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 790bd0519d7..129e15a0bf2 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5537,9 +5537,7 @@ class BaseLLMHTTPHandler: import websockets from websockets.asyncio.client import ClientConnection - url = self._append_query_params( - provider_config.get_complete_url(api_base, model, api_key), query_params - ) + url = provider_config.get_complete_url(api_base, model, api_key) headers = provider_config.validate_environment( headers=headers, model=model, diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 74f6cd4d831..e153d00e6ab 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -103,6 +103,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): # bypassing spend and budget accounting. self._pending_usage_metadata: Optional[dict] = None + def _include_function_response_id(self) -> bool: + """Google AI Studio Gemini 3.5+ accepts ``id`` on functionResponses; Vertex AI rejects it.""" + return True + @staticmethod def _usage_detail_alias(details: Any, defaults: Dict[str, int]) -> Dict[str, Any]: if not isinstance(details, dict): @@ -604,10 +608,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) # Build Gemini toolResponse format - function_response = { - "id": call_id, - "response": output_dict, - } + function_response: dict[str, Any] = {"response": output_dict} + if self._include_function_response_id() and call_id: + function_response["id"] = call_id if function_name: function_response["name"] = function_name diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index d6441db7856..1fe9f15c9f0 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -32,6 +32,9 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): self._project = project self._location = location + def _include_function_response_id(self) -> bool: + return False + # ------------------------------------------------------------------ # URL # ------------------------------------------------------------------ diff --git a/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py b/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py index 1ebd704be34..1f171496cce 100644 --- a/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py @@ -346,3 +346,74 @@ def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog) "Vertex AI Realtime" in record.message and "session.update" in record.message for record in caplog.records ) + + +async def test_async_realtime_does_not_forward_client_query_params_to_vertex_backend( + monkeypatch, +): + """Regression: forwarding client ?model=/?intent= to the Vertex Live WSS URL causes 1007 errors. + + Exercises ``async_realtime`` end-to-end so that re-adding ``_append_query_params`` + (the reverted bug) would push ``model=``/``intent=`` onto the backend URL and fail here. + """ + import websockets + + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + + captured = {} + + def fake_connect(url, *args, **kwargs): + captured["url"] = url + raise RuntimeError("stop before establishing the backend connection") + + monkeypatch.setattr(websockets, "connect", fake_connect) + + await BaseLLMHTTPHandler().async_realtime( + model="gemini-live-2.5-flash-native-audio", + websocket=AsyncMock(), + logging_obj=MagicMock(), + provider_config=cfg, + headers={}, + query_params={ + "model": "gemini-live-2.5-flash-native-audio", + "intent": "chat", + }, + ) + + assert "?" not in captured["url"] + assert "model=" not in captured["url"] + assert "intent=" not in captured["url"] + + +def test_vertex_function_call_output_omits_id(): + """Regression: Vertex Live rejects ``id`` on toolResponse.functionResponses (1007).""" + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + cfg._tool_call_id_to_name["call_abc123"] = "terminate_call" + + messages = cfg.transform_realtime_request( + json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "function_call_output", + "call_id": "call_abc123", + "output": '{"status": "ok"}', + }, + } + ), + "gemini-live-2.5-flash-native-audio", + session_configuration_request="existing", + ) + + assert len(messages) == 1 + payload = json.loads(messages[0]) + function_response = payload["toolResponse"]["functionResponses"][0] + assert "id" not in function_response + assert function_response["name"] == "terminate_call" + assert function_response["response"] == {"status": "ok"} From c8a9618afdacdba7146e997804fc103b48319826 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 23 Jun 2026 09:05:13 -0700 Subject: [PATCH 05/29] feat: add opensandbox sandbox provider (#31024) * feat: add opensandbox sandbox provider * fix: harden opensandbox sandbox startup * fix: address opensandbox review feedback * fix: address opensandbox sandbox review feedback * fix: address sandbox parser nits * fix(ci): clear opensandbox gates * fix(review): require opensandbox api base * chore(ci): rerun pass-through check --- .github/workflows/test-unit-misc.yml | 1 + litellm/constants.py | 12 + .../llms/base_llm/sandbox/transformation.py | 19 +- litellm/llms/custom_httpx/llm_http_handler.py | 182 ++++- litellm/llms/e2b/sandbox/transformation.py | 30 +- litellm/llms/opensandbox/__init__.py | 1 + litellm/llms/opensandbox/sandbox/__init__.py | 1 + .../opensandbox/sandbox/transformation.py | 598 ++++++++++++++++ litellm/responses/main.py | 5 +- litellm/sandbox/main.py | 2 +- litellm/types/utils.py | 1 + litellm/utils.py | 5 + provider_endpoints_support.json | 17 + .../custom_httpx/test_llm_http_handler.py | 119 ++++ .../sandbox/test_opensandbox_sandbox.py | 647 ++++++++++++++++++ 15 files changed, 1598 insertions(+), 42 deletions(-) create mode 100644 litellm/llms/opensandbox/__init__.py create mode 100644 litellm/llms/opensandbox/sandbox/__init__.py create mode 100644 litellm/llms/opensandbox/sandbox/transformation.py create mode 100644 tests/test_litellm/sandbox/test_opensandbox_sandbox.py diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index a7363ac3b43..c29c2d632f2 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -33,6 +33,7 @@ jobs: tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough + tests/test_litellm/sandbox tests/test_litellm/vector_stores tests/test_litellm/test_*.py workers: 2 diff --git a/litellm/constants.py b/litellm/constants.py index 083e9a1241b..b3e971a8cc1 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -201,6 +201,18 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( # Provider-specific API base URLs XAI_API_BASE = "https://api.x.ai/v1" +OPEN_SANDBOX_API_BASE_ENV_VAR = "OPEN_SANDBOX_API_BASE" +OPEN_SANDBOX_API_KEY_ENV_VAR = "OPEN_SANDBOX_API_KEY" +OPEN_SANDBOX_DEFAULT_TEMPLATE = "opensandbox/code-interpreter:v1.1.0" +_OPEN_SANDBOX_FALLBACK_ENTRYPOINT = "/opt/code-interpreter/code-interpreter.sh" +OPEN_SANDBOX_DEFAULT_ENTRYPOINT = (_OPEN_SANDBOX_FALLBACK_ENTRYPOINT,) +OPEN_SANDBOX_DEFAULT_LANGUAGE = "python" +OPEN_SANDBOX_DEFAULT_CPU_LIMIT = "1" +OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT = "2Gi" +OPEN_SANDBOX_EXECD_PORT = 44772 +OPEN_SANDBOX_DEFAULT_TIMEOUT = 300 +OPEN_SANDBOX_READY_TIMEOUT = 30.0 +OPEN_SANDBOX_POLL_INTERVAL = 0.2 DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024) diff --git a/litellm/llms/base_llm/sandbox/transformation.py b/litellm/llms/base_llm/sandbox/transformation.py index 6ad945f47a3..1c012a15fdb 100644 --- a/litellm/llms/base_llm/sandbox/transformation.py +++ b/litellm/llms/base_llm/sandbox/transformation.py @@ -8,10 +8,14 @@ run code -> delete container; `code_interpreter_tool` combines all three. from typing import Any, Union +import httpx + from pydantic import Field, PrivateAttr from litellm.types.llms.base import LiteLLMPydanticObjectBase +SANDBOX_MAX_OUTPUT_BYTES = 10 * 1024 * 1024 + class ContainerHandle(LiteLLMPydanticObjectBase): """A live sandbox container. Carries everything needed to reach it again.""" @@ -53,7 +57,7 @@ class BaseSandboxConfig: *, template: str | None = None, timeout: int | None = None, - allow_internet_access: bool = True, + allow_internet_access: bool | None = None, api_key: str | None = None, **kwargs, ) -> ContainerHandle: @@ -77,3 +81,16 @@ class BaseSandboxConfig: **kwargs, ) -> bool: raise NotImplementedError("adelete_sandbox must be implemented by provider") + + async def _read_capped_lines(self, response: httpx.Response) -> list[str]: + lines: list[str] = [] + total = 0 + async for line in response.aiter_lines(): + total += len(line.encode("utf-8")) + if total > SANDBOX_MAX_OUTPUT_BYTES: + raise ValueError( + f"Sandbox output exceeded {SANDBOX_MAX_OUTPUT_BYTES} bytes; aborting " + "to avoid unbounded memory use." + ) + lines.append(line) + return lines diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 129e15a0bf2..138f2410c89 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1,5 +1,6 @@ import json import ssl +from functools import lru_cache from urllib.parse import parse_qs, urlencode, urlparse, urlunparse from typing import ( TYPE_CHECKING, @@ -13,6 +14,7 @@ from typing import ( Tuple, Union, cast, + get_type_hints, ) import httpx # type: ignore @@ -26,6 +28,7 @@ from litellm._logging import _redact_string, verbose_logger from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming +from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, @@ -101,6 +104,7 @@ from litellm.types.llms.openai import ( HttpxBinaryResponseContent, OpenAIFileObject, ResponseInputParam, + ResponsesAPIOptionalRequestParams, ResponsesAPIResponse, ) from litellm.types.rerank import RerankResponse @@ -135,6 +139,7 @@ from litellm.utils import ( ImageResponse, ModelResponse, ProviderConfigManager, + async_pre_call_deployment_hook, ) from .http_handler import get_shared_realtime_ssl_context @@ -184,6 +189,47 @@ def _google_genai_streaming_hidden_params( } +@lru_cache(maxsize=None) +def _responses_api_optional_request_param_names() -> frozenset[str]: + return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams).keys()) + + +def _custom_logger_callbacks(logging_obj: Any) -> list[Any]: + from litellm.integrations.custom_logger import CustomLogger + from litellm.litellm_core_utils.litellm_logging import ( + get_custom_logger_compatible_class, + ) + + dynamic_success_callbacks = getattr(logging_obj, "dynamic_success_callbacks", None) + callbacks = list(litellm.callbacks) + if isinstance(dynamic_success_callbacks, (list, tuple)): + callbacks.extend(dynamic_success_callbacks) + + custom_loggers: list[Any] = [] + for cb in callbacks: + if isinstance(cb, str): + resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type] + if resolved is None: + continue + cb = resolved + if isinstance(cb, CustomLogger): + custom_loggers.append(cb) + return custom_loggers + + +def _has_pre_call_deployment_hook(logging_obj: Any) -> bool: + from litellm.integrations.custom_logger import CustomLogger + + base_func = CustomLogger.async_pre_call_deployment_hook + for cb in _custom_logger_callbacks(logging_obj): + cb_func = getattr(type(cb), "async_pre_call_deployment_hook", base_func) + if getattr(cb_func, "__func__", cb_func) is not getattr( + base_func, "__func__", base_func + ): + return True + return False + + class BaseLLMHTTPHandler: async def _make_common_async_call( self, @@ -2224,12 +2270,92 @@ class BaseLLMHTTPHandler: ) raise ValueError("anthropic_messages_handler is not implemented for sync calls") + def _run_sync_responses_pre_call_deployment_hook( + self, + *, + model: str, + input: Union[str, ResponseInputParam], + custom_llm_provider: str, + response_api_optional_request_params: dict[str, Any], + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + ) -> tuple[ + str, + Union[str, ResponseInputParam], + str, + dict[str, Any], + GenericLiteLLMParams, + ]: + if not _has_pre_call_deployment_hook(logging_obj): + return ( + model, + input, + custom_llm_provider, + response_api_optional_request_params, + litellm_params, + ) + + modified_kwargs = run_async_function( + async_pre_call_deployment_hook, + { + **dict(litellm_params), + **response_api_optional_request_params, + "model": model, + "input": input, + "custom_llm_provider": custom_llm_provider, + }, + CallTypes.responses.value, + ) + if modified_kwargs is None: + return ( + model, + input, + custom_llm_provider, + response_api_optional_request_params, + litellm_params, + ) + + optional_param_names = _responses_api_optional_request_param_names() + updated_response_params = { + **response_api_optional_request_params, + **{ + key: value + for key, value in modified_kwargs.items() + if key in optional_param_names + }, + } + updated_litellm_params = GenericLiteLLMParams( + **{ + **dict(litellm_params), + **{ + key: value + for key, value in modified_kwargs.items() + if key not in optional_param_names + and key not in {"model", "input", "custom_llm_provider"} + }, + } + ) + return ( + str(modified_kwargs["model"]) if "model" in modified_kwargs else model, + cast( + Union[str, ResponseInputParam], + modified_kwargs["input"] if "input" in modified_kwargs else input, + ), + ( + str(modified_kwargs["custom_llm_provider"]) + if "custom_llm_provider" in modified_kwargs + else custom_llm_provider + ), + updated_response_params, + updated_litellm_params, + ) + def response_api_handler( self, model: str, input: Union[str, ResponseInputParam], responses_api_provider_config: BaseResponsesAPIConfig, - response_api_optional_request_params: Dict, + response_api_optional_request_params: dict[str, Any], custom_llm_provider: str, litellm_params: GenericLiteLLMParams, logging_obj: LiteLLMLoggingObj, @@ -2276,6 +2402,21 @@ class BaseLLMHTTPHandler: shared_session=shared_session, ) + ( + model, + input, + custom_llm_provider, + response_api_optional_request_params, + litellm_params, + ) = self._run_sync_responses_pre_call_deployment_hook( + model=model, + input=input, + custom_llm_provider=custom_llm_provider, + response_api_optional_request_params=response_api_optional_request_params, + litellm_params=litellm_params, + logging_obj=logging_obj, + ) + if client is None or not isinstance(client, HTTPHandler): sync_httpx_client = _get_httpx_client( params={"ssl_verify": litellm_params.get("ssl_verify", None)} @@ -2414,9 +2555,27 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, ) ) - # Responses agentic interception (e.g. code interpreter) runs the follow-up - # loop via the async hook, so it is async-only for now; the sync path returns - # the initial response unchanged. + + if self._has_agentic_completion_hook(logging_obj): + final_response = run_async_function( + self._call_agentic_completion_hooks, + response=initial_response, + model=model, + messages=( + input + if isinstance(input, list) + else [{"role": "user", "content": input}] + ), + anthropic_messages_provider_config=responses_api_provider_config, + anthropic_messages_optional_request_params=response_api_optional_request_params, + logging_obj=logging_obj, + stream=False, + custom_llm_provider=custom_llm_provider, + kwargs=dict(litellm_params), + api_surface="responses", + ) + return final_response if final_response is not None else initial_response + return initial_response async def async_response_api_handler( @@ -4772,22 +4931,9 @@ class BaseLLMHTTPHandler: agentic callback is detected too. """ from litellm.integrations.custom_logger import CustomLogger - from litellm.litellm_core_utils.litellm_logging import ( - get_custom_logger_compatible_class, - ) base_func = CustomLogger.async_should_run_agentic_loop - callbacks = litellm.callbacks + ( - getattr(logging_obj, "dynamic_success_callbacks", None) or [] - ) - for cb in callbacks: - if isinstance(cb, str): - resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type] - if resolved is None: - continue - cb = resolved - if not isinstance(cb, CustomLogger): - continue + for cb in _custom_logger_callbacks(logging_obj): cb_func = getattr(type(cb), "async_should_run_agentic_loop", base_func) if getattr(cb_func, "__func__", cb_func) is not getattr( base_func, "__func__", base_func diff --git a/litellm/llms/e2b/sandbox/transformation.py b/litellm/llms/e2b/sandbox/transformation.py index c279fab22ab..ecfc1642c97 100644 --- a/litellm/llms/e2b/sandbox/transformation.py +++ b/litellm/llms/e2b/sandbox/transformation.py @@ -16,6 +16,7 @@ from litellm.llms.base_llm.sandbox.transformation import ( BaseSandboxConfig, CodeExecutionResult, ContainerHandle, + SANDBOX_MAX_OUTPUT_BYTES, ) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -29,7 +30,7 @@ E2B_DEFAULT_TEMPLATE = "code-interpreter-v1" E2B_DEFAULT_DOMAIN = "e2b.app" JUPYTER_PORT = 49999 DEFAULT_SANDBOX_TIMEOUT = 300 -MAX_OUTPUT_BYTES = 10 * 1024 * 1024 +MAX_OUTPUT_BYTES = SANDBOX_MAX_OUTPUT_BYTES class E2BSandboxConfig(BaseSandboxConfig): @@ -49,7 +50,7 @@ class E2BSandboxConfig(BaseSandboxConfig): *, template: str | None = None, timeout: int | None = None, - allow_internet_access: bool = True, + allow_internet_access: bool | None = None, api_key: str | None = None, api_base: str | None = None, metadata: dict | None = None, @@ -62,7 +63,9 @@ class E2BSandboxConfig(BaseSandboxConfig): "templateID": template or E2B_DEFAULT_TEMPLATE, "timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT, "secure": True, - "allow_internet_access": allow_internet_access, + "allow_internet_access": ( + True if allow_internet_access is None else allow_internet_access + ), } if metadata: body["metadata"] = metadata @@ -168,20 +171,6 @@ class E2BSandboxConfig(BaseSandboxConfig): handle._hidden_params = {} return handle - @staticmethod - async def _read_capped_lines(response: httpx.Response) -> list[str]: - lines: list[str] = [] - total = 0 - async for line in response.aiter_lines(): - total += len(line.encode("utf-8")) - if total > MAX_OUTPUT_BYTES: - raise ValueError( - f"Sandbox output exceeded {MAX_OUTPUT_BYTES} bytes; aborting to " - "avoid unbounded memory use." - ) - lines.append(line) - return lines - @staticmethod def _parse_lines(lines: list[str]) -> CodeExecutionResult: def _try_parse(stripped: str): @@ -192,10 +181,9 @@ class E2BSandboxConfig(BaseSandboxConfig): messages = tuple( parsed - for stripped in (line.strip() for line in lines) - if stripped - for parsed in (_try_parse(stripped),) - if parsed is not None + for line in lines + if (stripped := line.strip()) + if (parsed := _try_parse(stripped)) is not None ) def of_type(message_type: str): diff --git a/litellm/llms/opensandbox/__init__.py b/litellm/llms/opensandbox/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/litellm/llms/opensandbox/__init__.py @@ -0,0 +1 @@ + diff --git a/litellm/llms/opensandbox/sandbox/__init__.py b/litellm/llms/opensandbox/sandbox/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/litellm/llms/opensandbox/sandbox/__init__.py @@ -0,0 +1 @@ + diff --git a/litellm/llms/opensandbox/sandbox/transformation.py b/litellm/llms/opensandbox/sandbox/transformation.py new file mode 100644 index 00000000000..dc9f8440d30 --- /dev/null +++ b/litellm/llms/opensandbox/sandbox/transformation.py @@ -0,0 +1,598 @@ +import asyncio +import json +import time +from typing import Union, cast + +import httpx + +from litellm.constants import ( + OPEN_SANDBOX_API_BASE_ENV_VAR, + OPEN_SANDBOX_API_KEY_ENV_VAR, + OPEN_SANDBOX_DEFAULT_CPU_LIMIT, + OPEN_SANDBOX_DEFAULT_ENTRYPOINT, + OPEN_SANDBOX_DEFAULT_LANGUAGE, + OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT, + OPEN_SANDBOX_DEFAULT_TEMPLATE, + OPEN_SANDBOX_DEFAULT_TIMEOUT, + OPEN_SANDBOX_EXECD_PORT, + OPEN_SANDBOX_POLL_INTERVAL, + OPEN_SANDBOX_READY_TIMEOUT, +) +from litellm.llms.base_llm.sandbox.transformation import ( + BaseSandboxConfig, + CodeExecutionResult, + ContainerHandle, + SANDBOX_MAX_OUTPUT_BYTES, +) +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + get_async_httpx_client, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.custom_http import httpxSpecialProvider + +DEFAULT_SANDBOX_TIMEOUT = OPEN_SANDBOX_DEFAULT_TIMEOUT +DEFAULT_READY_TIMEOUT = OPEN_SANDBOX_READY_TIMEOUT +DEFAULT_POLL_INTERVAL = OPEN_SANDBOX_POLL_INTERVAL +MAX_OUTPUT_BYTES = SANDBOX_MAX_OUTPUT_BYTES + + +class OpenSandboxSandboxConfig(BaseSandboxConfig): + def _http(self, client: AsyncHTTPHandler | None) -> AsyncHTTPHandler: + if client is not None: + return client + return get_async_httpx_client(llm_provider=httpxSpecialProvider.Sandbox) + + def validate_environment(self, api_key: str | None = None, **kwargs) -> str: + if api_key is not None: + return api_key + return get_secret_str(OPEN_SANDBOX_API_KEY_ENV_VAR) or "" + + async def acreate_sandbox( + self, + *, + template: str | None = None, + timeout: int | None = None, + allow_internet_access: bool | None = None, + api_key: str | None = None, + api_base: str | None = None, + metadata: dict[str, str] | None = None, + env_vars: dict[str, str] | None = None, + resource_limits: dict[str, str] | None = None, + resource_requests: dict[str, str] | None = None, + entrypoint: list[str] | tuple[str, ...] | None = None, + network_policy: dict[str, object] | None = None, + secure_access: bool = False, + use_server_proxy: bool = False, + ready_timeout: float | None = None, + poll_interval: float | None = None, + client: AsyncHTTPHandler | None = None, + **kwargs, + ) -> ContainerHandle: + key = self.validate_environment(api_key=api_key) + base = self._api_base(api_base) + ready_timeout_seconds = ( + float(ready_timeout) if ready_timeout is not None else DEFAULT_READY_TIMEOUT + ) + poll_interval_seconds = ( + float(poll_interval) if poll_interval is not None else DEFAULT_POLL_INTERVAL + ) + body = self._create_body( + template=template, + timeout=timeout, + allow_internet_access=allow_internet_access, + metadata=metadata, + env_vars=env_vars, + resource_limits=resource_limits, + resource_requests=resource_requests, + entrypoint=entrypoint, + network_policy=network_policy, + secure_access=secure_access, + ) + + response = cast( + httpx.Response, + await self._http(client).post( + url=f"{base}/sandboxes", + headers=self._lifecycle_headers(key), + json=body, + ), + ) + data = response.json() + sandbox_id = str(data["id"]) + + if self._sandbox_state(data) != "Running": + await self._wait_until_running( + sandbox_id=sandbox_id, + api_base=base, + headers=self._lifecycle_headers(key), + client=client, + ready_timeout=ready_timeout_seconds, + poll_interval=poll_interval_seconds, + ) + + endpoint, endpoint_headers = await self._wait_for_execd_endpoint( + sandbox_id=sandbox_id, + api_base=base, + headers=self._lifecycle_headers(key), + use_server_proxy=use_server_proxy, + client=client, + ready_timeout=ready_timeout_seconds, + poll_interval=poll_interval_seconds, + ) + + handle = ContainerHandle(id=sandbox_id, provider="opensandbox", domain=base) + handle._hidden_params = { + "api_base": base, + "api_key": key, + "execd_endpoint": endpoint, + "execd_headers": endpoint_headers, + "use_server_proxy": use_server_proxy, + } + return handle + + async def arun_code( + self, + *, + container: Union[ContainerHandle, str], + code: str, + api_key: str | None = None, + api_base: str | None = None, + language: str = OPEN_SANDBOX_DEFAULT_LANGUAGE, + use_server_proxy: bool = False, + ready_timeout: float | None = None, + poll_interval: float | None = None, + client: AsyncHTTPHandler | None = None, + **kwargs, + ) -> CodeExecutionResult: + handle = await self._ensure_handle( + container=container, + api_key=api_key, + api_base=api_base, + use_server_proxy=use_server_proxy, + ready_timeout=( + float(ready_timeout) + if ready_timeout is not None + else DEFAULT_READY_TIMEOUT + ), + poll_interval=( + float(poll_interval) + if poll_interval is not None + else DEFAULT_POLL_INTERVAL + ), + client=client, + ) + endpoint = str(handle._hidden_params["execd_endpoint"]) + endpoint_headers = self._as_str_dict(handle._hidden_params.get("execd_headers")) + base = str( + handle._hidden_params.get("api_base") + or handle.domain + or self._api_base(api_base) + ) + lines = await self._post_code( + url=f"{self._endpoint_base_url(endpoint, base)}/code", + headers={ + "Content-Type": "application/json", + "Accept": "text/event-stream", + "Cache-Control": "no-cache", + **endpoint_headers, + }, + body={ + "code": code, + "context": {"language": language}, + }, + client=client, + ) + return self._parse_lines(lines) + + async def adelete_sandbox( + self, + *, + container: Union[ContainerHandle, str], + api_key: str | None = None, + api_base: str | None = None, + client: AsyncHTTPHandler | None = None, + **kwargs, + ) -> bool: + handle = self._as_handle(container, api_base=api_base) + base = str(handle._hidden_params.get("api_base") or self._api_base(api_base)) + key = self._api_key(api_key=api_key, handle=handle) + try: + response = cast( + httpx.Response, + await self._http(client).delete( + url=f"{base}/sandboxes/{handle.id}", + headers=self._lifecycle_headers(key), + ), + ) + except httpx.HTTPStatusError as e: + if e.response.status_code == 404: + return False + raise + return 200 <= response.status_code < 300 + + async def _ensure_handle( + self, + *, + container: Union[ContainerHandle, str], + api_key: str | None, + api_base: str | None, + use_server_proxy: bool, + ready_timeout: float, + poll_interval: float, + client: AsyncHTTPHandler | None, + ) -> ContainerHandle: + handle = self._as_handle(container, api_base=api_base) + if handle._hidden_params.get("execd_endpoint"): + return handle + + base = str(handle._hidden_params.get("api_base") or self._api_base(api_base)) + key = self._api_key(api_key=api_key, handle=handle) + resolved_use_server_proxy = bool( + handle._hidden_params.get("use_server_proxy", use_server_proxy) + ) + endpoint, endpoint_headers = await self._wait_for_execd_endpoint( + sandbox_id=handle.id, + api_base=base, + headers=self._lifecycle_headers(key), + use_server_proxy=resolved_use_server_proxy, + client=client, + ready_timeout=ready_timeout, + poll_interval=poll_interval, + ) + handle.domain = base + handle._hidden_params = { + **handle._hidden_params, + "api_base": base, + "api_key": key, + "execd_endpoint": endpoint, + "execd_headers": endpoint_headers, + "use_server_proxy": resolved_use_server_proxy, + } + return handle + + async def _wait_until_running( + self, + *, + sandbox_id: str, + api_base: str, + headers: dict[str, str], + client: AsyncHTTPHandler | None, + ready_timeout: float, + poll_interval: float, + ) -> None: + deadline = time.monotonic() + ready_timeout + while True: + response = cast( + httpx.Response, + await self._http(client).get( + url=f"{api_base}/sandboxes/{sandbox_id}", + headers=headers, + ), + ) + data = response.json() + state = self._sandbox_state(data) + if state == "Running": + return + if state in {"Failed", "Stopping", "Terminated"}: + raise ValueError(f"OpenSandbox sandbox {sandbox_id} entered {state}") + if time.monotonic() >= deadline: + raise TimeoutError( + f"OpenSandbox sandbox {sandbox_id} was not Running within " + f"{ready_timeout} seconds" + ) + await asyncio.sleep(poll_interval) + + async def _wait_for_execd_endpoint( + self, + *, + sandbox_id: str, + api_base: str, + headers: dict[str, str], + use_server_proxy: bool, + client: AsyncHTTPHandler | None, + ready_timeout: float, + poll_interval: float, + ) -> tuple[str, dict[str, str]]: + deadline = time.monotonic() + ready_timeout + last_error: Exception | None = None + while True: + try: + return await self._get_execd_endpoint( + sandbox_id=sandbox_id, + api_base=api_base, + headers=headers, + use_server_proxy=use_server_proxy, + client=client, + ) + except httpx.HTTPStatusError as e: + if e.response.status_code != 404: + raise + last_error = e + except ValueError as e: + last_error = e + + if time.monotonic() >= deadline: + raise TimeoutError( + f"OpenSandbox execd endpoint for {sandbox_id} was not ready within " + f"{ready_timeout} seconds" + ) from last_error + await asyncio.sleep(poll_interval) + + async def _get_execd_endpoint( + self, + *, + sandbox_id: str, + api_base: str, + headers: dict[str, str], + use_server_proxy: bool, + client: AsyncHTTPHandler | None, + ) -> tuple[str, dict[str, str]]: + response = cast( + httpx.Response, + await self._http(client).get( + url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}", + headers=headers, + params={"use_server_proxy": use_server_proxy}, + ), + ) + data = response.json() + endpoint = data.get("endpoint") + if not endpoint: + raise ValueError( + f"OpenSandbox did not return an execd endpoint for {sandbox_id}" + ) + return str(endpoint), self._as_str_dict(data.get("headers")) + + async def _post_code( + self, + *, + url: str, + headers: dict[str, str], + body: dict[str, object], + client: AsyncHTTPHandler | None, + ) -> list[str]: + timeout = httpx.Timeout(connect=30.0, read=None, write=30.0, pool=None) + response = cast( + httpx.Response, + await self._http(client).post( + url=url, + headers=headers, + timeout=timeout, + json=body, + stream=True, + ), + ) + return await self._read_capped_lines(response) + + def _api_key(self, *, api_key: str | None, handle: ContainerHandle) -> str: + if api_key is not None: + return api_key + if "api_key" in handle._hidden_params: + return str(handle._hidden_params["api_key"]) + return self.validate_environment() + + @staticmethod + def _create_body( + *, + template: str | None, + timeout: int | None, + allow_internet_access: bool | None, + metadata: dict[str, str] | None, + env_vars: dict[str, str] | None, + resource_limits: dict[str, str] | None, + resource_requests: dict[str, str] | None, + entrypoint: list[str] | tuple[str, ...] | None, + network_policy: dict[str, object] | None, + secure_access: bool, + ) -> dict[str, object]: + body: dict[str, object] = { + "image": {"uri": template or OPEN_SANDBOX_DEFAULT_TEMPLATE}, + "entrypoint": list(entrypoint or OPEN_SANDBOX_DEFAULT_ENTRYPOINT), + "timeout": timeout if timeout is not None else DEFAULT_SANDBOX_TIMEOUT, + "resourceLimits": resource_limits + or OpenSandboxSandboxConfig._default_resource_limits(), + } + if metadata: + body["metadata"] = metadata + if env_vars: + body["env"] = env_vars + if resource_requests: + body["resourceRequests"] = resource_requests + if network_policy is not None: + body["networkPolicy"] = network_policy + elif allow_internet_access is not True: + body["networkPolicy"] = {"defaultAction": "deny", "egress": []} + if secure_access: + body["secureAccess"] = True + return body + + @staticmethod + def _default_resource_limits() -> dict[str, str]: + return { + "cpu": OPEN_SANDBOX_DEFAULT_CPU_LIMIT, + "memory": OPEN_SANDBOX_DEFAULT_MEMORY_LIMIT, + } + + @staticmethod + def _sandbox_state(data: object) -> str | None: + if not isinstance(data, dict): + return None + status = data.get("status") + if not isinstance(status, dict): + return None + state = status.get("state") + return str(state) if state is not None else None + + @staticmethod + def _as_str_dict(value: object) -> dict[str, str]: + if not isinstance(value, dict): + return {} + return {str(k): str(v) for k, v in value.items()} + + @staticmethod + def _api_base(api_base: str | None) -> str: + base = api_base or get_secret_str(OPEN_SANDBOX_API_BASE_ENV_VAR) + if not base: + raise ValueError( + "OpenSandbox api_base is required. Pass api_base or set " + f"{OPEN_SANDBOX_API_BASE_ENV_VAR}." + ) + return str(base).rstrip("/") + + @staticmethod + def _lifecycle_headers(api_key: str) -> dict[str, str]: + headers = {"Content-Type": "application/json"} + if api_key: + headers["OPEN-SANDBOX-API-KEY"] = api_key + return headers + + @staticmethod + def _endpoint_base_url(endpoint: str, api_base: str) -> str: + normalized_endpoint = endpoint.rstrip("/") + if normalized_endpoint.startswith(("http://", "https://")): + return normalized_endpoint + protocol = api_base.split("://", 1)[0] if "://" in api_base else "http" + return f"{protocol}://{normalized_endpoint}" + + @staticmethod + def _as_handle( + container: Union[ContainerHandle, str], *, api_base: str | None + ) -> ContainerHandle: + if isinstance(container, ContainerHandle): + return container + handle = ContainerHandle( + id=str(container), + provider="opensandbox", + domain=OpenSandboxSandboxConfig._api_base(api_base), + ) + handle._hidden_params = {} + return handle + + @staticmethod + def _parse_lines(lines: list[str]) -> CodeExecutionResult: + messages = tuple( + event + for line in lines + if (event := OpenSandboxSandboxConfig._parse_sse_line(line)) is not None + ) + + def of_type(message_type: str): + return (m for m in messages if m.get("type") == message_type) + + error = next( + (OpenSandboxSandboxConfig._normalize_error(m) for m in of_type("error")), + None, + ) + execution_count = next( + ( + OpenSandboxSandboxConfig._as_int(m.get("execution_count")) + for m in of_type("execution_count") + if OpenSandboxSandboxConfig._as_int(m.get("execution_count")) + is not None + ), + None, + ) + + return CodeExecutionResult( + stdout="".join(str(m.get("text", "")) for m in of_type("stdout")), + stderr="".join(str(m.get("text", "")) for m in of_type("stderr")), + results=[ + OpenSandboxSandboxConfig._normalize_result(m) for m in of_type("result") + ], + error=error, + execution_count=execution_count, + ) + + @staticmethod + def _parse_sse_line(line: str) -> dict[str, object] | None: + stripped = line.strip() + if not stripped or stripped.startswith( + ( + ":", + "event:", + "id:", + "retry:", + ) + ): + return None + data = stripped[5:].strip() if stripped.startswith("data:") else stripped + if not data: + return None + try: + parsed = json.loads(data) + except json.JSONDecodeError: + return None + if not isinstance(parsed, dict): + return None + if "type" not in parsed and "code" in parsed and "message" in parsed: + return { + "type": "error", + "error": { + "ename": str(parsed["code"]), + "evalue": str(parsed["message"]), + "traceback": [], + }, + } + return parsed + + @staticmethod + def _normalize_result(message: dict[str, object]) -> dict[str, object]: + results = message.get("results") + if isinstance(results, dict): + return {str(k): v for k, v in results.items()} + return { + str(k): v + for k, v in message.items() + if k not in {"type", "timestamp", "execution_count"} + } + + @staticmethod + def _normalize_error(message: dict[str, object]) -> dict[str, object]: + raw_error = message.get("error") + if isinstance(raw_error, dict): + name = OpenSandboxSandboxConfig._first_non_none_value( + raw_error, "ename", "name", default="" + ) + value = OpenSandboxSandboxConfig._first_non_none_value( + raw_error, "evalue", "value", default="" + ) + traceback = OpenSandboxSandboxConfig._first_non_none_value( + raw_error, "traceback", default=[] + ) + return { + "name": name, + "value": value, + "traceback": traceback, + } + return { + "name": OpenSandboxSandboxConfig._first_non_none_value( + message, "name", default="" + ), + "value": OpenSandboxSandboxConfig._first_non_none_value( + message, "value", "text", default="" + ), + "traceback": OpenSandboxSandboxConfig._first_non_none_value( + message, "traceback", default=[] + ), + } + + @staticmethod + def _as_int(value: object) -> int | None: + if isinstance(value, int): + return value + if isinstance(value, str): + try: + return int(value) + except ValueError: + return None + return None + + @staticmethod + def _first_non_none_value( + values: dict[str, object], *keys: str, default: object + ) -> object: + return next( + (values[key] for key in keys if key in values and values[key] is not None), + default, + ) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 34c9cdd3d1c..2c46baaada5 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -58,7 +58,10 @@ from litellm.llms.openai.data_residency import infer_openai_data_residency from litellm.secret_managers.main import get_secret_str from litellm.types.responses.main import * from litellm.types.router import GenericLiteLLMParams -from litellm.utils import ProviderConfigManager, client +from litellm.utils import ( + ProviderConfigManager, + client, +) if TYPE_CHECKING: from mcp.types import Tool as MCPTool diff --git a/litellm/sandbox/main.py b/litellm/sandbox/main.py index 45d3bffb4f9..76d9994c683 100644 --- a/litellm/sandbox/main.py +++ b/litellm/sandbox/main.py @@ -68,7 +68,7 @@ async def acreate_sandbox( provider: str, template: str | None = None, timeout: int | None = None, - allow_internet_access: bool = True, + allow_internet_access: bool | None = None, api_key: str | None = None, api_base: str | None = None, **kwargs, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 76854cd28f9..00b095f33ca 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3521,6 +3521,7 @@ class SandboxProviders(str, Enum): """ E2B = "e2b" + OPENSANDBOX = "opensandbox" class LiteLLMLoggingBaseClass: diff --git a/litellm/utils.py b/litellm/utils.py index a842e9e058d..5c3ab3e1490 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9731,9 +9731,14 @@ class ProviderConfigManager: Get sandbox (code execution) configuration for a given provider. """ from litellm.llms.e2b.sandbox.transformation import E2BSandboxConfig + from litellm.llms.opensandbox.sandbox.transformation import ( + OpenSandboxSandboxConfig, + ) if provider == SandboxProviders.E2B: return E2BSandboxConfig() + if provider == SandboxProviders.OPENSANDBOX: + return OpenSandboxSandboxConfig() return None @staticmethod diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 9991ff9e01e..b137ec59a1f 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -1833,6 +1833,23 @@ "text_completion": true } }, + "opensandbox": { + "display_name": "OpenSandbox (`opensandbox`)", + "url": "https://open-sandbox.ai/api/", + "endpoints": { + "chat_completions": false, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "sandbox": true + } + }, "openai_like": { "display_name": "OpenAI-like (`openai_like`)", "url": "https://docs.litellm.ai/docs/providers/openai_compatible", diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index f7d445d0788..64ae30daa70 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -10,13 +10,21 @@ sys.path.insert( 0, os.path.abspath("../../../..") ) # Adds the parent directory to the system path import litellm +from litellm.integrations.code_interpreter_interception.handler import ( + CodeInterpreterInterceptionLogger, + LITELLM_CODE_EXECUTION_TOOL_NAME, +) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import ( BaseLLMHTTPHandler, _google_genai_streaming_hidden_params, ) +from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams +_ACTIVE_KEY = "_code_interpreter_interception_active" +_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key" + def test_prepare_fake_stream_request(): # Initialize the BaseLLMHTTPHandler @@ -116,6 +124,117 @@ def test_response_api_handler_streams_when_provider_transform_adds_stream(): assert client.post.call_args.kwargs["json"]["stream"] is True +def test_response_api_handler_runs_agentic_hooks_in_sync_path(monkeypatch): + handler = BaseLLMHTTPHandler() + config = Mock() + config.validate_environment.return_value = {} + config.get_complete_url.return_value = "https://chatgpt.example.com/responses" + config.transform_responses_api_request.return_value = { + "model": "gpt-5", + "input": "hi", + } + config.sign_request.return_value = ({}, None) + initial_response = Mock() + final_response = Mock() + config.transform_response_api_response.return_value = initial_response + + client = HTTPHandler(client=httpx.Client()) + client.post = Mock( + return_value=httpx.Response( + 200, + request=httpx.Request("POST", "https://chatgpt.example.com/responses"), + ) + ) + logging_obj = Mock() + + monkeypatch.setattr(handler, "_has_agentic_completion_hook", Mock(return_value=True)) + hook_mock = AsyncMock(return_value=final_response) + monkeypatch.setattr(handler, "_call_agentic_completion_hooks", hook_mock) + + response = handler.response_api_handler( + model="gpt-5", + input="hi", + responses_api_provider_config=config, + response_api_optional_request_params={}, + custom_llm_provider="openai", + litellm_params=GenericLiteLLMParams(), + logging_obj=logging_obj, + client=client, + ) + + assert response is final_response + hook_mock.assert_awaited_once() + assert hook_mock.call_args.kwargs["api_surface"] == "responses" + assert hook_mock.call_args.kwargs["messages"] == [ + {"role": "user", "content": "hi"} + ] + + +def test_response_api_handler_runs_responses_pre_call_hook_before_transform(): + handler = BaseLLMHTTPHandler() + config = Mock() + config.validate_environment.return_value = {} + config.get_complete_url.return_value = "https://api.openai.com/v1/responses" + config.sign_request.return_value = ({}, None) + initial_response = ResponsesAPIResponse( + id="resp_1", + created_at=0, + output=[], + status="completed", + model="gpt-5", + ) + config.transform_response_api_response.return_value = initial_response + + def transform_responses_api_request(**kwargs): + return { + "model": kwargs["model"], + "input": kwargs["input"], + **kwargs["response_api_optional_request_params"], + } + + config.transform_responses_api_request.side_effect = transform_responses_api_request + client = HTTPHandler(client=httpx.Client()) + client.post = Mock( + return_value=httpx.Response( + 200, + request=httpx.Request("POST", "https://api.openai.com/v1/responses"), + ) + ) + logging_obj = Mock() + logging_obj.dynamic_success_callbacks = [] + + old_callbacks = list(litellm.callbacks) + litellm.callbacks = [CodeInterpreterInterceptionLogger()] + try: + response = handler.response_api_handler( + model="gpt-5", + input="use code", + responses_api_provider_config=config, + response_api_optional_request_params={ + "tools": [{"type": "code_interpreter", "container": {"type": "auto"}}] + }, + custom_llm_provider="openai", + litellm_params=GenericLiteLLMParams(api_key="sk-test"), + logging_obj=logging_obj, + client=client, + ) + finally: + litellm.callbacks = old_callbacks + + assert response is initial_response + transform_kwargs = config.transform_responses_api_request.call_args.kwargs + tools = transform_kwargs["response_api_optional_request_params"]["tools"] + assert not any(tool.get("type") == "code_interpreter" for tool in tools) + assert any( + tool.get("type") == "function" + and tool.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME + for tool in tools + ) + hook_litellm_params = transform_kwargs["litellm_params"] + assert hook_litellm_params.get(_ACTIVE_KEY) is True + assert hook_litellm_params.get(_SANDBOX_KEY) + + @pytest.mark.asyncio async def test_async_response_api_handler_streams_when_provider_transform_adds_stream(): handler = BaseLLMHTTPHandler() diff --git a/tests/test_litellm/sandbox/test_opensandbox_sandbox.py b/tests/test_litellm/sandbox/test_opensandbox_sandbox.py new file mode 100644 index 00000000000..0d7bcbe1e53 --- /dev/null +++ b/tests/test_litellm/sandbox/test_opensandbox_sandbox.py @@ -0,0 +1,647 @@ +import json + +import httpx +import pytest + +import litellm +from litellm.llms.base_llm.sandbox.transformation import ContainerHandle +from litellm.llms.opensandbox.sandbox.transformation import ( + MAX_OUTPUT_BYTES, + OPEN_SANDBOX_DEFAULT_TEMPLATE, + OpenSandboxSandboxConfig, +) +from litellm.utils import ProviderConfigManager + +TEST_API_BASE = "https://sandbox.test/v1" + + +def http_status_error(status_code, url="http://test"): + return httpx.HTTPStatusError( + f"status {status_code}", + request=httpx.Request("GET", url), + response=httpx.Response(status_code), + ) + + +def sse(data): + return f"data: {json.dumps(data)}" + + +class FakeResponse: + def __init__(self, *, json_data=None, lines=None, status_code=200): + self._json = json_data + self._lines = lines or [] + self.status_code = status_code + + def json(self): + return self._json + + def raise_for_status(self): + if self.status_code >= 400: + raise http_status_error(self.status_code) + + async def aiter_lines(self): + for line in self._lines: + yield line + + +class FakeHTTPClient: + def __init__( + self, + *, + create_json=None, + sandbox_states=None, + endpoint_json=None, + endpoint_responses=None, + execute_lines=None, + delete_status=204, + execute_raises=None, + ): + self.create_json = create_json or { + "id": "osb_123", + "status": {"state": "Running"}, + "createdAt": "2026-01-01T00:00:00Z", + "entrypoint": ["/opt/code-interpreter/code-interpreter.sh"], + } + self.sandbox_states = list( + sandbox_states + or [ + { + "id": "osb_123", + "status": {"state": "Running"}, + "createdAt": "2026-01-01T00:00:00Z", + "entrypoint": ["/opt/code-interpreter/code-interpreter.sh"], + } + ] + ) + self.endpoint_json = endpoint_json or { + "endpoint": "execd.local:44772", + "headers": {"X-EXECD-ACCESS-TOKEN": "execd-token"}, + } + self.endpoint_responses = ( + list(endpoint_responses) if endpoint_responses is not None else None + ) + self.execute_lines = execute_lines or [] + self.delete_status = delete_status + self.execute_raises = execute_raises + self.calls = [] + + async def post(self, url, headers=None, json=None, stream=False, **kwargs): + self.calls.append(("POST", url, headers, json, {"stream": stream})) + if url.endswith("/sandboxes"): + return FakeResponse(json_data=self.create_json) + if url.endswith("/code"): + if self.execute_raises is not None: + raise self.execute_raises + return FakeResponse(lines=self.execute_lines) + raise AssertionError(f"unexpected POST {url}") + + async def get(self, url, headers=None, params=None, **kwargs): + self.calls.append(("GET", url, headers, None, params)) + if "/endpoints/44772" in url: + if self.endpoint_responses is not None and self.endpoint_responses: + response = self.endpoint_responses.pop(0) + if isinstance(response, Exception): + raise response + if isinstance(response, FakeResponse): + return response + return FakeResponse(json_data=response) + return FakeResponse(json_data=self.endpoint_json) + if "/sandboxes/" in url: + state = self.sandbox_states.pop(0) + return FakeResponse(json_data=state) + raise AssertionError(f"unexpected GET {url}") + + async def delete(self, url, headers=None, **kwargs): + self.calls.append(("DELETE", url, headers, None, None)) + if not (200 <= self.delete_status < 300): + raise http_status_error(self.delete_status, url) + return FakeResponse(status_code=self.delete_status) + + +def test_parse_sse_lines_maps_output_result_count_and_error(): + lines = [ + sse({"type": "stdout", "text": "hello\n"}), + sse({"type": "stderr", "text": "warn\n"}), + sse({"type": "result", "results": {"text/plain": "4"}}), + sse({"type": "execution_count", "execution_count": 7}), + sse( + { + "type": "error", + "error": { + "ename": "ValueError", + "evalue": "bad", + "traceback": ["Traceback"], + }, + } + ), + ] + + result = OpenSandboxSandboxConfig._parse_lines(lines) + + assert result.stdout == "hello\n" + assert result.stderr == "warn\n" + assert result.results == [{"text/plain": "4"}] + assert result.execution_count == 7 + assert result.error == { + "name": "ValueError", + "value": "bad", + "traceback": ["Traceback"], + } + + +def test_parse_sse_lines_skips_non_json_and_control_lines(): + lines = [ + "event: message", + "not-json", + "", + sse({"type": "stdout", "text": "ok\n"}), + ] + + result = OpenSandboxSandboxConfig._parse_lines(lines) + + assert result.stdout == "ok\n" + assert result.error is None + + +def test_parse_sse_lines_maps_fallback_shapes(): + lines = [ + "data:", + sse(["not-a-dict"]), + sse({"code": "BadRequest", "message": "nope"}), + sse({"type": "result", "text/plain": "4"}), + sse({"type": "error", "name": "RuntimeError", "text": "boom"}), + sse({"type": "execution_count", "execution_count": "8"}), + ] + + result = OpenSandboxSandboxConfig._parse_lines(lines) + + assert result.results == [{"text/plain": "4"}] + assert result.execution_count == 8 + assert result.error == { + "name": "BadRequest", + "value": "nope", + "traceback": [], + } + fallback_error = OpenSandboxSandboxConfig._parse_lines( + [sse({"type": "error", "name": "RuntimeError", "text": "boom"})] + ) + assert fallback_error.error == { + "name": "RuntimeError", + "value": "boom", + "traceback": [], + } + empty_string_error = OpenSandboxSandboxConfig._parse_lines( + [ + sse( + { + "type": "error", + "error": { + "ename": "", + "name": "FallbackName", + "evalue": "", + "value": "fallback value", + "traceback": [], + }, + } + ) + ] + ) + assert empty_string_error.error == { + "name": "", + "value": "", + "traceback": [], + } + + +def test_static_helpers_cover_defaults_and_fallbacks(monkeypatch): + def fake_secret(key): + if key == "OPEN_SANDBOX_API_KEY": + return "env-key" + if key == "OPEN_SANDBOX_API_BASE": + return TEST_API_BASE + return None + + monkeypatch.setattr( + "litellm.llms.opensandbox.sandbox.transformation.get_secret_str", + fake_secret, + ) + config = OpenSandboxSandboxConfig() + handle = ContainerHandle(id="osb", provider="opensandbox", domain="http://x/v1") + + assert config.validate_environment() == "env-key" + assert config.validate_environment(api_key="") == "" + assert config._api_key(api_key=None, handle=handle) == "env-key" + + handle._hidden_params = {"api_key": "stored-key"} + assert config._api_key(api_key=None, handle=handle) == "stored-key" + assert config._http(None) is not None + + body = config._create_body( + template=None, + timeout=None, + allow_internet_access=False, + metadata=None, + env_vars=None, + resource_limits=None, + resource_requests=None, + entrypoint=None, + network_policy={"egress": [{"domain": "example.com"}]}, + secure_access=True, + ) + assert body["networkPolicy"] == {"egress": [{"domain": "example.com"}]} + assert body["secureAccess"] is True + + other_body = config._create_body( + template=None, + timeout=None, + allow_internet_access=False, + metadata=None, + env_vars=None, + resource_limits=None, + resource_requests=None, + entrypoint=None, + network_policy=None, + secure_access=False, + ) + assert body["resourceLimits"] is not other_body["resourceLimits"] + + assert config._sandbox_state(None) is None + assert config._sandbox_state({"status": "Running"}) is None + assert config._as_str_dict(None) == {} + assert config._endpoint_base_url("http://execd.local", "https://api/v1") == ( + "http://execd.local" + ) + assert config._api_base(None) == TEST_API_BASE + assert config._api_base("https://direct.test/v1/") == "https://direct.test/v1" + assert config._as_int("9") == 9 + assert config._as_int("nope") is None + assert config._as_int(None) is None + assert isinstance( + ProviderConfigManager.get_provider_sandbox_config("opensandbox"), + OpenSandboxSandboxConfig, + ) + + +def test_api_base_requires_kwarg_or_env(monkeypatch): + monkeypatch.setattr( + "litellm.llms.opensandbox.sandbox.transformation.get_secret_str", + lambda key: None, + ) + + with pytest.raises(ValueError, match="api_base is required"): + OpenSandboxSandboxConfig._api_base(None) + + +@pytest.mark.asyncio +async def test_create_posts_default_body_and_omits_empty_api_key(): + client = FakeHTTPClient() + + handle = await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", api_base=TEST_API_BASE, client=client + ) + + method, url, headers, body, _ = client.calls[0] + assert method == "POST" + assert url == f"{TEST_API_BASE}/sandboxes" + assert "OPEN-SANDBOX-API-KEY" not in headers + assert body["image"] == {"uri": OPEN_SANDBOX_DEFAULT_TEMPLATE} + assert body["entrypoint"] == ["/opt/code-interpreter/code-interpreter.sh"] + assert body["timeout"] == 300 + assert body["resourceLimits"] == {"cpu": "1", "memory": "2Gi"} + assert body["networkPolicy"] == {"defaultAction": "deny", "egress": []} + assert handle.id == "osb_123" + assert handle._hidden_params["execd_endpoint"] == "execd.local:44772" + + +@pytest.mark.asyncio +async def test_create_can_opt_into_internet_access(): + client = FakeHTTPClient() + + await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", + api_base=TEST_API_BASE, + allow_internet_access=True, + client=client, + ) + + _, _, _, body, _ = client.calls[0] + assert "networkPolicy" not in body + + +@pytest.mark.asyncio +async def test_create_custom_options_poll_and_endpoint_resolution(): + client = FakeHTTPClient( + create_json={ + "id": "osb_pending", + "status": {"state": "Pending"}, + "createdAt": "2026-01-01T00:00:00Z", + "entrypoint": ["/bin/sh"], + }, + sandbox_states=[ + { + "id": "osb_pending", + "status": {"state": "Running"}, + "createdAt": "2026-01-01T00:00:00Z", + "entrypoint": ["/bin/sh"], + } + ], + ) + + handle = await OpenSandboxSandboxConfig().acreate_sandbox( + template="custom/image:latest", + timeout=600, + allow_internet_access=False, + api_key="osb-key", + api_base="https://sandbox.example/v1", + metadata={"suite": "unit"}, + env_vars={"PYTHONUNBUFFERED": "1"}, + resource_limits={"cpu": "500m", "memory": "512Mi"}, + resource_requests={"cpu": "250m", "memory": "256Mi"}, + entrypoint=["/bin/sh", "-lc", "sleep 3600"], + use_server_proxy=True, + client=client, + ) + + _, create_url, create_headers, body, _ = client.calls[0] + _, poll_url, poll_headers, _, _ = client.calls[1] + _, endpoint_url, endpoint_headers, _, endpoint_params = client.calls[2] + + assert create_url == "https://sandbox.example/v1/sandboxes" + assert create_headers["OPEN-SANDBOX-API-KEY"] == "osb-key" + assert body["image"] == {"uri": "custom/image:latest"} + assert body["entrypoint"] == ["/bin/sh", "-lc", "sleep 3600"] + assert body["metadata"] == {"suite": "unit"} + assert body["env"] == {"PYTHONUNBUFFERED": "1"} + assert body["resourceLimits"] == {"cpu": "500m", "memory": "512Mi"} + assert body["resourceRequests"] == {"cpu": "250m", "memory": "256Mi"} + assert body["networkPolicy"] == {"defaultAction": "deny", "egress": []} + assert poll_url == "https://sandbox.example/v1/sandboxes/osb_pending" + assert poll_headers["OPEN-SANDBOX-API-KEY"] == "osb-key" + assert endpoint_url.endswith("/sandboxes/osb_pending/endpoints/44772") + assert endpoint_headers["OPEN-SANDBOX-API-KEY"] == "osb-key" + assert endpoint_params == {"use_server_proxy": True} + assert handle.id == "osb_pending" + + +@pytest.mark.asyncio +async def test_create_waits_across_pending_state(monkeypatch): + client = FakeHTTPClient( + create_json={ + "id": "osb_pending", + "status": {"state": "Pending"}, + "createdAt": "2026-01-01T00:00:00Z", + }, + sandbox_states=[ + {"id": "osb_pending", "status": {"state": "Pending"}}, + {"id": "osb_pending", "status": {"state": "Running"}}, + ], + ) + sleeps = [] + + async def fake_sleep(interval): + sleeps.append(interval) + + monkeypatch.setattr( + "litellm.llms.opensandbox.sandbox.transformation.asyncio.sleep", fake_sleep + ) + + handle = await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", + api_base=TEST_API_BASE, + ready_timeout=1, + poll_interval=0.01, + client=client, + ) + + assert handle.id == "osb_pending" + assert sleeps == [0.01] + + +@pytest.mark.asyncio +async def test_create_raises_for_terminal_state(): + client = FakeHTTPClient( + create_json={"id": "osb_failed", "status": {"state": "Pending"}}, + sandbox_states=[ + {"id": "osb_failed", "status": {"state": "Failed"}}, + ], + ) + + with pytest.raises(ValueError, match="entered Failed"): + await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", api_base=TEST_API_BASE, client=client + ) + + +@pytest.mark.asyncio +async def test_create_times_out_waiting_for_running(): + client = FakeHTTPClient( + create_json={"id": "osb_slow", "status": {"state": "Pending"}}, + sandbox_states=[ + {"id": "osb_slow", "status": {"state": "Pending"}}, + ], + ) + + with pytest.raises(TimeoutError, match="was not Running"): + await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", + api_base=TEST_API_BASE, + ready_timeout=0, + poll_interval=0, + client=client, + ) + + +@pytest.mark.asyncio +async def test_create_waits_for_endpoint_resolution(monkeypatch): + client = FakeHTTPClient( + endpoint_responses=[ + http_status_error(404, f"{TEST_API_BASE}/sandboxes/osb_123"), + { + "endpoint": "execd.local:44772", + "headers": {"X-EXECD-ACCESS-TOKEN": "execd-token"}, + }, + ], + ) + sleeps = [] + + async def fake_sleep(interval): + sleeps.append(interval) + + monkeypatch.setattr( + "litellm.llms.opensandbox.sandbox.transformation.asyncio.sleep", fake_sleep + ) + + handle = await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", + api_base=TEST_API_BASE, + ready_timeout=1, + poll_interval=0.01, + client=client, + ) + + endpoint_calls = [call for call in client.calls if "/endpoints/44772" in call[1]] + assert handle._hidden_params["execd_endpoint"] == "execd.local:44772" + assert len(endpoint_calls) == 2 + assert sleeps == [0.01] + + +@pytest.mark.asyncio +async def test_create_raises_when_endpoint_is_missing(): + client = FakeHTTPClient(endpoint_json={"headers": {"X": "y"}}) + + with pytest.raises(TimeoutError, match="execd endpoint.*not ready"): + await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", api_base=TEST_API_BASE, ready_timeout=0, client=client + ) + + +@pytest.mark.asyncio +async def test_create_reraises_non_404_endpoint_error(): + client = FakeHTTPClient(endpoint_responses=[http_status_error(500)]) + + with pytest.raises(httpx.HTTPStatusError): + await OpenSandboxSandboxConfig().acreate_sandbox( + api_key="", api_base=TEST_API_BASE, client=client + ) + + +@pytest.mark.asyncio +async def test_run_code_resolves_bare_id_and_posts_sse_request(): + client = FakeHTTPClient( + execute_lines=[ + sse({"type": "stdout", "text": "42\n"}), + ] + ) + + result = await OpenSandboxSandboxConfig().arun_code( + container="osb_bare", + code="print(6*7)", + language="python", + api_key="", + api_base="http://sandbox.local/v1", + client=client, + ) + + endpoint_call = client.calls[0] + run_call = client.calls[1] + assert endpoint_call[0] == "GET" + assert ( + endpoint_call[1] == "http://sandbox.local/v1/sandboxes/osb_bare/endpoints/44772" + ) + assert run_call[0] == "POST" + assert run_call[1] == "http://execd.local:44772/code" + assert run_call[2]["X-EXECD-ACCESS-TOKEN"] == "execd-token" + assert run_call[3] == { + "code": "print(6*7)", + "context": {"language": "python"}, + } + assert run_call[4] == {"stream": True} + assert result.stdout == "42\n" + + +@pytest.mark.asyncio +async def test_run_code_uses_https_for_scheme_less_endpoint_when_api_base_is_https(): + client = FakeHTTPClient() + handle = ContainerHandle( + id="osb_https", provider="opensandbox", domain="https://sandbox.example/v1" + ) + handle._hidden_params = { + "execd_endpoint": "execd.example/route/44772", + "execd_headers": {}, + } + + await OpenSandboxSandboxConfig().arun_code( + container=handle, code="print(1)", client=client + ) + + assert client.calls[0][1] == "https://execd.example/route/44772/code" + + +@pytest.mark.asyncio +async def test_run_code_aborts_on_output_over_cap(): + client = FakeHTTPClient(execute_lines=["x" * (MAX_OUTPUT_BYTES + 1)]) + handle = ContainerHandle(id="osb_big", provider="opensandbox", domain="http://x/v1") + handle._hidden_params = {"execd_endpoint": "execd.local:44772", "execd_headers": {}} + + with pytest.raises(ValueError, match="exceeded"): + await OpenSandboxSandboxConfig().arun_code( + container=handle, code="print('x')", client=client + ) + + +@pytest.mark.asyncio +async def test_delete_returns_false_on_404(): + client = FakeHTTPClient(delete_status=404) + + ok = await OpenSandboxSandboxConfig().adelete_sandbox( + container="osb_gone", + api_key="", + api_base="http://sandbox.local/v1", + client=client, + ) + + assert ok is False + + +@pytest.mark.asyncio +async def test_delete_reraises_non_404_http_error(): + client = FakeHTTPClient(delete_status=500) + + with pytest.raises(httpx.HTTPStatusError): + await OpenSandboxSandboxConfig().adelete_sandbox( + container="osb_err", + api_key="", + api_base="http://sandbox.local/v1", + client=client, + ) + + +@pytest.mark.asyncio +async def test_public_lifecycle_create_run_delete(): + client = FakeHTTPClient( + execute_lines=[ + sse({"type": "stdout", "text": "42\n"}), + ] + ) + + container = await litellm.acreate_sandbox( + provider="opensandbox", api_key="", api_base=TEST_API_BASE, client=client + ) + result = await litellm.arun_code( + provider="opensandbox", + container=container, + code="print(6*7)", + api_key="", + client=client, + ) + ok = await litellm.adelete_sandbox( + provider="opensandbox", + container=container, + api_key="", + client=client, + ) + + assert container.id == "osb_123" + assert result.stdout == "42\n" + assert ok is True + + +@pytest.mark.asyncio +async def test_code_interpreter_tool_deletes_even_when_run_raises(): + client = FakeHTTPClient(execute_raises=RuntimeError("boom")) + + with pytest.raises(RuntimeError, match="boom"): + await litellm.acode_interpreter_tool( + provider="opensandbox", + code="1/0", + api_key="", + api_base=TEST_API_BASE, + client=client, + ) + + assert [call[0] for call in client.calls] == ["POST", "GET", "POST", "DELETE"] + assert client.calls[0][1].endswith("/sandboxes") + assert client.calls[1][1].endswith("/endpoints/44772") + assert client.calls[2][1].endswith("/code") + assert client.calls[3][1].endswith("/sandboxes/osb_123") From 7020e1e5f75fe6f98480753825589345d027b3c7 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 23 Jun 2026 09:11:49 -0700 Subject: [PATCH 06/29] feat(mcp): add resolve_credentials dispatch skeleton (#31056) PR3 of the MCP v2 outbound-credential migration, stacked on the typed vocabulary. Adds resolver.py: UpstreamCredentialProvider.resolve_credentials dispatches on the declared AuthConfig variant with one arm per mode, a wildcard-free match plus an assert_never tail so a missing arm fails basedpyright's exhaustiveness gate. Every arm is a not_implemented stub returning a typed CredError; each mode's real body and seam land in follow-up PRs. Pure v2, no v1 imports, nothing wired onto a request path. --- .../outbound_credentials/__init__.py | 15 ++-- .../outbound_credentials/resolver.py | 70 +++++++++++++++++++ .../outbound_credentials/test_resolver.py | 57 +++++++++++++++ 3 files changed, 137 insertions(+), 5 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py index b357504979d..73166a45d6e 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/__init__.py @@ -1,16 +1,20 @@ """Typed upstream-credential resolution for MCP servers. -This subpackage houses the typed credential vocabulary and (in a later PR) the -``resolve_credentials`` dispatch. A server declares one per-mode config from the -``AuthConfig`` discriminated union; failures are modeled as values via :mod:`.result` -(``Result[T, CredError]``) rather than raised, so every seam is total. Nothing here is -wired onto a live request path yet. +This subpackage houses the typed credential vocabulary and the ``resolve_credentials`` +dispatch. A server declares one per-mode config from the ``AuthConfig`` discriminated union; +``UpstreamCredentialProvider.resolve_credentials`` selects one arm and returns an ``httpx.Auth`` +or a typed ``CredError``. Failures are modeled as values via :mod:`.result` (``Result[T, +CredError]``) rather than raised, so every seam is total. Nothing here is wired onto a live +request path yet. """ from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( NoOpAuth, StaticHeaderAuth, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import ( + UpstreamCredentialProvider, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Error, Ok, @@ -45,6 +49,7 @@ __all__ = [ "Result", "NoOpAuth", "StaticHeaderAuth", + "UpstreamCredentialProvider", "AuthSpecKind", "CredError", "Subject", diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py new file mode 100644 index 00000000000..7bcdb3e6529 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -0,0 +1,70 @@ +"""The one credential resolver: dispatch on the declared mode, fail closed. + +`resolve_credentials` selects exactly one arm off the server's typed `config` and either +produces an `httpx.Auth` or returns a typed `CredError`. The `match` is over the `AuthConfig` +variant, so each arm receives its own fully-typed config with no field-presence inference and +no precedence cascade. It is wildcard-free with an `assert_never` tail, so adding a mode without +an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly +at runtime instead of returning `None`. + +This skeleton ships every arm as a `not_implemented` stub. Each mode's real body, with its +injected seam, lands in its own follow-up PR; until then the arm returns a typed error rather +than silently producing no credential. Pure v2: no imports from v1. +""" + +from __future__ import annotations + +import httpx +from typing_extensions import assert_never + +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( + Error, + Result, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + ApiKeyConfig, + AuthorizationCodeConfig, + AuthSpecKind, + AwsSigV4Config, + ClientCredentialsConfig, + CredError, + NoneConfig, + PassthroughConfig, + ServerSpec, + Subject, + TokenExchangeConfig, +) + + +class UpstreamCredentialProvider: + """Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode. + + Collaborators (the per-mode credential stores and token fetchers) are injected as each arm + is built; the skeleton needs none, since every arm is a stub. + """ + + async def resolve_credentials( + self, subject: Subject, server: ServerSpec + ) -> Result[httpx.Auth, CredError]: + match server.config: + case NoneConfig(): + return _not_implemented(AuthSpecKind.none) + case ApiKeyConfig(): + return _not_implemented(AuthSpecKind.api_key) + case PassthroughConfig(): + return _not_implemented(AuthSpecKind.passthrough) + case ClientCredentialsConfig(): + return _not_implemented(AuthSpecKind.client_credentials) + case TokenExchangeConfig(): + return _not_implemented(AuthSpecKind.token_exchange) + case AuthorizationCodeConfig(): + return _not_implemented(AuthSpecKind.authorization_code) + case AwsSigV4Config(): + return _not_implemented(AuthSpecKind.aws_sigv4) + assert_never(server.config) + + +def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]: + return Error( + CredError.of_not_implemented(f"{kind.value}: resolver arm not implemented yet") + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py new file mode 100644 index 00000000000..7885617aa46 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -0,0 +1,57 @@ +"""Tests for the resolver dispatch skeleton. + +Every mode must reach its own arm and, until that arm is built, return a typed +`not_implemented` CredError rather than silently producing no credential. Parametrizing over +one config per mode also guards reachability: if a `case` were dropped, that mode would fall to +the `assert_never` tail and raise here instead of returning the stub. +""" + +import pytest +from pydantic import SecretStr + +from litellm.proxy._experimental.mcp_server.outbound_credentials import ( + ApiKeyConfig, + AuthorizationCodeConfig, + AuthSpecKind, + AwsSigV4Config, + ClientCredentialsConfig, + Error, + NoneConfig, + PassthroughConfig, + ServerSpec, + SharedKey, + Subject, + TokenExchangeConfig, + UpstreamCredentialProvider, +) + +_ONE_CONFIG_PER_MODE = [ + (AuthSpecKind.none, NoneConfig()), + (AuthSpecKind.api_key, ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")))), + (AuthSpecKind.passthrough, PassthroughConfig()), + (AuthSpecKind.client_credentials, ClientCredentialsConfig()), + (AuthSpecKind.token_exchange, TokenExchangeConfig()), + (AuthSpecKind.authorization_code, AuthorizationCodeConfig()), + (AuthSpecKind.aws_sigv4, AwsSigV4Config(region="us-east-1")), +] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind, config", _ONE_CONFIG_PER_MODE) +async def test_every_mode_reaches_its_arm_and_returns_not_implemented(kind, config): + spec = ServerSpec( + server_id="s", resource="https://upstream.example.com", config=config + ) + subject = Subject(tenant_id="", subject_id="") + + result = await UpstreamCredentialProvider().resolve_credentials(subject, spec) + + assert isinstance(result, Error) + assert result.error.tag == "not_implemented" + assert kind.value in result.error.summary + + +def test_all_seven_modes_are_covered(): + # Guards that the parametrization (and therefore the dispatch) spans every AuthSpecKind, so a + # newly added mode without a test row is caught here rather than slipping through. + assert {kind for kind, _ in _ONE_CONFIG_PER_MODE} == set(AuthSpecKind) From 02f63d20bb563f7ebbd5396d4be85db66a4f26e4 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 23 Jun 2026 19:16:15 +0200 Subject: [PATCH 07/29] fix(router): guard num_retries=None in async_function_with_retries (#30036) When num_retries reaches async_function_with_retries as None - e.g. a caller passes num_retries=None explicitly (dict.get() does not fall back on an existing None value), an auto_router/complexity_router path does not propagate it, or Router.update_settings(num_retries=None) is used - the comparison `if num_retries > 0:` raised: TypeError: '>' not supported between instances of 'NoneType' and 'int' This only surfaced when the underlying call failed with a retryable error (rate limit / connection / 5xx), so the real upstream error was masked by a confusing TypeError. Normalise an explicit num_retries=None to the router default in _update_kwargs_before_fallbacks (falling back to 0 when the router default is itself None, and preserving an explicit 0), and keep the matching guard at the single pop site in async_function_with_retries as the safety net for paths that bypass the setter. Adds regression tests to test_router_per_deployment_num_retries.py. Relates to #23316, #25889, #23699, #28126 Co-authored-by: Cursor Co-authored-by: yuneng-jiang --- litellm/router.py | 19 ++- .../test_router_per_deployment_num_retries.py | 133 +++++++++++++++++- 2 files changed, 149 insertions(+), 3 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index acce1c58a8b..c687b2d9f67 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3134,7 +3134,18 @@ class Router: - litellm_trace_id - metadata """ - kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries) + # Normalise an explicit num_retries=None to the router default here (dict.get() + # only falls back when the key is absent, not when its value is None), then to 0 + # if the router default is itself None - mirroring the guard in + # async_function_with_retries, which remains the safety net for paths that bypass + # this setter. + _req_num_retries = kwargs.get("num_retries") + if _req_num_retries is not None: + kwargs["num_retries"] = _req_num_retries + else: + kwargs["num_retries"] = ( + self.num_retries if self.num_retries is not None else 0 + ) kwargs.setdefault("litellm_trace_id", str(uuid.uuid4())) model_group_alias: Optional[str] = None if self._get_model_from_alias(model=model): @@ -6931,7 +6942,11 @@ class Router: "model_group_retry_policy", self.model_group_retry_policy ) model_group: Optional[str] = kwargs.get("model") - num_retries = kwargs.pop("num_retries") + num_retries = kwargs.pop("num_retries", None) + if num_retries is None: + # Fall back to the router setting (then 0) so the comparisons below never + # hit `None > int`, which would mask the real upstream error with a TypeError. + num_retries = self.num_retries if self.num_retries is not None else 0 ## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking _metadata: dict = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {} diff --git a/tests/test_litellm/test_router_per_deployment_num_retries.py b/tests/test_litellm/test_router_per_deployment_num_retries.py index 154ba579e4e..af2372616a6 100644 --- a/tests/test_litellm/test_router_per_deployment_num_retries.py +++ b/tests/test_litellm/test_router_per_deployment_num_retries.py @@ -4,8 +4,9 @@ GitHub Issue: #18968 - Per-deployment max_retries/num_retries in litellm_params """ import pytest -from unittest.mock import MagicMock, patch +from unittest.mock import patch +import litellm from litellm import Router @@ -188,3 +189,133 @@ class TestPerDeploymentNumRetries: # Verify num_retries was converted from string to int assert exc.num_retries == 6 + + +class TestNumRetriesNoneGuard: + """ + Regression tests for the num_retries=None TypeError in async_function_with_retries. + + When num_retries reaches async_function_with_retries as None - e.g. a caller passes + num_retries=None explicitly (dict.get() does not fall back on an existing None value), + an auto_router/complexity_router path does not propagate it, or + Router.update_settings(num_retries=None) is used - AND the underlying call fails with a + retryable error, the comparison `if num_retries > 0:` raised: + + TypeError: '>' not supported between instances of 'NoneType' and 'int' + + This masked the real upstream error (rate limit / connection / 5xx) behind a TypeError. + Related issues: #23316, #25889, #23699, #28126. + """ + + @staticmethod + def _mock_router(num_retries=2): + return Router( + model_list=[ + { + "model_name": "mock-model", + "litellm_params": { + "model": "gpt-4o-mini", + "mock_response": "ok", + }, + } + ], + num_retries=num_retries, + ) + + def test_update_kwargs_normalises_explicit_none_to_router_default(self): + """ + _update_kwargs_before_fallbacks must normalise an explicit num_retries=None to + the router default (not leave it as None), while preserving an explicit 0. + """ + router = self._mock_router(num_retries=4) + + # explicit None -> router default + kwargs = {"num_retries": None} + router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs) + assert kwargs["num_retries"] == 4 + + # explicit 0 is preserved (retries stay disabled) + kwargs = {"num_retries": 0} + router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs) + assert kwargs["num_retries"] == 0 + + # absent -> router default (unchanged behaviour) + kwargs = {} + router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs) + assert kwargs["num_retries"] == 4 + + # explicit None with router default also None -> 0 (mirrors the downstream guard) + router.num_retries = None # simulate update_settings(num_retries=None) (#28126) + kwargs = {"num_retries": None} + router._update_kwargs_before_fallbacks(model="mock-model", kwargs=kwargs) + assert kwargs["num_retries"] == 0 + + @pytest.mark.asyncio + async def test_acompletion_num_retries_none_does_not_raise_typeerror(self): + """ + Per-request num_retries=None + a retryable error must NOT raise TypeError. + The router falls back to its configured num_retries and retries the (transient) + error, so the request succeeds. + """ + router = self._mock_router(num_retries=2) + with patch("asyncio.sleep", return_value=None): + response = await router.acompletion( + model="mock-model", + messages=[{"role": "user", "content": "hi"}], + num_retries=None, # the trigger + mock_testing_rate_limit_error=True, # retryable error path + ) + assert response.choices[0].message.content == "ok" + + @pytest.mark.asyncio + async def test_async_function_with_retries_none_falls_back_to_zero(self): + """ + When both the per-request value AND the router-level setting are None + (e.g. after Router.update_settings(num_retries=None), #28126), num_retries must + fall back to 0 and the real retryable error must surface - not a TypeError. + """ + router = self._mock_router(num_retries=0) + router.num_retries = None # simulate update_settings(num_retries=None) + + async def failing_fn(*args, **kwargs): + raise litellm.RateLimitError( + message="boom", model="mock-model", llm_provider="openai" + ) + + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.RateLimitError): + await router.async_function_with_retries( + original_function=failing_fn, + model="mock-model", + messages=[{"role": "user", "content": "hi"}], + num_retries=None, + ) + + @pytest.mark.asyncio + async def test_async_function_with_retries_none_falls_back_to_router_default(self): + """ + A None per-request num_retries falls back to the router-level setting, so retries + still happen (original_function is invoked more than once) before the real error + is raised - proving None did not silently disable retries or crash. + """ + router = self._mock_router(num_retries=3) + calls = {"n": 0} + + async def failing_fn(*args, **kwargs): + calls["n"] += 1 + raise litellm.InternalServerError( + message="boom", model="mock-model", llm_provider="openai" + ) + + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.InternalServerError): + await router.async_function_with_retries( + original_function=failing_fn, + model="mock-model", + messages=[{"role": "user", "content": "hi"}], + metadata={}, # populated by acompletion in the real path; log_retry needs it + num_retries=None, + ) + + # 1 initial attempt + at least 1 retry -> proves None fell back to a positive int + assert calls["n"] >= 2 From a5b75e8bab4b0ed6fda317c1fd3884bfe747666a Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 23 Jun 2026 10:33:31 -0700 Subject: [PATCH 08/29] fix(ui): keep team Organization optional for proxy admins in single-org setups (#30861) The Create Team form auto-selected, disabled, and required the Organization field whenever exactly one organization existed, regardless of role. For a proxy admin the organization is optional, so single-org setups could not create a standalone team even though the field is presented as optional. Gate the single-org preselect, the disabled state, and the restrictive help text on the org-admin role so they apply only to org admins, who must scope a team to their organization. Proxy admins now keep an optional, clearable, empty organization field regardless of how many organizations exist, matching the multi-org behavior. The /team/new endpoint already accepts a null organization, so this was a UI-only restriction. --- .../src/components/OldTeams.test.tsx | 49 +++++++++++++++++++ .../src/components/OldTeams.tsx | 11 +++-- 2 files changed, 55 insertions(+), 5 deletions(-) diff --git a/ui/litellm-dashboard/src/components/OldTeams.test.tsx b/ui/litellm-dashboard/src/components/OldTeams.test.tsx index d777ba1b0dc..0b6e5786aaf 100644 --- a/ui/litellm-dashboard/src/components/OldTeams.test.tsx +++ b/ui/litellm-dashboard/src/components/OldTeams.test.tsx @@ -1097,3 +1097,52 @@ describe("OldTeams - delete team warning copy", () => { ); }); }); + +describe("OldTeams - LIT-2530 organization stays optional for proxy admin with a single org", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockTeamInfoView.mockClear(); + vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue(["gpt-4"]); + vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]); + vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] }); + vi.mocked(teamListCall).mockResolvedValue({ teams: [], total: 0, page: 1, page_size: 100, total_pages: 1 }); + vi.mocked(teamCreateCall).mockResolvedValue({ + team_id: "new-team-1", + team_alias: "No Org Team", + models: ["gpt-4"], + organization_id: null, + keys: [], + members_with_roles: [], + spend: 0, + }); + mockUseOrganizations.mockReturnValue({ + data: [{ organization_id: "org-1", organization_alias: "Org 1", models: [], members: [] }], + }); + }); + + it("creates a team with no organization when exactly one organization exists", async () => { + renderWithQueryClient(); + + const createButton = screen.getAllByRole("button", { name: /create team/i })[0]; + act(() => { + fireEvent.click(createButton); + }); + + await waitFor(() => { + expect(screen.getByLabelText(/team name/i)).toBeInTheDocument(); + }); + + fireEvent.change(screen.getByLabelText(/team name/i), { target: { value: "No Org Team" } }); + fireEvent.change(screen.getByTestId("create-team-models-select"), { target: { value: "gpt-4" } }); + + const submitButtons = screen.getAllByRole("button", { name: /create team/i }); + fireEvent.click(submitButtons[submitButtons.length - 1]); + + await waitFor(() => { + expect(teamCreateCall).toHaveBeenCalledWith( + "test-token", + expect.objectContaining({ team_alias: "No Org Team", organization_id: null }), + ); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/OldTeams.tsx b/ui/litellm-dashboard/src/components/OldTeams.tsx index adfec4bdf6a..be9015e3730 100644 --- a/ui/litellm-dashboard/src/components/OldTeams.tsx +++ b/ui/litellm-dashboard/src/components/OldTeams.tsx @@ -262,14 +262,15 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser useEffect(() => { if (isTeamModalVisible) { const adminOrgs = getAdminOrganizations(userRole, userID, organizations); + const isOrgAdmin = userRole !== "Admin"; - // If there's exactly one organization the user is admin for, preselect it - if (adminOrgs.length === 1) { + // Org admins must scope a team to an org, so with exactly one we preselect it. + // Proxy admins can create org-less teams, so the field stays optional regardless of org count. + if (isOrgAdmin && adminOrgs.length === 1) { const org = adminOrgs[0]; form.setFieldValue("organization_id", org.organization_id); setCurrentOrgForCreateTeam(org); } else { - // Reset the organization selection for multiple orgs form.setFieldValue("organization_id", currentOrg?.organization_id || null); setCurrentOrgForCreateTeam(currentOrg); } @@ -1132,7 +1133,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser : [] } help={ - isSingleOrg + isOrgAdmin && isSingleOrg ? "You can only create teams within this organization" : isOrgAdmin ? "required" @@ -1142,7 +1143,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser