feat: Enhance OpenAI provider configuration with extraBody support and apiType validation
This commit is contained in:
@@ -173,7 +173,7 @@ class ProviderConfig(Base):
|
||||
api_base: str | None = None
|
||||
api_type: Literal["auto", "chat_completions", "responses"] = "auto" # Request API surface
|
||||
extra_headers: dict[str, str] | None = None # Custom headers (e.g. APP-Code for AiHubMix)
|
||||
extra_body: dict[str, Any] | None = None # Extra fields merged into every request body
|
||||
extra_body: dict[str, Any] | None = None # Extra provider request fields; shape depends on provider/API surface
|
||||
|
||||
|
||||
class BedrockProviderConfig(ProviderConfig):
|
||||
@@ -224,6 +224,16 @@ class ProvidersConfig(Base):
|
||||
qianfan: ProviderConfig = Field(default_factory=ProviderConfig) # Qianfan (百度千帆)
|
||||
nvidia: ProviderConfig = Field(default_factory=ProviderConfig) # NVIDIA NIM (nvapi- keys)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_api_type_scope(self) -> "ProvidersConfig":
|
||||
for name in self.__class__.model_fields:
|
||||
if name == "openai":
|
||||
continue
|
||||
provider = getattr(self, name, None)
|
||||
if isinstance(provider, ProviderConfig) and provider.api_type != "auto":
|
||||
raise ValueError("providers.<name>.api_type is only supported for providers.openai")
|
||||
return self
|
||||
|
||||
|
||||
class HeartbeatConfig(Base):
|
||||
"""Heartbeat service configuration."""
|
||||
|
||||
@@ -98,7 +98,7 @@ def _make_provider_core(
|
||||
extra_headers=p.extra_headers if p else None,
|
||||
spec=spec,
|
||||
extra_body=p.extra_body if p else None,
|
||||
api_type=p.api_type if p else "auto",
|
||||
api_type=p.api_type if p and provider_name == "openai" else "auto",
|
||||
)
|
||||
|
||||
provider.generation = resolved.to_generation_settings()
|
||||
|
||||
@@ -274,6 +274,47 @@ def _deep_merge(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any
|
||||
return merged
|
||||
|
||||
|
||||
def _merge_unique_list(base: Any, override: Any) -> Any:
|
||||
"""Append list values while preserving order and removing duplicates."""
|
||||
if not isinstance(base, list) or not isinstance(override, list):
|
||||
return override
|
||||
result: list[Any] = []
|
||||
seen: set[str] = set()
|
||||
for value in [*base, *override]:
|
||||
try:
|
||||
key = json.dumps(value, sort_keys=True, ensure_ascii=False)
|
||||
except Exception:
|
||||
key = repr(value)
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
result.append(value)
|
||||
return result
|
||||
|
||||
|
||||
def _merge_responses_extra_body(
|
||||
body: dict[str, Any],
|
||||
extra_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Merge configured Responses API body fields without clobbering tools."""
|
||||
reserved = {"include", "tools"}
|
||||
regular_extra = {key: value for key, value in extra_body.items() if key not in reserved}
|
||||
merged = _deep_merge(body, regular_extra)
|
||||
|
||||
if "include" in extra_body:
|
||||
merged["include"] = _merge_unique_list(body.get("include"), extra_body["include"])
|
||||
|
||||
if "tools" in extra_body:
|
||||
current_tools = body.get("tools")
|
||||
configured_tools = extra_body["tools"]
|
||||
if isinstance(current_tools, list) and isinstance(configured_tools, list):
|
||||
merged["tools"] = [*current_tools, *configured_tools]
|
||||
else:
|
||||
merged["tools"] = configured_tools
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
class OpenAICompatProvider(LLMProvider):
|
||||
"""Unified provider for all OpenAI-compatible APIs.
|
||||
|
||||
@@ -296,7 +337,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
self.extra_headers = extra_headers or {}
|
||||
self._spec = spec
|
||||
self._extra_body = extra_body or {}
|
||||
self._api_type = api_type
|
||||
self._api_type = api_type if spec and spec.name == "openai" else "auto"
|
||||
|
||||
if api_key and spec and spec.env_key:
|
||||
self._setup_env(api_key, api_base)
|
||||
@@ -697,12 +738,9 @@ class OpenAICompatProvider(LLMProvider):
|
||||
if self._api_type == "chat_completions":
|
||||
return False
|
||||
if self._spec and self._spec.name not in ("openai", "github_copilot"):
|
||||
if self._api_type != "responses":
|
||||
return False
|
||||
return False
|
||||
if self._api_type == "responses":
|
||||
return self._responses_circuit_allows_probe(model, reasoning_effort)
|
||||
if self._spec and self._spec.name not in ("openai", "github_copilot"):
|
||||
return False
|
||||
if self._spec is None or self._spec.name != "github_copilot":
|
||||
if not _is_direct_openai_base(self._effective_base):
|
||||
return False
|
||||
@@ -815,6 +853,10 @@ class OpenAICompatProvider(LLMProvider):
|
||||
body["tools"] = convert_tools(tools)
|
||||
body["tool_choice"] = tool_choice or "auto"
|
||||
|
||||
extra_body = getattr(self, "_extra_body", {})
|
||||
if extra_body:
|
||||
body = _merge_responses_extra_body(body, extra_body)
|
||||
|
||||
return body
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -181,18 +181,18 @@ def settings_payload(*, requires_restart: bool = False) -> dict[str, Any]:
|
||||
provider_config = getattr(config.providers, spec.name, None)
|
||||
if provider_config is None or spec.is_oauth:
|
||||
continue
|
||||
providers.append(
|
||||
{
|
||||
"name": spec.name,
|
||||
"label": spec.label,
|
||||
"configured": _provider_configured_for_settings(spec, provider_config),
|
||||
"api_key_required": _provider_requires_api_key(spec),
|
||||
"api_key_hint": _mask_secret_hint(provider_config.api_key),
|
||||
"api_base": provider_config.api_base,
|
||||
"default_api_base": spec.default_api_base or None,
|
||||
"api_type": provider_config.api_type,
|
||||
}
|
||||
)
|
||||
row = {
|
||||
"name": spec.name,
|
||||
"label": spec.label,
|
||||
"configured": _provider_configured_for_settings(spec, provider_config),
|
||||
"api_key_required": _provider_requires_api_key(spec),
|
||||
"api_key_hint": _mask_secret_hint(provider_config.api_key),
|
||||
"api_base": provider_config.api_base,
|
||||
"default_api_base": spec.default_api_base or None,
|
||||
}
|
||||
if spec.name == "openai":
|
||||
row["api_type"] = provider_config.api_type
|
||||
providers.append(row)
|
||||
|
||||
search_config = config.tools.web.search
|
||||
image_config = config.tools.image_generation
|
||||
@@ -472,15 +472,16 @@ def update_provider_settings(query: QueryParams) -> dict[str, Any]:
|
||||
provider_config.api_base = api_base
|
||||
changed = True
|
||||
|
||||
if "api_type" in query or "apiType" in query:
|
||||
api_type = (_query_first_alias(query, "api_type", "apiType") or "").strip()
|
||||
try:
|
||||
parsed_api_type = type(provider_config)(api_type=api_type).api_type
|
||||
except Exception:
|
||||
raise WebUISettingsError("api_type must be auto, chat_completions, or responses") from None
|
||||
if provider_config.api_type != parsed_api_type:
|
||||
provider_config.api_type = parsed_api_type
|
||||
changed = True
|
||||
if "api_type" in query:
|
||||
if spec.name == "openai":
|
||||
api_type = (_query_first(query, "api_type") or "").strip()
|
||||
try:
|
||||
parsed_api_type = type(provider_config)(api_type=api_type).api_type
|
||||
except Exception:
|
||||
raise WebUISettingsError("api_type must be auto, chat_completions, or responses") from None
|
||||
if provider_config.api_type != parsed_api_type:
|
||||
provider_config.api_type = parsed_api_type
|
||||
changed = True
|
||||
|
||||
if changed:
|
||||
save_config(config)
|
||||
|
||||
Reference in New Issue
Block a user