mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-09 20:50:20 +08:00
refactor: 调度器迁移至独立模块,消除 services->api 反向依赖
- 将调度器相关模块从 src/services/cache/ 迁移到 src/services/scheduling/ - 下沉类型定义到 core 层: AccessRestrictions, ProviderAuthInfo, ParsedChunk/StreamStats, 视频工具函数 - 提取 thinking_cache 签名缓存到 core/api_format/conversion/ - 提取 provider 认证逻辑到 services/provider/auth - 提取遥测记录到 services/usage/telemetry - 提取 models 列表缓存到 services/cache/model_list_cache - 更新所有引用方的 import 路径及相关测试
This commit is contained in:
@@ -32,7 +32,7 @@ from src.models.database import (
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
|
||||
@@ -22,8 +22,8 @@ from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import ApiKey, User
|
||||
from src.services.cache.affinity_manager import get_affinity_manager
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, get_cache_aware_scheduler
|
||||
from src.services.scheduling.affinity_manager import get_affinity_manager
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, get_cache_aware_scheduler
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitoring: Cache"])
|
||||
@@ -1103,8 +1103,8 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
class AdminCacheConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.config.constants import ConcurrencyDefaults
|
||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
from src.services.scheduling.affinity_manager import CacheAffinityManager
|
||||
|
||||
# 获取动态预留管理器的配置
|
||||
reservation_manager = get_adaptive_reservation_manager()
|
||||
|
||||
@@ -643,7 +643,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
|
||||
if self.key in ("scheduling_mode", "provider_priority_mode"):
|
||||
try:
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.services.cache.aware_scheduler import get_cache_aware_scheduler
|
||||
from src.services.scheduling.aware_scheduler import get_cache_aware_scheduler
|
||||
|
||||
redis_client = get_redis_client_sync()
|
||||
# 从数据库读取两个调度配置的最新值,确保一致性
|
||||
|
||||
@@ -19,15 +19,18 @@ from sqlalchemy import tuple_
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.access_restrictions import AccessRestrictions
|
||||
from src.core.api_format.conversion.compatibility import is_format_compatible
|
||||
from src.core.cache_service import CacheService
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Model, Provider, ProviderEndpoint, User
|
||||
from src.models.database import Model, Provider, ProviderEndpoint
|
||||
from src.services.cache.model_list_cache import MODELS_LIST_CACHE_PREFIX as _CACHE_KEY_PREFIX
|
||||
from src.services.cache.model_list_cache import (
|
||||
invalidate_models_list_cache,
|
||||
)
|
||||
from src.services.model.availability import ModelAvailabilityQuery
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
|
||||
# 缓存 key 前缀
|
||||
_CACHE_KEY_PREFIX = "models:list"
|
||||
_CACHE_TTL = CacheTTL.MODEL # 300 秒
|
||||
|
||||
|
||||
@@ -70,21 +73,7 @@ async def _set_cached_models(
|
||||
logger.warning(f"[ModelsService] 缓存写入失败: {e}")
|
||||
|
||||
|
||||
async def invalidate_models_list_cache() -> None:
|
||||
"""
|
||||
清除所有 /v1/models 列表缓存
|
||||
|
||||
在模型创建、更新、删除时调用,确保模型列表实时更新
|
||||
"""
|
||||
try:
|
||||
# 使用通配符删除所有 models:list:* 缓存(包括多格式组合的 key)
|
||||
deleted = await CacheService.delete_pattern(f"{_CACHE_KEY_PREFIX}:*")
|
||||
if deleted > 0:
|
||||
logger.info(f"[ModelsService] 已清除 {deleted} 个 {_CACHE_KEY_PREFIX} 缓存")
|
||||
else:
|
||||
logger.debug(f"[ModelsService] 无 {_CACHE_KEY_PREFIX} 缓存需要清除")
|
||||
except Exception as e:
|
||||
logger.warning(f"[ModelsService] 清除缓存失败: {e}")
|
||||
__all__ = ["AccessRestrictions", "invalidate_models_list_cache", "ModelInfo"]
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -115,91 +104,7 @@ class ModelInfo:
|
||||
output_modalities: list[str] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccessRestrictions:
|
||||
"""API Key 或 User 的访问限制"""
|
||||
|
||||
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
|
||||
allowed_models: list[str] | None = None # 允许的模型名称列表
|
||||
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
|
||||
|
||||
@classmethod
|
||||
def from_api_key_and_user(cls, api_key: ApiKey | None, user: User | None) -> AccessRestrictions:
|
||||
"""
|
||||
从 API Key 和 User 合并访问限制
|
||||
|
||||
限制逻辑:
|
||||
- API Key 的限制优先于 User 的限制
|
||||
- 如果 API Key 有限制,使用 API Key 的限制
|
||||
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
|
||||
- 两者都无限制则返回空限制
|
||||
"""
|
||||
allowed_providers: list[str] | None = None
|
||||
allowed_models: list[str] | None = None
|
||||
allowed_api_formats: list[str] | None = None
|
||||
|
||||
# 优先使用 API Key 的限制
|
||||
if api_key:
|
||||
if api_key.allowed_providers is not None:
|
||||
allowed_providers = api_key.allowed_providers
|
||||
if api_key.allowed_models is not None:
|
||||
allowed_models = api_key.allowed_models
|
||||
if api_key.allowed_api_formats is not None:
|
||||
allowed_api_formats = api_key.allowed_api_formats
|
||||
|
||||
# 如果 API Key 没有限制,检查 User 的限制
|
||||
if user:
|
||||
if allowed_providers is None and user.allowed_providers is not None:
|
||||
allowed_providers = user.allowed_providers
|
||||
if allowed_models is None and user.allowed_models is not None:
|
||||
allowed_models = user.allowed_models
|
||||
if allowed_api_formats is None and user.allowed_api_formats is not None:
|
||||
allowed_api_formats = user.allowed_api_formats
|
||||
|
||||
return cls(
|
||||
allowed_providers=allowed_providers,
|
||||
allowed_models=allowed_models,
|
||||
allowed_api_formats=allowed_api_formats,
|
||||
)
|
||||
|
||||
def is_api_format_allowed(self, api_format: str) -> bool:
|
||||
"""
|
||||
检查 API 格式是否被允许
|
||||
|
||||
Args:
|
||||
api_format: endpoint signature(如 "openai:chat")
|
||||
|
||||
Returns:
|
||||
True 如果格式被允许,False 否则
|
||||
"""
|
||||
if self.allowed_api_formats is None:
|
||||
return True
|
||||
target = normalize_endpoint_signature(api_format)
|
||||
allowed = {normalize_endpoint_signature(f) for f in self.allowed_api_formats if f}
|
||||
return target in allowed
|
||||
|
||||
def is_model_allowed(self, model_id: str, provider_id: str) -> bool:
|
||||
"""
|
||||
检查模型是否被允许访问
|
||||
|
||||
Args:
|
||||
model_id: 模型 ID
|
||||
provider_id: Provider ID
|
||||
|
||||
Returns:
|
||||
True 如果模型被允许,False 否则
|
||||
"""
|
||||
# 检查 Provider 限制
|
||||
if self.allowed_providers is not None:
|
||||
if provider_id not in self.allowed_providers:
|
||||
return False
|
||||
|
||||
# 检查模型限制
|
||||
if self.allowed_models is not None:
|
||||
if model_id not in self.allowed_models:
|
||||
return False
|
||||
|
||||
return True
|
||||
# AccessRestrictions -- re-export from src.core.access_restrictions (see __all__)
|
||||
|
||||
|
||||
def _normalize_api_formats(
|
||||
|
||||
@@ -45,8 +45,8 @@ from sqlalchemy.orm import Session
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.system.audit import audit_service
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.usage.telemetry import MessageTelemetry # re-export
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
@@ -55,286 +55,8 @@ if TYPE_CHECKING:
|
||||
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
|
||||
|
||||
class MessageTelemetry:
|
||||
"""
|
||||
负责记录 Usage/Audit,避免处理器里重复代码。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, db: Session, user: Any, api_key: Any, request_id: str, client_ip: str
|
||||
) -> None:
|
||||
self.db = db
|
||||
self.user = user
|
||||
self.api_key = api_key
|
||||
self.request_id = request_id
|
||||
self.client_ip = client_ip
|
||||
|
||||
async def calculate_cost(
|
||||
self,
|
||||
provider: str,
|
||||
model: str,
|
||||
*,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
) -> float:
|
||||
input_price, output_price = await UsageService.get_model_price_async(
|
||||
self.db, provider, model
|
||||
)
|
||||
_, _, _, _, _, _, total_cost = UsageService.calculate_cost(
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
input_price,
|
||||
output_price,
|
||||
cache_creation_tokens,
|
||||
cache_read_tokens,
|
||||
*await UsageService.get_cache_prices_async(self.db, provider, model, input_price),
|
||||
)
|
||||
return total_cost
|
||||
|
||||
async def record_success(
|
||||
self,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
response_body: Any,
|
||||
response_headers: dict[str, Any],
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
is_stream: bool = False,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 时间指标
|
||||
first_byte_time_ms: int | None = None, # 首字时间/TTFB
|
||||
# Provider 侧追踪信息(用于记录真实成本)
|
||||
provider_id: str | None = None,
|
||||
provider_endpoint_id: str | None = None,
|
||||
provider_api_key_id: str | None = None,
|
||||
api_format: str | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: str | None = None, # 端点原生 API 格式
|
||||
has_format_conversion: bool = False, # 是否发生了格式转换
|
||||
# 模型映射信息
|
||||
target_model: str | None = None,
|
||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||
response_metadata: dict[str, Any] | None = None,
|
||||
# 请求元数据(用于性能与调试记录)
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
) -> float:
|
||||
metadata = response_metadata
|
||||
if request_metadata:
|
||||
merged = dict(request_metadata)
|
||||
if response_metadata:
|
||||
merged.setdefault("response", response_metadata)
|
||||
metadata = merged
|
||||
|
||||
usage = await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
provider=provider,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
request_type="chat",
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms, # 传递首字时间
|
||||
status_code=status_code,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers or {},
|
||||
response_headers=response_headers,
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body,
|
||||
request_id=self.request_id,
|
||||
# Provider 侧追踪信息(用于记录真实成本)
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=provider_endpoint_id,
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
# 模型映射信息
|
||||
target_model=target_model,
|
||||
# Provider 响应元数据/请求元数据
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
total_cost = float(getattr(usage, "total_cost_usd", 0.0) or 0.0)
|
||||
|
||||
if self.user and self.api_key:
|
||||
audit_service.log_api_request(
|
||||
db=self.db,
|
||||
user_id=self.user.id,
|
||||
api_key_id=self.api_key.id,
|
||||
request_id=self.request_id,
|
||||
model=model,
|
||||
provider=provider,
|
||||
success=True,
|
||||
ip_address=self.client_ip,
|
||||
status_code=status_code,
|
||||
input_tokens=getattr(usage, "input_tokens", input_tokens),
|
||||
output_tokens=getattr(usage, "output_tokens", output_tokens),
|
||||
cost_usd=total_cost,
|
||||
)
|
||||
|
||||
return total_cost
|
||||
|
||||
async def record_failure(
|
||||
self,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
# 模型映射信息
|
||||
target_model: str | None = None,
|
||||
# 请求元数据(用于性能与调试记录)
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录失败请求
|
||||
|
||||
注意:Provider 链路信息(provider_id, endpoint_id, key_id)不在此处记录,
|
||||
因为 RequestCandidate 表已经记录了完整的请求链路追踪信息。
|
||||
|
||||
Args:
|
||||
input_tokens: 预估输入 tokens(来自 message_start,用于中断请求的成本估算)
|
||||
output_tokens: 预估输出 tokens(来自已收到的内容)
|
||||
cache_creation_tokens: 缓存创建 tokens
|
||||
cache_read_tokens: 缓存读取 tokens
|
||||
response_body: 响应体(如果有部分响应)
|
||||
response_headers: 响应头(Provider 返回的原始响应头)
|
||||
client_response_headers: 返回给客户端的响应头
|
||||
target_model: 映射后的目标模型名(如果发生了映射)
|
||||
"""
|
||||
provider_name = provider or "unknown"
|
||||
if provider_name == "unknown":
|
||||
logger.warning(
|
||||
f"[Telemetry] Recording failure with unknown provider (request_id={self.request_id})"
|
||||
)
|
||||
|
||||
await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
request_type="chat",
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=error_message,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers or {},
|
||||
response_headers=response_headers or {},
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body or {"error": error_message},
|
||||
request_id=self.request_id,
|
||||
# 模型映射信息
|
||||
target_model=target_model,
|
||||
# 请求元数据
|
||||
metadata=request_metadata,
|
||||
)
|
||||
|
||||
async def record_cancelled(
|
||||
self,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
response_time_ms: int,
|
||||
first_byte_time_ms: int | None,
|
||||
status_code: int,
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
target_model: str | None = None,
|
||||
# 请求元数据(用于性能与调试记录)
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录客户端取消的请求
|
||||
|
||||
客户端主动断开连接不算系统失败,使用 cancelled 状态。
|
||||
"""
|
||||
provider_name = provider or "unknown"
|
||||
|
||||
await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
request_type="chat",
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code,
|
||||
status="cancelled",
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers or {},
|
||||
response_headers=response_headers or {},
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body or {},
|
||||
request_id=self.request_id,
|
||||
target_model=target_model,
|
||||
metadata=request_metadata,
|
||||
)
|
||||
# MessageTelemetry -- re-export from src.services.usage.telemetry (see import above)
|
||||
__all__ = ["MessageTelemetry", "MessageHandlerProtocol", "AdapterDetectorType"]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
|
||||
@@ -13,8 +13,8 @@ from src.api.handlers.base.utils import get_format_converter_registry
|
||||
from src.core.exceptions import ThinkingSignatureException, UpstreamClientException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.transport import get_vertex_ai_effective_format
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
|
||||
def _get_error_status_code(e: Exception, default: int = 400) -> int:
|
||||
|
||||
@@ -75,7 +75,6 @@ from src.models.database import (
|
||||
ProviderEndpoint,
|
||||
User,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
@@ -85,6 +84,7 @@ from src.services.provider.stream_policy import (
|
||||
from src.services.provider.transport import (
|
||||
build_provider_url,
|
||||
)
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
|
||||
|
||||
@@ -49,7 +49,7 @@ if TYPE_CHECKING:
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -38,7 +38,6 @@ from src.core.exceptions import (
|
||||
ProviderTimeoutException,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
@@ -46,6 +45,7 @@ from src.services.provider.stream_policy import (
|
||||
resolve_upstream_is_stream,
|
||||
)
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.utils.sse_parser import SSEEventParser
|
||||
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
||||
|
||||
@@ -29,7 +29,6 @@ from src.core.exceptions import (
|
||||
ThinkingSignatureException,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
@@ -37,6 +36,7 @@ from src.services.provider.stream_policy import (
|
||||
resolve_upstream_is_stream,
|
||||
)
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
|
||||
|
||||
@@ -693,36 +693,22 @@ class GeminiCliResponseParser(GeminiResponseParser):
|
||||
self.api_format = "gemini:cli"
|
||||
|
||||
|
||||
# 解析器注册表
|
||||
_PARSERS: dict[str, type[ResponseParser]] = {
|
||||
"claude:chat": ClaudeResponseParser,
|
||||
"claude:cli": ClaudeCliResponseParser,
|
||||
"openai:chat": OpenAIResponseParser,
|
||||
"openai:cli": OpenAICliResponseParser,
|
||||
"gemini:chat": GeminiResponseParser,
|
||||
"gemini:cli": GeminiCliResponseParser,
|
||||
}
|
||||
# 注册解析器到 core 层注册表(供 services 层通过 format_id 获取)
|
||||
from src.core.stream_types import get_parser_for_format, register_parser
|
||||
|
||||
|
||||
def get_parser_for_format(format_id: str) -> ResponseParser:
|
||||
"""
|
||||
根据格式 ID 获取 ResponseParser
|
||||
def register_default_parsers() -> None:
|
||||
register_parser("claude:chat", ClaudeResponseParser)
|
||||
register_parser("claude:cli", ClaudeCliResponseParser)
|
||||
register_parser("openai:chat", OpenAIResponseParser)
|
||||
register_parser("openai:cli", OpenAICliResponseParser)
|
||||
register_parser("gemini:chat", GeminiResponseParser)
|
||||
register_parser("gemini:cli", GeminiCliResponseParser)
|
||||
|
||||
Args:
|
||||
format_id: endpoint signature,如 "claude:chat", "openai:cli"
|
||||
|
||||
Returns:
|
||||
ResponseParser 实例
|
||||
|
||||
Raises:
|
||||
KeyError: 格式不存在
|
||||
"""
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
normalized = normalize_signature_key(format_id)
|
||||
if normalized not in _PARSERS:
|
||||
raise KeyError(f"Unknown format: {normalized}")
|
||||
return _PARSERS[normalized]()
|
||||
# 模块加载时自动注册(保证 import parsers 即可用,测试也不需要手动初始化)
|
||||
# main.py lifespan 中的显式调用是冗余但无害的安全保障(dict 覆盖幂等)
|
||||
register_default_parsers()
|
||||
|
||||
|
||||
__all__ = [
|
||||
@@ -732,6 +718,7 @@ __all__ = [
|
||||
"ClaudeCliResponseParser",
|
||||
"GeminiResponseParser",
|
||||
"GeminiCliResponseParser",
|
||||
"register_default_parsers",
|
||||
"get_parser_for_format",
|
||||
"is_cli_format",
|
||||
]
|
||||
|
||||
@@ -14,17 +14,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import object_session
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.core.api_format import (
|
||||
UPSTREAM_DROP_HEADERS,
|
||||
HeaderBuilder,
|
||||
@@ -32,99 +25,9 @@ from src.core.api_format import (
|
||||
make_signature_key,
|
||||
)
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
|
||||
from src.core.provider_auth_types import ProviderAuthInfo
|
||||
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Service Account 认证结果类型
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProviderAuthInfo:
|
||||
"""Provider 认证信息(用于 Service Account 等异步认证场景)"""
|
||||
|
||||
auth_header: str
|
||||
auth_value: str
|
||||
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
|
||||
decrypted_auth_config: dict[str, Any] | None = None
|
||||
|
||||
def as_tuple(self) -> tuple[str, str]:
|
||||
"""返回 (auth_header, auth_value) 元组"""
|
||||
return (self.auth_header, self.auth_value)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh helpers
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _acquire_refresh_lock(key_id: str) -> tuple[Any, bool]:
|
||||
"""尝试获取 OAuth refresh 分布式锁。
|
||||
|
||||
返回 ``(redis_client | None, got_lock)``。调用方在刷新完成后
|
||||
必须调用 :func:`_release_refresh_lock` 释放锁。
|
||||
"""
|
||||
redis = await get_redis_client(require_redis=False)
|
||||
lock_key = f"provider_oauth_refresh_lock:{key_id}"
|
||||
got_lock = False
|
||||
if redis is not None:
|
||||
try:
|
||||
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
|
||||
except Exception:
|
||||
got_lock = False
|
||||
return redis, got_lock
|
||||
|
||||
|
||||
async def _release_refresh_lock(redis: Any, key_id: str) -> None:
|
||||
"""释放 OAuth refresh 分布式锁(best-effort)。"""
|
||||
if redis is not None:
|
||||
try:
|
||||
await redis.delete(f"provider_oauth_refresh_lock:{key_id}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _persist_refreshed_token(
|
||||
key: Any,
|
||||
access_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> None:
|
||||
"""将刷新后的 access_token 和 auth_config 持久化到数据库。"""
|
||||
key.api_key = crypto_service.encrypt(access_token)
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
sess = object_session(key)
|
||||
if sess is not None:
|
||||
sess.add(key)
|
||||
sess.commit()
|
||||
else:
|
||||
logger.warning(
|
||||
"[OAUTH_REFRESH] key {} refreshed but cannot persist (no session); "
|
||||
"next request will refresh again",
|
||||
key.id,
|
||||
)
|
||||
|
||||
|
||||
def _get_proxy_config(key: Any, endpoint: Any = None) -> Any:
|
||||
"""获取有效代理配置(Key 级别优先于 Provider 级别)。"""
|
||||
try:
|
||||
from src.services.proxy_node.resolver import resolve_effective_proxy
|
||||
|
||||
provider = getattr(key, "provider", None) or (
|
||||
getattr(endpoint, "provider", None) if endpoint else None
|
||||
)
|
||||
provider_proxy = getattr(provider, "proxy", None)
|
||||
key_proxy = getattr(key, "proxy", None)
|
||||
return resolve_effective_proxy(provider_proxy, key_proxy)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
from src.services.provider.auth import get_provider_auth
|
||||
|
||||
# ==============================================================================
|
||||
# 统一的头部配置常量
|
||||
@@ -1003,303 +906,3 @@ def build_passthrough_request(
|
||||
endpoint,
|
||||
key,
|
||||
)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh logic (Kiro / Generic)
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _refresh_kiro_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Kiro OAuth refresh: validate + call Kiro-specific refresh endpoint."""
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import (
|
||||
refresh_access_token,
|
||||
validate_refresh_token,
|
||||
)
|
||||
|
||||
cfg = KiroAuthConfig.from_dict(token_meta or {})
|
||||
if not (cfg.refresh_token or "").strip():
|
||||
raise InvalidRequestException(
|
||||
"Kiro auth_config missing refresh_token; please re-import credentials."
|
||||
)
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
access_token, new_cfg = await refresh_access_token(
|
||||
cfg,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
new_meta = new_cfg.to_dict()
|
||||
new_meta["updated_at"] = int(time.time())
|
||||
|
||||
_persist_refreshed_token(key, access_token, new_meta)
|
||||
return new_meta
|
||||
|
||||
|
||||
async def _refresh_generic_oauth_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
template: Any,
|
||||
provider_type: str,
|
||||
refresh_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Generic OAuth refresh via template (Codex, Antigravity, ClaudeCode, etc.)."""
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
scopes = getattr(template.oauth, "scopes", None) or []
|
||||
scope_str = " ".join(scopes) if scopes else ""
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if scope_str:
|
||||
body["scope"] = scope_str
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if scope_str:
|
||||
form["scope"] = scope_str
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=30.0,
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
_persist_refreshed_token(key, access_token, token_meta)
|
||||
else:
|
||||
logger.warning(
|
||||
"OAuth token refresh failed: provider={}, key_id={}, status={}",
|
||||
provider_type,
|
||||
getattr(key, "id", "?"),
|
||||
resp.status_code,
|
||||
)
|
||||
|
||||
return token_meta
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Service Account 认证支持
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def get_provider_auth(
|
||||
endpoint: "ProviderEndpoint",
|
||||
key: "ProviderAPIKey",
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> ProviderAuthInfo | None:
|
||||
"""
|
||||
获取 Provider 的认证信息
|
||||
|
||||
对于标准 API Key,返回 None(由 build_headers 自动处理)。
|
||||
对于 Service Account,异步获取 Access Token 并返回认证信息。
|
||||
|
||||
Args:
|
||||
endpoint: 端点配置
|
||||
key: Provider API Key
|
||||
|
||||
Returns:
|
||||
Service Account 场景: ProviderAuthInfo 对象(包含认证信息和解密后的配置)
|
||||
API Key 场景: None(由 build_headers 处理)
|
||||
|
||||
Raises:
|
||||
InvalidRequestException: 认证配置无效或认证失败
|
||||
"""
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
|
||||
auth_type = getattr(key, "auth_type", "api_key")
|
||||
|
||||
if auth_type == "oauth":
|
||||
# OAuth token 保存在 key.api_key(加密),refresh_token/expires_at 等在 auth_config(加密 JSON)中。
|
||||
# 在请求前做一次懒刷新:接近过期时刷新 access_token,并用 Redis lock 避免并发风暴。
|
||||
encrypted_auth_config = getattr(key, "auth_config", None)
|
||||
if encrypted_auth_config:
|
||||
try:
|
||||
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
||||
token_meta = json.loads(decrypted_config)
|
||||
except Exception:
|
||||
token_meta = {}
|
||||
else:
|
||||
token_meta = {}
|
||||
|
||||
expires_at = token_meta.get("expires_at")
|
||||
refresh_token = token_meta.get("refresh_token")
|
||||
provider_type = str(token_meta.get("provider_type") or "")
|
||||
cached_access_token = str(token_meta.get("access_token") or "").strip()
|
||||
|
||||
# 120s skew (or force refresh when upstream returns 401)
|
||||
should_refresh = False
|
||||
try:
|
||||
if expires_at is not None:
|
||||
should_refresh = int(time.time()) >= int(expires_at) - 120
|
||||
except Exception:
|
||||
should_refresh = False
|
||||
|
||||
if force_refresh:
|
||||
should_refresh = True
|
||||
|
||||
# Kiro 特殊处理:如果没有缓存的 access_token 或 key.api_key 是占位符,强制刷新
|
||||
if provider_type == "kiro" and not should_refresh:
|
||||
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:
|
||||
try:
|
||||
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
||||
from src.core.provider_templates.types import ProviderType
|
||||
|
||||
try:
|
||||
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
|
||||
except Exception:
|
||||
template = None
|
||||
|
||||
redis, got_lock = await _acquire_refresh_lock(key.id)
|
||||
if got_lock or redis is None:
|
||||
try:
|
||||
if provider_type == ProviderType.KIRO.value:
|
||||
token_meta = await _refresh_kiro_token(key, endpoint, token_meta)
|
||||
elif template:
|
||||
token_meta = await _refresh_generic_oauth_token(
|
||||
key, endpoint, template, provider_type, refresh_token, token_meta
|
||||
)
|
||||
finally:
|
||||
if got_lock:
|
||||
await _release_refresh_lock(redis, key.id)
|
||||
except Exception:
|
||||
# 刷新失败不阻断请求;后续由上游返回 401 再触发管理端处理
|
||||
pass
|
||||
|
||||
# 获取最终使用的 access_token
|
||||
# Kiro 优先使用 token_meta 中缓存的 access_token(刷新后会更新到 token_meta)
|
||||
if provider_type == "kiro":
|
||||
refreshed_token = str(token_meta.get("access_token") or "").strip()
|
||||
effective_token = refreshed_token or crypto_service.decrypt(key.api_key)
|
||||
else:
|
||||
effective_token = crypto_service.decrypt(key.api_key)
|
||||
|
||||
decrypted_auth_config: dict[str, Any] | None = None
|
||||
if isinstance(token_meta, dict) and token_meta:
|
||||
decrypted_auth_config = token_meta
|
||||
|
||||
return ProviderAuthInfo(
|
||||
auth_header="Authorization",
|
||||
auth_value=f"Bearer {effective_token}",
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
if auth_type == "vertex_ai":
|
||||
from src.core.vertex_auth import VertexAuthError, VertexAuthService
|
||||
|
||||
try:
|
||||
# 优先从 auth_config 读取,兼容从 api_key 读取(过渡期)
|
||||
encrypted_auth_config = getattr(key, "auth_config", None)
|
||||
if encrypted_auth_config:
|
||||
# auth_config 可能是加密字符串或未加密的 dict
|
||||
if isinstance(encrypted_auth_config, dict):
|
||||
# 已经是 dict,直接使用(兼容未加密存储的情况)
|
||||
sa_json = encrypted_auth_config
|
||||
else:
|
||||
# 是加密字符串,需要解密
|
||||
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
||||
sa_json = json.loads(decrypted_config)
|
||||
else:
|
||||
# 兼容旧数据:从 api_key 读取
|
||||
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||
# 检查是否是占位符(表示 auth_config 丢失)
|
||||
if decrypted_key == "__placeholder__":
|
||||
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
|
||||
sa_json = json.loads(decrypted_key)
|
||||
|
||||
if not isinstance(sa_json, dict):
|
||||
raise InvalidRequestException("Service Account JSON 无效,请重新添加该密钥。")
|
||||
|
||||
# 获取 Access Token
|
||||
service = VertexAuthService(sa_json)
|
||||
access_token = await service.get_access_token()
|
||||
|
||||
# Vertex AI 使用 Bearer token
|
||||
return ProviderAuthInfo(
|
||||
auth_header="Authorization",
|
||||
auth_value=f"Bearer {access_token}",
|
||||
decrypted_auth_config=sa_json,
|
||||
)
|
||||
except InvalidRequestException:
|
||||
raise
|
||||
except VertexAuthError as e:
|
||||
raise InvalidRequestException(f"Vertex AI 认证失败:{e}")
|
||||
except json.JSONDecodeError:
|
||||
raise InvalidRequestException("Service Account JSON 格式无效,请重新添加该密钥。")
|
||||
except Exception:
|
||||
raise InvalidRequestException("Vertex AI 认证失败,请检查 Key 的 auth_config")
|
||||
|
||||
# 其他认证类型可在此扩展
|
||||
# elif auth_type == "oauth2":
|
||||
# ...
|
||||
|
||||
# 标准 API Key:返回 None,由 build_headers 处理
|
||||
return None
|
||||
|
||||
@@ -1,176 +1,19 @@
|
||||
"""
|
||||
响应解析器基类 - 定义统一的响应解析接口
|
||||
响应解析器基类 - re-export from src.core.stream_types
|
||||
|
||||
实际定义已下沉到 src/core/stream_types.py,此文件保留向后兼容的 re-export。
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from src.core.stream_types import (
|
||||
ParsedChunk,
|
||||
ParsedResponse,
|
||||
ResponseParser,
|
||||
StreamStats,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParsedChunk:
|
||||
"""解析后的流式数据块"""
|
||||
|
||||
# 原始数据
|
||||
raw_line: str
|
||||
event_type: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
# 提取的内容
|
||||
text_delta: str = ""
|
||||
is_done: bool = False
|
||||
is_error: bool = False
|
||||
error_message: str | None = None
|
||||
|
||||
# 使用量信息(通常在最后一个 chunk 中)
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 响应 ID
|
||||
response_id: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamStats:
|
||||
"""流式响应统计信息"""
|
||||
|
||||
# 计数
|
||||
chunk_count: int = 0
|
||||
data_count: int = 0
|
||||
|
||||
# Token 使用量
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 内容
|
||||
collected_text: str = ""
|
||||
response_id: str | None = None
|
||||
|
||||
# 状态
|
||||
has_completion: bool = False
|
||||
status_code: int = 200
|
||||
error_message: str | None = None
|
||||
|
||||
# Provider 信息
|
||||
provider_name: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
|
||||
# 响应头和完整响应
|
||||
response_headers: dict[str, str] = field(default_factory=dict)
|
||||
final_response: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParsedResponse:
|
||||
"""解析后的非流式响应"""
|
||||
|
||||
# 原始响应
|
||||
raw_response: dict[str, Any]
|
||||
status_code: int
|
||||
|
||||
# 提取的内容
|
||||
text_content: str = ""
|
||||
response_id: str | None = None
|
||||
|
||||
# 使用量
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 错误信息
|
||||
is_error: bool = False
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
|
||||
embedded_status_code: int | None = None
|
||||
|
||||
|
||||
class ResponseParser(ABC):
|
||||
"""
|
||||
响应解析器基类
|
||||
|
||||
定义统一的接口来解析不同 API 格式的响应。
|
||||
子类需要实现具体的解析逻辑。
|
||||
"""
|
||||
|
||||
# 解析器名称(用于日志)
|
||||
name: str = "base"
|
||||
|
||||
# 支持的 API 格式
|
||||
api_format: str = "UNKNOWN"
|
||||
|
||||
@abstractmethod
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
Args:
|
||||
line: SSE 行数据
|
||||
stats: 流统计对象(会被更新)
|
||||
|
||||
Returns:
|
||||
解析后的数据块,如果行不包含有效数据则返回 None
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
"""
|
||||
解析非流式响应
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
status_code: HTTP 状态码
|
||||
|
||||
Returns:
|
||||
解析后的响应对象
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从响应中提取 token 使用量
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
|
||||
Returns:
|
||||
包含 input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens 的字典
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
|
||||
Returns:
|
||||
提取的文本内容
|
||||
"""
|
||||
pass
|
||||
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断响应是否为错误响应
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
|
||||
Returns:
|
||||
是否为错误响应
|
||||
"""
|
||||
return "error" in response
|
||||
|
||||
def create_stats(self) -> StreamStats:
|
||||
"""创建新的流统计对象"""
|
||||
return StreamStats()
|
||||
__all__ = [
|
||||
"ParsedChunk",
|
||||
"ParsedResponse",
|
||||
"ResponseParser",
|
||||
"StreamStats",
|
||||
]
|
||||
|
||||
@@ -6,7 +6,6 @@ Video Handler 基类
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
@@ -19,75 +18,20 @@ from src.config.settings import config
|
||||
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.core.video_utils import (
|
||||
extract_short_id_from_operation,
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import httpx
|
||||
|
||||
from src.services.candidate.submit import SubmitOutcome
|
||||
|
||||
# 敏感信息匹配正则(预编译提升性能)
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def sanitize_error_message(message: str, max_length: int = 200) -> str:
|
||||
"""
|
||||
移除错误消息中可能包含的敏感信息
|
||||
|
||||
Args:
|
||||
message: 原始错误消息
|
||||
max_length: 最大长度,默认 200
|
||||
|
||||
Returns:
|
||||
脱敏后的消息
|
||||
"""
|
||||
if not message:
|
||||
return "Request failed"
|
||||
# 先脱敏再截断,确保敏感信息不会因截断位置而泄露
|
||||
sanitized = _SENSITIVE_PATTERN.sub("[REDACTED]", message)
|
||||
return sanitized[:max_length]
|
||||
|
||||
|
||||
def extract_short_id_from_operation(operation_id: str) -> str:
|
||||
"""
|
||||
从 operation ID 中提取短 ID
|
||||
|
||||
我们对外暴露的 operation name 格式是:
|
||||
- models/{model}/operations/{short_id}
|
||||
|
||||
此函数提取最后一部分作为 short_id,用于在数据库中查找任务。
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID(如 "models/veo-3.1/operations/abc123")
|
||||
|
||||
Returns:
|
||||
short_id(如 "abc123")
|
||||
"""
|
||||
# 格式: models/{model}/operations/{short_id}
|
||||
# 或者直接是 short_id
|
||||
if "/" in operation_id:
|
||||
# 提取最后一部分
|
||||
return operation_id.rsplit("/", 1)[-1]
|
||||
return operation_id
|
||||
|
||||
|
||||
def normalize_gemini_operation_id(operation_id: str) -> str:
|
||||
"""
|
||||
规范化 Gemini operation ID(保留用于向后兼容)
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID
|
||||
|
||||
Returns:
|
||||
规范化后的 operation ID(原样返回)
|
||||
"""
|
||||
return operation_id
|
||||
|
||||
|
||||
class VideoHandlerBase(ABC):
|
||||
"""视频处理器基类"""
|
||||
|
||||
@@ -7,13 +7,9 @@ Gemini 图像生成模型请求适配
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.core.video_utils import is_image_gen_model
|
||||
|
||||
def is_image_gen_model(model: str | None) -> bool:
|
||||
"""判断是否为图像生成模型(模式匹配,覆盖 gemini-*-image / imagen-* 系列)"""
|
||||
if not model:
|
||||
return False
|
||||
m = model.lower()
|
||||
return "image" in m and ("gemini" in m or "imagen" in m)
|
||||
__all__ = ["is_image_gen_model", "adapt_request_for_image_gen"]
|
||||
|
||||
|
||||
def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
@@ -41,7 +41,7 @@ from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
|
||||
@@ -39,7 +39,7 @@ from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
|
||||
@@ -36,9 +36,9 @@ from src.core.logger import logger
|
||||
from src.database import create_session
|
||||
from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User
|
||||
from src.services.auth.service import AuthService
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
|
||||
from src.services.provider.transport import redact_url_for_log
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -153,6 +153,14 @@ class RPMDefaults:
|
||||
# confidence 自然衰减速率:每分钟衰减的比例
|
||||
CONFIDENCE_DECAY_PER_MINUTE = 0.005 # 每分钟 -0.5%,约 200 分钟(~3.3h)从 1.0 衰减到 0
|
||||
|
||||
# === RPM 计数器时间窗口配置 ===
|
||||
# RPM 计数时间窗口(秒)
|
||||
RPM_BUCKET_SECONDS = 60
|
||||
# Redis key 过期时间(秒),需覆盖当前分钟与边界
|
||||
RPM_KEY_TTL_SECONDS = 120
|
||||
# 内存模式清理间隔(秒)
|
||||
RPM_CLEANUP_INTERVAL_SECONDS = 300
|
||||
|
||||
|
||||
# 向后兼容别名
|
||||
ConcurrencyDefaults = RPMDefaults
|
||||
|
||||
@@ -135,6 +135,19 @@ class Config:
|
||||
# CACHE_RESERVATION_RATIO: 缓存用户预留比例(默认 10%,新用户可用 90%)
|
||||
self.cache_reservation_ratio = float(os.getenv("CACHE_RESERVATION_RATIO", "0.1"))
|
||||
|
||||
# RPM 计数器时间窗口配置
|
||||
from src.config.constants import RPMDefaults
|
||||
|
||||
self.rpm_bucket_seconds = int(
|
||||
os.getenv("RPM_BUCKET_SECONDS", str(RPMDefaults.RPM_BUCKET_SECONDS))
|
||||
)
|
||||
self.rpm_key_ttl_seconds = int(
|
||||
os.getenv("RPM_KEY_TTL_SECONDS", str(RPMDefaults.RPM_KEY_TTL_SECONDS))
|
||||
)
|
||||
self.rpm_cleanup_interval_seconds = int(
|
||||
os.getenv("RPM_CLEANUP_INTERVAL_SECONDS", str(RPMDefaults.RPM_CLEANUP_INTERVAL_SECONDS))
|
||||
)
|
||||
|
||||
# 限流降级策略配置
|
||||
# RATE_LIMIT_FAIL_OPEN: 当限流服务(Redis)异常时的行为
|
||||
#
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
"""
|
||||
API Key / User 访问限制数据类型。
|
||||
|
||||
从 api/base/models_service.py 下沉到 core 层,
|
||||
消除 services→api 的反向依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
from src.core.logger import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ApiKey, User
|
||||
|
||||
|
||||
def _safe_normalize_signature(value: str) -> str:
|
||||
"""归一化 endpoint signature,解析失败时原样返回(小写)。"""
|
||||
try:
|
||||
return normalize_signature_key(value)
|
||||
except ValueError:
|
||||
logger.warning("[AccessRestrictions] 无法归一化 API 格式 '{}', 原样使用小写形式", value)
|
||||
return value.strip().lower()
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccessRestrictions:
|
||||
"""API Key 或 User 的访问限制"""
|
||||
|
||||
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
|
||||
allowed_models: list[str] | None = None # 允许的模型名称列表
|
||||
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
|
||||
|
||||
@classmethod
|
||||
def from_api_key_and_user(cls, api_key: ApiKey | None, user: User | None) -> AccessRestrictions:
|
||||
"""
|
||||
从 API Key 和 User 合并访问限制
|
||||
|
||||
限制逻辑:
|
||||
- API Key 的限制优先于 User 的限制
|
||||
- 如果 API Key 有限制,使用 API Key 的限制
|
||||
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
|
||||
- 两者都无限制则返回空限制
|
||||
"""
|
||||
allowed_providers: list[str] | None = None
|
||||
allowed_models: list[str] | None = None
|
||||
allowed_api_formats: list[str] | None = None
|
||||
|
||||
# 优先使用 API Key 的限制
|
||||
if api_key:
|
||||
if api_key.allowed_providers is not None:
|
||||
allowed_providers = api_key.allowed_providers
|
||||
if api_key.allowed_models is not None:
|
||||
allowed_models = api_key.allowed_models
|
||||
if api_key.allowed_api_formats is not None:
|
||||
allowed_api_formats = api_key.allowed_api_formats
|
||||
|
||||
# 如果 API Key 没有限制,检查 User 的限制
|
||||
if user:
|
||||
if allowed_providers is None and user.allowed_providers is not None:
|
||||
allowed_providers = user.allowed_providers
|
||||
if allowed_models is None and user.allowed_models is not None:
|
||||
allowed_models = user.allowed_models
|
||||
if allowed_api_formats is None and user.allowed_api_formats is not None:
|
||||
allowed_api_formats = user.allowed_api_formats
|
||||
|
||||
return cls(
|
||||
allowed_providers=allowed_providers,
|
||||
allowed_models=allowed_models,
|
||||
allowed_api_formats=allowed_api_formats,
|
||||
)
|
||||
|
||||
def is_api_format_allowed(self, api_format: str) -> bool:
|
||||
"""
|
||||
检查 API 格式是否被允许
|
||||
|
||||
Args:
|
||||
api_format: endpoint signature(如 "openai:chat")
|
||||
|
||||
Returns:
|
||||
True 如果格式被允许,False 否则
|
||||
"""
|
||||
if self.allowed_api_formats is None:
|
||||
return True
|
||||
target = _safe_normalize_signature(api_format)
|
||||
allowed = {_safe_normalize_signature(f) for f in self.allowed_api_formats if f}
|
||||
return target in allowed
|
||||
|
||||
def is_model_allowed(self, model_id: str, provider_id: str) -> bool:
|
||||
"""
|
||||
检查模型是否被允许访问
|
||||
|
||||
Args:
|
||||
model_id: 模型 ID
|
||||
provider_id: Provider ID
|
||||
|
||||
Returns:
|
||||
True 如果模型被允许,False 否则
|
||||
"""
|
||||
# 检查 Provider 限制
|
||||
if self.allowed_providers is not None:
|
||||
if provider_id not in self.allowed_providers:
|
||||
return False
|
||||
|
||||
# 检查模型限制
|
||||
if self.allowed_models is not None:
|
||||
if model_id not in self.allowed_models:
|
||||
return False
|
||||
|
||||
return True
|
||||
@@ -1425,7 +1425,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
)
|
||||
|
||||
try:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import (
|
||||
from src.core.api_format.conversion.thinking_cache import (
|
||||
signature_cache,
|
||||
)
|
||||
|
||||
@@ -1565,7 +1565,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE
|
||||
|
||||
try:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
|
||||
cached_or_dummy = signature_cache.get_or_dummy(model, text_val)
|
||||
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
"""Thinking block signature cache (triple-layer).
|
||||
|
||||
三层缓存设计:
|
||||
Layer 1: tool_use_id -> thoughtSignature (工具调用签名恢复)
|
||||
Layer 2: signature -> model_family (跨模型兼容校验)
|
||||
Layer 3: session_id -> latest signature (会话级签名追踪 + rewind 检测)
|
||||
|
||||
同时保留原有的 model:text -> signature 兼容层。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE
|
||||
|
||||
# 签名最小长度阈值
|
||||
MIN_SIGNATURE_LENGTH = 50
|
||||
|
||||
# TTL: 2 小时
|
||||
_SIGNATURE_TTL_SECONDS = 2 * 60 * 60
|
||||
|
||||
# 各层缓存上限
|
||||
_TOOL_CACHE_LIMIT = 500
|
||||
_FAMILY_CACHE_LIMIT = 200
|
||||
_SESSION_CACHE_LIMIT = 1000
|
||||
_TEXT_CACHE_LIMIT = 1000
|
||||
|
||||
|
||||
class _CacheEntry:
|
||||
"""带时间戳的缓存条目,支持 TTL 过期。"""
|
||||
|
||||
__slots__ = ("data", "created_at")
|
||||
|
||||
def __init__(self, data: Any) -> None:
|
||||
self.data = data
|
||||
self.created_at: float = time.monotonic()
|
||||
|
||||
def is_expired(self, now: float | None = None) -> bool:
|
||||
return ((now or time.monotonic()) - self.created_at) > _SIGNATURE_TTL_SECONDS
|
||||
|
||||
|
||||
class _SessionEntry:
|
||||
"""Session 层缓存数据,包含消息计数用于 rewind 检测。"""
|
||||
|
||||
__slots__ = ("signature", "message_count")
|
||||
|
||||
def __init__(self, signature: str, message_count: int) -> None:
|
||||
self.signature = signature
|
||||
self.message_count = message_count
|
||||
|
||||
|
||||
class ThinkingSignatureCache:
|
||||
"""Triple-layer thinking signature cache.
|
||||
|
||||
Layer 1 (tool): tool_use_id -> thoughtSignature
|
||||
当客户端(如 OpenCode) 在 tool_result 中丢弃了 signature 时用于恢复。
|
||||
|
||||
Layer 2 (family): signature -> model_family
|
||||
防止跨模型签名污染(Claude 签名不能用在 Gemini 上)。
|
||||
|
||||
Layer 3 (session): session_id -> latest signature + message_count
|
||||
会话级追踪,支持 rewind 检测(用户删除消息后不会注入来自"未来"的签名)。
|
||||
|
||||
Legacy (text): SHA256(model + text) -> signature
|
||||
向后兼容的 get_or_dummy() 接口。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._tool_sigs: dict[str, _CacheEntry] = {}
|
||||
self._families: dict[str, _CacheEntry] = {}
|
||||
self._sessions: dict[str, _CacheEntry] = {}
|
||||
self._text_sigs: dict[str, _CacheEntry] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# ===== Layer 1: Tool Use ID -> Signature =====
|
||||
|
||||
def cache_tool_signature(self, tool_use_id: str, signature: str) -> None:
|
||||
"""缓存工具调用对应的 thinking signature。"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
with self._lock:
|
||||
self._tool_sigs[tool_use_id] = _CacheEntry(signature)
|
||||
if len(self._tool_sigs) > _TOOL_CACHE_LIMIT:
|
||||
self._prune(self._tool_sigs, limit=_TOOL_CACHE_LIMIT)
|
||||
|
||||
def get_tool_signature(self, tool_use_id: str) -> str | None:
|
||||
"""查找工具调用对应的 signature。"""
|
||||
with self._lock:
|
||||
entry = self._tool_sigs.get(tool_use_id)
|
||||
if entry is None:
|
||||
return None
|
||||
if entry.is_expired():
|
||||
self._tool_sigs.pop(tool_use_id, None)
|
||||
return None
|
||||
return entry.data
|
||||
|
||||
# ===== Layer 2: Signature -> Model Family =====
|
||||
|
||||
def cache_thinking_family(self, signature: str, family: str) -> None:
|
||||
"""记录 signature 所属的模型家族。"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
with self._lock:
|
||||
self._families[signature] = _CacheEntry(family)
|
||||
if len(self._families) > _FAMILY_CACHE_LIMIT:
|
||||
self._prune(self._families, limit=_FAMILY_CACHE_LIMIT)
|
||||
|
||||
def get_signature_family(self, signature: str) -> str | None:
|
||||
"""查找 signature 所属的模型家族。"""
|
||||
with self._lock:
|
||||
entry = self._families.get(signature)
|
||||
if entry is None:
|
||||
return None
|
||||
if entry.is_expired():
|
||||
self._families.pop(signature, None)
|
||||
return None
|
||||
return entry.data
|
||||
|
||||
# ===== Layer 3: Session ID -> Latest Signature =====
|
||||
|
||||
def cache_session_signature(
|
||||
self, session_id: str, signature: str, message_count: int = 0
|
||||
) -> None:
|
||||
"""存储会话的最新 thinking signature。
|
||||
|
||||
Rewind 检测:当 message_count 小于已缓存值时,说明用户删除了消息,
|
||||
强制更新签名以避免注入来自"未来"的签名。
|
||||
"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
|
||||
with self._lock:
|
||||
existing = self._sessions.get(session_id)
|
||||
should_store = True
|
||||
|
||||
if existing and not existing.is_expired():
|
||||
entry: _SessionEntry = existing.data
|
||||
if message_count < entry.message_count:
|
||||
# Rewind detected: 用户删除了消息,强制更新
|
||||
pass
|
||||
elif message_count == entry.message_count:
|
||||
# 同一轮消息:仅当新签名更长(更完整)时才替换
|
||||
should_store = len(signature) > len(entry.signature)
|
||||
# else: 正常递增,更新
|
||||
|
||||
if should_store:
|
||||
self._sessions[session_id] = _CacheEntry(_SessionEntry(signature, message_count))
|
||||
if len(self._sessions) > _SESSION_CACHE_LIMIT:
|
||||
self._prune(self._sessions, limit=_SESSION_CACHE_LIMIT)
|
||||
|
||||
def get_session_signature(self, session_id: str) -> str | None:
|
||||
"""获取会话的最新 thinking signature。"""
|
||||
with self._lock:
|
||||
entry = self._sessions.get(session_id)
|
||||
if entry is None:
|
||||
return None
|
||||
if entry.is_expired():
|
||||
self._sessions.pop(session_id, None)
|
||||
return None
|
||||
return entry.data.signature
|
||||
|
||||
# ===== Legacy: model:text -> signature(向后兼容) =====
|
||||
|
||||
def get_or_dummy(self, model: str, thinking_text: str) -> str | None:
|
||||
"""Legacy: 根据 model + thinking_text 查找 signature。
|
||||
|
||||
Gemini 模型在未命中时返回 DUMMY_THOUGHT_SIGNATURE(跳过验证)。
|
||||
"""
|
||||
key = self._text_key(model, thinking_text)
|
||||
with self._lock:
|
||||
entry = self._text_sigs.get(key)
|
||||
if entry is not None:
|
||||
if entry.is_expired():
|
||||
self._text_sigs.pop(key, None)
|
||||
else:
|
||||
return entry.data
|
||||
if str(model).startswith("gemini-"):
|
||||
return DUMMY_THOUGHT_SIGNATURE
|
||||
return None
|
||||
|
||||
def cache(self, model: str, thinking_text: str, signature: str) -> None:
|
||||
"""Legacy: 缓存 model + thinking_text -> signature。"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
|
||||
key = self._text_key(model, thinking_text)
|
||||
with self._lock:
|
||||
if key in self._text_sigs:
|
||||
self._text_sigs[key] = _CacheEntry(signature)
|
||||
return
|
||||
|
||||
if len(self._text_sigs) >= _TEXT_CACHE_LIMIT:
|
||||
# FIFO 淘汰 1/4
|
||||
evict_n = max(1, _TEXT_CACHE_LIMIT // 4)
|
||||
for k in list(self._text_sigs.keys())[:evict_n]:
|
||||
self._text_sigs.pop(k, None)
|
||||
|
||||
self._text_sigs[key] = _CacheEntry(signature)
|
||||
|
||||
# ===== Utilities =====
|
||||
|
||||
@staticmethod
|
||||
def _text_key(model: str, thinking_text: str) -> str:
|
||||
content = f"{model}\x00{thinking_text}"
|
||||
return hashlib.sha256(content.encode("utf-8")).hexdigest()[:32]
|
||||
|
||||
@staticmethod
|
||||
def _prune(d: dict[str, _CacheEntry], *, limit: int | None = None) -> None:
|
||||
"""Remove expired entries and optionally enforce a size limit."""
|
||||
now = time.monotonic()
|
||||
expired = [k for k, v in d.items() if v.is_expired(now)]
|
||||
for k in expired:
|
||||
d.pop(k, None)
|
||||
|
||||
if limit is None or len(d) <= limit:
|
||||
return
|
||||
|
||||
excess = len(d) - limit
|
||||
for k, _entry in sorted(d.items(), key=lambda kv: kv[1].created_at)[:excess]:
|
||||
d.pop(k, None)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""清空所有缓存层(用于测试或手动重置)。"""
|
||||
with self._lock:
|
||||
self._tool_sigs.clear()
|
||||
self._families.clear()
|
||||
self._sessions.clear()
|
||||
self._text_sigs.clear()
|
||||
|
||||
|
||||
signature_cache = ThinkingSignatureCache()
|
||||
|
||||
__all__ = ["ThinkingSignatureCache", "signature_cache", "MIN_SIGNATURE_LENGTH"]
|
||||
@@ -8,7 +8,7 @@ from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.core.modules.base import (
|
||||
@@ -22,6 +22,17 @@ if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
class ConfigBackend(Protocol):
|
||||
"""模块配置读写后端协议。
|
||||
|
||||
通过 ``ModuleRegistry.set_config_backend()`` 在应用启动时注入实现,
|
||||
使 core 层无需在运行时 import services 层。
|
||||
"""
|
||||
|
||||
def get_config(self, db: Any, key: str, default: Any = None) -> Any: ...
|
||||
def set_config(self, db: Any, key: str, value: Any, description: Any = None) -> Any: ...
|
||||
|
||||
|
||||
class ModuleRegistry:
|
||||
"""
|
||||
模块注册中心 - 单例模式
|
||||
@@ -34,11 +45,17 @@ class ModuleRegistry:
|
||||
"""
|
||||
|
||||
_instance: ModuleRegistry | None = None
|
||||
_config_backend: ConfigBackend | None = None
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._modules: dict[str, ModuleDefinition] = {}
|
||||
self._initialized: set[str] = set()
|
||||
|
||||
@classmethod
|
||||
def set_config_backend(cls, backend: ConfigBackend) -> None:
|
||||
"""注入配置读写后端,消除 core→services 的运行时依赖"""
|
||||
cls._config_backend = backend
|
||||
|
||||
@classmethod
|
||||
def get_instance(cls) -> ModuleRegistry:
|
||||
"""获取单例实例"""
|
||||
@@ -50,6 +67,7 @@ class ModuleRegistry:
|
||||
def reset_instance(cls) -> None:
|
||||
"""重置单例(仅用于测试)"""
|
||||
cls._instance = None
|
||||
cls._config_backend = None
|
||||
|
||||
def register(self, module: ModuleDefinition) -> None:
|
||||
"""
|
||||
@@ -123,6 +141,15 @@ class ModuleRegistry:
|
||||
|
||||
# ========== 启用状态检查(运行级)==========
|
||||
|
||||
def _get_config_backend(self) -> ConfigBackend:
|
||||
"""获取配置后端(优先使用已注入的,兜底 lazy import)"""
|
||||
if self._config_backend is not None:
|
||||
return self._config_backend
|
||||
# 兜底: 未注入时使用 lazy import(向后兼容独立脚本/测试场景)
|
||||
from src.services.system.config import SystemConfigService # noqa: lazy fallback
|
||||
|
||||
return SystemConfigService # type: ignore[return-value]
|
||||
|
||||
def is_enabled(self, name: str, db: Session) -> bool:
|
||||
"""
|
||||
检查模块是否运行启用(数据库配置)
|
||||
@@ -131,10 +158,8 @@ class ModuleRegistry:
|
||||
name: 模块名称
|
||||
db: 数据库会话
|
||||
"""
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
config_key = f"module.{name}.enabled"
|
||||
value = SystemConfigService.get_config(db, config_key, default=False)
|
||||
value = self._get_config_backend().get_config(db, config_key, default=False)
|
||||
return bool(value)
|
||||
|
||||
def set_enabled(self, name: str, enabled: bool, db: Session) -> None:
|
||||
@@ -146,15 +171,13 @@ class ModuleRegistry:
|
||||
enabled: 是否启用
|
||||
db: 数据库会话
|
||||
"""
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
if name not in self._modules:
|
||||
raise ValueError(f"Module [{name}] not registered")
|
||||
|
||||
config_key = f"module.{name}.enabled"
|
||||
module = self._modules[name]
|
||||
description = f"模块 [{module.metadata.display_name}] 启用状态"
|
||||
SystemConfigService.set_config(db, config_key, enabled, description)
|
||||
self._get_config_backend().set_config(db, config_key, enabled, description)
|
||||
|
||||
# ========== 激活状态检查 ==========
|
||||
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""
|
||||
Provider 认证相关的数据类型。
|
||||
|
||||
从 api/handlers/base/request_builder.py 下沉到 core 层,
|
||||
消除 services→api 的反向依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProviderAuthInfo:
|
||||
"""Provider 认证信息(用于 Service Account 等异步认证场景)"""
|
||||
|
||||
auth_header: str
|
||||
auth_value: str
|
||||
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
|
||||
decrypted_auth_config: dict[str, Any] | None = None
|
||||
|
||||
def as_tuple(self) -> tuple[str, str]:
|
||||
"""返回 (auth_header, auth_value) 元组"""
|
||||
return (self.auth_header, self.auth_value)
|
||||
@@ -10,7 +10,6 @@ import jwt
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.core.logger import logger
|
||||
from src.core.provider_types import ProviderType
|
||||
from src.services.proxy_node.resolver import build_proxy_url
|
||||
|
||||
_ANTHROPIC_TOKEN_URL = "https://console.anthropic.com/v1/oauth/token"
|
||||
_GOOGLE_USERINFO_URL = "https://www.googleapis.com/oauth2/v1/userinfo?alt=json"
|
||||
@@ -22,6 +21,8 @@ def _coerce_proxy_url(proxy_config: dict[str, Any] | None) -> str | None:
|
||||
try:
|
||||
if not proxy_config.get("enabled", True):
|
||||
return None
|
||||
from src.services.proxy_node.resolver import build_proxy_url # lazy: core→services
|
||||
|
||||
return build_proxy_url(proxy_config)
|
||||
except Exception:
|
||||
return None
|
||||
@@ -363,11 +364,10 @@ async def enrich_auth_config(
|
||||
"""Enrich auth_config with non-secret metadata (email/account_id).
|
||||
|
||||
各 provider 的 enrichment 逻辑通过 register_auth_enricher 注册。
|
||||
注意: ensure_providers_bootstrapped() 在应用启动时(main.py lifespan)已显式调用。
|
||||
"""
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
pt = normalize_provider_type(provider_type)
|
||||
enricher = _auth_enrichers.get(pt)
|
||||
if enricher:
|
||||
|
||||
@@ -15,9 +15,9 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
|
||||
from src.core.provider_templates.types import ProviderType
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
PROD_BASE_URL as ANTIGRAVITY_PROD_URL,
|
||||
)
|
||||
|
||||
# Antigravity 生产环境 URL(唯一定义点,services 层通过 re-export 引用)
|
||||
ANTIGRAVITY_PROD_URL = "https://cloudcode-pa.googleapis.com"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
"""
|
||||
响应解析器基类与流式统计类型。
|
||||
|
||||
从 api/handlers/base/response_parser.py 下沉到 core 层,
|
||||
消除 services→api 的反向依赖。同时提供 parser 注册表,
|
||||
允许 API 层注册具体实现,services 层通过 format_id 获取实例。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParsedChunk:
|
||||
"""解析后的流式数据块"""
|
||||
|
||||
# 原始数据
|
||||
raw_line: str
|
||||
event_type: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
# 提取的内容
|
||||
text_delta: str = ""
|
||||
is_done: bool = False
|
||||
is_error: bool = False
|
||||
error_message: str | None = None
|
||||
|
||||
# 使用量信息(通常在最后一个 chunk 中)
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 响应 ID
|
||||
response_id: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamStats:
|
||||
"""流式响应统计信息"""
|
||||
|
||||
# 计数
|
||||
chunk_count: int = 0
|
||||
data_count: int = 0
|
||||
|
||||
# Token 使用量
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 内容
|
||||
collected_text: str = ""
|
||||
response_id: str | None = None
|
||||
|
||||
# 状态
|
||||
has_completion: bool = False
|
||||
status_code: int = 200
|
||||
error_message: str | None = None
|
||||
|
||||
# Provider 信息
|
||||
provider_name: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
|
||||
# 响应头和完整响应
|
||||
response_headers: dict[str, str] = field(default_factory=dict)
|
||||
final_response: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParsedResponse:
|
||||
"""解析后的非流式响应"""
|
||||
|
||||
# 原始响应
|
||||
raw_response: dict[str, Any]
|
||||
status_code: int
|
||||
|
||||
# 提取的内容
|
||||
text_content: str = ""
|
||||
response_id: str | None = None
|
||||
|
||||
# 使用量
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 错误信息
|
||||
is_error: bool = False
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
|
||||
embedded_status_code: int | None = None
|
||||
|
||||
|
||||
class ResponseParser(ABC):
|
||||
"""
|
||||
响应解析器基类
|
||||
|
||||
定义统一的接口来解析不同 API 格式的响应。
|
||||
子类需要实现具体的解析逻辑。
|
||||
"""
|
||||
|
||||
# 解析器名称(用于日志)
|
||||
name: str = "base"
|
||||
|
||||
# 支持的 API 格式
|
||||
api_format: str = "UNKNOWN"
|
||||
|
||||
@abstractmethod
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
Args:
|
||||
line: SSE 行数据
|
||||
stats: 流统计对象(会被更新)
|
||||
|
||||
Returns:
|
||||
解析后的数据块,如果行不包含有效数据则返回 None
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
"""
|
||||
解析非流式响应
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
status_code: HTTP 状态码
|
||||
|
||||
Returns:
|
||||
解析后的响应对象
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从响应中提取 token 使用量
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
|
||||
Returns:
|
||||
包含 input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens 的字典
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
|
||||
Returns:
|
||||
提取的文本内容
|
||||
"""
|
||||
pass
|
||||
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断响应是否为错误响应
|
||||
|
||||
Args:
|
||||
response: 响应 JSON
|
||||
|
||||
Returns:
|
||||
是否为错误响应
|
||||
"""
|
||||
return "error" in response
|
||||
|
||||
def create_stats(self) -> StreamStats:
|
||||
"""创建新的流统计对象"""
|
||||
return StreamStats()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Parser 注册表 -- API 层注册具体实现,services 层通过 format_id 获取
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_PARSER_REGISTRY: dict[str, type[ResponseParser]] = {}
|
||||
|
||||
|
||||
def register_parser(format_id: str, parser_class: type[ResponseParser]) -> None:
|
||||
"""注册一个格式对应的 ResponseParser 实现"""
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
normalized = normalize_signature_key(format_id)
|
||||
_PARSER_REGISTRY[normalized] = parser_class
|
||||
|
||||
|
||||
def get_parser_for_format(format_id: str) -> ResponseParser:
|
||||
"""
|
||||
根据格式 ID 获取 ResponseParser 实例
|
||||
|
||||
Args:
|
||||
format_id: endpoint signature,如 "claude:chat", "openai:cli"
|
||||
|
||||
Returns:
|
||||
ResponseParser 实例
|
||||
|
||||
Raises:
|
||||
KeyError: 格式不存在
|
||||
"""
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
if not _PARSER_REGISTRY:
|
||||
raise KeyError(
|
||||
f"Parser registry is empty when looking up '{format_id}'. "
|
||||
"Ensure parsers are registered at startup (import src.api.handlers.base.parsers)."
|
||||
)
|
||||
|
||||
normalized = normalize_signature_key(format_id)
|
||||
if normalized not in _PARSER_REGISTRY:
|
||||
raise KeyError(f"Unknown format: {normalized}")
|
||||
return _PARSER_REGISTRY[normalized]()
|
||||
+17
-7
@@ -10,6 +10,7 @@ from __future__ import annotations
|
||||
import json
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
@@ -102,13 +103,17 @@ class VertexAuthService:
|
||||
}
|
||||
return jwt.encode(payload, self.private_key, algorithm="RS256")
|
||||
|
||||
async def get_access_token(self) -> str:
|
||||
async def get_access_token(self, *, httpx_client_kwargs: dict[str, Any] | None = None) -> str:
|
||||
"""
|
||||
获取 Access Token(带 LRU 缓存)
|
||||
|
||||
如果缓存中有有效的 Token(距离过期超过 60 秒),直接返回。
|
||||
否则重新获取 Token。缓存采用 LRU 策略,超过 100 个条目时淘汰最旧的。
|
||||
|
||||
Args:
|
||||
httpx_client_kwargs: 传给 httpx.AsyncClient 的额外参数(如代理配置)。
|
||||
调用者(services 层)负责构建,core 层不关心代理细节。
|
||||
|
||||
Returns:
|
||||
Access Token 字符串
|
||||
|
||||
@@ -129,10 +134,10 @@ class VertexAuthService:
|
||||
try:
|
||||
signed_jwt = self._create_jwt()
|
||||
|
||||
# 使用系统默认代理(Vertex AI token endpoint 是外部服务)
|
||||
from src.services.proxy_node.resolver import build_proxy_client_kwargs
|
||||
|
||||
async with httpx.AsyncClient(**build_proxy_client_kwargs(timeout=30)) as client:
|
||||
client_kwargs = (
|
||||
httpx_client_kwargs if httpx_client_kwargs is not None else {"timeout": 30}
|
||||
)
|
||||
async with httpx.AsyncClient(**client_kwargs) as client:
|
||||
resp = await client.post(
|
||||
self.TOKEN_URL,
|
||||
data={
|
||||
@@ -186,16 +191,21 @@ class VertexAuthService:
|
||||
cls._token_cache.clear()
|
||||
|
||||
|
||||
async def get_vertex_access_token(service_account_json: str) -> tuple[str, str]:
|
||||
async def get_vertex_access_token(
|
||||
service_account_json: str,
|
||||
*,
|
||||
httpx_client_kwargs: dict[str, Any] | None = None,
|
||||
) -> tuple[str, str]:
|
||||
"""
|
||||
便捷函数:获取 Vertex AI Access Token 和 Project ID
|
||||
|
||||
Args:
|
||||
service_account_json: Service Account JSON 字符串
|
||||
httpx_client_kwargs: 传给 httpx.AsyncClient 的额外参数(如代理配置)
|
||||
|
||||
Returns:
|
||||
(access_token, project_id) 元组
|
||||
"""
|
||||
service = VertexAuthService(service_account_json)
|
||||
token = await service.get_access_token()
|
||||
token = await service.get_access_token(httpx_client_kwargs=httpx_client_kwargs)
|
||||
return token, service.project_id
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
"""
|
||||
视频/图像相关的纯工具函数。
|
||||
|
||||
从 api/handlers 层下沉到 core 层,消除 services→api 的反向依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
# 敏感信息匹配正则(预编译提升性能)
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def sanitize_error_message(message: str, max_length: int = 200) -> str:
|
||||
"""
|
||||
移除错误消息中可能包含的敏感信息
|
||||
|
||||
Args:
|
||||
message: 原始错误消息
|
||||
max_length: 最大长度,默认 200
|
||||
|
||||
Returns:
|
||||
脱敏后的消息
|
||||
"""
|
||||
if not message:
|
||||
return "Request failed"
|
||||
# 先脱敏再截断,确保敏感信息不会因截断位置而泄露
|
||||
sanitized = _SENSITIVE_PATTERN.sub("[REDACTED]", message)
|
||||
return sanitized[:max_length]
|
||||
|
||||
|
||||
def extract_short_id_from_operation(operation_id: str) -> str:
|
||||
"""
|
||||
从 operation ID 中提取短 ID
|
||||
|
||||
我们对外暴露的 operation name 格式是:
|
||||
- models/{model}/operations/{short_id}
|
||||
|
||||
此函数提取最后一部分作为 short_id,用于在数据库中查找任务。
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID(如 "models/veo-3.1/operations/abc123")
|
||||
|
||||
Returns:
|
||||
short_id(如 "abc123")
|
||||
"""
|
||||
# 格式: models/{model}/operations/{short_id}
|
||||
# 或者直接是 short_id
|
||||
if "/" in operation_id:
|
||||
# 提取最后一部分
|
||||
return operation_id.rsplit("/", 1)[-1]
|
||||
return operation_id
|
||||
|
||||
|
||||
def normalize_gemini_operation_id(operation_id: str) -> str:
|
||||
"""
|
||||
规范化 Gemini operation ID(保留用于向后兼容)
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID
|
||||
|
||||
Returns:
|
||||
规范化后的 operation ID(原样返回)
|
||||
"""
|
||||
return operation_id
|
||||
|
||||
|
||||
def is_image_gen_model(model: str | None) -> bool:
|
||||
"""判断是否为图像生成模型(模式匹配,覆盖 gemini-*-image / imagen-* 系列)"""
|
||||
if not model:
|
||||
return False
|
||||
m = model.lower()
|
||||
return "image" in m and ("gemini" in m or "imagen" in m)
|
||||
+17
@@ -176,6 +176,12 @@ async def lifespan(app: FastAPI) -> Any:
|
||||
from src.modules import ALL_MODULES
|
||||
|
||||
module_registry = get_module_registry()
|
||||
|
||||
# 注入配置后端,消除 core/modules→services 的运行时 lazy import
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
module_registry.set_config_backend(SystemConfigService) # type: ignore[arg-type]
|
||||
|
||||
for module in ALL_MODULES:
|
||||
module_registry.register(module)
|
||||
|
||||
@@ -195,6 +201,17 @@ async def lifespan(app: FastAPI) -> Any:
|
||||
|
||||
logger.info(f"功能模块初始化完成: {len(available_modules)}/{len(ALL_MODULES)} 个模块可用")
|
||||
|
||||
# 显式 bootstrap provider plugins(注册 envelope/enricher 等)
|
||||
# 使 core/provider_oauth_utils 不需要在运行时 lazy import services 层
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
|
||||
# 显式触发 parsers 注册(使 core/stream_types 不需要 lazy import api 层)
|
||||
from src.api.handlers.base.parsers import register_default_parsers
|
||||
|
||||
register_default_parsers()
|
||||
|
||||
logger.info(f"服务启动成功: http://{config.host}:{config.port}")
|
||||
logger.info("=" * 60)
|
||||
|
||||
|
||||
Vendored
+3
-7
@@ -1,14 +1,10 @@
|
||||
"""
|
||||
缓存服务模块
|
||||
"""通用缓存模块。
|
||||
|
||||
包含缓存后端、缓存亲和性、缓存同步等功能。
|
||||
包含缓存后端、缓存失效与缓存同步等能力(backend/sync/*_cache)。
|
||||
|
||||
注意:由于循环依赖问题,部分类需要直接从子模块导入:
|
||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
调度/候选/缓存亲和性相关逻辑已迁移到 `src.services.scheduling`。
|
||||
"""
|
||||
|
||||
# 只导出不会导致循环依赖的基础类
|
||||
from src.services.cache.backend import BaseCacheBackend, LocalCache, RedisCache, get_cache_backend
|
||||
|
||||
__all__ = [
|
||||
|
||||
Vendored
+2
-2
@@ -54,7 +54,7 @@ class CacheInvalidationService:
|
||||
logger.error(f"[CacheInvalidation] 失效 ModelCacheService 缓存失败: {e}")
|
||||
|
||||
# 4. 清除 /v1/models 列表缓存
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.services.cache.model_list_cache import invalidate_models_list_cache
|
||||
|
||||
try:
|
||||
await invalidate_models_list_cache()
|
||||
@@ -79,7 +79,7 @@ class CacheInvalidationService:
|
||||
self._refresh_provider_cache(provider_id)
|
||||
|
||||
# 清除 /v1/models 列表缓存(allowed_models 变更会影响模型可用性)
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.services.cache.model_list_cache import invalidate_models_list_cache
|
||||
|
||||
try:
|
||||
await invalidate_models_list_cache()
|
||||
|
||||
+31
@@ -0,0 +1,31 @@
|
||||
"""
|
||||
/v1/models 列表缓存管理。
|
||||
|
||||
从 api/base/models_service.py 迁移到 services 层,
|
||||
消除 services→api 的反向依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from src.core.cache_service import CacheService
|
||||
from src.core.logger import logger
|
||||
|
||||
# 缓存 key 前缀(models_service.py 也使用此常量)
|
||||
MODELS_LIST_CACHE_PREFIX = "models:list"
|
||||
|
||||
|
||||
async def invalidate_models_list_cache() -> None:
|
||||
"""
|
||||
清除所有 /v1/models 列表缓存
|
||||
|
||||
在模型创建、更新、删除时调用,确保模型列表实时更新
|
||||
"""
|
||||
try:
|
||||
# 使用通配符删除所有 models:list:* 缓存(包括多格式组合的 key)
|
||||
deleted = await CacheService.delete_pattern(f"{MODELS_LIST_CACHE_PREFIX}:*")
|
||||
if deleted > 0:
|
||||
logger.info("[ModelsService] 已清除 {} 个 {} 缓存", deleted, MODELS_LIST_CACHE_PREFIX)
|
||||
else:
|
||||
logger.debug("[ModelsService] 无 {} 缓存需要清除", MODELS_LIST_CACHE_PREFIX)
|
||||
except Exception as e:
|
||||
logger.warning("[ModelsService] 清除缓存失败: {}", e)
|
||||
@@ -13,9 +13,9 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.database import RequestCandidate
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.task.exceptions import StreamProbeError
|
||||
from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult
|
||||
from src.services.task.schema import ExecutionResult
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
CANDIDATE_KEY_SCHEMA_VERSION = "1.0"
|
||||
|
||||
|
||||
@@ -5,8 +5,8 @@ from typing import Any
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import ApiKey
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
||||
from src.services.orchestration.error_classifier import ErrorClassifier
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
from .recorder import CandidateRecorder
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any, Protocol, runtime_checkable
|
||||
import httpx
|
||||
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
|
||||
@@ -928,6 +928,7 @@ class ModelCostService:
|
||||
|
||||
# 获取对应 API 格式的 Adapter 实例来计算成本
|
||||
# 优先检查 Chat Adapter,然后检查 CLI Adapter
|
||||
# TODO(arch): 引入 adapter 能力注册表,消除 services->api 依赖
|
||||
from src.api.handlers.base.chat_adapter_base import get_adapter_instance
|
||||
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_instance
|
||||
|
||||
|
||||
@@ -10,13 +10,13 @@ from sqlalchemy import and_
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.models.api import ModelCreate, ModelResponse, ModelUpdate
|
||||
from src.models.database import Model, Provider
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
from src.services.cache.model_list_cache import invalidate_models_list_cache
|
||||
|
||||
|
||||
class ModelService:
|
||||
|
||||
@@ -169,6 +169,7 @@ def merge_upstream_metadata(
|
||||
|
||||
def get_adapter_for_format(api_format: str) -> type | None:
|
||||
"""根据 API 格式获取对应的 Adapter 类"""
|
||||
# TODO(arch): 引入 adapter 能力注册表,消除 services->api 依赖
|
||||
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
|
||||
|
||||
|
||||
@@ -13,8 +13,8 @@ from sqlalchemy.orm import Session
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
class CandidateResolver:
|
||||
|
||||
@@ -25,9 +25,9 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.orchestration.error_handler import ErrorHandlerService
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
class ErrorAction(Enum):
|
||||
|
||||
@@ -22,11 +22,11 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
class ErrorHandlerService:
|
||||
|
||||
@@ -11,9 +11,9 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
from src.services.request.executor import RequestExecutor
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
class RequestDispatcher:
|
||||
|
||||
@@ -11,7 +11,9 @@ import re
|
||||
import threading
|
||||
|
||||
# ============== API 端点 ==============
|
||||
PROD_BASE_URL = "https://cloudcode-pa.googleapis.com"
|
||||
# 唯一定义在 core 层,此处 re-export 保持向后兼容
|
||||
from src.core.provider_templates.fixed_providers import ANTIGRAVITY_PROD_URL as PROD_BASE_URL
|
||||
|
||||
DAILY_BASE_URL = "https://daily-cloudcode-pa.googleapis.com"
|
||||
SANDBOX_BASE_URL = "https://daily-cloudcode-pa.sandbox.googleapis.com"
|
||||
|
||||
@@ -75,8 +77,7 @@ URL_UNAVAILABLE_TTL_SECONDS = 300 # 5 分钟
|
||||
# ============== Thinking Signature ==============
|
||||
# 统一从 core 层导入,避免多处定义
|
||||
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE # noqa: E402
|
||||
|
||||
MIN_SIGNATURE_LENGTH = 50 # 与 Antigravity-Manager 对齐
|
||||
from src.core.api_format.conversion.thinking_cache import MIN_SIGNATURE_LENGTH # noqa: E402, F401
|
||||
|
||||
# ============== Thinking Budget ==============
|
||||
THINKING_BUDGET_AUTO_CAP = 24576
|
||||
|
||||
@@ -521,7 +521,7 @@ def _inject_thought_signatures(inner_request: dict[str, Any], session_id: str |
|
||||
"""
|
||||
|
||||
try:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
except Exception:
|
||||
return
|
||||
|
||||
@@ -881,7 +881,7 @@ def wrap_v1internal_request(
|
||||
13. 注入 sessionId
|
||||
14. 构建 v1internal 信封
|
||||
"""
|
||||
from src.api.handlers.gemini.image_gen import is_image_gen_model
|
||||
from src.core.video_utils import is_image_gen_model
|
||||
|
||||
inner_request = dict(gemini_request)
|
||||
inner_request.pop("model", None)
|
||||
@@ -979,7 +979,7 @@ def cache_thought_signatures(model: str, response: dict[str, Any]) -> None:
|
||||
同时缓存到 legacy (text) 层和 tool (Layer 1) 层。
|
||||
"""
|
||||
try:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
except Exception:
|
||||
return
|
||||
|
||||
|
||||
@@ -1,239 +1,13 @@
|
||||
"""Antigravity thinking block signature cache (triple-layer).
|
||||
"""Backward-compatible re-export for Antigravity thinking signature cache.
|
||||
|
||||
与 Antigravity-Manager 对齐的三层缓存设计:
|
||||
Layer 1: tool_use_id → thoughtSignature (工具调用签名恢复)
|
||||
Layer 2: signature → model_family (跨模型兼容校验)
|
||||
Layer 3: session_id → latest signature (会话级签名追踪 + rewind 检测)
|
||||
|
||||
同时保留原有的 model:text → signature 兼容层。
|
||||
The implementation moved to `src.core.api_format.conversion.thinking_cache` to eliminate
|
||||
core → services reverse dependencies.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
DUMMY_THOUGHT_SIGNATURE,
|
||||
from src.core.api_format.conversion.thinking_cache import (
|
||||
MIN_SIGNATURE_LENGTH,
|
||||
ThinkingSignatureCache,
|
||||
signature_cache,
|
||||
)
|
||||
|
||||
# TTL: 2 小时(与 Antigravity-Manager 对齐)
|
||||
_SIGNATURE_TTL_SECONDS = 2 * 60 * 60
|
||||
|
||||
# 各层缓存上限
|
||||
_TOOL_CACHE_LIMIT = 500
|
||||
_FAMILY_CACHE_LIMIT = 200
|
||||
_SESSION_CACHE_LIMIT = 1000
|
||||
_TEXT_CACHE_LIMIT = 1000
|
||||
|
||||
|
||||
class _CacheEntry:
|
||||
"""带时间戳的缓存条目,支持 TTL 过期。"""
|
||||
|
||||
__slots__ = ("data", "created_at")
|
||||
|
||||
def __init__(self, data: Any) -> None:
|
||||
self.data = data
|
||||
self.created_at: float = time.monotonic()
|
||||
|
||||
def is_expired(self, now: float | None = None) -> bool:
|
||||
return ((now or time.monotonic()) - self.created_at) > _SIGNATURE_TTL_SECONDS
|
||||
|
||||
|
||||
class _SessionEntry:
|
||||
"""Session 层缓存数据,包含消息计数用于 rewind 检测。"""
|
||||
|
||||
__slots__ = ("signature", "message_count")
|
||||
|
||||
def __init__(self, signature: str, message_count: int) -> None:
|
||||
self.signature = signature
|
||||
self.message_count = message_count
|
||||
|
||||
|
||||
class ThinkingSignatureCache:
|
||||
"""Triple-layer thinking signature cache.
|
||||
|
||||
Layer 1 (tool): tool_use_id → thoughtSignature
|
||||
当客户端(如 OpenCode) 在 tool_result 中丢弃了 signature 时用于恢复。
|
||||
|
||||
Layer 2 (family): signature → model_family
|
||||
防止跨模型签名污染(Claude 签名不能用在 Gemini 上)。
|
||||
|
||||
Layer 3 (session): session_id → latest signature + message_count
|
||||
会话级追踪,支持 rewind 检测(用户删除消息后不会注入来自"未来"的签名)。
|
||||
|
||||
Legacy (text): SHA256(model + text) → signature
|
||||
向后兼容的 get_or_dummy() 接口。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._tool_sigs: dict[str, _CacheEntry] = {}
|
||||
self._families: dict[str, _CacheEntry] = {}
|
||||
self._sessions: dict[str, _CacheEntry] = {}
|
||||
self._text_sigs: dict[str, _CacheEntry] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# ===== Layer 1: Tool Use ID → Signature =====
|
||||
|
||||
def cache_tool_signature(self, tool_use_id: str, signature: str) -> None:
|
||||
"""缓存工具调用对应的 thinking signature。"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
with self._lock:
|
||||
self._tool_sigs[tool_use_id] = _CacheEntry(signature)
|
||||
if len(self._tool_sigs) > _TOOL_CACHE_LIMIT:
|
||||
self._prune(self._tool_sigs, limit=_TOOL_CACHE_LIMIT)
|
||||
|
||||
def get_tool_signature(self, tool_use_id: str) -> str | None:
|
||||
"""查找工具调用对应的 signature。"""
|
||||
with self._lock:
|
||||
entry = self._tool_sigs.get(tool_use_id)
|
||||
if entry is None:
|
||||
return None
|
||||
if entry.is_expired():
|
||||
self._tool_sigs.pop(tool_use_id, None)
|
||||
return None
|
||||
return entry.data
|
||||
|
||||
# ===== Layer 2: Signature → Model Family =====
|
||||
|
||||
def cache_thinking_family(self, signature: str, family: str) -> None:
|
||||
"""记录 signature 所属的模型家族。"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
with self._lock:
|
||||
self._families[signature] = _CacheEntry(family)
|
||||
if len(self._families) > _FAMILY_CACHE_LIMIT:
|
||||
self._prune(self._families, limit=_FAMILY_CACHE_LIMIT)
|
||||
|
||||
def get_signature_family(self, signature: str) -> str | None:
|
||||
"""查找 signature 所属的模型家族。"""
|
||||
with self._lock:
|
||||
entry = self._families.get(signature)
|
||||
if entry is None:
|
||||
return None
|
||||
if entry.is_expired():
|
||||
self._families.pop(signature, None)
|
||||
return None
|
||||
return entry.data
|
||||
|
||||
# ===== Layer 3: Session ID → Latest Signature =====
|
||||
|
||||
def cache_session_signature(
|
||||
self, session_id: str, signature: str, message_count: int = 0
|
||||
) -> None:
|
||||
"""存储会话的最新 thinking signature。
|
||||
|
||||
Rewind 检测:当 message_count 小于已缓存值时,说明用户删除了消息,
|
||||
强制更新签名以避免注入来自"未来"的签名。
|
||||
"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
|
||||
with self._lock:
|
||||
existing = self._sessions.get(session_id)
|
||||
should_store = True
|
||||
|
||||
if existing and not existing.is_expired():
|
||||
entry: _SessionEntry = existing.data
|
||||
if message_count < entry.message_count:
|
||||
# Rewind detected: 用户删除了消息,强制更新
|
||||
pass
|
||||
elif message_count == entry.message_count:
|
||||
# 同一轮消息:仅当新签名更长(更完整)时才替换
|
||||
should_store = len(signature) > len(entry.signature)
|
||||
# else: 正常递增,更新
|
||||
|
||||
if should_store:
|
||||
self._sessions[session_id] = _CacheEntry(_SessionEntry(signature, message_count))
|
||||
if len(self._sessions) > _SESSION_CACHE_LIMIT:
|
||||
self._prune(self._sessions, limit=_SESSION_CACHE_LIMIT)
|
||||
|
||||
def get_session_signature(self, session_id: str) -> str | None:
|
||||
"""获取会话的最新 thinking signature。"""
|
||||
with self._lock:
|
||||
entry = self._sessions.get(session_id)
|
||||
if entry is None:
|
||||
return None
|
||||
if entry.is_expired():
|
||||
self._sessions.pop(session_id, None)
|
||||
return None
|
||||
return entry.data.signature
|
||||
|
||||
# ===== Legacy: model:text → signature(向后兼容) =====
|
||||
|
||||
def get_or_dummy(self, model: str, thinking_text: str) -> str | None:
|
||||
"""Legacy: 根据 model + thinking_text 查找 signature。
|
||||
|
||||
Gemini 模型在未命中时返回 DUMMY_THOUGHT_SIGNATURE(跳过验证)。
|
||||
"""
|
||||
key = self._text_key(model, thinking_text)
|
||||
with self._lock:
|
||||
entry = self._text_sigs.get(key)
|
||||
if entry is not None:
|
||||
if entry.is_expired():
|
||||
self._text_sigs.pop(key, None)
|
||||
else:
|
||||
return entry.data
|
||||
if str(model).startswith("gemini-"):
|
||||
return DUMMY_THOUGHT_SIGNATURE
|
||||
return None
|
||||
|
||||
def cache(self, model: str, thinking_text: str, signature: str) -> None:
|
||||
"""Legacy: 缓存 model + thinking_text → signature。"""
|
||||
if len(signature) < MIN_SIGNATURE_LENGTH:
|
||||
return
|
||||
|
||||
key = self._text_key(model, thinking_text)
|
||||
with self._lock:
|
||||
if key in self._text_sigs:
|
||||
self._text_sigs[key] = _CacheEntry(signature)
|
||||
return
|
||||
|
||||
if len(self._text_sigs) >= _TEXT_CACHE_LIMIT:
|
||||
# FIFO 淘汰 1/4
|
||||
evict_n = max(1, _TEXT_CACHE_LIMIT // 4)
|
||||
for k in list(self._text_sigs.keys())[:evict_n]:
|
||||
self._text_sigs.pop(k, None)
|
||||
|
||||
self._text_sigs[key] = _CacheEntry(signature)
|
||||
|
||||
# ===== Utilities =====
|
||||
|
||||
@staticmethod
|
||||
def _text_key(model: str, thinking_text: str) -> str:
|
||||
# 使用 \x00 作为分隔符避免 model 中含 ':' 时的歧义
|
||||
content = f"{model}\x00{thinking_text}"
|
||||
return hashlib.sha256(content.encode("utf-8")).hexdigest()[:32]
|
||||
|
||||
@staticmethod
|
||||
def _prune(d: dict[str, _CacheEntry], *, limit: int | None = None) -> None:
|
||||
"""Remove expired entries and optionally enforce a size limit."""
|
||||
now = time.monotonic()
|
||||
expired = [k for k, v in d.items() if v.is_expired(now)]
|
||||
for k in expired:
|
||||
d.pop(k, None)
|
||||
|
||||
if limit is None or len(d) <= limit:
|
||||
return
|
||||
|
||||
# Evict oldest entries by created_at.
|
||||
excess = len(d) - limit
|
||||
for k, _entry in sorted(d.items(), key=lambda kv: kv[1].created_at)[:excess]:
|
||||
d.pop(k, None)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""清空所有缓存层(用于测试或手动重置)。"""
|
||||
with self._lock:
|
||||
self._tool_sigs.clear()
|
||||
self._families.clear()
|
||||
self._sessions.clear()
|
||||
self._text_sigs.clear()
|
||||
|
||||
|
||||
signature_cache = ThinkingSignatureCache()
|
||||
|
||||
__all__ = ["ThinkingSignatureCache", "signature_cache"]
|
||||
__all__ = ["MIN_SIGNATURE_LENGTH", "ThinkingSignatureCache", "signature_cache"]
|
||||
|
||||
@@ -0,0 +1,394 @@
|
||||
"""
|
||||
Provider 认证逻辑(OAuth / Service Account / Vertex AI)。
|
||||
|
||||
从 api/handlers/base/request_builder.py 迁移到 services 层,
|
||||
消除 services→api 的反向依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from sqlalchemy.orm import object_session
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.core.provider_auth_types import ProviderAuthInfo
|
||||
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh helpers
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _acquire_refresh_lock(key_id: str) -> tuple[Any, bool]:
|
||||
"""尝试获取 OAuth refresh 分布式锁。
|
||||
|
||||
返回 ``(redis_client | None, got_lock)``。调用方在刷新完成后
|
||||
必须调用 :func:`_release_refresh_lock` 释放锁。
|
||||
"""
|
||||
redis = await get_redis_client(require_redis=False)
|
||||
lock_key = f"provider_oauth_refresh_lock:{key_id}"
|
||||
got_lock = False
|
||||
if redis is not None:
|
||||
try:
|
||||
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
|
||||
except Exception:
|
||||
got_lock = False
|
||||
return redis, got_lock
|
||||
|
||||
|
||||
async def _release_refresh_lock(redis: Any, key_id: str) -> None:
|
||||
"""释放 OAuth refresh 分布式锁(best-effort)。"""
|
||||
if redis is not None:
|
||||
try:
|
||||
await redis.delete(f"provider_oauth_refresh_lock:{key_id}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _persist_refreshed_token(
|
||||
key: Any,
|
||||
access_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> None:
|
||||
"""将刷新后的 access_token 和 auth_config 持久化到数据库。"""
|
||||
key.api_key = crypto_service.encrypt(access_token)
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
sess = object_session(key)
|
||||
if sess is not None:
|
||||
sess.add(key)
|
||||
sess.commit()
|
||||
else:
|
||||
logger.warning(
|
||||
"[OAUTH_REFRESH] key {} refreshed but cannot persist (no session); "
|
||||
"next request will refresh again",
|
||||
key.id,
|
||||
)
|
||||
|
||||
|
||||
def _get_proxy_config(key: Any, endpoint: Any = None) -> Any:
|
||||
"""获取有效代理配置(Key 级别优先于 Provider 级别)。"""
|
||||
try:
|
||||
from src.services.proxy_node.resolver import resolve_effective_proxy
|
||||
|
||||
provider = getattr(key, "provider", None) or (
|
||||
getattr(endpoint, "provider", None) if endpoint else None
|
||||
)
|
||||
provider_proxy = getattr(provider, "proxy", None)
|
||||
key_proxy = getattr(key, "proxy", None)
|
||||
return resolve_effective_proxy(provider_proxy, key_proxy)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Provider-specific refresh implementations
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _refresh_kiro_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Kiro OAuth refresh: validate + call Kiro-specific refresh endpoint."""
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import (
|
||||
refresh_access_token,
|
||||
validate_refresh_token,
|
||||
)
|
||||
|
||||
cfg = KiroAuthConfig.from_dict(token_meta or {})
|
||||
if not (cfg.refresh_token or "").strip():
|
||||
raise InvalidRequestException(
|
||||
"Kiro auth_config missing refresh_token; please re-import credentials."
|
||||
)
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
access_token, new_cfg = await refresh_access_token(
|
||||
cfg,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
new_meta = new_cfg.to_dict()
|
||||
new_meta["updated_at"] = int(time.time())
|
||||
|
||||
_persist_refreshed_token(key, access_token, new_meta)
|
||||
return new_meta
|
||||
|
||||
|
||||
async def _refresh_generic_oauth_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
template: Any,
|
||||
provider_type: str,
|
||||
refresh_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Generic OAuth refresh via template (Codex, Antigravity, ClaudeCode, etc.)."""
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
scopes = getattr(template.oauth, "scopes", None) or []
|
||||
scope_str = " ".join(scopes) if scopes else ""
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if scope_str:
|
||||
body["scope"] = scope_str
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if scope_str:
|
||||
form["scope"] = scope_str
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=30.0,
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
_persist_refreshed_token(key, access_token, token_meta)
|
||||
else:
|
||||
logger.warning(
|
||||
"OAuth token refresh failed: provider={}, key_id={}, status={}",
|
||||
provider_type,
|
||||
getattr(key, "id", "?"),
|
||||
resp.status_code,
|
||||
)
|
||||
|
||||
return token_meta
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Service Account 认证支持
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def get_provider_auth(
|
||||
endpoint: "ProviderEndpoint",
|
||||
key: "ProviderAPIKey",
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> ProviderAuthInfo | None:
|
||||
"""
|
||||
获取 Provider 的认证信息
|
||||
|
||||
对于标准 API Key,返回 None(由 build_headers 自动处理)。
|
||||
对于 Service Account,异步获取 Access Token 并返回认证信息。
|
||||
|
||||
Args:
|
||||
endpoint: 端点配置
|
||||
key: Provider API Key
|
||||
|
||||
Returns:
|
||||
Service Account 场景: ProviderAuthInfo 对象(包含认证信息和解密后的配置)
|
||||
API Key 场景: None(由 build_headers 处理)
|
||||
|
||||
Raises:
|
||||
InvalidRequestException: 认证配置无效或认证失败
|
||||
"""
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
|
||||
auth_type = getattr(key, "auth_type", "api_key")
|
||||
|
||||
if auth_type == "oauth":
|
||||
# OAuth token 保存在 key.api_key(加密),refresh_token/expires_at 等在 auth_config(加密 JSON)中。
|
||||
# 在请求前做一次懒刷新:接近过期时刷新 access_token,并用 Redis lock 避免并发风暴。
|
||||
encrypted_auth_config = getattr(key, "auth_config", None)
|
||||
if encrypted_auth_config:
|
||||
try:
|
||||
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
||||
token_meta = json.loads(decrypted_config)
|
||||
except Exception:
|
||||
token_meta = {}
|
||||
else:
|
||||
token_meta = {}
|
||||
|
||||
expires_at = token_meta.get("expires_at")
|
||||
refresh_token = token_meta.get("refresh_token")
|
||||
provider_type = str(token_meta.get("provider_type") or "")
|
||||
cached_access_token = str(token_meta.get("access_token") or "").strip()
|
||||
|
||||
# 120s skew (or force refresh when upstream returns 401)
|
||||
should_refresh = False
|
||||
try:
|
||||
if expires_at is not None:
|
||||
should_refresh = int(time.time()) >= int(expires_at) - 120
|
||||
except Exception:
|
||||
should_refresh = False
|
||||
|
||||
if force_refresh:
|
||||
should_refresh = True
|
||||
|
||||
# Kiro 特殊处理:如果没有缓存的 access_token 或 key.api_key 是占位符,强制刷新
|
||||
if provider_type == "kiro" and not should_refresh:
|
||||
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:
|
||||
try:
|
||||
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
||||
from src.core.provider_templates.types import ProviderType
|
||||
|
||||
try:
|
||||
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
|
||||
except Exception:
|
||||
template = None
|
||||
|
||||
redis, got_lock = await _acquire_refresh_lock(key.id)
|
||||
if got_lock or redis is None:
|
||||
try:
|
||||
if provider_type == ProviderType.KIRO.value:
|
||||
token_meta = await _refresh_kiro_token(key, endpoint, token_meta)
|
||||
elif template:
|
||||
token_meta = await _refresh_generic_oauth_token(
|
||||
key, endpoint, template, provider_type, refresh_token, token_meta
|
||||
)
|
||||
finally:
|
||||
if got_lock:
|
||||
await _release_refresh_lock(redis, key.id)
|
||||
except Exception:
|
||||
# 刷新失败不阻断请求;后续由上游返回 401 再触发管理端处理
|
||||
pass
|
||||
|
||||
# 获取最终使用的 access_token
|
||||
# Kiro 优先使用 token_meta 中缓存的 access_token(刷新后会更新到 token_meta)
|
||||
if provider_type == "kiro":
|
||||
refreshed_token = str(token_meta.get("access_token") or "").strip()
|
||||
effective_token = refreshed_token or crypto_service.decrypt(key.api_key)
|
||||
else:
|
||||
effective_token = crypto_service.decrypt(key.api_key)
|
||||
|
||||
decrypted_auth_config: dict[str, Any] | None = None
|
||||
if isinstance(token_meta, dict) and token_meta:
|
||||
decrypted_auth_config = token_meta
|
||||
|
||||
return ProviderAuthInfo(
|
||||
auth_header="Authorization",
|
||||
auth_value=f"Bearer {effective_token}",
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
if auth_type == "vertex_ai":
|
||||
from src.core.vertex_auth import VertexAuthError, VertexAuthService
|
||||
|
||||
try:
|
||||
# 优先从 auth_config 读取,兼容从 api_key 读取(过渡期)
|
||||
encrypted_auth_config = getattr(key, "auth_config", None)
|
||||
if encrypted_auth_config:
|
||||
# auth_config 可能是加密字符串或未加密的 dict
|
||||
if isinstance(encrypted_auth_config, dict):
|
||||
# 已经是 dict,直接使用(兼容未加密存储的情况)
|
||||
sa_json = encrypted_auth_config
|
||||
else:
|
||||
# 是加密字符串,需要解密
|
||||
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
||||
sa_json = json.loads(decrypted_config)
|
||||
else:
|
||||
# 兼容旧数据:从 api_key 读取
|
||||
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||
# 检查是否是占位符(表示 auth_config 丢失)
|
||||
if decrypted_key == "__placeholder__":
|
||||
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
|
||||
sa_json = json.loads(decrypted_key)
|
||||
|
||||
if not isinstance(sa_json, dict):
|
||||
raise InvalidRequestException("Service Account JSON 无效,请重新添加该密钥。")
|
||||
|
||||
# 获取 Access Token(注入代理配置,core 层不依赖 services)
|
||||
from src.services.proxy_node.resolver import build_proxy_client_kwargs
|
||||
|
||||
service = VertexAuthService(sa_json)
|
||||
access_token = await service.get_access_token(
|
||||
httpx_client_kwargs=build_proxy_client_kwargs(timeout=30),
|
||||
)
|
||||
|
||||
# Vertex AI 使用 Bearer token
|
||||
return ProviderAuthInfo(
|
||||
auth_header="Authorization",
|
||||
auth_value=f"Bearer {access_token}",
|
||||
decrypted_auth_config=sa_json,
|
||||
)
|
||||
except InvalidRequestException:
|
||||
raise
|
||||
except VertexAuthError as e:
|
||||
raise InvalidRequestException(f"Vertex AI 认证失败:{e}")
|
||||
except json.JSONDecodeError:
|
||||
raise InvalidRequestException("Service Account JSON 格式无效,请重新添加该密钥。")
|
||||
except Exception:
|
||||
raise InvalidRequestException("Vertex AI 认证失败,请检查 Key 的 auth_config")
|
||||
|
||||
# 其他认证类型可在此扩展
|
||||
# elif auth_type == "oauth2":
|
||||
# ...
|
||||
|
||||
# 标准 API Key:返回 None,由 build_headers 处理
|
||||
return None
|
||||
@@ -64,7 +64,7 @@ async def resolve_oauth_access_token(
|
||||
"""
|
||||
|
||||
# Local import to avoid circular imports during app startup.
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.services.provider.auth import get_provider_auth
|
||||
|
||||
# Build detached key-like objects for get_provider_auth().
|
||||
provider_obj = (
|
||||
|
||||
@@ -28,8 +28,6 @@ class ConcurrencyManager:
|
||||
|
||||
_instance: ConcurrencyManager | None = None
|
||||
_redis: aioredis.Redis | None = None
|
||||
_key_rpm_bucket_seconds: int = 60
|
||||
_key_rpm_key_ttl_seconds: int = 120 # 2 分钟,足够覆盖当前分钟与边界
|
||||
|
||||
def __new__(cls) -> "ConcurrencyManager":
|
||||
"""单例模式"""
|
||||
@@ -42,13 +40,18 @@ class ConcurrencyManager:
|
||||
if hasattr(self, "_memory_initialized"):
|
||||
return
|
||||
|
||||
from src.config.settings import config
|
||||
|
||||
self._key_rpm_bucket_seconds: int = config.rpm_bucket_seconds
|
||||
self._key_rpm_key_ttl_seconds: int = config.rpm_key_ttl_seconds
|
||||
|
||||
self._memory_lock: asyncio.Lock = asyncio.Lock()
|
||||
# Key RPM 计数器:{key_id: (bucket, count)},bucket = floor(now / 60)
|
||||
self._memory_key_rpm_counts: dict[str, tuple[int, int]] = {}
|
||||
self._owns_redis: bool = False
|
||||
self._last_cleanup_bucket: int = 0 # 上次清理时的 bucket,用于定期清理过期数据
|
||||
self._last_cleanup_time: float = 0 # 上次清理的时间戳,用于强制定期清理
|
||||
self._cleanup_interval_seconds: int = 300 # 强制清理间隔(5 分钟)
|
||||
self._cleanup_interval_seconds: int = config.rpm_cleanup_interval_seconds
|
||||
self._cleanup_task: asyncio.Task | None = None # 后台清理任务
|
||||
|
||||
# 内存模式下的最大条目限制,防止内存泄漏(支持环境变量覆盖)
|
||||
@@ -130,16 +133,14 @@ class ConcurrencyManager:
|
||||
self._redis = None
|
||||
self._owns_redis = False
|
||||
|
||||
@classmethod
|
||||
def _get_rpm_bucket(cls, now_ts: float | None = None) -> int:
|
||||
def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
|
||||
"""获取当前 RPM 计数桶(按分钟)"""
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
return int(ts // cls._key_rpm_bucket_seconds)
|
||||
return int(ts // self._key_rpm_bucket_seconds)
|
||||
|
||||
@classmethod
|
||||
def _get_key_key(cls, key_id: str, bucket: int | None = None) -> str:
|
||||
def _get_key_key(self, key_id: str, bucket: int | None = None) -> str:
|
||||
"""获取 ProviderAPIKey RPM 计数的 Redis Key(按分钟桶)"""
|
||||
b = bucket if bucket is not None else cls._get_rpm_bucket()
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:key:{key_id}:{b}"
|
||||
|
||||
def _get_memory_key_rpm_count(self, key_id: str, bucket: int) -> int:
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
"""调度系统(候选/排序/亲和性/并发检查)。
|
||||
|
||||
原 `src.services.cache` 中与调度相关的代码已迁移到此包;
|
||||
`src.services.cache` 现在只保留通用缓存 backend/sync/*_cache。
|
||||
"""
|
||||
|
||||
from src.services.scheduling.affinity_manager import CacheAffinityManager, get_affinity_manager
|
||||
from src.services.scheduling.aware_scheduler import (
|
||||
CacheAwareScheduler,
|
||||
ConcurrencySnapshot,
|
||||
ProviderCandidate,
|
||||
get_cache_aware_scheduler,
|
||||
)
|
||||
from src.services.scheduling.candidate_builder import CandidateBuilder
|
||||
from src.services.scheduling.candidate_sorter import CandidateSorter
|
||||
from src.services.scheduling.concurrency_checker import ConcurrencyChecker
|
||||
from src.services.scheduling.scheduling_config import SchedulingConfig
|
||||
|
||||
__all__ = [
|
||||
"CacheAffinityManager",
|
||||
"CandidateBuilder",
|
||||
"CandidateSorter",
|
||||
"CacheAwareScheduler",
|
||||
"ConcurrencyChecker",
|
||||
"ConcurrencySnapshot",
|
||||
"ProviderCandidate",
|
||||
"SchedulingConfig",
|
||||
"get_affinity_manager",
|
||||
"get_cache_aware_scheduler",
|
||||
]
|
||||
+44
-27
@@ -31,7 +31,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.services.scheduling.protocols import (
|
||||
CacheAffinityManagerProtocol,
|
||||
CandidateBuilderProtocol,
|
||||
CandidateSorterProtocol,
|
||||
ConcurrencyCheckerProtocol,
|
||||
)
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -47,32 +55,32 @@ from src.models.database import (
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.cache.affinity_manager import (
|
||||
CacheAffinityManager,
|
||||
get_affinity_manager,
|
||||
)
|
||||
from src.services.cache.candidate_builder import (
|
||||
CandidateBuilder,
|
||||
)
|
||||
from src.services.cache.candidate_builder import (
|
||||
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
|
||||
)
|
||||
from src.services.cache.candidate_sorter import CandidateSorter
|
||||
from src.services.cache.concurrency_checker import ConcurrencyChecker
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
from src.services.cache.restriction_checker import get_effective_restrictions
|
||||
from src.services.cache.scheduling_config import SchedulingConfig
|
||||
from src.services.cache.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
|
||||
from src.services.cache.schemas import ProviderCandidate as ProviderCandidate # re-export
|
||||
from src.services.cache.utils import affinity_hash as _affinity_hash # re-export compat
|
||||
from src.services.cache.utils import (
|
||||
release_db_connection_before_await,
|
||||
)
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.rate_limit.adaptive_reservation import (
|
||||
get_adaptive_reservation_manager,
|
||||
)
|
||||
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
|
||||
from src.services.scheduling.affinity_manager import (
|
||||
CacheAffinityManager,
|
||||
get_affinity_manager,
|
||||
)
|
||||
from src.services.scheduling.candidate_builder import (
|
||||
CandidateBuilder,
|
||||
)
|
||||
from src.services.scheduling.candidate_builder import (
|
||||
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
|
||||
)
|
||||
from src.services.scheduling.candidate_sorter import CandidateSorter
|
||||
from src.services.scheduling.concurrency_checker import ConcurrencyChecker
|
||||
from src.services.scheduling.restriction_checker import get_effective_restrictions
|
||||
from src.services.scheduling.scheduling_config import SchedulingConfig
|
||||
from src.services.scheduling.schemas import ConcurrencySnapshot as ConcurrencySnapshot # re-export
|
||||
from src.services.scheduling.schemas import ProviderCandidate as ProviderCandidate # re-export
|
||||
from src.services.scheduling.utils import affinity_hash as _affinity_hash # re-export compat
|
||||
from src.services.scheduling.utils import (
|
||||
release_db_connection_before_await,
|
||||
)
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
|
||||
@@ -102,6 +110,11 @@ class CacheAwareScheduler:
|
||||
redis_client: Any | None = None,
|
||||
priority_mode: str | None = None,
|
||||
scheduling_mode: str | None = None,
|
||||
*,
|
||||
candidate_builder: CandidateBuilderProtocol | None = None,
|
||||
candidate_sorter: CandidateSorterProtocol | None = None,
|
||||
concurrency_checker: ConcurrencyCheckerProtocol | None = None,
|
||||
affinity_manager: CacheAffinityManagerProtocol | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
初始化调度器
|
||||
@@ -117,9 +130,9 @@ class CacheAwareScheduler:
|
||||
self.redis = redis_client
|
||||
self._config = SchedulingConfig(priority_mode, scheduling_mode)
|
||||
|
||||
# 异步子组件(将在第一次使用时初始化)
|
||||
self._affinity_manager: CacheAffinityManager | None = None
|
||||
self._concurrency_checker: ConcurrencyChecker | None = None
|
||||
# 异步子组件(将在第一次使用时初始化,可通过构造函数注入)
|
||||
self._affinity_manager: CacheAffinityManagerProtocol | None = affinity_manager
|
||||
self._concurrency_checker: ConcurrencyCheckerProtocol | None = concurrency_checker
|
||||
|
||||
self._metrics: dict[str, Any] = {
|
||||
"total_batches": 0,
|
||||
@@ -139,9 +152,13 @@ class CacheAwareScheduler:
|
||||
"last_reservation_result": None,
|
||||
}
|
||||
|
||||
# 初始化子模块(不传 self,解除反向引用)
|
||||
self._candidate_sorter = CandidateSorter(self._config)
|
||||
self._candidate_builder = CandidateBuilder(self._candidate_sorter)
|
||||
# 初始化子模块(不传 self,解除反向引用,可通过构造函数注入)
|
||||
self._candidate_sorter: CandidateSorterProtocol = candidate_sorter or CandidateSorter(
|
||||
self._config
|
||||
)
|
||||
self._candidate_builder: CandidateBuilderProtocol = candidate_builder or CandidateBuilder(
|
||||
self._candidate_sorter
|
||||
)
|
||||
|
||||
# ── 属性代理(保持外部访问兼容性)──────────────────────────
|
||||
|
||||
+6
-6
@@ -28,15 +28,15 @@ from src.models.database import (
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.cache.quota_skipper import is_key_quota_exhausted
|
||||
from src.services.cache.utils import release_db_connection_before_await
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.scheduling.quota_skipper import is_key_quota_exhausted
|
||||
from src.services.scheduling.utils import release_db_connection_before_await
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import GlobalModel
|
||||
from src.services.cache.candidate_sorter import CandidateSorter
|
||||
from src.services.cache.schemas import ProviderCandidate
|
||||
from src.services.scheduling.protocols import CandidateSorterProtocol
|
||||
from src.services.scheduling.schemas import ProviderCandidate
|
||||
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
|
||||
@@ -60,7 +60,7 @@ def _sort_endpoints_by_family_priority(
|
||||
class CandidateBuilder:
|
||||
"""候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。"""
|
||||
|
||||
def __init__(self, candidate_sorter: CandidateSorter) -> None:
|
||||
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
|
||||
self._sorter = candidate_sorter
|
||||
|
||||
def _query_providers(
|
||||
@@ -380,7 +380,7 @@ class CandidateBuilder:
|
||||
Returns:
|
||||
候选列表
|
||||
"""
|
||||
from src.services.cache.schemas import ProviderCandidate
|
||||
from src.services.scheduling.schemas import ProviderCandidate
|
||||
|
||||
candidates: list[ProviderCandidate] = []
|
||||
client_format_str = normalize_endpoint_signature(client_format)
|
||||
+3
-3
@@ -13,15 +13,15 @@ import random
|
||||
from collections import defaultdict
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from src.services.cache.scheduling_config import SchedulingConfig
|
||||
from src.services.cache.utils import affinity_hash
|
||||
from src.services.scheduling.scheduling_config import SchedulingConfig
|
||||
from src.services.scheduling.utils import affinity_hash
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.services.cache.schemas import ProviderCandidate
|
||||
from src.services.scheduling.schemas import ProviderCandidate
|
||||
|
||||
|
||||
class CandidateSorter:
|
||||
+1
-1
@@ -11,9 +11,9 @@ from typing import Any
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.services.cache.schemas import ConcurrencySnapshot
|
||||
from src.services.rate_limit.adaptive_reservation import AdaptiveReservationManager
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
from src.services.scheduling.schemas import ConcurrencySnapshot
|
||||
|
||||
|
||||
class ConcurrencyChecker:
|
||||
@@ -0,0 +1,137 @@
|
||||
"""调度/候选子组件的协议接口。
|
||||
|
||||
目的:
|
||||
- 用 `Protocol` 固化 CacheAwareScheduler 的子组件契约
|
||||
- 便于单测注入 stub/mocks,减少对具体实现类的耦合
|
||||
|
||||
说明:这里的协议面向“调度器内部协作”,因此保留了部分 `_` 前缀方法。
|
||||
后续如果要对外暴露更稳定的 API,可再抽出无下划线的 facade。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import GlobalModel, Provider, ProviderAPIKey
|
||||
from src.services.scheduling.affinity_manager import CacheAffinity
|
||||
from src.services.scheduling.schemas import ConcurrencySnapshot, ProviderCandidate
|
||||
|
||||
|
||||
class CandidateSorterProtocol(Protocol):
|
||||
def _apply_priority_mode_sort(
|
||||
self,
|
||||
candidates: list[ProviderCandidate],
|
||||
db: Session,
|
||||
affinity_key: str | None = None,
|
||||
api_format: str | None = None,
|
||||
) -> list[ProviderCandidate]: ...
|
||||
|
||||
def _apply_load_balance(
|
||||
self, candidates: list[ProviderCandidate], api_format: str | None = None
|
||||
) -> list[ProviderCandidate]: ...
|
||||
|
||||
def shuffle_keys_by_internal_priority(
|
||||
self,
|
||||
keys: list[ProviderAPIKey],
|
||||
affinity_key: str | None = None,
|
||||
use_random: bool = False,
|
||||
) -> list[ProviderAPIKey]: ...
|
||||
|
||||
|
||||
class CandidateBuilderProtocol(Protocol):
|
||||
def _query_providers(
|
||||
self,
|
||||
db: Session,
|
||||
provider_offset: int = 0,
|
||||
provider_limit: int | None = None,
|
||||
) -> list[Provider]: ...
|
||||
|
||||
async def _build_candidates(
|
||||
self,
|
||||
db: Session,
|
||||
providers: list[Provider],
|
||||
client_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str | None,
|
||||
model_mappings: list[str] | None = None,
|
||||
max_candidates: int | None = None,
|
||||
is_stream: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
global_conversion_enabled: bool = True,
|
||||
) -> list[ProviderCandidate]: ...
|
||||
|
||||
async def _check_model_support(
|
||||
self,
|
||||
db: Session,
|
||||
provider: Provider,
|
||||
model_name: str,
|
||||
api_format: str | None = None,
|
||||
is_stream: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
) -> tuple[bool, str | None, list[str] | None, set[str] | None]: ...
|
||||
|
||||
async def _check_model_support_for_global_model(
|
||||
self,
|
||||
db: Session,
|
||||
provider: Provider,
|
||||
global_model: GlobalModel,
|
||||
model_name: str,
|
||||
api_format: str | None = None,
|
||||
is_stream: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
) -> tuple[bool, str | None, list[str] | None, set[str] | None]: ...
|
||||
|
||||
def _check_key_availability(
|
||||
self,
|
||||
key: ProviderAPIKey,
|
||||
api_format: str | None,
|
||||
model_name: str,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
model_mappings: list[str] | None = None,
|
||||
candidate_models: set[str] | None = None,
|
||||
*,
|
||||
provider_type: str | None = None,
|
||||
) -> tuple[bool, str | None, str | None]: ...
|
||||
|
||||
|
||||
class ConcurrencyCheckerProtocol(Protocol):
|
||||
async def check_available(
|
||||
self,
|
||||
key: ProviderAPIKey,
|
||||
is_cached_user: bool = False,
|
||||
) -> tuple[bool, ConcurrencySnapshot]: ...
|
||||
|
||||
def get_reservation_stats(self) -> dict[str, Any]: ...
|
||||
|
||||
|
||||
class CacheAffinityManagerProtocol(Protocol):
|
||||
async def get_affinity(
|
||||
self, affinity_key: str, api_format: str, model_name: str
|
||||
) -> CacheAffinity | None: ...
|
||||
|
||||
async def set_affinity(
|
||||
self,
|
||||
affinity_key: str,
|
||||
provider_id: str,
|
||||
endpoint_id: str,
|
||||
key_id: str,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
supports_caching: bool = True,
|
||||
ttl: int | None = None,
|
||||
) -> None: ...
|
||||
|
||||
async def invalidate_affinity(
|
||||
self,
|
||||
affinity_key: str,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
key_id: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
endpoint_id: str | None = None,
|
||||
) -> None: ...
|
||||
|
||||
def get_stats(self) -> dict[str, Any]: ...
|
||||
@@ -72,7 +72,9 @@ class CacheWarmupService:
|
||||
"""预热管理员仪表盘统计缓存"""
|
||||
db = None
|
||||
try:
|
||||
from src.api.dashboard.routes import AdminDashboardStatsAdapter
|
||||
from src.api.dashboard.routes import ( # TODO(arch): 提取 dashboard 统计计算到 services 层
|
||||
AdminDashboardStatsAdapter,
|
||||
)
|
||||
from src.models.database import User as DBUser
|
||||
|
||||
db = create_session()
|
||||
@@ -128,7 +130,9 @@ class CacheWarmupService:
|
||||
"""预热每日统计缓存"""
|
||||
db = None
|
||||
try:
|
||||
from src.api.dashboard.routes import DashboardDailyStatsAdapter
|
||||
from src.api.dashboard.routes import ( # TODO(arch): 提取 dashboard 统计计算到 services 层
|
||||
DashboardDailyStatsAdapter,
|
||||
)
|
||||
from src.models.database import User as DBUser
|
||||
|
||||
db = create_session()
|
||||
|
||||
@@ -16,11 +16,6 @@ from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import (
|
||||
@@ -33,8 +28,14 @@ from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.core.provider_auth_types import ProviderAuthInfo
|
||||
from src.core.video_utils import (
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.database import create_session
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
|
||||
from src.services.provider.auth import get_provider_auth
|
||||
from src.services.task.service import TaskService
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any, AsyncIterator, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
|
||||
class AttemptKind(str, Enum):
|
||||
|
||||
@@ -3,8 +3,8 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.candidate.schema import CandidateKey
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
from .protocol import AttemptKind, AttemptResult
|
||||
|
||||
|
||||
@@ -30,10 +30,6 @@ from src.models.database import (
|
||||
User,
|
||||
VideoTask,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import (
|
||||
CacheAwareScheduler,
|
||||
get_cache_aware_scheduler,
|
||||
)
|
||||
from src.services.candidate.failover import FailoverEngine
|
||||
from src.services.candidate.policy import RetryPolicy, SkipPolicy
|
||||
from src.services.candidate.recorder import CandidateRecorder
|
||||
@@ -43,6 +39,10 @@ from src.services.orchestration.request_dispatcher import RequestDispatcher
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
from src.services.request.result import RequestMetadata
|
||||
from src.services.scheduling.aware_scheduler import (
|
||||
CacheAwareScheduler,
|
||||
get_cache_aware_scheduler,
|
||||
)
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.task.context import TaskMode
|
||||
from src.services.task.exceptions import TaskNotFoundError
|
||||
@@ -1560,7 +1560,6 @@ class TaskService:
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.core.api_format import (
|
||||
build_upstream_headers_for_endpoint,
|
||||
@@ -1569,6 +1568,7 @@ class TaskService:
|
||||
)
|
||||
from src.core.api_format.conversion.internal_video import VideoStatus
|
||||
from src.core.crypto import crypto_service
|
||||
from src.services.provider.auth import get_provider_auth
|
||||
from src.services.provider.transport import build_provider_url
|
||||
|
||||
try:
|
||||
|
||||
@@ -28,7 +28,7 @@ from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.utils import filter_proxy_response_headers
|
||||
from src.core.api_format import filter_response_headers as filter_proxy_response_headers
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, User
|
||||
from src.services.request.result import RequestResult
|
||||
|
||||
@@ -13,10 +13,9 @@ from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.parsers import get_parser_for_format
|
||||
from src.api.handlers.base.response_parser import StreamStats
|
||||
from src.core.exceptions import EmptyStreamException
|
||||
from src.core.logger import logger
|
||||
from src.core.stream_types import StreamStats, get_parser_for_format
|
||||
from src.database.database import create_session
|
||||
from src.models.database import ApiKey, User
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
@@ -0,0 +1,299 @@
|
||||
"""
|
||||
消息遥测记录器。
|
||||
|
||||
从 api/handlers/base/base_handler.py 迁移到 services 层,
|
||||
消除 services→api 的反向依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.services.system.audit import audit_service
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
class MessageTelemetry:
|
||||
"""
|
||||
负责记录 Usage/Audit,避免处理器里重复代码。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, db: Session, user: Any, api_key: Any, request_id: str, client_ip: str
|
||||
) -> None:
|
||||
self.db = db
|
||||
self.user = user
|
||||
self.api_key = api_key
|
||||
self.request_id = request_id
|
||||
self.client_ip = client_ip
|
||||
|
||||
async def calculate_cost(
|
||||
self,
|
||||
provider: str,
|
||||
model: str,
|
||||
*,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
) -> float:
|
||||
input_price, output_price = await UsageService.get_model_price_async(
|
||||
self.db, provider, model
|
||||
)
|
||||
_, _, _, _, _, _, total_cost = UsageService.calculate_cost(
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
input_price,
|
||||
output_price,
|
||||
cache_creation_tokens,
|
||||
cache_read_tokens,
|
||||
*await UsageService.get_cache_prices_async(self.db, provider, model, input_price),
|
||||
)
|
||||
return total_cost
|
||||
|
||||
async def record_success(
|
||||
self,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
response_body: Any,
|
||||
response_headers: dict[str, Any],
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
is_stream: bool = False,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 时间指标
|
||||
first_byte_time_ms: int | None = None, # 首字时间/TTFB
|
||||
# Provider 侧追踪信息(用于记录真实成本)
|
||||
provider_id: str | None = None,
|
||||
provider_endpoint_id: str | None = None,
|
||||
provider_api_key_id: str | None = None,
|
||||
api_format: str | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: str | None = None, # 端点原生 API 格式
|
||||
has_format_conversion: bool = False, # 是否发生了格式转换
|
||||
# 模型映射信息
|
||||
target_model: str | None = None,
|
||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||
response_metadata: dict[str, Any] | None = None,
|
||||
# 请求元数据(用于性能与调试记录)
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
) -> float:
|
||||
metadata = response_metadata
|
||||
if request_metadata:
|
||||
merged = dict(request_metadata)
|
||||
if response_metadata:
|
||||
merged.setdefault("response", response_metadata)
|
||||
metadata = merged
|
||||
|
||||
usage = await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
provider=provider,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
request_type="chat",
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms, # 传递首字时间
|
||||
status_code=status_code,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers or {},
|
||||
response_headers=response_headers,
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body,
|
||||
request_id=self.request_id,
|
||||
# Provider 侧追踪信息(用于记录真实成本)
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=provider_endpoint_id,
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
# 模型映射信息
|
||||
target_model=target_model,
|
||||
# Provider 响应元数据/请求元数据
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
total_cost = float(getattr(usage, "total_cost_usd", 0.0) or 0.0)
|
||||
|
||||
if self.user and self.api_key:
|
||||
audit_service.log_api_request(
|
||||
db=self.db,
|
||||
user_id=self.user.id,
|
||||
api_key_id=self.api_key.id,
|
||||
request_id=self.request_id,
|
||||
model=model,
|
||||
provider=provider,
|
||||
success=True,
|
||||
ip_address=self.client_ip,
|
||||
status_code=status_code,
|
||||
input_tokens=getattr(usage, "input_tokens", input_tokens),
|
||||
output_tokens=getattr(usage, "output_tokens", output_tokens),
|
||||
cost_usd=total_cost,
|
||||
)
|
||||
|
||||
return total_cost
|
||||
|
||||
async def record_failure(
|
||||
self,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
# 模型映射信息
|
||||
target_model: str | None = None,
|
||||
# 请求元数据(用于性能与调试记录)
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录失败请求
|
||||
|
||||
注意:Provider 链路信息(provider_id, endpoint_id, key_id)不在此处记录,
|
||||
因为 RequestCandidate 表已经记录了完整的请求链路追踪信息。
|
||||
|
||||
Args:
|
||||
input_tokens: 预估输入 tokens(来自 message_start,用于中断请求的成本估算)
|
||||
output_tokens: 预估输出 tokens(来自已收到的内容)
|
||||
cache_creation_tokens: 缓存创建 tokens
|
||||
cache_read_tokens: 缓存读取 tokens
|
||||
response_body: 响应体(如果有部分响应)
|
||||
response_headers: 响应头(Provider 返回的原始响应头)
|
||||
client_response_headers: 返回给客户端的响应头
|
||||
target_model: 映射后的目标模型名(如果发生了映射)
|
||||
"""
|
||||
provider_name = provider or "unknown"
|
||||
if provider_name == "unknown":
|
||||
logger.warning(
|
||||
"[Telemetry] Recording failure with unknown provider (request_id={})",
|
||||
self.request_id,
|
||||
)
|
||||
|
||||
await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
request_type="chat",
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=error_message,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers or {},
|
||||
response_headers=response_headers or {},
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body or {"error": error_message},
|
||||
request_id=self.request_id,
|
||||
# 模型映射信息
|
||||
target_model=target_model,
|
||||
# 请求元数据
|
||||
metadata=request_metadata,
|
||||
)
|
||||
|
||||
async def record_cancelled(
|
||||
self,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
response_time_ms: int,
|
||||
first_byte_time_ms: int | None,
|
||||
status_code: int,
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
target_model: str | None = None,
|
||||
# 请求元数据(用于性能与调试记录)
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录客户端取消的请求
|
||||
|
||||
客户端主动断开连接不算系统失败,使用 cancelled 状态。
|
||||
"""
|
||||
provider_name = provider or "unknown"
|
||||
|
||||
await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
request_type="chat",
|
||||
api_format=api_format,
|
||||
endpoint_api_format=endpoint_api_format,
|
||||
has_format_conversion=has_format_conversion,
|
||||
is_stream=is_stream,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=first_byte_time_ms,
|
||||
status_code=status_code,
|
||||
status="cancelled",
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
provider_request_headers=provider_request_headers or {},
|
||||
response_headers=response_headers or {},
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body or {},
|
||||
request_id=self.request_id,
|
||||
target_model=target_model,
|
||||
metadata=request_metadata,
|
||||
)
|
||||
@@ -8,11 +8,11 @@ import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.base_handler import MessageTelemetry
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.services.usage.events import UsageEventType, build_usage_event
|
||||
from src.services.usage.telemetry import MessageTelemetry
|
||||
|
||||
|
||||
class TelemetryWriter(ABC):
|
||||
|
||||
@@ -416,7 +416,7 @@ class UserService:
|
||||
通过 GlobalModel + Model 关联查询用户可用模型
|
||||
逻辑:使用 AccessRestrictions 统一处理 allowed_providers 和 allowed_models 限制
|
||||
"""
|
||||
from src.api.base.models_service import AccessRestrictions
|
||||
from src.core.access_restrictions import AccessRestrictions
|
||||
|
||||
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)
|
||||
|
||||
@@ -2,10 +2,12 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
|
||||
from src.services.provider.adapters.antigravity.envelope import (
|
||||
@@ -110,7 +112,9 @@ def test_handle_sse_event_unwraps_for_antigravity() -> None:
|
||||
}
|
||||
|
||||
with patch.object(handler, "_process_event_data") as mock_process:
|
||||
handler._handle_sse_event(ctx, None, json.dumps(v1_data), record_chunk=False)
|
||||
cast(CliHandlerProtocol, handler)._handle_sse_event(
|
||||
ctx, None, json.dumps(v1_data), record_chunk=False
|
||||
)
|
||||
|
||||
assert mock_process.call_count == 1
|
||||
passed_data = mock_process.call_args[0][2]
|
||||
@@ -120,7 +124,7 @@ def test_handle_sse_event_unwraps_for_antigravity() -> None:
|
||||
|
||||
|
||||
def test_handle_sse_event_caches_thought_signature_for_antigravity() -> None:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
|
||||
signature_cache.clear()
|
||||
|
||||
@@ -148,7 +152,9 @@ def test_handle_sse_event_caches_thought_signature_for_antigravity() -> None:
|
||||
}
|
||||
|
||||
with patch.object(handler, "_process_event_data") as _mock_process:
|
||||
handler._handle_sse_event(ctx, None, json.dumps(payload), record_chunk=False)
|
||||
cast(CliHandlerProtocol, handler)._handle_sse_event(
|
||||
ctx, None, json.dumps(payload), record_chunk=False
|
||||
)
|
||||
|
||||
assert signature_cache.get_or_dummy("claude-sonnet-4-5", "t1") == long_sig
|
||||
|
||||
@@ -206,7 +212,7 @@ async def test_antigravity_forces_conversion_path_in_stream_with_prefetch() -> N
|
||||
|
||||
|
||||
def test_wrap_v1internal_request_injects_thought_signature_from_tool_cache() -> None:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
|
||||
signature_cache.clear()
|
||||
|
||||
@@ -234,7 +240,7 @@ def test_wrap_v1internal_request_injects_thought_signature_from_tool_cache() ->
|
||||
|
||||
|
||||
def test_wrap_v1internal_request_injects_session_signature_when_tool_cache_missing() -> None:
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
|
||||
signature_cache.clear()
|
||||
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
def _make_candidate(
|
||||
@@ -31,15 +33,15 @@ def _make_candidate(
|
||||
)
|
||||
|
||||
return ProviderCandidate(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
provider=cast(Provider, provider),
|
||||
endpoint=cast(ProviderEndpoint, endpoint),
|
||||
key=cast(ProviderAPIKey, key),
|
||||
is_cached=False,
|
||||
is_skipped=is_skipped,
|
||||
skip_reason="unhealthy" if is_skipped else None,
|
||||
needs_conversion=needs_conversion,
|
||||
provider_api_format="openai:chat",
|
||||
) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
def _make_db() -> MagicMock:
|
||||
@@ -71,11 +71,11 @@ async def test_list_all_candidates_returns_provider_batch_count_even_when_candid
|
||||
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||
with patch.object(scheduler._candidate_builder, "_query_providers", return_value=providers):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
new=AsyncMock(return_value=global_model),
|
||||
):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
return_value=True,
|
||||
):
|
||||
candidates, global_model_id, provider_batch_count = (
|
||||
@@ -108,11 +108,11 @@ async def test_list_all_candidates_returns_zero_provider_batch_count_when_provid
|
||||
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||
with patch.object(scheduler._candidate_builder, "_query_providers", return_value=[]):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
new=AsyncMock(return_value=global_model),
|
||||
):
|
||||
with patch(
|
||||
"src.services.cache.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
return_value=True,
|
||||
):
|
||||
candidates, global_model_id, provider_batch_count = (
|
||||
|
||||
@@ -2,8 +2,8 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.thinking_cache import ThinkingSignatureCache
|
||||
from src.services.provider.adapters.antigravity.constants import DUMMY_THOUGHT_SIGNATURE
|
||||
from src.services.provider.adapters.antigravity.signature_cache import ThinkingSignatureCache
|
||||
|
||||
# 测试用签名(需 >= MIN_SIGNATURE_LENGTH=50)
|
||||
_SIG_A = "a" * 60
|
||||
@@ -54,11 +54,11 @@ def test_tool_signature_short_ignored() -> None:
|
||||
|
||||
|
||||
def test_tool_signature_cache_enforces_limit(monkeypatch: Any) -> None:
|
||||
import src.services.provider.adapters.antigravity.signature_cache as sc_mod
|
||||
import src.core.api_format.conversion.thinking_cache as tc_mod
|
||||
|
||||
# Use a small limit to make eviction deterministic in tests.
|
||||
monkeypatch.setattr(sc_mod, "_TOOL_CACHE_LIMIT", 3)
|
||||
cache = sc_mod.ThinkingSignatureCache()
|
||||
monkeypatch.setattr(tc_mod, "_TOOL_CACHE_LIMIT", 3)
|
||||
cache = tc_mod.ThinkingSignatureCache()
|
||||
cache.cache_tool_signature("toolu_1", _SIG_A)
|
||||
cache.cache_tool_signature("toolu_2", _SIG_B)
|
||||
cache.cache_tool_signature("toolu_3", _SIG_C)
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from src.api.handlers.base.utils import get_format_converter_registry
|
||||
from src.core.api_format.conversion.thinking_cache import signature_cache
|
||||
from src.services.provider.adapters.antigravity.constants import DUMMY_THOUGHT_SIGNATURE
|
||||
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
|
||||
|
||||
|
||||
def _reset_sig_cache() -> None:
|
||||
|
||||
@@ -15,7 +15,7 @@ from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
|
||||
from src.models.database import GlobalModel, Model, Provider
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
class TestCheckModelSupportForGlobalModel:
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
def _make_key(
|
||||
@@ -20,7 +20,7 @@ def _make_key(
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_kiro_quota_remaining_zero_skips(_mock_cb: MagicMock) -> None:
|
||||
@@ -39,7 +39,7 @@ def test_kiro_quota_remaining_zero_skips(_mock_cb: MagicMock) -> None:
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_kiro_quota_remaining_positive_allows(_mock_cb: MagicMock) -> None:
|
||||
@@ -58,7 +58,7 @@ def test_kiro_quota_remaining_positive_allows(_mock_cb: MagicMock) -> None:
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_codex_weekly_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
@@ -85,7 +85,7 @@ def test_codex_weekly_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_codex_5h_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
@@ -111,7 +111,7 @@ def test_codex_5h_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_codex_ignores_code_review_quota(_mock_cb: MagicMock) -> None:
|
||||
@@ -138,7 +138,7 @@ def test_codex_ignores_code_review_quota(_mock_cb: MagicMock) -> None:
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_antigravity_model_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
@@ -166,7 +166,7 @@ def test_antigravity_model_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_antigravity_other_model_not_exhausted_allows(_mock_cb: MagicMock) -> None:
|
||||
@@ -194,7 +194,7 @@ def test_antigravity_other_model_not_exhausted_allows(_mock_cb: MagicMock) -> No
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.cache.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_antigravity_quota_uses_mapping_matched_model(_mock_cb: MagicMock) -> None:
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.orchestration.candidate_resolver import CandidateResolver
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
class _FakeScheduler:
|
||||
@@ -89,19 +90,19 @@ def _make_global_key_candidate(*, key_id: str, priority: int) -> ProviderCandida
|
||||
global_priority_by_format={"openai:chat": priority},
|
||||
)
|
||||
return ProviderCandidate(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
provider=cast(Provider, provider),
|
||||
endpoint=cast(ProviderEndpoint, endpoint),
|
||||
key=cast(ProviderAPIKey, key),
|
||||
needs_conversion=False,
|
||||
provider_api_format="openai:chat",
|
||||
) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_candidate_resolver_pagination_continues_on_empty_candidate_batch() -> None:
|
||||
db = MagicMock()
|
||||
scheduler = _FakeScheduler()
|
||||
resolver = CandidateResolver(db=db, cache_scheduler=scheduler) # type: ignore[arg-type]
|
||||
resolver = CandidateResolver(db=db, cache_scheduler=cast(CacheAwareScheduler, scheduler))
|
||||
|
||||
candidates, global_model_id = await resolver.fetch_candidates(
|
||||
api_format="openai:chat",
|
||||
|
||||
@@ -50,7 +50,7 @@ async def test_submit_with_failover_skips_http_500_then_succeeds(
|
||||
lambda *_args, **_kwargs: "provider",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.cache.aware_scheduler.get_cache_aware_scheduler",
|
||||
"src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
@@ -106,7 +106,7 @@ async def test_submit_with_failover_stops_on_client_error(monkeypatch: pytest.Mo
|
||||
lambda *_args, **_kwargs: "provider",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.cache.aware_scheduler.get_cache_aware_scheduler",
|
||||
"src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
@@ -152,7 +152,7 @@ async def test_submit_with_failover_no_eligible_candidates_due_to_auth_type(
|
||||
lambda *_args, **_kwargs: "provider",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.cache.aware_scheduler.get_cache_aware_scheduler",
|
||||
"src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
@@ -194,7 +194,7 @@ async def test_submit_with_failover_filters_missing_billing_rule(
|
||||
lambda *_args, **_kwargs: "provider",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.cache.aware_scheduler.get_cache_aware_scheduler",
|
||||
"src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
|
||||
@@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, MagicMock
|
||||
import pytest
|
||||
|
||||
from src.core.api_format.conversion import register_default_normalizers
|
||||
from src.services.cache.aware_scheduler import (
|
||||
from src.services.scheduling.aware_scheduler import (
|
||||
CacheAwareScheduler,
|
||||
_sort_endpoints_by_family_priority,
|
||||
)
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from src.services.cache.candidate_sorter import CandidateSorter
|
||||
from src.services.cache.scheduling_config import SchedulingConfig
|
||||
from src.services.cache.schemas import ProviderCandidate
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.scheduling.candidate_sorter import CandidateSorter
|
||||
from src.services.scheduling.scheduling_config import SchedulingConfig
|
||||
from src.services.scheduling.schemas import ProviderCandidate
|
||||
|
||||
|
||||
def _make_candidate(
|
||||
@@ -28,12 +30,12 @@ def _make_candidate(
|
||||
global_priority_by_format={"openai:chat": global_priority},
|
||||
)
|
||||
return ProviderCandidate(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
provider=cast(Provider, provider),
|
||||
endpoint=cast(ProviderEndpoint, endpoint),
|
||||
key=cast(ProviderAPIKey, key),
|
||||
needs_conversion=needs_conversion,
|
||||
provider_api_format="openai:chat",
|
||||
) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_priority_sort_global_key_does_not_demote_when_global_keep_priority_enabled() -> None:
|
||||
@@ -59,7 +61,7 @@ def test_priority_sort_global_key_does_not_demote_when_global_keep_priority_enab
|
||||
|
||||
# 全局 keep_priority_on_conversion=True:不做 needs_conversion 降级分组,纯按 global_priority 排序
|
||||
with patch(
|
||||
"src.services.cache.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
"src.services.scheduling.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
return_value=True,
|
||||
):
|
||||
result = sorter._apply_priority_mode_sort([exact, demoted], db, None, "openai:chat")
|
||||
@@ -90,7 +92,7 @@ def test_priority_sort_global_key_demotes_convertible_when_global_keep_priority_
|
||||
|
||||
# 全局 keep_priority_on_conversion=False:需要降级的 convertible 候选整体排后
|
||||
with patch(
|
||||
"src.services.cache.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
"src.services.scheduling.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
return_value=False,
|
||||
):
|
||||
result = sorter._apply_priority_mode_sort([exact, demoted], db, None, "openai:chat")
|
||||
@@ -126,7 +128,7 @@ def test_priority_sort_global_key_provider_keep_priority_overrides_demotion_grou
|
||||
)
|
||||
|
||||
with patch(
|
||||
"src.services.cache.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
"src.services.scheduling.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
return_value=False,
|
||||
):
|
||||
result = sorter._apply_priority_mode_sort(
|
||||
@@ -163,7 +165,7 @@ def test_priority_sort_provider_mode_demotes_convertible_when_global_keep_priori
|
||||
)
|
||||
|
||||
with patch(
|
||||
"src.services.cache.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
"src.services.scheduling.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
return_value=False,
|
||||
):
|
||||
result = sorter._apply_priority_mode_sort([demoted, exact], db, None, "openai:chat")
|
||||
@@ -193,7 +195,7 @@ def test_priority_sort_provider_mode_does_not_demote_when_global_keep_priority_e
|
||||
)
|
||||
|
||||
with patch(
|
||||
"src.services.cache.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
"src.services.scheduling.candidate_sorter.SystemConfigService.is_keep_priority_on_conversion",
|
||||
return_value=True,
|
||||
):
|
||||
result = sorter._apply_priority_mode_sort([demoted, exact], db, None, "openai:chat")
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
|
||||
def _make_candidate(
|
||||
|
||||
Reference in New Issue
Block a user