Files
daily_stock_analysis/tests/test_llm_param_recovery.py
mumu 4aad40cdd4 fix: 增强 LLM 参数适配层 (#1317)
* fix: add llm generation parameter adaptation

* fix(review-feedback-1317): 补充官方来源链接或明确写成基于当前兼容测试的保守策略,并说明验证环境

* fix(review-feedback-1317): 补充官方文档/公告链接、当前依赖/运行时验证依据,以及 provider override 的预期扩展方式

* feat: 增强 LLM 参数错误自愈 (#1318)

* fix: add llm generation parameter adaptation

* fix: add llm parameter error recovery

* fix(review-feedback-1318): Scope recovery cache to the endpoint and making recovery caching

* fix(review-feedback-1318): Avoid forcing 1.0 for all default-only temperature errors

* fix(review-feedback-1318): 基于目标分支解决冲突,并确保解决后 diff 仍只包含本 PR 声称的 LLM 参数自愈增量

* fix(review-feedback-1318): Scope legacy Router recoveries to their endpoint

* fix(review-feedback-1318): 解决冲突并基于解决后的最终 diff 重新确认测试结果

* fix(review-feedback-1318): 基于目标 base 解决冲突,再重新确认 diff 与 CI

* fix(review-feedback-1318): 解决冲突后重新确认 diff 与 CI

* fix(review-feedback-1318): 补充覆盖 Analyzer legacy multi-key Router 的回归测试
2026-05-16 21:31:28 +08:00

203 lines
6.0 KiB
Python

# -*- coding: utf-8 -*-
"""Tests for LiteLLM generation-parameter recovery."""
from src.llm.errors import (
call_litellm_with_param_recovery,
classify_litellm_generation_param_error,
)
from src.llm.generation_params import (
apply_litellm_generation_params,
clear_litellm_generation_param_recovery_cache,
)
def test_temperature_default_only_error_sets_temperature_to_one() -> None:
recovery = classify_litellm_generation_param_error(
RuntimeError(
"Unsupported value: 'temperature' does not support 0.7 with this model. "
"Only the default (1.0) value is supported."
)
)
assert recovery is not None
assert recovery.set_params == {"temperature": 1.0}
assert recovery.omit_params == ()
def test_temperature_default_only_error_uses_named_default_value() -> None:
recovery = classify_litellm_generation_param_error(
RuntimeError(
"Unsupported value: 'temperature' does not support 1.0 with this model. "
"Only `0.6` is allowed."
)
)
assert recovery is not None
assert recovery.set_params == {"temperature": 0.6}
assert recovery.omit_params == ()
def test_temperature_default_only_error_without_named_value_omits_temperature() -> None:
recovery = classify_litellm_generation_param_error(
RuntimeError(
"Unsupported value: 'temperature' does not support 0.7 with this model. "
"Only the default value is supported."
)
)
assert recovery is not None
assert recovery.set_params == {}
assert recovery.omit_params == ("temperature",)
def test_unsupported_temperature_error_retries_once_and_caches_recovery() -> None:
clear_litellm_generation_param_recovery_cache()
calls = []
def _call(kwargs):
calls.append(dict(kwargs))
if len(calls) == 1:
raise RuntimeError("Unsupported parameter: temperature is not supported")
return "ok"
result = call_litellm_with_param_recovery(
_call,
model="openai/custom-temp-locked",
call_kwargs={
"model": "openai/custom-temp-locked",
"messages": [],
"temperature": 0.7,
},
)
future_kwargs = apply_litellm_generation_params(
{"model": "openai/custom-temp-locked", "messages": []},
"openai/custom-temp-locked",
0.7,
)
assert result == "ok"
assert calls[0]["temperature"] == 0.7
assert "temperature" not in calls[1]
assert "temperature" not in future_kwargs
def test_recovery_cache_is_scoped_to_api_base() -> None:
clear_litellm_generation_param_recovery_cache()
calls = []
def _call(kwargs):
calls.append(dict(kwargs))
if len(calls) == 1:
raise RuntimeError("Unsupported parameter: temperature is not supported")
return "ok"
result = call_litellm_with_param_recovery(
_call,
model="openai/shared-model",
call_kwargs={
"model": "openai/shared-model",
"messages": [],
"api_base": "https://strict.example/v1",
"temperature": 0.7,
},
)
strict_kwargs = apply_litellm_generation_params(
{"model": "openai/shared-model", "messages": [], "api_base": "https://strict.example/v1"},
"openai/shared-model",
0.7,
)
flexible_kwargs = apply_litellm_generation_params(
{"model": "openai/shared-model", "messages": [], "api_base": "https://flex.example/v1"},
"openai/shared-model",
0.7,
)
assert result == "ok"
assert "temperature" not in strict_kwargs
assert flexible_kwargs["temperature"] == 0.7
def test_recovery_cache_skips_ambiguous_router_endpoints() -> None:
clear_litellm_generation_param_recovery_cache()
model_list = [
{
"model_name": "openai/shared-model",
"litellm_params": {
"model": "openai/shared-model",
"api_base": "https://strict.example/v1",
},
},
{
"model_name": "openai/shared-model",
"litellm_params": {
"model": "openai/shared-model",
"api_base": "https://flex.example/v1",
},
},
]
calls = []
def _call(kwargs):
calls.append(dict(kwargs))
if len(calls) == 1:
raise RuntimeError("Unsupported parameter: temperature is not supported")
return "ok"
result = call_litellm_with_param_recovery(
_call,
model="openai/shared-model",
call_kwargs={"model": "openai/shared-model", "messages": [], "temperature": 0.7},
model_list=model_list,
)
future_kwargs = apply_litellm_generation_params(
{"model": "openai/shared-model", "messages": []},
"openai/shared-model",
0.7,
model_list=model_list,
)
assert result == "ok"
assert future_kwargs["temperature"] == 0.7
def test_streaming_retry_does_not_cache_before_stream_is_consumed() -> None:
clear_litellm_generation_param_recovery_cache()
calls = []
def _broken_stream():
raise RuntimeError("stream failed during iteration")
yield # pragma: no cover
def _call(kwargs):
calls.append(dict(kwargs))
if len(calls) == 1:
raise RuntimeError("Unsupported parameter: temperature is not supported")
return _broken_stream()
stream = call_litellm_with_param_recovery(
_call,
model="openai/stream-model",
call_kwargs={
"model": "openai/stream-model",
"messages": [],
"temperature": 0.7,
"stream": True,
},
cache_recovery=False,
)
try:
list(stream)
except RuntimeError:
pass
else: # pragma: no cover
raise AssertionError("stream should fail during iteration")
future_kwargs = apply_litellm_generation_params(
{"model": "openai/stream-model", "messages": []},
"openai/stream-model",
0.7,
)
assert "temperature" not in calls[1]
assert future_kwargs["temperature"] == 0.7