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:
fawney19
2026-02-16 11:00:48 +08:00
parent 4dc401677d
commit 63870931af
85 changed files with 1907 additions and 1463 deletions

View File

@@ -32,7 +32,7 @@ from src.models.database import (
ProviderAPIKey, ProviderAPIKey,
ProviderEndpoint, ProviderEndpoint,
) )
from src.services.cache.aware_scheduler import CacheAwareScheduler from src.services.scheduling.aware_scheduler import CacheAwareScheduler
from src.services.system.config import SystemConfigService from src.services.system.config import SystemConfigService
router = APIRouter(prefix="/global", tags=["Admin - Global Models"]) router = APIRouter(prefix="/global", tags=["Admin - Global Models"])

View File

@@ -22,8 +22,8 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.database import get_db from src.database import get_db
from src.models.database import ApiKey, User from src.models.database import ApiKey, User
from src.services.cache.affinity_manager import get_affinity_manager from src.services.scheduling.affinity_manager import get_affinity_manager
from src.services.cache.aware_scheduler import CacheAwareScheduler, get_cache_aware_scheduler from src.services.scheduling.aware_scheduler import CacheAwareScheduler, get_cache_aware_scheduler
from src.services.system.config import SystemConfigService from src.services.system.config import SystemConfigService
router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitoring: Cache"]) router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitoring: Cache"])
@@ -1103,8 +1103,8 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
class AdminCacheConfigAdapter(AdminApiAdapter): class AdminCacheConfigAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.config.constants import ConcurrencyDefaults 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.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
from src.services.scheduling.affinity_manager import CacheAffinityManager
# 获取动态预留管理器的配置 # 获取动态预留管理器的配置
reservation_manager = get_adaptive_reservation_manager() reservation_manager = get_adaptive_reservation_manager()

View File

@@ -643,7 +643,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
if self.key in ("scheduling_mode", "provider_priority_mode"): if self.key in ("scheduling_mode", "provider_priority_mode"):
try: try:
from src.clients.redis_client import get_redis_client_sync 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() redis_client = get_redis_client_sync()
# 从数据库读取两个调度配置的最新值,确保一致性 # 从数据库读取两个调度配置的最新值,确保一致性

View File

@@ -19,15 +19,18 @@ from sqlalchemy import tuple_
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from src.config.constants import CacheTTL 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.api_format.conversion.compatibility import is_format_compatible
from src.core.cache_service import CacheService from src.core.cache_service import CacheService
from src.core.logger import logger 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.model.availability import ModelAvailabilityQuery
from src.services.provider.format import normalize_endpoint_signature from src.services.provider.format import normalize_endpoint_signature
# 缓存 key 前缀
_CACHE_KEY_PREFIX = "models:list"
_CACHE_TTL = CacheTTL.MODEL # 300 秒 _CACHE_TTL = CacheTTL.MODEL # 300 秒
@@ -70,21 +73,7 @@ async def _set_cached_models(
logger.warning(f"[ModelsService] 缓存写入失败: {e}") logger.warning(f"[ModelsService] 缓存写入失败: {e}")
async def invalidate_models_list_cache() -> None: __all__ = ["AccessRestrictions", "invalidate_models_list_cache", "ModelInfo"]
"""
清除所有 /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}")
@dataclass @dataclass
@@ -115,91 +104,7 @@ class ModelInfo:
output_modalities: list[str] | None = None output_modalities: list[str] | None = None
@dataclass # AccessRestrictions -- re-export from src.core.access_restrictions (see __all__)
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
def _normalize_api_formats( def _normalize_api_formats(

View File

@@ -45,8 +45,8 @@ from sqlalchemy.orm import Session
from src.clients.redis_client import get_redis_client_sync from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger from src.core.logger import logger
from src.services.provider.format import normalize_endpoint_signature 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.service import UsageService
from src.services.usage.telemetry import MessageTelemetry # re-export
if TYPE_CHECKING: if TYPE_CHECKING:
from src.api.handlers.base.stream_context import StreamContext 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]] type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
class MessageTelemetry: # MessageTelemetry -- re-export from src.services.usage.telemetry (see import above)
""" __all__ = ["MessageTelemetry", "MessageHandlerProtocol", "AdapterDetectorType"]
负责记录 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,
)
@runtime_checkable @runtime_checkable

View File

@@ -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.exceptions import ThinkingSignatureException, UpstreamClientException
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ProviderAPIKey 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.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: def _get_error_status_code(e: Exception, default: int = 400) -> int:

View File

@@ -75,7 +75,6 @@ from src.models.database import (
ProviderEndpoint, ProviderEndpoint,
User, User,
) )
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.behavior import get_provider_behavior from src.services.provider.behavior import get_provider_behavior
from src.services.provider.stream_policy import ( from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream, enforce_stream_mode_for_upstream,
@@ -85,6 +84,7 @@ from src.services.provider.stream_policy import (
from src.services.provider.transport import ( from src.services.provider.transport import (
build_provider_url, build_provider_url,
) )
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.system.config import SystemConfigService from src.services.system.config import SystemConfigService

View File

@@ -49,7 +49,7 @@ if TYPE_CHECKING:
from src.api.handlers.base.chat_handler_base import ChatHandlerBase from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint 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 @dataclass

View File

@@ -38,7 +38,6 @@ from src.core.exceptions import (
ProviderTimeoutException, ProviderTimeoutException,
) )
from src.core.logger import logger 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.behavior import get_provider_behavior
from src.services.provider.stream_policy import ( from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream, enforce_stream_mode_for_upstream,
@@ -46,6 +45,7 @@ from src.services.provider.stream_policy import (
resolve_upstream_is_stream, resolve_upstream_is_stream,
) )
from src.services.provider.transport import build_provider_url 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.services.system.config import SystemConfigService
from src.utils.sse_parser import SSEEventParser from src.utils.sse_parser import SSEEventParser
from src.utils.timeout import read_first_chunk_with_ttfb_timeout from src.utils.timeout import read_first_chunk_with_ttfb_timeout

View File

@@ -29,7 +29,6 @@ from src.core.exceptions import (
ThinkingSignatureException, ThinkingSignatureException,
) )
from src.core.logger import logger 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.behavior import get_provider_behavior
from src.services.provider.stream_policy import ( from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream, enforce_stream_mode_for_upstream,
@@ -37,6 +36,7 @@ from src.services.provider.stream_policy import (
resolve_upstream_is_stream, resolve_upstream_is_stream,
) )
from src.services.provider.transport import build_provider_url from src.services.provider.transport import build_provider_url
from src.services.scheduling.aware_scheduler import ProviderCandidate
if TYPE_CHECKING: if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol from src.api.handlers.base.cli_protocol import CliHandlerProtocol

View File

@@ -693,36 +693,22 @@ class GeminiCliResponseParser(GeminiResponseParser):
self.api_format = "gemini:cli" self.api_format = "gemini:cli"
# 解析器注册表 # 注册解析器到 core 层注册表(供 services 层通过 format_id 获取)
_PARSERS: dict[str, type[ResponseParser]] = { from src.core.stream_types import get_parser_for_format, register_parser
"claude:chat": ClaudeResponseParser,
"claude:cli": ClaudeCliResponseParser,
"openai:chat": OpenAIResponseParser,
"openai:cli": OpenAICliResponseParser,
"gemini:chat": GeminiResponseParser,
"gemini:cli": GeminiCliResponseParser,
}
def get_parser_for_format(format_id: str) -> ResponseParser: def register_default_parsers() -> None:
""" register_parser("claude:chat", ClaudeResponseParser)
根据格式 ID 获取 ResponseParser 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: # 模块加载时自动注册(保证 import parsers 即可用,测试也不需要手动初始化)
ResponseParser 实例 # main.py lifespan 中的显式调用是冗余但无害的安全保障dict 覆盖幂等)
register_default_parsers()
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]()
__all__ = [ __all__ = [
@@ -732,6 +718,7 @@ __all__ = [
"ClaudeCliResponseParser", "ClaudeCliResponseParser",
"GeminiResponseParser", "GeminiResponseParser",
"GeminiCliResponseParser", "GeminiCliResponseParser",
"register_default_parsers",
"get_parser_for_format", "get_parser_for_format",
"is_cli_format", "is_cli_format",
] ]

View File

@@ -14,17 +14,10 @@
from __future__ import annotations from __future__ import annotations
import copy import copy
import json
import re import re
import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass from typing import Any
from typing import TYPE_CHECKING, Any
import httpx
from sqlalchemy.orm import object_session
from src.clients.redis_client import get_redis_client
from src.core.api_format import ( from src.core.api_format import (
UPSTREAM_DROP_HEADERS, UPSTREAM_DROP_HEADERS,
HeaderBuilder, HeaderBuilder,
@@ -32,99 +25,9 @@ from src.core.api_format import (
make_signature_key, make_signature_key,
) )
from src.core.crypto import crypto_service 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
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
from src.services.provider.auth import get_provider_auth
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
# ============================================================================== # ==============================================================================
# 统一的头部配置常量 # 统一的头部配置常量
@@ -1003,303 +906,3 @@ def build_passthrough_request(
endpoint, endpoint,
key, 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

View File

@@ -1,176 +1,19 @@
""" """
响应解析器基类 - 定义统一的响应解析接口 响应解析器基类 - re-export from src.core.stream_types
实际定义已下沉到 src/core/stream_types.py此文件保留向后兼容的 re-export。
""" """
from abc import ABC, abstractmethod from src.core.stream_types import (
from dataclasses import dataclass, field ParsedChunk,
from typing import Any ParsedResponse,
ResponseParser,
StreamStats,
)
__all__ = [
@dataclass "ParsedChunk",
class ParsedChunk: "ParsedResponse",
"""解析后的流式数据块""" "ResponseParser",
"StreamStats",
# 原始数据 ]
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()

View File

@@ -6,7 +6,6 @@ Video Handler 基类
from __future__ import annotations from __future__ import annotations
import re
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any 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.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
from src.core.exceptions import ProviderNotAvailableException from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger 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.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult 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: if TYPE_CHECKING:
import httpx import httpx
from src.services.candidate.submit import SubmitOutcome 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): class VideoHandlerBase(ABC):
"""视频处理器基类""" """视频处理器基类"""

View File

@@ -7,13 +7,9 @@ Gemini 图像生成模型请求适配
from typing import Any from typing import Any
from src.core.video_utils import is_image_gen_model
def is_image_gen_model(model: str | None) -> bool: __all__ = ["is_image_gen_model", "adapt_request_for_image_gen"]
"""判断是否为图像生成模型(模式匹配,覆盖 gemini-*-image / imagen-* 系列)"""
if not model:
return False
m = model.lower()
return "image" in m and ("gemini" in m or "imagen" in m)
def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]: def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]:

View File

@@ -41,7 +41,7 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService 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 from src.services.usage.service import UsageService

View File

@@ -39,7 +39,7 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService 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 from src.services.usage.service import UsageService

View File

@@ -36,9 +36,9 @@ from src.core.logger import logger
from src.database import create_session from src.database import create_session
from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User
from src.services.auth.service import AuthService 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.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.provider.transport import redact_url_for_log
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
@dataclass @dataclass

View File

@@ -153,6 +153,14 @@ class RPMDefaults:
# confidence 自然衰减速率:每分钟衰减的比例 # confidence 自然衰减速率:每分钟衰减的比例
CONFIDENCE_DECAY_PER_MINUTE = 0.005 # 每分钟 -0.5%,约 200 分钟(~3.3h)从 1.0 衰减到 0 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 ConcurrencyDefaults = RPMDefaults

View File

@@ -135,6 +135,19 @@ class Config:
# CACHE_RESERVATION_RATIO: 缓存用户预留比例(默认 10%,新用户可用 90% # CACHE_RESERVATION_RATIO: 缓存用户预留比例(默认 10%,新用户可用 90%
self.cache_reservation_ratio = float(os.getenv("CACHE_RESERVATION_RATIO", "0.1")) 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异常时的行为 # RATE_LIMIT_FAIL_OPEN: 当限流服务Redis异常时的行为
# #

View File

@@ -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

View File

@@ -1425,7 +1425,7 @@ class GeminiNormalizer(FormatNormalizer):
) )
try: try:
from src.services.provider.adapters.antigravity.signature_cache import ( from src.core.api_format.conversion.thinking_cache import (
signature_cache, signature_cache,
) )
@@ -1565,7 +1565,7 @@ class GeminiNormalizer(FormatNormalizer):
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE
try: 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) cached_or_dummy = signature_cache.get_or_dummy(model, text_val)

View File

@@ -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"]

View File

@@ -8,7 +8,7 @@ from __future__ import annotations
import importlib.util import importlib.util
import os import os
from typing import TYPE_CHECKING from typing import TYPE_CHECKING, Any, Protocol
from src.core.logger import logger from src.core.logger import logger
from src.core.modules.base import ( from src.core.modules.base import (
@@ -22,6 +22,17 @@ if TYPE_CHECKING:
from sqlalchemy.orm import Session 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: class ModuleRegistry:
""" """
模块注册中心 - 单例模式 模块注册中心 - 单例模式
@@ -34,11 +45,17 @@ class ModuleRegistry:
""" """
_instance: ModuleRegistry | None = None _instance: ModuleRegistry | None = None
_config_backend: ConfigBackend | None = None
def __init__(self) -> None: def __init__(self) -> None:
self._modules: dict[str, ModuleDefinition] = {} self._modules: dict[str, ModuleDefinition] = {}
self._initialized: set[str] = set() self._initialized: set[str] = set()
@classmethod
def set_config_backend(cls, backend: ConfigBackend) -> None:
"""注入配置读写后端,消除 core→services 的运行时依赖"""
cls._config_backend = backend
@classmethod @classmethod
def get_instance(cls) -> ModuleRegistry: def get_instance(cls) -> ModuleRegistry:
"""获取单例实例""" """获取单例实例"""
@@ -50,6 +67,7 @@ class ModuleRegistry:
def reset_instance(cls) -> None: def reset_instance(cls) -> None:
"""重置单例(仅用于测试)""" """重置单例(仅用于测试)"""
cls._instance = None cls._instance = None
cls._config_backend = None
def register(self, module: ModuleDefinition) -> 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: def is_enabled(self, name: str, db: Session) -> bool:
""" """
检查模块是否运行启用(数据库配置) 检查模块是否运行启用(数据库配置)
@@ -131,10 +158,8 @@ class ModuleRegistry:
name: 模块名称 name: 模块名称
db: 数据库会话 db: 数据库会话
""" """
from src.services.system.config import SystemConfigService
config_key = f"module.{name}.enabled" 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) return bool(value)
def set_enabled(self, name: str, enabled: bool, db: Session) -> None: def set_enabled(self, name: str, enabled: bool, db: Session) -> None:
@@ -146,15 +171,13 @@ class ModuleRegistry:
enabled: 是否启用 enabled: 是否启用
db: 数据库会话 db: 数据库会话
""" """
from src.services.system.config import SystemConfigService
if name not in self._modules: if name not in self._modules:
raise ValueError(f"Module [{name}] not registered") raise ValueError(f"Module [{name}] not registered")
config_key = f"module.{name}.enabled" config_key = f"module.{name}.enabled"
module = self._modules[name] module = self._modules[name]
description = f"模块 [{module.metadata.display_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)
# ========== 激活状态检查 ========== # ========== 激活状态检查 ==========

View File

@@ -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)

View File

@@ -10,7 +10,6 @@ import jwt
from src.clients.http_client import HTTPClientPool from src.clients.http_client import HTTPClientPool
from src.core.logger import logger from src.core.logger import logger
from src.core.provider_types import ProviderType 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" _ANTHROPIC_TOKEN_URL = "https://console.anthropic.com/v1/oauth/token"
_GOOGLE_USERINFO_URL = "https://www.googleapis.com/oauth2/v1/userinfo?alt=json" _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: try:
if not proxy_config.get("enabled", True): if not proxy_config.get("enabled", True):
return None return None
from src.services.proxy_node.resolver import build_proxy_url # lazy: core→services
return build_proxy_url(proxy_config) return build_proxy_url(proxy_config)
except Exception: except Exception:
return None return None
@@ -363,11 +364,10 @@ async def enrich_auth_config(
"""Enrich auth_config with non-secret metadata (email/account_id). """Enrich auth_config with non-secret metadata (email/account_id).
各 provider 的 enrichment 逻辑通过 register_auth_enricher 注册。 各 provider 的 enrichment 逻辑通过 register_auth_enricher 注册。
注意: ensure_providers_bootstrapped() 在应用启动时(main.py lifespan)已显式调用。
""" """
from src.core.provider_types import normalize_provider_type 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) pt = normalize_provider_type(provider_type)
enricher = _auth_enrichers.get(pt) enricher = _auth_enrichers.get(pt)
if enricher: if enricher:

View File

@@ -15,9 +15,9 @@ from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from src.core.provider_templates.types import ProviderType 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) @dataclass(frozen=True, slots=True)

224
src/core/stream_types.py Normal file
View File

@@ -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]()

View File

@@ -10,6 +10,7 @@ from __future__ import annotations
import json import json
import time import time
from collections import OrderedDict from collections import OrderedDict
from typing import Any
import httpx import httpx
import jwt import jwt
@@ -102,13 +103,17 @@ class VertexAuthService:
} }
return jwt.encode(payload, self.private_key, algorithm="RS256") 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 缓存) 获取 Access Token带 LRU 缓存)
如果缓存中有有效的 Token距离过期超过 60 秒),直接返回。 如果缓存中有有效的 Token距离过期超过 60 秒),直接返回。
否则重新获取 Token。缓存采用 LRU 策略,超过 100 个条目时淘汰最旧的。 否则重新获取 Token。缓存采用 LRU 策略,超过 100 个条目时淘汰最旧的。
Args:
httpx_client_kwargs: 传给 httpx.AsyncClient 的额外参数(如代理配置)。
调用者services 层负责构建core 层不关心代理细节。
Returns: Returns:
Access Token 字符串 Access Token 字符串
@@ -129,10 +134,10 @@ class VertexAuthService:
try: try:
signed_jwt = self._create_jwt() signed_jwt = self._create_jwt()
# 使用系统默认代理Vertex AI token endpoint 是外部服务) client_kwargs = (
from src.services.proxy_node.resolver import build_proxy_client_kwargs httpx_client_kwargs if httpx_client_kwargs is not None else {"timeout": 30}
)
async with httpx.AsyncClient(**build_proxy_client_kwargs(timeout=30)) as client: async with httpx.AsyncClient(**client_kwargs) as client:
resp = await client.post( resp = await client.post(
self.TOKEN_URL, self.TOKEN_URL,
data={ data={
@@ -186,16 +191,21 @@ class VertexAuthService:
cls._token_cache.clear() 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 便捷函数:获取 Vertex AI Access Token 和 Project ID
Args: Args:
service_account_json: Service Account JSON 字符串 service_account_json: Service Account JSON 字符串
httpx_client_kwargs: 传给 httpx.AsyncClient 的额外参数(如代理配置)
Returns: Returns:
(access_token, project_id) 元组 (access_token, project_id) 元组
""" """
service = VertexAuthService(service_account_json) 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 return token, service.project_id

77
src/core/video_utils.py Normal file
View File

@@ -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)

View File

@@ -176,6 +176,12 @@ async def lifespan(app: FastAPI) -> Any:
from src.modules import ALL_MODULES from src.modules import ALL_MODULES
module_registry = get_module_registry() 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: for module in ALL_MODULES:
module_registry.register(module) module_registry.register(module)
@@ -195,6 +201,17 @@ async def lifespan(app: FastAPI) -> Any:
logger.info(f"功能模块初始化完成: {len(available_modules)}/{len(ALL_MODULES)} 个模块可用") 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(f"服务启动成功: http://{config.host}:{config.port}")
logger.info("=" * 60) logger.info("=" * 60)

View File

@@ -1,14 +1,10 @@
""" """通用缓存模块。
缓存服务模块
包含缓存后端、缓存亲和性、缓存同步等能。 包含缓存后端、缓存失效与缓存同步等能backend/sync/*_cache
注意:由于循环依赖问题,部分类需要直接从子模块导入: 调度/候选/缓存亲和性相关逻辑已迁移到 `src.services.scheduling`。
from src.services.cache.affinity_manager import CacheAffinityManager
from src.services.cache.aware_scheduler import CacheAwareScheduler
""" """
# 只导出不会导致循环依赖的基础类
from src.services.cache.backend import BaseCacheBackend, LocalCache, RedisCache, get_cache_backend from src.services.cache.backend import BaseCacheBackend, LocalCache, RedisCache, get_cache_backend
__all__ = [ __all__ = [

View File

@@ -54,7 +54,7 @@ class CacheInvalidationService:
logger.error(f"[CacheInvalidation] 失效 ModelCacheService 缓存失败: {e}") logger.error(f"[CacheInvalidation] 失效 ModelCacheService 缓存失败: {e}")
# 4. 清除 /v1/models 列表缓存 # 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: try:
await invalidate_models_list_cache() await invalidate_models_list_cache()
@@ -79,7 +79,7 @@ class CacheInvalidationService:
self._refresh_provider_cache(provider_id) self._refresh_provider_cache(provider_id)
# 清除 /v1/models 列表缓存allowed_models 变更会影响模型可用性) # 清除 /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: try:
await invalidate_models_list_cache() await invalidate_models_list_cache()

31
src/services/cache/model_list_cache.py vendored Normal file
View File

@@ -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)

View File

@@ -13,9 +13,9 @@ from sqlalchemy.orm import Session
from src.core.logger import logger from src.core.logger import logger
from src.models.database import RequestCandidate 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.orchestration.error_classifier import ErrorAction, ErrorClassifier
from src.services.request.candidate import RequestCandidateService 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.exceptions import StreamProbeError
from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult from src.services.task.protocol import AttemptFunc, AttemptKind, AttemptResult
from src.services.task.schema import ExecutionResult from src.services.task.schema import ExecutionResult

View File

@@ -3,7 +3,7 @@ from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any 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" CANDIDATE_KEY_SCHEMA_VERSION = "1.0"

View File

@@ -5,8 +5,8 @@ from typing import Any
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from src.models.database import ApiKey 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.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 src.services.system.config import SystemConfigService
from .recorder import CandidateRecorder from .recorder import CandidateRecorder

View File

@@ -6,7 +6,7 @@ from typing import Any, Protocol, runtime_checkable
import httpx import httpx
from src.services.billing.rule_service import BillingRuleLookupResult 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 @runtime_checkable

View File

@@ -928,6 +928,7 @@ class ModelCostService:
# 获取对应 API 格式的 Adapter 实例来计算成本 # 获取对应 API 格式的 Adapter 实例来计算成本
# 优先检查 Chat Adapter然后检查 CLI 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.chat_adapter_base import get_adapter_instance
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_instance from src.api.handlers.base.cli_adapter_base import get_cli_adapter_instance

View File

@@ -10,13 +10,13 @@ from sqlalchemy import and_
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session 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.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger from src.core.logger import logger
from src.models.api import ModelCreate, ModelResponse, ModelUpdate from src.models.api import ModelCreate, ModelResponse, ModelUpdate
from src.models.database import Model, Provider from src.models.database import Model, Provider
from src.services.cache.invalidation import get_cache_invalidation_service from src.services.cache.invalidation import get_cache_invalidation_service
from src.services.cache.model_cache import ModelCacheService from src.services.cache.model_cache import ModelCacheService
from src.services.cache.model_list_cache import invalidate_models_list_cache
class ModelService: class ModelService:

View File

@@ -169,6 +169,7 @@ def merge_upstream_metadata(
def get_adapter_for_format(api_format: str) -> type | None: def get_adapter_for_format(api_format: str) -> type | None:
"""根据 API 格式获取对应的 Adapter 类""" """根据 API 格式获取对应的 Adapter 类"""
# TODO(arch): 引入 adapter 能力注册表,消除 services->api 依赖
from src.api.handlers.base.chat_adapter_base import get_adapter_class 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 from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class

View File

@@ -13,8 +13,8 @@ from sqlalchemy.orm import Session
from src.core.exceptions import ProviderNotAvailableException from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey 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.provider.format import normalize_endpoint_signature
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class CandidateResolver: class CandidateResolver:

View File

@@ -25,9 +25,9 @@ from src.core.exceptions import (
) )
from src.core.logger import logger from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint 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.orchestration.error_handler import ErrorHandlerService
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
class ErrorAction(Enum): class ErrorAction(Enum):

View File

@@ -22,11 +22,11 @@ from src.core.exceptions import (
) )
from src.core.logger import logger from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint 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.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature 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.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
class ErrorHandlerService: class ErrorHandlerService:

View File

@@ -11,9 +11,9 @@ from sqlalchemy.orm import Session
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey 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.candidate import RequestCandidateService
from src.services.request.executor import RequestExecutor from src.services.request.executor import RequestExecutor
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class RequestDispatcher: class RequestDispatcher:

View File

@@ -11,7 +11,9 @@ import re
import threading import threading
# ============== API 端点 ============== # ============== 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" DAILY_BASE_URL = "https://daily-cloudcode-pa.googleapis.com"
SANDBOX_BASE_URL = "https://daily-cloudcode-pa.sandbox.googleapis.com" SANDBOX_BASE_URL = "https://daily-cloudcode-pa.sandbox.googleapis.com"
@@ -75,8 +77,7 @@ URL_UNAVAILABLE_TTL_SECONDS = 300 # 5 分钟
# ============== Thinking Signature ============== # ============== Thinking Signature ==============
# 统一从 core 层导入,避免多处定义 # 统一从 core 层导入,避免多处定义
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE # noqa: E402 from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE # noqa: E402
from src.core.api_format.conversion.thinking_cache import MIN_SIGNATURE_LENGTH # noqa: E402, F401
MIN_SIGNATURE_LENGTH = 50 # 与 Antigravity-Manager 对齐
# ============== Thinking Budget ============== # ============== Thinking Budget ==============
THINKING_BUDGET_AUTO_CAP = 24576 THINKING_BUDGET_AUTO_CAP = 24576

View File

@@ -521,7 +521,7 @@ def _inject_thought_signatures(inner_request: dict[str, Any], session_id: str |
""" """
try: 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: except Exception:
return return
@@ -881,7 +881,7 @@ def wrap_v1internal_request(
13. 注入 sessionId 13. 注入 sessionId
14. 构建 v1internal 信封 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 = dict(gemini_request)
inner_request.pop("model", None) 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) 层。 同时缓存到 legacy (text) 层和 tool (Layer 1) 层。
""" """
try: 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: except Exception:
return return

View File

@@ -1,239 +1,13 @@
"""Antigravity thinking block signature cache (triple-layer). """Backward-compatible re-export for Antigravity thinking signature cache.
与 Antigravity-Manager 对齐的三层缓存设计: The implementation moved to `src.core.api_format.conversion.thinking_cache` to eliminate
Layer 1: tool_use_id → thoughtSignature (工具调用签名恢复) core → services reverse dependencies.
Layer 2: signature → model_family (跨模型兼容校验)
Layer 3: session_id → latest signature (会话级签名追踪 + rewind 检测)
同时保留原有的 model:text → signature 兼容层。
""" """
from __future__ import annotations from src.core.api_format.conversion.thinking_cache import (
import hashlib
import threading
import time
from typing import Any
from src.services.provider.adapters.antigravity.constants import (
DUMMY_THOUGHT_SIGNATURE,
MIN_SIGNATURE_LENGTH, MIN_SIGNATURE_LENGTH,
ThinkingSignatureCache,
signature_cache,
) )
# TTL: 2 小时(与 Antigravity-Manager 对齐) __all__ = ["MIN_SIGNATURE_LENGTH", "ThinkingSignatureCache", "signature_cache"]
_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"]

View File

@@ -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

View File

@@ -64,7 +64,7 @@ async def resolve_oauth_access_token(
""" """
# Local import to avoid circular imports during app startup. # 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(). # Build detached key-like objects for get_provider_auth().
provider_obj = ( provider_obj = (

View File

@@ -28,8 +28,6 @@ class ConcurrencyManager:
_instance: ConcurrencyManager | None = None _instance: ConcurrencyManager | None = None
_redis: aioredis.Redis | 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": def __new__(cls) -> "ConcurrencyManager":
"""单例模式""" """单例模式"""
@@ -42,13 +40,18 @@ class ConcurrencyManager:
if hasattr(self, "_memory_initialized"): if hasattr(self, "_memory_initialized"):
return 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() self._memory_lock: asyncio.Lock = asyncio.Lock()
# Key RPM 计数器:{key_id: (bucket, count)}bucket = floor(now / 60) # Key RPM 计数器:{key_id: (bucket, count)}bucket = floor(now / 60)
self._memory_key_rpm_counts: dict[str, tuple[int, int]] = {} self._memory_key_rpm_counts: dict[str, tuple[int, int]] = {}
self._owns_redis: bool = False self._owns_redis: bool = False
self._last_cleanup_bucket: int = 0 # 上次清理时的 bucket用于定期清理过期数据 self._last_cleanup_bucket: int = 0 # 上次清理时的 bucket用于定期清理过期数据
self._last_cleanup_time: float = 0 # 上次清理的时间戳,用于强制定期清理 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 # 后台清理任务 self._cleanup_task: asyncio.Task | None = None # 后台清理任务
# 内存模式下的最大条目限制,防止内存泄漏(支持环境变量覆盖) # 内存模式下的最大条目限制,防止内存泄漏(支持环境变量覆盖)
@@ -130,16 +133,14 @@ class ConcurrencyManager:
self._redis = None self._redis = None
self._owns_redis = False self._owns_redis = False
@classmethod def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
def _get_rpm_bucket(cls, now_ts: float | None = None) -> int:
"""获取当前 RPM 计数桶(按分钟)""" """获取当前 RPM 计数桶(按分钟)"""
ts = now_ts if now_ts is not None else time.time() 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(self, key_id: str, bucket: int | None = None) -> str:
def _get_key_key(cls, key_id: str, bucket: int | None = None) -> str:
"""获取 ProviderAPIKey RPM 计数的 Redis Key按分钟桶""" """获取 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}" return f"rpm:key:{key_id}:{b}"
def _get_memory_key_rpm_count(self, key_id: str, bucket: int) -> int: def _get_memory_key_rpm_count(self, key_id: str, bucket: int) -> int:

View File

@@ -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",
]

View File

@@ -31,7 +31,15 @@
from __future__ import annotations from __future__ import annotations
import time 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 from sqlalchemy.orm import Session
@@ -47,32 +55,32 @@ from src.models.database import (
ProviderAPIKey, ProviderAPIKey,
ProviderEndpoint, 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.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.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import ( from src.services.rate_limit.adaptive_reservation import (
get_adaptive_reservation_manager, get_adaptive_reservation_manager,
) )
from src.services.rate_limit.concurrency_manager import get_concurrency_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 from src.services.system.config import SystemConfigService
@@ -102,6 +110,11 @@ class CacheAwareScheduler:
redis_client: Any | None = None, redis_client: Any | None = None,
priority_mode: str | None = None, priority_mode: str | None = None,
scheduling_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: ) -> None:
""" """
初始化调度器 初始化调度器
@@ -117,9 +130,9 @@ class CacheAwareScheduler:
self.redis = redis_client self.redis = redis_client
self._config = SchedulingConfig(priority_mode, scheduling_mode) self._config = SchedulingConfig(priority_mode, scheduling_mode)
# 异步子组件(将在第一次使用时初始化) # 异步子组件(将在第一次使用时初始化,可通过构造函数注入
self._affinity_manager: CacheAffinityManager | None = None self._affinity_manager: CacheAffinityManagerProtocol | None = affinity_manager
self._concurrency_checker: ConcurrencyChecker | None = None self._concurrency_checker: ConcurrencyCheckerProtocol | None = concurrency_checker
self._metrics: dict[str, Any] = { self._metrics: dict[str, Any] = {
"total_batches": 0, "total_batches": 0,
@@ -139,9 +152,13 @@ class CacheAwareScheduler:
"last_reservation_result": None, "last_reservation_result": None,
} }
# 初始化子模块(不传 self解除反向引用 # 初始化子模块(不传 self解除反向引用,可通过构造函数注入
self._candidate_sorter = CandidateSorter(self._config) self._candidate_sorter: CandidateSorterProtocol = candidate_sorter or CandidateSorter(
self._candidate_builder = CandidateBuilder(self._candidate_sorter) self._config
)
self._candidate_builder: CandidateBuilderProtocol = candidate_builder or CandidateBuilder(
self._candidate_sorter
)
# ── 属性代理(保持外部访问兼容性)────────────────────────── # ── 属性代理(保持外部访问兼容性)──────────────────────────

View File

@@ -28,15 +28,15 @@ from src.models.database import (
ProviderAPIKey, ProviderAPIKey,
ProviderEndpoint, 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.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature 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: if TYPE_CHECKING:
from src.models.database import GlobalModel from src.models.database import GlobalModel
from src.services.cache.candidate_sorter import CandidateSorter from src.services.scheduling.protocols import CandidateSorterProtocol
from src.services.cache.schemas import ProviderCandidate from src.services.scheduling.schemas import ProviderCandidate
from src.services.cache.model_cache import ModelCacheService from src.services.cache.model_cache import ModelCacheService
@@ -60,7 +60,7 @@ def _sort_endpoints_by_family_priority(
class CandidateBuilder: class CandidateBuilder:
"""候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。""" """候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。"""
def __init__(self, candidate_sorter: CandidateSorter) -> None: def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
self._sorter = candidate_sorter self._sorter = candidate_sorter
def _query_providers( def _query_providers(
@@ -380,7 +380,7 @@ class CandidateBuilder:
Returns: Returns:
候选列表 候选列表
""" """
from src.services.cache.schemas import ProviderCandidate from src.services.scheduling.schemas import ProviderCandidate
candidates: list[ProviderCandidate] = [] candidates: list[ProviderCandidate] = []
client_format_str = normalize_endpoint_signature(client_format) client_format_str = normalize_endpoint_signature(client_format)

View File

@@ -13,15 +13,15 @@ import random
from collections import defaultdict from collections import defaultdict
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from src.services.cache.scheduling_config import SchedulingConfig from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.cache.utils import affinity_hash from src.services.scheduling.utils import affinity_hash
from src.services.system.config import SystemConfigService from src.services.system.config import SystemConfigService
if TYPE_CHECKING: if TYPE_CHECKING:
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey from src.models.database import ProviderAPIKey
from src.services.cache.schemas import ProviderCandidate from src.services.scheduling.schemas import ProviderCandidate
class CandidateSorter: class CandidateSorter:

View File

@@ -11,9 +11,9 @@ from typing import Any
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ProviderAPIKey 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_reservation import AdaptiveReservationManager
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.scheduling.schemas import ConcurrencySnapshot
class ConcurrencyChecker: class ConcurrencyChecker:

View File

@@ -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]: ...

View File

@@ -72,7 +72,9 @@ class CacheWarmupService:
"""预热管理员仪表盘统计缓存""" """预热管理员仪表盘统计缓存"""
db = None db = None
try: 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 from src.models.database import User as DBUser
db = create_session() db = create_session()
@@ -128,7 +130,9 @@ class CacheWarmupService:
"""预热每日统计缓存""" """预热每日统计缓存"""
db = None db = None
try: 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 from src.models.database import User as DBUser
db = create_session() db = create_session()

View File

@@ -16,11 +16,6 @@ from typing import Any
from sqlalchemy.orm import Session 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.clients.http_client import HTTPClientPool
from src.config.settings import config from src.config.settings import config
from src.core.api_format import ( 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.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service from src.core.crypto import crypto_service
from src.core.logger import logger 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.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.provider.auth import get_provider_auth
from src.services.task.service import TaskService from src.services.task.service import TaskService

View File

@@ -6,7 +6,7 @@ from typing import Any, AsyncIterator, Protocol, runtime_checkable
import httpx import httpx
from src.services.cache.aware_scheduler import ProviderCandidate from src.services.scheduling.aware_scheduler import ProviderCandidate
class AttemptKind(str, Enum): class AttemptKind(str, Enum):

View File

@@ -3,8 +3,8 @@ from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any from typing import Any
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.candidate.schema import CandidateKey from src.services.candidate.schema import CandidateKey
from src.services.scheduling.aware_scheduler import ProviderCandidate
from .protocol import AttemptKind, AttemptResult from .protocol import AttemptKind, AttemptResult

View File

@@ -30,10 +30,6 @@ from src.models.database import (
User, User,
VideoTask, VideoTask,
) )
from src.services.cache.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.candidate.failover import FailoverEngine from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryPolicy, SkipPolicy from src.services.candidate.policy import RetryPolicy, SkipPolicy
from src.services.candidate.recorder import CandidateRecorder 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.provider.format import normalize_endpoint_signature
from src.services.request.candidate import RequestCandidateService from src.services.request.candidate import RequestCandidateService
from src.services.request.result import RequestMetadata 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.system.config import SystemConfigService
from src.services.task.context import TaskMode from src.services.task.context import TaskMode
from src.services.task.exceptions import TaskNotFoundError from src.services.task.exceptions import TaskNotFoundError
@@ -1560,7 +1560,6 @@ class TaskService:
from fastapi import HTTPException from fastapi import HTTPException
from src.api.handlers.base.request_builder import get_provider_auth
from src.clients.http_client import HTTPClientPool from src.clients.http_client import HTTPClientPool
from src.core.api_format import ( from src.core.api_format import (
build_upstream_headers_for_endpoint, build_upstream_headers_for_endpoint,
@@ -1569,6 +1568,7 @@ class TaskService:
) )
from src.core.api_format.conversion.internal_video import VideoStatus from src.core.api_format.conversion.internal_video import VideoStatus
from src.core.crypto import crypto_service 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 from src.services.provider.transport import build_provider_url
try: try:

View File

@@ -28,7 +28,7 @@ from typing import Any
from sqlalchemy.orm import Session 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.core.logger import logger
from src.models.database import ApiKey, User from src.models.database import ApiKey, User
from src.services.request.result import RequestResult from src.services.request.result import RequestResult

View File

@@ -13,10 +13,9 @@ from typing import Any
from sqlalchemy.orm import Session 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.exceptions import EmptyStreamException
from src.core.logger import logger 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.database.database import create_session
from src.models.database import ApiKey, User from src.models.database import ApiKey, User
from src.services.usage.service import UsageService from src.services.usage.service import UsageService

View File

@@ -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,
)

View File

@@ -8,11 +8,11 @@ import json
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any from typing import Any
from src.api.handlers.base.base_handler import MessageTelemetry
from src.clients.redis_client import get_redis_client from src.clients.redis_client import get_redis_client
from src.config.settings import config from src.config.settings import config
from src.core.logger import logger from src.core.logger import logger
from src.services.usage.events import UsageEventType, build_usage_event from src.services.usage.events import UsageEventType, build_usage_event
from src.services.usage.telemetry import MessageTelemetry
class TelemetryWriter(ABC): class TelemetryWriter(ABC):

View File

@@ -416,7 +416,7 @@ class UserService:
通过 GlobalModel + Model 关联查询用户可用模型 通过 GlobalModel + Model 关联查询用户可用模型
逻辑:使用 AccessRestrictions 统一处理 allowed_providers 和 allowed_models 限制 逻辑:使用 AccessRestrictions 统一处理 allowed_providers 和 allowed_models 限制
""" """
from src.api.base.models_service import AccessRestrictions from src.core.access_restrictions import AccessRestrictions
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致) # 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user) restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)

View File

@@ -2,10 +2,12 @@ from __future__ import annotations
import json import json
from types import SimpleNamespace from types import SimpleNamespace
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.api.handlers.base.stream_context import StreamContext from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
from src.services.provider.adapters.antigravity.envelope import ( 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: 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 assert mock_process.call_count == 1
passed_data = mock_process.call_args[0][2] 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: 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() 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: 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 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: 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() 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: 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() signature_cache.clear()

View File

@@ -1,11 +1,13 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest 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( def _make_candidate(
@@ -31,15 +33,15 @@ def _make_candidate(
) )
return ProviderCandidate( return ProviderCandidate(
provider=provider, provider=cast(Provider, provider),
endpoint=endpoint, endpoint=cast(ProviderEndpoint, endpoint),
key=key, key=cast(ProviderAPIKey, key),
is_cached=False, is_cached=False,
is_skipped=is_skipped, is_skipped=is_skipped,
skip_reason="unhealthy" if is_skipped else None, skip_reason="unhealthy" if is_skipped else None,
needs_conversion=needs_conversion, needs_conversion=needs_conversion,
provider_api_format="openai:chat", provider_api_format="openai:chat",
) # type: ignore[arg-type] )
@pytest.mark.asyncio @pytest.mark.asyncio

View File

@@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from src.services.cache.aware_scheduler import CacheAwareScheduler from src.services.scheduling.aware_scheduler import CacheAwareScheduler
def _make_db() -> MagicMock: 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, "_ensure_initialized", new=AsyncMock(return_value=None)):
with patch.object(scheduler._candidate_builder, "_query_providers", return_value=providers): with patch.object(scheduler._candidate_builder, "_query_providers", return_value=providers):
with patch( 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), new=AsyncMock(return_value=global_model),
): ):
with patch( 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, return_value=True,
): ):
candidates, global_model_id, provider_batch_count = ( 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, "_ensure_initialized", new=AsyncMock(return_value=None)):
with patch.object(scheduler._candidate_builder, "_query_providers", return_value=[]): with patch.object(scheduler._candidate_builder, "_query_providers", return_value=[]):
with patch( 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), new=AsyncMock(return_value=global_model),
): ):
with patch( 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, return_value=True,
): ):
candidates, global_model_id, provider_batch_count = ( candidates, global_model_id, provider_batch_count = (

View File

@@ -2,8 +2,8 @@ from __future__ import annotations
from typing import Any 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.constants import DUMMY_THOUGHT_SIGNATURE
from src.services.provider.adapters.antigravity.signature_cache import ThinkingSignatureCache
# 测试用签名(需 >= MIN_SIGNATURE_LENGTH=50 # 测试用签名(需 >= MIN_SIGNATURE_LENGTH=50
_SIG_A = "a" * 60 _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: 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. # Use a small limit to make eviction deterministic in tests.
monkeypatch.setattr(sc_mod, "_TOOL_CACHE_LIMIT", 3) monkeypatch.setattr(tc_mod, "_TOOL_CACHE_LIMIT", 3)
cache = sc_mod.ThinkingSignatureCache() cache = tc_mod.ThinkingSignatureCache()
cache.cache_tool_signature("toolu_1", _SIG_A) cache.cache_tool_signature("toolu_1", _SIG_A)
cache.cache_tool_signature("toolu_2", _SIG_B) cache.cache_tool_signature("toolu_2", _SIG_B)
cache.cache_tool_signature("toolu_3", _SIG_C) cache.cache_tool_signature("toolu_3", _SIG_C)

View File

@@ -1,8 +1,8 @@
from __future__ import annotations from __future__ import annotations
from src.api.handlers.base.utils import get_format_converter_registry 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.constants import DUMMY_THOUGHT_SIGNATURE
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
def _reset_sig_cache() -> None: def _reset_sig_cache() -> None:

View File

@@ -15,7 +15,7 @@ from unittest.mock import MagicMock, patch
import pytest import pytest
from src.models.database import GlobalModel, Model, Provider 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: class TestCheckModelSupportForGlobalModel:

View File

@@ -2,7 +2,7 @@ from __future__ import annotations
from unittest.mock import MagicMock, patch 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( def _make_key(
@@ -20,7 +20,7 @@ def _make_key(
@patch( @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), return_value=(True, None),
) )
def test_kiro_quota_remaining_zero_skips(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_kiro_quota_remaining_positive_allows(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_codex_weekly_quota_exhausted_skips(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_codex_5h_quota_exhausted_skips(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_codex_ignores_code_review_quota(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_antigravity_model_quota_exhausted_skips(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_antigravity_other_model_not_exhausted_allows(_mock_cb: MagicMock) -> 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( @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), return_value=(True, None),
) )
def test_antigravity_quota_uses_mapping_matched_model(_mock_cb: MagicMock) -> None: def test_antigravity_quota_uses_mapping_matched_model(_mock_cb: MagicMock) -> None:

View File

@@ -1,13 +1,14 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any from typing import Any, cast
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest 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.orchestration.candidate_resolver import CandidateResolver
from src.services.scheduling.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class _FakeScheduler: class _FakeScheduler:
@@ -89,19 +90,19 @@ def _make_global_key_candidate(*, key_id: str, priority: int) -> ProviderCandida
global_priority_by_format={"openai:chat": priority}, global_priority_by_format={"openai:chat": priority},
) )
return ProviderCandidate( return ProviderCandidate(
provider=provider, provider=cast(Provider, provider),
endpoint=endpoint, endpoint=cast(ProviderEndpoint, endpoint),
key=key, key=cast(ProviderAPIKey, key),
needs_conversion=False, needs_conversion=False,
provider_api_format="openai:chat", provider_api_format="openai:chat",
) # type: ignore[arg-type] )
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_candidate_resolver_pagination_continues_on_empty_candidate_batch() -> None: async def test_candidate_resolver_pagination_continues_on_empty_candidate_batch() -> None:
db = MagicMock() db = MagicMock()
scheduler = _FakeScheduler() 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( candidates, global_model_id = await resolver.fetch_candidates(
api_format="openai:chat", api_format="openai:chat",

View File

@@ -50,7 +50,7 @@ async def test_submit_with_failover_skips_http_500_then_succeeds(
lambda *_args, **_kwargs: "provider", lambda *_args, **_kwargs: "provider",
) )
monkeypatch.setattr( monkeypatch.setattr(
"src.services.cache.aware_scheduler.get_cache_aware_scheduler", "src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
AsyncMock(return_value=None), AsyncMock(return_value=None),
) )
monkeypatch.setattr( monkeypatch.setattr(
@@ -106,7 +106,7 @@ async def test_submit_with_failover_stops_on_client_error(monkeypatch: pytest.Mo
lambda *_args, **_kwargs: "provider", lambda *_args, **_kwargs: "provider",
) )
monkeypatch.setattr( monkeypatch.setattr(
"src.services.cache.aware_scheduler.get_cache_aware_scheduler", "src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
AsyncMock(return_value=None), AsyncMock(return_value=None),
) )
monkeypatch.setattr( monkeypatch.setattr(
@@ -152,7 +152,7 @@ async def test_submit_with_failover_no_eligible_candidates_due_to_auth_type(
lambda *_args, **_kwargs: "provider", lambda *_args, **_kwargs: "provider",
) )
monkeypatch.setattr( monkeypatch.setattr(
"src.services.cache.aware_scheduler.get_cache_aware_scheduler", "src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
AsyncMock(return_value=None), AsyncMock(return_value=None),
) )
monkeypatch.setattr( monkeypatch.setattr(
@@ -194,7 +194,7 @@ async def test_submit_with_failover_filters_missing_billing_rule(
lambda *_args, **_kwargs: "provider", lambda *_args, **_kwargs: "provider",
) )
monkeypatch.setattr( monkeypatch.setattr(
"src.services.cache.aware_scheduler.get_cache_aware_scheduler", "src.services.scheduling.aware_scheduler.get_cache_aware_scheduler",
AsyncMock(return_value=None), AsyncMock(return_value=None),
) )
monkeypatch.setattr( monkeypatch.setattr(

View File

@@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from src.core.api_format.conversion import register_default_normalizers 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, CacheAwareScheduler,
_sort_endpoints_by_family_priority, _sort_endpoints_by_family_priority,
) )

View File

@@ -1,11 +1,13 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from typing import cast
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from src.services.cache.candidate_sorter import CandidateSorter from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.cache.scheduling_config import SchedulingConfig from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.cache.schemas import ProviderCandidate from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.scheduling.schemas import ProviderCandidate
def _make_candidate( def _make_candidate(
@@ -28,12 +30,12 @@ def _make_candidate(
global_priority_by_format={"openai:chat": global_priority}, global_priority_by_format={"openai:chat": global_priority},
) )
return ProviderCandidate( return ProviderCandidate(
provider=provider, provider=cast(Provider, provider),
endpoint=endpoint, endpoint=cast(ProviderEndpoint, endpoint),
key=key, key=cast(ProviderAPIKey, key),
needs_conversion=needs_conversion, needs_conversion=needs_conversion,
provider_api_format="openai:chat", provider_api_format="openai:chat",
) # type: ignore[arg-type] )
def test_priority_sort_global_key_does_not_demote_when_global_keep_priority_enabled() -> None: 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 排序 # 全局 keep_priority_on_conversion=True不做 needs_conversion 降级分组,纯按 global_priority 排序
with patch( 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, return_value=True,
): ):
result = sorter._apply_priority_mode_sort([exact, demoted], db, None, "openai:chat") 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 候选整体排后 # 全局 keep_priority_on_conversion=False需要降级的 convertible 候选整体排后
with patch( 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, return_value=False,
): ):
result = sorter._apply_priority_mode_sort([exact, demoted], db, None, "openai:chat") 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( 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, return_value=False,
): ):
result = sorter._apply_priority_mode_sort( result = sorter._apply_priority_mode_sort(
@@ -163,7 +165,7 @@ def test_priority_sort_provider_mode_demotes_convertible_when_global_keep_priori
) )
with patch( 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, return_value=False,
): ):
result = sorter._apply_priority_mode_sort([demoted, exact], db, None, "openai:chat") 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( 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, return_value=True,
): ):
result = sorter._apply_priority_mode_sort([demoted, exact], db, None, "openai:chat") result = sorter._apply_priority_mode_sort([demoted, exact], db, None, "openai:chat")

View File

@@ -2,7 +2,7 @@ from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from src.services.cache.aware_scheduler import ProviderCandidate from src.services.scheduling.aware_scheduler import ProviderCandidate
def _make_candidate( def _make_candidate(