mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
fix: 错误响应读取移至连接关闭前,ModelMapper 缓存改为模块级共享
1. chat_handler_base/cli_stream_mixin: 将 _extract_error_text 提前到 response_ctx.__aexit__ 之前执行,避免连接关闭后无法读取错误响应体 2. ModelMapperMiddleware: 实例级缓存改为模块级共享缓存,消除多实例 间缓存不一致问题;缓存失效服务改为直接调用静态方法
This commit is contained in:
@@ -1287,6 +1287,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
response_ctx = None
|
response_ctx = None
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
|
||||||
|
|
||||||
|
error_text = await ChatSyncExecutor(self)._extract_error_text(e)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if response_ctx is not None:
|
if response_ctx is not None:
|
||||||
await response_ctx.__aexit__(None, None, None)
|
await response_ctx.__aexit__(None, None, None)
|
||||||
@@ -1294,10 +1298,6 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
pass
|
pass
|
||||||
finally:
|
finally:
|
||||||
response_ctx = None
|
response_ctx = None
|
||||||
|
|
||||||
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
|
|
||||||
|
|
||||||
error_text = await ChatSyncExecutor(self)._extract_error_text(e)
|
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}"
|
f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -800,6 +800,8 @@ class CliStreamMixin:
|
|||||||
response_ctx = None
|
response_ctx = None
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
error_text = await self._extract_error_text(e)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if response_ctx is not None:
|
if response_ctx is not None:
|
||||||
await response_ctx.__aexit__(None, None, None)
|
await response_ctx.__aexit__(None, None, None)
|
||||||
@@ -807,8 +809,6 @@ class CliStreamMixin:
|
|||||||
pass
|
pass
|
||||||
finally:
|
finally:
|
||||||
response_ctx = None
|
response_ctx = None
|
||||||
|
|
||||||
error_text = await self._extract_error_text(e)
|
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
|
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
|
||||||
)
|
)
|
||||||
|
|||||||
24
src/services/cache/invalidation.py
vendored
24
src/services/cache/invalidation.py
vendored
@@ -15,12 +15,11 @@ class CacheInvalidationService:
|
|||||||
"""缓存失效服务"""
|
"""缓存失效服务"""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._model_mappers = []
|
pass
|
||||||
|
|
||||||
def register_model_mapper(self, model_mapper: Any) -> None:
|
def register_model_mapper(self, model_mapper: Any) -> None:
|
||||||
"""注册 ModelMapper 实例"""
|
"""注册 ModelMapper 实例(已弃用,保留兼容性,不执行任何操作)"""
|
||||||
if model_mapper not in self._model_mappers:
|
pass
|
||||||
self._model_mappers.append(model_mapper)
|
|
||||||
|
|
||||||
async def on_global_model_changed(
|
async def on_global_model_changed(
|
||||||
self, model_name: str, global_model_id: str | None = None
|
self, model_name: str, global_model_id: str | None = None
|
||||||
@@ -39,9 +38,10 @@ class CacheInvalidationService:
|
|||||||
|
|
||||||
clear_regex_cache()
|
clear_regex_cache()
|
||||||
|
|
||||||
# 2. 清空 ModelMapper 缓存
|
# 2. 清空 ModelMapper 共享缓存
|
||||||
for mapper in self._model_mappers:
|
from src.services.model.mapper import ModelMapperMiddleware
|
||||||
mapper.clear_cache()
|
|
||||||
|
ModelMapperMiddleware.clear_cache()
|
||||||
|
|
||||||
# 3. 清空 ModelCacheService 缓存
|
# 3. 清空 ModelCacheService 缓存
|
||||||
from src.services.cache.model_cache import ModelCacheService
|
from src.services.cache.model_cache import ModelCacheService
|
||||||
@@ -88,13 +88,15 @@ class CacheInvalidationService:
|
|||||||
|
|
||||||
def _refresh_provider_cache(self, provider_id: str) -> None:
|
def _refresh_provider_cache(self, provider_id: str) -> None:
|
||||||
"""刷新指定 Provider 的 ModelMapper 缓存"""
|
"""刷新指定 Provider 的 ModelMapper 缓存"""
|
||||||
for mapper in self._model_mappers:
|
from src.services.model.mapper import ModelMapperMiddleware
|
||||||
mapper.refresh_cache(provider_id)
|
|
||||||
|
ModelMapperMiddleware.refresh_cache(provider_id)
|
||||||
|
|
||||||
def clear_all_caches(self) -> None:
|
def clear_all_caches(self) -> None:
|
||||||
"""清空所有缓存"""
|
"""清空所有缓存"""
|
||||||
for mapper in self._model_mappers:
|
from src.services.model.mapper import ModelMapperMiddleware
|
||||||
mapper.clear_cache()
|
|
||||||
|
ModelMapperMiddleware.clear_cache()
|
||||||
|
|
||||||
|
|
||||||
# 全局单例
|
# 全局单例
|
||||||
|
|||||||
@@ -13,6 +13,9 @@ from src.models.claude import ClaudeMessagesRequest
|
|||||||
from src.models.database import GlobalModel, Model, Provider, ProviderEndpoint
|
from src.models.database import GlobalModel, Model, Provider, ProviderEndpoint
|
||||||
from src.services.cache.model_cache import ModelCacheService
|
from src.services.cache.model_cache import ModelCacheService
|
||||||
|
|
||||||
|
# 模块级共享缓存,所有 ModelMapperMiddleware 实例共用
|
||||||
|
_shared_cache = SyncLRUCache(max_size=1000, ttl=300)
|
||||||
|
|
||||||
|
|
||||||
class ModelMapperMiddleware:
|
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:
|
Args:
|
||||||
db: 数据库会话
|
db: 数据库会话
|
||||||
cache_max_size: 缓存最大容量(默认 1000)
|
|
||||||
cache_ttl: 缓存过期时间(秒,默认 300)
|
|
||||||
"""
|
"""
|
||||||
self.db = db
|
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(
|
async def apply_mapping(
|
||||||
self, request: ClaudeMessagesRequest, provider: Provider
|
self, request: ClaudeMessagesRequest, provider: Provider
|
||||||
@@ -104,8 +92,8 @@ class ModelMapperMiddleware:
|
|||||||
|
|
||||||
# 检查缓存(使用规范化后的名称)
|
# 检查缓存(使用规范化后的名称)
|
||||||
cache_key = f"{provider_id}:{normalized_name}"
|
cache_key = f"{provider_id}:{normalized_name}"
|
||||||
if cache_key in self._cache:
|
if cache_key in _shared_cache:
|
||||||
return self._cache[cache_key]
|
return _shared_cache[cache_key]
|
||||||
|
|
||||||
mapping = None
|
mapping = None
|
||||||
|
|
||||||
@@ -113,7 +101,7 @@ class ModelMapperMiddleware:
|
|||||||
|
|
||||||
if not global_model or not global_model.is_active:
|
if not global_model or not global_model.is_active:
|
||||||
logger.debug(f"GlobalModel not found or inactive: {normalized_name}")
|
logger.debug(f"GlobalModel not found or inactive: {normalized_name}")
|
||||||
self._cache[cache_key] = None
|
_shared_cache[cache_key] = None
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# 步骤 2: 查找该 Provider 是否有实现这个 GlobalModel 的 Model(使用缓存)
|
# 步骤 2: 查找该 Provider 是否有实现这个 GlobalModel 的 Model(使用缓存)
|
||||||
@@ -140,7 +128,7 @@ class ModelMapperMiddleware:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 缓存结果
|
# 缓存结果
|
||||||
self._cache[cache_key] = mapping
|
_shared_cache[cache_key] = mapping
|
||||||
|
|
||||||
return mapping
|
return mapping
|
||||||
|
|
||||||
@@ -221,12 +209,14 @@ class ModelMapperMiddleware:
|
|||||||
|
|
||||||
return True, None
|
return True, None
|
||||||
|
|
||||||
def clear_cache(self) -> None:
|
@staticmethod
|
||||||
"""清空缓存"""
|
def clear_cache() -> None:
|
||||||
self._cache.clear()
|
"""清空共享缓存"""
|
||||||
|
_shared_cache.clear()
|
||||||
logger.debug("Model mapping cache cleared")
|
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)
|
provider_id: 如果指定,只刷新该提供商的缓存 (UUID)
|
||||||
"""
|
"""
|
||||||
if provider_id:
|
if provider_id:
|
||||||
# 清除特定提供商的缓存
|
|
||||||
keys_to_remove = [
|
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:
|
for key in keys_to_remove:
|
||||||
del self._cache[key]
|
del _shared_cache[key]
|
||||||
logger.debug(f"Refreshed cache for provider {provider_id}")
|
logger.debug(f"Refreshed cache for provider {provider_id}")
|
||||||
else:
|
else:
|
||||||
# 清空所有缓存
|
ModelMapperMiddleware.clear_cache()
|
||||||
self.clear_cache()
|
|
||||||
|
|
||||||
|
|
||||||
class ModelRoutingMiddleware:
|
class ModelRoutingMiddleware:
|
||||||
|
|||||||
Reference in New Issue
Block a user