mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10: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
|
||||
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}"
|
||||
)
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
24
src/services/cache/invalidation.py
vendored
24
src/services/cache/invalidation.py
vendored
@@ -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()
|
||||
|
||||
|
||||
# 全局单例
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user