feat(ai): handle temperature incompatibility for gpt-5 and o-series models, add model picker UI
Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
+91
-49
@@ -11,6 +11,7 @@ See the Configuration Guide for full details on each provider's settings.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -19,6 +20,50 @@ from app.config import settings
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _resolve_temperature(model: str, requested: float) -> Optional[float]:
|
||||
"""Return a temperature value compatible with the given model, or ``None`` to omit it.
|
||||
|
||||
Certain model families have restrictions on the ``temperature`` parameter:
|
||||
|
||||
* **o-series reasoning models** (``o1``, ``o3``, ``o4``, …) – do not accept
|
||||
a ``temperature`` argument at all. Return ``None`` so callers can skip the
|
||||
parameter entirely.
|
||||
* **gpt-5 family** (``gpt-5``, ``gpt-5-nano``, ``gpt-5-codex``, …) – only
|
||||
``temperature=1`` is accepted; passing ``0`` raises a 400 error. Return
|
||||
``1`` and emit a debug log so the caller is aware of the coercion.
|
||||
* All other models – return the requested value unchanged.
|
||||
|
||||
The model string may include a provider prefix (e.g. ``openai/gpt-4o``);
|
||||
only the part after the last ``/`` is examined.
|
||||
|
||||
Args:
|
||||
model: Model identifier (may include a provider prefix).
|
||||
requested: The temperature the caller wants to use.
|
||||
|
||||
Returns:
|
||||
A compatible temperature float, or ``None`` if temperature should be
|
||||
omitted from the API call.
|
||||
"""
|
||||
bare = model.lower().split("/")[-1]
|
||||
|
||||
# o-series reasoning models (o1, o3, o4 …) do not support temperature
|
||||
if re.match(r"^o\d+(-|$)", bare):
|
||||
logger.debug("Dropping temperature parameter for reasoning model '%s' (not supported)", model)
|
||||
return None
|
||||
|
||||
# gpt-5 family only supports temperature=1
|
||||
if bare.startswith("gpt-5"):
|
||||
if requested != 1.0:
|
||||
logger.debug(
|
||||
"Coercing temperature from %s to 1 for model '%s' (only temperature=1 is supported)",
|
||||
requested,
|
||||
model,
|
||||
)
|
||||
return 1.0
|
||||
|
||||
return requested
|
||||
|
||||
|
||||
def _require_text_content(content: Optional[str]) -> str:
|
||||
"""Raise a clear error if the AI response contains no text content.
|
||||
|
||||
@@ -101,12 +146,12 @@ class OpenAIProvider(AIProvider):
|
||||
temperature: float = 0,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
completion = self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model, "messages": messages}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
completion = self._client.chat.completions.create(**call_kwargs)
|
||||
_content = completion.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -130,12 +175,12 @@ class AzureOpenAIProvider(AIProvider):
|
||||
temperature: float = 0,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
completion = self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model, "messages": messages}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
completion = self._client.chat.completions.create(**call_kwargs)
|
||||
_content = completion.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -161,13 +206,12 @@ class AnthropicProvider(AIProvider):
|
||||
import litellm
|
||||
|
||||
model_name = model if model.startswith("anthropic/") else f"anthropic/{model}"
|
||||
response = litellm.completion(
|
||||
model=model_name,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
api_key=self._api_key,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model_name, "messages": messages, "api_key": self._api_key}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
response = litellm.completion(**call_kwargs)
|
||||
_content = response.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -193,13 +237,12 @@ class GeminiProvider(AIProvider):
|
||||
import litellm
|
||||
|
||||
model_name = model if model.startswith("gemini/") else f"gemini/{model}"
|
||||
response = litellm.completion(
|
||||
model=model_name,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
api_key=self._api_key,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model_name, "messages": messages, "api_key": self._api_key}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
response = litellm.completion(**call_kwargs)
|
||||
_content = response.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -235,12 +278,12 @@ class OllamaProvider(AIProvider):
|
||||
temperature: float = 0,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
completion = self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model, "messages": messages}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
completion = self._client.chat.completions.create(**call_kwargs)
|
||||
_content = completion.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -269,12 +312,12 @@ class OpenRouterProvider(AIProvider):
|
||||
temperature: float = 0,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
completion = self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model, "messages": messages}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
completion = self._client.chat.completions.create(**call_kwargs)
|
||||
_content = completion.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -335,12 +378,12 @@ class PortkeyProvider(AIProvider):
|
||||
temperature: float = 0,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
completion = self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
**kwargs,
|
||||
)
|
||||
call_kwargs: Dict[str, Any] = {"model": model, "messages": messages}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
call_kwargs["temperature"] = safe_temp
|
||||
call_kwargs.update(kwargs)
|
||||
completion = self._client.chat.completions.create(**call_kwargs)
|
||||
_content = completion.choices[0].message.content
|
||||
return _require_text_content(_content)
|
||||
|
||||
@@ -371,11 +414,10 @@ class LiteLLMProvider(AIProvider):
|
||||
) -> str:
|
||||
import litellm
|
||||
|
||||
completion_kwargs: Dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
}
|
||||
completion_kwargs: Dict[str, Any] = {"model": model, "messages": messages}
|
||||
safe_temp = _resolve_temperature(model, temperature)
|
||||
if safe_temp is not None:
|
||||
completion_kwargs["temperature"] = safe_temp
|
||||
if self._api_key:
|
||||
completion_kwargs["api_key"] = self._api_key
|
||||
if self._api_base:
|
||||
|
||||
@@ -153,10 +153,33 @@ SETTING_METADATA = {
|
||||
"openai_model": {
|
||||
"category": "AI Services",
|
||||
"description": "Fallback model name used when AI_MODEL is not set (e.g. gpt-4o-mini)",
|
||||
"type": "string",
|
||||
"type": "model_picker",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
"suggested_models": [
|
||||
"gpt-4o",
|
||||
"gpt-4o-mini",
|
||||
"gpt-4-turbo",
|
||||
"gpt-4",
|
||||
"gpt-3.5-turbo",
|
||||
"o1",
|
||||
"o1-mini",
|
||||
"o3",
|
||||
"o3-mini",
|
||||
"gpt-5",
|
||||
"gpt-5-nano",
|
||||
"claude-3-5-sonnet-20241022",
|
||||
"claude-3-5-haiku-20241022",
|
||||
"claude-3-opus-20240229",
|
||||
"gemini-1.5-pro",
|
||||
"gemini-1.5-flash",
|
||||
"gemini-2.0-flash-exp",
|
||||
"llama3.2",
|
||||
"qwen2.5:7b",
|
||||
"phi3",
|
||||
"mistral",
|
||||
],
|
||||
},
|
||||
"ai_provider": {
|
||||
"category": "AI Services",
|
||||
@@ -170,10 +193,33 @@ SETTING_METADATA = {
|
||||
"ai_model": {
|
||||
"category": "AI Services",
|
||||
"description": "Model name for the selected provider (overrides OPENAI_MODEL). E.g. gpt-4o, claude-3-5-sonnet-20241022, gemini-1.5-pro, llama3.2",
|
||||
"type": "string",
|
||||
"type": "model_picker",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
"suggested_models": [
|
||||
"gpt-4o",
|
||||
"gpt-4o-mini",
|
||||
"gpt-4-turbo",
|
||||
"gpt-4",
|
||||
"gpt-3.5-turbo",
|
||||
"o1",
|
||||
"o1-mini",
|
||||
"o3",
|
||||
"o3-mini",
|
||||
"gpt-5",
|
||||
"gpt-5-nano",
|
||||
"claude-3-5-sonnet-20241022",
|
||||
"claude-3-5-haiku-20241022",
|
||||
"claude-3-opus-20240229",
|
||||
"gemini-1.5-pro",
|
||||
"gemini-1.5-flash",
|
||||
"gemini-2.0-flash-exp",
|
||||
"llama3.2",
|
||||
"qwen2.5:7b",
|
||||
"phi3",
|
||||
"mistral",
|
||||
],
|
||||
},
|
||||
"anthropic_api_key": {
|
||||
"category": "AI Services",
|
||||
|
||||
@@ -139,6 +139,29 @@
|
||||
Enable {{ setting.key.replace('_', ' ').title() }}
|
||||
</label>
|
||||
</div>
|
||||
{% elif setting.metadata.type == 'model_picker' %}
|
||||
<!-- Model Picker: free-text input with datalist of common models -->
|
||||
<div class="relative">
|
||||
<input
|
||||
type="text"
|
||||
id="{{ setting.key }}"
|
||||
name="{{ setting.key }}"
|
||||
list="{{ setting.key }}_models"
|
||||
x-model="formData['{{ setting.key }}']"
|
||||
class="setting-input w-full px-3 py-2 border border-gray-300 rounded-md shadow-sm focus:outline-none focus:ring-blue-500 focus:border-blue-500"
|
||||
placeholder="Select a common model or type a custom name…"
|
||||
autocomplete="off"
|
||||
/>
|
||||
<datalist id="{{ setting.key }}_models">
|
||||
{% for m in setting.metadata.suggested_models %}
|
||||
<option value="{{ m }}">{{ m }}</option>
|
||||
{% endfor %}
|
||||
</datalist>
|
||||
<p class="text-xs text-gray-400 mt-1">
|
||||
<i class="fas fa-info-circle mr-1"></i>
|
||||
Pick from the list or type any model name supported by your provider.
|
||||
</p>
|
||||
</div>
|
||||
{% elif setting.metadata.options %}
|
||||
<!-- Dropdown Select for fields with a fixed list of values -->
|
||||
<select
|
||||
|
||||
@@ -16,6 +16,7 @@ from app.utils.ai_provider import (
|
||||
OpenRouterProvider,
|
||||
PortkeyProvider,
|
||||
_require_text_content,
|
||||
_resolve_temperature,
|
||||
get_ai_provider,
|
||||
)
|
||||
|
||||
@@ -74,6 +75,77 @@ class TestRequireTextContent:
|
||||
_require_text_content(None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_temperature helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestResolveTemperature:
|
||||
"""Tests for the _resolve_temperature compatibility helper."""
|
||||
|
||||
def test_regular_model_returns_requested_temperature(self):
|
||||
"""Standard models return the temperature unchanged."""
|
||||
assert _resolve_temperature("gpt-4o", 0) == 0
|
||||
assert _resolve_temperature("gpt-4o-mini", 0.7) == 0.7
|
||||
|
||||
def test_gpt4_model_returns_requested_temperature(self):
|
||||
"""gpt-4 models are not gpt-5, so temperature is returned as-is."""
|
||||
assert _resolve_temperature("gpt-4-turbo", 0) == 0
|
||||
|
||||
def test_gpt5_model_forces_temperature_1(self):
|
||||
"""gpt-5 models only accept temperature=1; any other value is coerced."""
|
||||
assert _resolve_temperature("gpt-5", 0) == 1.0
|
||||
|
||||
def test_gpt5_nano_forces_temperature_1(self):
|
||||
"""gpt-5-nano (gpt-5 variant) gets temperature coerced to 1."""
|
||||
assert _resolve_temperature("gpt-5-nano", 0) == 1.0
|
||||
|
||||
def test_gpt5_codex_forces_temperature_1(self):
|
||||
"""gpt-5-codex (gpt-5 variant) gets temperature coerced to 1."""
|
||||
assert _resolve_temperature("gpt-5-codex", 0) == 1.0
|
||||
|
||||
def test_gpt5_already_at_1_unchanged(self):
|
||||
"""gpt-5 with temperature=1 returns 1 (no unnecessary log noise)."""
|
||||
assert _resolve_temperature("gpt-5", 1.0) == 1.0
|
||||
|
||||
def test_o1_returns_none(self):
|
||||
"""o1 reasoning model does not support temperature; None is returned."""
|
||||
assert _resolve_temperature("o1", 0) is None
|
||||
|
||||
def test_o1_mini_returns_none(self):
|
||||
"""o1-mini returns None (temperature not supported)."""
|
||||
assert _resolve_temperature("o1-mini", 0) is None
|
||||
|
||||
def test_o1_preview_returns_none(self):
|
||||
"""o1-preview returns None (temperature not supported)."""
|
||||
assert _resolve_temperature("o1-preview", 0) is None
|
||||
|
||||
def test_o3_returns_none(self):
|
||||
"""o3 reasoning model does not support temperature."""
|
||||
assert _resolve_temperature("o3", 0) is None
|
||||
|
||||
def test_o3_mini_returns_none(self):
|
||||
"""o3-mini returns None (temperature not supported)."""
|
||||
assert _resolve_temperature("o3-mini", 0) is None
|
||||
|
||||
def test_o4_mini_returns_none(self):
|
||||
"""o4-mini returns None (temperature not supported)."""
|
||||
assert _resolve_temperature("o4-mini", 0) is None
|
||||
|
||||
def test_provider_prefix_is_stripped(self):
|
||||
"""Provider prefix (e.g. 'openai/') is ignored when matching."""
|
||||
assert _resolve_temperature("openai/gpt-5-nano", 0) == 1.0
|
||||
assert _resolve_temperature("openai/o1-mini", 0) is None
|
||||
assert _resolve_temperature("openai/gpt-4o", 0) == 0
|
||||
|
||||
def test_model_names_are_case_insensitive(self):
|
||||
"""Model matching is case-insensitive."""
|
||||
assert _resolve_temperature("GPT-5", 0) == 1.0
|
||||
assert _resolve_temperature("O1", 0) is None
|
||||
assert _resolve_temperature("O3-Mini", 0) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Abstract base class
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user