diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index f5509819cd..3c455ab71b 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -2235,6 +2235,11 @@ "type": "bool", "hint": "关闭 Ollama 思考模式。", }, + "reasoning_key": { + "description": "思考内容字段名", + "type": "string", + "hint": "从响应中提取思考内容的字段名。默认 reasoning_content(DeepSeek/Moonshot/阿里百炼等);OpenRouter 及其兼容中转渠道使用 reasoning。", + }, "custom_extra_body": { "description": "自定义请求体参数", "type": "dict", diff --git a/astrbot/core/provider/sources/openai_source.py b/astrbot/core/provider/sources/openai_source.py index f7870b7137..ac45e4f8d5 100644 --- a/astrbot/core/provider/sources/openai_source.py +++ b/astrbot/core/provider/sources/openai_source.py @@ -1,1447 +1,1464 @@ -import asyncio -import copy -import inspect -import json -import random -import re -from collections.abc import AsyncGenerator -from typing import Any, Literal - -import httpx -from openai import AsyncAzureOpenAI, AsyncOpenAI -from openai._exceptions import NotFoundError -from openai.lib.streaming.chat._completions import ChatCompletionStreamState -from openai.types.chat.chat_completion import ChatCompletion -from openai.types.chat.chat_completion_chunk import ChatCompletionChunk -from openai.types.completion_usage import CompletionUsage - -import astrbot.core.message.components as Comp -from astrbot import logger -from astrbot.api.provider import Provider -from astrbot.core.agent.message import ( - AudioURLPart, - ContentPart, - ImageURLPart, - Message, - TextPart, -) -from astrbot.core.agent.tool import ToolSet -from astrbot.core.exceptions import EmptyModelOutputError -from astrbot.core.message.message_event_result import MessageChain -from astrbot.core.provider.entities import LLMResponse, TokenUsage, ToolCallsResult -from astrbot.core.utils.media_utils import ( - describe_media_ref, - resolve_media_ref_to_base64_data, -) -from astrbot.core.utils.network_utils import ( - create_proxy_client, - is_connection_error, - log_connection_failure, -) -from astrbot.core.utils.string_utils import normalize_and_dedupe_strings - -from ..register import register_provider_adapter -from .request_retry import retry_provider_request - - -@register_provider_adapter( - "openai_chat_completion", - "OpenAI API Chat Completion 提供商适配器", -) -class ProviderOpenAIOfficial(Provider): - _ERROR_TEXT_CANDIDATE_MAX_CHARS = 4096 - - @classmethod - def _truncate_error_text_candidate(cls, text: str) -> str: - if len(text) <= cls._ERROR_TEXT_CANDIDATE_MAX_CHARS: - return text - return text[: cls._ERROR_TEXT_CANDIDATE_MAX_CHARS] - - @staticmethod - def _safe_json_dump(value: Any) -> str | None: - try: - return json.dumps(value, ensure_ascii=False, default=str) - except Exception: - return None - - def _get_image_moderation_error_patterns(self) -> list[str]: - """Return configured moderation patterns (case-insensitive substring match, not regex).""" - configured = self.provider_config.get("image_moderation_error_patterns", []) - patterns: list[str] = [] - if isinstance(configured, str): - configured = [configured] - if isinstance(configured, list): - for pattern in configured: - if not isinstance(pattern, str): - continue - pattern = pattern.strip() - if pattern: - patterns.append(pattern) - return patterns - - @staticmethod - def _extract_error_text_candidates(error: Exception) -> list[str]: - candidates: list[str] = [] - - def _append_candidate(candidate: Any): - if candidate is None: - return - text = str(candidate).strip() - if not text: - return - candidates.append( - ProviderOpenAIOfficial._truncate_error_text_candidate(text) - ) - - _append_candidate(str(error)) - - body = getattr(error, "body", None) - if isinstance(body, dict): - err_obj = body.get("error") - body_text = ProviderOpenAIOfficial._safe_json_dump( - {"error": err_obj} if isinstance(err_obj, dict) else body - ) - _append_candidate(body_text) - if isinstance(err_obj, dict): - for field in ("message", "type", "code", "param"): - value = err_obj.get(field) - if value is not None: - _append_candidate(value) - elif isinstance(body, str): - _append_candidate(body) - - response = getattr(error, "response", None) - if response is not None: - response_text = getattr(response, "text", None) - if isinstance(response_text, str): - _append_candidate(response_text) - - return normalize_and_dedupe_strings(candidates) - - def _is_content_moderated_upload_error(self, error: Exception) -> bool: - patterns = [ - pattern.lower() for pattern in self._get_image_moderation_error_patterns() - ] - if not patterns: - return False - candidates = [ - candidate.lower() - for candidate in self._extract_error_text_candidates(error) - ] - for pattern in patterns: - if any(pattern in candidate for candidate in candidates): - return True - return False - - @staticmethod - def _context_contains_image(contexts: list[dict]) -> bool: - for context in contexts: - content = context.get("content") - if not isinstance(content, list): - continue - for item in content: - if isinstance(item, dict) and item.get("type") in { - "image_url", - "audio_url", - }: - return True - return False - - def _is_invalid_attachment_error(self, error: Exception) -> bool: - body = getattr(error, "body", None) - code: str | None = None - message: str | None = None - if isinstance(body, dict): - err_obj = body.get("error") - if isinstance(err_obj, dict): - raw_code = err_obj.get("code") - raw_message = err_obj.get("message") - code = raw_code.lower() if isinstance(raw_code, str) else None - message = raw_message.lower() if isinstance(raw_message, str) else None - - if code == "invalid_attachment": - return True - - text_sources: list[str] = [] - if message: - text_sources.append(message) - if code: - text_sources.append(code) - text_sources.extend(map(str, self._extract_error_text_candidates(error))) - - error_text = " ".join(text.lower() for text in text_sources if text) - if "invalid_attachment" in error_text: - return True - if "download attachment" in error_text and "404" in error_text: - return True - return False - - async def _image_ref_to_data_url( - self, - image_ref: str, - *, - mode: Literal["safe", "strict"] = "safe", - ) -> str | None: - image_data = await resolve_media_ref_to_base64_data( - image_ref, - media_type="image", - strict=mode == "strict", - ) - return image_data.to_data_url() if image_data else None - - async def _resolve_image_part( - self, - image_url: str, - *, - image_detail: str | None = None, - ) -> dict | None: - image_data = await self._image_ref_to_data_url(image_url, mode="safe") - if not image_data: - logger.warning("图片预处理结果为空,将忽略。") - return None - image_payload = {"url": image_data} - - if image_detail: - image_payload["detail"] = image_detail - return { - "type": "image_url", - "image_url": image_payload, - } - - def _extract_image_part_info(self, part: dict) -> tuple[str | None, str | None]: - if not isinstance(part, dict) or part.get("type") != "image_url": - return None, None - - image_url_data = part.get("image_url") - if not isinstance(image_url_data, dict): - logger.warning("图片内容块格式无效,将保留原始内容。") - return None, None - - url = image_url_data.get("url") - if not isinstance(url, str) or not url: - logger.warning("图片内容块缺少有效 URL,将保留原始内容。") - return None, None - - image_detail = image_url_data.get("detail") - if not isinstance(image_detail, str): - image_detail = None - return url, image_detail - - def _extract_audio_part_info(self, part: dict) -> str | None: - if not isinstance(part, dict) or part.get("type") != "audio_url": - return None - - audio_url_data = part.get("audio_url") - if not isinstance(audio_url_data, dict): - logger.warning("音频内容块格式无效,将保留原始内容。") - return None - - url = audio_url_data.get("url") - if not isinstance(url, str) or not url: - logger.warning("音频内容块缺少有效路径,将保留原始内容。") - return None - - return url - - async def _resolve_audio_part(self, audio_ref: str) -> dict | None: - try: - audio_data = await resolve_media_ref_to_base64_data( - audio_ref, - media_type="audio", - strict=True, - ) - except Exception as exc: - logger.warning("音频预处理失败,将忽略。错误: %s", exc) - return None - - if not audio_data or not audio_data.format: - logger.warning("音频预处理结果为空,将忽略。") - return None - - return { - "type": "input_audio", - "input_audio": { - "data": audio_data.base64_data, - "format": audio_data.format, - }, - } - - async def _transform_content_part(self, part: dict) -> dict: - if not isinstance(part, dict): - return part - - if part.get("type") == "image_url": - url, image_detail = self._extract_image_part_info(part) - if not url: - return part - - try: - resolved_part = await self._resolve_image_part( - url, image_detail=image_detail - ) - except Exception as exc: - logger.warning( - "图片 %s 预处理失败,将保留原始内容。错误: %s", - url, - exc, - ) - return part - - return resolved_part or part - - if part.get("type") == "audio_url": - audio_ref = self._extract_audio_part_info(part) - if not audio_ref: - return part - resolved_part = await self._resolve_audio_part(audio_ref) - return resolved_part or part - - return part - - async def _materialize_message_image_parts(self, message: dict) -> dict: - content = message.get("content") - if not isinstance(content, list): - return {**message} - - new_content = [await self._transform_content_part(part) for part in content] - return {**message, "content": new_content} - - async def _materialize_context_image_parts( - self, context_query: list[dict] - ) -> list[dict]: - return [ - await self._materialize_message_image_parts(message) - for message in context_query - ] - - async def _fallback_to_text_only_and_retry( - self, - payloads: dict, - context_query: list, - chosen_key: str, - available_api_keys: list[str], - func_tool: ToolSet | None, - reason: str, - *, - image_fallback_used: bool = False, - ) -> tuple: - logger.warning( - "检测到图片请求失败(%s),已移除图片并重试(保留文本内容)。", - reason, - ) - new_contexts = await self._remove_image_from_context(context_query) - payloads["messages"] = new_contexts - return ( - False, - chosen_key, - available_api_keys, - payloads, - new_contexts, - func_tool, - image_fallback_used, - ) - - def _create_http_client(self, provider_config: dict) -> httpx.AsyncClient: - """创建带代理的 HTTP 客户端""" - proxy = provider_config.get("proxy", "") - httpx_module: Any = httpx - try: - from openai import _base_client as openai_base_client - - httpx_module = getattr(openai_base_client, "httpx", httpx) - except ImportError: - pass - return create_proxy_client("OpenAI", proxy, httpx_module=httpx_module) - - def __init__(self, provider_config, provider_settings) -> None: - super().__init__(provider_config, provider_settings) - self.chosen_api_key = None - self.api_keys: list = super().get_keys() - self.chosen_api_key = self.api_keys[0] if len(self.api_keys) > 0 else None - self.timeout = provider_config.get("timeout", 120) - self.custom_headers = provider_config.get("custom_headers", {}) - if isinstance(self.timeout, str): - self.timeout = int(self.timeout) - - if not isinstance(self.custom_headers, dict) or not self.custom_headers: - self.custom_headers = None - else: - for key in self.custom_headers: - self.custom_headers[key] = str(self.custom_headers[key]) - - if "api_version" in provider_config: - # Using Azure OpenAI API - self.client = AsyncAzureOpenAI( - api_key=self.chosen_api_key, - api_version=provider_config.get("api_version", None), - default_headers=self.custom_headers, - base_url=provider_config.get("api_base", ""), - timeout=self.timeout, - http_client=self._create_http_client(provider_config), - ) - else: - # Using OpenAI Official API - self.client = AsyncOpenAI( - api_key=self.chosen_api_key, - base_url=provider_config.get("api_base", None), - default_headers=self.custom_headers, - timeout=self.timeout, - http_client=self._create_http_client(provider_config), - ) - - self.default_params = inspect.signature( - self.client.chat.completions.create, - ).parameters.keys() - - model = provider_config.get("model", "unknown") - self.set_model(model) - - self.reasoning_key = "reasoning_content" - - def _ollama_disable_thinking_enabled(self) -> bool: - value = self.provider_config.get("ollama_disable_thinking", False) - if isinstance(value, str): - return value.strip().lower() in {"1", "true", "yes", "on"} - return bool(value) - - def _apply_provider_specific_request_overrides( - self, - payloads: dict[str, Any], - extra_body: dict[str, Any], - ) -> None: - provider = self.provider_config.get("provider") - model = str(payloads.get("model", "")).lower() - - # NVIDIA's hosted MiniMax M3 endpoint can return empty choices when - # max_tokens is omitted (#9206). Scope the compatibility default to - # that model; other NVIDIA models have different token limits. - if ( - provider == "nvidia" - and model == "minimaxai/minimax-m3" - and "max_tokens" not in payloads - and "max_tokens" not in extra_body - ): - payloads["max_tokens"] = 8192 - - if provider != "ollama": - return - if not self._ollama_disable_thinking_enabled(): - return - - # Ollama's OpenAI-compatible endpoint reliably maps reasoning_effort=none - # to think=false, while direct think=false passthrough is not stable. - extra_body.pop("reasoning", None) - extra_body.pop("think", None) - extra_body["reasoning_effort"] = "none" - - async def get_models(self): - try: - models_str = [] - models = await retry_provider_request( - "OpenAI", - lambda: self.client.models.list(), - ) - models = sorted(models.data, key=lambda x: x.id) - for model in models: - models_str.append(model.id) - return models_str - except NotFoundError as e: - raise Exception(f"获取模型列表失败:{e}") - - @staticmethod - def _sanitize_assistant_messages(payloads: dict) -> None: - """在请求发送前过滤/规范化空的 assistant 消息。 - - 严格 API(Moonshot、DeepSeek Reasoner 等)会在 assistant 消息同时缺少 - ``content`` 和 ``tool_calls`` 时返回 400。把 ``""`` / ``None`` / ``[]`` - 都视作空内容:无 tool_calls 时整条过滤掉;有 tool_calls 时将 content - 设为 ``None`` 以符合 OpenAI 规范。就地修改 ``payloads["messages"]``。 - """ - messages = payloads.get("messages") - if not isinstance(messages, list): - return - - def _is_empty(content: Any) -> bool: - return content is None or content == "" or content == [] - - cleaned: list[Any] = [] - for idx, msg in enumerate(messages): - if not isinstance(msg, dict) or msg.get("role") != "assistant": - cleaned.append(msg) - continue - - content = msg.get("content") - tool_calls = msg.get("tool_calls") - reasoning_content = msg.get("reasoning_content") - - if _is_empty(content) and not tool_calls: - if not reasoning_content: - # 三者全空,真正的垃圾消息,丢弃 - logger.debug( - f"过滤第 {idx} 条空 assistant 消息 (无 content | tool_calls | reasoning_content)" - ) - continue - else: - # ⭐ 有 reasoning_content 但没有 content 和 tool_calls - # 不能丢(推理模型需要 reasoning 历史) - # 但 API 要求 content 或 tool_calls 至少有一个 - # → 设空字符串占位,满足校验 - msg["content"] = "" - - elif _is_empty(content) and tool_calls: - msg["content"] = None # 有 tool_calls,按 OpenAI 规范 - - cleaned.append(msg) - - # Drop orphaned or duplicate tool messages whose assistant(tool_calls) - # was removed by context truncation / compression. - pending_tool_call_ids: set[str] = set() - final: list = [] - removed_tool_messages = 0 - for msg in cleaned: - if not isinstance(msg, dict): - final.append(msg) - pending_tool_call_ids = set() - continue - role = msg.get("role") - if role == "assistant" and msg.get("tool_calls"): - pending_tool_call_ids = { - tc["id"] - for tc in msg["tool_calls"] - if isinstance(tc, dict) and "id" in tc - } - final.append(msg) - elif role == "tool": - tool_call_id = msg.get("tool_call_id") - if tool_call_id in pending_tool_call_ids: - final.append(msg) - pending_tool_call_ids.remove(tool_call_id) - else: - removed_tool_messages += 1 - else: - pending_tool_call_ids = set() - final.append(msg) - if removed_tool_messages: - logger.debug( - "Filtered %d orphaned or duplicate tool message(s)", - removed_tool_messages, - ) - payloads["messages"] = final - - async def _query( - self, - payloads: dict, - tools: ToolSet | None, - *, - request_max_retries: int | None = None, - ) -> LLMResponse: - if tools: - model = payloads.get("model", "").lower() - omit_empty_param_field = "gemini" in model - tool_list = tools.get_func_desc_openai_style( - omit_empty_parameter_field=omit_empty_param_field, - ) - if tool_list: - payloads["tools"] = tool_list - payloads["tool_choice"] = payloads.get("tool_choice", "auto") - - # 不在默认参数中的参数放在 extra_body 中 - extra_body = {} - to_del = [] - for key in payloads: - if key not in self.default_params: - extra_body[key] = payloads[key] - to_del.append(key) - for key in to_del: - del payloads[key] - - # 读取并合并 custom_extra_body 配置 - custom_extra_body = self.provider_config.get("custom_extra_body", {}) - if isinstance(custom_extra_body, dict): - extra_body.update(custom_extra_body) - self._apply_provider_specific_request_overrides(payloads, extra_body) - - model = payloads.get("model", "").lower() - - self._sanitize_assistant_messages(payloads) - - completion = await retry_provider_request( - "OpenAI", - lambda: self.client.chat.completions.create( - **payloads, - stream=False, - extra_body=extra_body, - ), - max_attempts=request_max_retries, - ) - - if not isinstance(completion, ChatCompletion): - raise Exception( - f"API 返回的 completion 类型错误:{type(completion)}: {completion}。", - ) - - logger.debug(f"completion: {completion}") - - llm_response = await self._parse_openai_completion(completion, tools) - - return llm_response - - async def _query_stream( - self, - payloads: dict, - tools: ToolSet | None, - *, - request_max_retries: int | None = None, - ) -> AsyncGenerator[LLMResponse, None]: - """流式查询API,逐步返回结果""" - if tools: - model = payloads.get("model", "").lower() - omit_empty_param_field = "gemini" in model - tool_list = tools.get_func_desc_openai_style( - omit_empty_parameter_field=omit_empty_param_field, - ) - if tool_list: - payloads["tools"] = tool_list - payloads["tool_choice"] = payloads.get("tool_choice", "auto") - - # 不在默认参数中的参数放在 extra_body 中 - extra_body = {} - - # 读取并合并 custom_extra_body 配置 - custom_extra_body = self.provider_config.get("custom_extra_body", {}) - if isinstance(custom_extra_body, dict): - extra_body.update(custom_extra_body) - - to_del = [] - for key in payloads: - if key not in self.default_params: - extra_body[key] = payloads[key] - to_del.append(key) - for key in to_del: - del payloads[key] - self._apply_provider_specific_request_overrides(payloads, extra_body) - - self._sanitize_assistant_messages(payloads) - - stream = await retry_provider_request( - "OpenAI", - lambda: self.client.chat.completions.create( - **payloads, - stream=True, - extra_body=extra_body, - stream_options={"include_usage": True}, - ), - max_attempts=request_max_retries, - ) - - llm_response = LLMResponse("assistant", is_chunk=True) - - state = ChatCompletionStreamState() - - async for chunk in stream: - choice = chunk.choices[0] if chunk.choices else None - delta = choice.delta if choice else None - - if delta and (dtcs := delta.tool_calls): - for idx, tc in enumerate(dtcs): - # siliconflow workaround - if tc.function and tc.function.arguments: - tc.type = "function" - # Fix for #6661: Add missing 'index' field to tool_call deltas - # Gemini and some OpenAI-compatible proxies omit this field - if not hasattr(tc, "index") or tc.index is None: - tc.index = idx - # 跳过 delta=None 的 chunk,避免 SDK 内部 _convert_initial_chunk_into_snapshot - # 第 747 行 choice.delta.to_dict() 抛出 NoneType 错误。 - # refs: AstrBot#6689 / openai-python#5069 / #5047 - # 例外:流末尾的 usage chunk(choices=[],delta=None 但有 usage 数据) - # 需要传给 state,否则最终 completion 会丢失 usage 信息 - if delta is not None or chunk.usage: - try: - state.handle_chunk(chunk) - except Exception as e: - logger.error("Saving chunk state error: " + str(e)) - # logger.debug(f"chunk delta: {delta}") - # handle the content delta - reasoning = self._extract_reasoning_content(chunk) - _y = False - llm_response.id = chunk.id - llm_response.reasoning_content = None - llm_response.completion_text = "" - if reasoning is not None: - llm_response.reasoning_content = reasoning - _y = True - if delta and delta.content: - # Don't strip streaming chunks to preserve spaces between words - completion_text = self._normalize_content(delta.content, strip=False) - llm_response.result_chain = MessageChain( - chain=[Comp.Plain(completion_text)], - ) - _y = True - if chunk.usage: - llm_response.usage = self._extract_usage(chunk.usage) - elif choice and (choice_usage := getattr(choice, "usage", None)): - # Workaround for some providers that only return usage in choices[].usage, e.g. MoonshotAI - # See https://github.com/AstrBotDevs/AstrBot/issues/6614 - llm_response.usage = self._extract_usage(choice_usage) - state.current_completion_snapshot.usage = choice_usage - if _y: - yield llm_response - - try: - final_completion = state.get_final_completion() - llm_response = await self._parse_openai_completion(final_completion, tools) - yield llm_response - except Exception as e: - logger.error("get_final_completion error: " + str(e)) - # 流式内容已通过 yield 发出,记录错误后正常结束即可 - return - - def _extract_reasoning_content( - self, - completion: ChatCompletion | ChatCompletionChunk, - ) -> str | None: - """Extract reasoning content from OpenAI ChatCompletion if available.""" - - def _get_reasoning_attr(obj: Any) -> str | None: - fields_set = getattr(obj, "model_fields_set", None) - if isinstance(fields_set, set) and self.reasoning_key in fields_set: - attr = getattr(obj, self.reasoning_key, "") - return "" if attr is None else str(attr) - attr = getattr(obj, self.reasoning_key, None) - return None if attr is None else str(attr) - - if not completion.choices: - return None - if isinstance(completion, ChatCompletion): - choice = completion.choices[0] - reasoning_attr = _get_reasoning_attr(choice.message) - elif isinstance(completion, ChatCompletionChunk): - delta = completion.choices[0].delta - reasoning_attr = _get_reasoning_attr(delta) - else: - return None - return reasoning_attr - - def _extract_usage(self, usage: CompletionUsage | dict) -> TokenUsage: - ptd = getattr(usage, "prompt_tokens_details", None) - cached = getattr(ptd, "cached_tokens", 0) if ptd else 0 - cached = ( - cached if isinstance(cached, int) else 0 - ) # ptd.cached_tokens 可能为None - prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0 # 安全 - completion_tokens = getattr(usage, "completion_tokens", 0) or 0 - cached = cached or 0 - prompt_tokens = prompt_tokens or 0 - completion_tokens = completion_tokens or 0 - return TokenUsage( - input_other=prompt_tokens - cached, - input_cached=cached, - output=completion_tokens, - ) - - @staticmethod - def _normalize_content(raw_content: Any, strip: bool = True) -> str: - """Normalize content from various formats to plain string. - - Some LLM providers return content as list[dict] format - like [{'type': 'text', 'text': '...'}] instead of - plain string. This method handles both formats. - - Args: - raw_content: The raw content from LLM response, can be str, list, dict, or other. - strip: Whether to strip whitespace from the result. Set to False for - streaming chunks to preserve spaces between words. - - Returns: - Normalized plain text string. - """ - # Handle dict format (e.g., {"type": "text", "text": "..."}) - if isinstance(raw_content, dict): - if "text" in raw_content: - text_val = raw_content.get("text", "") - return str(text_val) if text_val is not None else "" - # For other dict formats, return empty string and log - logger.warning(f"Unexpected dict format content: {raw_content}") - return "" - - if isinstance(raw_content, list): - # Check if this looks like OpenAI content-part format - # Only process if at least one item has {'type': 'text', 'text': ...} structure - has_content_part = any( - isinstance(part, dict) and part.get("type") == "text" - for part in raw_content - ) - if has_content_part: - text_parts = [] - for part in raw_content: - if isinstance(part, dict) and part.get("type") == "text": - text_val = part.get("text", "") - # Coerce to str in case text is null or non-string - text_parts.append(str(text_val) if text_val is not None else "") - return "".join(text_parts) - # Not content-part format, return string representation - return str(raw_content) - - if isinstance(raw_content, str): - content = raw_content.strip() if strip else raw_content - # Check if the string is a JSON-encoded list (e.g., "[{'type': 'text', ...}]") - # This can happen when streaming concatenates content that was originally list format - # Only check if it looks like a complete JSON array (requires strip for check) - check_content = raw_content.strip() - if ( - check_content.startswith("[") - and check_content.endswith("]") - and len(check_content) < 8192 - ): - try: - # First try standard JSON parsing - parsed = json.loads(check_content) - except json.JSONDecodeError: - # If that fails, try parsing as Python literal (handles single quotes) - # This is safer than blind replace("'", '"') which corrupts apostrophes - try: - import ast - - parsed = ast.literal_eval(check_content) - except (ValueError, SyntaxError): - parsed = None - - if isinstance(parsed, list): - # Only convert if it matches OpenAI content-part schema - # i.e., at least one item has {'type': 'text', 'text': ...} - has_content_part = any( - isinstance(part, dict) and part.get("type") == "text" - for part in parsed - ) - if has_content_part: - text_parts = [] - for part in parsed: - if isinstance(part, dict) and part.get("type") == "text": - text_val = part.get("text", "") - # Coerce to str in case text is null or non-string - text_parts.append( - str(text_val) if text_val is not None else "" - ) - if text_parts: - return "".join(text_parts) - return content - - # Fallback for other types (int, float, etc.) - return str(raw_content) if raw_content is not None else "" - - async def _parse_openai_completion( - self, completion: ChatCompletion, tools: ToolSet | None - ) -> LLMResponse: - """Parse OpenAI ChatCompletion into LLMResponse""" - llm_response = LLMResponse("assistant") - - # workaround for #9374 - if not completion.choices: - data = getattr(completion, "data", None) - if isinstance(data, dict): - try: - completion = ChatCompletion.model_validate(data) - except (TypeError, ValueError): - pass - - if not completion.choices: - raise EmptyModelOutputError( - f"OpenAI completion has no choices. response_id={completion.id}" - ) - choice = completion.choices[0] - - # parse the text completion - if choice.message.content is not None: - completion_text = self._normalize_content(choice.message.content) - # specially, some providers may set tags around reasoning content in the completion text, - # we use regex to remove them, and store then in reasoning_content field - reasoning_pattern = re.compile(r"(.*?)", re.DOTALL) - matches = reasoning_pattern.findall(completion_text) - if matches: - llm_response.reasoning_content = "\n".join( - [match.strip() for match in matches], - ) - completion_text = reasoning_pattern.sub("", completion_text).strip() - # Also clean up orphan tags that may leak from some models - completion_text = re.sub(r"\s*$", "", completion_text).strip() - llm_response.result_chain = MessageChain().message(completion_text) - elif refusal := getattr(choice.message, "refusal", None): - refusal_text = self._normalize_content(refusal) - if refusal_text: - llm_response.result_chain = MessageChain().message(refusal_text) - - # parse the reasoning content if any - # the priority is higher than the tag extraction - reasoning_content = self._extract_reasoning_content(completion) - if reasoning_content is not None: - llm_response.reasoning_content = reasoning_content - - # parse tool calls if any - if choice.message.tool_calls and tools is not None: - args_ls = [] - func_name_ls = [] - tool_call_ids = [] - tool_call_extra_content_dict = {} - for tool_call in choice.message.tool_calls: - if isinstance(tool_call, str): - # workaround for #1359 - tool_call = json.loads(tool_call) - if tools is None: - # 工具集未提供 - # Should be unreachable - raise Exception("工具集未提供") - - if tool_call.type == "function": - # workaround for #1454 - if isinstance(tool_call.function.arguments, str): - try: - args = json.loads(tool_call.function.arguments) - except json.JSONDecodeError as e: - logger.error(f"解析参数失败: {e}") - args = {} - else: - args = tool_call.function.arguments - # Some API may return None for tools with no parameters - if args is None: - args = {} - args_ls.append(args) - func_name_ls.append(tool_call.function.name) - tool_call_ids.append(tool_call.id) - - # gemini-2.5 / gemini-3 series extra_content handling - extra_content = getattr(tool_call, "extra_content", None) - if extra_content is not None: - tool_call_extra_content_dict[tool_call.id] = extra_content - - llm_response.role = "tool" - llm_response.tools_call_args = args_ls - llm_response.tools_call_name = func_name_ls - llm_response.tools_call_ids = tool_call_ids - llm_response.tools_call_extra_content = tool_call_extra_content_dict - # specially handle finish reason - if choice.finish_reason == "content_filter": - raise Exception( - "API 返回的 completion 由于内容安全过滤被拒绝(非 AstrBot)。", - ) - has_text_output = bool((llm_response.completion_text or "").strip()) - has_reasoning_output = bool((llm_response.reasoning_content or "").strip()) - if ( - not has_text_output - and not has_reasoning_output - and not llm_response.tools_call_args - ): - logger.error(f"OpenAI completion has no usable output: {completion}.") - raise EmptyModelOutputError( - "OpenAI completion has no usable output. " - f"response_id={completion.id}, finish_reason={choice.finish_reason}" - ) - - llm_response.raw_completion = completion - llm_response.id = completion.id - - llm_response.usage = ( - self._extract_usage(completion.usage) if completion.usage else TokenUsage() - ) - - return llm_response - - async def _prepare_chat_payload( - self, - prompt: str | None, - image_urls: list[str] | None = None, - audio_urls: list[str] | None = None, - contexts: list[dict] | list[Message] | None = None, - system_prompt: str | None = None, - tool_calls_result: ToolCallsResult | list[ToolCallsResult] | None = None, - model: str | None = None, - extra_user_content_parts: list[ContentPart] | None = None, - **kwargs, - ) -> tuple: - """准备聊天所需的有效载荷和上下文""" - if contexts is None: - contexts = [] - new_record = None - if prompt is not None: - new_record = await self.assemble_context( - prompt or "", - image_urls, - audio_urls, - extra_user_content_parts, - ) - context_query = copy.deepcopy(self._ensure_message_to_dicts(contexts)) - if new_record: - context_query.append(new_record) - if system_prompt: - context_query.insert(0, {"role": "system", "content": system_prompt}) - - for part in context_query: - if "_no_save" in part: - del part["_no_save"] - - # tool calls result - if tool_calls_result: - if isinstance(tool_calls_result, ToolCallsResult): - context_query.extend(tool_calls_result.to_openai_messages()) - else: - for tcr in tool_calls_result: - context_query.extend(tcr.to_openai_messages()) - - if self._context_contains_image(context_query): - context_query = await self._materialize_context_image_parts(context_query) - - model = model or self.get_model() - - payloads = {"messages": context_query, "model": model} - - self._finally_convert_payload(payloads) - - return payloads, context_query - - def _finally_convert_payload(self, payloads: dict) -> None: - """Finally convert the payload. Such as think part conversion, tool inject.""" - model = payloads.get("model", "").lower() - is_gemini = "gemini" in model - _deepseek_v4_markers = ("deepseek-v4-pro", "deepseek-v4-flash", "deepseek-v4") - is_deepseek_v4_reasoning = ( - any(marker in model for marker in _deepseek_v4_markers) - or "api.deepseek.com" in self.client.base_url.host - ) - # deepseek-chat and deepseek-reasoner now point to V4 models (per official website) - - # MiMo 推理模型(MiMo-V2.5-Pro / MiMo-V2.5 / MiMo-V2-Pro / MiMo-V2-Omni / MiMo-V2-Flash) - # 要求 assistant 历史消息必须回传 reasoning_content,否则返回 400 - mimo_reasoning_models = { - "mimo-v2.5-pro", - "mimo-v2.5", - "mimo-v2-pro", - "mimo-v2-omni", - "mimo-v2-flash", - } - is_mimo_reasoning = model in mimo_reasoning_models - for message in payloads.get("messages", []): - if message.get("role") == "assistant" and isinstance( - message.get("content"), list - ): - reasoning_content = "" - reasoning_content_present = False - new_content = [] # not including think part - for part in message["content"]: - if part.get("type") == "think": - reasoning_content_present = True - reasoning_content += str(part.get("think")) - else: - new_content.append(part) - # Some providers (Grok, etc.) reject empty content lists. - # When all parts were think blocks, fall back to None. - message["content"] = new_content or None - if reasoning_content_present: - message["reasoning_content"] = reasoning_content - - if ( - message.get("role") == "assistant" - and is_deepseek_v4_reasoning - and "reasoning_content" not in message - ): - # DeepSeek v4 reasoning models require the field on assistant - # history messages, even when the reasoning content is empty. - message["reasoning_content"] = "" - - if ( - message.get("role") == "assistant" - and is_mimo_reasoning - and "reasoning_content" not in message - ): - # MiMo 推理模型要求 assistant 历史消息回传 reasoning_content, - # 缺失时 API 返回 400。参见 MiMo 官方文档。 - message["reasoning_content"] = "" - - # Gemini 的 function_response 要求 google.protobuf.Struct(即 JSON 对象), - # 纯文本会触发 400 Invalid argument,需要包一层 JSON。 - if is_gemini and message.get("role") == "tool": - content = message.get("content", "") - if isinstance(content, str): - try: - json.loads(content) - except (json.JSONDecodeError, ValueError): - message["content"] = json.dumps( - {"result": content}, ensure_ascii=False - ) - - async def _handle_api_error( - self, - e: Exception, - payloads: dict, - context_query: list, - func_tool: ToolSet | None, - chosen_key: str, - available_api_keys: list[str], - retry_cnt: int, - max_retries: int, - image_fallback_used: bool = False, - ) -> tuple: - """处理API错误并尝试恢复""" - if "429" in str(e): - logger.warning( - f"API 调用过于频繁,尝试使用其他 Key 重试。当前 Key: {chosen_key[:12]}", - ) - # 最后一次不等待 - if retry_cnt < max_retries - 1: - await asyncio.sleep(1) - if chosen_key in available_api_keys: - available_api_keys.remove(chosen_key) - if len(available_api_keys) > 0: - chosen_key = random.choice(available_api_keys) - return ( - False, - chosen_key, - available_api_keys, - payloads, - context_query, - func_tool, - image_fallback_used, - ) - raise e - if "maximum context length" in str(e) or "context length" in str(e).lower(): - logger.warning( - f"上下文长度超过限制。尝试弹出最早的记录然后重试。当前记录条数: {len(context_query)}", - ) - await self.pop_record(context_query) - payloads["messages"] = context_query - return ( - False, - chosen_key, - available_api_keys, - payloads, - context_query, - func_tool, - image_fallback_used, - ) - if "The model is not a VLM" in str(e): # siliconcloud - if image_fallback_used or not self._context_contains_image(context_query): - raise e - # 尝试删除所有 image - return await self._fallback_to_text_only_and_retry( - payloads, - context_query, - chosen_key, - available_api_keys, - func_tool, - "model_not_vlm", - image_fallback_used=True, - ) - if self._is_content_moderated_upload_error(e): - if image_fallback_used or not self._context_contains_image(context_query): - raise e - return await self._fallback_to_text_only_and_retry( - payloads, - context_query, - chosen_key, - available_api_keys, - func_tool, - "image_content_moderated", - image_fallback_used=True, - ) - if self._is_invalid_attachment_error(e): - if image_fallback_used or not self._context_contains_image(context_query): - raise e - return await self._fallback_to_text_only_and_retry( - payloads, - context_query, - chosen_key, - available_api_keys, - func_tool, - "invalid_attachment", - image_fallback_used=True, - ) - - if ( - "Function calling is not enabled" in str(e) - or ("tool" in str(e).lower() and "support" in str(e).lower()) - or ("function" in str(e).lower() and "support" in str(e).lower()) - ): - # openai, ollama, gemini openai, siliconcloud 的错误提示与 code 不统一,只能通过字符串匹配 - logger.warning( - f"{self.get_model()} 不支持函数工具调用,已自动去除,不影响使用。如需永久关闭,可前往 WebUI 中关闭工具调用。", - ) - payloads.pop("tools", None) - return ( - False, - chosen_key, - available_api_keys, - payloads, - context_query, - None, - image_fallback_used, - ) - # logger.error(f"发生了错误。Provider 配置如下: {self.provider_config}") - - if is_connection_error(e): - proxy = self.provider_config.get("proxy", "") - log_connection_failure("OpenAI", e, proxy) - - raise e - - async def text_chat( - self, - prompt=None, - session_id=None, - image_urls=None, - audio_urls=None, - func_tool=None, - contexts=None, - system_prompt=None, - tool_calls_result=None, - model=None, - extra_user_content_parts=None, - tool_choice: Literal["auto", "required"] = "auto", - request_max_retries: int | None = None, - **kwargs, - ) -> LLMResponse: - payloads, context_query = await self._prepare_chat_payload( - prompt, - image_urls, - audio_urls, - contexts, - system_prompt, - tool_calls_result, - model=model, - extra_user_content_parts=extra_user_content_parts, - **kwargs, - ) - if func_tool and not func_tool.empty(): - payloads["tool_choice"] = tool_choice - - llm_response = None - max_retries = 10 - available_api_keys = self.api_keys.copy() - chosen_key = random.choice(available_api_keys) - image_fallback_used = False - - last_exception = None - retry_cnt = 0 - for retry_cnt in range(max_retries): - try: - self.client.api_key = chosen_key - llm_response = await self._query( - payloads, - func_tool, - request_max_retries=request_max_retries, - ) - break - except Exception as e: - last_exception = e - ( - success, - chosen_key, - available_api_keys, - payloads, - context_query, - func_tool, - image_fallback_used, - ) = await self._handle_api_error( - e, - payloads, - context_query, - func_tool, - chosen_key, - available_api_keys, - retry_cnt, - max_retries, - image_fallback_used=image_fallback_used, - ) - if success: - break - - if retry_cnt == max_retries - 1 or llm_response is None: - logger.error(f"API 调用失败,重试 {max_retries} 次仍然失败。") - if last_exception is None: - raise Exception("未知错误") - raise last_exception - return llm_response - - async def text_chat_stream( - self, - prompt=None, - session_id=None, - image_urls=None, - audio_urls=None, - func_tool=None, - contexts=None, - system_prompt=None, - tool_calls_result=None, - model=None, - tool_choice: Literal["auto", "required"] = "auto", - request_max_retries: int | None = None, - **kwargs, - ) -> AsyncGenerator[LLMResponse, None]: - """流式对话,与服务商交互并逐步返回结果""" - payloads, context_query = await self._prepare_chat_payload( - prompt, - image_urls, - audio_urls, - contexts, - system_prompt, - tool_calls_result, - model=model, - **kwargs, - ) - if func_tool and not func_tool.empty(): - payloads["tool_choice"] = tool_choice - - max_retries = 10 - available_api_keys = self.api_keys.copy() - chosen_key = random.choice(available_api_keys) - image_fallback_used = False - - last_exception = None - retry_cnt = 0 - for retry_cnt in range(max_retries): - try: - self.client.api_key = chosen_key - async for response in self._query_stream( - payloads, - func_tool, - request_max_retries=request_max_retries, - ): - yield response - break - except Exception as e: - last_exception = e - ( - success, - chosen_key, - available_api_keys, - payloads, - context_query, - func_tool, - image_fallback_used, - ) = await self._handle_api_error( - e, - payloads, - context_query, - func_tool, - chosen_key, - available_api_keys, - retry_cnt, - max_retries, - image_fallback_used=image_fallback_used, - ) - if success: - break - - if retry_cnt == max_retries - 1: - logger.error(f"API 调用失败,重试 {max_retries} 次仍然失败。") - if last_exception is None: - raise Exception("未知错误") - raise last_exception - - async def _remove_image_from_context(self, contexts: list): - """从上下文中删除所有带有 image 的记录""" - new_contexts = [] - - for context in contexts: - if "content" in context and isinstance(context["content"], list): - # continue - new_content = [] - for item in context["content"]: - if isinstance(item, dict) and "image_url" in item: - continue - new_content.append(item) - if not new_content: - # 用户只发了图片 - new_content = [{"type": "text", "text": "[图片]"}] - context["content"] = new_content - new_contexts.append(context) - return new_contexts - - def get_current_key(self) -> str: - return self.client.api_key - - def get_keys(self) -> list[str]: - return self.api_keys - - def set_key(self, key) -> None: - self.client.api_key = key - - async def assemble_context( - self, - text: str, - image_urls: list[str] | None = None, - audio_urls: list[str] | None = None, - extra_user_content_parts: list[ContentPart] | None = None, - ) -> dict: - """组装成符合 OpenAI 格式的 role 为 user 的消息段""" - - # 构建内容块列表 - content_blocks = [] - - # 1. 用户原始发言(OpenAI 建议:用户发言在前) - if text: - content_blocks.append({"type": "text", "text": text}) - elif image_urls: - # 如果没有文本但有图片,添加占位文本 - content_blocks.append({"type": "text", "text": "[Image]"}) - elif audio_urls: - content_blocks.append({"type": "text", "text": "[Audio]"}) - elif extra_user_content_parts: - # 如果只有额外内容块,也需要添加占位文本 - content_blocks.append({"type": "text", "text": " "}) - - # 2. 额外的内容块(系统提醒、指令等) - if extra_user_content_parts: - for part in extra_user_content_parts: - if isinstance(part, TextPart): - content_blocks.append({"type": "text", "text": part.text}) - elif isinstance(part, ImageURLPart): - image_part = await self._resolve_image_part( - part.image_url.url, - ) - if image_part: - content_blocks.append(image_part) - elif isinstance(part, AudioURLPart): - audio_part = await self._resolve_audio_part(part.audio_url.url) - if audio_part: - content_blocks.append(audio_part) - else: - raise ValueError(f"不支持的额外内容块类型: {type(part)}") - - # 3. 图片内容 - if image_urls: - for image_url in image_urls: - image_part = await self._resolve_image_part(image_url) - if image_part: - content_blocks.append(image_part) - - if audio_urls: - for audio_path in audio_urls: - audio_part = await self._resolve_audio_part(audio_path) - if audio_part: - content_blocks.append(audio_part) - - # 如果只有主文本且没有额外内容块和图片,返回简单格式以保持向后兼容 - if ( - text - and not extra_user_content_parts - and not image_urls - and not audio_urls - and len(content_blocks) == 1 - and content_blocks[0]["type"] == "text" - ): - return {"role": "user", "content": content_blocks[0]["text"]} - - # 否则返回多模态格式 - return {"role": "user", "content": content_blocks} - - async def encode_image_bs64(self, image_url: str) -> str: - """将图片转换为 base64""" - image_data = await self._image_ref_to_data_url(image_url, mode="strict") - if image_data is None: - raise RuntimeError( - f"Failed to encode image data: {describe_media_ref(image_url)}" - ) - return image_data - - async def terminate(self): - if self.client: - await self.client.close() +import asyncio +import copy +import inspect +import json +import random +import re +from collections.abc import AsyncGenerator +from typing import Any, Literal + +import httpx +from openai import AsyncAzureOpenAI, AsyncOpenAI +from openai._exceptions import NotFoundError +from openai.lib.streaming.chat._completions import ChatCompletionStreamState +from openai.types.chat.chat_completion import ChatCompletion +from openai.types.chat.chat_completion_chunk import ChatCompletionChunk +from openai.types.completion_usage import CompletionUsage + +import astrbot.core.message.components as Comp +from astrbot import logger +from astrbot.api.provider import Provider +from astrbot.core.agent.message import ( + AudioURLPart, + ContentPart, + ImageURLPart, + Message, + TextPart, +) +from astrbot.core.agent.tool import ToolSet +from astrbot.core.exceptions import EmptyModelOutputError +from astrbot.core.message.message_event_result import MessageChain +from astrbot.core.provider.entities import LLMResponse, TokenUsage, ToolCallsResult +from astrbot.core.utils.media_utils import ( + describe_media_ref, + resolve_media_ref_to_base64_data, +) +from astrbot.core.utils.network_utils import ( + create_proxy_client, + is_connection_error, + log_connection_failure, +) +from astrbot.core.utils.string_utils import normalize_and_dedupe_strings + +from ..register import register_provider_adapter +from .request_retry import retry_provider_request + + +@register_provider_adapter( + "openai_chat_completion", + "OpenAI API Chat Completion 提供商适配器", +) +class ProviderOpenAIOfficial(Provider): + _ERROR_TEXT_CANDIDATE_MAX_CHARS = 4096 + + @classmethod + def _truncate_error_text_candidate(cls, text: str) -> str: + if len(text) <= cls._ERROR_TEXT_CANDIDATE_MAX_CHARS: + return text + return text[: cls._ERROR_TEXT_CANDIDATE_MAX_CHARS] + + @staticmethod + def _safe_json_dump(value: Any) -> str | None: + try: + return json.dumps(value, ensure_ascii=False, default=str) + except Exception: + return None + + def _get_image_moderation_error_patterns(self) -> list[str]: + """Return configured moderation patterns (case-insensitive substring match, not regex).""" + configured = self.provider_config.get("image_moderation_error_patterns", []) + patterns: list[str] = [] + if isinstance(configured, str): + configured = [configured] + if isinstance(configured, list): + for pattern in configured: + if not isinstance(pattern, str): + continue + pattern = pattern.strip() + if pattern: + patterns.append(pattern) + return patterns + + @staticmethod + def _extract_error_text_candidates(error: Exception) -> list[str]: + candidates: list[str] = [] + + def _append_candidate(candidate: Any): + if candidate is None: + return + text = str(candidate).strip() + if not text: + return + candidates.append( + ProviderOpenAIOfficial._truncate_error_text_candidate(text) + ) + + _append_candidate(str(error)) + + body = getattr(error, "body", None) + if isinstance(body, dict): + err_obj = body.get("error") + body_text = ProviderOpenAIOfficial._safe_json_dump( + {"error": err_obj} if isinstance(err_obj, dict) else body + ) + _append_candidate(body_text) + if isinstance(err_obj, dict): + for field in ("message", "type", "code", "param"): + value = err_obj.get(field) + if value is not None: + _append_candidate(value) + elif isinstance(body, str): + _append_candidate(body) + + response = getattr(error, "response", None) + if response is not None: + response_text = getattr(response, "text", None) + if isinstance(response_text, str): + _append_candidate(response_text) + + return normalize_and_dedupe_strings(candidates) + + def _is_content_moderated_upload_error(self, error: Exception) -> bool: + patterns = [ + pattern.lower() for pattern in self._get_image_moderation_error_patterns() + ] + if not patterns: + return False + candidates = [ + candidate.lower() + for candidate in self._extract_error_text_candidates(error) + ] + for pattern in patterns: + if any(pattern in candidate for candidate in candidates): + return True + return False + + @staticmethod + def _context_contains_image(contexts: list[dict]) -> bool: + for context in contexts: + content = context.get("content") + if not isinstance(content, list): + continue + for item in content: + if isinstance(item, dict) and item.get("type") in { + "image_url", + "audio_url", + }: + return True + return False + + def _is_invalid_attachment_error(self, error: Exception) -> bool: + body = getattr(error, "body", None) + code: str | None = None + message: str | None = None + if isinstance(body, dict): + err_obj = body.get("error") + if isinstance(err_obj, dict): + raw_code = err_obj.get("code") + raw_message = err_obj.get("message") + code = raw_code.lower() if isinstance(raw_code, str) else None + message = raw_message.lower() if isinstance(raw_message, str) else None + + if code == "invalid_attachment": + return True + + text_sources: list[str] = [] + if message: + text_sources.append(message) + if code: + text_sources.append(code) + text_sources.extend(map(str, self._extract_error_text_candidates(error))) + + error_text = " ".join(text.lower() for text in text_sources if text) + if "invalid_attachment" in error_text: + return True + if "download attachment" in error_text and "404" in error_text: + return True + return False + + async def _image_ref_to_data_url( + self, + image_ref: str, + *, + mode: Literal["safe", "strict"] = "safe", + ) -> str | None: + image_data = await resolve_media_ref_to_base64_data( + image_ref, + media_type="image", + strict=mode == "strict", + ) + return image_data.to_data_url() if image_data else None + + async def _resolve_image_part( + self, + image_url: str, + *, + image_detail: str | None = None, + ) -> dict | None: + image_data = await self._image_ref_to_data_url(image_url, mode="safe") + if not image_data: + logger.warning("图片预处理结果为空,将忽略。") + return None + image_payload = {"url": image_data} + + if image_detail: + image_payload["detail"] = image_detail + return { + "type": "image_url", + "image_url": image_payload, + } + + def _extract_image_part_info(self, part: dict) -> tuple[str | None, str | None]: + if not isinstance(part, dict) or part.get("type") != "image_url": + return None, None + + image_url_data = part.get("image_url") + if not isinstance(image_url_data, dict): + logger.warning("图片内容块格式无效,将保留原始内容。") + return None, None + + url = image_url_data.get("url") + if not isinstance(url, str) or not url: + logger.warning("图片内容块缺少有效 URL,将保留原始内容。") + return None, None + + image_detail = image_url_data.get("detail") + if not isinstance(image_detail, str): + image_detail = None + return url, image_detail + + def _extract_audio_part_info(self, part: dict) -> str | None: + if not isinstance(part, dict) or part.get("type") != "audio_url": + return None + + audio_url_data = part.get("audio_url") + if not isinstance(audio_url_data, dict): + logger.warning("音频内容块格式无效,将保留原始内容。") + return None + + url = audio_url_data.get("url") + if not isinstance(url, str) or not url: + logger.warning("音频内容块缺少有效路径,将保留原始内容。") + return None + + return url + + async def _resolve_audio_part(self, audio_ref: str) -> dict | None: + try: + audio_data = await resolve_media_ref_to_base64_data( + audio_ref, + media_type="audio", + strict=True, + ) + except Exception as exc: + logger.warning("音频预处理失败,将忽略。错误: %s", exc) + return None + + if not audio_data or not audio_data.format: + logger.warning("音频预处理结果为空,将忽略。") + return None + + return { + "type": "input_audio", + "input_audio": { + "data": audio_data.base64_data, + "format": audio_data.format, + }, + } + + async def _transform_content_part(self, part: dict) -> dict: + if not isinstance(part, dict): + return part + + if part.get("type") == "image_url": + url, image_detail = self._extract_image_part_info(part) + if not url: + return part + + try: + resolved_part = await self._resolve_image_part( + url, image_detail=image_detail + ) + except Exception as exc: + logger.warning( + "图片 %s 预处理失败,将保留原始内容。错误: %s", + url, + exc, + ) + return part + + return resolved_part or part + + if part.get("type") == "audio_url": + audio_ref = self._extract_audio_part_info(part) + if not audio_ref: + return part + resolved_part = await self._resolve_audio_part(audio_ref) + return resolved_part or part + + return part + + async def _materialize_message_image_parts(self, message: dict) -> dict: + content = message.get("content") + if not isinstance(content, list): + return {**message} + + new_content = [await self._transform_content_part(part) for part in content] + return {**message, "content": new_content} + + async def _materialize_context_image_parts( + self, context_query: list[dict] + ) -> list[dict]: + return [ + await self._materialize_message_image_parts(message) + for message in context_query + ] + + async def _fallback_to_text_only_and_retry( + self, + payloads: dict, + context_query: list, + chosen_key: str, + available_api_keys: list[str], + func_tool: ToolSet | None, + reason: str, + *, + image_fallback_used: bool = False, + ) -> tuple: + logger.warning( + "检测到图片请求失败(%s),已移除图片并重试(保留文本内容)。", + reason, + ) + new_contexts = await self._remove_image_from_context(context_query) + payloads["messages"] = new_contexts + return ( + False, + chosen_key, + available_api_keys, + payloads, + new_contexts, + func_tool, + image_fallback_used, + ) + + def _create_http_client(self, provider_config: dict) -> httpx.AsyncClient: + """创建带代理的 HTTP 客户端""" + proxy = provider_config.get("proxy", "") + httpx_module: Any = httpx + try: + from openai import _base_client as openai_base_client + + httpx_module = getattr(openai_base_client, "httpx", httpx) + except ImportError: + pass + return create_proxy_client("OpenAI", proxy, httpx_module=httpx_module) + + def __init__(self, provider_config, provider_settings) -> None: + super().__init__(provider_config, provider_settings) + self.chosen_api_key = None + self.api_keys: list = super().get_keys() + self.chosen_api_key = self.api_keys[0] if len(self.api_keys) > 0 else None + self.timeout = provider_config.get("timeout", 120) + self.custom_headers = provider_config.get("custom_headers", {}) + if isinstance(self.timeout, str): + self.timeout = int(self.timeout) + + if not isinstance(self.custom_headers, dict) or not self.custom_headers: + self.custom_headers = None + else: + for key in self.custom_headers: + self.custom_headers[key] = str(self.custom_headers[key]) + + if "api_version" in provider_config: + # Using Azure OpenAI API + self.client = AsyncAzureOpenAI( + api_key=self.chosen_api_key, + api_version=provider_config.get("api_version", None), + default_headers=self.custom_headers, + base_url=provider_config.get("api_base", ""), + timeout=self.timeout, + http_client=self._create_http_client(provider_config), + ) + else: + # Using OpenAI Official API + self.client = AsyncOpenAI( + api_key=self.chosen_api_key, + base_url=provider_config.get("api_base", None), + default_headers=self.custom_headers, + timeout=self.timeout, + http_client=self._create_http_client(provider_config), + ) + + self.default_params = inspect.signature( + self.client.chat.completions.create, + ).parameters.keys() + + model = provider_config.get("model", "unknown") + self.set_model(model) + + # Different upstream/relay channels name the thinking field + # differently (e.g. reasoning_content for DeepSeek/Moonshot, + # reasoning for OpenRouter-style relays), so allow overriding it + # via the provider config. See issue #9783. + self.reasoning_key = ( + provider_config.get("reasoning_key") or "reasoning_content" + ) + + def _ollama_disable_thinking_enabled(self) -> bool: + value = self.provider_config.get("ollama_disable_thinking", False) + if isinstance(value, str): + return value.strip().lower() in {"1", "true", "yes", "on"} + return bool(value) + + def _apply_provider_specific_request_overrides( + self, + payloads: dict[str, Any], + extra_body: dict[str, Any], + ) -> None: + provider = self.provider_config.get("provider") + model = str(payloads.get("model", "")).lower() + + # NVIDIA's hosted MiniMax M3 endpoint can return empty choices when + # max_tokens is omitted (#9206). Scope the compatibility default to + # that model; other NVIDIA models have different token limits. + if ( + provider == "nvidia" + and model == "minimaxai/minimax-m3" + and "max_tokens" not in payloads + and "max_tokens" not in extra_body + ): + payloads["max_tokens"] = 8192 + + if provider != "ollama": + return + if not self._ollama_disable_thinking_enabled(): + return + + # Ollama's OpenAI-compatible endpoint reliably maps reasoning_effort=none + # to think=false, while direct think=false passthrough is not stable. + extra_body.pop("reasoning", None) + extra_body.pop("think", None) + extra_body["reasoning_effort"] = "none" + + async def get_models(self): + try: + models_str = [] + models = await retry_provider_request( + "OpenAI", + lambda: self.client.models.list(), + ) + models = sorted(models.data, key=lambda x: x.id) + for model in models: + models_str.append(model.id) + return models_str + except NotFoundError as e: + raise Exception(f"获取模型列表失败:{e}") + + @staticmethod + def _sanitize_assistant_messages( + payloads: dict, reasoning_key: str = "reasoning_content" + ) -> None: + """在请求发送前过滤/规范化空的 assistant 消息。 + + 严格 API(Moonshot、DeepSeek Reasoner 等)会在 assistant 消息同时缺少 + ``content`` 和 ``tool_calls`` 时返回 400。把 ``""`` / ``None`` / ``[]`` + 都视作空内容:无 tool_calls 时整条过滤掉;有 tool_calls 时将 content + 设为 ``None`` 以符合 OpenAI 规范。就地修改 ``payloads["messages"]``。 + + ``reasoning_key`` 是思考历史的字段名(可经 provider 配置 ``reasoning_key`` + 覆盖,issue #9783);同时兼容默认 ``reasoning_content`` 的历史存量消息。 + """ + messages = payloads.get("messages") + if not isinstance(messages, list): + return + + def _is_empty(content: Any) -> bool: + return content is None or content == "" or content == [] + + cleaned: list[Any] = [] + for idx, msg in enumerate(messages): + if not isinstance(msg, dict) or msg.get("role") != "assistant": + cleaned.append(msg) + continue + + content = msg.get("content") + tool_calls = msg.get("tool_calls") + # Follow the configured reasoning key (#9783); fall back to the + # default key so history saved by older versions is not dropped. + reasoning_content = msg.get(reasoning_key) or msg.get( + "reasoning_content" + ) + + if _is_empty(content) and not tool_calls: + if not reasoning_content: + # 三者全空,真正的垃圾消息,丢弃 + logger.debug( + f"过滤第 {idx} 条空 assistant 消息 (无 content | tool_calls | reasoning_content)" + ) + continue + else: + # ⭐ 有 reasoning_content 但没有 content 和 tool_calls + # 不能丢(推理模型需要 reasoning 历史) + # 但 API 要求 content 或 tool_calls 至少有一个 + # → 设空字符串占位,满足校验 + msg["content"] = "" + + elif _is_empty(content) and tool_calls: + msg["content"] = None # 有 tool_calls,按 OpenAI 规范 + + cleaned.append(msg) + + # Drop orphaned or duplicate tool messages whose assistant(tool_calls) + # was removed by context truncation / compression. + pending_tool_call_ids: set[str] = set() + final: list = [] + removed_tool_messages = 0 + for msg in cleaned: + if not isinstance(msg, dict): + final.append(msg) + pending_tool_call_ids = set() + continue + role = msg.get("role") + if role == "assistant" and msg.get("tool_calls"): + pending_tool_call_ids = { + tc["id"] + for tc in msg["tool_calls"] + if isinstance(tc, dict) and "id" in tc + } + final.append(msg) + elif role == "tool": + tool_call_id = msg.get("tool_call_id") + if tool_call_id in pending_tool_call_ids: + final.append(msg) + pending_tool_call_ids.remove(tool_call_id) + else: + removed_tool_messages += 1 + else: + pending_tool_call_ids = set() + final.append(msg) + if removed_tool_messages: + logger.debug( + "Filtered %d orphaned or duplicate tool message(s)", + removed_tool_messages, + ) + payloads["messages"] = final + + async def _query( + self, + payloads: dict, + tools: ToolSet | None, + *, + request_max_retries: int | None = None, + ) -> LLMResponse: + if tools: + model = payloads.get("model", "").lower() + omit_empty_param_field = "gemini" in model + tool_list = tools.get_func_desc_openai_style( + omit_empty_parameter_field=omit_empty_param_field, + ) + if tool_list: + payloads["tools"] = tool_list + payloads["tool_choice"] = payloads.get("tool_choice", "auto") + + # 不在默认参数中的参数放在 extra_body 中 + extra_body = {} + to_del = [] + for key in payloads: + if key not in self.default_params: + extra_body[key] = payloads[key] + to_del.append(key) + for key in to_del: + del payloads[key] + + # 读取并合并 custom_extra_body 配置 + custom_extra_body = self.provider_config.get("custom_extra_body", {}) + if isinstance(custom_extra_body, dict): + extra_body.update(custom_extra_body) + self._apply_provider_specific_request_overrides(payloads, extra_body) + + model = payloads.get("model", "").lower() + + self._sanitize_assistant_messages(payloads, self.reasoning_key) + + completion = await retry_provider_request( + "OpenAI", + lambda: self.client.chat.completions.create( + **payloads, + stream=False, + extra_body=extra_body, + ), + max_attempts=request_max_retries, + ) + + if not isinstance(completion, ChatCompletion): + raise Exception( + f"API 返回的 completion 类型错误:{type(completion)}: {completion}。", + ) + + logger.debug(f"completion: {completion}") + + llm_response = await self._parse_openai_completion(completion, tools) + + return llm_response + + async def _query_stream( + self, + payloads: dict, + tools: ToolSet | None, + *, + request_max_retries: int | None = None, + ) -> AsyncGenerator[LLMResponse, None]: + """流式查询API,逐步返回结果""" + if tools: + model = payloads.get("model", "").lower() + omit_empty_param_field = "gemini" in model + tool_list = tools.get_func_desc_openai_style( + omit_empty_parameter_field=omit_empty_param_field, + ) + if tool_list: + payloads["tools"] = tool_list + payloads["tool_choice"] = payloads.get("tool_choice", "auto") + + # 不在默认参数中的参数放在 extra_body 中 + extra_body = {} + + # 读取并合并 custom_extra_body 配置 + custom_extra_body = self.provider_config.get("custom_extra_body", {}) + if isinstance(custom_extra_body, dict): + extra_body.update(custom_extra_body) + + to_del = [] + for key in payloads: + if key not in self.default_params: + extra_body[key] = payloads[key] + to_del.append(key) + for key in to_del: + del payloads[key] + self._apply_provider_specific_request_overrides(payloads, extra_body) + + self._sanitize_assistant_messages(payloads, self.reasoning_key) + + stream = await retry_provider_request( + "OpenAI", + lambda: self.client.chat.completions.create( + **payloads, + stream=True, + extra_body=extra_body, + stream_options={"include_usage": True}, + ), + max_attempts=request_max_retries, + ) + + llm_response = LLMResponse("assistant", is_chunk=True) + + state = ChatCompletionStreamState() + + async for chunk in stream: + choice = chunk.choices[0] if chunk.choices else None + delta = choice.delta if choice else None + + if delta and (dtcs := delta.tool_calls): + for idx, tc in enumerate(dtcs): + # siliconflow workaround + if tc.function and tc.function.arguments: + tc.type = "function" + # Fix for #6661: Add missing 'index' field to tool_call deltas + # Gemini and some OpenAI-compatible proxies omit this field + if not hasattr(tc, "index") or tc.index is None: + tc.index = idx + # 跳过 delta=None 的 chunk,避免 SDK 内部 _convert_initial_chunk_into_snapshot + # 第 747 行 choice.delta.to_dict() 抛出 NoneType 错误。 + # refs: AstrBot#6689 / openai-python#5069 / #5047 + # 例外:流末尾的 usage chunk(choices=[],delta=None 但有 usage 数据) + # 需要传给 state,否则最终 completion 会丢失 usage 信息 + if delta is not None or chunk.usage: + try: + state.handle_chunk(chunk) + except Exception as e: + logger.error("Saving chunk state error: " + str(e)) + # logger.debug(f"chunk delta: {delta}") + # handle the content delta + reasoning = self._extract_reasoning_content(chunk) + _y = False + llm_response.id = chunk.id + llm_response.reasoning_content = None + llm_response.completion_text = "" + if reasoning is not None: + llm_response.reasoning_content = reasoning + _y = True + if delta and delta.content: + # Don't strip streaming chunks to preserve spaces between words + completion_text = self._normalize_content(delta.content, strip=False) + llm_response.result_chain = MessageChain( + chain=[Comp.Plain(completion_text)], + ) + _y = True + if chunk.usage: + llm_response.usage = self._extract_usage(chunk.usage) + elif choice and (choice_usage := getattr(choice, "usage", None)): + # Workaround for some providers that only return usage in choices[].usage, e.g. MoonshotAI + # See https://github.com/AstrBotDevs/AstrBot/issues/6614 + llm_response.usage = self._extract_usage(choice_usage) + state.current_completion_snapshot.usage = choice_usage + if _y: + yield llm_response + + try: + final_completion = state.get_final_completion() + llm_response = await self._parse_openai_completion(final_completion, tools) + yield llm_response + except Exception as e: + logger.error("get_final_completion error: " + str(e)) + # 流式内容已通过 yield 发出,记录错误后正常结束即可 + return + + def _extract_reasoning_content( + self, + completion: ChatCompletion | ChatCompletionChunk, + ) -> str | None: + """Extract reasoning content from OpenAI ChatCompletion if available.""" + + def _get_reasoning_attr(obj: Any) -> str | None: + fields_set = getattr(obj, "model_fields_set", None) + if isinstance(fields_set, set) and self.reasoning_key in fields_set: + attr = getattr(obj, self.reasoning_key, "") + return "" if attr is None else str(attr) + attr = getattr(obj, self.reasoning_key, None) + return None if attr is None else str(attr) + + if not completion.choices: + return None + if isinstance(completion, ChatCompletion): + choice = completion.choices[0] + reasoning_attr = _get_reasoning_attr(choice.message) + elif isinstance(completion, ChatCompletionChunk): + delta = completion.choices[0].delta + reasoning_attr = _get_reasoning_attr(delta) + else: + return None + return reasoning_attr + + def _extract_usage(self, usage: CompletionUsage | dict) -> TokenUsage: + ptd = getattr(usage, "prompt_tokens_details", None) + cached = getattr(ptd, "cached_tokens", 0) if ptd else 0 + cached = ( + cached if isinstance(cached, int) else 0 + ) # ptd.cached_tokens 可能为None + prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0 # 安全 + completion_tokens = getattr(usage, "completion_tokens", 0) or 0 + cached = cached or 0 + prompt_tokens = prompt_tokens or 0 + completion_tokens = completion_tokens or 0 + return TokenUsage( + input_other=prompt_tokens - cached, + input_cached=cached, + output=completion_tokens, + ) + + @staticmethod + def _normalize_content(raw_content: Any, strip: bool = True) -> str: + """Normalize content from various formats to plain string. + + Some LLM providers return content as list[dict] format + like [{'type': 'text', 'text': '...'}] instead of + plain string. This method handles both formats. + + Args: + raw_content: The raw content from LLM response, can be str, list, dict, or other. + strip: Whether to strip whitespace from the result. Set to False for + streaming chunks to preserve spaces between words. + + Returns: + Normalized plain text string. + """ + # Handle dict format (e.g., {"type": "text", "text": "..."}) + if isinstance(raw_content, dict): + if "text" in raw_content: + text_val = raw_content.get("text", "") + return str(text_val) if text_val is not None else "" + # For other dict formats, return empty string and log + logger.warning(f"Unexpected dict format content: {raw_content}") + return "" + + if isinstance(raw_content, list): + # Check if this looks like OpenAI content-part format + # Only process if at least one item has {'type': 'text', 'text': ...} structure + has_content_part = any( + isinstance(part, dict) and part.get("type") == "text" + for part in raw_content + ) + if has_content_part: + text_parts = [] + for part in raw_content: + if isinstance(part, dict) and part.get("type") == "text": + text_val = part.get("text", "") + # Coerce to str in case text is null or non-string + text_parts.append(str(text_val) if text_val is not None else "") + return "".join(text_parts) + # Not content-part format, return string representation + return str(raw_content) + + if isinstance(raw_content, str): + content = raw_content.strip() if strip else raw_content + # Check if the string is a JSON-encoded list (e.g., "[{'type': 'text', ...}]") + # This can happen when streaming concatenates content that was originally list format + # Only check if it looks like a complete JSON array (requires strip for check) + check_content = raw_content.strip() + if ( + check_content.startswith("[") + and check_content.endswith("]") + and len(check_content) < 8192 + ): + try: + # First try standard JSON parsing + parsed = json.loads(check_content) + except json.JSONDecodeError: + # If that fails, try parsing as Python literal (handles single quotes) + # This is safer than blind replace("'", '"') which corrupts apostrophes + try: + import ast + + parsed = ast.literal_eval(check_content) + except (ValueError, SyntaxError): + parsed = None + + if isinstance(parsed, list): + # Only convert if it matches OpenAI content-part schema + # i.e., at least one item has {'type': 'text', 'text': ...} + has_content_part = any( + isinstance(part, dict) and part.get("type") == "text" + for part in parsed + ) + if has_content_part: + text_parts = [] + for part in parsed: + if isinstance(part, dict) and part.get("type") == "text": + text_val = part.get("text", "") + # Coerce to str in case text is null or non-string + text_parts.append( + str(text_val) if text_val is not None else "" + ) + if text_parts: + return "".join(text_parts) + return content + + # Fallback for other types (int, float, etc.) + return str(raw_content) if raw_content is not None else "" + + async def _parse_openai_completion( + self, completion: ChatCompletion, tools: ToolSet | None + ) -> LLMResponse: + """Parse OpenAI ChatCompletion into LLMResponse""" + llm_response = LLMResponse("assistant") + + # workaround for #9374 + if not completion.choices: + data = getattr(completion, "data", None) + if isinstance(data, dict): + try: + completion = ChatCompletion.model_validate(data) + except (TypeError, ValueError): + pass + + if not completion.choices: + raise EmptyModelOutputError( + f"OpenAI completion has no choices. response_id={completion.id}" + ) + choice = completion.choices[0] + + # parse the text completion + if choice.message.content is not None: + completion_text = self._normalize_content(choice.message.content) + # specially, some providers may set tags around reasoning content in the completion text, + # we use regex to remove them, and store then in reasoning_content field + reasoning_pattern = re.compile(r"(.*?)", re.DOTALL) + matches = reasoning_pattern.findall(completion_text) + if matches: + llm_response.reasoning_content = "\n".join( + [match.strip() for match in matches], + ) + completion_text = reasoning_pattern.sub("", completion_text).strip() + # Also clean up orphan tags that may leak from some models + completion_text = re.sub(r"\s*$", "", completion_text).strip() + llm_response.result_chain = MessageChain().message(completion_text) + elif refusal := getattr(choice.message, "refusal", None): + refusal_text = self._normalize_content(refusal) + if refusal_text: + llm_response.result_chain = MessageChain().message(refusal_text) + + # parse the reasoning content if any + # the priority is higher than the tag extraction + reasoning_content = self._extract_reasoning_content(completion) + if reasoning_content is not None: + llm_response.reasoning_content = reasoning_content + + # parse tool calls if any + if choice.message.tool_calls and tools is not None: + args_ls = [] + func_name_ls = [] + tool_call_ids = [] + tool_call_extra_content_dict = {} + for tool_call in choice.message.tool_calls: + if isinstance(tool_call, str): + # workaround for #1359 + tool_call = json.loads(tool_call) + if tools is None: + # 工具集未提供 + # Should be unreachable + raise Exception("工具集未提供") + + if tool_call.type == "function": + # workaround for #1454 + if isinstance(tool_call.function.arguments, str): + try: + args = json.loads(tool_call.function.arguments) + except json.JSONDecodeError as e: + logger.error(f"解析参数失败: {e}") + args = {} + else: + args = tool_call.function.arguments + # Some API may return None for tools with no parameters + if args is None: + args = {} + args_ls.append(args) + func_name_ls.append(tool_call.function.name) + tool_call_ids.append(tool_call.id) + + # gemini-2.5 / gemini-3 series extra_content handling + extra_content = getattr(tool_call, "extra_content", None) + if extra_content is not None: + tool_call_extra_content_dict[tool_call.id] = extra_content + + llm_response.role = "tool" + llm_response.tools_call_args = args_ls + llm_response.tools_call_name = func_name_ls + llm_response.tools_call_ids = tool_call_ids + llm_response.tools_call_extra_content = tool_call_extra_content_dict + # specially handle finish reason + if choice.finish_reason == "content_filter": + raise Exception( + "API 返回的 completion 由于内容安全过滤被拒绝(非 AstrBot)。", + ) + has_text_output = bool((llm_response.completion_text or "").strip()) + has_reasoning_output = bool((llm_response.reasoning_content or "").strip()) + if ( + not has_text_output + and not has_reasoning_output + and not llm_response.tools_call_args + ): + logger.error(f"OpenAI completion has no usable output: {completion}.") + raise EmptyModelOutputError( + "OpenAI completion has no usable output. " + f"response_id={completion.id}, finish_reason={choice.finish_reason}" + ) + + llm_response.raw_completion = completion + llm_response.id = completion.id + + llm_response.usage = ( + self._extract_usage(completion.usage) if completion.usage else TokenUsage() + ) + + return llm_response + + async def _prepare_chat_payload( + self, + prompt: str | None, + image_urls: list[str] | None = None, + audio_urls: list[str] | None = None, + contexts: list[dict] | list[Message] | None = None, + system_prompt: str | None = None, + tool_calls_result: ToolCallsResult | list[ToolCallsResult] | None = None, + model: str | None = None, + extra_user_content_parts: list[ContentPart] | None = None, + **kwargs, + ) -> tuple: + """准备聊天所需的有效载荷和上下文""" + if contexts is None: + contexts = [] + new_record = None + if prompt is not None: + new_record = await self.assemble_context( + prompt or "", + image_urls, + audio_urls, + extra_user_content_parts, + ) + context_query = copy.deepcopy(self._ensure_message_to_dicts(contexts)) + if new_record: + context_query.append(new_record) + if system_prompt: + context_query.insert(0, {"role": "system", "content": system_prompt}) + + for part in context_query: + if "_no_save" in part: + del part["_no_save"] + + # tool calls result + if tool_calls_result: + if isinstance(tool_calls_result, ToolCallsResult): + context_query.extend(tool_calls_result.to_openai_messages()) + else: + for tcr in tool_calls_result: + context_query.extend(tcr.to_openai_messages()) + + if self._context_contains_image(context_query): + context_query = await self._materialize_context_image_parts(context_query) + + model = model or self.get_model() + + payloads = {"messages": context_query, "model": model} + + self._finally_convert_payload(payloads) + + return payloads, context_query + + def _finally_convert_payload(self, payloads: dict) -> None: + """Finally convert the payload. Such as think part conversion, tool inject.""" + model = payloads.get("model", "").lower() + is_gemini = "gemini" in model + _deepseek_v4_markers = ("deepseek-v4-pro", "deepseek-v4-flash", "deepseek-v4") + is_deepseek_v4_reasoning = ( + any(marker in model for marker in _deepseek_v4_markers) + or "api.deepseek.com" in self.client.base_url.host + ) + # deepseek-chat and deepseek-reasoner now point to V4 models (per official website) + + # MiMo 推理模型(MiMo-V2.5-Pro / MiMo-V2.5 / MiMo-V2-Pro / MiMo-V2-Omni / MiMo-V2-Flash) + # 要求 assistant 历史消息必须回传 reasoning_content,否则返回 400 + mimo_reasoning_models = { + "mimo-v2.5-pro", + "mimo-v2.5", + "mimo-v2-pro", + "mimo-v2-omni", + "mimo-v2-flash", + } + is_mimo_reasoning = model in mimo_reasoning_models + for message in payloads.get("messages", []): + if message.get("role") == "assistant" and isinstance( + message.get("content"), list + ): + reasoning_content = "" + reasoning_content_present = False + new_content = [] # not including think part + for part in message["content"]: + if part.get("type") == "think": + reasoning_content_present = True + reasoning_content += str(part.get("think")) + else: + new_content.append(part) + # Some providers (Grok, etc.) reject empty content lists. + # When all parts were think blocks, fall back to None. + message["content"] = new_content or None + if reasoning_content_present: + # Emit the thinking history under the configured key so + # channels that expect `reasoning` (issue #9783) accept it. + message[self.reasoning_key] = reasoning_content + + if ( + message.get("role") == "assistant" + and is_deepseek_v4_reasoning + and "reasoning_content" not in message + ): + # DeepSeek v4 reasoning models require the field on assistant + # history messages, even when the reasoning content is empty. + message["reasoning_content"] = "" + + if ( + message.get("role") == "assistant" + and is_mimo_reasoning + and "reasoning_content" not in message + ): + # MiMo 推理模型要求 assistant 历史消息回传 reasoning_content, + # 缺失时 API 返回 400。参见 MiMo 官方文档。 + message["reasoning_content"] = "" + + # Gemini 的 function_response 要求 google.protobuf.Struct(即 JSON 对象), + # 纯文本会触发 400 Invalid argument,需要包一层 JSON。 + if is_gemini and message.get("role") == "tool": + content = message.get("content", "") + if isinstance(content, str): + try: + json.loads(content) + except (json.JSONDecodeError, ValueError): + message["content"] = json.dumps( + {"result": content}, ensure_ascii=False + ) + + async def _handle_api_error( + self, + e: Exception, + payloads: dict, + context_query: list, + func_tool: ToolSet | None, + chosen_key: str, + available_api_keys: list[str], + retry_cnt: int, + max_retries: int, + image_fallback_used: bool = False, + ) -> tuple: + """处理API错误并尝试恢复""" + if "429" in str(e): + logger.warning( + f"API 调用过于频繁,尝试使用其他 Key 重试。当前 Key: {chosen_key[:12]}", + ) + # 最后一次不等待 + if retry_cnt < max_retries - 1: + await asyncio.sleep(1) + if chosen_key in available_api_keys: + available_api_keys.remove(chosen_key) + if len(available_api_keys) > 0: + chosen_key = random.choice(available_api_keys) + return ( + False, + chosen_key, + available_api_keys, + payloads, + context_query, + func_tool, + image_fallback_used, + ) + raise e + if "maximum context length" in str(e) or "context length" in str(e).lower(): + logger.warning( + f"上下文长度超过限制。尝试弹出最早的记录然后重试。当前记录条数: {len(context_query)}", + ) + await self.pop_record(context_query) + payloads["messages"] = context_query + return ( + False, + chosen_key, + available_api_keys, + payloads, + context_query, + func_tool, + image_fallback_used, + ) + if "The model is not a VLM" in str(e): # siliconcloud + if image_fallback_used or not self._context_contains_image(context_query): + raise e + # 尝试删除所有 image + return await self._fallback_to_text_only_and_retry( + payloads, + context_query, + chosen_key, + available_api_keys, + func_tool, + "model_not_vlm", + image_fallback_used=True, + ) + if self._is_content_moderated_upload_error(e): + if image_fallback_used or not self._context_contains_image(context_query): + raise e + return await self._fallback_to_text_only_and_retry( + payloads, + context_query, + chosen_key, + available_api_keys, + func_tool, + "image_content_moderated", + image_fallback_used=True, + ) + if self._is_invalid_attachment_error(e): + if image_fallback_used or not self._context_contains_image(context_query): + raise e + return await self._fallback_to_text_only_and_retry( + payloads, + context_query, + chosen_key, + available_api_keys, + func_tool, + "invalid_attachment", + image_fallback_used=True, + ) + + if ( + "Function calling is not enabled" in str(e) + or ("tool" in str(e).lower() and "support" in str(e).lower()) + or ("function" in str(e).lower() and "support" in str(e).lower()) + ): + # openai, ollama, gemini openai, siliconcloud 的错误提示与 code 不统一,只能通过字符串匹配 + logger.warning( + f"{self.get_model()} 不支持函数工具调用,已自动去除,不影响使用。如需永久关闭,可前往 WebUI 中关闭工具调用。", + ) + payloads.pop("tools", None) + return ( + False, + chosen_key, + available_api_keys, + payloads, + context_query, + None, + image_fallback_used, + ) + # logger.error(f"发生了错误。Provider 配置如下: {self.provider_config}") + + if is_connection_error(e): + proxy = self.provider_config.get("proxy", "") + log_connection_failure("OpenAI", e, proxy) + + raise e + + async def text_chat( + self, + prompt=None, + session_id=None, + image_urls=None, + audio_urls=None, + func_tool=None, + contexts=None, + system_prompt=None, + tool_calls_result=None, + model=None, + extra_user_content_parts=None, + tool_choice: Literal["auto", "required"] = "auto", + request_max_retries: int | None = None, + **kwargs, + ) -> LLMResponse: + payloads, context_query = await self._prepare_chat_payload( + prompt, + image_urls, + audio_urls, + contexts, + system_prompt, + tool_calls_result, + model=model, + extra_user_content_parts=extra_user_content_parts, + **kwargs, + ) + if func_tool and not func_tool.empty(): + payloads["tool_choice"] = tool_choice + + llm_response = None + max_retries = 10 + available_api_keys = self.api_keys.copy() + chosen_key = random.choice(available_api_keys) + image_fallback_used = False + + last_exception = None + retry_cnt = 0 + for retry_cnt in range(max_retries): + try: + self.client.api_key = chosen_key + llm_response = await self._query( + payloads, + func_tool, + request_max_retries=request_max_retries, + ) + break + except Exception as e: + last_exception = e + ( + success, + chosen_key, + available_api_keys, + payloads, + context_query, + func_tool, + image_fallback_used, + ) = await self._handle_api_error( + e, + payloads, + context_query, + func_tool, + chosen_key, + available_api_keys, + retry_cnt, + max_retries, + image_fallback_used=image_fallback_used, + ) + if success: + break + + if retry_cnt == max_retries - 1 or llm_response is None: + logger.error(f"API 调用失败,重试 {max_retries} 次仍然失败。") + if last_exception is None: + raise Exception("未知错误") + raise last_exception + return llm_response + + async def text_chat_stream( + self, + prompt=None, + session_id=None, + image_urls=None, + audio_urls=None, + func_tool=None, + contexts=None, + system_prompt=None, + tool_calls_result=None, + model=None, + tool_choice: Literal["auto", "required"] = "auto", + request_max_retries: int | None = None, + **kwargs, + ) -> AsyncGenerator[LLMResponse, None]: + """流式对话,与服务商交互并逐步返回结果""" + payloads, context_query = await self._prepare_chat_payload( + prompt, + image_urls, + audio_urls, + contexts, + system_prompt, + tool_calls_result, + model=model, + **kwargs, + ) + if func_tool and not func_tool.empty(): + payloads["tool_choice"] = tool_choice + + max_retries = 10 + available_api_keys = self.api_keys.copy() + chosen_key = random.choice(available_api_keys) + image_fallback_used = False + + last_exception = None + retry_cnt = 0 + for retry_cnt in range(max_retries): + try: + self.client.api_key = chosen_key + async for response in self._query_stream( + payloads, + func_tool, + request_max_retries=request_max_retries, + ): + yield response + break + except Exception as e: + last_exception = e + ( + success, + chosen_key, + available_api_keys, + payloads, + context_query, + func_tool, + image_fallback_used, + ) = await self._handle_api_error( + e, + payloads, + context_query, + func_tool, + chosen_key, + available_api_keys, + retry_cnt, + max_retries, + image_fallback_used=image_fallback_used, + ) + if success: + break + + if retry_cnt == max_retries - 1: + logger.error(f"API 调用失败,重试 {max_retries} 次仍然失败。") + if last_exception is None: + raise Exception("未知错误") + raise last_exception + + async def _remove_image_from_context(self, contexts: list): + """从上下文中删除所有带有 image 的记录""" + new_contexts = [] + + for context in contexts: + if "content" in context and isinstance(context["content"], list): + # continue + new_content = [] + for item in context["content"]: + if isinstance(item, dict) and "image_url" in item: + continue + new_content.append(item) + if not new_content: + # 用户只发了图片 + new_content = [{"type": "text", "text": "[图片]"}] + context["content"] = new_content + new_contexts.append(context) + return new_contexts + + def get_current_key(self) -> str: + return self.client.api_key + + def get_keys(self) -> list[str]: + return self.api_keys + + def set_key(self, key) -> None: + self.client.api_key = key + + async def assemble_context( + self, + text: str, + image_urls: list[str] | None = None, + audio_urls: list[str] | None = None, + extra_user_content_parts: list[ContentPart] | None = None, + ) -> dict: + """组装成符合 OpenAI 格式的 role 为 user 的消息段""" + + # 构建内容块列表 + content_blocks = [] + + # 1. 用户原始发言(OpenAI 建议:用户发言在前) + if text: + content_blocks.append({"type": "text", "text": text}) + elif image_urls: + # 如果没有文本但有图片,添加占位文本 + content_blocks.append({"type": "text", "text": "[Image]"}) + elif audio_urls: + content_blocks.append({"type": "text", "text": "[Audio]"}) + elif extra_user_content_parts: + # 如果只有额外内容块,也需要添加占位文本 + content_blocks.append({"type": "text", "text": " "}) + + # 2. 额外的内容块(系统提醒、指令等) + if extra_user_content_parts: + for part in extra_user_content_parts: + if isinstance(part, TextPart): + content_blocks.append({"type": "text", "text": part.text}) + elif isinstance(part, ImageURLPart): + image_part = await self._resolve_image_part( + part.image_url.url, + ) + if image_part: + content_blocks.append(image_part) + elif isinstance(part, AudioURLPart): + audio_part = await self._resolve_audio_part(part.audio_url.url) + if audio_part: + content_blocks.append(audio_part) + else: + raise ValueError(f"不支持的额外内容块类型: {type(part)}") + + # 3. 图片内容 + if image_urls: + for image_url in image_urls: + image_part = await self._resolve_image_part(image_url) + if image_part: + content_blocks.append(image_part) + + if audio_urls: + for audio_path in audio_urls: + audio_part = await self._resolve_audio_part(audio_path) + if audio_part: + content_blocks.append(audio_part) + + # 如果只有主文本且没有额外内容块和图片,返回简单格式以保持向后兼容 + if ( + text + and not extra_user_content_parts + and not image_urls + and not audio_urls + and len(content_blocks) == 1 + and content_blocks[0]["type"] == "text" + ): + return {"role": "user", "content": content_blocks[0]["text"]} + + # 否则返回多模态格式 + return {"role": "user", "content": content_blocks} + + async def encode_image_bs64(self, image_url: str) -> str: + """将图片转换为 base64""" + image_data = await self._image_ref_to_data_url(image_url, mode="strict") + if image_data is None: + raise RuntimeError( + f"Failed to encode image data: {describe_media_ref(image_url)}" + ) + return image_data + + async def terminate(self): + if self.client: + await self.client.close() diff --git a/tests/test_openai_source.py b/tests/test_openai_source.py index 911b76131f..06fe6f4ffe 100644 --- a/tests/test_openai_source.py +++ b/tests/test_openai_source.py @@ -1,2203 +1,2328 @@ -import base64 -import builtins -from io import BytesIO -from types import SimpleNamespace - -import httpx -import pytest -from openai.types.chat.chat_completion import ChatCompletion -from openai.types.chat.chat_completion_chunk import ChatCompletionChunk -from PIL import Image as PILImage - -import astrbot.core.provider.sources.openai_source as openai_source_module -import astrbot.core.provider.sources.request_retry as request_retry -from astrbot.core.exceptions import EmptyModelOutputError -from astrbot.core.provider.entities import LLMResponse -from astrbot.core.provider.sources.groq_source import ProviderGroq -from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial -from astrbot.core.utils.media_utils import ResolvedMediaData, file_uri_to_path - - -class _ErrorWithBody(Exception): - def __init__(self, message: str, body: dict): - super().__init__(message) - self.body = body - - -class _ErrorWithResponse(Exception): - def __init__(self, message: str, response_text: str): - super().__init__(message) - self.response = SimpleNamespace(text=response_text) - - -def _make_provider(overrides: dict | None = None) -> ProviderOpenAIOfficial: - provider_config = { - "id": "test-openai", - "type": "openai_chat_completion", - "model": "gpt-4o-mini", - "key": ["test-key"], - } - if overrides: - provider_config.update(overrides) - return ProviderOpenAIOfficial( - provider_config=provider_config, - provider_settings={}, - ) - - -def _make_groq_provider(overrides: dict | None = None) -> ProviderGroq: - provider_config = { - "id": "test-groq", - "type": "groq_chat_completion", - "model": "qwen/qwen3-32b", - "key": ["test-key"], - } - if overrides: - provider_config.update(overrides) - return ProviderGroq( - provider_config=provider_config, - provider_settings={}, - ) - - -def test_create_http_client_uses_openai_httpx_module(monkeypatch): - captured: dict[str, object] = {} - fake_httpx_module = object() - - from openai import _base_client as openai_base_client - - monkeypatch.setattr( - openai_base_client, - "httpx", - fake_httpx_module, - raising=False, - ) - - def fake_create_proxy_client( - provider_label: str, - proxy: str | None = None, - headers: dict[str, str] | None = None, - verify=None, - httpx_module=None, - ): - captured["httpx_module"] = httpx_module - return object() - - monkeypatch.setattr( - openai_source_module, - "create_proxy_client", - fake_create_proxy_client, - ) - - provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) - provider._create_http_client({"proxy": ""}) - - assert captured["httpx_module"] is fake_httpx_module - - -def test_create_http_client_falls_back_to_global_httpx_module(monkeypatch): - captured: dict[str, object] = {} - - def fake_create_proxy_client( - provider_label: str, - proxy: str | None = None, - headers: dict[str, str] | None = None, - verify=None, - httpx_module=None, - ): - captured["httpx_module"] = httpx_module - return object() - - real_import = builtins.__import__ - - def fake_import(name, globals=None, locals=None, fromlist=(), level=0): - if name == "openai" and fromlist: - raise ImportError("missing openai._base_client") - return real_import(name, globals, locals, fromlist, level) - - monkeypatch.setattr( - openai_source_module, - "create_proxy_client", - fake_create_proxy_client, - ) - monkeypatch.setattr(builtins, "__import__", fake_import) - - provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) - provider._create_http_client({"proxy": ""}) - - assert captured["httpx_module"] is openai_source_module.httpx - - -@pytest.mark.asyncio -async def test_get_models_retries_transient_request_error(monkeypatch): - monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0) - monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0) - - class FakeModels: - def __init__(self): - self.calls = 0 - - async def list(self): - self.calls += 1 - if self.calls == 1: - raise httpx.ConnectError("temporary connection failure") - return SimpleNamespace( - data=[ - SimpleNamespace(id="gpt-b"), - SimpleNamespace(id="gpt-a"), - ] - ) - - models = FakeModels() - provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) - provider.client = SimpleNamespace(models=models) - - assert await provider.get_models() == ["gpt-a", "gpt-b"] - assert models.calls == 2 - - -@pytest.mark.asyncio -async def test_text_chat_passes_request_max_retries_to_query(): - captured: dict[str, object] = {} - - provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) - provider.api_keys = ["test-key"] - provider.client = SimpleNamespace(api_key=None) - - async def fake_prepare_chat_payload(*args, **kwargs): - return {"messages": [], "model": "gpt-4o-mini"}, [] - - async def fake_query(payloads, func_tool, *, request_max_retries=None): - captured["request_max_retries"] = request_max_retries - return LLMResponse(role="assistant", completion_text="ok") - - provider._prepare_chat_payload = fake_prepare_chat_payload - provider._query = fake_query - - await provider.text_chat(prompt="hello", request_max_retries=2) - - assert captured["request_max_retries"] == 2 - - -@pytest.mark.asyncio -async def test_handle_api_error_content_moderated_removes_images(): - provider = _make_provider( - {"image_moderation_error_patterns": ["file:content-moderated"]} - ) - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - - success, *_rest = await provider._handle_api_error( - Exception("Content is moderated [WKE=file:content-moderated]"), - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - - assert success is False - updated_context = payloads["messages"] - assert isinstance(updated_context, list) - assert updated_context[0]["content"] == [{"type": "text", "text": "hello"}] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_model_not_vlm_removes_images_and_retries_text_only(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - - success, *_rest = await provider._handle_api_error( - Exception("The model is not a VLM and cannot process images"), - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - - assert success is False - updated_context = payloads["messages"] - assert isinstance(updated_context, list) - assert updated_context[0]["content"] == [{"type": "text", "text": "hello"}] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_model_not_vlm_after_fallback_raises(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - - with pytest.raises(Exception, match="not a VLM"): - await provider._handle_api_error( - Exception("The model is not a VLM and cannot process images"), - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=1, - max_retries=10, - image_fallback_used=True, - ) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_content_moderated_with_unserializable_body(): - provider = _make_provider({"image_moderation_error_patterns": ["blocked"]}) - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - err = _ErrorWithBody( - "upstream error", - {"error": {"message": "blocked"}, "raw": object()}, - ) - - success, *_rest = await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - assert success is False - assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] - finally: - await provider.terminate() - - -def test_extract_error_text_candidates_truncates_long_response_text(): - long_text = "x" * 20000 - err = _ErrorWithResponse("upstream error", long_text) - candidates = ProviderOpenAIOfficial._extract_error_text_candidates(err) - assert candidates - assert max(len(candidate) for candidate in candidates) <= ( - ProviderOpenAIOfficial._ERROR_TEXT_CANDIDATE_MAX_CHARS - ) - - -@pytest.mark.asyncio -async def test_openai_payload_keeps_reasoning_content_in_assistant_history(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "assistant", - "content": [ - {"type": "think", "think": "step 1"}, - {"type": "text", "text": "final answer"}, - ], - } - ] - } - - provider._finally_convert_payload(payloads) - - assistant_message = payloads["messages"][0] - assert assistant_message["content"] == [ - {"type": "text", "text": "final answer"} - ] - assert assistant_message["reasoning_content"] == "step 1" - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_groq_payload_drops_reasoning_content_from_assistant_history(): - provider = _make_groq_provider() - try: - payloads = { - "messages": [ - { - "role": "assistant", - "content": [ - {"type": "think", "think": "step 1"}, - {"type": "text", "text": "final answer"}, - ], - } - ] - } - - provider._finally_convert_payload(payloads) - - assistant_message = payloads["messages"][0] - assert assistant_message["content"] == [ - {"type": "text", "text": "final answer"} - ] - assert "reasoning_content" not in assistant_message - assert "reasoning" not in assistant_message - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_content_moderated_without_images_raises(): - provider = _make_provider( - {"image_moderation_error_patterns": ["file:content-moderated"]} - ) - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [{"type": "text", "text": "hello"}], - } - ] - } - context_query = payloads["messages"] - err = Exception("Content is moderated [WKE=file:content-moderated]") - - with pytest.raises(Exception, match="content-moderated"): - await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_content_moderated_detects_structured_body(): - provider = _make_provider( - {"image_moderation_error_patterns": ["content_moderated"]} - ) - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - err = _ErrorWithBody( - "upstream error", - {"error": {"code": "content_moderated", "message": "blocked"}}, - ) - - success, *_rest = await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - assert success is False - assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_content_moderated_supports_custom_patterns(): - provider = _make_provider( - {"image_moderation_error_patterns": ["blocked_by_policy_code_123"]} - ) - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - err = Exception("upstream: blocked_by_policy_code_123") - - success, *_rest = await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - assert success is False - assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_content_moderated_without_patterns_raises(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - err = Exception("Content is moderated [WKE=file:content-moderated]") - - with pytest.raises(Exception, match="content-moderated"): - await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_unknown_image_error_raises(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - - with pytest.raises(Exception, match="unknown provider image upload error"): - await provider._handle_api_error( - Exception("some unknown provider image upload error"), - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_invalid_attachment_removes_images_and_retries_text_only(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - err = _ErrorWithBody( - "upstream error", - { - "error": { - "code": "INVALID_ATTACHMENT", - "message": "download attachment: unexpected status 404", - } - }, - ) - - success, *_rest = await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - - assert success is False - assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_invalid_attachment_without_images_raises(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [{"type": "text", "text": "hello"}], - } - ] - } - context_query = payloads["messages"] - err = _ErrorWithBody( - "upstream error", - { - "error": { - "code": "INVALID_ATTACHMENT", - "message": "download attachment: unexpected status 404", - } - }, - ) - - with pytest.raises(_ErrorWithBody, match="upstream error"): - await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=0, - max_retries=10, - ) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_handle_api_error_invalid_attachment_after_fallback_raises(): - provider = _make_provider() - try: - payloads = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "hello"}, - { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - }, - ], - } - ] - } - context_query = payloads["messages"] - err = _ErrorWithBody( - "upstream error", - { - "error": { - "code": "INVALID_ATTACHMENT", - "message": "download attachment: unexpected status 404", - } - }, - ) - - with pytest.raises(_ErrorWithBody, match="upstream error"): - await provider._handle_api_error( - err, - payloads=payloads, - context_query=context_query, - func_tool=None, - chosen_key="test-key", - available_api_keys=["test-key"], - retry_cnt=1, - max_retries=10, - image_fallback_used=True, - ) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_materializes_context_http_image_urls(monkeypatch): - provider = _make_provider() - try: - - async def fake_resolve_media_ref_to_base64_data( - media_ref: str, - *, - media_type: str, - strict: bool = False, - ) -> ResolvedMediaData: - assert media_ref == "https://example.com/quoted.png" - assert media_type == "image" - assert strict is False - return ResolvedMediaData(base64_data="abcd", mime_type="image/png") - - monkeypatch.setattr( - openai_source_module, - "resolve_media_ref_to_base64_data", - fake_resolve_media_ref_to_base64_data, - ) - - contexts = [ - { - "role": "user", - "metadata": {"source": "quoted"}, - "content": [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/quoted.png", - "id": "ctx-img", - "detail": "high", - }, - }, - ], - } - ] - - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=contexts, - ) - - assert payloads["messages"][0]["content"] == [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,abcd", - "detail": "high", - }, - }, - ] - assert payloads["messages"][0]["content"][1]["image_url"].get("id") is None - assert contexts[0]["content"][1]["image_url"] == { - "url": "https://example.com/quoted.png", - "id": "ctx-img", - "detail": "high", - } - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_skips_materialization_for_text_only_context( - monkeypatch, -): - provider = _make_provider() - try: - - async def fail_if_called(_context_query): - raise AssertionError("materialization should be skipped") - - monkeypatch.setattr( - provider, "_materialize_context_image_parts", fail_if_called - ) - - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=[{"role": "user", "content": "hello"}], - ) - - assert payloads["messages"] == [{"role": "user", "content": "hello"}] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_skips_materialization_for_text_only_parts( - monkeypatch, -): - provider = _make_provider() - try: - - async def fail_if_called(_context_query): - raise AssertionError("materialization should be skipped") - - monkeypatch.setattr( - provider, "_materialize_context_image_parts", fail_if_called - ) - - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=[ - { - "role": "user", - "content": [{"type": "text", "text": "hello"}], - } - ], - ) - - assert payloads["messages"] == [ - { - "role": "user", - "content": [{"type": "text", "text": "hello"}], - } - ] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_materializes_context_http_image_urls_with_detected_mime( - monkeypatch, tmp_path -): - provider = _make_provider() - try: - image_path = tmp_path / "quoted-image.png" - PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) - - async def fake_download(url: str, target_path: str) -> None: - assert url == "https://example.com/quoted.png" - with open(target_path, "wb") as f: - f.write(image_path.read_bytes()) - - monkeypatch.setattr( - "astrbot.core.utils.media_utils.download_file", - fake_download, - ) - - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=[ - { - "role": "user", - "content": [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/quoted.png", - }, - }, - ], - } - ], - ) - - image_payload = payloads["messages"][0]["content"][1]["image_url"] - assert image_payload["url"].startswith("data:image/png;base64,") - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_materializes_context_file_uri_image_urls(tmp_path): - provider = _make_provider() - try: - image_path = tmp_path / "quoted-image.png" - PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) - - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=[ - { - "role": "user", - "content": [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": image_path.as_uri(), - }, - }, - ], - } - ], - ) - - image_payload = payloads["messages"][0]["content"][1]["image_url"] - assert image_payload["url"].startswith("data:image/png;base64,") - finally: - await provider.terminate() - - -def test_file_uri_to_path_preserves_windows_drive_letter(): - assert file_uri_to_path("file:///C:/tmp/quoted-image.png") == ( - "C:/tmp/quoted-image.png" - ) - - -def test_file_uri_to_path_preserves_windows_netloc_drive_letter(): - assert file_uri_to_path("file://C:/tmp/quoted-image.png") == ( - "C:/tmp/quoted-image.png" - ) - - -def test_file_uri_to_path_preserves_remote_netloc_as_unc_path(): - assert file_uri_to_path("file://server/share/quoted-image.png") == ( - "//server/share/quoted-image.png" - ) - - -@pytest.mark.asyncio -async def test_resolve_image_part_rejects_invalid_local_file(tmp_path): - provider = _make_provider() - try: - invalid_file = tmp_path / "not-image.txt" - invalid_file.write_text("not an image") - - assert await provider._resolve_image_part(str(invalid_file)) is None - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_resolve_image_part_rejects_invalid_file_uri(tmp_path): - provider = _make_provider() - try: - invalid_file = tmp_path / "not-image.txt" - invalid_file.write_text("not an image") - - assert await provider._resolve_image_part(invalid_file.as_uri()) is None - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_image_ref_to_data_url_mode_controls_invalid_file_behavior(tmp_path): - provider = _make_provider() - try: - invalid_file = tmp_path / "not-image.txt" - invalid_file.write_text("not an image") - - assert ( - await provider._image_ref_to_data_url(str(invalid_file), mode="safe") - is None - ) - with pytest.raises(ValueError, match="Invalid image file"): - await provider._image_ref_to_data_url(str(invalid_file), mode="strict") - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_materialize_context_image_parts_returns_new_messages(monkeypatch): - provider = _make_provider() - try: - context_query = [ - { - "role": "user", - "metadata": {"source": "quoted"}, - "content": [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/quoted.png", - "detail": "high", - }, - }, - ], - }, - {"role": "assistant", "content": "plain text"}, - ] - - async def fake_resolve(image_url: str, *, image_detail: str | None = None): - assert image_url == "https://example.com/quoted.png" - assert image_detail == "high" - return { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,abcd", - "detail": "high", - }, - } - - monkeypatch.setattr(provider, "_resolve_image_part", fake_resolve) - - materialized = await provider._materialize_context_image_parts(context_query) - - assert materialized is not context_query - assert materialized[0] is not context_query[0] - assert materialized[0]["metadata"] is context_query[0]["metadata"] - assert materialized[0]["content"][0] is context_query[0]["content"][0] - assert ( - materialized[0]["content"][1]["image_url"]["url"] - == "data:image/png;base64,abcd" - ) - assert ( - context_query[0]["content"][1]["image_url"]["url"] - == "https://example.com/quoted.png" - ) - assert materialized[1] is not context_query[1] - assert materialized[1]["content"] == "plain text" - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_encode_image_bs64_missing_file_raises(tmp_path): - provider = _make_provider() - try: - missing_path = tmp_path / "missing-image.png" - with pytest.raises(FileNotFoundError): - await provider.encode_image_bs64(str(missing_path)) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_encode_image_bs64_invalid_file_raises(tmp_path): - provider = _make_provider() - try: - invalid_file = tmp_path / "not-image.txt" - invalid_file.write_text("not an image") - - with pytest.raises(ValueError, match="Invalid image file"): - await provider.encode_image_bs64(str(invalid_file)) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_encode_image_bs64_supports_base64_scheme(): - provider = _make_provider() - try: - image_data = await provider.encode_image_bs64("base64://abcd") - - assert image_data == "data:image/jpeg;base64,abcd" - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_encode_image_bs64_supports_file_uri(tmp_path): - provider = _make_provider() - try: - image_path = tmp_path / "quoted-image.png" - PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) - - image_data = await provider.encode_image_bs64(image_path.as_uri()) - - assert image_data.startswith("data:image/png;base64,") - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_resolve_image_part_supports_base64_scheme(): - provider = _make_provider() - try: - assert await provider._resolve_image_part("base64://abcd") == { - "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,abcd"}, - } - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_resolve_image_part_preserves_base64_png_mime_type(): - provider = _make_provider() - try: - image_buffer = BytesIO() - PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save( - image_buffer, - format="PNG", - ) - image_base64 = base64.b64encode(image_buffer.getvalue()).decode("ascii") - - image_part = await provider._resolve_image_part(f"base64://{image_base64}") - - assert image_part == { - "type": "image_url", - "image_url": {"url": f"data:image/png;base64,{image_base64}"}, - } - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_materializes_context_localhost_file_uri_image_urls( - tmp_path, -): - provider = _make_provider() - try: - image_path = tmp_path / "quoted-image.png" - PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) - - localhost_uri = f"file://localhost{image_path.as_posix()}" - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=[ - { - "role": "user", - "content": [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": localhost_uri, - }, - }, - ], - } - ], - ) - - image_payload = payloads["messages"][0]["content"][1]["image_url"] - assert image_payload["url"].startswith("data:image/png;base64,") - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_resolve_audio_part_supports_data_audio_uri(tmp_path, monkeypatch): - monkeypatch.setattr( - "astrbot.core.utils.media_utils.get_astrbot_temp_path", - lambda: str(tmp_path), - ) - provider = _make_provider() - try: - audio_bytes = b"RIFF\x24\x00\x00\x00WAVEfmt " + b"\x00" * 16 - audio_ref = f"data:audio/wav;base64,{base64.b64encode(audio_bytes).decode()}" - - audio_part = await provider._resolve_audio_part(audio_ref) - - assert audio_part == { - "type": "input_audio", - "input_audio": { - "data": base64.b64encode(audio_bytes).decode("utf-8"), - "format": "wav", - }, - } - assert not list(tmp_path.iterdir()) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_resolve_audio_part_supports_base64_scheme(tmp_path, monkeypatch): - monkeypatch.setattr( - "astrbot.core.utils.media_utils.get_astrbot_temp_path", - lambda: str(tmp_path), - ) - provider = _make_provider() - try: - audio_bytes = b"RIFF\x24\x00\x00\x00WAVEfmt " + b"\x00" * 16 - audio_ref = f"base64://{base64.b64encode(audio_bytes).decode()}" - - audio_part = await provider._resolve_audio_part(audio_ref) - - assert audio_part == { - "type": "input_audio", - "input_audio": { - "data": base64.b64encode(audio_bytes).decode("utf-8"), - "format": "wav", - }, - } - assert not list(tmp_path.iterdir()) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_audio_preprocess_failure_does_not_log_media_ref(monkeypatch): - provider = _make_provider() - captured: dict[str, object] = {} - - async def fake_resolve_media_ref_to_base64_data(*args, **kwargs): - raise ValueError("boom") - - def fake_warning(message, *args, **kwargs): - captured["message"] = message - captured["args"] = args - - monkeypatch.setattr( - openai_source_module, - "resolve_media_ref_to_base64_data", - fake_resolve_media_ref_to_base64_data, - ) - monkeypatch.setattr(openai_source_module.logger, "warning", fake_warning) - - try: - audio_ref = "data:audio/wav;base64," + "A" * 1000 - - assert await provider._resolve_audio_part(audio_ref) is None - - assert captured["message"] == "音频预处理失败,将忽略。错误: %s" - assert len(captured["args"]) == 1 - assert str(captured["args"][0]) == "boom" - rendered_log_args = f"{captured['message']} {captured['args']}" - assert audio_ref not in rendered_log_args - assert "data:audio" not in rendered_log_args - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_prepare_chat_payload_keeps_original_context_image_when_materialization_fails( - monkeypatch, -): - provider = _make_provider() - try: - - async def fake_resolve_media_ref_to_base64_data( - media_ref: str, - *, - media_type: str, - strict: bool = False, - ) -> None: - assert media_ref == "https://example.com/expired.png" - assert media_type == "image" - assert strict is False - return None - - monkeypatch.setattr( - openai_source_module, - "resolve_media_ref_to_base64_data", - fake_resolve_media_ref_to_base64_data, - ) - - payloads, _ = await provider._prepare_chat_payload( - prompt=None, - contexts=[ - { - "role": "user", - "content": [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/expired.png", - }, - }, - ], - } - ], - ) - - assert payloads["messages"][0]["content"] == [ - {"type": "text", "text": "look"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/expired.png", - }, - }, - ] - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_apply_provider_specific_request_overrides_disables_ollama_thinking(): - provider = _make_provider( - { - "provider": "ollama", - "ollama_disable_thinking": True, - } - ) - try: - extra_body = { - "reasoning": {"effort": "high"}, - "reasoning_effort": "low", - "think": True, - "temperature": 0.2, - } - - provider._apply_provider_specific_request_overrides({}, extra_body) - - assert extra_body["reasoning_effort"] == "none" - assert "reasoning" not in extra_body - assert "think" not in extra_body - assert extra_body["temperature"] == 0.2 - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_provider_specific_request_overrides_sets_minimax_m3_max_tokens(): - provider = _make_provider({"provider": "nvidia"}) - try: - payloads = {"model": "minimaxai/minimax-m3"} - extra_body = {"temperature": 0.2} - - provider._apply_provider_specific_request_overrides(payloads, extra_body) - - assert payloads["max_tokens"] == 8192 - assert extra_body == {"temperature": 0.2} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_minimax_m3_max_tokens_preserves_custom_extra_body_value(): - provider = _make_provider({"provider": "nvidia"}) - try: - payloads = {"model": "minimaxai/minimax-m3"} - extra_body = {"max_tokens": 4096} - - provider._apply_provider_specific_request_overrides(payloads, extra_body) - - assert "max_tokens" not in payloads - assert extra_body["max_tokens"] == 4096 - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_minimax_m3_max_tokens_preserves_standard_payload_value(): - provider = _make_provider({"provider": "nvidia"}) - try: - payloads = { - "model": "minimaxai/minimax-m3", - "max_tokens": 2048, - } - extra_body = {} - - provider._apply_provider_specific_request_overrides(payloads, extra_body) - - assert payloads["max_tokens"] == 2048 - assert extra_body == {} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_nvidia_request_does_not_set_max_tokens_for_other_models(): - provider = _make_provider({"provider": "nvidia"}) - try: - payloads = {"model": "nvidia/usdcode"} - extra_body = {} - - provider._apply_provider_specific_request_overrides(payloads, extra_body) - - assert "max_tokens" not in payloads - assert "max_tokens" not in extra_body - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_injects_reasoning_effort_none_for_ollama(monkeypatch): - provider = _make_provider( - { - "provider": "ollama", - "ollama_disable_thinking": True, - "custom_extra_body": { - "reasoning": {"effort": "high"}, - "temperature": 0.1, - }, - } - ) - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "qwen3.5:4b", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - await provider._query( - payloads={ - "model": "qwen3.5:4b", - "messages": [{"role": "user", "content": "hello"}], - }, - tools=None, - ) - - extra_body = captured_kwargs["extra_body"] - assert extra_body["reasoning_effort"] == "none" - assert "reasoning" not in extra_body - assert extra_body["temperature"] == 0.1 - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_parse_openai_completion_raises_empty_model_output_error(): - provider = _make_provider() - try: - completion = ChatCompletion.model_validate( - { - "id": "chatcmpl-empty", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": None, - "refusal": None, - "tool_calls": None, - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 0, - "total_tokens": 1, - }, - } - ) - - with pytest.raises(EmptyModelOutputError): - await provider._parse_openai_completion(completion, tools=None) - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_parse_openai_completion_reads_nested_data_choices(): - provider = _make_provider() - try: - completion = ChatCompletion.model_construct( - id=None, - object="chat.completion", - created=None, - model=None, - choices=None, - data={ - "id": "gen_test", - "object": "chat.completion", - "created": 0, - "model": "deepseek/deepseek-v4-flash", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "PONG", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 12, - "completion_tokens": 38, - "total_tokens": 50, - }, - }, - ) - - response = await provider._parse_openai_completion(completion, tools=None) - - assert response.completion_text == "PONG" - assert response.id == "gen_test" - assert response.usage is not None - assert response.usage.input_other == 12 - assert response.usage.output == 38 - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_stream_extracts_usage_from_empty_choices_chunk(monkeypatch): - provider = _make_provider() - try: - chunks = [ - ChatCompletionChunk.model_validate( - { - "id": "chatcmpl-stream", - "object": "chat.completion.chunk", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "delta": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": None, - } - ], - } - ), - ChatCompletionChunk.model_validate( - { - "id": "chatcmpl-stream", - "object": "chat.completion.chunk", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "delta": {}, - "finish_reason": "stop", - } - ], - } - ), - ChatCompletionChunk.model_validate( - { - "id": "chatcmpl-stream", - "object": "chat.completion.chunk", - "created": 0, - "model": "gpt-4o-mini", - "choices": [], - "usage": { - "prompt_tokens": 2550, - "completion_tokens": 125, - "total_tokens": 2675, - "prompt_tokens_details": { - "cached_tokens": 2488, - }, - }, - } - ), - ] - - async def fake_stream(): - for chunk in chunks: - yield chunk - - async def fake_create(**kwargs): - return fake_stream() - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - responses = [ - response - async for response in provider._query_stream( - payloads={ - "model": "gpt-4o-mini", - "messages": [{"role": "user", "content": "hello"}], - }, - tools=None, - ) - ] - - final_response = responses[-1] - assert final_response.completion_text == "ok" - assert final_response.usage is not None - assert final_response.usage.input_other == 62 - assert final_response.usage.input_cached == 2488 - assert final_response.usage.output == 125 - finally: - await provider.terminate() - - -def test_sanitize_assistant_messages_removes_orphaned_tool_messages(): - payloads = { - "messages": [ - {"role": "user", "content": "hello"}, - { - "role": "tool", - "tool_call_id": "missing_call", - "content": "stale result", - }, - {"role": "user", "content": "continue"}, - ] - } - - ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) - - assert payloads["messages"] == [ - {"role": "user", "content": "hello"}, - {"role": "user", "content": "continue"}, - ] - - -def test_sanitize_assistant_messages_keeps_valid_tool_messages_only(): - payloads = { - "messages": [ - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_00", - "type": "function", - "function": {"name": "search", "arguments": "{}"}, - } - ], - }, - {"role": "tool", "tool_call_id": "call_00", "content": "one"}, - { - "role": "tool", - "tool_call_id": "", - "content": "empty id should not be valid", - }, - ] - } - - ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) - - assert payloads["messages"] == [ - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_00", - "type": "function", - "function": {"name": "search", "arguments": "{}"}, - } - ], - }, - {"role": "tool", "tool_call_id": "call_00", "content": "one"}, - ] - - -def test_sanitize_assistant_messages_removes_stale_duplicate_tool_message(): - payloads = { - "messages": [ - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_00", - "type": "function", - "function": {"name": "search", "arguments": "{}"}, - } - ], - }, - {"role": "tool", "tool_call_id": "call_00", "content": "one"}, - { - "role": "tool", - "tool_call_id": "call_00", - "content": "stale duplicate", - }, - {"role": "assistant", "content": "done"}, - ] - } - - ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) - - assert payloads["messages"] == [ - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_00", - "type": "function", - "function": {"name": "search", "arguments": "{}"}, - } - ], - }, - {"role": "tool", "tool_call_id": "call_00", "content": "one"}, - {"role": "assistant", "content": "done"}, - ] - - -def test_sanitize_assistant_messages_resets_tool_ids_after_non_tool_message(): - payloads = { - "messages": [ - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_00", - "type": "function", - "function": {"name": "search", "arguments": "{}"}, - } - ], - }, - {"role": "user", "content": "new turn"}, - { - "role": "tool", - "tool_call_id": "call_00", - "content": "stale late result", - }, - ] - } - - ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) - - assert payloads["messages"] == [ - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_00", - "type": "function", - "function": {"name": "search", "arguments": "{}"}, - } - ], - }, - {"role": "user", "content": "new turn"}, - ] - - -@pytest.mark.asyncio -async def test_query_filters_empty_assistant_message_without_tool_calls(monkeypatch): - """Test that empty assistant messages without tool_calls are filtered out.""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "user", "content": "hello"}, - {"role": "assistant", "content": ""}, # Should be filtered - {"role": "user", "content": "world"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - # The empty assistant message should be filtered out - messages = captured_kwargs["messages"] - assert len(messages) == 2 - assert messages[0] == {"role": "user", "content": "hello"} - assert messages[1] == {"role": "user", "content": "world"} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_filters_null_content_assistant_message_without_tool_calls( - monkeypatch, -): - """Test that assistant messages with null content and no tool_calls are filtered.""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "user", "content": "hello"}, - {"role": "assistant", "content": None}, # Should be filtered - {"role": "user", "content": "world"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - # The null content assistant message should be filtered out - messages = captured_kwargs["messages"] - assert len(messages) == 2 - assert messages[0] == {"role": "user", "content": "hello"} - assert messages[1] == {"role": "user", "content": "world"} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_converts_empty_content_to_none_with_tool_calls(monkeypatch): - """Test that empty content with tool_calls is converted to None (OpenAI spec).""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "user", "content": "hello"}, - { - "role": "assistant", - "content": "", - "tool_calls": [ - { - "id": "call-123", - "type": "function", - "function": {"name": "test", "arguments": "{}"}, - } - ], - }, - {"role": "user", "content": "world"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - # The assistant message with tool_calls should be kept but content set to None - messages = captured_kwargs["messages"] - assert len(messages) == 3 - assert messages[1]["role"] == "assistant" - assert messages[1]["content"] is None - assert messages[1]["tool_calls"] is not None - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_keeps_valid_assistant_message_with_content(monkeypatch): - """Test that valid assistant messages with content are kept.""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "user", "content": "hello"}, - {"role": "assistant", "content": "response"}, - {"role": "user", "content": "world"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - # All messages should be kept - messages = captured_kwargs["messages"] - assert len(messages) == 3 - assert messages[1] == {"role": "assistant", "content": "response"} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_keeps_assistant_message_with_tool_calls_and_none_content( - monkeypatch, -): - """Test that assistant messages with tool_calls and None content are kept.""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "user", "content": "hello"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call-123", - "type": "function", - "function": {"name": "test", "arguments": "{}"}, - } - ], - }, - {"role": "user", "content": "world"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - # The assistant message with tool_calls should be kept - messages = captured_kwargs["messages"] - assert len(messages) == 3 - assert messages[1]["role"] == "assistant" - assert messages[1]["content"] is None - assert messages[1]["tool_calls"] is not None - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_does_not_filter_user_or_system_messages(monkeypatch): - """Test that user and system messages are not affected by the filter.""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "ok", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "system", "content": ""}, # Empty system message - {"role": "user", "content": ""}, # Empty user message - {"role": "assistant", "content": ""}, # Should be filtered - {"role": "user", "content": "hello"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - # Only assistant message should be filtered - messages = captured_kwargs["messages"] - assert len(messages) == 3 - assert messages[0] == {"role": "system", "content": ""} - assert messages[1] == {"role": "user", "content": ""} - assert messages[2] == {"role": "user", "content": "hello"} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_stream_filters_empty_assistant_message(monkeypatch): - """Regression for #7721: streaming path must also filter empty assistant messages. - - Previously only ``_query`` sanitized the payload; ``_query_stream`` forwarded - the raw history and strict providers (e.g. DeepSeek Reasoner) returned 400 on - the next turn after a tool call whose assistant entry had reasoning only. - """ - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_stream(): - yield ChatCompletionChunk.model_validate( - { - "id": "chatcmpl-stream", - "object": "chat.completion.chunk", - "created": 0, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": {"role": "assistant", "content": "ok"}, - "finish_reason": "stop", - } - ], - } - ) - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return fake_stream() - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "deepseek-reasoner", - "messages": [ - {"role": "user", "content": "hello"}, - {"role": "assistant", "content": ""}, # should be filtered - {"role": "user", "content": "world"}, - ], - } - - async for _ in provider._query_stream(payloads=payloads, tools=None): - pass - - messages = captured_kwargs["messages"] - assert len(messages) == 2 - assert messages[0] == {"role": "user", "content": "hello"} - assert messages[1] == {"role": "user", "content": "world"} - finally: - await provider.terminate() - - -@pytest.mark.asyncio -async def test_query_filters_empty_list_content_assistant_message(monkeypatch): - """Empty-list content (``content == []``) must also be filtered, not just ``""`` / ``None``.""" - provider = _make_provider() - try: - captured_kwargs = {} - - async def fake_create(**kwargs): - captured_kwargs.update(kwargs) - return ChatCompletion.model_validate( - { - "id": "chatcmpl-test", - "object": "chat.completion", - "created": 0, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "ok"}, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 1, - "completion_tokens": 1, - "total_tokens": 2, - }, - } - ) - - monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) - - payloads = { - "model": "gpt-4o-mini", - "messages": [ - {"role": "user", "content": "hi"}, - {"role": "assistant", "content": []}, # should be filtered - {"role": "user", "content": "again"}, - ], - } - - await provider._query(payloads=payloads, tools=None) - - messages = captured_kwargs["messages"] - assert len(messages) == 2 - assert messages[0] == {"role": "user", "content": "hi"} - assert messages[1] == {"role": "user", "content": "again"} - finally: - await provider.terminate() +import base64 +import builtins +from io import BytesIO +from types import SimpleNamespace + +import httpx +import pytest +from openai.types.chat.chat_completion import ( + ChatCompletion, + ChatCompletionMessage, + Choice, +) +from openai.types.chat.chat_completion_chunk import ChatCompletionChunk +from PIL import Image as PILImage + +import astrbot.core.provider.sources.openai_source as openai_source_module +import astrbot.core.provider.sources.request_retry as request_retry +from astrbot.core.exceptions import EmptyModelOutputError +from astrbot.core.provider.entities import LLMResponse +from astrbot.core.provider.sources.groq_source import ProviderGroq +from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial +from astrbot.core.utils.media_utils import ResolvedMediaData, file_uri_to_path + + +class _ErrorWithBody(Exception): + def __init__(self, message: str, body: dict): + super().__init__(message) + self.body = body + + +class _ErrorWithResponse(Exception): + def __init__(self, message: str, response_text: str): + super().__init__(message) + self.response = SimpleNamespace(text=response_text) + + +def _make_provider(overrides: dict | None = None) -> ProviderOpenAIOfficial: + provider_config = { + "id": "test-openai", + "type": "openai_chat_completion", + "model": "gpt-4o-mini", + "key": ["test-key"], + } + if overrides: + provider_config.update(overrides) + return ProviderOpenAIOfficial( + provider_config=provider_config, + provider_settings={}, + ) + + +def _make_groq_provider(overrides: dict | None = None) -> ProviderGroq: + provider_config = { + "id": "test-groq", + "type": "groq_chat_completion", + "model": "qwen/qwen3-32b", + "key": ["test-key"], + } + if overrides: + provider_config.update(overrides) + return ProviderGroq( + provider_config=provider_config, + provider_settings={}, + ) + + +def test_create_http_client_uses_openai_httpx_module(monkeypatch): + captured: dict[str, object] = {} + fake_httpx_module = object() + + from openai import _base_client as openai_base_client + + monkeypatch.setattr( + openai_base_client, + "httpx", + fake_httpx_module, + raising=False, + ) + + def fake_create_proxy_client( + provider_label: str, + proxy: str | None = None, + headers: dict[str, str] | None = None, + verify=None, + httpx_module=None, + ): + captured["httpx_module"] = httpx_module + return object() + + monkeypatch.setattr( + openai_source_module, + "create_proxy_client", + fake_create_proxy_client, + ) + + provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) + provider._create_http_client({"proxy": ""}) + + assert captured["httpx_module"] is fake_httpx_module + + +def test_create_http_client_falls_back_to_global_httpx_module(monkeypatch): + captured: dict[str, object] = {} + + def fake_create_proxy_client( + provider_label: str, + proxy: str | None = None, + headers: dict[str, str] | None = None, + verify=None, + httpx_module=None, + ): + captured["httpx_module"] = httpx_module + return object() + + real_import = builtins.__import__ + + def fake_import(name, globals=None, locals=None, fromlist=(), level=0): + if name == "openai" and fromlist: + raise ImportError("missing openai._base_client") + return real_import(name, globals, locals, fromlist, level) + + monkeypatch.setattr( + openai_source_module, + "create_proxy_client", + fake_create_proxy_client, + ) + monkeypatch.setattr(builtins, "__import__", fake_import) + + provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) + provider._create_http_client({"proxy": ""}) + + assert captured["httpx_module"] is openai_source_module.httpx + + +@pytest.mark.asyncio +async def test_get_models_retries_transient_request_error(monkeypatch): + monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0) + monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0) + + class FakeModels: + def __init__(self): + self.calls = 0 + + async def list(self): + self.calls += 1 + if self.calls == 1: + raise httpx.ConnectError("temporary connection failure") + return SimpleNamespace( + data=[ + SimpleNamespace(id="gpt-b"), + SimpleNamespace(id="gpt-a"), + ] + ) + + models = FakeModels() + provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) + provider.client = SimpleNamespace(models=models) + + assert await provider.get_models() == ["gpt-a", "gpt-b"] + assert models.calls == 2 + + +@pytest.mark.asyncio +async def test_text_chat_passes_request_max_retries_to_query(): + captured: dict[str, object] = {} + + provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) + provider.api_keys = ["test-key"] + provider.client = SimpleNamespace(api_key=None) + + async def fake_prepare_chat_payload(*args, **kwargs): + return {"messages": [], "model": "gpt-4o-mini"}, [] + + async def fake_query(payloads, func_tool, *, request_max_retries=None): + captured["request_max_retries"] = request_max_retries + return LLMResponse(role="assistant", completion_text="ok") + + provider._prepare_chat_payload = fake_prepare_chat_payload + provider._query = fake_query + + await provider.text_chat(prompt="hello", request_max_retries=2) + + assert captured["request_max_retries"] == 2 + + +@pytest.mark.asyncio +async def test_handle_api_error_content_moderated_removes_images(): + provider = _make_provider( + {"image_moderation_error_patterns": ["file:content-moderated"]} + ) + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + + success, *_rest = await provider._handle_api_error( + Exception("Content is moderated [WKE=file:content-moderated]"), + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + + assert success is False + updated_context = payloads["messages"] + assert isinstance(updated_context, list) + assert updated_context[0]["content"] == [{"type": "text", "text": "hello"}] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_model_not_vlm_removes_images_and_retries_text_only(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + + success, *_rest = await provider._handle_api_error( + Exception("The model is not a VLM and cannot process images"), + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + + assert success is False + updated_context = payloads["messages"] + assert isinstance(updated_context, list) + assert updated_context[0]["content"] == [{"type": "text", "text": "hello"}] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_model_not_vlm_after_fallback_raises(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + + with pytest.raises(Exception, match="not a VLM"): + await provider._handle_api_error( + Exception("The model is not a VLM and cannot process images"), + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=1, + max_retries=10, + image_fallback_used=True, + ) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_content_moderated_with_unserializable_body(): + provider = _make_provider({"image_moderation_error_patterns": ["blocked"]}) + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + err = _ErrorWithBody( + "upstream error", + {"error": {"message": "blocked"}, "raw": object()}, + ) + + success, *_rest = await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + assert success is False + assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] + finally: + await provider.terminate() + + +def test_extract_error_text_candidates_truncates_long_response_text(): + long_text = "x" * 20000 + err = _ErrorWithResponse("upstream error", long_text) + candidates = ProviderOpenAIOfficial._extract_error_text_candidates(err) + assert candidates + assert max(len(candidate) for candidate in candidates) <= ( + ProviderOpenAIOfficial._ERROR_TEXT_CANDIDATE_MAX_CHARS + ) + + +@pytest.mark.asyncio +async def test_openai_payload_keeps_reasoning_content_in_assistant_history(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "assistant", + "content": [ + {"type": "think", "think": "step 1"}, + {"type": "text", "text": "final answer"}, + ], + } + ] + } + + provider._finally_convert_payload(payloads) + + assistant_message = payloads["messages"][0] + assert assistant_message["content"] == [ + {"type": "text", "text": "final answer"} + ] + assert assistant_message["reasoning_content"] == "step 1" + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_groq_payload_drops_reasoning_content_from_assistant_history(): + provider = _make_groq_provider() + try: + payloads = { + "messages": [ + { + "role": "assistant", + "content": [ + {"type": "think", "think": "step 1"}, + {"type": "text", "text": "final answer"}, + ], + } + ] + } + + provider._finally_convert_payload(payloads) + + assistant_message = payloads["messages"][0] + assert assistant_message["content"] == [ + {"type": "text", "text": "final answer"} + ] + assert "reasoning_content" not in assistant_message + assert "reasoning" not in assistant_message + finally: + await provider.terminate() + + +def _make_reasoning_completion( + thinking_field: str, +) -> ChatCompletion: + kwargs = {thinking_field: "thoughts"} + message = ChatCompletionMessage(role="assistant", content="answer", **kwargs) + return ChatCompletion( + id="chatcmpl-reasoning-test", + choices=[Choice(index=0, finish_reason="stop", message=message)], + created=1, + model="test-model", + object="chat.completion", + ) + + +@pytest.mark.asyncio +async def test_reasoning_key_configurable_extracts_alias_field(): + """reasoning_key 可配置:按配置字段名提取思考内容 (#9783)""" + provider = _make_provider({"reasoning_key": "reasoning"}) + try: + completion = _make_reasoning_completion("reasoning") + assert provider._extract_reasoning_content(completion) == "thoughts" + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_reasoning_key_default_keeps_reasoning_content_behavior(): + """默认 key 仍为 reasoning_content,不取 reasoning 别名(行为不回归)""" + provider = _make_provider() + try: + aliased = _make_reasoning_completion("reasoning") + assert provider._extract_reasoning_content(aliased) is None + standard = _make_reasoning_completion("reasoning_content") + assert provider._extract_reasoning_content(standard) == "thoughts" + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_reasoning_key_applied_to_assistant_history_payload(): + """配置 reasoning_key 后,历史 think 内容写入配置的字段名 (#9783)""" + provider = _make_provider({"reasoning_key": "reasoning"}) + try: + payloads = { + "model": "test-model", + "messages": [ + { + "role": "assistant", + "content": [ + {"type": "think", "think": "step 1"}, + {"type": "text", "text": "final answer"}, + ], + } + ], + } + + provider._finally_convert_payload(payloads) + + assistant_message = payloads["messages"][0] + assert assistant_message["reasoning"] == "step 1" + assert "reasoning_content" not in assistant_message + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_sanitize_assistant_messages_uses_configured_reasoning_key(): + """reasoning_key='reasoning' 时,think-only 历史不被误删为空消息 (#9783, PR #9829 review)""" + provider = _make_provider({"reasoning_key": "reasoning"}) + try: + payloads = { + "model": "test-model", + "messages": [ + { + "role": "assistant", + "content": [{"type": "think", "think": "step 1"}], + } + ], + } + + provider._finally_convert_payload(payloads) + provider._sanitize_assistant_messages(payloads, provider.reasoning_key) + + assert len(payloads["messages"]) == 1 + assistant_message = payloads["messages"][0] + assert assistant_message["reasoning"] == "step 1" + assert assistant_message["content"] == "" + finally: + await provider.terminate() + + +def test_sanitize_assistant_messages_keeps_custom_key_only_history(): + """自定义 reasoning_key 下:该字段或默认字段的思考历史都保留,真空消息仍丢弃""" + payloads = { + "messages": [ + { + "role": "assistant", + "content": None, + "reasoning": "thinking under custom key", + }, + { + "role": "assistant", + "content": None, + "reasoning_content": "legacy thinking under default key", + }, + {"role": "assistant", "content": None}, + ] + } + + ProviderOpenAIOfficial._sanitize_assistant_messages(payloads, "reasoning") + + assert len(payloads["messages"]) == 2 + assert payloads["messages"][0]["reasoning"] == "thinking under custom key" + assert payloads["messages"][0]["content"] == "" + assert ( + payloads["messages"][1]["reasoning_content"] + == "legacy thinking under default key" + ) + assert payloads["messages"][1]["content"] == "" + + +@pytest.mark.asyncio +async def test_handle_api_error_content_moderated_without_images_raises(): + provider = _make_provider( + {"image_moderation_error_patterns": ["file:content-moderated"]} + ) + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [{"type": "text", "text": "hello"}], + } + ] + } + context_query = payloads["messages"] + err = Exception("Content is moderated [WKE=file:content-moderated]") + + with pytest.raises(Exception, match="content-moderated"): + await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_content_moderated_detects_structured_body(): + provider = _make_provider( + {"image_moderation_error_patterns": ["content_moderated"]} + ) + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + err = _ErrorWithBody( + "upstream error", + {"error": {"code": "content_moderated", "message": "blocked"}}, + ) + + success, *_rest = await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + assert success is False + assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_content_moderated_supports_custom_patterns(): + provider = _make_provider( + {"image_moderation_error_patterns": ["blocked_by_policy_code_123"]} + ) + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + err = Exception("upstream: blocked_by_policy_code_123") + + success, *_rest = await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + assert success is False + assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_content_moderated_without_patterns_raises(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + err = Exception("Content is moderated [WKE=file:content-moderated]") + + with pytest.raises(Exception, match="content-moderated"): + await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_unknown_image_error_raises(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + + with pytest.raises(Exception, match="unknown provider image upload error"): + await provider._handle_api_error( + Exception("some unknown provider image upload error"), + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_invalid_attachment_removes_images_and_retries_text_only(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + err = _ErrorWithBody( + "upstream error", + { + "error": { + "code": "INVALID_ATTACHMENT", + "message": "download attachment: unexpected status 404", + } + }, + ) + + success, *_rest = await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + + assert success is False + assert payloads["messages"][0]["content"] == [{"type": "text", "text": "hello"}] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_invalid_attachment_without_images_raises(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [{"type": "text", "text": "hello"}], + } + ] + } + context_query = payloads["messages"] + err = _ErrorWithBody( + "upstream error", + { + "error": { + "code": "INVALID_ATTACHMENT", + "message": "download attachment: unexpected status 404", + } + }, + ) + + with pytest.raises(_ErrorWithBody, match="upstream error"): + await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=0, + max_retries=10, + ) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_handle_api_error_invalid_attachment_after_fallback_raises(): + provider = _make_provider() + try: + payloads = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello"}, + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + }, + ], + } + ] + } + context_query = payloads["messages"] + err = _ErrorWithBody( + "upstream error", + { + "error": { + "code": "INVALID_ATTACHMENT", + "message": "download attachment: unexpected status 404", + } + }, + ) + + with pytest.raises(_ErrorWithBody, match="upstream error"): + await provider._handle_api_error( + err, + payloads=payloads, + context_query=context_query, + func_tool=None, + chosen_key="test-key", + available_api_keys=["test-key"], + retry_cnt=1, + max_retries=10, + image_fallback_used=True, + ) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_materializes_context_http_image_urls(monkeypatch): + provider = _make_provider() + try: + + async def fake_resolve_media_ref_to_base64_data( + media_ref: str, + *, + media_type: str, + strict: bool = False, + ) -> ResolvedMediaData: + assert media_ref == "https://example.com/quoted.png" + assert media_type == "image" + assert strict is False + return ResolvedMediaData(base64_data="abcd", mime_type="image/png") + + monkeypatch.setattr( + openai_source_module, + "resolve_media_ref_to_base64_data", + fake_resolve_media_ref_to_base64_data, + ) + + contexts = [ + { + "role": "user", + "metadata": {"source": "quoted"}, + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/quoted.png", + "id": "ctx-img", + "detail": "high", + }, + }, + ], + } + ] + + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=contexts, + ) + + assert payloads["messages"][0]["content"] == [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,abcd", + "detail": "high", + }, + }, + ] + assert payloads["messages"][0]["content"][1]["image_url"].get("id") is None + assert contexts[0]["content"][1]["image_url"] == { + "url": "https://example.com/quoted.png", + "id": "ctx-img", + "detail": "high", + } + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_skips_materialization_for_text_only_context( + monkeypatch, +): + provider = _make_provider() + try: + + async def fail_if_called(_context_query): + raise AssertionError("materialization should be skipped") + + monkeypatch.setattr( + provider, "_materialize_context_image_parts", fail_if_called + ) + + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=[{"role": "user", "content": "hello"}], + ) + + assert payloads["messages"] == [{"role": "user", "content": "hello"}] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_skips_materialization_for_text_only_parts( + monkeypatch, +): + provider = _make_provider() + try: + + async def fail_if_called(_context_query): + raise AssertionError("materialization should be skipped") + + monkeypatch.setattr( + provider, "_materialize_context_image_parts", fail_if_called + ) + + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=[ + { + "role": "user", + "content": [{"type": "text", "text": "hello"}], + } + ], + ) + + assert payloads["messages"] == [ + { + "role": "user", + "content": [{"type": "text", "text": "hello"}], + } + ] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_materializes_context_http_image_urls_with_detected_mime( + monkeypatch, tmp_path +): + provider = _make_provider() + try: + image_path = tmp_path / "quoted-image.png" + PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) + + async def fake_download(url: str, target_path: str) -> None: + assert url == "https://example.com/quoted.png" + with open(target_path, "wb") as f: + f.write(image_path.read_bytes()) + + monkeypatch.setattr( + "astrbot.core.utils.media_utils.download_file", + fake_download, + ) + + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/quoted.png", + }, + }, + ], + } + ], + ) + + image_payload = payloads["messages"][0]["content"][1]["image_url"] + assert image_payload["url"].startswith("data:image/png;base64,") + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_materializes_context_file_uri_image_urls(tmp_path): + provider = _make_provider() + try: + image_path = tmp_path / "quoted-image.png" + PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) + + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": image_path.as_uri(), + }, + }, + ], + } + ], + ) + + image_payload = payloads["messages"][0]["content"][1]["image_url"] + assert image_payload["url"].startswith("data:image/png;base64,") + finally: + await provider.terminate() + + +def test_file_uri_to_path_preserves_windows_drive_letter(): + assert file_uri_to_path("file:///C:/tmp/quoted-image.png") == ( + "C:/tmp/quoted-image.png" + ) + + +def test_file_uri_to_path_preserves_windows_netloc_drive_letter(): + assert file_uri_to_path("file://C:/tmp/quoted-image.png") == ( + "C:/tmp/quoted-image.png" + ) + + +def test_file_uri_to_path_preserves_remote_netloc_as_unc_path(): + assert file_uri_to_path("file://server/share/quoted-image.png") == ( + "//server/share/quoted-image.png" + ) + + +@pytest.mark.asyncio +async def test_resolve_image_part_rejects_invalid_local_file(tmp_path): + provider = _make_provider() + try: + invalid_file = tmp_path / "not-image.txt" + invalid_file.write_text("not an image") + + assert await provider._resolve_image_part(str(invalid_file)) is None + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_resolve_image_part_rejects_invalid_file_uri(tmp_path): + provider = _make_provider() + try: + invalid_file = tmp_path / "not-image.txt" + invalid_file.write_text("not an image") + + assert await provider._resolve_image_part(invalid_file.as_uri()) is None + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_image_ref_to_data_url_mode_controls_invalid_file_behavior(tmp_path): + provider = _make_provider() + try: + invalid_file = tmp_path / "not-image.txt" + invalid_file.write_text("not an image") + + assert ( + await provider._image_ref_to_data_url(str(invalid_file), mode="safe") + is None + ) + with pytest.raises(ValueError, match="Invalid image file"): + await provider._image_ref_to_data_url(str(invalid_file), mode="strict") + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_materialize_context_image_parts_returns_new_messages(monkeypatch): + provider = _make_provider() + try: + context_query = [ + { + "role": "user", + "metadata": {"source": "quoted"}, + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/quoted.png", + "detail": "high", + }, + }, + ], + }, + {"role": "assistant", "content": "plain text"}, + ] + + async def fake_resolve(image_url: str, *, image_detail: str | None = None): + assert image_url == "https://example.com/quoted.png" + assert image_detail == "high" + return { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,abcd", + "detail": "high", + }, + } + + monkeypatch.setattr(provider, "_resolve_image_part", fake_resolve) + + materialized = await provider._materialize_context_image_parts(context_query) + + assert materialized is not context_query + assert materialized[0] is not context_query[0] + assert materialized[0]["metadata"] is context_query[0]["metadata"] + assert materialized[0]["content"][0] is context_query[0]["content"][0] + assert ( + materialized[0]["content"][1]["image_url"]["url"] + == "data:image/png;base64,abcd" + ) + assert ( + context_query[0]["content"][1]["image_url"]["url"] + == "https://example.com/quoted.png" + ) + assert materialized[1] is not context_query[1] + assert materialized[1]["content"] == "plain text" + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_encode_image_bs64_missing_file_raises(tmp_path): + provider = _make_provider() + try: + missing_path = tmp_path / "missing-image.png" + with pytest.raises(FileNotFoundError): + await provider.encode_image_bs64(str(missing_path)) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_encode_image_bs64_invalid_file_raises(tmp_path): + provider = _make_provider() + try: + invalid_file = tmp_path / "not-image.txt" + invalid_file.write_text("not an image") + + with pytest.raises(ValueError, match="Invalid image file"): + await provider.encode_image_bs64(str(invalid_file)) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_encode_image_bs64_supports_base64_scheme(): + provider = _make_provider() + try: + image_data = await provider.encode_image_bs64("base64://abcd") + + assert image_data == "data:image/jpeg;base64,abcd" + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_encode_image_bs64_supports_file_uri(tmp_path): + provider = _make_provider() + try: + image_path = tmp_path / "quoted-image.png" + PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) + + image_data = await provider.encode_image_bs64(image_path.as_uri()) + + assert image_data.startswith("data:image/png;base64,") + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_resolve_image_part_supports_base64_scheme(): + provider = _make_provider() + try: + assert await provider._resolve_image_part("base64://abcd") == { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,abcd"}, + } + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_resolve_image_part_preserves_base64_png_mime_type(): + provider = _make_provider() + try: + image_buffer = BytesIO() + PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save( + image_buffer, + format="PNG", + ) + image_base64 = base64.b64encode(image_buffer.getvalue()).decode("ascii") + + image_part = await provider._resolve_image_part(f"base64://{image_base64}") + + assert image_part == { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{image_base64}"}, + } + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_materializes_context_localhost_file_uri_image_urls( + tmp_path, +): + provider = _make_provider() + try: + image_path = tmp_path / "quoted-image.png" + PILImage.new("RGBA", (1, 1), (255, 0, 0, 255)).save(image_path) + + localhost_uri = f"file://localhost{image_path.as_posix()}" + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": localhost_uri, + }, + }, + ], + } + ], + ) + + image_payload = payloads["messages"][0]["content"][1]["image_url"] + assert image_payload["url"].startswith("data:image/png;base64,") + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_resolve_audio_part_supports_data_audio_uri(tmp_path, monkeypatch): + monkeypatch.setattr( + "astrbot.core.utils.media_utils.get_astrbot_temp_path", + lambda: str(tmp_path), + ) + provider = _make_provider() + try: + audio_bytes = b"RIFF\x24\x00\x00\x00WAVEfmt " + b"\x00" * 16 + audio_ref = f"data:audio/wav;base64,{base64.b64encode(audio_bytes).decode()}" + + audio_part = await provider._resolve_audio_part(audio_ref) + + assert audio_part == { + "type": "input_audio", + "input_audio": { + "data": base64.b64encode(audio_bytes).decode("utf-8"), + "format": "wav", + }, + } + assert not list(tmp_path.iterdir()) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_resolve_audio_part_supports_base64_scheme(tmp_path, monkeypatch): + monkeypatch.setattr( + "astrbot.core.utils.media_utils.get_astrbot_temp_path", + lambda: str(tmp_path), + ) + provider = _make_provider() + try: + audio_bytes = b"RIFF\x24\x00\x00\x00WAVEfmt " + b"\x00" * 16 + audio_ref = f"base64://{base64.b64encode(audio_bytes).decode()}" + + audio_part = await provider._resolve_audio_part(audio_ref) + + assert audio_part == { + "type": "input_audio", + "input_audio": { + "data": base64.b64encode(audio_bytes).decode("utf-8"), + "format": "wav", + }, + } + assert not list(tmp_path.iterdir()) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_audio_preprocess_failure_does_not_log_media_ref(monkeypatch): + provider = _make_provider() + captured: dict[str, object] = {} + + async def fake_resolve_media_ref_to_base64_data(*args, **kwargs): + raise ValueError("boom") + + def fake_warning(message, *args, **kwargs): + captured["message"] = message + captured["args"] = args + + monkeypatch.setattr( + openai_source_module, + "resolve_media_ref_to_base64_data", + fake_resolve_media_ref_to_base64_data, + ) + monkeypatch.setattr(openai_source_module.logger, "warning", fake_warning) + + try: + audio_ref = "data:audio/wav;base64," + "A" * 1000 + + assert await provider._resolve_audio_part(audio_ref) is None + + assert captured["message"] == "音频预处理失败,将忽略。错误: %s" + assert len(captured["args"]) == 1 + assert str(captured["args"][0]) == "boom" + rendered_log_args = f"{captured['message']} {captured['args']}" + assert audio_ref not in rendered_log_args + assert "data:audio" not in rendered_log_args + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_prepare_chat_payload_keeps_original_context_image_when_materialization_fails( + monkeypatch, +): + provider = _make_provider() + try: + + async def fake_resolve_media_ref_to_base64_data( + media_ref: str, + *, + media_type: str, + strict: bool = False, + ) -> None: + assert media_ref == "https://example.com/expired.png" + assert media_type == "image" + assert strict is False + return None + + monkeypatch.setattr( + openai_source_module, + "resolve_media_ref_to_base64_data", + fake_resolve_media_ref_to_base64_data, + ) + + payloads, _ = await provider._prepare_chat_payload( + prompt=None, + contexts=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/expired.png", + }, + }, + ], + } + ], + ) + + assert payloads["messages"][0]["content"] == [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/expired.png", + }, + }, + ] + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_apply_provider_specific_request_overrides_disables_ollama_thinking(): + provider = _make_provider( + { + "provider": "ollama", + "ollama_disable_thinking": True, + } + ) + try: + extra_body = { + "reasoning": {"effort": "high"}, + "reasoning_effort": "low", + "think": True, + "temperature": 0.2, + } + + provider._apply_provider_specific_request_overrides({}, extra_body) + + assert extra_body["reasoning_effort"] == "none" + assert "reasoning" not in extra_body + assert "think" not in extra_body + assert extra_body["temperature"] == 0.2 + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_provider_specific_request_overrides_sets_minimax_m3_max_tokens(): + provider = _make_provider({"provider": "nvidia"}) + try: + payloads = {"model": "minimaxai/minimax-m3"} + extra_body = {"temperature": 0.2} + + provider._apply_provider_specific_request_overrides(payloads, extra_body) + + assert payloads["max_tokens"] == 8192 + assert extra_body == {"temperature": 0.2} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_minimax_m3_max_tokens_preserves_custom_extra_body_value(): + provider = _make_provider({"provider": "nvidia"}) + try: + payloads = {"model": "minimaxai/minimax-m3"} + extra_body = {"max_tokens": 4096} + + provider._apply_provider_specific_request_overrides(payloads, extra_body) + + assert "max_tokens" not in payloads + assert extra_body["max_tokens"] == 4096 + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_minimax_m3_max_tokens_preserves_standard_payload_value(): + provider = _make_provider({"provider": "nvidia"}) + try: + payloads = { + "model": "minimaxai/minimax-m3", + "max_tokens": 2048, + } + extra_body = {} + + provider._apply_provider_specific_request_overrides(payloads, extra_body) + + assert payloads["max_tokens"] == 2048 + assert extra_body == {} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_nvidia_request_does_not_set_max_tokens_for_other_models(): + provider = _make_provider({"provider": "nvidia"}) + try: + payloads = {"model": "nvidia/usdcode"} + extra_body = {} + + provider._apply_provider_specific_request_overrides(payloads, extra_body) + + assert "max_tokens" not in payloads + assert "max_tokens" not in extra_body + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_injects_reasoning_effort_none_for_ollama(monkeypatch): + provider = _make_provider( + { + "provider": "ollama", + "ollama_disable_thinking": True, + "custom_extra_body": { + "reasoning": {"effort": "high"}, + "temperature": 0.1, + }, + } + ) + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "qwen3.5:4b", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + await provider._query( + payloads={ + "model": "qwen3.5:4b", + "messages": [{"role": "user", "content": "hello"}], + }, + tools=None, + ) + + extra_body = captured_kwargs["extra_body"] + assert extra_body["reasoning_effort"] == "none" + assert "reasoning" not in extra_body + assert extra_body["temperature"] == 0.1 + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_parse_openai_completion_raises_empty_model_output_error(): + provider = _make_provider() + try: + completion = ChatCompletion.model_validate( + { + "id": "chatcmpl-empty", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": None, + "refusal": None, + "tool_calls": None, + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 0, + "total_tokens": 1, + }, + } + ) + + with pytest.raises(EmptyModelOutputError): + await provider._parse_openai_completion(completion, tools=None) + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_parse_openai_completion_reads_nested_data_choices(): + provider = _make_provider() + try: + completion = ChatCompletion.model_construct( + id=None, + object="chat.completion", + created=None, + model=None, + choices=None, + data={ + "id": "gen_test", + "object": "chat.completion", + "created": 0, + "model": "deepseek/deepseek-v4-flash", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "PONG", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 12, + "completion_tokens": 38, + "total_tokens": 50, + }, + }, + ) + + response = await provider._parse_openai_completion(completion, tools=None) + + assert response.completion_text == "PONG" + assert response.id == "gen_test" + assert response.usage is not None + assert response.usage.input_other == 12 + assert response.usage.output == 38 + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_stream_extracts_usage_from_empty_choices_chunk(monkeypatch): + provider = _make_provider() + try: + chunks = [ + ChatCompletionChunk.model_validate( + { + "id": "chatcmpl-stream", + "object": "chat.completion.chunk", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "delta": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": None, + } + ], + } + ), + ChatCompletionChunk.model_validate( + { + "id": "chatcmpl-stream", + "object": "chat.completion.chunk", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "delta": {}, + "finish_reason": "stop", + } + ], + } + ), + ChatCompletionChunk.model_validate( + { + "id": "chatcmpl-stream", + "object": "chat.completion.chunk", + "created": 0, + "model": "gpt-4o-mini", + "choices": [], + "usage": { + "prompt_tokens": 2550, + "completion_tokens": 125, + "total_tokens": 2675, + "prompt_tokens_details": { + "cached_tokens": 2488, + }, + }, + } + ), + ] + + async def fake_stream(): + for chunk in chunks: + yield chunk + + async def fake_create(**kwargs): + return fake_stream() + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + responses = [ + response + async for response in provider._query_stream( + payloads={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + }, + tools=None, + ) + ] + + final_response = responses[-1] + assert final_response.completion_text == "ok" + assert final_response.usage is not None + assert final_response.usage.input_other == 62 + assert final_response.usage.input_cached == 2488 + assert final_response.usage.output == 125 + finally: + await provider.terminate() + + +def test_sanitize_assistant_messages_removes_orphaned_tool_messages(): + payloads = { + "messages": [ + {"role": "user", "content": "hello"}, + { + "role": "tool", + "tool_call_id": "missing_call", + "content": "stale result", + }, + {"role": "user", "content": "continue"}, + ] + } + + ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) + + assert payloads["messages"] == [ + {"role": "user", "content": "hello"}, + {"role": "user", "content": "continue"}, + ] + + +def test_sanitize_assistant_messages_keeps_valid_tool_messages_only(): + payloads = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_00", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_00", "content": "one"}, + { + "role": "tool", + "tool_call_id": "", + "content": "empty id should not be valid", + }, + ] + } + + ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) + + assert payloads["messages"] == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_00", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_00", "content": "one"}, + ] + + +def test_sanitize_assistant_messages_removes_stale_duplicate_tool_message(): + payloads = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_00", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_00", "content": "one"}, + { + "role": "tool", + "tool_call_id": "call_00", + "content": "stale duplicate", + }, + {"role": "assistant", "content": "done"}, + ] + } + + ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) + + assert payloads["messages"] == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_00", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_00", "content": "one"}, + {"role": "assistant", "content": "done"}, + ] + + +def test_sanitize_assistant_messages_resets_tool_ids_after_non_tool_message(): + payloads = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_00", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "user", "content": "new turn"}, + { + "role": "tool", + "tool_call_id": "call_00", + "content": "stale late result", + }, + ] + } + + ProviderOpenAIOfficial._sanitize_assistant_messages(payloads) + + assert payloads["messages"] == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_00", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "user", "content": "new turn"}, + ] + + +@pytest.mark.asyncio +async def test_query_filters_empty_assistant_message_without_tool_calls(monkeypatch): + """Test that empty assistant messages without tool_calls are filtered out.""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": ""}, # Should be filtered + {"role": "user", "content": "world"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + # The empty assistant message should be filtered out + messages = captured_kwargs["messages"] + assert len(messages) == 2 + assert messages[0] == {"role": "user", "content": "hello"} + assert messages[1] == {"role": "user", "content": "world"} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_filters_null_content_assistant_message_without_tool_calls( + monkeypatch, +): + """Test that assistant messages with null content and no tool_calls are filtered.""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": None}, # Should be filtered + {"role": "user", "content": "world"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + # The null content assistant message should be filtered out + messages = captured_kwargs["messages"] + assert len(messages) == 2 + assert messages[0] == {"role": "user", "content": "hello"} + assert messages[1] == {"role": "user", "content": "world"} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_converts_empty_content_to_none_with_tool_calls(monkeypatch): + """Test that empty content with tool_calls is converted to None (OpenAI spec).""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call-123", + "type": "function", + "function": {"name": "test", "arguments": "{}"}, + } + ], + }, + {"role": "user", "content": "world"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + # The assistant message with tool_calls should be kept but content set to None + messages = captured_kwargs["messages"] + assert len(messages) == 3 + assert messages[1]["role"] == "assistant" + assert messages[1]["content"] is None + assert messages[1]["tool_calls"] is not None + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_keeps_valid_assistant_message_with_content(monkeypatch): + """Test that valid assistant messages with content are kept.""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "response"}, + {"role": "user", "content": "world"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + # All messages should be kept + messages = captured_kwargs["messages"] + assert len(messages) == 3 + assert messages[1] == {"role": "assistant", "content": "response"} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_keeps_assistant_message_with_tool_calls_and_none_content( + monkeypatch, +): + """Test that assistant messages with tool_calls and None content are kept.""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-123", + "type": "function", + "function": {"name": "test", "arguments": "{}"}, + } + ], + }, + {"role": "user", "content": "world"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + # The assistant message with tool_calls should be kept + messages = captured_kwargs["messages"] + assert len(messages) == 3 + assert messages[1]["role"] == "assistant" + assert messages[1]["content"] is None + assert messages[1]["tool_calls"] is not None + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_does_not_filter_user_or_system_messages(monkeypatch): + """Test that user and system messages are not affected by the filter.""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "ok", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "system", "content": ""}, # Empty system message + {"role": "user", "content": ""}, # Empty user message + {"role": "assistant", "content": ""}, # Should be filtered + {"role": "user", "content": "hello"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + # Only assistant message should be filtered + messages = captured_kwargs["messages"] + assert len(messages) == 3 + assert messages[0] == {"role": "system", "content": ""} + assert messages[1] == {"role": "user", "content": ""} + assert messages[2] == {"role": "user", "content": "hello"} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_stream_filters_empty_assistant_message(monkeypatch): + """Regression for #7721: streaming path must also filter empty assistant messages. + + Previously only ``_query`` sanitized the payload; ``_query_stream`` forwarded + the raw history and strict providers (e.g. DeepSeek Reasoner) returned 400 on + the next turn after a tool call whose assistant entry had reasoning only. + """ + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_stream(): + yield ChatCompletionChunk.model_validate( + { + "id": "chatcmpl-stream", + "object": "chat.completion.chunk", + "created": 0, + "model": "deepseek-reasoner", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + } + ) + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return fake_stream() + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "deepseek-reasoner", + "messages": [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": ""}, # should be filtered + {"role": "user", "content": "world"}, + ], + } + + async for _ in provider._query_stream(payloads=payloads, tools=None): + pass + + messages = captured_kwargs["messages"] + assert len(messages) == 2 + assert messages[0] == {"role": "user", "content": "hello"} + assert messages[1] == {"role": "user", "content": "world"} + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_query_filters_empty_list_content_assistant_message(monkeypatch): + """Empty-list content (``content == []``) must also be filtered, not just ``""`` / ``None``.""" + provider = _make_provider() + try: + captured_kwargs = {} + + async def fake_create(**kwargs): + captured_kwargs.update(kwargs) + return ChatCompletion.model_validate( + { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + ) + + monkeypatch.setattr(provider.client.chat.completions, "create", fake_create) + + payloads = { + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": []}, # should be filtered + {"role": "user", "content": "again"}, + ], + } + + await provider._query(payloads=payloads, tools=None) + + messages = captured_kwargs["messages"] + assert len(messages) == 2 + assert messages[0] == {"role": "user", "content": "hi"} + assert messages[1] == {"role": "user", "content": "again"} + finally: + await provider.terminate()