fix: 错误响应读取移至连接关闭前,ModelMapper 缓存改为模块级共享

1. chat_handler_base/cli_stream_mixin: 将 _extract_error_text 提前到
   response_ctx.__aexit__ 之前执行,避免连接关闭后无法读取错误响应体
2. ModelMapperMiddleware: 实例级缓存改为模块级共享缓存,消除多实例
   间缓存不一致问题;缓存失效服务改为直接调用静态方法
This commit is contained in:
fawney19
2026-03-08 22:18:10 +08:00
parent 8be9601963
commit f5f7a23bb0
4 changed files with 36 additions and 46 deletions

View File

@@ -1287,6 +1287,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
response_ctx = None
continue
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
error_text = await ChatSyncExecutor(self)._extract_error_text(e)
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
@@ -1294,10 +1298,6 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
pass
finally:
response_ctx = None
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
error_text = await ChatSyncExecutor(self)._extract_error_text(e)
logger.error(
f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}"
)

View File

@@ -800,6 +800,8 @@ class CliStreamMixin:
response_ctx = None
continue
error_text = await self._extract_error_text(e)
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
@@ -807,8 +809,6 @@ class CliStreamMixin:
pass
finally:
response_ctx = None
error_text = await self._extract_error_text(e)
logger.error(
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
)

View File

@@ -15,12 +15,11 @@ class CacheInvalidationService:
"""缓存失效服务"""
def __init__(self) -> None:
self._model_mappers = []
pass
def register_model_mapper(self, model_mapper: Any) -> None:
"""注册 ModelMapper 实例"""
if model_mapper not in self._model_mappers:
self._model_mappers.append(model_mapper)
"""注册 ModelMapper 实例(已弃用,保留兼容性,不执行任何操作)"""
pass
async def on_global_model_changed(
self, model_name: str, global_model_id: str | None = None
@@ -39,9 +38,10 @@ class CacheInvalidationService:
clear_regex_cache()
# 2. 清空 ModelMapper 缓存
for mapper in self._model_mappers:
mapper.clear_cache()
# 2. 清空 ModelMapper 共享缓存
from src.services.model.mapper import ModelMapperMiddleware
ModelMapperMiddleware.clear_cache()
# 3. 清空 ModelCacheService 缓存
from src.services.cache.model_cache import ModelCacheService
@@ -88,13 +88,15 @@ class CacheInvalidationService:
def _refresh_provider_cache(self, provider_id: str) -> None:
"""刷新指定 Provider 的 ModelMapper 缓存"""
for mapper in self._model_mappers:
mapper.refresh_cache(provider_id)
from src.services.model.mapper import ModelMapperMiddleware
ModelMapperMiddleware.refresh_cache(provider_id)
def clear_all_caches(self) -> None:
"""清空所有缓存"""
for mapper in self._model_mappers:
mapper.clear_cache()
from src.services.model.mapper import ModelMapperMiddleware
ModelMapperMiddleware.clear_cache()
# 全局单例

View File

@@ -13,6 +13,9 @@ from src.models.claude import ClaudeMessagesRequest
from src.models.database import GlobalModel, Model, Provider, ProviderEndpoint
from src.services.cache.model_cache import ModelCacheService
# 模块级共享缓存,所有 ModelMapperMiddleware 实例共用
_shared_cache = SyncLRUCache(max_size=1000, ttl=300)
class ModelMapperMiddleware:
"""
@@ -20,29 +23,14 @@ class ModelMapperMiddleware:
负责将用户请求的模型名映射到提供商的实际模型名
"""
def __init__(self, db: Session, cache_max_size: int = 1000, cache_ttl: int = 300):
def __init__(self, db: Session):
"""
初始化模型映射中间件
Args:
db: 数据库会话
cache_max_size: 缓存最大容量(默认 1000
cache_ttl: 缓存过期时间(秒,默认 300
"""
self.db = db
self._cache = SyncLRUCache(max_size=cache_max_size, ttl=cache_ttl)
logger.debug(f"[ModelMapper] 初始化max_size={cache_max_size}, ttl={cache_ttl}s")
# 注册到缓存失效服务
try:
from src.services.cache.invalidation import get_cache_invalidation_service
cache_service = get_cache_invalidation_service()
cache_service.register_model_mapper(self)
logger.debug("[ModelMapper] 已注册到缓存失效服务")
except Exception as e:
logger.warning(f"[ModelMapper] 注册缓存失效服务失败: {e}")
async def apply_mapping(
self, request: ClaudeMessagesRequest, provider: Provider
@@ -104,8 +92,8 @@ class ModelMapperMiddleware:
# 检查缓存(使用规范化后的名称)
cache_key = f"{provider_id}:{normalized_name}"
if cache_key in self._cache:
return self._cache[cache_key]
if cache_key in _shared_cache:
return _shared_cache[cache_key]
mapping = None
@@ -113,7 +101,7 @@ class ModelMapperMiddleware:
if not global_model or not global_model.is_active:
logger.debug(f"GlobalModel not found or inactive: {normalized_name}")
self._cache[cache_key] = None
_shared_cache[cache_key] = None
return None
# 步骤 2: 查找该 Provider 是否有实现这个 GlobalModel 的 Model使用缓存
@@ -140,7 +128,7 @@ class ModelMapperMiddleware:
)
# 缓存结果
self._cache[cache_key] = mapping
_shared_cache[cache_key] = mapping
return mapping
@@ -221,12 +209,14 @@ class ModelMapperMiddleware:
return True, None
def clear_cache(self) -> None:
"""清空缓存"""
self._cache.clear()
@staticmethod
def clear_cache() -> None:
"""清空共享缓存"""
_shared_cache.clear()
logger.debug("Model mapping cache cleared")
def refresh_cache(self, provider_id: str | None = None) -> None:
@staticmethod
def refresh_cache(provider_id: str | None = None) -> None:
"""
刷新缓存
@@ -234,16 +224,14 @@ class ModelMapperMiddleware:
provider_id: 如果指定,只刷新该提供商的缓存 (UUID)
"""
if provider_id:
# 清除特定提供商的缓存
keys_to_remove = [
key for key in self._cache.keys() if key.startswith(f"{provider_id}:")
key for key in _shared_cache.keys() if key.startswith(f"{provider_id}:")
]
for key in keys_to_remove:
del self._cache[key]
del _shared_cache[key]
logger.debug(f"Refreshed cache for provider {provider_id}")
else:
# 清空所有缓存
self.clear_cache()
ModelMapperMiddleware.clear_cache()
class ModelRoutingMiddleware: