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()