refactor: 优化 kiro token 刷新逻辑和 build_all_format_configs

- kiro auth: 拆分无缓存 token 和占位符 key 的判断逻辑,避免不必要的解密操作
- kiro auth: 简化 effective_token 获取逻辑
- build_all_format_configs: 只对实际配置了端点的格式构建请求配置,
  不再用某个端点的 base_url 尝试其他未配置的格式
- get_adapter_for_format: 简化为单行表达式
- 修复 f-string 日志为 loguru 风格占位符
- 新增 build_all_format_configs 单元测试
This commit is contained in:
fawney19
2026-02-09 10:11:38 +08:00
parent 1b8e73bb5c
commit 8673ed1459
3 changed files with 106 additions and 59 deletions

View File

@@ -1047,8 +1047,9 @@ async def get_provider_auth(
# Kiro 特殊处理:如果没有缓存的 access_token 或 key.api_key 是占位符,强制刷新
if provider_type == "kiro" and not should_refresh:
decrypted_api_key = crypto_service.decrypt(key.api_key)
if not cached_access_token or decrypted_api_key == "__placeholder__":
if not cached_access_token:
should_refresh = True
elif crypto_service.decrypt(key.api_key) == "__placeholder__":
should_refresh = True
if should_refresh and refresh_token and provider_type:
@@ -1079,16 +1080,11 @@ async def get_provider_auth(
# 获取最终使用的 access_token
# Kiro 优先使用 token_meta 中缓存的 access_token刷新后会更新到 token_meta
effective_access_token: str
if provider_type == "kiro":
# Kiro: 优先使用 token_meta 中的 access_token回退到 key.api_key
cached_token = str(token_meta.get("access_token") or "").strip()
if cached_token:
effective_access_token = cached_token
else:
effective_access_token = crypto_service.decrypt(key.api_key)
refreshed_token = str(token_meta.get("access_token") or "").strip()
effective_token = refreshed_token or crypto_service.decrypt(key.api_key)
else:
effective_access_token = crypto_service.decrypt(key.api_key)
effective_token = crypto_service.decrypt(key.api_key)
decrypted_auth_config: dict[str, Any] | None = None
if isinstance(token_meta, dict) and token_meta:
@@ -1096,7 +1092,7 @@ async def get_provider_auth(
return ProviderAuthInfo(
auth_header="Authorization",
auth_value=f"Bearer {effective_access_token}",
auth_value=f"Bearer {effective_token}",
decrypted_auth_config=decrypted_auth_config,
)
if auth_type == "vertex_ai":

View File

@@ -140,13 +140,7 @@ def get_adapter_for_format(api_format: str) -> type | None:
from src.api.handlers.base.chat_adapter_base import get_adapter_class
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class
adapter_class = get_adapter_class(api_format)
if adapter_class:
return adapter_class
cli_adapter_class = get_cli_adapter_class(api_format)
if cli_adapter_class:
return cli_adapter_class
return None
return get_adapter_class(api_format) or get_cli_adapter_class(api_format)
def build_all_format_configs(
@@ -156,8 +150,8 @@ def build_all_format_configs(
"""
构建所有 API 格式的端点配置
从基础 endpoint signature 列表构建配置,如果该格式有专门的端点配置则使用
否则使用基础端点的 base_url 尝试。
只对实际配置了端点的格式构建请求配置,不同端点的 base_url 可能不同
不应使用某个端点的 base_url 尝试其他格式
Args:
api_key_value: 解密后的 API Key
@@ -169,44 +163,17 @@ def build_all_format_configs(
if not format_to_endpoint:
return []
# 获取任意一个端点的 base_url 作为基础(用于尝试所有格式)
# 优先使用 OPENAI 格式的端点,因为它最通用
base_endpoint = (
format_to_endpoint.get("openai:chat")
or format_to_endpoint.get("claude:chat")
or format_to_endpoint.get("gemini:chat")
or next(iter(format_to_endpoint.values()))
)
base_url = base_endpoint.base_url
extra_headers = get_extra_headers_from_endpoint(base_endpoint)
# 只对基础 API 格式获取模型CLI 格式使用相同的上游 API
endpoint_configs: list[dict] = []
for fmt in MODEL_FETCH_FORMATS:
fmt_value = fmt
# 如果该格式有专门的端点配置,使用其 base_url 和 headers
if fmt_value in format_to_endpoint:
ep = format_to_endpoint[fmt_value]
endpoint_configs.append(
{
"api_key": api_key_value,
"base_url": ep.base_url,
"api_format": fmt_value,
"extra_headers": get_extra_headers_from_endpoint(ep),
}
)
else:
# 没有专门配置,使用基础端点的 base_url 尝试
endpoint_configs.append(
{
"api_key": api_key_value,
"base_url": base_url,
"api_format": fmt_value,
"extra_headers": extra_headers,
}
)
return endpoint_configs
return [
{
"api_key": api_key_value,
"base_url": ep.base_url,
"api_format": fmt,
"extra_headers": get_extra_headers_from_endpoint(ep),
}
for fmt in MODEL_FETCH_FORMATS
if (ep := format_to_endpoint.get(fmt)) is not None
]
async def fetch_models_from_endpoints(
@@ -255,10 +222,10 @@ async def fetch_models_from_endpoints(
success = error is None
return models, error, success
except httpx.TimeoutException:
logger.warning(f"获取 {api_format} 模型超时")
logger.warning("获取 {} 模型超时", api_format)
return [], f"{api_format}: timeout", False
except Exception:
logger.exception(f"获取 {api_format} 模型出错")
logger.exception("获取 {} 模型出错", api_format)
return [], f"{api_format}: error", False
async with httpx.AsyncClient(timeout=timeout, verify=get_ssl_context()) as client: