mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor: 全局适配 ApiFamily/EndpointKind 结构化标识体系
将新的 (ApiFamily, EndpointKind) / `family:kind` 签名体系应用到整个代码库: - API Handlers: 所有 adapter/handler 使用新的签名格式 - Services: provider, model, usage, cache, auth 等服务层适配 - Database: ProviderEndpoint 新增 api_family/endpoint_kind 字段 - Frontend: Provider 管理、Usage 表格等组件适配 - Tests: 更新所有相关测试用例
This commit is contained in:
@@ -11,7 +11,6 @@ from datetime import datetime, timezone
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
# 安全策略配置:当 Redis 不可用时的行为
|
||||
# True = fail-closed(安全优先,拒绝访问)
|
||||
# False = fail-open(可用性优先,允许访问)
|
||||
@@ -71,7 +70,12 @@ class JWTBlacklistService:
|
||||
await redis_client.setex(redis_key, ttl_seconds, reason)
|
||||
|
||||
token_fp = JWTBlacklistService._get_token_hash(token)[:12]
|
||||
logger.info("Token 已加入黑名单: token_fp={} (原因: {}, TTL: {}s)", token_fp, reason, ttl_seconds)
|
||||
logger.info(
|
||||
"Token 已加入黑名单: token_fp={} (原因: {}, TTL: {}s)",
|
||||
token_fp,
|
||||
reason,
|
||||
ttl_seconds,
|
||||
)
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
@@ -94,7 +98,9 @@ class JWTBlacklistService:
|
||||
if redis_client is None:
|
||||
# Redis 不可用时,根据安全策略决定行为
|
||||
if BLACKLIST_FAIL_CLOSED:
|
||||
logger.warning("Redis 不可用,采用 fail-closed 策略拒绝访问(可通过 JWT_BLACKLIST_FAIL_CLOSED=false 改变)")
|
||||
logger.warning(
|
||||
"Redis 不可用,采用 fail-closed 策略拒绝访问(可通过 JWT_BLACKLIST_FAIL_CLOSED=false 改变)"
|
||||
)
|
||||
return True # 返回 True 表示在黑名单中,拒绝访问
|
||||
else:
|
||||
logger.debug("Redis 不可用,采用 fail-open 策略允许访问")
|
||||
|
||||
@@ -172,7 +172,9 @@ class LDAPService:
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def authenticate_with_config(config: dict[str, Any], username: str, password: str) -> dict | None:
|
||||
def authenticate_with_config(
|
||||
config: dict[str, Any], username: str, password: str
|
||||
) -> dict | None:
|
||||
"""
|
||||
LDAP bind 验证
|
||||
|
||||
@@ -186,7 +188,7 @@ class LDAPService:
|
||||
"""
|
||||
try:
|
||||
import ldap3
|
||||
from ldap3 import Server, Connection, SUBTREE
|
||||
from ldap3 import SUBTREE, Connection, Server
|
||||
from ldap3.core.exceptions import LDAPBindError, LDAPSocketOpenError
|
||||
except ImportError:
|
||||
logger.error("ldap3 库未安装")
|
||||
@@ -272,9 +274,7 @@ class LDAPService:
|
||||
|
||||
# 提取用户属性(优先用 LDAP 提供的值,不合法则回退默认)
|
||||
ldap_username = _get_attr_value(user_entry, config["username_attr"], username)
|
||||
email = _get_attr_value(
|
||||
user_entry, config["email_attr"], f"{username}@ldap.local"
|
||||
)
|
||||
email = _get_attr_value(user_entry, config["email_attr"], f"{username}@ldap.local")
|
||||
display_name = _get_attr_value(user_entry, config["display_name_attr"], username)
|
||||
|
||||
logger.info(f"LDAP 认证成功: {username}")
|
||||
@@ -315,7 +315,7 @@ class LDAPService:
|
||||
"""
|
||||
try:
|
||||
import ldap3
|
||||
from ldap3 import Server, Connection
|
||||
from ldap3 import Connection, Server
|
||||
except ImportError:
|
||||
return False, "ldap3 库未安装"
|
||||
|
||||
|
||||
@@ -3,4 +3,3 @@
|
||||
from .service import OAuthService
|
||||
|
||||
__all__ = ["OAuthService"]
|
||||
|
||||
|
||||
@@ -98,7 +98,9 @@ class OAuthProviderBase(ABC):
|
||||
timeout_seconds: float = 5.0,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> httpx.Response:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()
|
||||
) as client:
|
||||
return await client.post(url, data=data, headers=headers)
|
||||
|
||||
async def _http_get(
|
||||
@@ -108,5 +110,7 @@ class OAuthProviderBase(ABC):
|
||||
timeout_seconds: float = 5.0,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> httpx.Response:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(timeout_seconds), verify=get_ssl_context()
|
||||
) as client:
|
||||
return await client.get(url, headers=headers)
|
||||
|
||||
@@ -29,4 +29,3 @@ class OAuthFlowError(Exception):
|
||||
super().__init__(error_code)
|
||||
self.error_code = error_code
|
||||
self.detail = detail
|
||||
|
||||
|
||||
@@ -3,4 +3,3 @@
|
||||
from .linuxdo import LinuxDoOAuthProvider
|
||||
|
||||
__all__ = ["LinuxDoOAuthProvider"]
|
||||
|
||||
|
||||
@@ -93,4 +93,3 @@ def get_oauth_provider_registry() -> OAuthProviderRegistry:
|
||||
if _registry is None:
|
||||
_registry = OAuthProviderRegistry()
|
||||
return _registry
|
||||
|
||||
|
||||
@@ -17,13 +17,13 @@ from src.core.logger import logger
|
||||
from src.core.modules import get_module_registry
|
||||
from src.models.database import OAuthProvider, User, UserOAuthLink
|
||||
from src.services.auth.ldap import LDAPService
|
||||
from src.services.auth.oauth.base import OAuthProviderBase
|
||||
from src.services.auth.oauth.models import OAuthFlowError, OAuthUserInfo
|
||||
from src.services.auth.oauth.registry import get_oauth_provider_registry
|
||||
from src.services.auth.oauth.state import consume_oauth_state, create_oauth_state
|
||||
from src.services.auth.service import AuthService
|
||||
from src.services.cache.user_cache import UserCacheService
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.auth.oauth.base import OAuthProviderBase
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
|
||||
|
||||
@@ -34,7 +34,9 @@ class OAuthService:
|
||||
def _require_module_active(db: Session) -> None:
|
||||
registry = get_module_registry()
|
||||
if not registry.is_active("oauth", db):
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="OAuth 模块未启用")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="OAuth 模块未启用"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_provider_impl(provider_type: str) -> OAuthProviderBase | None:
|
||||
@@ -138,7 +140,9 @@ class OAuthService:
|
||||
if redis is None:
|
||||
raise HTTPException(status_code=503, detail="Redis 不可用")
|
||||
|
||||
state = await create_oauth_state(redis, provider_type=provider_type, action="login", user_id=None)
|
||||
state = await create_oauth_state(
|
||||
redis, provider_type=provider_type, action="login", user_id=None
|
||||
)
|
||||
return provider.get_authorization_url(config, state)
|
||||
|
||||
@staticmethod
|
||||
@@ -163,7 +167,9 @@ class OAuthService:
|
||||
if redis is None:
|
||||
raise HTTPException(status_code=503, detail="Redis 不可用")
|
||||
|
||||
state = await create_oauth_state(redis, provider_type=provider_type, action="bind", user_id=user.id)
|
||||
state = await create_oauth_state(
|
||||
redis, provider_type=provider_type, action="bind", user_id=user.id
|
||||
)
|
||||
return provider.get_authorization_url(config, state)
|
||||
|
||||
@staticmethod
|
||||
@@ -319,7 +325,9 @@ class OAuthService:
|
||||
|
||||
if state_data.action == "bind":
|
||||
try:
|
||||
await OAuthService._handle_bind(db, user_id=state_data.user_id or "", config=config, oauth_user=oauth_user)
|
||||
await OAuthService._handle_bind(
|
||||
db, user_id=state_data.user_id or "", config=config, oauth_user=oauth_user
|
||||
)
|
||||
except OAuthFlowError as exc:
|
||||
return OAuthService._build_frontend_error_redirect(
|
||||
frontend_callback_url, error_code=exc.error_code, error_detail=exc.detail
|
||||
@@ -348,7 +356,10 @@ class OAuthService:
|
||||
}
|
||||
)
|
||||
refresh_token = AuthService.create_refresh_token(
|
||||
data={"user_id": user.id, "created_at": user.created_at.isoformat() if user.created_at else None}
|
||||
data={
|
||||
"user_id": user.id,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
)
|
||||
|
||||
return OAuthService._build_frontend_login_success_redirect(
|
||||
@@ -356,7 +367,9 @@ class OAuthService:
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _handle_login(db: Session, *, config: OAuthProvider, oauth_user: OAuthUserInfo) -> User:
|
||||
async def _handle_login(
|
||||
db: Session, *, config: OAuthProvider, oauth_user: OAuthUserInfo
|
||||
) -> User:
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# 1) 已绑定账号:直接登录
|
||||
@@ -381,7 +394,9 @@ class OAuthService:
|
||||
return linked_user
|
||||
|
||||
# 2) 未绑定账号:可能需要新建用户(受注册开关控制)
|
||||
enable_registration = SystemConfigService.get_config(db, "enable_registration", default=False)
|
||||
enable_registration = SystemConfigService.get_config(
|
||||
db, "enable_registration", default=False
|
||||
)
|
||||
if not enable_registration:
|
||||
raise OAuthFlowError("registration_disabled")
|
||||
|
||||
@@ -399,7 +414,11 @@ class OAuthService:
|
||||
raise OAuthFlowError("email_is_ldap")
|
||||
raise OAuthFlowError("email_is_oauth")
|
||||
|
||||
base_username = oauth_user.username or (email.split("@", 1)[0] if email else None) or f"user_{uuid.uuid4().hex[:8]}"
|
||||
base_username = (
|
||||
oauth_user.username
|
||||
or (email.split("@", 1)[0] if email else None)
|
||||
or f"user_{uuid.uuid4().hex[:8]}"
|
||||
)
|
||||
default_quota = SystemConfigService.get_config(db, "default_user_quota_usd", default=10.0)
|
||||
|
||||
# 生成唯一用户名 + 创建用户(简单重试)
|
||||
@@ -469,7 +488,9 @@ class OAuthService:
|
||||
existing_link.last_login_at = now
|
||||
db.commit()
|
||||
assert existing_user.id is not None
|
||||
await UserCacheService.invalidate_user_cache(existing_user.id, existing_user.email)
|
||||
await UserCacheService.invalidate_user_cache(
|
||||
existing_user.id, existing_user.email
|
||||
)
|
||||
return existing_user
|
||||
raise OAuthFlowError("oauth_already_bound")
|
||||
raise OAuthFlowError("provider_error", "link_create_failed")
|
||||
@@ -724,7 +745,11 @@ class OAuthService:
|
||||
OAuthService._validate_redirect_uri(data.redirect_uri)
|
||||
|
||||
# 覆盖端点:必须 https 且 hostname 命中 provider 白名单
|
||||
for field_name in ("authorization_url_override", "token_url_override", "userinfo_url_override"):
|
||||
for field_name in (
|
||||
"authorization_url_override",
|
||||
"token_url_override",
|
||||
"userinfo_url_override",
|
||||
):
|
||||
value = getattr(data, field_name)
|
||||
if value:
|
||||
OAuthService._validate_url_override(provider, value)
|
||||
@@ -784,7 +809,9 @@ class OAuthService:
|
||||
|
||||
async def _reachable(url: str) -> bool:
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()
|
||||
) as client:
|
||||
await client.get(url)
|
||||
return True
|
||||
except Exception:
|
||||
@@ -799,7 +826,9 @@ class OAuthService:
|
||||
if cfg.client_secret_encrypted:
|
||||
# 使用无效 code 做一次 token 请求(仅做粗略判定)
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(5.0), verify=get_ssl_context()
|
||||
) as client:
|
||||
resp = await client.post(
|
||||
token_url,
|
||||
data={
|
||||
@@ -859,7 +888,9 @@ class OAuthService:
|
||||
|
||||
async def _reachable(url: str) -> bool:
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(5.0), follow_redirects=False, verify=get_ssl_context()
|
||||
) as client:
|
||||
await client.get(url)
|
||||
return True
|
||||
except Exception:
|
||||
@@ -874,7 +905,9 @@ class OAuthService:
|
||||
if client_secret:
|
||||
# 使用无效 code 做一次 token 请求(仅做粗略判定)
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(5.0), verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(5.0), verify=get_ssl_context()
|
||||
) as client:
|
||||
resp = await client.post(
|
||||
token_url,
|
||||
data={
|
||||
@@ -926,7 +959,10 @@ class OAuthService:
|
||||
if not link:
|
||||
raise InvalidRequestException("未绑定该 Provider")
|
||||
|
||||
total_links = db.query(func.count(UserOAuthLink.id)).filter(UserOAuthLink.user_id == user.id).scalar() or 0
|
||||
total_links = (
|
||||
db.query(func.count(UserOAuthLink.id)).filter(UserOAuthLink.user_id == user.id).scalar()
|
||||
or 0
|
||||
)
|
||||
|
||||
if user.auth_source == AuthSource.OAUTH and total_links <= 1:
|
||||
raise InvalidRequestException("OAUTH 用户必须至少保留一个 OAuth 绑定")
|
||||
@@ -935,7 +971,11 @@ class OAuthService:
|
||||
if user.auth_source == AuthSource.LOCAL and not user.password_hash and total_links <= 1:
|
||||
raise InvalidRequestException("请先设置密码后再解绑")
|
||||
|
||||
if LDAPService.is_ldap_exclusive(db) and user.auth_source == AuthSource.LOCAL and user.role != UserRole.ADMIN:
|
||||
if (
|
||||
LDAPService.is_ldap_exclusive(db)
|
||||
and user.auth_source == AuthSource.LOCAL
|
||||
and user.role != UserRole.ADMIN
|
||||
):
|
||||
# ldap_exclusive=true 时,普通本地用户解绑最后一个 OAuth 会锁死(密码登录被禁用)
|
||||
if total_links <= 1:
|
||||
raise InvalidRequestException("当前处于 LDAP 专属模式,解绑后将无法登录")
|
||||
|
||||
@@ -3,9 +3,9 @@ from __future__ import annotations
|
||||
import json
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, cast
|
||||
from collections.abc import Awaitable
|
||||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
|
||||
@@ -20,20 +20,20 @@ from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.config import config
|
||||
from src.core.logger import logger
|
||||
from src.core.enums import AuthSource
|
||||
from src.core.exceptions import ForbiddenException
|
||||
from src.core.logger import logger
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ManagementToken
|
||||
|
||||
from src.models.database import ApiKey, User, UserRole
|
||||
from src.services.auth.jwt_blacklist import JWTBlacklistService
|
||||
from src.services.auth.ldap import LDAPService
|
||||
from src.services.cache.user_cache import UserCacheService
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
|
||||
|
||||
# API Key last_used_at 更新节流配置
|
||||
# 同一个 API Key 在此时间间隔内只会更新一次 last_used_at
|
||||
_LAST_USED_UPDATE_INTERVAL = 60 # 秒
|
||||
@@ -243,9 +243,8 @@ class AuthService:
|
||||
# 登录校验必须读取密码哈希,不能使用不包含 password_hash 的缓存对象
|
||||
# 支持邮箱或用户名登录
|
||||
from sqlalchemy import or_
|
||||
user = db.query(User).filter(
|
||||
or_(User.email == email, User.username == email)
|
||||
).first()
|
||||
|
||||
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
|
||||
|
||||
if not user:
|
||||
logger.warning(f"登录失败 - 用户不存在: {email}")
|
||||
@@ -295,7 +294,9 @@ class AuthService:
|
||||
注意:使用 with_for_update() 防止并发首次登录创建重复用户
|
||||
"""
|
||||
ldap_dn = (ldap_user.get("ldap_dn") or "").strip() or None
|
||||
ldap_username = (ldap_user.get("ldap_username") or ldap_user.get("username") or "").strip() or None
|
||||
ldap_username = (
|
||||
ldap_user.get("ldap_username") or ldap_user.get("username") or ""
|
||||
).strip() or None
|
||||
email = ldap_user["email"]
|
||||
|
||||
# 优先用稳定标识查找,避免邮箱变更/用户名冲突导致重复建号
|
||||
@@ -333,11 +334,7 @@ class AuthService:
|
||||
|
||||
# 同步邮箱(LDAP 侧邮箱变更时更新;若新邮箱已被占用则拒绝)
|
||||
if user.email != email:
|
||||
email_taken = (
|
||||
db.query(User)
|
||||
.filter(User.email == email, User.id != user.id)
|
||||
.first()
|
||||
)
|
||||
email_taken = db.query(User).filter(User.email == email, User.id != user.id).first()
|
||||
if email_taken:
|
||||
logger.warning(f"LDAP 登录拒绝 - 新邮箱已被占用: {email}")
|
||||
return None
|
||||
@@ -370,7 +367,9 @@ class AuthService:
|
||||
logger.info(f"LDAP 用户名冲突,使用新用户名: {ldap_user['username']} -> {username}")
|
||||
|
||||
# 读取系统配置的默认配额
|
||||
default_quota = SystemConfigService.get_config(db, "default_user_quota_usd", default=10.0)
|
||||
default_quota = SystemConfigService.get_config(
|
||||
db, "default_user_quota_usd", default=10.0
|
||||
)
|
||||
|
||||
# 创建新用户
|
||||
user = User(
|
||||
@@ -408,7 +407,9 @@ class AuthService:
|
||||
logger.error(f"LDAP 用户创建失败(用户名冲突重试耗尽): {username}")
|
||||
return None
|
||||
username = f"{base_username}_ldap_{int(time.time())}{uuid.uuid4().hex[:4]}"
|
||||
logger.warning(f"LDAP 用户创建用户名冲突,重试 ({attempt + 1}/{max_retries}): {username}")
|
||||
logger.warning(
|
||||
f"LDAP 用户创建用户名冲突,重试 ({attempt + 1}/{max_retries}): {username}"
|
||||
)
|
||||
else:
|
||||
# 其他约束冲突,不重试
|
||||
logger.error(f"LDAP 用户创建失败 - 未知数据库约束冲突: {e}")
|
||||
@@ -458,8 +459,10 @@ class AuthService:
|
||||
if not is_balance_ok:
|
||||
# 获取剩余余额用于日志
|
||||
remaining_balance = ApiKeyService.get_remaining_balance(key_record)
|
||||
logger.warning(f"API认证失败 - 余额不足 "
|
||||
f"(已用: ${key_record.balance_used_usd:.4f}, 剩余: ${remaining_balance:.4f})")
|
||||
logger.warning(
|
||||
f"API认证失败 - 余额不足 "
|
||||
f"(已用: ${key_record.balance_used_usd:.4f}, 剩余: ${remaining_balance:.4f})"
|
||||
)
|
||||
return None
|
||||
|
||||
# 获取用户
|
||||
@@ -492,7 +495,9 @@ class AuthService:
|
||||
|
||||
# 检查美元配额
|
||||
if user.used_usd + estimated_cost > user.quota_usd:
|
||||
logger.warning(f"用户配额不足: {user.email} (已用: ${user.used_usd:.2f}, 配额: ${user.quota_usd:.2f})")
|
||||
logger.warning(
|
||||
f"用户配额不足: {user.email} (已用: ${user.used_usd:.2f}, 配额: ${user.quota_usd:.2f})"
|
||||
)
|
||||
return False
|
||||
|
||||
return True
|
||||
@@ -509,7 +514,9 @@ class AuthService:
|
||||
if role_rank.get(user.role, -1) >= role_rank.get(required_role, 999):
|
||||
return True
|
||||
|
||||
logger.warning(f"权限不足: 用户 {user.email} 角色 {user.role.value} < 需要 {required_role.value}")
|
||||
logger.warning(
|
||||
f"权限不足: 用户 {user.email} 角色 {user.role.value} < 需要 {required_role.value}"
|
||||
)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
@@ -668,9 +675,7 @@ class AuthService:
|
||||
|
||||
# 检查 IP 白名单
|
||||
if not token_record.is_ip_allowed(client_ip):
|
||||
logger.warning(
|
||||
f"Management Token IP 限制 - Token: {token_record.id}, IP: {client_ip}"
|
||||
)
|
||||
logger.warning(f"Management Token IP 限制 - Token: {token_record.id}, IP: {client_ip}")
|
||||
AuditService.log_event(
|
||||
db=db,
|
||||
event_type=AuditEventType.MANAGEMENT_TOKEN_IP_BLOCKED,
|
||||
|
||||
55
src/services/cache/affinity_manager.py
vendored
55
src/services/cache/affinity_manager.py
vendored
@@ -31,7 +31,6 @@ from src.config.constants import CacheTTL
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
|
||||
class CacheAffinity(NamedTuple):
|
||||
"""缓存亲和性信息"""
|
||||
|
||||
@@ -78,7 +77,9 @@ class CacheAffinityManager:
|
||||
# 默认缓存TTL(秒)- 使用统一常量
|
||||
DEFAULT_CACHE_TTL = CacheTTL.CACHE_AFFINITY
|
||||
|
||||
def __init__(self, redis_client: Any | None = None, default_ttl: int = DEFAULT_CACHE_TTL) -> None:
|
||||
def __init__(
|
||||
self, redis_client: Any | None = None, default_ttl: int = DEFAULT_CACHE_TTL
|
||||
) -> None:
|
||||
"""
|
||||
初始化缓存亲和性管理器
|
||||
|
||||
@@ -114,7 +115,9 @@ class CacheAffinityManager:
|
||||
if self.redis:
|
||||
logger.debug("CacheAffinityManager: 使用Redis存储")
|
||||
else:
|
||||
logger.debug("CacheAffinityManager: Redis不可用,回退到内存存储(仅适用于单实例/开发环境)")
|
||||
logger.debug(
|
||||
"CacheAffinityManager: Redis不可用,回退到内存存储(仅适用于单实例/开发环境)"
|
||||
)
|
||||
|
||||
def _is_memory_backend(self) -> bool:
|
||||
"""是否处于内存模式"""
|
||||
@@ -172,8 +175,7 @@ class CacheAffinityManager:
|
||||
清理的条目数量
|
||||
"""
|
||||
expired_keys = [
|
||||
key for key, (expire_at, _) in self._l1_cache.items()
|
||||
if current_time > expire_at
|
||||
key for key, (expire_at, _) in self._l1_cache.items() if current_time > expire_at
|
||||
]
|
||||
for key in expired_keys:
|
||||
self._l1_cache.pop(key, None)
|
||||
@@ -181,8 +183,7 @@ class CacheAffinityManager:
|
||||
# 如果缓存仍然过大,按过期时间排序移除最旧的条目
|
||||
if len(self._l1_cache) > self._l1_max_size:
|
||||
sorted_items = sorted(
|
||||
self._l1_cache.items(),
|
||||
key=lambda x: x[1][0] # 按 expire_at 排序
|
||||
self._l1_cache.items(), key=lambda x: x[1][0] # 按 expire_at 排序
|
||||
)
|
||||
# 移除最旧的 20% 条目
|
||||
remove_count = len(self._l1_cache) - int(self._l1_max_size * 0.8)
|
||||
@@ -191,7 +192,9 @@ class CacheAffinityManager:
|
||||
expired_keys.extend([k for k, _ in sorted_items[:remove_count]])
|
||||
|
||||
if expired_keys:
|
||||
logger.debug(f"L1 缓存清理: 移除 {len(expired_keys)} 个条目,当前 {len(self._l1_cache)} 个")
|
||||
logger.debug(
|
||||
f"L1 缓存清理: 移除 {len(expired_keys)} 个条目,当前 {len(self._l1_cache)} 个"
|
||||
)
|
||||
|
||||
return len(expired_keys)
|
||||
|
||||
@@ -371,16 +374,20 @@ class CacheAffinityManager:
|
||||
or existing_affinity.key_id != key_id
|
||||
):
|
||||
self._stats["key_switches"] += 1
|
||||
logger.debug(f"Key {affinity_key[:8]}... 在 {api_format} 格式下切换后端: "
|
||||
logger.debug(
|
||||
f"Key {affinity_key[:8]}... 在 {api_format} 格式下切换后端: "
|
||||
f"[{existing_affinity.provider_id[:8]}.../{existing_affinity.endpoint_id[:8]}.../"
|
||||
f"{existing_affinity.key_id[:8]}...] → "
|
||||
f"[{provider_id[:8]}.../{endpoint_id[:8]}.../{key_id[:8]}...], 重置计数器")
|
||||
f"[{provider_id[:8]}.../{endpoint_id[:8]}.../{key_id[:8]}...], 重置计数器"
|
||||
)
|
||||
created_at = current_time
|
||||
request_count = 1
|
||||
else:
|
||||
logger.debug(f"刷新缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
logger.debug(
|
||||
f"刷新缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
f"provider={provider_id[:8]}..., endpoint={endpoint_id[:8]}..., "
|
||||
f"provider_key={key_id[:8]}..., ttl+={ttl}s")
|
||||
f"provider_key={key_id[:8]}..., ttl+={ttl}s"
|
||||
)
|
||||
else:
|
||||
created_at = current_time
|
||||
request_count = 1
|
||||
@@ -399,9 +406,11 @@ class CacheAffinityManager:
|
||||
|
||||
await self._save_affinity_dict(cache_key, ttl, affinity_dict)
|
||||
|
||||
logger.debug(f"设置缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
logger.debug(
|
||||
f"设置缓存亲和性: key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
f"model={model_name}, provider={provider_id[:8]}..., endpoint={endpoint_id[:8]}..., "
|
||||
f"provider_key={key_id[:8]}..., ttl={ttl}s")
|
||||
f"provider_key={key_id[:8]}..., ttl={ttl}s"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.exception(f"设置缓存亲和性失败: {e}")
|
||||
|
||||
@@ -443,8 +452,10 @@ class CacheAffinityManager:
|
||||
should_invalidate = False
|
||||
|
||||
if not should_invalidate:
|
||||
logger.debug(f"跳过失效: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
f"model={model_name}, 过滤条件不匹配 (key={key_id}, provider={provider_id}, endpoint={endpoint_id})")
|
||||
logger.debug(
|
||||
f"跳过失效: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
f"model={model_name}, 过滤条件不匹配 (key={key_id}, provider={provider_id}, endpoint={endpoint_id})"
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
@@ -454,10 +465,12 @@ class CacheAffinityManager:
|
||||
|
||||
self._stats["cache_invalidations"] += 1
|
||||
|
||||
logger.debug(f"失效缓存亲和性: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
logger.debug(
|
||||
f"失效缓存亲和性: affinity_key={affinity_key[:8]}..., api_format={api_format}, "
|
||||
f"model={model_name}, provider={existing_affinity.provider_id[:8]}..., "
|
||||
f"endpoint={existing_affinity.endpoint_id[:8]}..., "
|
||||
f"provider_key={existing_affinity.key_id[:8]}...")
|
||||
f"provider_key={existing_affinity.key_id[:8]}..."
|
||||
)
|
||||
except Exception as e:
|
||||
logger.exception(f"删除缓存亲和性失败: {e}")
|
||||
|
||||
@@ -493,8 +506,10 @@ class CacheAffinityManager:
|
||||
self._stats["cache_invalidations"] += 1
|
||||
|
||||
if invalidated_count > 0:
|
||||
logger.debug(f"批量失效Provider缓存亲和性: provider={provider_id[:8]}..., "
|
||||
f"失效数量={invalidated_count}")
|
||||
logger.debug(
|
||||
f"批量失效Provider缓存亲和性: provider={provider_id[:8]}..., "
|
||||
f"失效数量={invalidated_count}"
|
||||
)
|
||||
|
||||
return invalidated_count
|
||||
|
||||
|
||||
195
src/services/cache/aware_scheduler.py
vendored
195
src/services/cache/aware_scheduler.py
vendored
@@ -30,7 +30,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
import hashlib
|
||||
import random
|
||||
import re
|
||||
@@ -40,7 +39,8 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.enums import EndpointKind
|
||||
from src.core.api_format.signature import make_signature_key, parse_signature_key
|
||||
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import (
|
||||
@@ -60,7 +60,7 @@ from src.services.cache.affinity_manager import (
|
||||
)
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.provider.format import normalize_api_format
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.rate_limit.adaptive_reservation import (
|
||||
AdaptiveReservationManager,
|
||||
get_adaptive_reservation_manager,
|
||||
@@ -194,7 +194,7 @@ class CacheAwareScheduler:
|
||||
self,
|
||||
db: Session,
|
||||
affinity_key: str,
|
||||
api_format: str | APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
excluded_endpoints: list[str] | None = None,
|
||||
excluded_keys: list[str] | None = None,
|
||||
@@ -222,14 +222,14 @@ class CacheAwareScheduler:
|
||||
excluded_endpoints_set = set(excluded_endpoints or [])
|
||||
excluded_keys_set = set(excluded_keys or [])
|
||||
|
||||
normalized_format = normalize_api_format(api_format)
|
||||
normalized_format = normalize_endpoint_signature(api_format)
|
||||
|
||||
logger.debug(
|
||||
f"[CacheAwareScheduler] select_with_cache_affinity: "
|
||||
f"affinity_key={affinity_key[:8]}..., api_format={normalized_format.value}, model={model_name}"
|
||||
f"affinity_key={affinity_key[:8]}..., api_format={normalized_format}, model={model_name}"
|
||||
)
|
||||
|
||||
self._metrics["last_api_format"] = normalized_format.value
|
||||
self._metrics["last_api_format"] = normalized_format
|
||||
self._metrics["last_model_name"] = model_name
|
||||
|
||||
provider_offset = 0
|
||||
@@ -290,22 +290,17 @@ class CacheAwareScheduler:
|
||||
f"并发状态[{snapshot.describe()}]"
|
||||
)
|
||||
|
||||
if key.cache_ttl_minutes > 0 and global_model_id:
|
||||
ttl = key.cache_ttl_minutes * 60 if key.cache_ttl_minutes > 0 else None
|
||||
api_format_str = (
|
||||
normalized_format.value
|
||||
if isinstance(normalized_format, APIFormat)
|
||||
else normalized_format
|
||||
)
|
||||
await self.set_cache_affinity(
|
||||
affinity_key=affinity_key,
|
||||
provider_id=str(provider.id),
|
||||
endpoint_id=str(endpoint.id),
|
||||
key_id=str(key.id),
|
||||
api_format=api_format_str,
|
||||
global_model_id=global_model_id,
|
||||
ttl=int(ttl) if ttl is not None else None,
|
||||
)
|
||||
if key.cache_ttl_minutes > 0 and global_model_id:
|
||||
ttl = key.cache_ttl_minutes * 60 if key.cache_ttl_minutes > 0 else None
|
||||
await self.set_cache_affinity(
|
||||
affinity_key=affinity_key,
|
||||
provider_id=str(provider.id),
|
||||
endpoint_id=str(endpoint.id),
|
||||
key_id=str(key.id),
|
||||
api_format=normalized_format,
|
||||
global_model_id=global_model_id,
|
||||
ttl=int(ttl) if ttl is not None else None,
|
||||
)
|
||||
|
||||
if is_cached_user:
|
||||
self._metrics["cache_hits"] += 1
|
||||
@@ -549,7 +544,7 @@ class CacheAwareScheduler:
|
||||
async def list_all_candidates(
|
||||
self,
|
||||
db: Session,
|
||||
api_format: str | APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str | None = None,
|
||||
user_api_key: ApiKey | None = None,
|
||||
@@ -584,7 +579,13 @@ class CacheAwareScheduler:
|
||||
"""
|
||||
await self._ensure_initialized()
|
||||
|
||||
target_format = normalize_api_format(api_format)
|
||||
target_format = normalize_endpoint_signature(api_format)
|
||||
|
||||
logger.debug(
|
||||
"[Scheduler] list_all_candidates: model=%s, api_format=%s",
|
||||
model_name,
|
||||
target_format,
|
||||
)
|
||||
|
||||
# 0. 解析 model_name 到 GlobalModel(支持直接匹配和映射名匹配,使用 ModelCacheService)
|
||||
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
||||
@@ -595,6 +596,12 @@ class CacheAwareScheduler:
|
||||
logger.warning(f"GlobalModel not found: {model_name}")
|
||||
raise ModelNotSupportedException(model=model_name)
|
||||
|
||||
logger.debug(
|
||||
"[Scheduler] GlobalModel resolved: id=%s, name=%s",
|
||||
global_model.id,
|
||||
global_model.name,
|
||||
)
|
||||
|
||||
# 使用 GlobalModel.id 作为缓存亲和性的模型标识,确保映射名和规范名都能命中同一个缓存
|
||||
global_model_id: str = str(global_model.id)
|
||||
|
||||
@@ -613,11 +620,10 @@ class CacheAwareScheduler:
|
||||
|
||||
# 0.1 检查 API 格式是否被允许
|
||||
if allowed_api_formats is not None:
|
||||
# 统一转为大写比较,兼容数据库中存储的大小写
|
||||
allowed_upper = {f.upper() for f in allowed_api_formats}
|
||||
if target_format.value.upper() not in allowed_upper:
|
||||
allowed_norm = {normalize_endpoint_signature(f) for f in allowed_api_formats if f}
|
||||
if target_format not in allowed_norm:
|
||||
logger.debug(
|
||||
f"API Key {user_api_key.id[:8] if user_api_key else 'N/A'}... 不允许使用 API 格式 {target_format.value}, "
|
||||
f"API Key {user_api_key.id[:8] if user_api_key else 'N/A'}... 不允许使用 API 格式 {target_format}, "
|
||||
f"允许的格式: {allowed_api_formats}"
|
||||
)
|
||||
return [], global_model_id
|
||||
@@ -642,6 +648,20 @@ class CacheAwareScheduler:
|
||||
provider_limit=provider_limit,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"[Scheduler] Found %d active providers",
|
||||
len(providers),
|
||||
)
|
||||
for p in providers:
|
||||
logger.debug(
|
||||
"[Scheduler] Provider: id=%s, name=%s, is_active=%s, endpoints=%d, models=%d",
|
||||
p.id[:8] if p.id else "N/A",
|
||||
p.name,
|
||||
p.is_active,
|
||||
len(p.endpoints) if p.endpoints else 0,
|
||||
len(p.models) if p.models else 0,
|
||||
)
|
||||
|
||||
if not providers:
|
||||
return [], global_model_id
|
||||
|
||||
@@ -660,6 +680,7 @@ class CacheAwareScheduler:
|
||||
|
||||
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤)
|
||||
from src.config.settings import config
|
||||
|
||||
global_conversion_enabled = config.format_conversion_enabled
|
||||
candidates = await self._build_candidates(
|
||||
db=db,
|
||||
@@ -675,7 +696,7 @@ class CacheAwareScheduler:
|
||||
)
|
||||
|
||||
# 3. 应用优先级模式排序
|
||||
candidates = self._apply_priority_mode_sort(candidates, affinity_key, target_format.value)
|
||||
candidates = self._apply_priority_mode_sort(candidates, affinity_key, target_format)
|
||||
|
||||
# 更新指标
|
||||
self._metrics["total_candidates"] += len(candidates)
|
||||
@@ -683,7 +704,7 @@ class CacheAwareScheduler:
|
||||
|
||||
logger.debug(
|
||||
f"预先获取到 {len(candidates)} 个可用组合 "
|
||||
f"(api_format={target_format.value}, model={model_name})"
|
||||
f"(api_format={target_format}, model={model_name})"
|
||||
)
|
||||
|
||||
# 4. 根据调度模式应用不同的排序策略
|
||||
@@ -698,7 +719,7 @@ class CacheAwareScheduler:
|
||||
)
|
||||
elif self.scheduling_mode == self.SCHEDULING_MODE_LOAD_BALANCE:
|
||||
# 负载均衡模式:忽略缓存,同优先级内随机轮换
|
||||
candidates = self._apply_load_balance(candidates, target_format.value)
|
||||
candidates = self._apply_load_balance(candidates, target_format)
|
||||
for candidate in candidates:
|
||||
candidate.is_cached = False
|
||||
else:
|
||||
@@ -877,11 +898,14 @@ class CacheAwareScheduler:
|
||||
|
||||
mapping_api_formats = raw.get("api_formats")
|
||||
if api_format and mapping_api_formats:
|
||||
if (
|
||||
isinstance(mapping_api_formats, list)
|
||||
and api_format not in mapping_api_formats
|
||||
):
|
||||
continue
|
||||
# 新模式:endpoint signature(family:kind),按小写 canonical 比较
|
||||
if isinstance(mapping_api_formats, list):
|
||||
target = str(api_format).strip().lower()
|
||||
allowed = {
|
||||
str(fmt).strip().lower() for fmt in mapping_api_formats if fmt
|
||||
}
|
||||
if target not in allowed:
|
||||
continue
|
||||
|
||||
provider_model_names.add(name.strip())
|
||||
|
||||
@@ -982,7 +1006,7 @@ class CacheAwareScheduler:
|
||||
self,
|
||||
db: Session,
|
||||
providers: list[Provider],
|
||||
client_format: APIFormat,
|
||||
client_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str | None,
|
||||
model_mappings: list[str] | None = None,
|
||||
@@ -1014,9 +1038,21 @@ class CacheAwareScheduler:
|
||||
from src.core.api_format.conversion.compatibility import is_format_compatible
|
||||
|
||||
candidates: list[ProviderCandidate] = []
|
||||
client_format_str = client_format.value
|
||||
client_format_str = normalize_endpoint_signature(client_format)
|
||||
client_sig = parse_signature_key(client_format_str)
|
||||
client_family, client_kind = client_sig.api_family, client_sig.endpoint_kind
|
||||
# chat/cli 互相可回退(用于同协议族下的端点变体),video/image 等不跨类回退
|
||||
if client_kind in {EndpointKind.CHAT, EndpointKind.CLI}:
|
||||
allowed_kinds = {EndpointKind.CHAT, EndpointKind.CLI}
|
||||
else:
|
||||
allowed_kinds = {client_kind}
|
||||
|
||||
for provider in providers:
|
||||
logger.debug(
|
||||
"[Scheduler] Checking provider: %s, endpoints=%d",
|
||||
provider.name,
|
||||
len(provider.endpoints) if provider.endpoints else 0,
|
||||
)
|
||||
# 按端点格式分别判断兼容性与模型/Key 可用性:
|
||||
# - 同格式端点优先(needs_conversion=False)
|
||||
# - 跨格式端点次之(needs_conversion=True)
|
||||
@@ -1026,14 +1062,62 @@ class CacheAwareScheduler:
|
||||
exact_candidates: list[ProviderCandidate] = []
|
||||
convertible_candidates: list[ProviderCandidate] = []
|
||||
|
||||
for endpoint in provider.endpoints:
|
||||
if not endpoint.is_active:
|
||||
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
|
||||
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
|
||||
# - chat/cli 请求允许互相回退(优先同 kind)
|
||||
# - video 等请求只允许同 kind
|
||||
endpoints = list(provider.endpoints or [])
|
||||
allowed_kind_values = {k.value for k in allowed_kinds}
|
||||
preferred: list[ProviderEndpoint] = []
|
||||
preferred_other_family: list[ProviderEndpoint] = []
|
||||
fallback: list[ProviderEndpoint] = []
|
||||
fallback_other_family: list[ProviderEndpoint] = []
|
||||
|
||||
for ep in endpoints:
|
||||
if not getattr(ep, "is_active", False):
|
||||
continue
|
||||
|
||||
endpoint_format_str = (
|
||||
endpoint.api_format
|
||||
if isinstance(endpoint.api_format, str)
|
||||
else endpoint.api_format.value
|
||||
raw_family = getattr(ep, "api_family", None)
|
||||
raw_kind = getattr(ep, "endpoint_kind", None)
|
||||
if not isinstance(raw_family, str) or not raw_family.strip():
|
||||
continue
|
||||
if not isinstance(raw_kind, str) or not raw_kind.strip():
|
||||
continue
|
||||
|
||||
ep_family = raw_family.strip().lower()
|
||||
ep_kind = raw_kind.strip().lower()
|
||||
|
||||
if allowed_kind_values and ep_kind not in allowed_kind_values:
|
||||
continue
|
||||
|
||||
same_family = ep_family == client_family.value
|
||||
same_kind = ep_kind == client_kind.value
|
||||
if same_kind and same_family:
|
||||
preferred.append(ep)
|
||||
elif same_kind:
|
||||
preferred_other_family.append(ep)
|
||||
elif same_family:
|
||||
fallback.append(ep)
|
||||
else:
|
||||
fallback_other_family.append(ep)
|
||||
|
||||
endpoints = preferred + preferred_other_family + fallback + fallback_other_family
|
||||
|
||||
for endpoint in endpoints:
|
||||
logger.debug(
|
||||
"[Scheduler] Checking endpoint: family=%s, kind=%s, is_active=%s, base_url=%s",
|
||||
getattr(endpoint, "api_family", None),
|
||||
getattr(endpoint, "endpoint_kind", None),
|
||||
getattr(endpoint, "is_active", None),
|
||||
(endpoint.base_url[:50] if endpoint.base_url else "N/A"),
|
||||
)
|
||||
if not endpoint.is_active:
|
||||
logger.debug("[Scheduler] Endpoint skipped: not active")
|
||||
continue
|
||||
|
||||
endpoint_format_str = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
|
||||
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
|
||||
@@ -1043,6 +1127,13 @@ class CacheAwareScheduler:
|
||||
is_stream,
|
||||
global_conversion_enabled,
|
||||
)
|
||||
logger.debug(
|
||||
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, reason=%s",
|
||||
client_format_str,
|
||||
endpoint_format_str,
|
||||
is_compatible,
|
||||
_compat_reason,
|
||||
)
|
||||
if not is_compatible:
|
||||
continue
|
||||
|
||||
@@ -1059,6 +1150,13 @@ class CacheAwareScheduler:
|
||||
supports_model, skip_reason, _model_caps, provider_model_names = (
|
||||
model_support_cache[endpoint_format_str]
|
||||
)
|
||||
logger.debug(
|
||||
"[Scheduler] Model support: provider=%s, model=%s, supports=%s, reason=%s",
|
||||
provider.name,
|
||||
model_name,
|
||||
supports_model,
|
||||
skip_reason,
|
||||
)
|
||||
if not supports_model:
|
||||
logger.debug(
|
||||
f"Provider {provider.name} 端点 {endpoint_format_str} 不支持模型 {model_name}: {skip_reason}"
|
||||
@@ -1111,7 +1209,7 @@ class CacheAwareScheduler:
|
||||
skip_reason=key_skip_reason,
|
||||
mapping_matched_model=mapping_matched_model,
|
||||
needs_conversion=needs_conversion,
|
||||
provider_api_format=str(endpoint_format_str or "").upper(),
|
||||
provider_api_format=str(endpoint_format_str or ""),
|
||||
)
|
||||
|
||||
if needs_conversion:
|
||||
@@ -1132,7 +1230,7 @@ class CacheAwareScheduler:
|
||||
self,
|
||||
candidates: list[ProviderCandidate],
|
||||
affinity_key: str,
|
||||
api_format: APIFormat,
|
||||
api_format: str,
|
||||
global_model_id: str,
|
||||
) -> list[ProviderCandidate]:
|
||||
"""
|
||||
@@ -1151,7 +1249,7 @@ class CacheAwareScheduler:
|
||||
"""
|
||||
try:
|
||||
# 查询该亲和性标识符在当前 API 格式和模型下的缓存亲和性
|
||||
api_format_str = api_format.value if isinstance(api_format, APIFormat) else api_format
|
||||
api_format_str = str(api_format)
|
||||
affinity = await self._affinity_manager.get_affinity(
|
||||
affinity_key, api_format_str, global_model_id
|
||||
)
|
||||
@@ -1164,6 +1262,7 @@ class CacheAwareScheduler:
|
||||
|
||||
# 判断候选是否应该被降级(用于分组)
|
||||
from src.config.settings import config
|
||||
|
||||
global_keep_priority = config.keep_priority_on_conversion
|
||||
|
||||
def should_demote(c: ProviderCandidate) -> bool:
|
||||
|
||||
1
src/services/cache/backend.py
vendored
1
src/services/cache/backend.py
vendored
@@ -20,7 +20,6 @@ from collections import OrderedDict
|
||||
from typing import Any
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
from src.core.logger import logger
|
||||
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.logger import logger
|
||||
|
||||
2
src/services/cache/invalidation.py
vendored
2
src/services/cache/invalidation.py
vendored
@@ -4,10 +4,10 @@
|
||||
统一管理各种缓存的失效逻辑
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
|
||||
8
src/services/cache/model_cache.py
vendored
8
src/services/cache/model_cache.py
vendored
@@ -493,9 +493,7 @@ class ModelCacheService:
|
||||
model_mapping_conflict_total.inc()
|
||||
|
||||
# 按名称排序确保确定性
|
||||
result_global_model = sorted(
|
||||
mapping_matches, key=lambda gm: gm.name or ""
|
||||
)[0]
|
||||
result_global_model = sorted(mapping_matches, key=lambda gm: gm.name or "")[0]
|
||||
global_model_dict = ModelCacheService._global_model_to_dict(result_global_model)
|
||||
await CacheService.set(
|
||||
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
|
||||
@@ -509,9 +507,7 @@ class ModelCacheService:
|
||||
# 6. 完全未找到
|
||||
resolution_method = "not_found"
|
||||
# 未找到匹配,缓存负结果
|
||||
await CacheService.set(
|
||||
cache_key, "NOT_FOUND", ttl_seconds=ModelCacheService.CACHE_TTL
|
||||
)
|
||||
await CacheService.set(cache_key, "NOT_FOUND", ttl_seconds=ModelCacheService.CACHE_TTL)
|
||||
logger.debug(f"GlobalModel 未找到(映射解析): {normalized_name}")
|
||||
return None
|
||||
|
||||
|
||||
25
src/services/cache/provider_cache.py
vendored
25
src/services/cache/provider_cache.py
vendored
@@ -5,7 +5,6 @@ Provider 缓存服务 - 减少 Provider 和 ProviderAPIKey 查询
|
||||
这些数据在 UsageService.record_usage() 中被频繁查询但变化不频繁。
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -44,9 +43,9 @@ class ProviderCacheService:
|
||||
计算后的 rate_multiplier
|
||||
"""
|
||||
if api_format and rate_multipliers:
|
||||
format_upper = api_format.upper()
|
||||
if format_upper in rate_multipliers:
|
||||
return float(rate_multipliers[format_upper])
|
||||
format_key = str(api_format).strip().lower()
|
||||
if format_key in rate_multipliers:
|
||||
return float(rate_multipliers[format_key])
|
||||
return 1.0
|
||||
|
||||
@staticmethod
|
||||
@@ -67,13 +66,15 @@ class ProviderCacheService:
|
||||
rate_multiplier 或 None(如果找不到)
|
||||
"""
|
||||
# 缓存键包含 api_format
|
||||
format_suffix = api_format.upper() if api_format else "default"
|
||||
format_suffix = str(api_format).strip().lower() if api_format else "default"
|
||||
cache_key = f"provider_api_key:rate_multiplier:{provider_api_key_id}:{format_suffix}"
|
||||
|
||||
# 1. 尝试从缓存获取
|
||||
cached_data = await CacheService.get(cache_key)
|
||||
if cached_data is not None:
|
||||
logger.debug(f"ProviderAPIKey rate_multiplier 缓存命中: {provider_api_key_id[:8]}... format={format_suffix}")
|
||||
logger.debug(
|
||||
f"ProviderAPIKey rate_multiplier 缓存命中: {provider_api_key_id[:8]}... format={format_suffix}"
|
||||
)
|
||||
# 缓存的 "NOT_FOUND" 表示数据库中不存在
|
||||
if cached_data == "NOT_FOUND":
|
||||
return None
|
||||
@@ -95,7 +96,9 @@ class ProviderCacheService:
|
||||
await CacheService.set(
|
||||
cache_key, rate_multiplier, ttl_seconds=ProviderCacheService.CACHE_TTL
|
||||
)
|
||||
logger.debug(f"ProviderAPIKey rate_multiplier 已缓存: {provider_api_key_id[:8]}... format={format_suffix} value={rate_multiplier}")
|
||||
logger.debug(
|
||||
f"ProviderAPIKey rate_multiplier 已缓存: {provider_api_key_id[:8]}... format={format_suffix} value={rate_multiplier}"
|
||||
)
|
||||
return rate_multiplier
|
||||
else:
|
||||
# 缓存负结果
|
||||
@@ -133,9 +136,7 @@ class ProviderCacheService:
|
||||
await CacheService.delete(cache_key)
|
||||
|
||||
# 2. 缓存未命中,查询数据库
|
||||
provider = (
|
||||
db.query(Provider.billing_type).filter(Provider.id == provider_id).first()
|
||||
)
|
||||
provider = db.query(Provider.billing_type).filter(Provider.id == provider_id).first()
|
||||
|
||||
# 3. 写入缓存
|
||||
if provider:
|
||||
@@ -196,7 +197,9 @@ class ProviderCacheService:
|
||||
async def invalidate_provider_api_key_cache(provider_api_key_id: str) -> None:
|
||||
"""清除 ProviderAPIKey 缓存(包括所有 API 格式的缓存)"""
|
||||
# 使用模式匹配删除所有格式的缓存
|
||||
await CacheService.delete_pattern(f"provider_api_key:rate_multiplier:{provider_api_key_id}:*")
|
||||
await CacheService.delete_pattern(
|
||||
f"provider_api_key:rate_multiplier:{provider_api_key_id}:*"
|
||||
)
|
||||
logger.debug(f"ProviderAPIKey 缓存已清除: {provider_api_key_id[:8]}...")
|
||||
|
||||
@staticmethod
|
||||
|
||||
14
src/services/cache/sync.py
vendored
14
src/services/cache/sync.py
vendored
@@ -11,14 +11,12 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
from src.core.logger import logger
|
||||
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.logger import logger
|
||||
@@ -69,9 +67,11 @@ class CacheSyncService:
|
||||
self._listener_task = asyncio.create_task(self._listen())
|
||||
self._running = True
|
||||
|
||||
logger.info("[CacheSync] 缓存同步服务已启动,订阅频道: "
|
||||
logger.info(
|
||||
"[CacheSync] 缓存同步服务已启动,订阅频道: "
|
||||
f"{self.CHANNEL_GLOBAL_MODEL}, "
|
||||
f"{self.CHANNEL_MODEL}, {self.CHANNEL_CLEAR_ALL}")
|
||||
f"{self.CHANNEL_MODEL}, {self.CHANNEL_CLEAR_ALL}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"[CacheSync] 启动失败: {e}")
|
||||
raise
|
||||
@@ -167,7 +167,9 @@ class CacheSyncService:
|
||||
_cache_sync_service: CacheSyncService | None = None
|
||||
|
||||
|
||||
async def get_cache_sync_service(redis_client: aioredis.Redis | None = None) -> CacheSyncService | None:
|
||||
async def get_cache_sync_service(
|
||||
redis_client: aioredis.Redis | None = None,
|
||||
) -> CacheSyncService | None:
|
||||
"""
|
||||
获取缓存同步服务实例
|
||||
|
||||
|
||||
2
src/services/cache/user_cache.py
vendored
2
src/services/cache/user_cache.py
vendored
@@ -19,10 +19,10 @@
|
||||
await UserCacheService.invalidate_user_cache(user_id, email)
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.constants import CacheTTL
|
||||
|
||||
@@ -9,16 +9,15 @@
|
||||
5. 显式传入 (用于重试升级)
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
from src.core.api_format import get_header_value
|
||||
from src.core.key_capabilities import (
|
||||
CAPABILITY_DEFINITIONS,
|
||||
CapabilityConfigMode,
|
||||
get_user_configurable_capabilities,
|
||||
)
|
||||
from src.core.api_format import get_header_value
|
||||
from src.core.logger import logger
|
||||
|
||||
# Adapter 检测器类型:接受 headers 和可选的 request_body,返回能力需求字典
|
||||
@@ -113,9 +112,7 @@ class CapabilityResolver:
|
||||
# 只有尚未设置的能力才从 Adapter 检测
|
||||
if cap_name not in requirements:
|
||||
requirements[cap_name] = cap_value
|
||||
logger.debug(
|
||||
f"[CapabilityResolver] 从 Adapter 检测到 {cap_name}={cap_value}"
|
||||
)
|
||||
logger.debug(f"[CapabilityResolver] 从 Adapter 检测到 {cap_name}={cap_value}")
|
||||
|
||||
# 5. 显式覆盖(重试时使用)
|
||||
if explicit_requirements:
|
||||
|
||||
@@ -64,7 +64,9 @@ class EmailSenderService:
|
||||
"smtp_use_tls": SystemConfigService.get_config(db, "smtp_use_tls", default=True),
|
||||
"smtp_use_ssl": SystemConfigService.get_config(db, "smtp_use_ssl", default=False),
|
||||
"smtp_from_email": SystemConfigService.get_config(db, "smtp_from_email"),
|
||||
"smtp_from_name": SystemConfigService.get_config(db, "smtp_from_name", default="Aether"),
|
||||
"smtp_from_name": SystemConfigService.get_config(
|
||||
db, "smtp_from_name", default="Aether"
|
||||
),
|
||||
}
|
||||
return config
|
||||
|
||||
@@ -143,7 +145,11 @@ class EmailSenderService:
|
||||
|
||||
# 发送邮件
|
||||
return await EmailSenderService._send_email(
|
||||
config=config, to_email=to_email, subject=subject, html_body=html_body, text_body=text_body
|
||||
config=config,
|
||||
to_email=to_email,
|
||||
subject=subject,
|
||||
html_body=html_body,
|
||||
text_body=text_body,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -419,9 +419,7 @@ class EmailTemplate:
|
||||
return EmailTemplate.html_to_text(html)
|
||||
|
||||
@staticmethod
|
||||
def get_subject(
|
||||
template_type: str = "verification", db: Session | None = None
|
||||
) -> str:
|
||||
def get_subject(template_type: str = "verification", db: Session | None = None) -> str:
|
||||
"""
|
||||
获取邮件主题
|
||||
|
||||
|
||||
@@ -63,7 +63,7 @@ def _extract_file_name_from_uri(file_uri: str) -> str | None:
|
||||
# 完整 URL 格式
|
||||
if "/files/" in file_uri:
|
||||
idx = file_uri.rfind("/files/")
|
||||
return file_uri[idx + 1:] # 提取 files/xxx 部分
|
||||
return file_uri[idx + 1 :] # 提取 files/xxx 部分
|
||||
# 短格式
|
||||
if file_uri.startswith("files/"):
|
||||
return file_uri
|
||||
|
||||
@@ -21,7 +21,6 @@ from sqlalchemy.orm import Session
|
||||
from src.core.logger import logger
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, RequestCandidate
|
||||
|
||||
|
||||
# 缓存配置
|
||||
CACHE_TTL_SECONDS = 30 # 缓存 30 秒
|
||||
CACHE_KEY_PREFIX = "health:endpoint:"
|
||||
@@ -31,6 +30,7 @@ def _get_redis_client() -> Any:
|
||||
"""获取 Redis 客户端,失败返回 None"""
|
||||
try:
|
||||
from src.clients.redis_client import redis_client
|
||||
|
||||
return redis_client
|
||||
except Exception:
|
||||
return None
|
||||
@@ -77,15 +77,19 @@ class EndpointHealthService:
|
||||
|
||||
# 批量查询所有密钥(通过 provider_id 关联)
|
||||
all_keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.provider_id.in_(all_provider_ids))
|
||||
.all()
|
||||
) if all_provider_ids else []
|
||||
(
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.provider_id.in_(all_provider_ids))
|
||||
.all()
|
||||
)
|
||||
if all_provider_ids
|
||||
else []
|
||||
)
|
||||
|
||||
# 按 api_format 分组密钥(通过 api_formats 字段)
|
||||
keys_by_format: dict[str, list[ProviderAPIKey]] = defaultdict(list)
|
||||
for key in all_keys:
|
||||
for fmt in (key.api_formats or []):
|
||||
for fmt in key.api_formats or []:
|
||||
keys_by_format[fmt].append(key)
|
||||
|
||||
# 按 API 格式聚合
|
||||
@@ -157,11 +161,14 @@ class EndpointHealthService:
|
||||
result = []
|
||||
|
||||
for api_format, stats in format_stats.items():
|
||||
timeline_data = timeline_data_map.get(api_format, {
|
||||
"timeline": ["unknown"] * 100,
|
||||
"time_range_start": None,
|
||||
"time_range_end": None,
|
||||
})
|
||||
timeline_data = timeline_data_map.get(
|
||||
api_format,
|
||||
{
|
||||
"timeline": ["unknown"] * 100,
|
||||
"time_range_start": None,
|
||||
"time_range_end": None,
|
||||
},
|
||||
)
|
||||
timeline = timeline_data["timeline"]
|
||||
time_range_start = timeline_data.get("time_range_start")
|
||||
time_range_end = timeline_data.get("time_range_end")
|
||||
@@ -271,28 +278,22 @@ class EndpointHealthService:
|
||||
final_statuses = ["success", "failed", "skipped"]
|
||||
|
||||
segment_expr = func.floor(
|
||||
func.extract('epoch', RequestCandidate.created_at - start_time) / segment_seconds
|
||||
).label('segment_idx')
|
||||
func.extract("epoch", RequestCandidate.created_at - start_time) / segment_seconds
|
||||
).label("segment_idx")
|
||||
|
||||
candidate_stats = (
|
||||
db.query(
|
||||
RequestCandidate.key_id,
|
||||
segment_expr,
|
||||
func.count(RequestCandidate.id).label('total_count'),
|
||||
func.sum(
|
||||
case(
|
||||
(RequestCandidate.status == "success", 1),
|
||||
else_=0
|
||||
)
|
||||
).label('success_count'),
|
||||
func.sum(
|
||||
case(
|
||||
(RequestCandidate.status == "failed", 1),
|
||||
else_=0
|
||||
)
|
||||
).label('failed_count'),
|
||||
func.min(RequestCandidate.created_at).label('min_time'),
|
||||
func.max(RequestCandidate.created_at).label('max_time'),
|
||||
func.count(RequestCandidate.id).label("total_count"),
|
||||
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
|
||||
"success_count"
|
||||
),
|
||||
func.sum(case((RequestCandidate.status == "failed", 1), else_=0)).label(
|
||||
"failed_count"
|
||||
),
|
||||
func.min(RequestCandidate.created_at).label("min_time"),
|
||||
func.max(RequestCandidate.created_at).label("max_time"),
|
||||
)
|
||||
.filter(
|
||||
RequestCandidate.key_id.in_(all_key_ids),
|
||||
@@ -311,13 +312,17 @@ class EndpointHealthService:
|
||||
key_to_format[key_id] = api_format
|
||||
|
||||
# 按 api_format 和 segment 聚合数据
|
||||
format_segment_data: dict[str, dict[int, dict]] = defaultdict(lambda: defaultdict(lambda: {
|
||||
"total": 0,
|
||||
"success": 0,
|
||||
"failed": 0,
|
||||
"min_time": None,
|
||||
"max_time": None,
|
||||
}))
|
||||
format_segment_data: dict[str, dict[int, dict]] = defaultdict(
|
||||
lambda: defaultdict(
|
||||
lambda: {
|
||||
"total": 0,
|
||||
"success": 0,
|
||||
"failed": 0,
|
||||
"min_time": None,
|
||||
"max_time": None,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
for row in candidate_stats:
|
||||
key_id = row.key_id
|
||||
@@ -454,24 +459,34 @@ class EndpointHealthService:
|
||||
db, format_key_mapping, now, lookback_hours, segments
|
||||
)
|
||||
|
||||
return result.get("_single", {
|
||||
"timeline": ["unknown"] * 100,
|
||||
"time_range_start": None,
|
||||
"time_range_end": None,
|
||||
})
|
||||
return result.get(
|
||||
"_single",
|
||||
{
|
||||
"timeline": ["unknown"] * 100,
|
||||
"time_range_start": None,
|
||||
"time_range_end": None,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _format_display_name(api_format: str) -> str:
|
||||
"""格式化 API 格式的显示名称"""
|
||||
format_names = {
|
||||
"CLAUDE": "Claude API",
|
||||
"CLAUDE_CLI": "Claude CLI",
|
||||
"CLAUDE_COMPATIBLE": "Claude 兼容",
|
||||
"OPENAI": "OpenAI API",
|
||||
"OPENAI_CLI": "OpenAI CLI",
|
||||
"OPENAI_COMPATIBLE": "OpenAI 兼容",
|
||||
}
|
||||
return format_names.get(api_format, api_format)
|
||||
raw = str(api_format or "").strip()
|
||||
normalized = raw.lower()
|
||||
if ":" not in normalized:
|
||||
return raw or api_format
|
||||
|
||||
fam, kind = normalized.split(":", 1)
|
||||
fam_label = {"claude": "Claude", "openai": "OpenAI", "gemini": "Gemini"}.get(fam, fam)
|
||||
kind_label = {
|
||||
"chat": "",
|
||||
"cli": "CLI",
|
||||
"video": "Video",
|
||||
"image": "Image",
|
||||
}.get(kind, kind)
|
||||
if not kind_label:
|
||||
return fam_label
|
||||
return f"{fam_label} {kind_label}"
|
||||
|
||||
@staticmethod
|
||||
def _get_from_cache(key: str) -> list[dict[str, Any]] | None:
|
||||
|
||||
@@ -230,9 +230,9 @@ class HealthMonitor:
|
||||
|
||||
if state == CircuitState.HALF_OPEN:
|
||||
# 半开状态:记录成功
|
||||
circuit_data["half_open_successes"] = int(
|
||||
circuit_data.get("half_open_successes") or 0
|
||||
) + 1
|
||||
circuit_data["half_open_successes"] = (
|
||||
int(circuit_data.get("half_open_successes") or 0) + 1
|
||||
)
|
||||
|
||||
if circuit_data["half_open_successes"] >= cls.HALF_OPEN_SUCCESS_THRESHOLD:
|
||||
# 达到成功阈值,关闭熔断器
|
||||
@@ -357,9 +357,9 @@ class HealthMonitor:
|
||||
|
||||
if state == CircuitState.HALF_OPEN:
|
||||
# 半开状态:记录失败
|
||||
circuit_data["half_open_failures"] = int(
|
||||
circuit_data.get("half_open_failures") or 0
|
||||
) + 1
|
||||
circuit_data["half_open_failures"] = (
|
||||
int(circuit_data.get("half_open_failures") or 0) + 1
|
||||
)
|
||||
|
||||
if circuit_data["half_open_failures"] >= cls.HALF_OPEN_FAILURE_THRESHOLD:
|
||||
# 达到失败阈值,重新打开熔断器
|
||||
@@ -581,9 +581,7 @@ class HealthMonitor:
|
||||
return cls._get_status_from_circuit_data(circuit_data)
|
||||
|
||||
@classmethod
|
||||
def _get_status_from_circuit_data(
|
||||
cls, circuit_data: dict[str, Any]
|
||||
) -> tuple[bool, str | None]:
|
||||
def _get_status_from_circuit_data(cls, circuit_data: dict[str, Any]) -> tuple[bool, str | None]:
|
||||
"""从熔断器数据获取状态描述"""
|
||||
if not circuit_data.get("open"):
|
||||
return True, None
|
||||
@@ -662,9 +660,7 @@ class HealthMonitor:
|
||||
result["health_score"] = float(health_data.get("health_score") or 1.0)
|
||||
result["error_rate"] = cls._calculate_error_rate_from_window(window, now_ts)
|
||||
result["window_size"] = len(valid_window)
|
||||
result["consecutive_failures"] = int(
|
||||
health_data.get("consecutive_failures") or 0
|
||||
)
|
||||
result["consecutive_failures"] = int(health_data.get("consecutive_failures") or 0)
|
||||
result["last_failure_at"] = health_data.get("last_failure_at")
|
||||
result["circuit_breaker"] = {
|
||||
"state": cls._get_circuit_state_from_data(circuit_data, now),
|
||||
@@ -678,7 +674,7 @@ class HealthMonitor:
|
||||
else:
|
||||
# 返回所有格式的健康度数据
|
||||
formats_health = {}
|
||||
for fmt in (key.api_formats or []):
|
||||
for fmt in key.api_formats or []:
|
||||
health_data = health_by_format.get(fmt, _default_health_data())
|
||||
circuit_data = circuit_by_format.get(fmt, _default_circuit_data())
|
||||
window = health_data.get("request_results_window") or []
|
||||
@@ -688,9 +684,7 @@ class HealthMonitor:
|
||||
"health_score": float(health_data.get("health_score") or 1.0),
|
||||
"error_rate": cls._calculate_error_rate_from_window(window, now_ts),
|
||||
"window_size": len(valid_window),
|
||||
"consecutive_failures": int(
|
||||
health_data.get("consecutive_failures") or 0
|
||||
),
|
||||
"consecutive_failures": int(health_data.get("consecutive_failures") or 0),
|
||||
"last_failure_at": health_data.get("last_failure_at"),
|
||||
"circuit_breaker": {
|
||||
"state": cls._get_circuit_state_from_data(circuit_data, now),
|
||||
@@ -701,9 +695,7 @@ class HealthMonitor:
|
||||
"half_open_successes": int(
|
||||
circuit_data.get("half_open_successes") or 0
|
||||
),
|
||||
"half_open_failures": int(
|
||||
circuit_data.get("half_open_failures") or 0
|
||||
),
|
||||
"half_open_failures": int(circuit_data.get("half_open_failures") or 0),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -711,9 +703,7 @@ class HealthMonitor:
|
||||
|
||||
# 计算整体健康度(取最低值)
|
||||
if formats_health:
|
||||
result["health_score"] = min(
|
||||
h["health_score"] for h in formats_health.values()
|
||||
)
|
||||
result["health_score"] = min(h["health_score"] for h in formats_health.values())
|
||||
result["any_circuit_open"] = any(
|
||||
h["circuit_breaker"]["open"] for h in formats_health.values()
|
||||
)
|
||||
@@ -731,9 +721,7 @@ class HealthMonitor:
|
||||
def get_endpoint_health(cls, db: Session, endpoint_id: str) -> dict[str, Any] | None:
|
||||
"""获取 Endpoint 健康状态"""
|
||||
try:
|
||||
endpoint = (
|
||||
db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
|
||||
)
|
||||
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
|
||||
if not endpoint:
|
||||
return None
|
||||
|
||||
@@ -782,6 +770,12 @@ class HealthMonitor:
|
||||
db.rollback()
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def reset_open_circuit_count(cls) -> None:
|
||||
"""重置进程级别的熔断计数器(批量恢复后调用)。"""
|
||||
cls._open_circuit_keys = 0
|
||||
health_open_circuits.set(0)
|
||||
|
||||
@classmethod
|
||||
def manually_enable(cls, db: Session, key_id: str | None = None) -> bool:
|
||||
"""手动启用 Key"""
|
||||
@@ -915,18 +909,14 @@ class HealthMonitor:
|
||||
# ==================== 便捷方法 ====================
|
||||
|
||||
@classmethod
|
||||
def get_health_score(
|
||||
cls, key: ProviderAPIKey, api_format: str | None = None
|
||||
) -> float:
|
||||
def get_health_score(cls, key: ProviderAPIKey, api_format: str | None = None) -> float:
|
||||
"""获取指定格式的健康度分数"""
|
||||
if not api_format:
|
||||
# 返回所有格式中的最低健康度
|
||||
health_by_format = key.health_by_format or {}
|
||||
if not health_by_format:
|
||||
return 1.0
|
||||
return min(
|
||||
float(h.get("health_score") or 1.0) for h in health_by_format.values()
|
||||
)
|
||||
return min(float(h.get("health_score") or 1.0) for h in health_by_format.values())
|
||||
|
||||
health_data = cls._get_health_data(key, api_format)
|
||||
return float(health_data.get("health_score") or 1.0)
|
||||
|
||||
@@ -2,9 +2,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import ipaddress
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -159,9 +159,7 @@ class ManagementTokenService:
|
||||
ValueError: 如果名称已存在或超过数量限制
|
||||
"""
|
||||
# 检查用户 Token 数量限制
|
||||
token_count = (
|
||||
db.query(ManagementToken).filter(ManagementToken.user_id == user_id).count()
|
||||
)
|
||||
token_count = db.query(ManagementToken).filter(ManagementToken.user_id == user_id).count()
|
||||
max_tokens = config.management_token_max_per_user
|
||||
if token_count >= max_tokens:
|
||||
raise ValueError(f"已达到 Token 数量上限({max_tokens})")
|
||||
@@ -335,9 +333,7 @@ class ManagementTokenService:
|
||||
return token
|
||||
|
||||
@staticmethod
|
||||
def delete_token(
|
||||
db: Session, token_id: str, user_id: str | None = None
|
||||
) -> bool:
|
||||
def delete_token(db: Session, token_id: str, user_id: str | None = None) -> bool:
|
||||
"""删除 Token
|
||||
|
||||
Args:
|
||||
|
||||
@@ -4,10 +4,10 @@
|
||||
包含模型管理、成本计算等功能。
|
||||
"""
|
||||
|
||||
from src.services.model.availability import ModelAvailabilityQuery
|
||||
from src.services.model.cost import ModelCostService
|
||||
from src.services.model.fetch_scheduler import ModelFetchScheduler, get_model_fetch_scheduler
|
||||
from src.services.model.global_model import GlobalModelService
|
||||
from src.services.model.availability import ModelAvailabilityQuery
|
||||
from src.services.model.service import ModelService
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -8,9 +8,7 @@
|
||||
- API Key/User 的请求级访问限制由 models_service.AccessRestrictions 处理
|
||||
"""
|
||||
|
||||
|
||||
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy import or_, tuple_
|
||||
from sqlalchemy.orm import Query, Session, contains_eager
|
||||
|
||||
from src.core.logger import logger
|
||||
@@ -21,6 +19,7 @@ from src.models.database import (
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
|
||||
|
||||
class ModelAvailabilityQuery:
|
||||
@@ -86,23 +85,45 @@ class ModelAvailabilityQuery:
|
||||
Returns:
|
||||
{provider_id: {format1, format2, ...}}
|
||||
"""
|
||||
target_formats = {f.upper() for f in api_formats}
|
||||
target_pairs: list[tuple[str, str]] = []
|
||||
for fmt in api_formats:
|
||||
if not fmt:
|
||||
continue
|
||||
try:
|
||||
norm = normalize_endpoint_signature(fmt)
|
||||
fam, kind = norm.split(":", 1)
|
||||
if fam and kind:
|
||||
target_pairs.append((fam, kind))
|
||||
except Exception:
|
||||
continue
|
||||
if not target_pairs:
|
||||
return {}
|
||||
|
||||
endpoint_rows = (
|
||||
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
|
||||
db.query(
|
||||
ProviderEndpoint.provider_id,
|
||||
ProviderEndpoint.api_family,
|
||||
ProviderEndpoint.endpoint_kind,
|
||||
)
|
||||
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||
.filter(
|
||||
Provider.is_active.is_(True),
|
||||
ProviderEndpoint.api_format.in_(list(target_formats)),
|
||||
ProviderEndpoint.is_active.is_(True),
|
||||
ProviderEndpoint.api_family.isnot(None),
|
||||
ProviderEndpoint.endpoint_kind.isnot(None),
|
||||
tuple_(ProviderEndpoint.api_family, ProviderEndpoint.endpoint_kind).in_(
|
||||
target_pairs
|
||||
),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
provider_to_formats: dict[str, set[str]] = {}
|
||||
for provider_id, fmt in endpoint_rows:
|
||||
if provider_id and fmt:
|
||||
provider_to_formats.setdefault(provider_id, set()).add(str(fmt).upper())
|
||||
for provider_id, fam, kind in endpoint_rows:
|
||||
if provider_id and fam and kind:
|
||||
provider_to_formats.setdefault(provider_id, set()).add(
|
||||
normalize_endpoint_signature(f"{fam}:{kind}")
|
||||
)
|
||||
|
||||
return provider_to_formats
|
||||
|
||||
@@ -123,7 +144,7 @@ class ModelAvailabilityQuery:
|
||||
if not provider_ids:
|
||||
return set()
|
||||
|
||||
target_formats = {f.upper() for f in api_formats}
|
||||
target_formats = {normalize_endpoint_signature(f) for f in api_formats if f}
|
||||
|
||||
key_rows = (
|
||||
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
|
||||
@@ -144,17 +165,24 @@ class ModelAvailabilityQuery:
|
||||
continue
|
||||
|
||||
# 类型兜底:key_formats 是 JSON 字段
|
||||
if not isinstance(key_formats, list):
|
||||
if key_formats is not None:
|
||||
logger.warning(
|
||||
f"[ModelAvailability] Key api_formats 类型异常, "
|
||||
f"provider_id={provider_id}, type={type(key_formats).__name__}"
|
||||
)
|
||||
if key_formats is None:
|
||||
# None = 全支持(兼容历史数据)
|
||||
key_formats_norm = set(endpoint_formats)
|
||||
elif not isinstance(key_formats, list):
|
||||
logger.warning(
|
||||
"[ModelAvailability] Key api_formats 类型异常, provider_id=%s, type=%s",
|
||||
provider_id,
|
||||
type(key_formats).__name__,
|
||||
)
|
||||
continue
|
||||
else:
|
||||
key_formats_norm = {
|
||||
normalize_endpoint_signature(str(f))
|
||||
for f in key_formats
|
||||
if isinstance(f, str) and f
|
||||
}
|
||||
|
||||
key_formats_upper: set[str] = {str(f).upper() for f in key_formats if isinstance(f, str)}
|
||||
|
||||
if key_formats_upper & endpoint_formats & target_formats:
|
||||
if key_formats_norm & endpoint_formats & target_formats:
|
||||
available_provider_ids.add(provider_id)
|
||||
|
||||
return available_provider_ids
|
||||
@@ -179,7 +207,7 @@ class ModelAvailabilityQuery:
|
||||
if not provider_ids:
|
||||
return {}
|
||||
|
||||
target_formats = {f.upper() for f in api_formats}
|
||||
target_formats = {normalize_endpoint_signature(f) for f in api_formats if f}
|
||||
|
||||
key_rows = (
|
||||
db.query(
|
||||
@@ -205,16 +233,23 @@ class ModelAvailabilityQuery:
|
||||
continue
|
||||
|
||||
# 类型兜底:key_formats
|
||||
if not isinstance(key_formats, list):
|
||||
if key_formats is not None:
|
||||
logger.warning(
|
||||
f"[ModelAvailability] Key api_formats 类型异常, "
|
||||
f"key_id={key_id}, type={type(key_formats).__name__}"
|
||||
)
|
||||
if key_formats is None:
|
||||
key_formats_norm = set(endpoint_formats)
|
||||
elif not isinstance(key_formats, list):
|
||||
logger.warning(
|
||||
"[ModelAvailability] Key api_formats 类型异常, key_id=%s, type=%s",
|
||||
key_id,
|
||||
type(key_formats).__name__,
|
||||
)
|
||||
continue
|
||||
else:
|
||||
key_formats_norm = {
|
||||
normalize_endpoint_signature(str(f))
|
||||
for f in key_formats
|
||||
if isinstance(f, str) and f
|
||||
}
|
||||
|
||||
key_formats_upper: set[str] = {str(f).upper() for f in key_formats if isinstance(f, str)}
|
||||
usable_formats = key_formats_upper & endpoint_formats & target_formats
|
||||
usable_formats = key_formats_norm & endpoint_formats & target_formats
|
||||
if not usable_formats:
|
||||
continue
|
||||
|
||||
@@ -259,4 +294,3 @@ class ModelAvailabilityQuery:
|
||||
query = query.filter(Model.provider_id.in_(provider_ids))
|
||||
|
||||
return query
|
||||
|
||||
|
||||
@@ -17,13 +17,13 @@ from sqlalchemy.orm import Session
|
||||
from src.core.logger import logger
|
||||
from src.models.database import GlobalModel, Model, Provider
|
||||
|
||||
|
||||
ProviderRef = str | Provider | None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TieredPriceResult:
|
||||
"""阶梯计费价格查询结果"""
|
||||
|
||||
input_price_per_1m: float
|
||||
output_price_per_1m: float
|
||||
cache_creation_price_per_1m: float | None = None
|
||||
@@ -34,6 +34,7 @@ class TieredPriceResult:
|
||||
@dataclass
|
||||
class CostBreakdown:
|
||||
"""成本明细"""
|
||||
|
||||
input_cost: float
|
||||
output_cost: float
|
||||
cache_creation_cost: float
|
||||
@@ -58,10 +59,7 @@ class ModelCostService:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def get_tier_for_tokens(
|
||||
tiered_pricing: dict,
|
||||
total_input_tokens: int
|
||||
) -> dict | None:
|
||||
def get_tier_for_tokens(tiered_pricing: dict, total_input_tokens: int) -> dict | None:
|
||||
"""
|
||||
根据总输入 token 数确定价格阶梯。
|
||||
|
||||
@@ -89,8 +87,7 @@ class ModelCostService:
|
||||
|
||||
@staticmethod
|
||||
def get_cache_read_price_for_ttl(
|
||||
tier: dict,
|
||||
cache_ttl_minutes: int | None = None
|
||||
tier: dict, cache_ttl_minutes: int | None = None
|
||||
) -> float | None:
|
||||
"""
|
||||
根据缓存 TTL 获取缓存读取价格。
|
||||
@@ -121,9 +118,7 @@ class ModelCostService:
|
||||
# 使用默认的缓存读取价格
|
||||
return tier.get("cache_read_price_per_1m")
|
||||
|
||||
async def get_tiered_pricing_async(
|
||||
self, provider: ProviderRef, model: str
|
||||
) -> dict | None:
|
||||
async def get_tiered_pricing_async(self, provider: ProviderRef, model: str) -> dict | None:
|
||||
"""
|
||||
异步获取模型的阶梯计费配置。
|
||||
|
||||
@@ -186,21 +181,18 @@ class ModelCostService:
|
||||
if model_obj:
|
||||
# 判断定价来源
|
||||
if model_obj.tiered_pricing is not None:
|
||||
result = {
|
||||
"pricing": model_obj.tiered_pricing,
|
||||
"source": "provider"
|
||||
}
|
||||
result = {"pricing": model_obj.tiered_pricing, "source": "provider"}
|
||||
elif global_model.default_tiered_pricing is not None:
|
||||
result = {
|
||||
"pricing": global_model.default_tiered_pricing,
|
||||
"source": "global"
|
||||
"source": "global",
|
||||
}
|
||||
else:
|
||||
# Provider 没有实现该模型,直接使用 GlobalModel 的默认阶梯配置
|
||||
if global_model.default_tiered_pricing is not None:
|
||||
result = {
|
||||
"pricing": global_model.default_tiered_pricing,
|
||||
"source": "global"
|
||||
"source": "global",
|
||||
}
|
||||
|
||||
self._tiered_pricing_cache[cache_key] = result
|
||||
@@ -250,10 +242,16 @@ class ModelCostService:
|
||||
if model_obj.tiered_pricing is not None:
|
||||
result = {"pricing": model_obj.tiered_pricing, "source": "provider"}
|
||||
elif global_model.default_tiered_pricing is not None:
|
||||
result = {"pricing": global_model.default_tiered_pricing, "source": "global"}
|
||||
result = {
|
||||
"pricing": global_model.default_tiered_pricing,
|
||||
"source": "global",
|
||||
}
|
||||
else:
|
||||
if global_model.default_tiered_pricing is not None:
|
||||
result = {"pricing": global_model.default_tiered_pricing, "source": "global"}
|
||||
result = {
|
||||
"pricing": global_model.default_tiered_pricing,
|
||||
"source": "global",
|
||||
}
|
||||
|
||||
self._tiered_pricing_cache[cache_key] = result
|
||||
return result.get("pricing") if result else None
|
||||
@@ -322,8 +320,10 @@ class ModelCostService:
|
||||
else:
|
||||
input_price = model_obj.get_effective_input_price()
|
||||
output_price = model_obj.get_effective_output_price()
|
||||
logger.debug(f"找到模型价格配置: {provider_name}/{model} "
|
||||
f"(输入: ${input_price}/M, 输出: ${output_price}/M)")
|
||||
logger.debug(
|
||||
f"找到模型价格配置: {provider_name}/{model} "
|
||||
f"(输入: ${input_price}/M, 输出: ${output_price}/M)"
|
||||
)
|
||||
else:
|
||||
# Provider 没有实现该模型,直接使用 GlobalModel 的默认价格
|
||||
tiered = global_model.default_tiered_pricing
|
||||
@@ -334,8 +334,10 @@ class ModelCostService:
|
||||
else:
|
||||
input_price = 0.0
|
||||
output_price = 0.0
|
||||
logger.debug(f"使用 GlobalModel 默认价格: {provider_name}/{model} "
|
||||
f"(输入: ${input_price}/M, 输出: ${output_price}/M)")
|
||||
logger.debug(
|
||||
f"使用 GlobalModel 默认价格: {provider_name}/{model} "
|
||||
f"(输入: ${input_price}/M, 输出: ${output_price}/M)"
|
||||
)
|
||||
|
||||
# 如果没有找到价格配置,使用 0.0 并记录警告
|
||||
if input_price is None:
|
||||
@@ -348,7 +350,9 @@ class ModelCostService:
|
||||
# 异步检查按次计费价格
|
||||
price_per_request = await self.get_request_price_async(provider, model)
|
||||
if price_per_request is None or price_per_request == 0.0:
|
||||
logger.warning(f"未找到模型价格配置: {provider_name}/{model},请在 GlobalModel 中配置价格")
|
||||
logger.warning(
|
||||
f"未找到模型价格配置: {provider_name}/{model},请在 GlobalModel 中配置价格"
|
||||
)
|
||||
|
||||
self._price_cache[cache_key] = {"input": input_price, "output": output_price}
|
||||
return input_price, output_price
|
||||
|
||||
@@ -110,9 +110,7 @@ def _get_upstream_models_cache_key(provider_id: str, api_key_id: str) -> str:
|
||||
return f"upstream_models:{provider_id}:{api_key_id}"
|
||||
|
||||
|
||||
async def get_upstream_models_from_cache(
|
||||
provider_id: str, api_key_id: str
|
||||
) -> list[dict] | None:
|
||||
async def get_upstream_models_from_cache(provider_id: str, api_key_id: str) -> list[dict] | None:
|
||||
"""从缓存获取上游模型列表"""
|
||||
cache_key = _get_upstream_models_cache_key(provider_id, api_key_id)
|
||||
cached = await CacheService.get(cache_key)
|
||||
|
||||
@@ -597,9 +597,7 @@ class GlobalModelService:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.error(
|
||||
f"Failed to auto-disassociate Provider {provider.name}: {e}"
|
||||
)
|
||||
logger.error(f"Failed to auto-disassociate Provider {provider.name}: {e}")
|
||||
# 清空 success,记录整体错误
|
||||
results["success"] = []
|
||||
results["errors"].append(
|
||||
|
||||
@@ -68,18 +68,20 @@ class ModelMapperMiddleware:
|
||||
original_model = request.model
|
||||
request.model = mapping.model.select_provider_model_name()
|
||||
|
||||
logger.debug(f"Applied model mapping for provider {provider.name}: "
|
||||
f"{original_model} -> {request.model}")
|
||||
logger.debug(
|
||||
f"Applied model mapping for provider {provider.name}: "
|
||||
f"{original_model} -> {request.model}"
|
||||
)
|
||||
else:
|
||||
# 没有找到映射,使用原始模型名
|
||||
logger.debug(f"No model mapping found for {source_model} with provider {provider.name}, "
|
||||
f"forwarding with original model name")
|
||||
logger.debug(
|
||||
f"No model mapping found for {source_model} with provider {provider.name}, "
|
||||
f"forwarding with original model name"
|
||||
)
|
||||
|
||||
return request
|
||||
|
||||
async def get_mapping(
|
||||
self, source_model: str, provider_id: str
|
||||
) -> object | None:
|
||||
async def get_mapping(self, source_model: str, provider_id: str) -> object | None:
|
||||
"""
|
||||
获取模型映射
|
||||
|
||||
@@ -129,8 +131,10 @@ class ModelMapperMiddleware:
|
||||
},
|
||||
)()
|
||||
|
||||
logger.debug(f"Found model mapping: {source_model} -> {model.provider_model_name} "
|
||||
f"(provider={provider_id[:8]}...)")
|
||||
logger.debug(
|
||||
f"Found model mapping: {source_model} -> {model.provider_model_name} "
|
||||
f"(provider={provider_id[:8]}...)"
|
||||
)
|
||||
|
||||
# 缓存结果
|
||||
self._cache[cache_key] = mapping
|
||||
@@ -275,6 +279,15 @@ class ModelRoutingMiddleware:
|
||||
选中的提供商,如果没有找到返回None
|
||||
"""
|
||||
request_prefix = f"ID:{request_id} | " if request_id else ""
|
||||
allowed_norm: set[str] | None = None
|
||||
if allowed_api_formats:
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
|
||||
allowed_norm = {
|
||||
normalize_endpoint_signature(str(fmt))
|
||||
for fmt in allowed_api_formats
|
||||
if isinstance(fmt, str) and fmt
|
||||
}
|
||||
|
||||
# 1. 如果指定了提供商,直接使用
|
||||
if preferred_provider:
|
||||
@@ -286,18 +299,26 @@ class ModelRoutingMiddleware:
|
||||
|
||||
if provider:
|
||||
# 检查API格式 - 从 endpoints 中检查
|
||||
if allowed_api_formats:
|
||||
if allowed_norm:
|
||||
has_matching_endpoint = any(
|
||||
ep.is_active and ep.api_format and ep.api_format in allowed_api_formats
|
||||
ep.is_active
|
||||
and ep.api_format
|
||||
and str(ep.api_format).strip().lower() in allowed_norm
|
||||
for ep in provider.endpoints
|
||||
)
|
||||
if not has_matching_endpoint:
|
||||
logger.warning(f"Specified provider {provider.name} has no active endpoints with allowed API formats ({allowed_api_formats})")
|
||||
logger.warning(
|
||||
f"Specified provider {provider.name} has no active endpoints with allowed API formats ({allowed_api_formats})"
|
||||
)
|
||||
else:
|
||||
logger.debug(f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}")
|
||||
logger.debug(
|
||||
f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}"
|
||||
)
|
||||
return provider
|
||||
else:
|
||||
logger.debug(f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}")
|
||||
logger.debug(
|
||||
f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}"
|
||||
)
|
||||
return provider
|
||||
else:
|
||||
logger.warning(f"Specified provider {preferred_provider} not found or inactive")
|
||||
@@ -305,12 +326,12 @@ class ModelRoutingMiddleware:
|
||||
# 2. 查找优先级最高的活动提供商
|
||||
query = self.db.query(Provider).filter(Provider.is_active == True)
|
||||
|
||||
if allowed_api_formats:
|
||||
if allowed_norm:
|
||||
query = (
|
||||
query.join(ProviderEndpoint)
|
||||
.filter(
|
||||
ProviderEndpoint.is_active == True,
|
||||
ProviderEndpoint.api_format.in_(allowed_api_formats),
|
||||
ProviderEndpoint.api_format.in_(sorted(allowed_norm)),
|
||||
)
|
||||
.distinct()
|
||||
)
|
||||
@@ -318,11 +339,15 @@ class ModelRoutingMiddleware:
|
||||
best_provider = query.order_by(Provider.provider_priority.asc(), Provider.id.asc()).first()
|
||||
|
||||
if best_provider:
|
||||
logger.debug(f" └─ {request_prefix}使用优先级最高提供商: {best_provider.name} (priority:{best_provider.provider_priority}) | 模型:{model_name}")
|
||||
logger.debug(
|
||||
f" └─ {request_prefix}使用优先级最高提供商: {best_provider.name} (priority:{best_provider.provider_priority}) | 模型:{model_name}"
|
||||
)
|
||||
return best_provider
|
||||
|
||||
if allowed_api_formats:
|
||||
logger.error(f"No active providers found with allowed API formats {allowed_api_formats}.")
|
||||
logger.error(
|
||||
f"No active providers found with allowed API formats {allowed_api_formats}."
|
||||
)
|
||||
else:
|
||||
logger.error("No active providers found.")
|
||||
return None
|
||||
@@ -392,13 +417,15 @@ class ModelRoutingMiddleware:
|
||||
# 按总价格排序
|
||||
cheapest = min(
|
||||
models_with_providers,
|
||||
key=lambda x: x[1].get_effective_input_price() + x[1].get_effective_output_price()
|
||||
key=lambda x: x[1].get_effective_input_price() + x[1].get_effective_output_price(),
|
||||
)
|
||||
|
||||
provider = cheapest[0]
|
||||
model = cheapest[1]
|
||||
|
||||
logger.debug(f"Selected cheapest provider {provider.name} for model {model_name} "
|
||||
f"(input: ${model.get_effective_input_price()}/M, output: ${model.get_effective_output_price()}/M)")
|
||||
logger.debug(
|
||||
f"Selected cheapest provider {provider.name} for model {model_name} "
|
||||
f"(input: ${model.get_effective_input_price()}/M, output: ${model.get_effective_output_price()}/M)"
|
||||
)
|
||||
|
||||
return provider
|
||||
|
||||
@@ -16,6 +16,7 @@ from dataclasses import dataclass
|
||||
@dataclass
|
||||
class UsageTokens:
|
||||
"""请求的 token 使用量"""
|
||||
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_creation_input_tokens: int = 0
|
||||
@@ -25,6 +26,7 @@ class UsageTokens:
|
||||
@dataclass
|
||||
class PricingConfig:
|
||||
"""价格配置"""
|
||||
|
||||
input_price_per_1m: float = 0.0
|
||||
output_price_per_1m: float = 0.0
|
||||
cache_creation_price_per_1m: float | None = None
|
||||
@@ -37,6 +39,7 @@ class PricingConfig:
|
||||
@dataclass
|
||||
class CostResult:
|
||||
"""计费结果"""
|
||||
|
||||
input_cost: float = 0.0
|
||||
output_cost: float = 0.0
|
||||
cache_creation_cost: float = 0.0
|
||||
|
||||
@@ -10,16 +10,15 @@ from sqlalchemy import and_
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.models.api import ModelCreate, ModelResponse, ModelUpdate
|
||||
from src.models.database import Model, Provider
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
|
||||
|
||||
|
||||
class ModelService:
|
||||
"""模型管理服务"""
|
||||
|
||||
@@ -76,7 +75,9 @@ class ModelService:
|
||||
.first()
|
||||
)
|
||||
|
||||
logger.info(f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}")
|
||||
logger.info(
|
||||
f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}"
|
||||
)
|
||||
|
||||
# 清除 Redis 缓存(异步执行,不阻塞返回)
|
||||
# 重要:新增模型可能需要清除 resolver 的 NOT_FOUND 负缓存(global_model:resolve:*),
|
||||
@@ -182,12 +183,16 @@ class ModelService:
|
||||
|
||||
# 添加调试日志
|
||||
logger.debug(f"更新模型 {model_id} 收到的数据: {update_data}")
|
||||
logger.debug(f"更新前的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}")
|
||||
logger.debug(
|
||||
f"更新前的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
|
||||
)
|
||||
|
||||
for field, value in update_data.items():
|
||||
setattr(model, field, value)
|
||||
|
||||
logger.debug(f"更新后的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}")
|
||||
logger.debug(
|
||||
f"更新后的 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
|
||||
)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
@@ -228,7 +233,9 @@ class ModelService:
|
||||
# 清除 /v1/models 列表缓存
|
||||
asyncio.create_task(invalidate_models_list_cache())
|
||||
|
||||
logger.info(f"更新模型成功: id={model_id}, 最终 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}")
|
||||
logger.info(
|
||||
f"更新模型成功: id={model_id}, 最终 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
|
||||
)
|
||||
return model
|
||||
except IntegrityError as e:
|
||||
db.rollback()
|
||||
@@ -260,8 +267,10 @@ class ModelService:
|
||||
)
|
||||
|
||||
if other_implementations == 0:
|
||||
logger.warning(f"警告:删除模型 {model_id}(Provider: {model.provider_id[:8]}...)后,"
|
||||
f"GlobalModel '{model.global_model_id}' 将没有任何活跃的关联提供商")
|
||||
logger.warning(
|
||||
f"警告:删除模型 {model_id}(Provider: {model.provider_id[:8]}...)后,"
|
||||
f"GlobalModel '{model.global_model_id}' 将没有任何活跃的关联提供商"
|
||||
)
|
||||
|
||||
# 保存缓存清除所需的信息(删除后无法访问)
|
||||
cache_info = {
|
||||
@@ -290,13 +299,17 @@ class ModelService:
|
||||
# 清除内存缓存
|
||||
if cache_info["provider_id"] and cache_info["global_model_id"]:
|
||||
cache_service = get_cache_invalidation_service()
|
||||
cache_service.on_model_changed(cache_info["provider_id"], cache_info["global_model_id"])
|
||||
cache_service.on_model_changed(
|
||||
cache_info["provider_id"], cache_info["global_model_id"]
|
||||
)
|
||||
|
||||
# 清除 /v1/models 列表缓存
|
||||
asyncio.create_task(invalidate_models_list_cache())
|
||||
|
||||
logger.info(f"删除模型成功: id={model_id}, provider_model_name={cache_info['provider_model_name']}, "
|
||||
f"global_model_id={cache_info['global_model_id'][:8] if cache_info['global_model_id'] else 'None'}...")
|
||||
logger.info(
|
||||
f"删除模型成功: id={model_id}, provider_model_name={cache_info['provider_model_name']}, "
|
||||
f"global_model_id={cache_info['global_model_id'][:8] if cache_info['global_model_id'] else 'None'}..."
|
||||
)
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.error(f"删除模型失败: {str(e)}")
|
||||
|
||||
@@ -10,7 +10,7 @@ import asyncio
|
||||
|
||||
import httpx
|
||||
|
||||
from src.core.api_format import APIFormat, get_extra_headers_from_endpoint
|
||||
from src.core.api_format import get_extra_headers_from_endpoint
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ProviderEndpoint
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
@@ -18,8 +18,8 @@ from src.utils.ssl_utils import get_ssl_context
|
||||
# 并发请求限制
|
||||
MAX_CONCURRENT_REQUESTS = 5
|
||||
|
||||
# 只对这些基础格式获取模型列表,CLI 格式使用相同的上游 API
|
||||
MODEL_FETCH_FORMATS = [APIFormat.OPENAI, APIFormat.CLAUDE, APIFormat.GEMINI]
|
||||
# 只对这些基础 endpoint signature 获取模型列表,CLI 使用相同的上游 API
|
||||
MODEL_FETCH_FORMATS = ["openai:chat", "claude:chat", "gemini:chat"]
|
||||
|
||||
|
||||
def _get_adapter_for_format(api_format: str) -> type | None:
|
||||
@@ -43,7 +43,7 @@ def build_all_format_configs(
|
||||
"""
|
||||
构建所有 API 格式的端点配置
|
||||
|
||||
从所有 APIFormat 枚举值构建配置,如果该格式有专门的端点配置则使用,
|
||||
从基础 endpoint signature 列表构建配置,如果该格式有专门的端点配置则使用,
|
||||
否则使用基础端点的 base_url 尝试。
|
||||
|
||||
Args:
|
||||
@@ -59,9 +59,9 @@ def build_all_format_configs(
|
||||
# 获取任意一个端点的 base_url 作为基础(用于尝试所有格式)
|
||||
# 优先使用 OPENAI 格式的端点,因为它最通用
|
||||
base_endpoint = (
|
||||
format_to_endpoint.get("OPENAI")
|
||||
or format_to_endpoint.get("CLAUDE")
|
||||
or format_to_endpoint.get("GEMINI")
|
||||
format_to_endpoint.get("openai:chat")
|
||||
or format_to_endpoint.get("claude:chat")
|
||||
or format_to_endpoint.get("gemini:chat")
|
||||
or next(iter(format_to_endpoint.values()))
|
||||
)
|
||||
base_url = base_endpoint.base_url
|
||||
@@ -70,7 +70,7 @@ def build_all_format_configs(
|
||||
# 只对基础 API 格式获取模型,CLI 格式使用相同的上游 API
|
||||
endpoint_configs: list[dict] = []
|
||||
for fmt in MODEL_FETCH_FORMATS:
|
||||
fmt_value = fmt.value
|
||||
fmt_value = fmt
|
||||
# 如果该格式有专门的端点配置,使用其 base_url 和 headers
|
||||
if fmt_value in format_to_endpoint:
|
||||
ep = format_to_endpoint[fmt_value]
|
||||
@@ -115,9 +115,7 @@ async def fetch_models_from_endpoints(
|
||||
has_success = False
|
||||
semaphore = asyncio.Semaphore(MAX_CONCURRENT_REQUESTS)
|
||||
|
||||
async def fetch_one(
|
||||
client: httpx.AsyncClient, config: dict
|
||||
) -> tuple[list, str | None, bool]:
|
||||
async def fetch_one(client: httpx.AsyncClient, config: dict) -> tuple[list, str | None, bool]:
|
||||
base_url = config["base_url"]
|
||||
if not base_url:
|
||||
return [], None, False
|
||||
|
||||
@@ -10,14 +10,12 @@ from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
|
||||
|
||||
|
||||
class CandidateResolver:
|
||||
"""
|
||||
候选解析器 - 负责获取和排序可用的 Provider 组合
|
||||
@@ -45,7 +43,7 @@ class CandidateResolver:
|
||||
|
||||
async def fetch_candidates(
|
||||
self,
|
||||
api_format: APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey | None = None,
|
||||
@@ -78,6 +76,12 @@ class CandidateResolver:
|
||||
provider_batch_size = 20
|
||||
global_model_id: str | None = None
|
||||
|
||||
logger.debug(
|
||||
"[CandidateResolver] fetch_candidates starting: model=%s, api_format=%s",
|
||||
model_name,
|
||||
api_format,
|
||||
)
|
||||
|
||||
while True:
|
||||
candidates, resolved_global_model_id = await self.cache_scheduler.list_all_candidates(
|
||||
db=self.db,
|
||||
@@ -91,6 +95,12 @@ class CandidateResolver:
|
||||
capability_requirements=capability_requirements,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"[CandidateResolver] list_all_candidates batch: offset=%d, returned=%d candidates",
|
||||
provider_offset,
|
||||
len(candidates),
|
||||
)
|
||||
|
||||
if resolved_global_model_id and global_model_id is None:
|
||||
global_model_id = resolved_global_model_id
|
||||
|
||||
@@ -100,6 +110,11 @@ class CandidateResolver:
|
||||
all_candidates.extend(candidates)
|
||||
provider_offset += provider_batch_size
|
||||
|
||||
logger.debug(
|
||||
"[CandidateResolver] fetch_candidates completed: total=%d candidates",
|
||||
len(all_candidates),
|
||||
)
|
||||
|
||||
if not all_candidates:
|
||||
logger.error(f" [{request_id}] 没有找到任何可用的 Provider/Endpoint/Key 组合")
|
||||
request_type = "流式" if is_stream else "非流式"
|
||||
@@ -196,7 +211,9 @@ class CandidateResolver:
|
||||
candidate_record_map[(candidate_index, 0)] = record_id
|
||||
else:
|
||||
# max_retries 已从 Endpoint 迁移到 Provider(Endpoint 仍可能保留旧字段用于兼容)
|
||||
max_retries_for_candidate = int(provider.max_retries or 2) if candidate.is_cached else 1
|
||||
max_retries_for_candidate = (
|
||||
int(provider.max_retries or 2) if candidate.is_cached else 1
|
||||
)
|
||||
|
||||
for retry_index in range(max_retries_for_candidate):
|
||||
record_id = str(uuid.uuid4())
|
||||
@@ -226,7 +243,9 @@ class CandidateResolver:
|
||||
)
|
||||
self.db.flush()
|
||||
|
||||
logger.debug(f" [{request_id}] 批量插入完成: {len(candidate_records_to_insert)} 条记录")
|
||||
logger.debug(
|
||||
f" [{request_id}] 批量插入完成: {len(candidate_records_to_insert)} 条记录"
|
||||
)
|
||||
|
||||
return candidate_record_map
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ from typing import Any
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.signature import make_signature_key
|
||||
from src.core.exceptions import (
|
||||
ConcurrencyLimitError,
|
||||
ProviderAuthException,
|
||||
@@ -28,7 +28,7 @@ from src.core.logger import logger
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.provider.format import normalize_api_format
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type
|
||||
|
||||
@@ -561,7 +561,7 @@ class ErrorClassifier:
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
affinity_key: str,
|
||||
api_format: str | APIFormat,
|
||||
api_format: str,
|
||||
global_model_id: str,
|
||||
request_id: str | None,
|
||||
captured_key_concurrent: int | None,
|
||||
@@ -618,18 +618,11 @@ class ErrorClassifier:
|
||||
extra_data["error_response"] = error_response_text
|
||||
|
||||
# client_format:用于缓存亲和性/缓存失效(用户视角)
|
||||
client_format_str = (
|
||||
normalize_api_format(api_format).value
|
||||
if isinstance(api_format, (str, APIFormat))
|
||||
else str(api_format)
|
||||
)
|
||||
client_format_str = normalize_endpoint_signature(api_format)
|
||||
# provider_format:用于健康度/熔断 bucket(Provider 真实端点格式)
|
||||
provider_api_format = getattr(endpoint, "api_format", None)
|
||||
provider_format_str = (
|
||||
provider_api_format.value
|
||||
if isinstance(provider_api_format, APIFormat)
|
||||
else str(provider_api_format or client_format_str)
|
||||
).upper()
|
||||
fam = str(getattr(endpoint, "api_family", "")).strip().lower()
|
||||
kind = str(getattr(endpoint, "endpoint_kind", "")).strip().lower()
|
||||
provider_format_str = make_signature_key(fam, kind) if fam and kind else client_format_str
|
||||
|
||||
# 处理客户端请求错误(不应重试,不失效缓存,不记录健康失败)
|
||||
if isinstance(converted_error, UpstreamClientException):
|
||||
@@ -704,7 +697,7 @@ class ErrorClassifier:
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
affinity_key: str,
|
||||
api_format: str | APIFormat,
|
||||
api_format: str,
|
||||
global_model_id: str,
|
||||
captured_key_concurrent: int | None,
|
||||
elapsed_ms: int | None,
|
||||
@@ -737,18 +730,11 @@ class ErrorClassifier:
|
||||
)
|
||||
|
||||
# client_format:用于缓存亲和性/缓存失效(用户视角)
|
||||
client_format_str = (
|
||||
normalize_api_format(api_format).value
|
||||
if isinstance(api_format, (str, APIFormat))
|
||||
else str(api_format)
|
||||
)
|
||||
client_format_str = normalize_endpoint_signature(api_format)
|
||||
# provider_format:用于健康度/熔断 bucket(Provider 真实端点格式)
|
||||
provider_api_format = getattr(endpoint, "api_format", None)
|
||||
provider_format_str = (
|
||||
provider_api_format.value
|
||||
if isinstance(provider_api_format, APIFormat)
|
||||
else str(provider_api_format or client_format_str)
|
||||
).upper()
|
||||
fam = str(getattr(endpoint, "api_family", "")).strip().lower()
|
||||
kind = str(getattr(endpoint, "endpoint_kind", "")).strip().lower()
|
||||
provider_format_str = make_signature_key(fam, kind) if fam and kind else client_format_str
|
||||
|
||||
# 处理限流错误
|
||||
if isinstance(error, ProviderRateLimitException) and key:
|
||||
|
||||
@@ -21,19 +21,16 @@
|
||||
- 本类作为协调者,组合使用上述组件
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, NoReturn
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any, NoReturn
|
||||
|
||||
import httpx
|
||||
from redis import Redis
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
from src.core.error_utils import extract_error_message
|
||||
from src.core.exceptions import (
|
||||
@@ -51,7 +48,7 @@ from src.services.cache.aware_scheduler import (
|
||||
get_cache_aware_scheduler,
|
||||
)
|
||||
from src.services.message.thinking_rectifier import ThinkingRectifier
|
||||
from src.services.provider.format import normalize_api_format
|
||||
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.concurrency_manager import get_concurrency_manager
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
@@ -172,7 +169,7 @@ class FallbackOrchestrator:
|
||||
|
||||
async def _fetch_all_candidates(
|
||||
self,
|
||||
api_format: APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey | None = None,
|
||||
@@ -256,7 +253,7 @@ class FallbackOrchestrator:
|
||||
user_api_key: ApiKey,
|
||||
request_func: Callable[..., Any],
|
||||
request_id: str | None,
|
||||
api_format: APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
global_model_id: str,
|
||||
@@ -343,8 +340,11 @@ class FallbackOrchestrator:
|
||||
if not config.thinking_rectifier_enabled:
|
||||
logger.info(f" [{request_id}] Thinking 错误:整流器已禁用,终止重试")
|
||||
self._mark_thinking_error_failed(
|
||||
candidate_record_id, converted_error, elapsed_ms,
|
||||
captured_key_concurrent, serializable_extra_data
|
||||
candidate_record_id,
|
||||
converted_error,
|
||||
elapsed_ms,
|
||||
captured_key_concurrent,
|
||||
serializable_extra_data,
|
||||
)
|
||||
raise converted_error
|
||||
|
||||
@@ -352,8 +352,11 @@ class FallbackOrchestrator:
|
||||
if request_body_ref is None:
|
||||
logger.warning(f" [{request_id}] Thinking 错误:无法获取请求体引用,终止重试")
|
||||
self._mark_thinking_error_failed(
|
||||
candidate_record_id, converted_error, elapsed_ms,
|
||||
captured_key_concurrent, serializable_extra_data
|
||||
candidate_record_id,
|
||||
converted_error,
|
||||
elapsed_ms,
|
||||
captured_key_concurrent,
|
||||
serializable_extra_data,
|
||||
)
|
||||
raise converted_error
|
||||
|
||||
@@ -361,8 +364,11 @@ class FallbackOrchestrator:
|
||||
if request_body_ref.get("_rectified", False):
|
||||
logger.warning(f" [{request_id}] Thinking 错误:已整流仍失败,终止重试")
|
||||
self._mark_thinking_error_failed(
|
||||
candidate_record_id, converted_error, elapsed_ms,
|
||||
captured_key_concurrent, {**serializable_extra_data, "rectified": True}
|
||||
candidate_record_id,
|
||||
converted_error,
|
||||
elapsed_ms,
|
||||
captured_key_concurrent,
|
||||
{**serializable_extra_data, "rectified": True},
|
||||
)
|
||||
raise converted_error
|
||||
|
||||
@@ -383,8 +389,11 @@ class FallbackOrchestrator:
|
||||
# 标记当前尝试为失败(整流前的状态)
|
||||
# 注意:整流后重试会复用此记录 ID,成功时会覆盖为 success 状态
|
||||
self._mark_thinking_error_failed(
|
||||
candidate_record_id, converted_error, elapsed_ms,
|
||||
captured_key_concurrent, {**serializable_extra_data, "rectified": True}
|
||||
candidate_record_id,
|
||||
converted_error,
|
||||
elapsed_ms,
|
||||
captured_key_concurrent,
|
||||
{**serializable_extra_data, "rectified": True},
|
||||
)
|
||||
|
||||
# 返回 continue:在当前候选的重试循环中继续,使用整流后的请求体重试
|
||||
@@ -392,8 +401,11 @@ class FallbackOrchestrator:
|
||||
else:
|
||||
logger.warning(f" [{request_id}] Thinking 错误:无可整流内容")
|
||||
self._mark_thinking_error_failed(
|
||||
candidate_record_id, converted_error, elapsed_ms,
|
||||
captured_key_concurrent, serializable_extra_data
|
||||
candidate_record_id,
|
||||
converted_error,
|
||||
elapsed_ms,
|
||||
captured_key_concurrent,
|
||||
serializable_extra_data,
|
||||
)
|
||||
raise converted_error
|
||||
|
||||
@@ -425,7 +437,7 @@ class FallbackOrchestrator:
|
||||
retry_index: int,
|
||||
max_retries_for_candidate: int,
|
||||
affinity_key: str,
|
||||
api_format: APIFormat,
|
||||
api_format: str,
|
||||
global_model_id: str,
|
||||
request_id: str | None,
|
||||
attempt: int,
|
||||
@@ -506,9 +518,7 @@ class FallbackOrchestrator:
|
||||
"provider_id": str(provider.id),
|
||||
"provider_endpoint_id": str(endpoint.id),
|
||||
"provider_api_key_id": str(key.id),
|
||||
"api_format": (
|
||||
api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
),
|
||||
"api_format": str(api_format),
|
||||
}
|
||||
raise client_error
|
||||
else:
|
||||
@@ -588,9 +598,7 @@ class FallbackOrchestrator:
|
||||
"provider_id": str(provider.id),
|
||||
"provider_endpoint_id": str(endpoint.id),
|
||||
"provider_api_key_id": str(key.id),
|
||||
"api_format": (
|
||||
api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
),
|
||||
"api_format": str(api_format),
|
||||
}
|
||||
raise converted_error
|
||||
|
||||
@@ -662,7 +670,7 @@ class FallbackOrchestrator:
|
||||
user_api_key: ApiKey,
|
||||
model_name: str,
|
||||
is_stream: bool,
|
||||
api_format_enum: APIFormat,
|
||||
api_format: str,
|
||||
) -> None:
|
||||
"""创建 pending 状态的使用记录(用于实时状态追踪)"""
|
||||
if not request_id:
|
||||
@@ -681,7 +689,7 @@ class FallbackOrchestrator:
|
||||
api_key=user_api_key,
|
||||
model=model_name,
|
||||
is_stream=is_stream,
|
||||
api_format=api_format_enum.value,
|
||||
api_format=api_format,
|
||||
)
|
||||
except Exception as e:
|
||||
# 创建 pending 记录失败不应阻塞请求
|
||||
@@ -694,7 +702,7 @@ class FallbackOrchestrator:
|
||||
user_api_key: ApiKey,
|
||||
request_func: Callable[..., Any],
|
||||
request_id: str | None,
|
||||
api_format_enum: APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
global_model_id: str,
|
||||
@@ -724,7 +732,7 @@ class FallbackOrchestrator:
|
||||
user_api_key=user_api_key,
|
||||
request_func=request_func,
|
||||
request_id=request_id,
|
||||
api_format_enum=api_format_enum,
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
global_model_id=global_model_id,
|
||||
@@ -735,9 +743,9 @@ class FallbackOrchestrator:
|
||||
)
|
||||
|
||||
if result["success"]:
|
||||
response: tuple[
|
||||
Any, str, str | None, str | None, str | None, str | None
|
||||
] = result["response"]
|
||||
response: tuple[Any, str, str | None, str | None, str | None, str | None] = result[
|
||||
"response"
|
||||
]
|
||||
return response
|
||||
|
||||
# 更新计数器和错误信息
|
||||
@@ -746,14 +754,12 @@ class FallbackOrchestrator:
|
||||
if result.get("error"):
|
||||
last_error = result["error"]
|
||||
if result.get("should_raise") and last_error is not None:
|
||||
self._attach_metadata_to_error(
|
||||
last_error, last_candidate, model_name, api_format_enum
|
||||
)
|
||||
self._attach_metadata_to_error(last_error, last_candidate, model_name, api_format)
|
||||
raise last_error
|
||||
|
||||
# 所有组合都已尝试完毕,全部失败
|
||||
self._raise_all_failed_exception(
|
||||
request_id, max_attempts, last_candidate, model_name, api_format_enum, last_error
|
||||
request_id, max_attempts, last_candidate, model_name, api_format, last_error
|
||||
)
|
||||
|
||||
async def _try_candidate_with_retries(
|
||||
@@ -764,7 +770,7 @@ class FallbackOrchestrator:
|
||||
user_api_key: ApiKey,
|
||||
request_func: Callable[..., Any],
|
||||
request_id: str | None,
|
||||
api_format_enum: APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
global_model_id: str,
|
||||
@@ -820,7 +826,7 @@ class FallbackOrchestrator:
|
||||
user_api_key=user_api_key,
|
||||
request_func=request_func,
|
||||
request_id=request_id,
|
||||
api_format=api_format_enum,
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
global_model_id=global_model_id,
|
||||
@@ -839,7 +845,7 @@ class FallbackOrchestrator:
|
||||
retry_index=retry_index,
|
||||
max_retries_for_candidate=max_retries_for_candidate,
|
||||
affinity_key=affinity_key,
|
||||
api_format=api_format_enum,
|
||||
api_format=api_format,
|
||||
global_model_id=global_model_id,
|
||||
request_id=request_id,
|
||||
attempt=attempt_counter,
|
||||
@@ -886,7 +892,7 @@ class FallbackOrchestrator:
|
||||
error: Exception | None,
|
||||
candidate: ProviderCandidate | None,
|
||||
model_name: str,
|
||||
api_format_enum: APIFormat,
|
||||
api_format: str,
|
||||
) -> None:
|
||||
"""附加 candidate 信息到异常,以便记录 usage"""
|
||||
if not error or not candidate:
|
||||
@@ -915,7 +921,7 @@ class FallbackOrchestrator:
|
||||
provider_api_key_id=(
|
||||
getattr(existing_metadata, "provider_api_key_id", None) or str(candidate.key.id)
|
||||
),
|
||||
api_format=api_format_enum.value,
|
||||
api_format=api_format,
|
||||
)
|
||||
# 使用 setattr 避免类型检查错误
|
||||
setattr(error, "request_metadata", metadata)
|
||||
@@ -926,7 +932,7 @@ class FallbackOrchestrator:
|
||||
max_attempts: int,
|
||||
last_candidate: ProviderCandidate | None,
|
||||
model_name: str,
|
||||
api_format_enum: APIFormat,
|
||||
api_format: str,
|
||||
last_error: Exception | None = None,
|
||||
) -> NoReturn:
|
||||
"""所有组合都失败时抛出异常"""
|
||||
@@ -940,7 +946,7 @@ class FallbackOrchestrator:
|
||||
"provider_id": str(last_candidate.provider.id),
|
||||
"provider_endpoint_id": str(last_candidate.endpoint.id),
|
||||
"provider_api_key_id": str(last_candidate.key.id),
|
||||
"api_format": api_format_enum.value,
|
||||
"api_format": api_format,
|
||||
}
|
||||
|
||||
# 提取上游错误响应
|
||||
@@ -988,7 +994,7 @@ class FallbackOrchestrator:
|
||||
|
||||
async def execute_with_fallback(
|
||||
self,
|
||||
api_format: str | APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
user_api_key: ApiKey,
|
||||
request_func: Callable[[Provider, ProviderEndpoint, ProviderAPIKey], Any],
|
||||
@@ -1002,7 +1008,7 @@ class FallbackOrchestrator:
|
||||
执行请求,并在失败时自动故障转移(缓存感知)
|
||||
|
||||
Args:
|
||||
api_format: API 格式(如 'CLAUDE', 'OPENAI')
|
||||
api_format: endpoint signature(如 'claude:chat', 'openai:cli')
|
||||
model_name: 模型名称
|
||||
user_api_key: 用户的 API Key对象
|
||||
request_func: 请求函数,接收 (provider, endpoint, key) 参数,返回响应
|
||||
@@ -1023,22 +1029,22 @@ class FallbackOrchestrator:
|
||||
# 准备执行上下文
|
||||
affinity_key = str(user_api_key.id)
|
||||
user_id = str(user_api_key.user_id)
|
||||
api_format_enum = normalize_api_format(api_format)
|
||||
api_format_norm = normalize_endpoint_signature(api_format)
|
||||
|
||||
logger.debug(
|
||||
f"[FallbackOrchestrator] execute_with_fallback 被调用: "
|
||||
f"api_format={api_format_enum.value}, model_name={model_name}, "
|
||||
f"api_format={api_format_norm}, model_name={model_name}, "
|
||||
f"request_id={request_id}, is_stream={is_stream}"
|
||||
)
|
||||
|
||||
# 创建 pending 状态的使用记录
|
||||
self._create_pending_usage_record(
|
||||
request_id, user_api_key, model_name, is_stream, api_format_enum
|
||||
request_id, user_api_key, model_name, is_stream, api_format_norm
|
||||
)
|
||||
|
||||
# 1. 收集所有候选(同时获取规范化的 global_model_id 用于缓存亲和性)
|
||||
all_candidates, global_model_id = await self._fetch_all_candidates(
|
||||
api_format=api_format_enum,
|
||||
api_format=api_format_norm,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
user_api_key=user_api_key,
|
||||
@@ -1064,7 +1070,7 @@ class FallbackOrchestrator:
|
||||
user_api_key=user_api_key,
|
||||
request_func=request_func,
|
||||
request_id=request_id,
|
||||
api_format_enum=api_format_enum,
|
||||
api_format=api_format_norm,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
global_model_id=global_model_id,
|
||||
|
||||
@@ -4,13 +4,11 @@
|
||||
负责执行单个候选请求
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||
@@ -18,7 +16,6 @@ from src.services.request.candidate import RequestCandidateService
|
||||
from src.services.request.executor import RequestExecutor
|
||||
|
||||
|
||||
|
||||
class RequestDispatcher:
|
||||
"""
|
||||
请求分发器 - 负责执行单个候选请求
|
||||
@@ -56,7 +53,7 @@ class RequestDispatcher:
|
||||
user_api_key: ApiKey,
|
||||
request_func: Callable[..., Any],
|
||||
request_id: str | None,
|
||||
api_format: APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
global_model_id: str,
|
||||
@@ -133,15 +130,12 @@ class RequestDispatcher:
|
||||
# 设置缓存亲和性
|
||||
if provider_supports_caching and self.cache_scheduler is not None:
|
||||
try:
|
||||
api_format_str = (
|
||||
api_format.value if isinstance(api_format, APIFormat) else api_format
|
||||
)
|
||||
await self.cache_scheduler.set_cache_affinity(
|
||||
affinity_key=affinity_key,
|
||||
provider_id=provider_id,
|
||||
endpoint_id=endpoint_id,
|
||||
key_id=key_id,
|
||||
api_format=api_format_str,
|
||||
api_format=api_format,
|
||||
global_model_id=global_model_id,
|
||||
ttl=provider_cache_ttl_seconds,
|
||||
)
|
||||
|
||||
@@ -4,12 +4,12 @@ Provider 服务模块
|
||||
包含 Provider 管理、格式处理、传输层等功能。
|
||||
"""
|
||||
|
||||
from src.services.provider.format import normalize_api_format
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.provider.service import ProviderService
|
||||
from src.services.provider.transport import build_provider_url
|
||||
|
||||
__all__ = [
|
||||
"ProviderService",
|
||||
"normalize_api_format",
|
||||
"normalize_endpoint_signature",
|
||||
"build_provider_url",
|
||||
]
|
||||
|
||||
@@ -1,18 +1,45 @@
|
||||
"""
|
||||
API 格式辅助函数,确保在调度/编排链路中使用统一的枚举值。
|
||||
Endpoint signature 辅助函数。
|
||||
|
||||
调度/编排链路使用 endpoint signature key:`family:kind`(如 "claude:chat", "openai:cli")。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from src.core.api_format.enums import ApiFamily, EndpointKind
|
||||
from src.core.api_format.signature import (
|
||||
EndpointSignature,
|
||||
make_signature_key,
|
||||
normalize_signature_key,
|
||||
)
|
||||
|
||||
DEFAULT_ENDPOINT_SIGNATURE: str = make_signature_key(ApiFamily.CLAUDE, EndpointKind.CHAT)
|
||||
|
||||
|
||||
from src.core.api_format import APIFormat, resolve_api_format
|
||||
|
||||
|
||||
def normalize_api_format(
|
||||
value: str | APIFormat | None, default: APIFormat = APIFormat.CLAUDE
|
||||
) -> APIFormat:
|
||||
def normalize_endpoint_signature(
|
||||
value: str | EndpointSignature | tuple[ApiFamily, EndpointKind] | None,
|
||||
*,
|
||||
default: str = DEFAULT_ENDPOINT_SIGNATURE,
|
||||
) -> str:
|
||||
"""
|
||||
将任意字符串/枚举值归一化为 APIFormat。
|
||||
未识别的值回退到默认枚举(默认 CLAUDE)。
|
||||
将任意输入归一化为 canonical signature key(`family:kind`,小写)。
|
||||
|
||||
不支持旧格式(如 "CLAUDE_CLI"),仅接受 `family:kind` 格式。
|
||||
解析失败时返回默认值。
|
||||
"""
|
||||
resolved = resolve_api_format(value)
|
||||
return resolved or default
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, EndpointSignature):
|
||||
return value.key
|
||||
if isinstance(value, tuple) and len(value) == 2:
|
||||
fam, kind = value
|
||||
if isinstance(fam, ApiFamily) and isinstance(kind, EndpointKind):
|
||||
return make_signature_key(fam, kind)
|
||||
return default
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return normalize_signature_key(value)
|
||||
except ValueError:
|
||||
# 如果解析失败,返回默认值
|
||||
return default
|
||||
return default
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
负责提供商选择、模型映射和请求处理
|
||||
"""
|
||||
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.models.database import GlobalModel, Model, Provider
|
||||
|
||||
@@ -13,8 +13,13 @@ import re
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from src.core.api_format import APIFormat, get_default_path, resolve_api_format
|
||||
from src.core.api_format import (
|
||||
EndpointKind,
|
||||
get_default_path_for_endpoint,
|
||||
make_signature_key,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint
|
||||
@@ -102,13 +107,29 @@ def build_provider_url(
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
|
||||
# 准备路径参数,添加 Gemini API 所需的 action 参数
|
||||
effective_path_params = dict(path_params) if path_params else {}
|
||||
# endpoint signature(新模式)
|
||||
raw_family = getattr(endpoint, "api_family", None)
|
||||
raw_kind = getattr(endpoint, "endpoint_kind", None)
|
||||
endpoint_sig = ""
|
||||
if isinstance(raw_family, str) and isinstance(raw_kind, str) and raw_family and raw_kind:
|
||||
endpoint_sig = make_signature_key(raw_family, raw_kind)
|
||||
else:
|
||||
# 兜底:允许 api_format 已直接存 signature key 的情况
|
||||
raw_format = getattr(endpoint, "api_format", None)
|
||||
if isinstance(raw_format, str) and ":" in raw_format:
|
||||
endpoint_sig = raw_format
|
||||
|
||||
# 为 Gemini API 格式自动添加 action 参数
|
||||
resolved_format = resolve_api_format(endpoint.api_format)
|
||||
if resolved_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||
if "action" not in effective_path_params:
|
||||
# endpoint_sig 为空时保持为空(更安全:默认路径回退到 "/",避免误判为 claude:chat)
|
||||
endpoint_sig = normalize_endpoint_signature(endpoint_sig) if endpoint_sig else ""
|
||||
|
||||
# 准备路径参数(Gemini chat/cli 需要 action)
|
||||
effective_path_params = dict(path_params) if path_params else {}
|
||||
if endpoint_sig.startswith("gemini:"):
|
||||
try:
|
||||
kind = EndpointKind(endpoint_sig.split(":", 1)[1])
|
||||
except Exception:
|
||||
kind = None
|
||||
if kind in {EndpointKind.CHAT, EndpointKind.CLI} and "action" not in effective_path_params:
|
||||
effective_path_params["action"] = (
|
||||
"streamGenerateContent" if is_stream else "generateContent"
|
||||
)
|
||||
@@ -124,7 +145,7 @@ def build_provider_url(
|
||||
pass
|
||||
else:
|
||||
# 使用 API 格式的默认路径
|
||||
path = _resolve_default_path(endpoint.api_format)
|
||||
path = _resolve_default_path(endpoint_sig)
|
||||
if effective_path_params:
|
||||
try:
|
||||
path = path.format(**effective_path_params)
|
||||
@@ -143,9 +164,9 @@ def build_provider_url(
|
||||
# 合并查询参数
|
||||
effective_query_params = dict(query_params) if query_params else {}
|
||||
|
||||
# Gemini 格式下清除可能存在的 key 参数(避免客户端传入的认证信息泄露到上游)
|
||||
# Gemini family 下清除可能存在的 key 参数(避免客户端传入的认证信息泄露到上游)
|
||||
# 上游认证始终使用 header 方式,不使用 URL 参数
|
||||
if resolved_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||
if endpoint_sig.startswith("gemini:"):
|
||||
effective_query_params.pop("key", None)
|
||||
# Gemini streamGenerateContent 官方支持 `?alt=sse` 返回 SSE(data: {...})。
|
||||
# 网关侧统一使用 SSE 输出,优先向上游请求 SSE 以减少解析分支;同时保留 JSON-array 兜底解析。
|
||||
@@ -161,16 +182,13 @@ def build_provider_url(
|
||||
return url
|
||||
|
||||
|
||||
def _resolve_default_path(api_format: str | None) -> str:
|
||||
"""
|
||||
根据 API 格式返回默认路径
|
||||
"""
|
||||
resolved = resolve_api_format(api_format)
|
||||
if resolved:
|
||||
return get_default_path(resolved)
|
||||
|
||||
logger.warning(f"Unknown api_format '{api_format}' for endpoint, fallback to '/'")
|
||||
return "/"
|
||||
def _resolve_default_path(endpoint_sig: str | None) -> str:
|
||||
"""根据 endpoint signature 返回默认路径。"""
|
||||
try:
|
||||
return get_default_path_for_endpoint(endpoint_sig or "")
|
||||
except Exception:
|
||||
logger.warning(f"Unknown endpoint signature '{endpoint_sig}' for endpoint, fallback to '/'")
|
||||
return "/"
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
@@ -179,15 +197,15 @@ def _resolve_default_path(api_format: str | None) -> str:
|
||||
|
||||
# Vertex AI 模型前缀到 API 格式的映射
|
||||
# 用于 auth_type=vertex_ai 时,根据模型名动态确定实际的请求/响应格式
|
||||
# 格式:前缀 -> APIFormat 值
|
||||
# 格式:前缀 -> endpoint signature(family:kind)
|
||||
VERTEX_AI_MODEL_FORMAT_MAPPING: dict[str, str] = {
|
||||
"claude-": "CLAUDE", # Anthropic Claude 模型
|
||||
"gemini-": "GEMINI", # Google Gemini 模型
|
||||
"imagen-": "GEMINI", # Google Imagen 模型(使用 Gemini 格式)
|
||||
"claude-": "claude:chat", # Anthropic Claude 模型
|
||||
"gemini-": "gemini:chat", # Google Gemini 模型
|
||||
"imagen-": "gemini:chat", # Google Imagen 模型(使用 Gemini chat 格式)
|
||||
}
|
||||
|
||||
# Vertex AI 默认 API 格式(当模型前缀不匹配时)
|
||||
VERTEX_AI_DEFAULT_FORMAT: str = "GEMINI"
|
||||
# Vertex AI 默认 endpoint signature(当模型前缀不匹配时)
|
||||
VERTEX_AI_DEFAULT_FORMAT: str = "gemini:chat"
|
||||
|
||||
|
||||
def get_vertex_ai_effective_format(
|
||||
@@ -220,7 +238,7 @@ def get_vertex_ai_effective_format(
|
||||
auth_config: 解密后的认证配置(可选),可包含 model_format_mapping 和 default_format
|
||||
|
||||
Returns:
|
||||
实际应使用的 API 格式(如 "CLAUDE", "GEMINI")
|
||||
实际应使用的 endpoint signature(如 "claude:chat", "gemini:chat")
|
||||
"""
|
||||
# 用户配置的模型-格式映射
|
||||
user_format_mapping: dict[str, str] = {}
|
||||
@@ -232,21 +250,39 @@ def get_vertex_ai_effective_format(
|
||||
|
||||
# 1. 用户配置:精确匹配
|
||||
if model in user_format_mapping:
|
||||
return user_format_mapping[model].upper()
|
||||
try:
|
||||
return normalize_endpoint_signature(user_format_mapping[model])
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Invalid vertex_ai model_format_mapping value for model '%s': %r",
|
||||
model,
|
||||
user_format_mapping[model],
|
||||
)
|
||||
|
||||
# 2. 用户配置:前缀匹配
|
||||
for prefix, api_format in user_format_mapping.items():
|
||||
if prefix.endswith("-") and model.startswith(prefix):
|
||||
return api_format.upper()
|
||||
try:
|
||||
return normalize_endpoint_signature(api_format)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Invalid vertex_ai model_format_mapping value for prefix '%s': %r",
|
||||
prefix,
|
||||
api_format,
|
||||
)
|
||||
break
|
||||
|
||||
# 3. 内置配置:前缀匹配
|
||||
for prefix, api_format in VERTEX_AI_MODEL_FORMAT_MAPPING.items():
|
||||
if model.startswith(prefix):
|
||||
return api_format
|
||||
return normalize_endpoint_signature(api_format)
|
||||
|
||||
# 4. 用户默认格式
|
||||
if user_default_format:
|
||||
return user_default_format.upper()
|
||||
try:
|
||||
return normalize_endpoint_signature(user_default_format)
|
||||
except Exception:
|
||||
logger.warning("Invalid vertex_ai default_format: %r", user_default_format)
|
||||
|
||||
# 5. 内置默认格式
|
||||
return VERTEX_AI_DEFAULT_FORMAT
|
||||
|
||||
@@ -163,8 +163,15 @@ class NewApiBalanceAction(BalanceAction):
|
||||
|
||||
# 检查是否是认证失败(未登录、无权限等)- Cookie 已失效
|
||||
auth_fail_indicators = [
|
||||
"未登录", "请登录", "login", "unauthorized", "无权限", "权限不足",
|
||||
"turnstile", "captcha", "验证码", # 需要人机验证
|
||||
"未登录",
|
||||
"请登录",
|
||||
"login",
|
||||
"unauthorized",
|
||||
"无权限",
|
||||
"权限不足",
|
||||
"turnstile",
|
||||
"captcha",
|
||||
"验证码", # 需要人机验证
|
||||
]
|
||||
is_auth_fail = any(ind in message.lower() for ind in auth_fail_indicators)
|
||||
if is_auth_fail:
|
||||
|
||||
@@ -2,12 +2,12 @@
|
||||
Provider 架构模块
|
||||
"""
|
||||
|
||||
from src.services.provider_ops.architectures.anyrouter import AnyrouterArchitecture
|
||||
from src.services.provider_ops.architectures.base import (
|
||||
ProviderArchitecture,
|
||||
ProviderConnector,
|
||||
VerifyResult,
|
||||
)
|
||||
from src.services.provider_ops.architectures.anyrouter import AnyrouterArchitecture
|
||||
from src.services.provider_ops.architectures.cubence import CubenceArchitecture
|
||||
from src.services.provider_ops.architectures.generic_api import GenericApiArchitecture
|
||||
from src.services.provider_ops.architectures.nekocode import NekoCodeArchitecture
|
||||
|
||||
@@ -11,7 +11,6 @@ from typing import Any
|
||||
import httpx
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
from src.services.provider_ops.actions import (
|
||||
AnyrouterBalanceAction,
|
||||
ProviderAction,
|
||||
@@ -22,6 +21,7 @@ from src.services.provider_ops.architectures.base import (
|
||||
VerifyResult,
|
||||
)
|
||||
from src.services.provider_ops.types import ConnectorAuthType, ProviderActionType
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
|
||||
# acw_sc__v2 算法常量
|
||||
_XOR_KEY = "3000176000856006061501533003690027800375"
|
||||
|
||||
@@ -3,23 +3,22 @@ Provider 架构抽象基类
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
import httpx
|
||||
|
||||
from src.services.provider_ops.actions.base import ProviderAction
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
from src.services.provider_ops.types import (
|
||||
ConnectorAuthType,
|
||||
ConnectorState,
|
||||
ConnectorStatus,
|
||||
ProviderActionType,
|
||||
)
|
||||
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
|
||||
# ==================== 连接器基类 ====================
|
||||
|
||||
@@ -488,8 +487,6 @@ class ProviderArchitecture(ABC):
|
||||
for a in self.supported_actions
|
||||
],
|
||||
"default_connector": (
|
||||
self.supported_connectors[0].auth_type.value
|
||||
if self.supported_connectors
|
||||
else None
|
||||
self.supported_connectors[0].auth_type.value if self.supported_connectors else None
|
||||
),
|
||||
}
|
||||
|
||||
@@ -80,7 +80,16 @@ class ProviderOpsService:
|
||||
"""
|
||||
|
||||
# 凭据中需要加密的字段
|
||||
SENSITIVE_FIELDS = {"api_key", "password", "session_token", "session_cookie", "token_cookie", "auth_cookie", "cookie_string", "cookie"}
|
||||
SENSITIVE_FIELDS = {
|
||||
"api_key",
|
||||
"password",
|
||||
"session_token",
|
||||
"session_cookie",
|
||||
"token_cookie",
|
||||
"auth_cookie",
|
||||
"cookie_string",
|
||||
"cookie",
|
||||
}
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
@@ -385,9 +394,7 @@ class ProviderOpsService:
|
||||
Returns:
|
||||
操作结果
|
||||
"""
|
||||
result = await self.execute_action(
|
||||
provider_id, ProviderActionType.QUERY_BALANCE, config
|
||||
)
|
||||
result = await self.execute_action(provider_id, ProviderActionType.QUERY_BALANCE, config)
|
||||
|
||||
# 成功或 auth_expired 时缓存(auth_expired 带有 cookie_expired 信息供前端显示警告)
|
||||
if result.status in (ActionStatus.SUCCESS, ActionStatus.AUTH_EXPIRED) and result.data:
|
||||
@@ -485,7 +492,9 @@ class ProviderOpsService:
|
||||
"response_time_ms": result.response_time_ms,
|
||||
}
|
||||
await CacheService.set(cache_key, cache_data, AUTH_FAILED_CACHE_TTL)
|
||||
logger.info(f"余额缓存已写入(认证失败): provider_id={provider_id}, message={result.message}")
|
||||
logger.info(
|
||||
f"余额缓存已写入(认证失败): provider_id={provider_id}, message={result.message}"
|
||||
)
|
||||
|
||||
async def _cache_balance(self, provider_id: str, result: ActionResult) -> None:
|
||||
"""缓存余额结果"""
|
||||
@@ -504,7 +513,9 @@ class ProviderOpsService:
|
||||
}
|
||||
|
||||
await CacheService.set(cache_key, cache_data, BALANCE_CACHE_TTL)
|
||||
logger.debug(f"余额缓存已写入: provider_id={provider_id}, extra={data.get('extra') if data else None}")
|
||||
logger.debug(
|
||||
f"余额缓存已写入: provider_id={provider_id}, extra={data.get('extra') if data else None}"
|
||||
)
|
||||
|
||||
async def _cache_balance_from_verify(
|
||||
self,
|
||||
@@ -633,7 +644,9 @@ class ProviderOpsService:
|
||||
if key in self.SENSITIVE_FIELDS and isinstance(value, str):
|
||||
if value: # 只加密非空值
|
||||
encrypted[key] = self.crypto.encrypt(value)
|
||||
logger.debug(f"加密字段 {key}: 原始长度={len(value)}, 加密后长度={len(encrypted[key])}")
|
||||
logger.debug(
|
||||
f"加密字段 {key}: 原始长度={len(value)}, 加密后长度={len(encrypted[key])}"
|
||||
)
|
||||
else:
|
||||
logger.warning(f"跳过空值字段 {key}")
|
||||
encrypted[key] = value
|
||||
@@ -705,8 +718,14 @@ class ProviderOpsService:
|
||||
if saved_config:
|
||||
saved_credentials = self._decrypt_credentials(saved_config.connector_credentials)
|
||||
sensitive_fields = [
|
||||
"api_key", "password", "session_token", "cookie_string", "cookie",
|
||||
"token_cookie", "auth_cookie", "session_cookie", # Cookie 认证字段
|
||||
"api_key",
|
||||
"password",
|
||||
"session_token",
|
||||
"cookie_string",
|
||||
"cookie",
|
||||
"token_cookie",
|
||||
"auth_cookie",
|
||||
"session_cookie", # Cookie 认证字段
|
||||
]
|
||||
|
||||
for field in sensitive_fields:
|
||||
@@ -738,11 +757,7 @@ class ProviderOpsService:
|
||||
if provider_ids is None:
|
||||
# 查询所有已配置的 Provider
|
||||
providers = self.db.query(Provider).filter(Provider.is_active.is_(True)).all()
|
||||
provider_ids = [
|
||||
p.id
|
||||
for p in providers
|
||||
if p.config and p.config.get("provider_ops")
|
||||
]
|
||||
provider_ids = [p.id for p in providers if p.config and p.config.get("provider_ops")]
|
||||
|
||||
if not provider_ids:
|
||||
return {}
|
||||
|
||||
@@ -4,8 +4,8 @@
|
||||
包含自适应 RPM 控制、并发管理、IP限流等功能。
|
||||
"""
|
||||
|
||||
from src.services.rate_limit.adaptive_rpm import AdaptiveConcurrencyManager # 向后兼容别名
|
||||
from src.services.rate_limit.adaptive_rpm import (
|
||||
AdaptiveConcurrencyManager, # 向后兼容别名
|
||||
AdaptiveRPMManager,
|
||||
get_adaptive_rpm_manager,
|
||||
)
|
||||
|
||||
@@ -17,7 +17,6 @@ from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
|
||||
from src.config.constants import AdaptiveReservationDefaults
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -165,9 +164,7 @@ class AdaptiveReservationManager:
|
||||
|
||||
return request_count
|
||||
|
||||
def _calculate_load_ratio(
|
||||
self, current_usage: int, effective_limit: int | None
|
||||
) -> float:
|
||||
def _calculate_load_ratio(self, current_usage: int, effective_limit: int | None) -> float:
|
||||
"""计算当前负载率"""
|
||||
if not effective_limit or effective_limit <= 0:
|
||||
return 0.0
|
||||
|
||||
@@ -110,9 +110,7 @@ class AdaptiveRPMManager:
|
||||
is_adaptive = key.rpm_limit is None
|
||||
|
||||
if not is_adaptive:
|
||||
logger.debug(
|
||||
f"Key {key.id} 设置了固定 RPM 限制 ({key.rpm_limit}),跳过自适应调整"
|
||||
)
|
||||
logger.debug(f"Key {key.id} 设置了固定 RPM 限制 ({key.rpm_limit}),跳过自适应调整")
|
||||
return int(key.rpm_limit) # type: ignore[arg-type]
|
||||
|
||||
# 更新429统计
|
||||
@@ -318,7 +316,7 @@ class AdaptiveRPMManager:
|
||||
|
||||
# 限制采样数量
|
||||
if len(samples) > self.UTILIZATION_WINDOW_SIZE:
|
||||
samples = samples[-self.UTILIZATION_WINDOW_SIZE:]
|
||||
samples = samples[-self.UTILIZATION_WINDOW_SIZE :]
|
||||
|
||||
# 更新到 key 对象
|
||||
key.utilization_samples = samples # type: ignore[assignment]
|
||||
@@ -538,7 +536,7 @@ class AdaptiveRPMManager:
|
||||
|
||||
# 保留最近N条记录
|
||||
if len(history) > self.MAX_HISTORY_RECORDS:
|
||||
history = history[-self.MAX_HISTORY_RECORDS:]
|
||||
history = history[-self.MAX_HISTORY_RECORDS :]
|
||||
|
||||
key.adjustment_history = history # type: ignore[assignment]
|
||||
|
||||
|
||||
@@ -10,12 +10,12 @@ RPM 限制管理器 - 支持 Redis 或内存的 Key 级别 RPM 限制
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import asyncio
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
|
||||
@@ -75,7 +75,9 @@ class ConcurrencyManager:
|
||||
if self._redis:
|
||||
logger.info("[OK] ConcurrencyManager 已复用全局 Redis 客户端")
|
||||
else:
|
||||
logger.warning("[WARN] Redis 不可用,RPM 限制降级为内存模式(仅在单实例环境下安全)")
|
||||
logger.warning(
|
||||
"[WARN] Redis 不可用,RPM 限制降级为内存模式(仅在单实例环境下安全)"
|
||||
)
|
||||
# 内存模式下启动后台清理任务
|
||||
self._start_background_cleanup()
|
||||
except Exception as e:
|
||||
@@ -191,14 +193,11 @@ class ConcurrencyManager:
|
||||
evict_count = max(1, self._max_memory_rpm_entries // 5)
|
||||
# 按 bucket(时间)排序,删除最旧的
|
||||
sorted_keys = sorted(
|
||||
self._memory_key_rpm_counts.items(),
|
||||
key=lambda x: x[1][0] # 按 bucket 排序
|
||||
self._memory_key_rpm_counts.items(), key=lambda x: x[1][0] # 按 bucket 排序
|
||||
)
|
||||
for k, _ in sorted_keys[:evict_count]:
|
||||
del self._memory_key_rpm_counts[k]
|
||||
logger.warning(
|
||||
f"[WARN] 内存 RPM 计数器达到上限,已淘汰 {evict_count} 个最旧条目"
|
||||
)
|
||||
logger.warning(f"[WARN] 内存 RPM 计数器达到上限,已淘汰 {evict_count} 个最旧条目")
|
||||
self._memory_key_rpm_counts[key_id] = (bucket, count)
|
||||
|
||||
def _cleanup_expired_memory_rpm_counts(self, current_bucket: int, force: bool = False) -> None:
|
||||
@@ -487,9 +486,7 @@ class ConcurrencyManager:
|
||||
from src.core.exceptions import ConcurrencyLimitError
|
||||
|
||||
user_type = "缓存用户" if is_cached_user else "新用户"
|
||||
raise ConcurrencyLimitError(
|
||||
f"RPM 限制已达上限: key={key_id}, 类型={user_type}"
|
||||
)
|
||||
raise ConcurrencyLimitError(f"RPM 限制已达上限: key={key_id}, 类型={user_type}")
|
||||
|
||||
# 记录开始时间和状态
|
||||
import time
|
||||
|
||||
@@ -147,14 +147,12 @@ class RateLimitDetector:
|
||||
and retry_after <= 30
|
||||
):
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
concurrent_reason = (
|
||||
f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
)
|
||||
|
||||
# 条件 B:无 remaining 头但 retry_after 很短(<= 5 秒)
|
||||
elif (
|
||||
requests_remaining is None
|
||||
and retry_after is not None
|
||||
and retry_after <= 5
|
||||
):
|
||||
elif requests_remaining is None and retry_after is not None and retry_after <= 5:
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
|
||||
|
||||
@@ -245,12 +243,10 @@ class RateLimitDetector:
|
||||
and retry_after <= 30
|
||||
):
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
elif (
|
||||
requests_remaining is None
|
||||
and retry_after is not None
|
||||
and retry_after <= 5
|
||||
):
|
||||
concurrent_reason = (
|
||||
f"remaining={requests_remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
)
|
||||
elif requests_remaining is None and retry_after is not None and retry_after <= 5:
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
|
||||
|
||||
@@ -330,11 +326,7 @@ class RateLimitDetector:
|
||||
):
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"remaining={remaining} > 0, retry_after={retry_after}s <= 30s"
|
||||
elif (
|
||||
remaining is None
|
||||
and retry_after is not None
|
||||
and retry_after <= 5
|
||||
):
|
||||
elif remaining is None and retry_after is not None and retry_after <= 5:
|
||||
is_likely_concurrent = True
|
||||
concurrent_reason = f"no remaining header, retry_after={retry_after}s <= 5s"
|
||||
|
||||
|
||||
@@ -12,7 +12,6 @@ from src.clients.redis_client import get_redis_client
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
|
||||
class IPRateLimiter:
|
||||
"""IP 速率限制服务"""
|
||||
|
||||
@@ -91,7 +90,9 @@ class IPRateLimiter:
|
||||
allowed = count <= rate_limit
|
||||
|
||||
if not allowed:
|
||||
logger.warning(f"IP 速率限制触发: {ip_address}, 类型: {endpoint_type}, 计数: {count}/{rate_limit}")
|
||||
logger.warning(
|
||||
f"IP 速率限制触发: {ip_address}, 类型: {endpoint_type}, 计数: {count}/{rate_limit}"
|
||||
)
|
||||
|
||||
return allowed, remaining, ttl
|
||||
|
||||
|
||||
@@ -351,7 +351,9 @@ class RequestCandidateService:
|
||||
候选自身的 TTFB(毫秒),如果计算失败则返回 global_first_byte_time_ms
|
||||
"""
|
||||
try:
|
||||
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
|
||||
candidate = (
|
||||
db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
|
||||
)
|
||||
if candidate and candidate.started_at:
|
||||
started_at = candidate.started_at
|
||||
if started_at.tzinfo is None:
|
||||
|
||||
@@ -5,18 +5,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.signature import make_signature_key
|
||||
from src.core.exceptions import ConcurrencyLimitError
|
||||
from src.core.logger import logger
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.provider.format import normalize_api_format
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
@@ -64,7 +63,7 @@ class RequestExecutor:
|
||||
user_api_key: Any,
|
||||
request_func: Callable[..., Any],
|
||||
request_id: str | None,
|
||||
api_format: str | APIFormat,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
is_stream: bool = False,
|
||||
) -> ExecutionResult:
|
||||
@@ -148,16 +147,11 @@ class RequestExecutor:
|
||||
|
||||
context.elapsed_ms = int((time.time() - context.start_time) * 1000)
|
||||
|
||||
provider_api_format = getattr(endpoint, "api_format", None)
|
||||
provider_format_str = (
|
||||
provider_api_format.value
|
||||
if isinstance(provider_api_format, APIFormat)
|
||||
else str(provider_api_format or "")
|
||||
)
|
||||
client_format_str = (
|
||||
api_format.value if isinstance(api_format, APIFormat) else str(api_format)
|
||||
)
|
||||
health_format = normalize_api_format(provider_format_str or client_format_str).value
|
||||
fam = str(getattr(endpoint, "api_family", "")).strip().lower()
|
||||
kind = str(getattr(endpoint, "endpoint_kind", "")).strip().lower()
|
||||
provider_format_str = make_signature_key(fam, kind) if fam and kind else ""
|
||||
client_format_str = normalize_endpoint_signature(api_format)
|
||||
health_format = provider_format_str or client_format_str
|
||||
|
||||
health_monitor.record_success(
|
||||
db=self.db,
|
||||
@@ -196,11 +190,7 @@ class RequestExecutor:
|
||||
extra_data={
|
||||
"is_cached_user": is_cached_user,
|
||||
"model_name": model_name,
|
||||
"api_format": (
|
||||
api_format.value
|
||||
if isinstance(api_format, APIFormat)
|
||||
else api_format
|
||||
),
|
||||
"api_format": api_format,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -14,10 +14,11 @@
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
|
||||
class RequestStatus(Enum):
|
||||
@@ -276,7 +277,9 @@ class RequestResult:
|
||||
provider_id=get_meta_value(existing_metadata, "provider_id"),
|
||||
provider_endpoint_id=get_meta_value(existing_metadata, "provider_endpoint_id"),
|
||||
provider_api_key_id=get_meta_value(existing_metadata, "provider_api_key_id"),
|
||||
provider_request_headers=get_meta_value(existing_metadata, "provider_request_headers", {}),
|
||||
provider_request_headers=get_meta_value(
|
||||
existing_metadata, "provider_request_headers", {}
|
||||
),
|
||||
provider_response_headers=get_meta_value(
|
||||
existing_metadata, "provider_response_headers", {}
|
||||
),
|
||||
|
||||
@@ -6,12 +6,12 @@
|
||||
|
||||
from src.services.system.announcement import AnnouncementService
|
||||
from src.services.system.audit import AuditService
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.system.maintenance_scheduler import CleanupScheduler # 兼容旧名称
|
||||
from src.services.system.maintenance_scheduler import (
|
||||
CleanupScheduler, # 兼容旧名称
|
||||
MaintenanceScheduler,
|
||||
get_maintenance_scheduler,
|
||||
)
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.system.scheduler import APP_TIMEZONE, TaskScheduler, get_scheduler
|
||||
from src.services.system.sync_stats import SyncStatsService
|
||||
|
||||
|
||||
@@ -14,7 +14,6 @@ from src.core.logger import logger
|
||||
from src.models.database import Announcement, AnnouncementRead, User, UserRole
|
||||
|
||||
|
||||
|
||||
class AnnouncementService:
|
||||
"""公告系统服务"""
|
||||
|
||||
|
||||
@@ -14,8 +14,6 @@ from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import AuditEventType, AuditLog
|
||||
|
||||
|
||||
|
||||
# 审计模型已移至 src/models/database.py
|
||||
|
||||
|
||||
|
||||
@@ -61,7 +61,9 @@ class CacheWarmupService:
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
if error_count > 0:
|
||||
logger.warning(f"缓存预热完成: {success_count}/3 成功, {error_count} 失败, 耗时 {elapsed:.2f}s")
|
||||
logger.warning(
|
||||
f"缓存预热完成: {success_count}/3 成功, {error_count} 失败, 耗时 {elapsed:.2f}s"
|
||||
)
|
||||
else:
|
||||
logger.info(f"缓存预热完成: {success_count}/3 成功, 耗时 {elapsed:.2f}s")
|
||||
|
||||
|
||||
@@ -227,7 +227,9 @@ class SystemConfigService:
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def set_config(db: Session, key: str, value: Any, description: str | None = None) -> SystemConfig:
|
||||
def set_config(
|
||||
db: Session, key: str, value: Any, description: str | None = None
|
||||
) -> SystemConfig:
|
||||
"""设置系统配置值"""
|
||||
config = db.query(SystemConfig).filter(SystemConfig.key == key).first()
|
||||
|
||||
|
||||
@@ -14,9 +14,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import delete
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -196,10 +196,12 @@ class MaintenanceScheduler:
|
||||
|
||||
logger.info("开始执行统计数据聚合...")
|
||||
|
||||
from src.models.database import StatsDaily, User as DBUser
|
||||
from src.services.system.scheduler import APP_TIMEZONE
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from src.models.database import StatsDaily
|
||||
from src.models.database import User as DBUser
|
||||
from src.services.system.scheduler import APP_TIMEZONE
|
||||
|
||||
# 使用业务时区计算日期,确保与定时任务触发时间一致
|
||||
# 定时任务在 Asia/Shanghai 凌晨 1 点触发,此时应聚合 Asia/Shanghai 的"昨天"
|
||||
app_tz = ZoneInfo(APP_TIMEZONE)
|
||||
@@ -227,9 +229,9 @@ class MaintenanceScheduler:
|
||||
from src.models.database import StatsDailyModel, StatsDailyProvider
|
||||
|
||||
yesterday_business_date = today_local.date() - timedelta(days=1)
|
||||
max_backfill_days: int = SystemConfigService.get_config(
|
||||
db, "max_stats_backfill_days", 30
|
||||
) or 30
|
||||
max_backfill_days: int = (
|
||||
SystemConfigService.get_config(db, "max_stats_backfill_days", 30) or 30
|
||||
)
|
||||
|
||||
# 计算回填检查的起始日期
|
||||
check_start_date = yesterday_business_date - timedelta(
|
||||
@@ -287,7 +289,9 @@ class MaintenanceScheduler:
|
||||
# 需要回填 StatsDailyProvider 的日期
|
||||
missing_provider_dates = all_dates - existing_provider_dates
|
||||
# 合并所有需要处理的日期
|
||||
dates_to_process = missing_daily_dates | missing_model_dates | missing_provider_dates
|
||||
dates_to_process = (
|
||||
missing_daily_dates | missing_model_dates | missing_provider_dates
|
||||
)
|
||||
|
||||
if dates_to_process:
|
||||
sorted_dates = sorted(dates_to_process)
|
||||
@@ -298,9 +302,7 @@ class MaintenanceScheduler:
|
||||
f"StatsDailyProvider 缺失 {len(missing_provider_dates)} 天)"
|
||||
)
|
||||
|
||||
users = (
|
||||
db.query(DBUser.id).filter(DBUser.is_active.is_(True)).all()
|
||||
)
|
||||
users = db.query(DBUser.id).filter(DBUser.is_active.is_(True)).all()
|
||||
|
||||
failed_dates = 0
|
||||
failed_users = 0
|
||||
@@ -496,15 +498,9 @@ class MaintenanceScheduler:
|
||||
|
||||
# 获取所有已配置 provider_ops 的活跃 Provider(只查询需要的字段)
|
||||
providers = (
|
||||
db.query(Provider.id, Provider.config)
|
||||
.filter(Provider.is_active.is_(True))
|
||||
.all()
|
||||
db.query(Provider.id, Provider.config).filter(Provider.is_active.is_(True)).all()
|
||||
)
|
||||
provider_ids = [
|
||||
p.id
|
||||
for p in providers
|
||||
if p.config and p.config.get("provider_ops")
|
||||
]
|
||||
provider_ids = [p.id for p in providers if p.config and p.config.get("provider_ops")]
|
||||
|
||||
if not provider_ids:
|
||||
logger.info("无已配置的 Provider,跳过签到任务")
|
||||
@@ -548,9 +544,7 @@ class MaintenanceScheduler:
|
||||
|
||||
# 统计结果
|
||||
success_count = sum(1 for _, success, _ in results if success)
|
||||
logger.info(
|
||||
f"Provider 签到完成: {success_count}/{len(provider_ids)} 成功"
|
||||
)
|
||||
logger.info(f"Provider 签到完成: {success_count}/{len(provider_ids)} 成功")
|
||||
|
||||
# 记录详细结果
|
||||
for provider_id, success, message in results:
|
||||
@@ -694,12 +688,12 @@ class MaintenanceScheduler:
|
||||
.values(
|
||||
request_body=null(),
|
||||
response_body=null(),
|
||||
request_body_compressed=compress_json(req_body)
|
||||
if req_body
|
||||
else None,
|
||||
response_body_compressed=compress_json(resp_body)
|
||||
if resp_body
|
||||
else None,
|
||||
request_body_compressed=(
|
||||
compress_json(req_body) if req_body else None
|
||||
),
|
||||
response_body_compressed=(
|
||||
compress_json(resp_body) if resp_body else None
|
||||
),
|
||||
)
|
||||
)
|
||||
if result.rowcount > 0:
|
||||
@@ -858,9 +852,7 @@ class MaintenanceScheduler:
|
||||
break
|
||||
|
||||
total_cleaned += rows_updated
|
||||
logger.debug(
|
||||
f"已清理 {rows_updated} 条记录的 header 字段,累计 {total_cleaned} 条"
|
||||
)
|
||||
logger.debug(f"已清理 {rows_updated} 条记录的 header 字段,累计 {total_cleaned} 条")
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
|
||||
@@ -43,9 +43,7 @@ def _get_business_day_range(date: datetime) -> tuple[datetime, datetime]:
|
||||
app_tz = ZoneInfo(APP_TIMEZONE)
|
||||
|
||||
# 取日期部分,构造业务时区的当天 00:00:00
|
||||
day_start_local = datetime(
|
||||
date.year, date.month, date.day, 0, 0, 0, tzinfo=app_tz
|
||||
)
|
||||
day_start_local = datetime(date.year, date.month, date.day, 0, 0, 0, tzinfo=app_tz)
|
||||
day_end_local = day_start_local + timedelta(days=1)
|
||||
|
||||
# 转换为 UTC
|
||||
@@ -91,11 +89,9 @@ class StatsAggregatorService:
|
||||
"unique_providers": 0,
|
||||
}
|
||||
|
||||
error_requests = (
|
||||
base_query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
)
|
||||
error_requests = base_query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
|
||||
aggregated = (
|
||||
db.query(
|
||||
@@ -159,13 +155,17 @@ class StatsAggregatorService:
|
||||
"error_requests": error_requests,
|
||||
"input_tokens": int(aggregated.input_tokens or 0) if aggregated else 0,
|
||||
"output_tokens": int(aggregated.output_tokens or 0) if aggregated else 0,
|
||||
"cache_creation_tokens": int(aggregated.cache_creation_tokens or 0) if aggregated else 0,
|
||||
"cache_creation_tokens": (
|
||||
int(aggregated.cache_creation_tokens or 0) if aggregated else 0
|
||||
),
|
||||
"cache_read_tokens": int(aggregated.cache_read_tokens or 0) if aggregated else 0,
|
||||
"total_cost": float(aggregated.total_cost or 0) if aggregated else 0.0,
|
||||
"actual_total_cost": float(aggregated.actual_total_cost or 0) if aggregated else 0.0,
|
||||
"input_cost": float(aggregated.input_cost or 0) if aggregated else 0.0,
|
||||
"output_cost": float(aggregated.output_cost or 0) if aggregated else 0.0,
|
||||
"cache_creation_cost": float(aggregated.cache_creation_cost or 0) if aggregated else 0.0,
|
||||
"cache_creation_cost": (
|
||||
float(aggregated.cache_creation_cost or 0) if aggregated else 0.0
|
||||
),
|
||||
"cache_read_cost": float(aggregated.cache_read_cost or 0) if aggregated else 0.0,
|
||||
"avg_response_time_ms": float(aggregated.avg_response_time or 0) if aggregated else 0.0,
|
||||
"fallback_count": fallback_count,
|
||||
@@ -219,7 +219,9 @@ class StatsAggregatorService:
|
||||
db.commit()
|
||||
|
||||
# 日志使用业务日期(输入参数),而不是 UTC 日期
|
||||
logger.info(f"[StatsAggregator] 聚合日期 {date.date()} 完成: {computed['total_requests']} 请求")
|
||||
logger.info(
|
||||
f"[StatsAggregator] 聚合日期 {date.date()} 完成: {computed['total_requests']} 请求"
|
||||
)
|
||||
return stats
|
||||
|
||||
@staticmethod
|
||||
@@ -259,16 +261,16 @@ class StatsAggregatorService:
|
||||
|
||||
existing = (
|
||||
db.query(StatsDailyModel)
|
||||
.filter(and_(StatsDailyModel.date == day_start, StatsDailyModel.model == stat.model))
|
||||
.filter(
|
||||
and_(StatsDailyModel.date == day_start, StatsDailyModel.model == stat.model)
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if existing:
|
||||
record = existing
|
||||
else:
|
||||
record = StatsDailyModel(
|
||||
id=str(uuid.uuid4()), date=day_start, model=stat.model
|
||||
)
|
||||
record = StatsDailyModel(id=str(uuid.uuid4()), date=day_start, model=stat.model)
|
||||
|
||||
record.total_requests = stat.total_requests or 0
|
||||
record.input_tokens = int(stat.input_tokens or 0)
|
||||
@@ -283,9 +285,7 @@ class StatsAggregatorService:
|
||||
results.append(record)
|
||||
|
||||
db.commit()
|
||||
logger.info(
|
||||
f"[StatsAggregator] 聚合日期 {date.date()} 模型统计完成: {len(results)} 个模型"
|
||||
)
|
||||
logger.info(f"[StatsAggregator] 聚合日期 {date.date()} 模型统计完成: {len(results)} 个模型")
|
||||
return results
|
||||
|
||||
@staticmethod
|
||||
@@ -322,7 +322,12 @@ class StatsAggregatorService:
|
||||
for stat in provider_stats:
|
||||
existing = (
|
||||
db.query(StatsDailyProvider)
|
||||
.filter(and_(StatsDailyProvider.date == day_start, StatsDailyProvider.provider_name == stat.provider_name))
|
||||
.filter(
|
||||
and_(
|
||||
StatsDailyProvider.date == day_start,
|
||||
StatsDailyProvider.provider_name == stat.provider_name,
|
||||
)
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
@@ -384,24 +389,28 @@ class StatsAggregatorService:
|
||||
date_utc = stat.date.astimezone(timezone.utc)
|
||||
date_str = date_utc.astimezone(app_tz).date().isoformat()
|
||||
|
||||
result.append({
|
||||
"date": date_str,
|
||||
"model": stat.model,
|
||||
"requests": stat.total_requests,
|
||||
"tokens": (
|
||||
stat.input_tokens + stat.output_tokens +
|
||||
stat.cache_creation_tokens + stat.cache_read_tokens
|
||||
),
|
||||
"cost": stat.total_cost,
|
||||
"avg_response_time": stat.avg_response_time_ms / 1000.0 if stat.avg_response_time_ms else 0,
|
||||
})
|
||||
result.append(
|
||||
{
|
||||
"date": date_str,
|
||||
"model": stat.model,
|
||||
"requests": stat.total_requests,
|
||||
"tokens": (
|
||||
stat.input_tokens
|
||||
+ stat.output_tokens
|
||||
+ stat.cache_creation_tokens
|
||||
+ stat.cache_read_tokens
|
||||
),
|
||||
"cost": stat.total_cost,
|
||||
"avg_response_time": (
|
||||
stat.avg_response_time_ms / 1000.0 if stat.avg_response_time_ms else 0
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def aggregate_user_daily_stats(
|
||||
db: Session, user_id: str, date: datetime
|
||||
) -> StatsUserDaily:
|
||||
def aggregate_user_daily_stats(db: Session, user_id: str, date: datetime) -> StatsUserDaily:
|
||||
"""聚合指定用户指定日期的统计数据"""
|
||||
# 将业务日期转换为 UTC 时间范围
|
||||
day_start, day_end = _get_business_day_range(date)
|
||||
@@ -443,11 +452,9 @@ class StatsAggregatorService:
|
||||
db.commit()
|
||||
return stats
|
||||
|
||||
error_requests = (
|
||||
base_query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
)
|
||||
error_requests = base_query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
|
||||
aggregated = (
|
||||
db.query(
|
||||
@@ -581,11 +588,9 @@ class StatsAggregatorService:
|
||||
"actual_total_cost": 0.0,
|
||||
}
|
||||
|
||||
error_requests = (
|
||||
base_query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
)
|
||||
error_requests = base_query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
|
||||
aggregated = (
|
||||
db.query(
|
||||
@@ -624,8 +629,7 @@ class StatsAggregatorService:
|
||||
|
||||
return {
|
||||
"total_requests": summary.all_time_requests + today_stats["total_requests"],
|
||||
"success_requests": summary.all_time_success_requests
|
||||
+ today_stats["success_requests"],
|
||||
"success_requests": summary.all_time_success_requests + today_stats["success_requests"],
|
||||
"error_requests": summary.all_time_error_requests + today_stats["error_requests"],
|
||||
"input_tokens": summary.all_time_input_tokens + today_stats["input_tokens"],
|
||||
"output_tokens": summary.all_time_output_tokens + today_stats["output_tokens"],
|
||||
|
||||
@@ -3,7 +3,6 @@ API密钥统计同步服务
|
||||
定期同步API密钥的统计数据,确保与实际使用记录一致
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import func
|
||||
@@ -13,7 +12,6 @@ from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Usage
|
||||
|
||||
|
||||
|
||||
class SyncStatsService:
|
||||
"""API密钥统计同步服务"""
|
||||
|
||||
@@ -75,12 +73,16 @@ class SyncStatsService:
|
||||
# 检查是否需要更新
|
||||
needs_update = False
|
||||
if api_key.total_requests != actual_requests:
|
||||
logger.info(f"API密钥 {api_key.id} 请求数不一致: {api_key.total_requests} -> {actual_requests}")
|
||||
logger.info(
|
||||
f"API密钥 {api_key.id} 请求数不一致: {api_key.total_requests} -> {actual_requests}"
|
||||
)
|
||||
api_key.total_requests = actual_requests
|
||||
needs_update = True
|
||||
|
||||
if abs(api_key.total_cost_usd - actual_cost) > 0.0001:
|
||||
logger.info(f"API密钥 {api_key.id} 费用不一致: {api_key.total_cost_usd} -> {actual_cost}")
|
||||
logger.info(
|
||||
f"API密钥 {api_key.id} 费用不一致: {api_key.total_cost_usd} -> {actual_cost}"
|
||||
)
|
||||
api_key.total_cost_usd = actual_cost
|
||||
needs_update = True
|
||||
|
||||
@@ -104,7 +106,9 @@ class SyncStatsService:
|
||||
|
||||
# 提交所有更改
|
||||
db.commit()
|
||||
logger.info(f"同步完成: 处理 {result['synced']} 个密钥, 更新 {result['updated']} 个, 错误 {result['errors']} 个")
|
||||
logger.info(
|
||||
f"同步完成: 处理 {result['synced']} 个密钥, 更新 {result['updated']} 个, 错误 {result['errors']} 个"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"同步统计数据时出错: {e}")
|
||||
|
||||
25
src/services/task/__init__.py
Normal file
25
src/services/task/__init__.py
Normal file
@@ -0,0 +1,25 @@
|
||||
"""
|
||||
异步任务服务层
|
||||
|
||||
提供视频/图片/音频等异步任务的:
|
||||
- 提交阶段故障转移(AsyncTaskOrchestrator)
|
||||
- 终态计费与 Usage 写入(VideoTelemetry 等)
|
||||
"""
|
||||
|
||||
from .orchestrator import (
|
||||
AllCandidatesFailedError,
|
||||
AsyncTaskOrchestrator,
|
||||
CandidateSubmissionError,
|
||||
CandidateUnsupportedError,
|
||||
SubmitOutcome,
|
||||
UpstreamClientRequestError,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AsyncTaskOrchestrator",
|
||||
"SubmitOutcome",
|
||||
"AllCandidatesFailedError",
|
||||
"UpstreamClientRequestError",
|
||||
"CandidateUnsupportedError",
|
||||
"CandidateSubmissionError",
|
||||
]
|
||||
3
src/services/task/impl/__init__.py
Normal file
3
src/services/task/impl/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
"""Task telemetry implementations for concrete task types (video/image/audio)."""
|
||||
|
||||
__all__ = []
|
||||
320
src/services/task/impl/video_telemetry.py
Normal file
320
src/services/task/impl/video_telemetry.py
Normal file
@@ -0,0 +1,320 @@
|
||||
"""
|
||||
VideoTelemetry(Phase3)
|
||||
|
||||
将 Video 异步任务的“终态计费 + Usage 写入 + required 缺失告警”从 poller 中抽离出来,
|
||||
便于未来 Image/Audio 复用相同框架。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.video_handler_base import sanitize_error_message
|
||||
from src.config.settings import config
|
||||
from src.core.api_format.conversion.internal_video import VideoStatus
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Provider, User, VideoTask
|
||||
from src.services.billing.dimension_collector_service import DimensionCollectorService
|
||||
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
|
||||
from src.services.billing.rule_service import BillingRuleService
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
class VideoTelemetry:
|
||||
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._formula_engine = FormulaEngine()
|
||||
|
||||
async def record_terminal_usage(self, task: VideoTask) -> None:
|
||||
"""
|
||||
为视频任务终态写入 Usage:
|
||||
- COMPLETED: 使用 FormulaEngine 计算 cost(或 no_rule / incomplete -> cost=0)
|
||||
- FAILED: cost=0
|
||||
|
||||
该方法可能会在 strict_mode 缺失 required 维度时将任务降级为 FAILED 并隐藏产物。
|
||||
"""
|
||||
request_id = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
request_id = task.request_metadata.get("request_id")
|
||||
request_id = request_id or task.id
|
||||
|
||||
# 计算异步任务总耗时(ms)
|
||||
response_time_ms = None
|
||||
if task.submitted_at and task.completed_at:
|
||||
delta = task.completed_at - task.submitted_at
|
||||
response_time_ms = int(delta.total_seconds() * 1000)
|
||||
|
||||
# 基础维度(无需 collectors 也可计费)
|
||||
base_dimensions: dict[str, Any] = {
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size or "",
|
||||
"retry_count": task.retry_count,
|
||||
}
|
||||
|
||||
# collectors 可用的 metadata(结构稳定,便于配置 path)
|
||||
collector_metadata: dict[str, Any] = {
|
||||
"task": {
|
||||
"id": task.id,
|
||||
"external_task_id": task.external_task_id,
|
||||
"model": task.model,
|
||||
"duration_seconds": task.duration_seconds,
|
||||
"resolution": task.resolution,
|
||||
"aspect_ratio": task.aspect_ratio,
|
||||
"size": task.size,
|
||||
"retry_count": task.retry_count,
|
||||
"video_size_bytes": task.video_size_bytes,
|
||||
},
|
||||
"result": {
|
||||
"video_url": task.video_url,
|
||||
"video_urls": task.video_urls or [],
|
||||
},
|
||||
}
|
||||
|
||||
# 维度采集:base + collectors 覆盖/补全
|
||||
dims = DimensionCollectorService(self.db).collect_dimensions(
|
||||
api_format=task.provider_api_format,
|
||||
task_type="video",
|
||||
request=task.original_request_body or {},
|
||||
response=(
|
||||
(task.request_metadata or {}).get("poll_raw_response")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
metadata=collector_metadata,
|
||||
base_dimensions=base_dimensions,
|
||||
)
|
||||
|
||||
# 取冻结的 rule_snapshot;若缺失则回退 DB 查找(兼容旧任务)
|
||||
rule_snapshot = None
|
||||
if isinstance(task.request_metadata, dict):
|
||||
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
|
||||
|
||||
billing_snapshot: dict[str, Any] = {
|
||||
"status": "complete",
|
||||
"missing_required": [],
|
||||
"strict_mode": config.billing_strict_mode,
|
||||
}
|
||||
cost = 0.0
|
||||
|
||||
if task.status == VideoStatus.FAILED.value:
|
||||
billing_snapshot["billed_reason"] = "task_failed"
|
||||
else:
|
||||
# COMPLETED:计算成本
|
||||
expression = None
|
||||
variables = None
|
||||
dimension_mappings = None
|
||||
rule_id = None
|
||||
rule_name = None
|
||||
rule_scope = None
|
||||
|
||||
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
|
||||
rule_id = rule_snapshot.get("rule_id")
|
||||
rule_name = rule_snapshot.get("rule_name")
|
||||
rule_scope = rule_snapshot.get("scope")
|
||||
expression = rule_snapshot.get("expression")
|
||||
variables = rule_snapshot.get("variables")
|
||||
dimension_mappings = rule_snapshot.get("dimension_mappings")
|
||||
else:
|
||||
lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=task.provider_id,
|
||||
model_name=task.model,
|
||||
task_type="video",
|
||||
)
|
||||
if lookup:
|
||||
rule = lookup.rule
|
||||
rule_id = rule.id
|
||||
rule_name = rule.name
|
||||
rule_scope = lookup.scope
|
||||
expression = rule.expression
|
||||
variables = rule.variables
|
||||
dimension_mappings = rule.dimension_mappings
|
||||
|
||||
if not expression:
|
||||
billing_snapshot["status"] = "no_rule"
|
||||
billing_snapshot["cost_breakdown"] = {"total": 0.0}
|
||||
logger.warning(
|
||||
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
task.provider_id,
|
||||
)
|
||||
else:
|
||||
billing_snapshot.update(
|
||||
{
|
||||
"rule_id": rule_id,
|
||||
"rule_name": rule_name,
|
||||
"rule_scope": rule_scope,
|
||||
"expression": expression,
|
||||
"variables": variables or {},
|
||||
}
|
||||
)
|
||||
try:
|
||||
result = self._formula_engine.evaluate(
|
||||
expression=expression,
|
||||
variables=variables or {},
|
||||
dimensions=dims,
|
||||
dimension_mappings=dimension_mappings or {},
|
||||
strict_mode=config.billing_strict_mode,
|
||||
)
|
||||
billing_snapshot["status"] = result.status
|
||||
billing_snapshot["missing_required"] = result.missing_required
|
||||
billing_snapshot["resolved_values"] = result.resolved_values
|
||||
if result.status == "complete":
|
||||
cost = result.cost
|
||||
else:
|
||||
logger.error(
|
||||
"Billing incomplete due to missing required dimensions "
|
||||
"(request_id=%s, model=%s, missing=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
result.missing_required,
|
||||
)
|
||||
cost = 0.0
|
||||
await self._maybe_alert_missing_required(
|
||||
model=task.model,
|
||||
missing_required=result.missing_required,
|
||||
)
|
||||
if result.error:
|
||||
billing_snapshot["error"] = result.error
|
||||
except BillingIncompleteError as exc:
|
||||
logger.error(
|
||||
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
|
||||
request_id,
|
||||
task.model,
|
||||
exc.missing_required,
|
||||
)
|
||||
billing_snapshot["status"] = "incomplete"
|
||||
billing_snapshot["missing_required"] = exc.missing_required
|
||||
billing_snapshot["resolved_values"] = {}
|
||||
billing_snapshot["error"] = "strict_mode_missing_required"
|
||||
cost = 0.0
|
||||
|
||||
# strict_mode=true:标记任务失败并隐藏产物,避免"免费放行"
|
||||
task.status = VideoStatus.FAILED.value
|
||||
task.error_code = "billing_incomplete"
|
||||
task.error_message = f"Missing required dimensions: {exc.missing_required}"
|
||||
task.video_url = None
|
||||
task.video_urls = None
|
||||
|
||||
await self._maybe_alert_missing_required(
|
||||
model=task.model,
|
||||
missing_required=exc.missing_required,
|
||||
)
|
||||
|
||||
billing_snapshot["cost_breakdown"] = {"total": cost}
|
||||
|
||||
# 将 billing_snapshot 回写到 task.request_metadata 便于对账(不会影响 usage 的单独存档)
|
||||
if task.request_metadata is None:
|
||||
task.request_metadata = {}
|
||||
if isinstance(task.request_metadata, dict):
|
||||
task.request_metadata["billing_snapshot"] = billing_snapshot
|
||||
|
||||
usage_metadata: dict[str, Any] = {
|
||||
"billing_snapshot": billing_snapshot,
|
||||
"dimensions": dims,
|
||||
"raw_response_ref": {
|
||||
"video_task_id": task.id,
|
||||
"field": "video_tasks.request_metadata.poll_raw_response",
|
||||
},
|
||||
}
|
||||
|
||||
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
|
||||
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
|
||||
api_key_obj = (
|
||||
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
|
||||
if task.api_key_id
|
||||
else None
|
||||
)
|
||||
provider_obj = (
|
||||
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
|
||||
if task.provider_id
|
||||
else None
|
||||
)
|
||||
provider_name = provider_obj.name if provider_obj else "unknown"
|
||||
|
||||
await UsageService.record_usage_with_custom_cost(
|
||||
db=self.db,
|
||||
user=user_obj,
|
||||
api_key=api_key_obj,
|
||||
provider=provider_name,
|
||||
model=task.model,
|
||||
request_type="video",
|
||||
total_cost_usd=cost,
|
||||
request_cost_usd=cost,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
api_format=task.client_api_format,
|
||||
endpoint_api_format=task.provider_api_format,
|
||||
has_format_conversion=bool(task.format_converted),
|
||||
is_stream=False,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=None,
|
||||
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
|
||||
error_message=(
|
||||
None
|
||||
if task.status == VideoStatus.COMPLETED.value
|
||||
else (task.error_message or task.error_code or "video_task_failed")
|
||||
),
|
||||
metadata=usage_metadata,
|
||||
request_headers=(
|
||||
(task.request_metadata or {}).get("request_headers")
|
||||
if isinstance(task.request_metadata, dict)
|
||||
else None
|
||||
),
|
||||
request_body=task.original_request_body,
|
||||
provider_request_headers=None,
|
||||
response_headers=None,
|
||||
client_response_headers=None,
|
||||
response_body=None,
|
||||
request_id=request_id,
|
||||
provider_id=task.provider_id,
|
||||
provider_endpoint_id=task.endpoint_id,
|
||||
provider_api_key_id=task.key_id,
|
||||
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
|
||||
target_model=None,
|
||||
)
|
||||
|
||||
async def _maybe_alert_missing_required(
|
||||
self, *, model: str, missing_required: list[str]
|
||||
) -> None:
|
||||
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
|
||||
if not missing_required:
|
||||
return
|
||||
if not self.redis:
|
||||
logger.error(
|
||||
"Missing required billing dimensions (model=%s): %s", model, missing_required
|
||||
)
|
||||
return
|
||||
|
||||
# 按小时 bucket 聚合
|
||||
now = datetime.now(timezone.utc)
|
||||
hour_bucket = now.strftime("%Y%m%d%H")
|
||||
for dim in missing_required:
|
||||
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
|
||||
try:
|
||||
count = await self.redis.incr(key)
|
||||
if count == 1:
|
||||
await self.redis.expire(key, 3700)
|
||||
if count >= 10:
|
||||
logger.warning(
|
||||
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
|
||||
model,
|
||||
dim,
|
||||
count,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["VideoTelemetry"]
|
||||
635
src/services/task/orchestrator.py
Normal file
635
src/services/task/orchestrator.py
Normal file
@@ -0,0 +1,635 @@
|
||||
"""
|
||||
AsyncTaskOrchestrator
|
||||
|
||||
提交阶段故障转移(多候选尝试):
|
||||
- 目标:拿到 external_task_id 后锁定 provider/endpoint/key,后续轮询不再切换。
|
||||
- 仅覆盖“提交阶段”;轮询阶段由各 task poller 使用已锁定的信息执行。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, RequestCandidate
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
|
||||
from src.services.orchestration.candidate_resolver import CandidateResolver
|
||||
from src.services.orchestration.error_classifier import ErrorClassifier
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _sanitize(message: str, max_length: int = 200) -> str:
|
||||
if not message:
|
||||
return "request_failed"
|
||||
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SubmitFunc(Protocol):
|
||||
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ExtractExternalTaskIdFunc(Protocol):
|
||||
def __call__(self, payload: dict[str, Any]) -> str | None: ...
|
||||
|
||||
|
||||
class UpstreamClientRequestError(RuntimeError):
|
||||
"""可判定为客户端请求问题(不应 failover)的上游错误。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
response: httpx.Response,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
) -> None:
|
||||
self.response = response
|
||||
self.candidate_keys = candidate_keys
|
||||
super().__init__(f"Upstream client error: HTTP {response.status_code}")
|
||||
|
||||
|
||||
class AllCandidatesFailedError(RuntimeError):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
reason: str,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
last_status_code: int | None = None,
|
||||
) -> None:
|
||||
self.reason = reason
|
||||
self.candidate_keys = candidate_keys
|
||||
self.last_status_code = last_status_code
|
||||
super().__init__(f"All candidates failed: {reason}")
|
||||
|
||||
|
||||
class CandidateUnsupportedError(RuntimeError):
|
||||
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
|
||||
|
||||
|
||||
class CandidateSubmissionError(RuntimeError):
|
||||
"""候选提交异常(网络/解密/解析等)。"""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SubmitOutcome:
|
||||
candidate: ProviderCandidate
|
||||
candidate_keys: list[dict[str, Any]]
|
||||
external_task_id: str
|
||||
rule_lookup: BillingRuleLookupResult | None
|
||||
upstream_payload: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class AsyncTaskOrchestrator:
|
||||
"""
|
||||
异步任务编排器:只负责提交阶段的候选遍历与错误处理策略。
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session, *, redis_client: Redis | None = None) -> None:
|
||||
self.db = db
|
||||
self.redis = redis_client
|
||||
self._candidate_resolver: CandidateResolver | None = None
|
||||
self._error_classifier: ErrorClassifier | None = None
|
||||
|
||||
self._cache_scheduler = None
|
||||
# 候选记录映射:{candidate_index: RequestCandidate}
|
||||
self._candidate_records: dict[int, RequestCandidate] = {}
|
||||
|
||||
def _create_candidate_records(
|
||||
self,
|
||||
candidates: list[ProviderCandidate],
|
||||
request_id: str | None,
|
||||
user_api_key: ApiKey,
|
||||
) -> dict[int, RequestCandidate]:
|
||||
"""
|
||||
为所有候选预创建 RequestCandidate 记录。
|
||||
|
||||
Args:
|
||||
candidates: 候选列表
|
||||
request_id: 请求 ID
|
||||
user_api_key: 用户 API Key
|
||||
|
||||
Returns:
|
||||
{candidate_index: RequestCandidate} 映射
|
||||
"""
|
||||
if not request_id:
|
||||
return {}
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
records: dict[int, RequestCandidate] = {}
|
||||
|
||||
for idx, cand in enumerate(candidates):
|
||||
record = RequestCandidate(
|
||||
id=str(uuid.uuid4()),
|
||||
request_id=request_id,
|
||||
candidate_index=idx,
|
||||
retry_index=0,
|
||||
user_id=user_api_key.user_id if user_api_key else None,
|
||||
api_key_id=user_api_key.id if user_api_key else None,
|
||||
provider_id=cand.provider.id,
|
||||
endpoint_id=cand.endpoint.id,
|
||||
key_id=cand.key.id,
|
||||
status="available",
|
||||
is_cached=bool(getattr(cand, "is_cached", False)),
|
||||
created_at=now,
|
||||
)
|
||||
self.db.add(record)
|
||||
records[idx] = record
|
||||
|
||||
try:
|
||||
self.db.flush()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Failed to create candidate records: %s",
|
||||
str(exc),
|
||||
)
|
||||
self.db.rollback()
|
||||
return {}
|
||||
|
||||
return records
|
||||
|
||||
def _update_candidate_record(
|
||||
self,
|
||||
idx: int,
|
||||
*,
|
||||
status: str,
|
||||
skip_reason: str | None = None,
|
||||
status_code: int | None = None,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
started_at: datetime | None = None,
|
||||
finished_at: datetime | None = None,
|
||||
) -> None:
|
||||
"""更新候选记录状态。"""
|
||||
record = self._candidate_records.get(idx)
|
||||
if not record:
|
||||
return
|
||||
|
||||
record.status = status
|
||||
if skip_reason is not None:
|
||||
record.skip_reason = skip_reason
|
||||
if status_code is not None:
|
||||
record.status_code = status_code
|
||||
if error_type is not None:
|
||||
record.error_type = error_type
|
||||
if error_message is not None:
|
||||
record.error_message = error_message
|
||||
if started_at is not None:
|
||||
record.started_at = started_at
|
||||
if finished_at is not None:
|
||||
record.finished_at = finished_at
|
||||
|
||||
try:
|
||||
self.db.flush()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Failed to update candidate record %d: %s",
|
||||
idx,
|
||||
str(exc),
|
||||
)
|
||||
|
||||
def _commit_candidate_records(self) -> None:
|
||||
"""提交候选记录到数据库。"""
|
||||
if not self._candidate_records:
|
||||
return
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Failed to commit candidate records: %s",
|
||||
str(exc),
|
||||
)
|
||||
self.db.rollback()
|
||||
|
||||
async def _ensure_initialized(self) -> None:
|
||||
if self._cache_scheduler is not None:
|
||||
return
|
||||
|
||||
# 使用 SystemConfigService 读取运行时调度策略(与 Chat/CLI 一致)
|
||||
priority_mode = SystemConfigService.get_config(
|
||||
self.db,
|
||||
"provider_priority_mode",
|
||||
"provider",
|
||||
)
|
||||
scheduling_mode = SystemConfigService.get_config(
|
||||
self.db,
|
||||
"scheduling_mode",
|
||||
"cache_affinity",
|
||||
)
|
||||
self._cache_scheduler = await get_cache_aware_scheduler(
|
||||
self.redis,
|
||||
priority_mode=priority_mode,
|
||||
scheduling_mode=scheduling_mode,
|
||||
)
|
||||
|
||||
self._candidate_resolver = CandidateResolver(
|
||||
db=self.db,
|
||||
cache_scheduler=self._cache_scheduler,
|
||||
)
|
||||
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
|
||||
|
||||
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
|
||||
"""
|
||||
判断某个上游 HTTP 错误是否为“客户端错误”(不应 failover)。
|
||||
|
||||
规则:
|
||||
- 401/403/429:一般是 key/权限/限流问题,优先 failover
|
||||
- 其他 4xx:若 ErrorClassifier 判断为客户端请求错误,则停止
|
||||
"""
|
||||
if status_code in (401, 403, 429):
|
||||
return False
|
||||
if 400 <= status_code < 500:
|
||||
assert self._error_classifier is not None
|
||||
return self._error_classifier.is_client_error(error_text)
|
||||
return False
|
||||
|
||||
async def submit_with_failover(
|
||||
self,
|
||||
*,
|
||||
api_format: str,
|
||||
model_name: str,
|
||||
affinity_key: str,
|
||||
user_api_key: ApiKey,
|
||||
request_id: str | None,
|
||||
task_type: str,
|
||||
submit_func: SubmitFunc,
|
||||
extract_external_task_id: ExtractExternalTaskIdFunc,
|
||||
supported_auth_types: set[str] | None = None,
|
||||
allow_format_conversion: bool = False,
|
||||
capability_requirements: dict[str, bool] | None = None,
|
||||
max_candidates: int | None = None,
|
||||
) -> SubmitOutcome:
|
||||
"""
|
||||
提交异步任务并在失败时自动尝试下一个候选,直到拿到 external_task_id。
|
||||
|
||||
Returns:
|
||||
SubmitOutcome(包含选中的候选 + external_task_id + candidate_keys + billing rule lookup)
|
||||
|
||||
Raises:
|
||||
UpstreamClientRequestError: 判定为客户端请求错误(不应 failover)
|
||||
ProviderNotAvailableException: 没有可用候选(调度器层面)
|
||||
AllCandidatesFailedError: 有候选但全部提交失败
|
||||
"""
|
||||
await self._ensure_initialized()
|
||||
assert self._candidate_resolver is not None
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] submit_with_failover: "
|
||||
"api_format=%s, model=%s, task_type=%s, request_id=%s",
|
||||
api_format,
|
||||
model_name,
|
||||
task_type,
|
||||
request_id,
|
||||
)
|
||||
|
||||
candidates, _global_model_id = await self._candidate_resolver.fetch_candidates(
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=affinity_key,
|
||||
user_api_key=user_api_key,
|
||||
request_id=request_id,
|
||||
is_stream=False,
|
||||
capability_requirements=capability_requirements,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] fetch_candidates returned %d candidates for model=%s",
|
||||
len(candidates),
|
||||
model_name,
|
||||
)
|
||||
|
||||
# 如果没有候选,直接抛出异常
|
||||
if not candidates:
|
||||
logger.error(
|
||||
"[AsyncTaskOrchestrator] No candidates returned from fetch_candidates for model=%s",
|
||||
model_name,
|
||||
)
|
||||
raise ProviderNotAvailableException("No candidates available")
|
||||
|
||||
if max_candidates is not None and max_candidates > 0:
|
||||
candidates = candidates[:max_candidates]
|
||||
|
||||
# 创建候选记录(用于链路追踪)
|
||||
self._candidate_records = self._create_candidate_records(
|
||||
candidates=candidates,
|
||||
request_id=request_id,
|
||||
user_api_key=user_api_key,
|
||||
)
|
||||
|
||||
candidate_keys: list[dict[str, Any]] = []
|
||||
eligible_count = 0
|
||||
last_status_code: int | None = None
|
||||
|
||||
for idx, cand in enumerate(candidates):
|
||||
submit_started_at = datetime.now(timezone.utc)
|
||||
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
|
||||
candidate_info: dict[str, Any] = {
|
||||
"index": idx,
|
||||
"provider_id": cand.provider.id,
|
||||
"provider_name": cand.provider.name,
|
||||
"endpoint_id": cand.endpoint.id,
|
||||
"key_id": cand.key.id,
|
||||
"key_name": cand.key.name,
|
||||
"auth_type": auth_type,
|
||||
"priority": getattr(cand.key, "priority", 0) or 0,
|
||||
"is_cached": bool(getattr(cand, "is_cached", False)),
|
||||
}
|
||||
candidate_keys.append(candidate_info)
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Checking candidate %d: provider=%s, is_skipped=%s, skip_reason=%s, needs_conversion=%s, auth_type=%s",
|
||||
idx,
|
||||
cand.provider.name,
|
||||
getattr(cand, "is_skipped", False),
|
||||
getattr(cand, "skip_reason", None),
|
||||
getattr(cand, "needs_conversion", False),
|
||||
auth_type,
|
||||
)
|
||||
|
||||
# 调度器层面标记为跳过(健康/熔断/并发等)
|
||||
if getattr(cand, "is_skipped", False):
|
||||
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
|
||||
candidate_info.update(
|
||||
{
|
||||
"skipped": True,
|
||||
"skip_reason": skip_reason,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d skipped: is_skipped=True, reason=%s",
|
||||
idx,
|
||||
cand.skip_reason,
|
||||
)
|
||||
continue
|
||||
|
||||
# 视频/图片等直连 upstream 的 handler 目前不支持跨格式转换
|
||||
if not allow_format_conversion and bool(getattr(cand, "needs_conversion", False)):
|
||||
candidate_info.update(
|
||||
{"skipped": True, "skip_reason": "format_conversion_not_supported"}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx, status="skipped", skip_reason="format_conversion_not_supported"
|
||||
)
|
||||
logger.info("[AsyncTaskOrchestrator] Candidate %d skipped: needs_conversion", idx)
|
||||
continue
|
||||
|
||||
# auth_type 过滤
|
||||
if supported_auth_types is not None and auth_type not in supported_auth_types:
|
||||
skip_reason = f"unsupported_auth_type:{auth_type}"
|
||||
candidate_info.update(
|
||||
{
|
||||
"skipped": True,
|
||||
"skip_reason": skip_reason,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d skipped: unsupported_auth_type=%s",
|
||||
idx,
|
||||
auth_type,
|
||||
)
|
||||
continue
|
||||
|
||||
# billing rule 过滤(可选)
|
||||
rule_lookup: BillingRuleLookupResult | None = None
|
||||
has_billing_rule = True
|
||||
if config.billing_require_rule:
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Checking billing rule for candidate %d (billing_require_rule=True)",
|
||||
idx,
|
||||
)
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=cand.provider.id,
|
||||
model_name=model_name,
|
||||
task_type=task_type,
|
||||
)
|
||||
has_billing_rule = rule_lookup is not None
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Billing rule lookup result: has_rule=%s",
|
||||
has_billing_rule,
|
||||
)
|
||||
if not has_billing_rule:
|
||||
candidate_info.update(
|
||||
{
|
||||
"has_billing_rule": False,
|
||||
"skipped": True,
|
||||
"skip_reason": "billing_rule_missing",
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx, status="skipped", skip_reason="billing_rule_missing"
|
||||
)
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d skipped: billing_rule_missing", idx
|
||||
)
|
||||
continue
|
||||
candidate_info["has_billing_rule"] = has_billing_rule
|
||||
|
||||
logger.info("[AsyncTaskOrchestrator] Candidate %d eligible, attempting submit", idx)
|
||||
eligible_count += 1
|
||||
|
||||
# 更新记录为 pending 状态(开始尝试)
|
||||
self._update_candidate_record(idx, status="pending", started_at=submit_started_at)
|
||||
|
||||
# 尝试提交
|
||||
try:
|
||||
response = await submit_func(cand)
|
||||
except Exception as exc:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
logger.error(
|
||||
"[AsyncTaskOrchestrator] Candidate %d submit exception: %s: %s",
|
||||
idx,
|
||||
type(exc).__name__,
|
||||
str(exc),
|
||||
)
|
||||
error_msg = _sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "exception",
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
error_type=type(exc).__name__,
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
continue
|
||||
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d submit response: status_code=%d",
|
||||
idx,
|
||||
response.status_code,
|
||||
)
|
||||
|
||||
last_status_code = int(getattr(response, "status_code", 0) or 0)
|
||||
|
||||
# 上游错误:决定是否停止
|
||||
if response.status_code >= 400:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
error_text = ""
|
||||
try:
|
||||
error_text = response.text or ""
|
||||
except Exception:
|
||||
error_text = ""
|
||||
error_msg = _sanitize(error_text)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "http_error",
|
||||
"status_code": response.status_code,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="http_error",
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
if self._should_stop_on_http_error(
|
||||
status_code=response.status_code, error_text=error_text
|
||||
):
|
||||
self._commit_candidate_records()
|
||||
raise UpstreamClientRequestError(
|
||||
response=response,
|
||||
candidate_keys=candidate_keys,
|
||||
)
|
||||
continue
|
||||
|
||||
# 解析任务 ID(200 但缺字段也视为失败并 failover)
|
||||
payload: dict[str, Any] | None = None
|
||||
try:
|
||||
data = response.json()
|
||||
if isinstance(data, dict):
|
||||
payload = data
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d response payload: %s",
|
||||
idx,
|
||||
str(payload)[:500] if payload else "None",
|
||||
)
|
||||
except Exception as exc:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
logger.error(
|
||||
"[AsyncTaskOrchestrator] Candidate %d invalid JSON: %s",
|
||||
idx,
|
||||
str(exc),
|
||||
)
|
||||
error_msg = _sanitize(str(exc))
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "invalid_json",
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": error_msg,
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="invalid_json",
|
||||
error_message=error_msg,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
continue
|
||||
|
||||
external_task_id = extract_external_task_id(payload or {})
|
||||
logger.info(
|
||||
"[AsyncTaskOrchestrator] Candidate %d extracted task_id: %s",
|
||||
idx,
|
||||
external_task_id,
|
||||
)
|
||||
if not external_task_id:
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
candidate_info.update(
|
||||
{
|
||||
"attempt_status": "empty_task_id",
|
||||
"error_message": "Upstream returned empty task id",
|
||||
}
|
||||
)
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="failed",
|
||||
status_code=response.status_code,
|
||||
error_type="empty_task_id",
|
||||
error_message="Upstream returned empty task id",
|
||||
finished_at=finished_at,
|
||||
)
|
||||
logger.warning(
|
||||
"[AsyncTaskOrchestrator] Candidate %d: empty task_id, payload keys: %s",
|
||||
idx,
|
||||
list(payload.keys()) if payload else [],
|
||||
)
|
||||
continue
|
||||
|
||||
# 成功
|
||||
finished_at = datetime.now(timezone.utc)
|
||||
candidate_info.update({"attempt_status": "success", "selected": True})
|
||||
self._update_candidate_record(
|
||||
idx,
|
||||
status="success",
|
||||
status_code=response.status_code,
|
||||
finished_at=finished_at,
|
||||
)
|
||||
self._commit_candidate_records()
|
||||
return SubmitOutcome(
|
||||
candidate=cand,
|
||||
candidate_keys=candidate_keys,
|
||||
external_task_id=str(external_task_id),
|
||||
rule_lookup=rule_lookup,
|
||||
upstream_payload=payload,
|
||||
)
|
||||
|
||||
# 没有任何候选可尝试
|
||||
if not candidates:
|
||||
raise ProviderNotAvailableException("No candidates available")
|
||||
|
||||
# 提交所有候选记录
|
||||
self._commit_candidate_records()
|
||||
|
||||
if eligible_count == 0:
|
||||
reason = "no_eligible_candidates"
|
||||
if config.billing_require_rule:
|
||||
reason = "no_candidate_with_billing_rule"
|
||||
raise AllCandidatesFailedError(
|
||||
reason=reason,
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
|
||||
raise AllCandidatesFailedError(
|
||||
reason="all_candidates_failed",
|
||||
candidate_keys=candidate_keys,
|
||||
last_status_code=last_status_code,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AsyncTaskOrchestrator",
|
||||
"SubmitOutcome",
|
||||
"AllCandidatesFailedError",
|
||||
"UpstreamClientRequestError",
|
||||
"CandidateUnsupportedError",
|
||||
"CandidateSubmissionError",
|
||||
]
|
||||
@@ -4,7 +4,6 @@ Usage Redis Streams consumer.
|
||||
高性能消费者实现,支持批量处理和单次提交多条记录。
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -104,9 +103,7 @@ async def ensure_usage_stream_group() -> None:
|
||||
id="0-0",
|
||||
mkstream=True,
|
||||
)
|
||||
logger.info(
|
||||
f"[usage-queue] Created consumer group {config.usage_queue_stream_group}"
|
||||
)
|
||||
logger.info(f"[usage-queue] Created consumer group {config.usage_queue_stream_group}")
|
||||
except ResponseError as exc:
|
||||
if "BUSYGROUP" in str(exc):
|
||||
return
|
||||
@@ -115,7 +112,7 @@ async def ensure_usage_stream_group() -> None:
|
||||
|
||||
class UsageQueueConsumer:
|
||||
"""Usage 队列消费者
|
||||
|
||||
|
||||
性能优化:
|
||||
- 缓存配置值避免重复属性访问
|
||||
- STREAMING 事件使用 pipeline 批量 ACK
|
||||
@@ -259,14 +256,14 @@ class UsageQueueConsumer:
|
||||
) -> None:
|
||||
"""批量处理 STREAMING 事件(状态更新)"""
|
||||
success_ids: list[str] = []
|
||||
|
||||
|
||||
for message_id, event in messages:
|
||||
try:
|
||||
await self._apply_streaming_event(event)
|
||||
success_ids.append(message_id)
|
||||
except Exception as exc:
|
||||
await self._handle_processing_error(redis_client, message_id, {}, exc)
|
||||
|
||||
|
||||
# 使用 pipeline 批量 ACK 成功处理的消息
|
||||
if success_ids:
|
||||
pipe = redis_client.pipeline()
|
||||
@@ -276,8 +273,8 @@ class UsageQueueConsumer:
|
||||
|
||||
async def _process_record_batch(
|
||||
self,
|
||||
redis_client: Any,
|
||||
messages: list[tuple[str, dict[str, Any], UsageEvent]],
|
||||
redis_client: Any,
|
||||
messages: list[tuple[str, dict[str, Any], UsageEvent]],
|
||||
) -> None:
|
||||
"""批量处理记录类型的事件"""
|
||||
db = create_session()
|
||||
@@ -306,7 +303,9 @@ class UsageQueueConsumer:
|
||||
|
||||
except Exception as exc:
|
||||
# 批量处理失败,回退到逐条处理(复用已创建的 db session)
|
||||
logger.warning(f"[usage-queue] Batch processing failed, falling back to individual: {exc}")
|
||||
logger.warning(
|
||||
f"[usage-queue] Batch processing failed, falling back to individual: {exc}"
|
||||
)
|
||||
try:
|
||||
db.rollback() # 清理批量失败的事务状态
|
||||
except Exception:
|
||||
@@ -330,7 +329,9 @@ class UsageQueueConsumer:
|
||||
else:
|
||||
await self._handle_processing_error(redis_client, message_id, fields, ie)
|
||||
except Exception as individual_exc:
|
||||
await self._handle_processing_error(redis_client, message_id, fields, individual_exc)
|
||||
await self._handle_processing_error(
|
||||
redis_client, message_id, fields, individual_exc
|
||||
)
|
||||
# 批量 ACK 成功处理的消息
|
||||
if success_ids:
|
||||
pipe = redis_client.pipeline()
|
||||
@@ -342,10 +343,10 @@ class UsageQueueConsumer:
|
||||
|
||||
async def _handle_processing_error(
|
||||
self,
|
||||
redis_client: Any,
|
||||
message_id: str,
|
||||
fields: dict[str, Any],
|
||||
error: Exception,
|
||||
redis_client: Any,
|
||||
message_id: str,
|
||||
fields: dict[str, Any],
|
||||
error: Exception,
|
||||
) -> None:
|
||||
retries = await self._get_delivery_count(redis_client, message_id)
|
||||
if retries >= self._max_retries:
|
||||
@@ -415,9 +416,7 @@ class UsageQueueConsumer:
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def _apply_record_event(
|
||||
self, event: UsageEvent, db: Session | None = None
|
||||
) -> None:
|
||||
async def _apply_record_event(self, event: UsageEvent, db: Session | None = None) -> None:
|
||||
"""处理记录类型事件(逐条写入,用于 fallback)
|
||||
|
||||
Args:
|
||||
@@ -501,9 +500,7 @@ class UsageQueueConsumer:
|
||||
pending_count = int(pending.get("pending", 0))
|
||||
elif isinstance(pending, (list, tuple)) and pending:
|
||||
pending_count = int(pending[0])
|
||||
logger.info(
|
||||
f"[usage-queue] backlog={stream_len} pending={pending_count}"
|
||||
)
|
||||
logger.info(f"[usage-queue] backlog={stream_len} pending={pending_count}")
|
||||
except Exception as exc:
|
||||
logger.debug(f"[usage-queue] metrics log failed: {exc}")
|
||||
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
Usage 事件定义与序列化工具(用于 Redis Streams)
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
||||
@@ -10,8 +10,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.core.logger import logger
|
||||
|
||||
@@ -36,7 +36,6 @@ from src.services.system.audit import audit_service
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
|
||||
class UsageRecorder:
|
||||
"""
|
||||
统一的 Usage 记录器
|
||||
@@ -144,9 +143,11 @@ class UsageRecorder:
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
logger.debug(f"[UsageRecorder] 成功记录: provider={metadata.provider}, "
|
||||
logger.debug(
|
||||
f"[UsageRecorder] 成功记录: provider={metadata.provider}, "
|
||||
f"model={metadata.model}, api_format={metadata.api_format}, "
|
||||
f"tokens={usage.input_tokens}+{usage.output_tokens}")
|
||||
f"tokens={usage.input_tokens}+{usage.output_tokens}"
|
||||
)
|
||||
|
||||
async def record_failure(
|
||||
self,
|
||||
@@ -213,9 +214,11 @@ class UsageRecorder:
|
||||
error_message=result.error_message,
|
||||
)
|
||||
|
||||
logger.debug(f"[UsageRecorder] 失败记录: provider={metadata.provider}, "
|
||||
logger.debug(
|
||||
f"[UsageRecorder] 失败记录: provider={metadata.provider}, "
|
||||
f"model={metadata.model}, api_format={metadata.api_format}, "
|
||||
f"status={result.status_code}, error={result.error_message[:100] if result.error_message else 'N/A'}")
|
||||
f"status={result.status_code}, error={result.error_message[:100] if result.error_message else 'N/A'}"
|
||||
)
|
||||
|
||||
async def record_from_exception(
|
||||
self,
|
||||
|
||||
@@ -12,7 +12,8 @@ from typing import Any
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.api_format.metadata import can_passthrough
|
||||
from src.core.api_format.metadata import can_passthrough_endpoint
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
from src.core.logger import logger
|
||||
from src.models.database import (
|
||||
ApiKey,
|
||||
@@ -2607,7 +2608,12 @@ class UsageService:
|
||||
|
||||
# 兼容历史数据:当 streaming 状态已拿到两个格式但 has_format_conversion 为空时,回填推断结果
|
||||
if has_format_conversion is None and api_format and endpoint_api_format:
|
||||
has_format_conversion = not can_passthrough(api_format, endpoint_api_format)
|
||||
client_raw = str(api_format).strip()
|
||||
endpoint_raw = str(endpoint_api_format).strip()
|
||||
if ":" in client_raw and ":" in endpoint_raw:
|
||||
client_fmt = normalize_signature_key(client_raw)
|
||||
endpoint_fmt = normalize_signature_key(endpoint_raw)
|
||||
has_format_conversion = not can_passthrough_endpoint(client_fmt, endpoint_fmt)
|
||||
|
||||
item: dict[str, Any] = {
|
||||
"id": r.id,
|
||||
|
||||
@@ -7,8 +7,8 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -21,7 +21,6 @@ from src.models.database import ApiKey, User
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
|
||||
class StreamUsageTracker:
|
||||
"""流式响应用量跟踪器"""
|
||||
|
||||
@@ -64,7 +63,7 @@ class StreamUsageTracker:
|
||||
provider_id: Provider ID(用于记录真实成本)
|
||||
provider_endpoint_id: Endpoint ID(用于记录真实成本)
|
||||
provider_api_key_id: API Key ID(用于记录真实成本)
|
||||
api_format: API 格式(CLAUDE, CLAUDE_CLI, OPENAI, OPENAI_CLI)
|
||||
api_format: endpoint signature(如 "claude:chat", "openai:cli")
|
||||
endpoint_api_format: 端点原生 API 格式
|
||||
has_format_conversion: 是否发生了格式转换
|
||||
"""
|
||||
@@ -85,7 +84,7 @@ class StreamUsageTracker:
|
||||
self.provider_api_key_id = provider_api_key_id
|
||||
|
||||
# API 格式和响应解析器
|
||||
self.api_format = api_format or "CLAUDE"
|
||||
self.api_format = api_format or "claude:chat"
|
||||
self.endpoint_api_format = endpoint_api_format
|
||||
self.has_format_conversion = has_format_conversion
|
||||
self.response_parser = get_parser_for_format(self.api_format)
|
||||
@@ -144,7 +143,9 @@ class StreamUsageTracker:
|
||||
"""
|
||||
self.status_code = status_code
|
||||
self.error_message = error_message
|
||||
logger.debug(f"ID:{self.request_id} | 流式响应错误状态已设置 | 状态码:{status_code} | 错误:{error_message[:100]}")
|
||||
logger.debug(
|
||||
f"ID:{self.request_id} | 流式响应错误状态已设置 | 状态码:{status_code} | 错误:{error_message[:100]}"
|
||||
)
|
||||
|
||||
def _update_complete_response(self, chunk: dict[str, Any]) -> None:
|
||||
"""根据响应块更新完整响应结构"""
|
||||
@@ -465,7 +466,9 @@ class StreamUsageTracker:
|
||||
messages = request_data.get("messages", [])
|
||||
self.input_tokens = self.estimate_input_tokens(messages)
|
||||
|
||||
logger.debug(f"ID:{self.request_id} | 开始跟踪流式响应 | 估算输入tokens:{self.input_tokens}")
|
||||
logger.debug(
|
||||
f"ID:{self.request_id} | 开始跟踪流式响应 | 估算输入tokens:{self.input_tokens}"
|
||||
)
|
||||
|
||||
chunk_count = 0
|
||||
first_byte_time_ms = None # 预先记录 TTFB,避免 yield 后计算不准确
|
||||
@@ -479,7 +482,9 @@ class StreamUsageTracker:
|
||||
if chunk_count == 1:
|
||||
# 计算 TTFB(使用请求原始开始时间或 track_stream 开始时间)
|
||||
base_time = self.request_start_time or self.start_time
|
||||
first_byte_time_ms = int((time.time() - base_time) * 1000) if base_time else None
|
||||
first_byte_time_ms = (
|
||||
int((time.time() - base_time) * 1000) if base_time else None
|
||||
)
|
||||
|
||||
# 先返回原始块给客户端,确保 TTFB 不受数据库操作影响
|
||||
yield chunk
|
||||
@@ -535,8 +540,10 @@ class StreamUsageTracker:
|
||||
# 流结束后记录使用量
|
||||
self.end_time = time.time()
|
||||
|
||||
logger.debug(f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
|
||||
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}")
|
||||
logger.debug(
|
||||
f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
|
||||
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}"
|
||||
)
|
||||
|
||||
# 检查是否收到了有效数据
|
||||
# 情况1: 收到了原始数据但无法解析为有效的SSE JSON
|
||||
@@ -564,7 +571,9 @@ class StreamUsageTracker:
|
||||
await self._record_usage()
|
||||
except Exception as e:
|
||||
# 如果记录失败,至少输出基本的汇总日志
|
||||
logger.exception(f"Failed to record stream usage for request {self.request_id}: {e}")
|
||||
logger.exception(
|
||||
f"Failed to record stream usage for request {self.request_id}: {e}"
|
||||
)
|
||||
# 尝试输出基本的汇总日志,使用多层防护
|
||||
try:
|
||||
# 计算响应时间,使用多层后备机制
|
||||
@@ -581,13 +590,17 @@ class StreamUsageTracker:
|
||||
total_response_time = 0
|
||||
|
||||
# 安全地输出汇总日志
|
||||
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
|
||||
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)")
|
||||
logger.info(
|
||||
f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
|
||||
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)"
|
||||
)
|
||||
except Exception as log_error:
|
||||
# 最后的防线:输出最简单的完成标记
|
||||
logger.error(f"Failed to output summary log: {log_error}")
|
||||
try:
|
||||
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成")
|
||||
logger.info(
|
||||
f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成"
|
||||
)
|
||||
except Exception:
|
||||
# 如果连最简单的日志都失败了,放弃
|
||||
pass
|
||||
@@ -686,7 +699,9 @@ class StreamUsageTracker:
|
||||
user, api_key = _load_user_and_key(db_for_usage)
|
||||
except InvalidRequestError:
|
||||
# 会话处于不可用状态,需要回滚并重新开始
|
||||
logger.warning(f"Session in invalid state for request {self.request_id}, rolling back and retrying")
|
||||
logger.warning(
|
||||
f"Session in invalid state for request {self.request_id}, rolling back and retrying"
|
||||
)
|
||||
try:
|
||||
db_for_usage.rollback()
|
||||
except Exception:
|
||||
@@ -703,7 +718,9 @@ class StreamUsageTracker:
|
||||
created_temp_session = True
|
||||
user, api_key = _load_user_and_key(db_for_usage)
|
||||
except Exception as session_error:
|
||||
logger.exception(f"Failed to recover from invalid session for request {self.request_id}: {session_error}")
|
||||
logger.exception(
|
||||
f"Failed to recover from invalid session for request {self.request_id}: {session_error}"
|
||||
)
|
||||
return
|
||||
|
||||
# 根据状态码确定请求状态
|
||||
@@ -774,8 +791,10 @@ class StreamUsageTracker:
|
||||
else:
|
||||
cost_str = "$0"
|
||||
|
||||
logger.info(f"{status_prefix} ID:{self.request_id} | {self.status_code} | 耗时:{total_response_time}ms | "
|
||||
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:{cost_str}")
|
||||
logger.info(
|
||||
f"{status_prefix} ID:{self.request_id} | {self.status_code} | 耗时:{total_response_time}ms | "
|
||||
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:{cost_str}"
|
||||
)
|
||||
|
||||
# 记录提供商结果用于动态权重调整
|
||||
# 记录提供商结果的健康监控已由 FallbackOrchestrator 自动处理
|
||||
@@ -936,7 +955,9 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
|
||||
messages = request_data.get("messages", [])
|
||||
self.input_tokens = self.estimate_input_tokens(messages)
|
||||
|
||||
logger.debug(f"ID:{self.request_id} | 开始跟踪流式响应(Enhanced) | 估算输入tokens:{self.input_tokens}")
|
||||
logger.debug(
|
||||
f"ID:{self.request_id} | 开始跟踪流式响应(Enhanced) | 估算输入tokens:{self.input_tokens}"
|
||||
)
|
||||
|
||||
chunk_count = 0
|
||||
first_byte_time_ms = None # 预先记录 TTFB,避免 yield 后计算不准确
|
||||
@@ -950,7 +971,9 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
|
||||
if chunk_count == 1:
|
||||
# 计算 TTFB(使用请求原始开始时间或 track_stream 开始时间)
|
||||
base_time = self.request_start_time or self.start_time
|
||||
first_byte_time_ms = int((time.time() - base_time) * 1000) if base_time else None
|
||||
first_byte_time_ms = (
|
||||
int((time.time() - base_time) * 1000) if base_time else None
|
||||
)
|
||||
|
||||
# 先返回原始块给客户端,确保 TTFB 不受数据库操作影响
|
||||
yield chunk
|
||||
@@ -997,8 +1020,10 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
|
||||
# 流结束后记录使用量
|
||||
self.end_time = time.time()
|
||||
|
||||
logger.debug(f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
|
||||
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}")
|
||||
logger.debug(
|
||||
f"ID:{self.request_id} | 流式响应结束 | 共处理{chunk_count}个chunks | "
|
||||
f"累积内容长度:{len(self.accumulated_content)} | 输出tokens:{self.output_tokens}"
|
||||
)
|
||||
|
||||
# 检查是否收到了有效数据
|
||||
# 情况1: 收到了原始数据但无法解析为有效的SSE JSON
|
||||
@@ -1026,7 +1051,9 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
|
||||
await self._record_usage()
|
||||
except Exception as e:
|
||||
# 如果记录失败,至少输出基本的汇总日志
|
||||
logger.exception(f"Failed to record stream usage for request {self.request_id}: {e}")
|
||||
logger.exception(
|
||||
f"Failed to record stream usage for request {self.request_id}: {e}"
|
||||
)
|
||||
# 尝试输出基本的汇总日志,使用多层防护
|
||||
try:
|
||||
# 计算响应时间,使用多层后备机制
|
||||
@@ -1043,13 +1070,17 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
|
||||
total_response_time = 0
|
||||
|
||||
# 安全地输出汇总日志
|
||||
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
|
||||
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)")
|
||||
logger.info(
|
||||
f"[请求完成] ID:{self.request_id or 'unknown'} | 200 | 耗时:{total_response_time}ms | "
|
||||
f"Token:输入{self.input_tokens}/输出{self.output_tokens} | 费用:未知(记录失败)"
|
||||
)
|
||||
except Exception as log_error:
|
||||
# 最后的防线:输出最简单的完成标记
|
||||
logger.error(f"Failed to output summary log: {log_error}")
|
||||
try:
|
||||
logger.info(f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成")
|
||||
logger.info(
|
||||
f"[请求完成] ID:{self.request_id or 'unknown'} | 记录失败但流已完成"
|
||||
)
|
||||
except Exception:
|
||||
# 如果连最简单的日志都失败了,放弃
|
||||
pass
|
||||
@@ -1096,7 +1127,7 @@ def create_stream_tracker(
|
||||
provider_id: Provider ID(用于记录真实成本)
|
||||
provider_endpoint_id: Endpoint ID(用于记录真实成本)
|
||||
provider_api_key_id: API Key ID(用于记录真实成本)
|
||||
api_format: API 格式(CLAUDE, CLAUDE_CLI, OPENAI, OPENAI_CLI)
|
||||
api_format: endpoint signature(如 "claude:chat", "openai:cli")
|
||||
endpoint_api_format: 端点原生 API 格式
|
||||
has_format_conversion: 是否发生了格式转换
|
||||
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
Telemetry writer abstraction for stream usage.
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
@@ -124,7 +123,7 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
"user_id": self.user_id,
|
||||
"api_key_id": self.api_key_id,
|
||||
}
|
||||
|
||||
|
||||
# 可选字段 - 只添加非 None/非默认值,减少 payload 大小
|
||||
# 注意:消费者端需要处理缺失字段的默认值
|
||||
if kwargs.get("provider"):
|
||||
@@ -133,7 +132,7 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
data["model"] = kwargs["model"]
|
||||
if kwargs.get("target_model"):
|
||||
data["target_model"] = kwargs["target_model"]
|
||||
|
||||
|
||||
# Token 计数 - 0 是常见值,但仍需传递
|
||||
input_tokens = kwargs.get("input_tokens", 0)
|
||||
output_tokens = kwargs.get("output_tokens", 0)
|
||||
@@ -141,7 +140,7 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
data["input_tokens"] = input_tokens
|
||||
if output_tokens:
|
||||
data["output_tokens"] = output_tokens
|
||||
|
||||
|
||||
# 缓存 token(cache_creation_tokens -> cache_creation_input_tokens 映射)
|
||||
cache_creation = kwargs.get("cache_creation_tokens", 0)
|
||||
cache_read = kwargs.get("cache_read_tokens", 0)
|
||||
@@ -149,20 +148,20 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
data["cache_creation_input_tokens"] = cache_creation
|
||||
if cache_read:
|
||||
data["cache_read_input_tokens"] = cache_read
|
||||
|
||||
|
||||
# 时间指标
|
||||
if kwargs.get("response_time_ms") is not None:
|
||||
data["response_time_ms"] = kwargs["response_time_ms"]
|
||||
if kwargs.get("first_byte_time_ms") is not None:
|
||||
data["first_byte_time_ms"] = kwargs["first_byte_time_ms"]
|
||||
|
||||
|
||||
# 状态信息
|
||||
status_code = kwargs.get("status_code", 200)
|
||||
if status_code != 200:
|
||||
data["status_code"] = status_code
|
||||
if kwargs.get("error_message"):
|
||||
data["error_message"] = kwargs["error_message"]
|
||||
|
||||
|
||||
# 格式信息
|
||||
request_type = kwargs.get("request_type", "chat")
|
||||
if request_type != "chat":
|
||||
@@ -173,11 +172,11 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
data["endpoint_api_format"] = kwargs["endpoint_api_format"]
|
||||
if kwargs.get("has_format_conversion"):
|
||||
data["has_format_conversion"] = True
|
||||
|
||||
|
||||
# 流式标记 - 默认 True,只记录 False
|
||||
if not kwargs.get("is_stream", True):
|
||||
data["is_stream"] = False
|
||||
|
||||
|
||||
# Provider 追踪
|
||||
if kwargs.get("provider_id"):
|
||||
data["provider_id"] = kwargs["provider_id"]
|
||||
@@ -185,7 +184,7 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
data["provider_endpoint_id"] = kwargs["provider_endpoint_id"]
|
||||
if kwargs.get("provider_api_key_id"):
|
||||
data["provider_api_key_id"] = kwargs["provider_api_key_id"]
|
||||
|
||||
|
||||
# 元数据
|
||||
if kwargs.get("metadata"):
|
||||
data["metadata"] = kwargs["metadata"]
|
||||
|
||||
@@ -15,7 +15,6 @@ from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Usage
|
||||
|
||||
|
||||
|
||||
class ApiKeyService:
|
||||
"""API密钥管理服务"""
|
||||
|
||||
@@ -86,8 +85,10 @@ class ApiKeyService:
|
||||
db.commit()
|
||||
db.refresh(api_key)
|
||||
|
||||
logger.info(f"创建API密钥: 用户ID {user_id}, 密钥名 {api_key.name}, "
|
||||
f"独立Key={is_standalone}, 初始余额={initial_balance_usd}")
|
||||
logger.info(
|
||||
f"创建API密钥: 用户ID {user_id}, 密钥名 {api_key.name}, "
|
||||
f"独立Key={is_standalone}, 初始余额={initial_balance_usd}"
|
||||
)
|
||||
return api_key, key # 返回密钥对象和明文密钥
|
||||
|
||||
@staticmethod
|
||||
@@ -256,7 +257,9 @@ class ApiKeyService:
|
||||
is_allowed = request_count < api_key.rate_limit
|
||||
|
||||
if not is_allowed:
|
||||
logger.warning(f"API密钥速率限制: Key ID {api_key.id}, 请求数 {request_count}/{api_key.rate_limit}")
|
||||
logger.warning(
|
||||
f"API密钥速率限制: Key ID {api_key.id}, 请求数 {request_count}/{api_key.rate_limit}"
|
||||
)
|
||||
|
||||
return is_allowed, api_key.rate_limit - request_count
|
||||
|
||||
@@ -289,7 +292,9 @@ class ApiKeyService:
|
||||
if amount_usd < 0:
|
||||
current = api_key.current_balance_usd or 0
|
||||
if abs(amount_usd) > current:
|
||||
logger.warning(f"余额扣除失败: 扣除金额 ${abs(amount_usd):.4f} 超过当前余额 ${current:.4f}")
|
||||
logger.warning(
|
||||
f"余额扣除失败: 扣除金额 ${abs(amount_usd):.4f} 超过当前余额 ${current:.4f}"
|
||||
)
|
||||
return None
|
||||
|
||||
# 调整当前余额
|
||||
@@ -303,8 +308,10 @@ class ApiKeyService:
|
||||
db.refresh(api_key)
|
||||
|
||||
action = "增加" if amount_usd > 0 else "扣除"
|
||||
logger.info(f"余额调整成功: Key ID {key_id}, {action} ${abs(amount_usd):.4f}, "
|
||||
f"新余额 ${api_key.current_balance_usd:.4f}")
|
||||
logger.info(
|
||||
f"余额调整成功: Key ID {key_id}, {action} ${abs(amount_usd):.4f}, "
|
||||
f"新余额 ${api_key.current_balance_usd:.4f}"
|
||||
)
|
||||
return api_key
|
||||
|
||||
@staticmethod
|
||||
@@ -338,14 +345,18 @@ class ApiKeyService:
|
||||
if should_delete:
|
||||
# 物理删除(Usage记录会保留,因为是 SET NULL)
|
||||
db.delete(api_key)
|
||||
logger.info(f"删除过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
|
||||
f"过期时间 {api_key.expires_at}")
|
||||
logger.info(
|
||||
f"删除过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
|
||||
f"过期时间 {api_key.expires_at}"
|
||||
)
|
||||
else:
|
||||
# 仅禁用
|
||||
api_key.is_active = False
|
||||
api_key.updated_at = now
|
||||
logger.info(f"禁用过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
|
||||
f"过期时间 {api_key.expires_at}")
|
||||
logger.info(
|
||||
f"禁用过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
|
||||
f"过期时间 {api_key.expires_at}"
|
||||
)
|
||||
count += 1
|
||||
|
||||
if count > 0:
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
用户偏好设置服务
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -12,7 +11,6 @@ from src.core.logger import logger
|
||||
from src.models.database import Provider, User, UserPreference
|
||||
|
||||
|
||||
|
||||
class PreferenceService:
|
||||
"""用户偏好设置服务"""
|
||||
|
||||
|
||||
@@ -18,7 +18,6 @@ from src.services.cache.user_cache import UserCacheService
|
||||
from src.utils.transaction_manager import retry_on_database_error, transactional
|
||||
|
||||
|
||||
|
||||
class UserService:
|
||||
"""用户管理服务"""
|
||||
|
||||
@@ -198,7 +197,12 @@ class UserService:
|
||||
]
|
||||
|
||||
# 允许设置为 None 的字段(表示无限制)
|
||||
nullable_fields = ["quota_usd", "allowed_providers", "allowed_api_formats", "allowed_models"]
|
||||
nullable_fields = [
|
||||
"quota_usd",
|
||||
"allowed_providers",
|
||||
"allowed_api_formats",
|
||||
"allowed_models",
|
||||
]
|
||||
|
||||
for field, value in kwargs.items():
|
||||
if field not in updatable_fields:
|
||||
@@ -442,7 +446,9 @@ class UserService:
|
||||
# 应用访问限制过滤
|
||||
filtered_models = []
|
||||
for model in all_models:
|
||||
model_name = model.global_model.name if model.global_model else model.provider_model_name
|
||||
model_name = (
|
||||
model.global_model.name if model.global_model else model.provider_model_name
|
||||
)
|
||||
# 使用 AccessRestrictions.is_model_allowed 检查模型是否可访问
|
||||
if restrictions.is_model_allowed(model_name, model.provider_id):
|
||||
filtered_models.append(model)
|
||||
|
||||
Reference in New Issue
Block a user