mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 03:39:49 +08:00
refactor: 共享请求管道、按需懒加载、流式内存护栏与连接池治理
- 抽取 ApiRequestPipeline 单例,44 个路由文件共享同一实例 - Handler/Adapter 模块级 __getattr__ 延迟导入,减少启动时间 - 新增 ensure_stream_buffer_limit() 流式内存护栏(16MB 单行 / 32MB 总量) - HTTP 空闲连接清理与 curl_cffi LRU 会话池 - ensure_providers_bootstrapped 按需引导指定 provider_types - Usage 事件序列化迁移至 msgpack,Redis codec 隔离 - 启动预热任务(/readyz 就绪门控)与优雅关闭 - 通知邮件模块独立开关与 SMTP 配置校验 - CryptoService DCL 线程安全修复 - 通知模块开关 DB 查询 30s 内存缓存 - /readyz 对 unknown 状态返回 503 - 预热关闭 5s 超时保护 - 预热适配器逐个 try-except 容错 - FormatConversionRegistry 哨兵模式防并发重复物化 - 流式缓冲检查无条件执行 Closes #230 Co-authored-by: AAEE86 <[email protected]>
This commit is contained in:
@@ -21,14 +21,14 @@ from sqlalchemy.orm import Session, load_only
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.database import get_db
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
|
||||
router = APIRouter(prefix="/api/admin/adaptive", tags=["Adaptive RPM"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ==================== Pydantic Models ====================
|
||||
|
||||
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config import config
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
@@ -67,7 +67,7 @@ def parse_expiry_date(date_str: str | None) -> datetime | None:
|
||||
|
||||
|
||||
router = APIRouter(prefix="/api/admin/api-keys", tags=["Admin - API Keys (Standalone)"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _serialize_standalone_key_item(api_key: ApiKey) -> dict[str, Any]:
|
||||
|
||||
@@ -19,7 +19,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.database import BillingRule, DimensionCollector
|
||||
@@ -27,7 +27,7 @@ from src.services.billing.formula_engine import SafeExpressionEvaluator, UnsafeE
|
||||
from src.services.billing.presets import BillingPresetService, PresetApplyMode, list_preset_packs
|
||||
|
||||
router = APIRouter(prefix="/api/admin/billing", tags=["Admin - Billing"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
_expr_validator = SafeExpressionEvaluator()
|
||||
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.database import ProviderAPIKey
|
||||
@@ -18,7 +18,7 @@ from src.models.endpoint_models import KeyRpmStatusResponse
|
||||
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
|
||||
|
||||
router = APIRouter(tags=["RPM Control"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/rpm/key/{key_id}", response_model=KeyRpmStatusResponse)
|
||||
|
||||
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
@@ -29,7 +29,7 @@ from src.models.endpoint_models import (
|
||||
HealthSummaryResponse,
|
||||
)
|
||||
from src.services.health.endpoint import EndpointHealthService
|
||||
from src.services.health.monitor import HealthMonitor, health_monitor
|
||||
from src.services.health.monitor import HealthMonitor, get_health_monitor
|
||||
|
||||
router = APIRouter(tags=["Endpoint Health"])
|
||||
|
||||
@@ -39,7 +39,7 @@ def _recover_key_health_sync(db: Session, key_id: str, api_format: str | None) -
|
||||
if not key:
|
||||
raise NotFoundException(f"Key {key_id} 不存在")
|
||||
|
||||
success = health_monitor.reset_health(db, key_id=key_id, api_format=api_format)
|
||||
success = get_health_monitor().reset_health(db, key_id=key_id, api_format=api_format)
|
||||
if not success:
|
||||
raise Exception("重置健康度失败")
|
||||
|
||||
@@ -93,7 +93,7 @@ def _format_str(api_format_enum: Any) -> str:
|
||||
return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
|
||||
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/health/summary", response_model=HealthSummaryResponse)
|
||||
@@ -285,7 +285,7 @@ async def recover_all_keys_health(
|
||||
|
||||
class AdminHealthSummaryAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
summary = health_monitor.get_all_health_status(context.db)
|
||||
summary = get_health_monitor().get_all_health_status(context.db)
|
||||
return HealthSummaryResponse(**summary)
|
||||
|
||||
|
||||
@@ -507,7 +507,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
|
||||
health_data = get_health_monitor().get_key_health(context.db, self.key_id, self.api_format)
|
||||
if not health_data:
|
||||
raise NotFoundException(f"Key {self.key_id} 不存在")
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.database import get_db
|
||||
from src.models.database import User
|
||||
from src.models.endpoint_models import (
|
||||
@@ -41,7 +41,7 @@ from src.services.provider_keys.key_quota_service import (
|
||||
from src.utils.auth_utils import require_admin
|
||||
|
||||
router = APIRouter(tags=["Provider Keys"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.put("/keys/{key_id}", response_model=EndpointAPIKeyResponse)
|
||||
|
||||
@@ -17,7 +17,7 @@ from sqlalchemy.orm.attributes import flag_modified
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.api_format.metadata import get_default_body_rules_for_endpoint
|
||||
from src.core.api_format.signature import parse_signature_key
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
@@ -34,7 +34,7 @@ from src.models.endpoint_models import (
|
||||
from src.services.provider.stream_policy import UpstreamStreamPolicy, parse_upstream_stream_policy
|
||||
|
||||
router = APIRouter(tags=["Endpoint Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def mask_proxy_password(proxy_config: dict | None) -> dict | None:
|
||||
|
||||
@@ -11,7 +11,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.enums import AuthSource
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
@@ -21,7 +21,7 @@ from src.models.database import AuditEventType, LDAPConfig, User, UserRole
|
||||
from src.services.system.audit import AuditService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/ldap", tags=["Admin - LDAP"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
# bcrypt 哈希格式正则:$2a$, $2b$, $2y$ + 2位cost + $ + 53字符(22位salt + 31位hash)
|
||||
BCRYPT_HASH_PATTERN = re.compile(r"^\$2[aby]\$\d{2}\$.{53}$")
|
||||
|
||||
@@ -12,14 +12,14 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.database import AuditEventType, ManagementToken, User
|
||||
from src.services.management_token import ManagementTokenService, token_to_dict
|
||||
|
||||
router = APIRouter(prefix="/api/admin/management-tokens", tags=["Admin - Management Tokens"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ============== 安全基类 ==============
|
||||
|
||||
@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.database import get_db
|
||||
from src.models.database import GlobalModel, Model
|
||||
from src.models.pydantic_models import (
|
||||
@@ -24,7 +24,7 @@ from src.models.pydantic_models import (
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/catalog", tags=["Admin - Model Catalog"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("", response_model=ModelCatalogResponse)
|
||||
|
||||
@@ -15,7 +15,7 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.pydantic_models import (
|
||||
@@ -32,7 +32,7 @@ from src.models.pydantic_models import (
|
||||
from src.services.model.global_model import GlobalModelService
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("", response_model=GlobalModelListResponse)
|
||||
|
||||
@@ -18,7 +18,7 @@ from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.crypto import CryptoService
|
||||
from src.core.model_permissions import (
|
||||
check_model_allowed_with_mappings,
|
||||
@@ -36,7 +36,7 @@ from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ========== Response Models ==========
|
||||
|
||||
@@ -11,13 +11,13 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.modules import ModuleStatus, get_module_registry
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/admin/modules", tags=["Admin - Modules"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ========== Response Models ==========
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pagination import PaginationMeta, build_pagination_payload
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
@@ -31,7 +31,7 @@ from src.utils.cache_decorator import cache_result
|
||||
from src.utils.database_helpers import escape_like_pattern
|
||||
|
||||
router = APIRouter(prefix="/api/admin/monitoring", tags=["Admin - Monitoring"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/audit-logs")
|
||||
|
||||
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pagination import build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
@@ -28,7 +28,7 @@ from src.services.scheduling.aware_scheduler import CacheAwareScheduler, get_cac
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitoring: Cache"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
REDIS_SCAN_BATCH_SIZE = 200
|
||||
REDIS_DELETE_BATCH_SIZE = 500
|
||||
|
||||
|
||||
@@ -15,14 +15,14 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.crypto import crypto_service
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/monitoring/trace", tags=["Admin - Monitoring: Trace"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
class CandidateResponse(BaseModel):
|
||||
|
||||
@@ -12,14 +12,14 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.api.serializers import serialize_payment_callback, serialize_payment_order
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
|
||||
from src.database import get_db, get_db_context
|
||||
from src.services.payment import PaymentService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/payments", tags=["Admin - Payments"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
class AdminPaymentOrderCreditPayload(BaseModel):
|
||||
|
||||
@@ -25,7 +25,7 @@ from sqlalchemy.orm import Session, load_only
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.core.logger import logger
|
||||
@@ -66,7 +66,7 @@ from .schemas import (
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api/admin/pool", tags=["pool-management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
+231
-147
@@ -17,7 +17,7 @@ import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy import update
|
||||
from sqlalchemy.orm import Session, joinedload, make_transient
|
||||
from sqlalchemy.orm import Session, joinedload, selectinload
|
||||
|
||||
from src.api.handlers.base.chat_adapter_base import get_adapter_class
|
||||
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class
|
||||
@@ -395,7 +395,20 @@ async def query_available_models(
|
||||
db.query(Provider)
|
||||
.options(
|
||||
joinedload(Provider.endpoints),
|
||||
joinedload(Provider.api_keys),
|
||||
selectinload(Provider.api_keys)
|
||||
.defer(ProviderAPIKey.note)
|
||||
.defer(ProviderAPIKey.last_error_msg)
|
||||
.defer(ProviderAPIKey.auto_fetch_models)
|
||||
.defer(ProviderAPIKey.locked_models)
|
||||
.defer(ProviderAPIKey.model_include_patterns)
|
||||
.defer(ProviderAPIKey.model_exclude_patterns)
|
||||
.defer(ProviderAPIKey.last_models_fetch_at)
|
||||
.defer(ProviderAPIKey.last_models_fetch_error)
|
||||
.defer(ProviderAPIKey.max_probe_interval_minutes)
|
||||
.defer(ProviderAPIKey.expires_at)
|
||||
.defer(ProviderAPIKey.adjustment_history)
|
||||
.defer(ProviderAPIKey.utilization_samples)
|
||||
.defer(ProviderAPIKey.upstream_metadata),
|
||||
)
|
||||
.filter(Provider.id == request.provider_id)
|
||||
.first()
|
||||
@@ -412,7 +425,7 @@ async def query_available_models(
|
||||
# 延迟导入避免循环依赖(与 upstream_fetcher.fetch_models_for_key 保持一致)
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
ensure_providers_bootstrapped(provider_types=[provider_type] if provider_type else None)
|
||||
has_custom_fetcher = UpstreamModelsFetcherRegistry.get(provider_type) is not None
|
||||
|
||||
if not format_to_endpoint and not has_custom_fetcher:
|
||||
@@ -479,22 +492,30 @@ async def query_available_models(
|
||||
error = f"Key {api_key.name or api_key.id}: {'; '.join(errors)}" if errors else None
|
||||
return unique_models, error, False # models, error, from_cache
|
||||
|
||||
# 并发执行所有 Key 的获取
|
||||
results = await asyncio.gather(*[fetch_for_key(key) for key in active_keys])
|
||||
|
||||
# 合并结果
|
||||
all_models: list = []
|
||||
all_errors: list[str] = []
|
||||
cache_hit_count = 0
|
||||
fetch_count = 0
|
||||
for models, error, from_cache in results:
|
||||
all_models.extend(models)
|
||||
if error:
|
||||
all_errors.append(error)
|
||||
if from_cache:
|
||||
cache_hit_count += 1
|
||||
else:
|
||||
fetch_count += 1
|
||||
|
||||
# 并发执行所有 Key 的获取(增量聚合,减少中间列表峰值)
|
||||
tasks = [asyncio.create_task(fetch_for_key(key)) for key in active_keys]
|
||||
try:
|
||||
for completed in asyncio.as_completed(tasks):
|
||||
models, error, from_cache = await completed
|
||||
all_models.extend(models)
|
||||
if error:
|
||||
all_errors.append(error)
|
||||
if from_cache:
|
||||
cache_hit_count += 1
|
||||
else:
|
||||
fetch_count += 1
|
||||
finally:
|
||||
pending_tasks = [task for task in tasks if not task.done()]
|
||||
for task in pending_tasks:
|
||||
task.cancel()
|
||||
if pending_tasks:
|
||||
await asyncio.gather(*pending_tasks, return_exceptions=True)
|
||||
|
||||
# 按 model id 聚合,合并所有 api_format 到 api_formats 数组
|
||||
unique_models = _aggregate_models_by_id(all_models)
|
||||
@@ -1677,7 +1698,6 @@ async def _run_concurrent_test(
|
||||
from src.services.candidate.recorder import CandidateRecorder
|
||||
from src.services.task.service import pool_on_error
|
||||
|
||||
semaphore = asyncio.Semaphore(max(1, concurrency))
|
||||
record_map = _precreate_concurrent_test_records(
|
||||
db=db,
|
||||
request_id=request_id,
|
||||
@@ -1685,123 +1705,163 @@ async def _run_concurrent_test(
|
||||
user=user,
|
||||
)
|
||||
|
||||
# 预加载所有候选的 provider/endpoint/key,避免每个 worker 重复查询
|
||||
_preloaded: dict[int, tuple[Provider, ProviderEndpoint, ProviderAPIKey]] = {}
|
||||
with create_session() as preload_db:
|
||||
provider_ids = {str(getattr(c.provider, "id", "") or "") for c in candidates}
|
||||
endpoint_ids = {str(getattr(c.endpoint, "id", "") or "") for c in candidates}
|
||||
key_ids = {str(getattr(c.key, "id", "") or "") for c in candidates}
|
||||
providers_by_id = {
|
||||
str(p.id): p
|
||||
for p in preload_db.query(Provider).filter(Provider.id.in_(provider_ids)).all()
|
||||
}
|
||||
endpoints_by_id = {
|
||||
str(e.id): e
|
||||
for e in preload_db.query(ProviderEndpoint)
|
||||
.filter(ProviderEndpoint.id.in_(endpoint_ids))
|
||||
.all()
|
||||
}
|
||||
keys_by_id = {
|
||||
str(k.id): k
|
||||
for k in preload_db.query(ProviderAPIKey).filter(ProviderAPIKey.id.in_(key_ids)).all()
|
||||
}
|
||||
_already_detached: set[int] = set()
|
||||
for idx, cand in enumerate(candidates):
|
||||
p = providers_by_id.get(str(getattr(cand.provider, "id", "") or ""))
|
||||
e = endpoints_by_id.get(str(getattr(cand.endpoint, "id", "") or ""))
|
||||
k = keys_by_id.get(str(getattr(cand.key, "id", "") or ""))
|
||||
if p is not None and e is not None and k is not None:
|
||||
# make_transient 将对象脱离 session 并保留已加载属性,
|
||||
# 避免 expired 状态导致跨协程访问时触发 lazy load 报错。
|
||||
# 同一个对象(多个 candidate 可能共享同一 provider/endpoint)
|
||||
# 只需处理一次。
|
||||
for obj in (p, e, k):
|
||||
obj_id = id(obj)
|
||||
if obj_id not in _already_detached:
|
||||
make_transient(obj)
|
||||
_already_detached.add(obj_id)
|
||||
_preloaded[idx] = (p, e, k)
|
||||
|
||||
success_payload: dict[str, Any] = {}
|
||||
success_event = asyncio.Event()
|
||||
candidate_recorder = CandidateRecorder(db)
|
||||
last_error: Exception | None = None
|
||||
|
||||
candidate_indexes = [
|
||||
candidate_index
|
||||
for candidate_index, candidate in enumerate(candidates)
|
||||
if not bool(getattr(candidate, "is_skipped", False))
|
||||
]
|
||||
candidate_identity_map: dict[int, tuple[str, str, str]] = {}
|
||||
for candidate_index in candidate_indexes:
|
||||
candidate = candidates[candidate_index]
|
||||
candidate_identity_map[candidate_index] = (
|
||||
str(getattr(candidate.provider, "id", "") or ""),
|
||||
str(getattr(candidate.endpoint, "id", "") or ""),
|
||||
str(getattr(candidate.key, "id", "") or ""),
|
||||
)
|
||||
|
||||
preloaded_runtime_objects: dict[int, tuple[Provider, ProviderEndpoint, ProviderAPIKey]] = {}
|
||||
providers_by_id: dict[str, Provider] = {}
|
||||
endpoints_by_id: dict[str, ProviderEndpoint] = {}
|
||||
keys_by_id: dict[str, ProviderAPIKey] = {}
|
||||
provider_ids = {
|
||||
provider_id for provider_id, _, _ in candidate_identity_map.values() if provider_id
|
||||
}
|
||||
endpoint_ids = {
|
||||
endpoint_id for _, endpoint_id, _ in candidate_identity_map.values() if endpoint_id
|
||||
}
|
||||
key_ids = {key_id for _, _, key_id in candidate_identity_map.values() if key_id}
|
||||
|
||||
if candidate_identity_map:
|
||||
with create_session() as preload_db:
|
||||
if provider_ids:
|
||||
providers_by_id = {
|
||||
str(provider.id): provider
|
||||
for provider in preload_db.query(Provider)
|
||||
.filter(Provider.id.in_(provider_ids))
|
||||
.all()
|
||||
}
|
||||
if endpoint_ids:
|
||||
endpoints_by_id = {
|
||||
str(endpoint.id): endpoint
|
||||
for endpoint in preload_db.query(ProviderEndpoint)
|
||||
.filter(ProviderEndpoint.id.in_(endpoint_ids))
|
||||
.all()
|
||||
}
|
||||
if key_ids:
|
||||
keys_by_id = {
|
||||
str(key.id): key
|
||||
for key in preload_db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.id.in_(key_ids))
|
||||
.all()
|
||||
}
|
||||
|
||||
for provider in providers_by_id.values():
|
||||
preload_db.expunge(provider)
|
||||
for endpoint in endpoints_by_id.values():
|
||||
preload_db.expunge(endpoint)
|
||||
for key in keys_by_id.values():
|
||||
preload_db.expunge(key)
|
||||
|
||||
for candidate_index, (provider_id, endpoint_id, key_id) in candidate_identity_map.items():
|
||||
if not provider_id or not endpoint_id or not key_id:
|
||||
continue
|
||||
local_provider = providers_by_id.get(provider_id)
|
||||
local_endpoint = endpoints_by_id.get(endpoint_id)
|
||||
local_key = keys_by_id.get(key_id)
|
||||
if local_provider is None or local_endpoint is None or local_key is None:
|
||||
continue
|
||||
preloaded_runtime_objects[candidate_index] = (local_provider, local_endpoint, local_key)
|
||||
|
||||
def _load_candidate_runtime_objects(
|
||||
candidate_index: int,
|
||||
) -> tuple[Provider, ProviderEndpoint, ProviderAPIKey]:
|
||||
loaded = preloaded_runtime_objects.get(candidate_index)
|
||||
if loaded is None:
|
||||
raise RuntimeError("并发测试目标不存在或已被删除")
|
||||
return loaded
|
||||
|
||||
async def _worker(candidate_index: int) -> dict[str, Any]:
|
||||
nonlocal last_error
|
||||
record_id = record_map[candidate_index]
|
||||
|
||||
started = False
|
||||
started_at = 0.0
|
||||
local_provider: Provider | None = None
|
||||
local_endpoint: ProviderEndpoint | None = None
|
||||
local_key: ProviderAPIKey | None = None
|
||||
|
||||
try:
|
||||
preloaded = _preloaded.get(candidate_index)
|
||||
if preloaded is None:
|
||||
raise RuntimeError("并发测试目标不存在或已被删除")
|
||||
local_provider, local_endpoint, local_key = preloaded
|
||||
if success_event.is_set() or await is_cancelled():
|
||||
_mark_concurrent_test_record_cancelled(record_id)
|
||||
return {"status": "cancelled"}
|
||||
|
||||
async with semaphore:
|
||||
if success_event.is_set() or await is_cancelled():
|
||||
_mark_concurrent_test_record_cancelled(record_id)
|
||||
return {"status": "cancelled"}
|
||||
local_provider, local_endpoint, local_key = _load_candidate_runtime_objects(
|
||||
candidate_index
|
||||
)
|
||||
|
||||
with create_session() as update_db:
|
||||
RequestCandidateService.mark_candidate_started(update_db, record_id)
|
||||
if success_event.is_set() or await is_cancelled():
|
||||
_mark_concurrent_test_record_cancelled(record_id)
|
||||
return {"status": "cancelled"}
|
||||
|
||||
started = True
|
||||
started_at = time.perf_counter()
|
||||
response, auth_type = await _execute_test_check(
|
||||
provider_obj=local_provider,
|
||||
with create_session() as update_db:
|
||||
RequestCandidateService.mark_candidate_started(update_db, record_id)
|
||||
|
||||
started = True
|
||||
started_at = time.perf_counter()
|
||||
response, auth_type = await _execute_test_check(
|
||||
provider_obj=local_provider,
|
||||
endpoint=local_endpoint,
|
||||
key=local_key,
|
||||
effective_model=effective_model_by_candidate_index.get(
|
||||
candidate_index,
|
||||
str(request_payload.get("model", "") or ""),
|
||||
),
|
||||
request_payload=request_payload,
|
||||
request_timeout=request_timeout,
|
||||
provider_type=provider_type,
|
||||
user=user,
|
||||
db=None,
|
||||
)
|
||||
elapsed_ms = max(0, int((time.perf_counter() - started_at) * 1000))
|
||||
|
||||
with create_session() as parse_db:
|
||||
parse_key = (
|
||||
parse_db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.id == str(getattr(local_key, "id", "") or ""))
|
||||
.first()
|
||||
)
|
||||
parsed = _extract_test_response_or_raise(
|
||||
response=response,
|
||||
endpoint=local_endpoint,
|
||||
key=local_key,
|
||||
effective_model=effective_model_by_candidate_index.get(
|
||||
candidate_index,
|
||||
str(request_payload.get("model", "") or ""),
|
||||
),
|
||||
request_payload=request_payload,
|
||||
request_timeout=request_timeout,
|
||||
provider_type=provider_type,
|
||||
user=user,
|
||||
db=None,
|
||||
provider_name=str(local_provider.name),
|
||||
auth_type=auth_type,
|
||||
api_key=parse_key or local_key,
|
||||
db=parse_db,
|
||||
)
|
||||
elapsed_ms = max(0, int((time.perf_counter() - started_at) * 1000))
|
||||
|
||||
with create_session() as parse_db:
|
||||
parse_key = (
|
||||
parse_db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.id == str(getattr(local_key, "id", "") or ""))
|
||||
.first()
|
||||
)
|
||||
parsed = _extract_test_response_or_raise(
|
||||
response=response,
|
||||
endpoint=local_endpoint,
|
||||
provider_name=str(local_provider.name),
|
||||
auth_type=auth_type,
|
||||
api_key=parse_key or local_key,
|
||||
db=parse_db,
|
||||
)
|
||||
with create_session() as update_db:
|
||||
RequestCandidateService.mark_candidate_success(
|
||||
db=update_db,
|
||||
candidate_id=record_id,
|
||||
status_code=200,
|
||||
latency_ms=elapsed_ms,
|
||||
)
|
||||
|
||||
with create_session() as update_db:
|
||||
RequestCandidateService.mark_candidate_success(
|
||||
db=update_db,
|
||||
candidate_id=record_id,
|
||||
status_code=200,
|
||||
latency_ms=elapsed_ms,
|
||||
)
|
||||
|
||||
if not success_event.is_set():
|
||||
success_payload.update(
|
||||
{
|
||||
"response": parsed,
|
||||
"candidate_index": candidate_index,
|
||||
"key_id": str(getattr(local_key, "id", "") or "") or None,
|
||||
}
|
||||
)
|
||||
success_event.set()
|
||||
return {"status": "success"}
|
||||
if not success_event.is_set():
|
||||
success_payload.update(
|
||||
{
|
||||
"response": parsed,
|
||||
"candidate_index": candidate_index,
|
||||
"key_id": str(getattr(local_key, "id", "") or "") or None,
|
||||
}
|
||||
)
|
||||
success_event.set()
|
||||
return {"status": "success"}
|
||||
except asyncio.CancelledError:
|
||||
if started or not success_event.is_set():
|
||||
_mark_concurrent_test_record_cancelled(record_id)
|
||||
@@ -1817,9 +1877,8 @@ async def _run_concurrent_test(
|
||||
elif isinstance(exc, EmbeddedErrorException):
|
||||
status_code = int(exc.error_code or 200)
|
||||
|
||||
loaded = _preloaded.get(candidate_index)
|
||||
if loaded is not None and status_code is not None:
|
||||
await pool_on_error(loaded[0], loaded[2], status_code, exc)
|
||||
if local_provider is not None and local_key is not None and status_code is not None:
|
||||
await pool_on_error(local_provider, local_key, status_code, exc)
|
||||
|
||||
with create_session() as update_db:
|
||||
RequestCandidateService.mark_candidate_failed(
|
||||
@@ -1843,45 +1902,70 @@ async def _run_concurrent_test(
|
||||
await asyncio.sleep(0.1)
|
||||
return False
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(_worker(candidate_index))
|
||||
for candidate_index, candidate in enumerate(candidates)
|
||||
if not bool(getattr(candidate, "is_skipped", False))
|
||||
]
|
||||
candidate_queue: asyncio.Queue[int] = asyncio.Queue()
|
||||
for candidate_index in candidate_indexes:
|
||||
candidate_queue.put_nowait(candidate_index)
|
||||
|
||||
def _drain_candidate_queue() -> None:
|
||||
while True:
|
||||
try:
|
||||
candidate_queue.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
else:
|
||||
candidate_queue.task_done()
|
||||
|
||||
async def _queue_worker() -> None:
|
||||
while not success_event.is_set():
|
||||
if await is_cancelled():
|
||||
return
|
||||
try:
|
||||
candidate_index = candidate_queue.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
return
|
||||
|
||||
try:
|
||||
result = await _worker(candidate_index)
|
||||
if result.get("status") == "success":
|
||||
_drain_candidate_queue()
|
||||
return
|
||||
if await is_cancelled():
|
||||
_drain_candidate_queue()
|
||||
return
|
||||
finally:
|
||||
candidate_queue.task_done()
|
||||
|
||||
worker_count = min(max(1, concurrency), len(candidate_indexes))
|
||||
workers = [asyncio.create_task(_queue_worker()) for _ in range(worker_count)]
|
||||
disconnect_task = asyncio.create_task(_watch_disconnect())
|
||||
pending: set[asyncio.Task[Any]] = set(tasks)
|
||||
pending.add(disconnect_task)
|
||||
queue_done_task = asyncio.create_task(candidate_queue.join())
|
||||
success_wait_task = asyncio.create_task(success_event.wait())
|
||||
|
||||
try:
|
||||
while pending:
|
||||
if pending == {disconnect_task}:
|
||||
disconnect_task.cancel()
|
||||
pending.clear()
|
||||
break
|
||||
|
||||
done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED)
|
||||
if disconnect_task in done and disconnect_task.result() is True:
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
break
|
||||
|
||||
for finished in done:
|
||||
if finished is disconnect_task:
|
||||
continue
|
||||
result = finished.result()
|
||||
if result.get("status") == "success":
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
pending.discard(disconnect_task)
|
||||
disconnect_task.cancel()
|
||||
break
|
||||
if success_event.is_set():
|
||||
break
|
||||
done, _ = await asyncio.wait(
|
||||
{disconnect_task, queue_done_task, success_wait_task},
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if disconnect_task in done and disconnect_task.result() is True:
|
||||
_drain_candidate_queue()
|
||||
elif success_wait_task in done and success_wait_task.result():
|
||||
_drain_candidate_queue()
|
||||
finally:
|
||||
await asyncio.gather(*pending, return_exceptions=True)
|
||||
if not disconnect_task.done():
|
||||
disconnect_task.cancel()
|
||||
await asyncio.gather(disconnect_task, return_exceptions=True)
|
||||
for worker in workers:
|
||||
if not worker.done():
|
||||
worker.cancel()
|
||||
if workers:
|
||||
await asyncio.gather(*workers, return_exceptions=True)
|
||||
|
||||
for task in (disconnect_task, queue_done_task, success_wait_task):
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(
|
||||
disconnect_task,
|
||||
queue_done_task,
|
||||
success_wait_task,
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
if success_event.is_set():
|
||||
_cancel_remaining_concurrent_test_records(request_id)
|
||||
|
||||
@@ -15,7 +15,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
@@ -24,7 +24,7 @@ from src.models.database import Provider
|
||||
from src.models.database_extensions import ProviderUsageTracking
|
||||
|
||||
router = APIRouter(prefix="/api/admin/provider-strategy", tags=["Provider Strategy"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
class ProviderBillingUpdate(BaseModel):
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session, joinedload
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
@@ -40,7 +40,7 @@ from src.models.pydantic_models import (
|
||||
from src.services.model.service import ModelService
|
||||
|
||||
router = APIRouter(tags=["Model Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
|
||||
|
||||
@@ -15,7 +15,7 @@ from sqlalchemy.orm import Session, load_only
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
@@ -35,7 +35,7 @@ from src.utils.cache_decorator import cache_result
|
||||
from .summary import _build_provider_summary
|
||||
|
||||
router = APIRouter(tags=["Provider CRUD"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# 映射预览配置(管理后台功能,限制宽松)
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session, load_only
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
@@ -44,7 +44,7 @@ from src.services.cache.provider_cache import ProviderCacheService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(tags=["Provider Summary"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/summary", response_model=ProviderSummaryPageResponse)
|
||||
|
||||
@@ -15,13 +15,13 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.database import get_db
|
||||
from src.services.proxy_node.service import ProxyNodeService, node_to_dict
|
||||
|
||||
router = APIRouter(prefix="/api/admin/proxy-nodes", tags=["Admin - Proxy Nodes"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -15,13 +15,13 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.adapter import ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.database import get_db
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
|
||||
router = APIRouter(prefix="/api/admin/security/ip", tags=["Admin - Security"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ========== Pydantic 模型 ==========
|
||||
|
||||
@@ -10,12 +10,12 @@ from typing import Any, Literal
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import and_, or_
|
||||
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.settings import config
|
||||
from src.models.database import Usage
|
||||
from src.services.system.time_range import TimeRangeParams
|
||||
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _apply_admin_default_range(
|
||||
|
||||
+175
-137
@@ -15,17 +15,14 @@ from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db, get_db_context
|
||||
from src.models.api import SystemSettingsRequest, SystemSettingsResponse
|
||||
from src.models.database import ApiKey, Provider, Usage, User
|
||||
from src.services.email.email_template import EmailTemplate
|
||||
from src.services.provider_ops.types import SENSITIVE_CREDENTIAL_FIELDS
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.wallet import WalletService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/admin/system", tags=["Admin - System"])
|
||||
@@ -35,6 +32,24 @@ CONFIG_SUPPORTED_VERSIONS = ("2.0", "2.1", "2.2")
|
||||
MAX_IMPORT_SIZE = 10 * 1024 * 1024 # 10MB
|
||||
|
||||
|
||||
def _email_template_service() -> Any:
|
||||
from src.services.email.email_template import EmailTemplate
|
||||
|
||||
return EmailTemplate
|
||||
|
||||
|
||||
def _system_config_service() -> Any:
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
return SystemConfigService
|
||||
|
||||
|
||||
def _wallet_service() -> Any:
|
||||
from src.services.wallet import WalletService
|
||||
|
||||
return WalletService
|
||||
|
||||
|
||||
def _get_version_from_git() -> str | None:
|
||||
"""从 git describe 获取版本号"""
|
||||
import subprocess
|
||||
@@ -254,7 +269,7 @@ async def check_update() -> Any:
|
||||
return _make_empty_response(f"检查更新失败: {str(e)}")
|
||||
|
||||
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/settings")
|
||||
@@ -562,17 +577,17 @@ async def purge_stats(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
class AdminGetSystemSettingsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
default_provider = SystemConfigService.get_default_provider(db)
|
||||
default_model = SystemConfigService.get_config(db, "default_model")
|
||||
default_provider = _system_config_service().get_default_provider(db)
|
||||
default_model = _system_config_service().get_config(db, "default_model")
|
||||
enable_usage_tracking = (
|
||||
SystemConfigService.get_config(db, "enable_usage_tracking", "true") == "true"
|
||||
_system_config_service().get_config(db, "enable_usage_tracking", "true") == "true"
|
||||
)
|
||||
|
||||
return SystemSettingsResponse(
|
||||
default_provider=default_provider,
|
||||
default_model=default_model,
|
||||
enable_usage_tracking=enable_usage_tracking,
|
||||
password_policy_level=SystemConfigService.get_password_policy_level(db),
|
||||
password_policy_level=_system_config_service().get_password_policy_level(db),
|
||||
)
|
||||
|
||||
|
||||
@@ -605,25 +620,27 @@ class AdminUpdateSystemSettingsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
if settings_request.default_provider:
|
||||
SystemConfigService.set_default_provider(db, settings_request.default_provider)
|
||||
_system_config_service().set_default_provider(db, settings_request.default_provider)
|
||||
else:
|
||||
SystemConfigService.delete_config(db, "default_provider")
|
||||
_system_config_service().delete_config(db, "default_provider")
|
||||
|
||||
if settings_request.default_model is not None:
|
||||
if settings_request.default_model:
|
||||
SystemConfigService.set_config(db, "default_model", settings_request.default_model)
|
||||
_system_config_service().set_config(
|
||||
db, "default_model", settings_request.default_model
|
||||
)
|
||||
else:
|
||||
SystemConfigService.delete_config(db, "default_model")
|
||||
_system_config_service().delete_config(db, "default_model")
|
||||
|
||||
if settings_request.enable_usage_tracking is not None:
|
||||
SystemConfigService.set_config(
|
||||
_system_config_service().set_config(
|
||||
db,
|
||||
"enable_usage_tracking",
|
||||
str(settings_request.enable_usage_tracking).lower(),
|
||||
)
|
||||
|
||||
if settings_request.password_policy_level is not None:
|
||||
SystemConfigService.set_config(
|
||||
_system_config_service().set_config(
|
||||
db,
|
||||
"password_policy_level",
|
||||
settings_request.password_policy_level,
|
||||
@@ -634,7 +651,7 @@ class AdminUpdateSystemSettingsAdapter(AdminApiAdapter):
|
||||
|
||||
class AdminGetAllConfigsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
return SystemConfigService.get_all_configs(context.db)
|
||||
return _system_config_service().get_all_configs(context.db)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -645,8 +662,8 @@ class AdminGetSystemConfigAdapter(AdminApiAdapter):
|
||||
SENSITIVE_KEYS = {"smtp_password"}
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
value = SystemConfigService.get_config(context.db, self.key)
|
||||
if value is None and self.key not in SystemConfigService.DEFAULT_CONFIGS:
|
||||
value = _system_config_service().get_config(context.db, self.key)
|
||||
if value is None and self.key not in _system_config_service().DEFAULT_CONFIGS:
|
||||
raise NotFoundException(f"配置项 '{self.key}' 不存在")
|
||||
# 对敏感配置,只返回是否已设置的标志,不返回实际值
|
||||
if self.key in self.SENSITIVE_KEYS:
|
||||
@@ -672,7 +689,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
|
||||
value = crypto_service.encrypt(value)
|
||||
|
||||
try:
|
||||
config = SystemConfigService.set_config(
|
||||
config = _system_config_service().set_config(
|
||||
context.db,
|
||||
self.key,
|
||||
value,
|
||||
@@ -699,12 +716,12 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
|
||||
|
||||
redis_client = get_redis_client_sync()
|
||||
# 从数据库读取两个调度配置的最新值,确保一致性
|
||||
priority_mode = SystemConfigService.get_config(
|
||||
priority_mode = _system_config_service().get_config(
|
||||
context.db,
|
||||
"provider_priority_mode",
|
||||
"provider",
|
||||
)
|
||||
scheduling_mode = SystemConfigService.get_config(
|
||||
scheduling_mode = _system_config_service().get_config(
|
||||
context.db,
|
||||
"scheduling_mode",
|
||||
"cache_affinity",
|
||||
@@ -739,7 +756,7 @@ class AdminDeleteSystemConfigAdapter(AdminApiAdapter):
|
||||
key: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
deleted = SystemConfigService.delete_config(context.db, self.key)
|
||||
deleted = _system_config_service().delete_config(context.db, self.key)
|
||||
if not deleted:
|
||||
raise NotFoundException(f"配置项 '{self.key}' 不存在")
|
||||
return {"message": f"配置项 '{self.key}' 已删除"}
|
||||
@@ -999,16 +1016,10 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
# 预建 global_model_id -> name 映射,避免导出 Model 时 N+1 查询
|
||||
gm_name_map: dict[str, str] = {gm.id: gm.name for gm in global_models}
|
||||
|
||||
# 导出 Providers 及其关联数据
|
||||
providers = (
|
||||
db.query(Provider)
|
||||
.options(
|
||||
selectinload(Provider.endpoints),
|
||||
selectinload(Provider.api_keys),
|
||||
selectinload(Provider.models),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
# 导出 Providers 及其关联数据(分批加载,避免全量 ORM 对象常驻内存)
|
||||
batch_size = 50
|
||||
provider_ids = [provider_id for (provider_id,) in db.query(Provider.id).all()]
|
||||
provider_order = {provider_id: idx for idx, provider_id in enumerate(provider_ids)}
|
||||
providers_data = []
|
||||
|
||||
def _normalize_created_at_for_sort(value: datetime | None) -> datetime:
|
||||
@@ -1016,73 +1027,96 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
return datetime.min.replace(tzinfo=timezone.utc)
|
||||
return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)
|
||||
|
||||
for provider in providers:
|
||||
# 导出 Endpoints
|
||||
endpoints = list(provider.endpoints)
|
||||
endpoints_data = [ep.to_export_dict() for ep in endpoints]
|
||||
provider_endpoint_formats = self._collect_provider_endpoint_formats(endpoints)
|
||||
|
||||
# 导出 Provider Keys(按 provider_id 归属,包含 api_formats)
|
||||
keys = sorted(
|
||||
provider.api_keys,
|
||||
key=lambda key: (
|
||||
key.internal_priority if key.internal_priority is not None else float("inf"),
|
||||
_normalize_created_at_for_sort(key.created_at),
|
||||
),
|
||||
)
|
||||
keys_data = []
|
||||
for key in keys:
|
||||
key_data = key.to_export_dict()
|
||||
key_formats = self._resolve_export_key_api_formats(
|
||||
key_data.get("api_formats"),
|
||||
provider_endpoint_formats,
|
||||
for offset in range(0, len(provider_ids), batch_size):
|
||||
batch_ids = provider_ids[offset : offset + batch_size]
|
||||
providers_batch = (
|
||||
db.query(Provider)
|
||||
.options(
|
||||
selectinload(Provider.endpoints),
|
||||
selectinload(Provider.api_keys),
|
||||
selectinload(Provider.models),
|
||||
)
|
||||
# 保持现有字段名 api_formats,并补充可读别名 supported_endpoints。
|
||||
key_data["api_formats"] = key_formats
|
||||
key_data["supported_endpoints"] = list(key_formats)
|
||||
# 解密 API Key
|
||||
try:
|
||||
key_data["api_key"] = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"API Key 解密失败: provider={}, key_id={}, api_formats={}",
|
||||
provider.name,
|
||||
key.id,
|
||||
key.api_formats,
|
||||
.filter(Provider.id.in_(batch_ids))
|
||||
.all()
|
||||
)
|
||||
providers_batch.sort(key=lambda item: provider_order.get(item.id, 0))
|
||||
|
||||
for provider in providers_batch:
|
||||
# 导出 Endpoints
|
||||
endpoints = list(provider.endpoints)
|
||||
endpoints_data = [ep.to_export_dict() for ep in endpoints]
|
||||
provider_endpoint_formats = self._collect_provider_endpoint_formats(endpoints)
|
||||
|
||||
# 导出 Provider Keys(按 provider_id 归属,包含 api_formats)
|
||||
keys = sorted(
|
||||
provider.api_keys,
|
||||
key=lambda key: (
|
||||
(
|
||||
key.internal_priority
|
||||
if key.internal_priority is not None
|
||||
else float("inf")
|
||||
),
|
||||
_normalize_created_at_for_sort(key.created_at),
|
||||
),
|
||||
)
|
||||
keys_data = []
|
||||
for key in keys:
|
||||
key_data = key.to_export_dict()
|
||||
key_formats = self._resolve_export_key_api_formats(
|
||||
key_data.get("api_formats"),
|
||||
provider_endpoint_formats,
|
||||
)
|
||||
key_data["api_key"] = ""
|
||||
# 解密 auth_config(OAuth 等认证配置)
|
||||
# 导出值为解密后的 JSON 字符串(非 dict),导入时需按字符串重新加密
|
||||
if key.auth_config:
|
||||
# 保持现有字段名 api_formats,并补充可读别名 supported_endpoints。
|
||||
key_data["api_formats"] = key_formats
|
||||
key_data["supported_endpoints"] = list(key_formats)
|
||||
# 解密 API Key
|
||||
try:
|
||||
key_data["auth_config"] = crypto_service.decrypt(key.auth_config)
|
||||
key_data["api_key"] = crypto_service.decrypt(key.api_key)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"auth_config 解密失败: provider={}, key_id={}",
|
||||
"API Key 解密失败: provider={}, key_id={}, api_formats={}",
|
||||
provider.name,
|
||||
key.id,
|
||||
key.api_formats,
|
||||
)
|
||||
pass # 解密失败则不导出 auth_config
|
||||
keys_data.append(key_data)
|
||||
key_data["api_key"] = ""
|
||||
# 解密 auth_config(OAuth 等认证配置)
|
||||
# 导出值为解密后的 JSON 字符串(非 dict),导入时需按字符串重新加密
|
||||
if key.auth_config:
|
||||
try:
|
||||
key_data["auth_config"] = crypto_service.decrypt(key.auth_config)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"auth_config 解密失败: provider={}, key_id={}",
|
||||
provider.name,
|
||||
key.id,
|
||||
)
|
||||
pass # 解密失败则不导出 auth_config
|
||||
keys_data.append(key_data)
|
||||
|
||||
# 导出 Provider Models
|
||||
# 注意:提供商模型(Model)必须关联全局模型(GlobalModel)才能参与路由
|
||||
# 导入时未关联 GlobalModel 的模型会被跳过,这是业务规则而非 bug
|
||||
models = list(provider.models)
|
||||
models_data = []
|
||||
for model in models:
|
||||
model_data = model.to_export_dict()
|
||||
# 追加关联的 GlobalModel 名称(导入时通过名称查找)
|
||||
model_data["global_model_name"] = gm_name_map.get(model.global_model_id)
|
||||
models_data.append(model_data)
|
||||
# 导出 Provider Models
|
||||
# 注意:提供商模型(Model)必须关联全局模型(GlobalModel)才能参与路由
|
||||
# 导入时未关联 GlobalModel 的模型会被跳过,这是业务规则而非 bug
|
||||
models = list(provider.models)
|
||||
models_data = []
|
||||
for model in models:
|
||||
model_data = model.to_export_dict()
|
||||
# 追加关联的 GlobalModel 名称(导入时通过名称查找)
|
||||
model_data["global_model_name"] = gm_name_map.get(model.global_model_id)
|
||||
models_data.append(model_data)
|
||||
|
||||
# 解密 Provider config 中的 credentials
|
||||
provider_data = provider.to_export_dict()
|
||||
provider_data["config"] = self._decrypt_provider_config(provider.config, crypto_service)
|
||||
provider_data["endpoints"] = endpoints_data
|
||||
provider_data["api_keys"] = keys_data
|
||||
provider_data["models"] = models_data
|
||||
providers_data.append(provider_data)
|
||||
# 解密 Provider config 中的 credentials
|
||||
provider_data = provider.to_export_dict()
|
||||
provider_data["config"] = self._decrypt_provider_config(
|
||||
provider.config, crypto_service
|
||||
)
|
||||
provider_data["endpoints"] = endpoints_data
|
||||
provider_data["api_keys"] = keys_data
|
||||
provider_data["models"] = models_data
|
||||
providers_data.append(provider_data)
|
||||
|
||||
# 每批完成后清空会话身份映射,降低导出峰值内存
|
||||
db.expunge_all()
|
||||
|
||||
# 导出 LDAP 配置
|
||||
from src.models.database import LDAPConfig
|
||||
@@ -2146,7 +2180,7 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
|
||||
wallet = None
|
||||
if db is not None and key.is_standalone:
|
||||
wallet = WalletService.get_wallet(db, api_key_id=key.id)
|
||||
wallet = _wallet_service().get_wallet(db, api_key_id=key.id)
|
||||
|
||||
data: dict[str, Any] = {
|
||||
"key_hash": key.key_hash,
|
||||
@@ -2162,7 +2196,7 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
"auto_delete_on_expiry": key.auto_delete_on_expiry,
|
||||
"total_requests": key.total_requests,
|
||||
"total_cost_usd": key.total_cost_usd,
|
||||
"wallet": WalletService.serialize_wallet_summary(wallet) if wallet else None,
|
||||
"wallet": _wallet_service().serialize_wallet_summary(wallet) if wallet else None,
|
||||
}
|
||||
|
||||
if key.key_encrypted:
|
||||
@@ -2188,19 +2222,24 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
|
||||
db = context.db
|
||||
|
||||
# 导出 Users(排除管理员)
|
||||
users = db.query(User).filter(User.is_deleted.is_(False), User.role != UserRole.ADMIN).all()
|
||||
wallet_service = _wallet_service()
|
||||
|
||||
# 导出 Users(排除管理员),预加载非独立余额 Key,避免 N+1
|
||||
users = (
|
||||
db.query(User)
|
||||
.options(selectinload(User.api_keys))
|
||||
.filter(User.is_deleted.is_(False), User.role != UserRole.ADMIN)
|
||||
.all()
|
||||
)
|
||||
wallet_map = wallet_service.get_wallets_by_user_ids(db, [user.id for user in users])
|
||||
users_data = []
|
||||
for user in users:
|
||||
wallet = WalletService.get_wallet(db, user_id=user.id)
|
||||
wallet = wallet_map.get(user.id)
|
||||
# 导出用户的 API Keys(排除独立余额Key,独立Key单独导出)
|
||||
api_keys = (
|
||||
db.query(ApiKey)
|
||||
.filter(ApiKey.user_id == user.id, ApiKey.is_standalone.is_(False))
|
||||
.all()
|
||||
)
|
||||
api_keys_data = [
|
||||
self._serialize_api_key(key, include_is_standalone=True) for key in api_keys
|
||||
self._serialize_api_key(key, include_is_standalone=True)
|
||||
for key in user.api_keys
|
||||
if not key.is_standalone
|
||||
]
|
||||
|
||||
users_data.append(
|
||||
@@ -2214,8 +2253,8 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
"allowed_api_formats": user.allowed_api_formats,
|
||||
"allowed_models": user.allowed_models,
|
||||
"model_capability_settings": user.model_capability_settings,
|
||||
"unlimited": WalletService.is_unlimited_wallet(wallet),
|
||||
"wallet": WalletService.serialize_wallet_summary(wallet) if wallet else None,
|
||||
"unlimited": wallet_service.is_unlimited_wallet(wallet),
|
||||
"wallet": (wallet_service.serialize_wallet_summary(wallet) if wallet else None),
|
||||
"is_active": user.is_active,
|
||||
"api_keys": api_keys_data,
|
||||
}
|
||||
@@ -2379,7 +2418,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
)
|
||||
existing_user.is_active = user_data.get("is_active", True)
|
||||
existing_user.updated_at = datetime.now(timezone.utc)
|
||||
wallet = WalletService.get_or_create_wallet(db, user=existing_user)
|
||||
wallet = _wallet_service().get_or_create_wallet(db, user=existing_user)
|
||||
if wallet is not None:
|
||||
wallet.limit_mode = wallet_limit_mode
|
||||
if wallet_payload:
|
||||
@@ -2415,7 +2454,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
)
|
||||
db.add(new_user)
|
||||
db.flush()
|
||||
wallet = WalletService.get_or_create_wallet(db, user=new_user)
|
||||
wallet = _wallet_service().get_or_create_wallet(db, user=new_user)
|
||||
if wallet is not None:
|
||||
wallet.limit_mode = wallet_limit_mode
|
||||
if wallet_payload:
|
||||
@@ -2454,7 +2493,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
if new_key:
|
||||
db.add(new_key)
|
||||
db.flush()
|
||||
wallet = WalletService.get_or_create_wallet(db, api_key=new_key)
|
||||
wallet = _wallet_service().get_or_create_wallet(db, api_key=new_key)
|
||||
wallet_payload = (
|
||||
key_data.get("wallet")
|
||||
if isinstance(key_data.get("wallet"), dict)
|
||||
@@ -2511,7 +2550,6 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
"""测试 SMTP 连接"""
|
||||
from src.core.crypto import crypto_service
|
||||
from src.services.email.email_sender import EmailSenderService
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
db = context.db
|
||||
payload = context.ensure_json_body() or {}
|
||||
@@ -2519,7 +2557,7 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
# 获取密码:优先使用前端传入的明文密码,否则从数据库获取并解密
|
||||
smtp_password = payload.get("smtp_password")
|
||||
if not smtp_password:
|
||||
encrypted_password = SystemConfigService.get_config(db, "smtp_password")
|
||||
encrypted_password = _system_config_service().get_config(db, "smtp_password")
|
||||
if encrypted_password:
|
||||
try:
|
||||
smtp_password = crypto_service.decrypt(encrypted_password, silent=True)
|
||||
@@ -2530,26 +2568,26 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
# 前端可传入未保存的配置,优先使用前端值,否则回退数据库
|
||||
config = {
|
||||
"smtp_host": payload.get("smtp_host")
|
||||
or SystemConfigService.get_config(db, "smtp_host"),
|
||||
or _system_config_service().get_config(db, "smtp_host"),
|
||||
"smtp_port": payload.get("smtp_port")
|
||||
or SystemConfigService.get_config(db, "smtp_port", default=587),
|
||||
or _system_config_service().get_config(db, "smtp_port", default=587),
|
||||
"smtp_user": payload.get("smtp_user")
|
||||
or SystemConfigService.get_config(db, "smtp_user"),
|
||||
or _system_config_service().get_config(db, "smtp_user"),
|
||||
"smtp_password": smtp_password,
|
||||
"smtp_use_tls": (
|
||||
payload.get("smtp_use_tls")
|
||||
if payload.get("smtp_use_tls") is not None
|
||||
else SystemConfigService.get_config(db, "smtp_use_tls", default=True)
|
||||
else _system_config_service().get_config(db, "smtp_use_tls", default=True)
|
||||
),
|
||||
"smtp_use_ssl": (
|
||||
payload.get("smtp_use_ssl")
|
||||
if payload.get("smtp_use_ssl") is not None
|
||||
else SystemConfigService.get_config(db, "smtp_use_ssl", default=False)
|
||||
else _system_config_service().get_config(db, "smtp_use_ssl", default=False)
|
||||
),
|
||||
"smtp_from_email": payload.get("smtp_from_email")
|
||||
or SystemConfigService.get_config(db, "smtp_from_email"),
|
||||
or _system_config_service().get_config(db, "smtp_from_email"),
|
||||
"smtp_from_name": payload.get("smtp_from_name")
|
||||
or SystemConfigService.get_config(db, "smtp_from_name", default="Aether"),
|
||||
or _system_config_service().get_config(db, "smtp_from_name", default="Aether"),
|
||||
}
|
||||
|
||||
# 验证必要配置
|
||||
@@ -2588,10 +2626,10 @@ class AdminGetEmailTemplatesAdapter(AdminApiAdapter):
|
||||
db = context.db
|
||||
templates = []
|
||||
|
||||
for template_type, type_info in EmailTemplate.TEMPLATE_TYPES.items():
|
||||
for template_type, type_info in _email_template_service().TEMPLATE_TYPES.items():
|
||||
# 获取自定义模板或默认模板
|
||||
template = EmailTemplate.get_template(db, template_type)
|
||||
default_template = EmailTemplate.get_default_template(template_type)
|
||||
template = _email_template_service().get_template(db, template_type)
|
||||
default_template = _email_template_service().get_default_template(template_type)
|
||||
|
||||
# 检查是否使用了自定义模板
|
||||
is_custom = (
|
||||
@@ -2621,13 +2659,13 @@ class AdminGetEmailTemplateAdapter(AdminApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# 验证模板类型
|
||||
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
|
||||
if self.template_type not in _email_template_service().TEMPLATE_TYPES:
|
||||
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
|
||||
|
||||
db = context.db
|
||||
type_info = EmailTemplate.TEMPLATE_TYPES[self.template_type]
|
||||
template = EmailTemplate.get_template(db, self.template_type)
|
||||
default_template = EmailTemplate.get_default_template(self.template_type)
|
||||
type_info = _email_template_service().TEMPLATE_TYPES[self.template_type]
|
||||
template = _email_template_service().get_template(db, self.template_type)
|
||||
default_template = _email_template_service().get_default_template(self.template_type)
|
||||
|
||||
is_custom = (
|
||||
template["subject"] != default_template["subject"]
|
||||
@@ -2654,7 +2692,7 @@ class AdminUpdateEmailTemplateAdapter(AdminApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# 验证模板类型
|
||||
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
|
||||
if self.template_type not in _email_template_service().TEMPLATE_TYPES:
|
||||
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
|
||||
|
||||
db = context.db
|
||||
@@ -2673,16 +2711,16 @@ class AdminUpdateEmailTemplateAdapter(AdminApiAdapter):
|
||||
|
||||
if subject is not None:
|
||||
if subject:
|
||||
SystemConfigService.set_config(db, subject_key, subject)
|
||||
_system_config_service().set_config(db, subject_key, subject)
|
||||
else:
|
||||
# 空字符串表示删除自定义值,恢复默认
|
||||
SystemConfigService.delete_config(db, subject_key)
|
||||
_system_config_service().delete_config(db, subject_key)
|
||||
|
||||
if html is not None:
|
||||
if html:
|
||||
SystemConfigService.set_config(db, html_key, html)
|
||||
_system_config_service().set_config(db, html_key, html)
|
||||
else:
|
||||
SystemConfigService.delete_config(db, html_key)
|
||||
_system_config_service().delete_config(db, html_key)
|
||||
|
||||
return {"message": "模板保存成功"}
|
||||
|
||||
@@ -2695,7 +2733,7 @@ class AdminPreviewEmailTemplateAdapter(AdminApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# 验证模板类型
|
||||
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
|
||||
if self.template_type not in _email_template_service().TEMPLATE_TYPES:
|
||||
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
|
||||
|
||||
db = context.db
|
||||
@@ -2704,17 +2742,17 @@ class AdminPreviewEmailTemplateAdapter(AdminApiAdapter):
|
||||
# 获取模板 HTML(优先使用请求体中的,否则使用数据库中的)
|
||||
html = payload.get("html")
|
||||
if not html:
|
||||
template = EmailTemplate.get_template(db, self.template_type)
|
||||
template = _email_template_service().get_template(db, self.template_type)
|
||||
html = template["html"]
|
||||
|
||||
# 获取预览变量
|
||||
type_info = EmailTemplate.TEMPLATE_TYPES[self.template_type]
|
||||
type_info = _email_template_service().TEMPLATE_TYPES[self.template_type]
|
||||
|
||||
# 构建预览变量,使用请求中的值或默认示例值
|
||||
preview_variables = {}
|
||||
default_values = {
|
||||
"app_name": SystemConfigService.get_config(db, "email_app_name")
|
||||
or SystemConfigService.get_config(db, "smtp_from_name", default="Aether"),
|
||||
"app_name": _system_config_service().get_config(db, "email_app_name")
|
||||
or _system_config_service().get_config(db, "smtp_from_name", default="Aether"),
|
||||
"code": "123456",
|
||||
"expire_minutes": "30",
|
||||
"email": "[email protected]",
|
||||
@@ -2725,7 +2763,7 @@ class AdminPreviewEmailTemplateAdapter(AdminApiAdapter):
|
||||
preview_variables[var] = payload.get(var, default_values.get(var, f"{{{{{var}}}}}"))
|
||||
|
||||
# 渲染模板
|
||||
rendered_html = EmailTemplate.render_template(html, preview_variables)
|
||||
rendered_html = _email_template_service().render_template(html, preview_variables)
|
||||
|
||||
return {
|
||||
"html": rendered_html,
|
||||
@@ -2741,7 +2779,7 @@ class AdminResetEmailTemplateAdapter(AdminApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# 验证模板类型
|
||||
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
|
||||
if self.template_type not in _email_template_service().TEMPLATE_TYPES:
|
||||
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
|
||||
|
||||
db = context.db
|
||||
@@ -2750,12 +2788,12 @@ class AdminResetEmailTemplateAdapter(AdminApiAdapter):
|
||||
subject_key = f"email_template_{self.template_type}_subject"
|
||||
html_key = f"email_template_{self.template_type}_html"
|
||||
|
||||
SystemConfigService.delete_config(db, subject_key)
|
||||
SystemConfigService.delete_config(db, html_key)
|
||||
_system_config_service().delete_config(db, subject_key)
|
||||
_system_config_service().delete_config(db, html_key)
|
||||
|
||||
# 返回默认模板
|
||||
default_template = EmailTemplate.get_default_template(self.template_type)
|
||||
type_info = EmailTemplate.TEMPLATE_TYPES[self.template_type]
|
||||
default_template = _email_template_service().get_default_template(self.template_type)
|
||||
type_info = _email_template_service().TEMPLATE_TYPES[self.template_type]
|
||||
|
||||
return {
|
||||
"message": "模板已重置为默认值",
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session, defer
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
@@ -34,7 +34,7 @@ from src.services.usage.service import UsageService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/admin/usage", tags=["Admin - Usage"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _apply_admin_default_range(
|
||||
|
||||
@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
@@ -29,7 +29,7 @@ from src.services.wallet import WalletService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/admin/users", tags=["Admin - Users"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
class _WalletSentinelType:
|
||||
|
||||
@@ -15,7 +15,7 @@ from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.api.dashboard.routes import DashboardAdapter
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.config.constants import CacheTTL
|
||||
@@ -27,7 +27,7 @@ from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/admin/video-tasks", tags=["Admin - Video Tasks"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("")
|
||||
|
||||
@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.api.serializers import (
|
||||
serialize_admin_wallet,
|
||||
serialize_admin_wallet_refund,
|
||||
@@ -24,7 +24,7 @@ from src.models.database import RefundRequest, Wallet, WalletTransaction
|
||||
from src.services.wallet import WalletService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/wallets", tags=["Admin - Wallets"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
class ManualRechargePayload(BaseModel):
|
||||
|
||||
@@ -13,7 +13,7 @@ from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.database import get_db
|
||||
from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
|
||||
@@ -22,7 +22,7 @@ from src.services.system.announcement import AnnouncementService
|
||||
from src.utils.auth_utils import authenticate_user_from_bearer_token
|
||||
|
||||
router = APIRouter(prefix="/api/announcements", tags=["Announcements"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ============== 公共端点(所有用户可访问) ==============
|
||||
|
||||
@@ -14,7 +14,7 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.core.logger import logger
|
||||
from src.core.validators import PasswordValidator
|
||||
@@ -93,7 +93,7 @@ def validate_email_suffix(db: Session, email: str) -> tuple[bool, str | None]:
|
||||
|
||||
router = APIRouter(prefix="/api/auth", tags=["Authentication"])
|
||||
security = HTTPBearer()
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# API端点
|
||||
|
||||
@@ -684,3 +684,11 @@ class ApiRequestPipeline:
|
||||
except Exception:
|
||||
return str(value)
|
||||
return str(value)
|
||||
|
||||
|
||||
_shared_pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def get_pipeline() -> ApiRequestPipeline:
|
||||
"""返回全局共享的无状态请求管道实例。"""
|
||||
return _shared_pipeline
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.enums import UserRole
|
||||
from src.database import get_db
|
||||
@@ -37,7 +37,7 @@ from src.services.wallet import WalletService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/dashboard", tags=["Dashboard"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def format_tokens(num: int) -> str:
|
||||
@@ -375,7 +375,9 @@ class AdminDashboardStatsAdapter(AdminApiAdapter):
|
||||
total_requests = int(monthly_stats.total_requests or 0) + requests_today
|
||||
error_requests = int(monthly_stats.error_requests or 0) + today_stats["error_requests"]
|
||||
total_cost = float(monthly_stats.total_cost or 0) + float(cost_today)
|
||||
total_actual_cost = float(monthly_stats.actual_total_cost or 0) + float(actual_cost_today)
|
||||
total_actual_cost = float(monthly_stats.actual_total_cost or 0) + float(
|
||||
actual_cost_today
|
||||
)
|
||||
total_tokens = int(monthly_stats.total_tokens or 0) + tokens_today
|
||||
cache_creation_tokens = (
|
||||
int(monthly_stats.cache_creation_tokens or 0) + cache_creation_today
|
||||
|
||||
@@ -15,6 +15,7 @@ from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.utils import (
|
||||
check_html_response,
|
||||
check_prefetched_response_error,
|
||||
ensure_stream_buffer_limit,
|
||||
)
|
||||
from src.config.constants import StreamDefaults
|
||||
from src.config.settings import config
|
||||
@@ -199,12 +200,22 @@ class CliPrefetchMixin:
|
||||
prefetched_chunks.append(first_chunk)
|
||||
total_prefetched_bytes += len(first_chunk)
|
||||
buffer += first_chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=str(provider.name),
|
||||
)
|
||||
|
||||
# 继续读取剩余的预读数据
|
||||
async for chunk in aiter:
|
||||
prefetched_chunks.append(chunk)
|
||||
total_prefetched_bytes += len(chunk)
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=str(provider.name),
|
||||
)
|
||||
|
||||
# 尝试按行解析缓冲区(SSE 格式)
|
||||
while b"\n" in buffer:
|
||||
@@ -381,6 +392,11 @@ class CliPrefetchMixin:
|
||||
# 先处理预读的字节块
|
||||
for chunk in prefetched_chunks:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
@@ -442,6 +458,11 @@ class CliPrefetchMixin:
|
||||
# 继续处理剩余的流数据(使用同一个迭代器)
|
||||
async for chunk in byte_iterator:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
|
||||
@@ -20,11 +20,9 @@ from src.api.handlers.base.base_handler import (
|
||||
from src.api.handlers.base.parsers import get_parser_for_format
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.upstream_stream_bridge import (
|
||||
aggregate_upstream_stream_to_internal_response,
|
||||
)
|
||||
from src.api.handlers.base.utils import (
|
||||
build_sse_headers,
|
||||
ensure_stream_buffer_limit,
|
||||
filter_proxy_response_headers,
|
||||
get_format_converter_registry,
|
||||
resolve_client_content_encoding,
|
||||
@@ -49,7 +47,6 @@ from src.services.provider.transport import build_provider_url
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.utils.sse_parser import SSEEventParser
|
||||
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
||||
|
||||
from .cli_sse_helpers import _format_converted_events_to_sse
|
||||
|
||||
@@ -269,8 +266,6 @@ class CliStreamMixin:
|
||||
client_content_encoding: str | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""执行流式请求并返回流生成器"""
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
|
||||
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
||||
ctx.parsed_chunks = []
|
||||
ctx.provider_parsed_chunks = []
|
||||
@@ -931,6 +926,11 @@ class CliStreamMixin:
|
||||
|
||||
async for chunk in chunk_source:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
|
||||
@@ -32,6 +32,7 @@ from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.utils import (
|
||||
check_html_response,
|
||||
check_prefetched_response_error,
|
||||
ensure_stream_buffer_limit,
|
||||
get_format_converter_registry,
|
||||
)
|
||||
from src.config.constants import StreamDefaults
|
||||
@@ -416,12 +417,22 @@ class StreamProcessor:
|
||||
if kiro_binary_stream:
|
||||
return prefetched_chunks
|
||||
buffer += first_chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=str(provider.name),
|
||||
)
|
||||
|
||||
# 继续读取剩余的预读数据
|
||||
async for chunk in aiter:
|
||||
prefetched_chunks.append(chunk)
|
||||
total_prefetched_bytes += len(chunk)
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=str(provider.name),
|
||||
)
|
||||
|
||||
# 尝试按行解析缓冲区
|
||||
while b"\n" in buffer:
|
||||
@@ -880,6 +891,11 @@ class StreamProcessor:
|
||||
if prefetched_chunks:
|
||||
for chunk in prefetched_chunks:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
try:
|
||||
@@ -916,6 +932,11 @@ class StreamProcessor:
|
||||
|
||||
async for chunk in byte_iterator:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
try:
|
||||
@@ -982,6 +1003,11 @@ class StreamProcessor:
|
||||
yield chunk
|
||||
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
@@ -1006,6 +1032,11 @@ class StreamProcessor:
|
||||
yield chunk
|
||||
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=ctx.provider_name,
|
||||
)
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
@@ -1246,6 +1277,8 @@ class StreamProcessor:
|
||||
async def create_smoothed_stream(
|
||||
self,
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
*,
|
||||
provider_name: str | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建平滑输出的流生成器
|
||||
@@ -1271,6 +1304,11 @@ class StreamProcessor:
|
||||
|
||||
async for chunk in stream_generator:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id,
|
||||
provider_name=provider_name,
|
||||
)
|
||||
|
||||
# 按双换行分割 SSE 事件(标准 SSE 格式)
|
||||
while b"\n\n" in buffer:
|
||||
@@ -1406,6 +1444,9 @@ async def create_smoothed_stream(
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
chunk_size: int = 20,
|
||||
delay_ms: int = 8,
|
||||
*,
|
||||
request_id: str | None = None,
|
||||
provider_name: str | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
独立的平滑流生成函数
|
||||
@@ -1420,7 +1461,12 @@ async def create_smoothed_stream(
|
||||
Yields:
|
||||
平滑处理后的响应数据块
|
||||
"""
|
||||
processor = _LightweightSmoother(chunk_size=chunk_size, delay_ms=delay_ms)
|
||||
processor = _LightweightSmoother(
|
||||
chunk_size=chunk_size,
|
||||
delay_ms=delay_ms,
|
||||
request_id=request_id,
|
||||
provider_name=provider_name,
|
||||
)
|
||||
async for chunk in processor.smooth(stream_generator):
|
||||
yield chunk
|
||||
|
||||
@@ -1432,9 +1478,17 @@ class _LightweightSmoother:
|
||||
只包含平滑输出所需的最小逻辑,不依赖 StreamProcessor 的其他功能。
|
||||
"""
|
||||
|
||||
def __init__(self, chunk_size: int = 20, delay_ms: int = 8) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
chunk_size: int = 20,
|
||||
delay_ms: int = 8,
|
||||
request_id: str | None = None,
|
||||
provider_name: str | None = None,
|
||||
) -> None:
|
||||
self.chunk_size = chunk_size
|
||||
self.delay_ms = delay_ms
|
||||
self.request_id = request_id
|
||||
self.provider_name = provider_name
|
||||
self._extractors: dict[str, ContentExtractor] = {}
|
||||
|
||||
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||
@@ -1468,6 +1522,11 @@ class _LightweightSmoother:
|
||||
|
||||
async for chunk in stream_generator:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=self.request_id or "unknown",
|
||||
provider_name=self.provider_name,
|
||||
)
|
||||
|
||||
while b"\n\n" in buffer:
|
||||
event_block, buffer = buffer.split(b"\n\n", 1)
|
||||
|
||||
@@ -17,7 +17,10 @@ from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.response_parser import ResponseParser
|
||||
from src.api.handlers.base.utils import get_format_converter_registry
|
||||
from src.api.handlers.base.utils import (
|
||||
ensure_stream_buffer_limit,
|
||||
get_format_converter_registry,
|
||||
)
|
||||
from src.core.api_format.conversion.internal import InternalResponse
|
||||
from src.core.api_format.conversion.stream_bridge import InternalStreamAggregator
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
@@ -184,6 +187,11 @@ async def aggregate_upstream_stream_to_internal_response(
|
||||
|
||||
async for chunk in byte_iter:
|
||||
buffer += chunk
|
||||
ensure_stream_buffer_limit(
|
||||
buffer,
|
||||
request_id=str(request_id or ""),
|
||||
provider_name=str(provider_name or "unknown"),
|
||||
)
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastapi.responses import JSONResponse, Response
|
||||
|
||||
from src.config.constants import StreamDefaults
|
||||
from src.core.api_format import filter_response_headers
|
||||
from src.core.api_format.headers import get_header_value
|
||||
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
||||
@@ -143,6 +144,40 @@ def check_html_response(line: str) -> bool:
|
||||
return lower_line.startswith("<!doctype") or lower_line.startswith("<html")
|
||||
|
||||
|
||||
def ensure_stream_buffer_limit(
|
||||
buffer: bytes,
|
||||
*,
|
||||
request_id: str,
|
||||
provider_name: str | None = None,
|
||||
) -> None:
|
||||
"""防止上游单行流式数据过大导致内存失控。"""
|
||||
total_limit = StreamDefaults.MAX_STREAM_BUFFER_TOTAL_BYTES
|
||||
if len(buffer) > total_limit:
|
||||
raise ProviderNotAvailableException(
|
||||
"上游流式响应异常:总缓冲区超过安全上限",
|
||||
provider_name=provider_name,
|
||||
upstream_status=502,
|
||||
upstream_response=(
|
||||
f"stream buffer total overflow: {len(buffer)} bytes > {total_limit}, "
|
||||
f"request_id={request_id}"
|
||||
),
|
||||
)
|
||||
|
||||
limit = StreamDefaults.MAX_STREAM_BUFFER_BYTES
|
||||
if len(buffer) <= limit:
|
||||
return
|
||||
# 允许 chunk 中存在大量完整行;只限制“最后一行未闭合缓冲”的体积。
|
||||
trailing_line = buffer.rsplit(b"\n", 1)[-1]
|
||||
if len(trailing_line) <= limit:
|
||||
return
|
||||
raise ProviderNotAvailableException(
|
||||
"上游流式响应异常:单行数据超过安全上限",
|
||||
provider_name=provider_name,
|
||||
upstream_status=502,
|
||||
upstream_response=f"stream buffer overflow: {len(buffer)} bytes > {limit}, request_id={request_id}",
|
||||
)
|
||||
|
||||
|
||||
def check_prefetched_response_error(
|
||||
prefetched_chunks: list,
|
||||
parser: Any,
|
||||
|
||||
@@ -1,17 +1,26 @@
|
||||
"""
|
||||
Claude Chat API 处理器
|
||||
"""
|
||||
"""Claude handler package (lazy exports)."""
|
||||
|
||||
from src.api.handlers.claude.adapter import (
|
||||
ClaudeChatAdapter,
|
||||
ClaudeTokenCountAdapter,
|
||||
build_claude_adapter,
|
||||
)
|
||||
from src.api.handlers.claude.handler import ClaudeChatHandler
|
||||
from __future__ import annotations
|
||||
|
||||
__all__ = [
|
||||
"ClaudeChatAdapter",
|
||||
"ClaudeTokenCountAdapter",
|
||||
"build_claude_adapter",
|
||||
"ClaudeChatHandler",
|
||||
]
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"ClaudeChatAdapter": (".adapter", "ClaudeChatAdapter"),
|
||||
"ClaudeTokenCountAdapter": (".adapter", "ClaudeTokenCountAdapter"),
|
||||
"build_claude_adapter": (".adapter", "build_claude_adapter"),
|
||||
"ClaudeChatHandler": (".handler", "ClaudeChatHandler"),
|
||||
}
|
||||
|
||||
__all__ = list(_LAZY_EXPORTS.keys())
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name not in _LAZY_EXPORTS:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
module_name, attr_name = _LAZY_EXPORTS[name]
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attr_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
@@ -17,7 +17,6 @@ from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_ad
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.core.api_format import ApiFamily, get_header_value
|
||||
from src.core.logger import logger
|
||||
from src.core.optimization_utils import TokenCounter
|
||||
from src.models.claude import ClaudeMessagesRequest, ClaudeTokenCountRequest
|
||||
|
||||
|
||||
@@ -100,6 +99,53 @@ def _detect_cache_1h_in_body(body: dict[str, Any]) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
_TOKEN_COUNTER_PLUGIN: Any = None
|
||||
|
||||
|
||||
def _get_token_counter() -> Any:
|
||||
global _TOKEN_COUNTER_PLUGIN # noqa: PLW0603
|
||||
if _TOKEN_COUNTER_PLUGIN is None:
|
||||
from src.plugins.token.tiktoken_counter import TiktokenCounterPlugin
|
||||
|
||||
_TOKEN_COUNTER_PLUGIN = TiktokenCounterPlugin(name="tiktoken")
|
||||
return _TOKEN_COUNTER_PLUGIN
|
||||
|
||||
|
||||
async def _count_text_tokens_with_fallback(text: str, model: str) -> int:
|
||||
"""使用 tiktoken 插件计数,失败时回退到轻量估算。"""
|
||||
if not text:
|
||||
return 0
|
||||
try:
|
||||
plugin = _get_token_counter()
|
||||
if plugin.enabled:
|
||||
return await plugin.count_tokens(text, model)
|
||||
except Exception as exc:
|
||||
logger.debug("tiktoken token 计数失败,使用估算回退: {}", exc)
|
||||
# 与旧实现保持一致:按字符估算
|
||||
return max(1, len(text) // 4)
|
||||
|
||||
|
||||
async def _count_messages_tokens_with_fallback(messages: list[dict[str, Any]], model: str) -> int:
|
||||
"""按历史逻辑统计 messages token(每条消息固定开销 + 内容 token)。"""
|
||||
total = 0
|
||||
for message in messages:
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
total += 4 # 角色与分隔符开销
|
||||
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, str):
|
||||
total += await _count_text_tokens_with_fallback(content, model)
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
if isinstance(item, dict):
|
||||
text = item.get("text")
|
||||
if isinstance(text, str):
|
||||
total += await _count_text_tokens_with_fallback(text, model)
|
||||
|
||||
return total
|
||||
|
||||
|
||||
@register_adapter
|
||||
class ClaudeChatAdapter(ChatAdapterBase):
|
||||
"""
|
||||
@@ -249,21 +295,24 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
||||
logger.error(f"Token count payload invalid: {e}")
|
||||
raise HTTPException(status_code=400, detail="Invalid token count payload") from e
|
||||
|
||||
token_counter = TokenCounter()
|
||||
total_tokens = 0
|
||||
|
||||
if request.system:
|
||||
if isinstance(request.system, str):
|
||||
total_tokens += token_counter.count_tokens(request.system, request.model)
|
||||
total_tokens += await _count_text_tokens_with_fallback(
|
||||
request.system, request.model
|
||||
)
|
||||
elif isinstance(request.system, list):
|
||||
for block in request.system:
|
||||
if hasattr(block, "text"):
|
||||
total_tokens += token_counter.count_tokens(block.text, request.model)
|
||||
total_tokens += await _count_text_tokens_with_fallback(
|
||||
block.text, request.model
|
||||
)
|
||||
|
||||
messages_dict = [
|
||||
msg.model_dump() if hasattr(msg, "model_dump") else msg for msg in request.messages
|
||||
]
|
||||
total_tokens += token_counter.count_messages_tokens(messages_dict, request.model)
|
||||
total_tokens += await _count_messages_tokens_with_fallback(messages_dict, request.model)
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="claude_token_count",
|
||||
|
||||
@@ -1,20 +1,28 @@
|
||||
"""
|
||||
Gemini API Handler 模块
|
||||
"""Gemini handler package (lazy exports)."""
|
||||
|
||||
提供 Gemini API 格式的请求处理
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from src.api.handlers.gemini.adapter import GeminiChatAdapter, build_gemini_adapter
|
||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
||||
from src.api.handlers.gemini.video_adapter import GeminiVeoAdapter
|
||||
from src.api.handlers.gemini.video_handler import GeminiVeoHandler
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
__all__ = [
|
||||
"GeminiChatAdapter",
|
||||
"GeminiChatHandler",
|
||||
"GeminiStreamParser",
|
||||
"build_gemini_adapter",
|
||||
"GeminiVeoAdapter",
|
||||
"GeminiVeoHandler",
|
||||
]
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"GeminiChatAdapter": (".adapter", "GeminiChatAdapter"),
|
||||
"build_gemini_adapter": (".adapter", "build_gemini_adapter"),
|
||||
"GeminiChatHandler": (".handler", "GeminiChatHandler"),
|
||||
"GeminiStreamParser": (".stream_parser", "GeminiStreamParser"),
|
||||
"GeminiVeoAdapter": (".video_adapter", "GeminiVeoAdapter"),
|
||||
"GeminiVeoHandler": (".video_handler", "GeminiVeoHandler"),
|
||||
}
|
||||
|
||||
__all__ = list(_LAZY_EXPORTS.keys())
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name not in _LAZY_EXPORTS:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
module_name, attr_name = _LAZY_EXPORTS[name]
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attr_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
@@ -1,12 +1,25 @@
|
||||
"""
|
||||
Gemini CLI 透传处理器
|
||||
"""
|
||||
"""Gemini CLI handler package (lazy exports)."""
|
||||
|
||||
from src.api.handlers.gemini_cli.adapter import GeminiCliAdapter, build_gemini_cli_adapter
|
||||
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
|
||||
from __future__ import annotations
|
||||
|
||||
__all__ = [
|
||||
"GeminiCliAdapter",
|
||||
"GeminiCliMessageHandler",
|
||||
"build_gemini_cli_adapter",
|
||||
]
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"GeminiCliAdapter": (".adapter", "GeminiCliAdapter"),
|
||||
"build_gemini_cli_adapter": (".adapter", "build_gemini_cli_adapter"),
|
||||
"GeminiCliMessageHandler": (".handler", "GeminiCliMessageHandler"),
|
||||
}
|
||||
|
||||
__all__ = list(_LAZY_EXPORTS.keys())
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name not in _LAZY_EXPORTS:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
module_name, attr_name = _LAZY_EXPORTS[name]
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attr_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
@@ -1,15 +1,26 @@
|
||||
"""
|
||||
OpenAI Chat API 处理器
|
||||
"""
|
||||
"""OpenAI handler package (lazy exports)."""
|
||||
|
||||
from src.api.handlers.openai.adapter import OpenAIChatAdapter
|
||||
from src.api.handlers.openai.handler import OpenAIChatHandler
|
||||
from src.api.handlers.openai.video_adapter import OpenAIVideoAdapter
|
||||
from src.api.handlers.openai.video_handler import OpenAIVideoHandler
|
||||
from __future__ import annotations
|
||||
|
||||
__all__ = [
|
||||
"OpenAIChatAdapter",
|
||||
"OpenAIChatHandler",
|
||||
"OpenAIVideoAdapter",
|
||||
"OpenAIVideoHandler",
|
||||
]
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"OpenAIChatAdapter": (".adapter", "OpenAIChatAdapter"),
|
||||
"OpenAIChatHandler": (".handler", "OpenAIChatHandler"),
|
||||
"OpenAIVideoAdapter": (".video_adapter", "OpenAIVideoAdapter"),
|
||||
"OpenAIVideoHandler": (".video_handler", "OpenAIVideoHandler"),
|
||||
}
|
||||
|
||||
__all__ = list(_LAZY_EXPORTS.keys())
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name not in _LAZY_EXPORTS:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
module_name, attr_name = _LAZY_EXPORTS[name]
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attr_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
@@ -1,12 +1,25 @@
|
||||
"""
|
||||
OpenAI CLI 透传处理器
|
||||
"""
|
||||
"""OpenAI CLI handler package (lazy exports)."""
|
||||
|
||||
from src.api.handlers.openai_cli.adapter import OpenAICliAdapter, OpenAICompactAdapter
|
||||
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
|
||||
from __future__ import annotations
|
||||
|
||||
__all__ = [
|
||||
"OpenAICliAdapter",
|
||||
"OpenAICompactAdapter",
|
||||
"OpenAICliMessageHandler",
|
||||
]
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"OpenAICliAdapter": (".adapter", "OpenAICliAdapter"),
|
||||
"OpenAICompactAdapter": (".adapter", "OpenAICompactAdapter"),
|
||||
"OpenAICliMessageHandler": (".handler", "OpenAICliMessageHandler"),
|
||||
}
|
||||
|
||||
__all__ = list(_LAZY_EXPORTS.keys())
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name not in _LAZY_EXPORTS:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
module_name, attr_name = _LAZY_EXPORTS[name]
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attr_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
@@ -12,14 +12,14 @@ from sqlalchemy.orm import Session
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_query
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import ApiKey, AuditLog
|
||||
from src.plugins.manager import get_plugin_manager
|
||||
|
||||
router = APIRouter(prefix="/api/monitoring", tags=["Monitoring"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/my-audit-logs")
|
||||
|
||||
@@ -10,7 +10,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.database import get_db
|
||||
from src.models.database import OAuthProvider
|
||||
@@ -18,7 +18,7 @@ from src.services.auth.oauth.registry import get_oauth_provider_registry
|
||||
from src.services.auth.oauth.service import OAuthService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/oauth", tags=["Admin - OAuth"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
class SupportedOAuthType(BaseModel):
|
||||
|
||||
@@ -9,7 +9,7 @@ from starlette.responses import RedirectResponse
|
||||
from src.api.base.adapter import ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.database import get_db
|
||||
from src.models.database import User
|
||||
@@ -17,7 +17,7 @@ from src.services.auth.oauth.service import OAuthService
|
||||
from src.services.auth.oauth.state import consume_oauth_bind_token, create_oauth_bind_token
|
||||
|
||||
router = APIRouter(prefix="/api/user/oauth", tags=["User - OAuth"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/bindable-providers")
|
||||
|
||||
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session, joinedload, load_only
|
||||
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
@@ -44,7 +44,7 @@ from src.services.system.config import SystemConfigService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/public", tags=["System Catalog"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/site-info")
|
||||
|
||||
@@ -12,15 +12,11 @@ from typing import Any
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.handlers.claude import (
|
||||
ClaudeTokenCountAdapter,
|
||||
build_claude_adapter,
|
||||
)
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(tags=["Claude API"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.post("/v1/messages")
|
||||
@@ -45,6 +41,8 @@ async def create_message(
|
||||
}
|
||||
```
|
||||
"""
|
||||
from src.api.handlers.claude import build_claude_adapter
|
||||
|
||||
adapter = build_claude_adapter(http_request)
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
@@ -67,6 +65,8 @@ async def count_tokens(
|
||||
|
||||
**认证方式**: x-api-key 请求头
|
||||
"""
|
||||
from src.api.handlers.claude import ClaudeTokenCountAdapter
|
||||
|
||||
adapter = ClaudeTokenCountAdapter()
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
|
||||
+16
-12
@@ -15,13 +15,11 @@ from typing import Any
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.handlers.gemini import build_gemini_adapter
|
||||
from src.api.handlers.gemini_cli import build_gemini_cli_adapter
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(tags=["Gemini API"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _is_cli_request(request: Request) -> bool:
|
||||
@@ -46,6 +44,18 @@ def _is_cli_request(request: Request) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _build_adapter_for_request(request: Request) -> Any:
|
||||
"""按请求类型懒加载 Gemini 适配器,降低模块导入开销。"""
|
||||
if _is_cli_request(request):
|
||||
from src.api.handlers.gemini_cli import build_gemini_cli_adapter
|
||||
|
||||
return build_gemini_cli_adapter()
|
||||
|
||||
from src.api.handlers.gemini import build_gemini_adapter
|
||||
|
||||
return build_gemini_adapter()
|
||||
|
||||
|
||||
@router.post("/v1beta/models/{model}:generateContent")
|
||||
async def generate_content(
|
||||
model: str,
|
||||
@@ -72,10 +82,7 @@ async def generate_content(
|
||||
- `model`: 模型名称,如 gemini-2.0-flash
|
||||
"""
|
||||
# 根据 user-agent 或 x-app header 选择适配器
|
||||
if _is_cli_request(http_request):
|
||||
adapter = build_gemini_cli_adapter()
|
||||
else:
|
||||
adapter = build_gemini_adapter()
|
||||
adapter = _build_adapter_for_request(http_request)
|
||||
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
@@ -109,10 +116,7 @@ async def stream_generate_content(
|
||||
注意: Gemini API 通过 URL 端点区分流式/非流式,不需要在请求体中添加 stream 字段
|
||||
"""
|
||||
# 根据 user-agent 或 x-app header 选择适配器
|
||||
if _is_cli_request(http_request):
|
||||
adapter = build_gemini_cli_adapter()
|
||||
else:
|
||||
adapter = build_gemini_adapter()
|
||||
adapter = _build_adapter_for_request(http_request)
|
||||
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
|
||||
@@ -13,13 +13,11 @@ from typing import Any
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.handlers.openai import OpenAIChatAdapter
|
||||
from src.api.handlers.openai_cli import OpenAICliAdapter, OpenAICompactAdapter
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(tags=["OpenAI API"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.post("/v1/chat/completions")
|
||||
@@ -45,6 +43,8 @@ async def create_chat_completion(
|
||||
|
||||
**支持的参数**: model, messages, stream, temperature, max_tokens 等标准 OpenAI 参数
|
||||
"""
|
||||
from src.api.handlers.openai import OpenAIChatAdapter
|
||||
|
||||
adapter = OpenAIChatAdapter()
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
@@ -68,6 +68,8 @@ async def create_responses_compact(
|
||||
|
||||
**认证方式**: Bearer Token(API Key 或 JWT Token)
|
||||
"""
|
||||
from src.api.handlers.openai_cli import OpenAICompactAdapter
|
||||
|
||||
adapter = OpenAICompactAdapter()
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
@@ -90,6 +92,8 @@ async def create_responses(
|
||||
|
||||
**认证方式**: Bearer Token(API Key 或 JWT Token)
|
||||
"""
|
||||
from src.api.handlers.openai_cli import OpenAICliAdapter
|
||||
|
||||
adapter = OpenAICliAdapter()
|
||||
return await pipeline.run(
|
||||
adapter=adapter,
|
||||
|
||||
@@ -9,13 +9,13 @@ from typing import Any
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.api.handlers.gemini.video_adapter import GeminiVeoAdapter
|
||||
from src.api.handlers.openai.video_adapter import OpenAIVideoAdapter
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(tags=["Video Generation"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# -------------------- OpenAI Sora compatible --------------------
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.database import AuditEventType
|
||||
@@ -25,7 +25,7 @@ from src.services.management_token import (
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api/me/management-tokens", tags=["Management Tokens"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ============== 安全基类 ==============
|
||||
|
||||
@@ -14,7 +14,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.enums import UserRole
|
||||
@@ -57,7 +57,7 @@ from src.services.wallet import WalletService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
router = APIRouter(prefix="/api/users/me", tags=["User Profile"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _calculate_token_cache_hit_rate(total_input_context: int, cache_read_tokens: int) -> float:
|
||||
|
||||
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.api.serializers import (
|
||||
safe_gateway_response,
|
||||
serialize_payment_order,
|
||||
@@ -39,7 +39,7 @@ from src.services.payment import PaymentService
|
||||
from src.services.wallet import WalletDailyUsageLedgerService, WalletService
|
||||
|
||||
router = APIRouter(prefix="/api/wallet", tags=["Wallet"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _create_recharge_order_sync(user_id: str, req: CreateRechargePayload) -> dict[str, Any]:
|
||||
|
||||
@@ -20,6 +20,8 @@ Design notes:
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -46,10 +48,21 @@ except ImportError:
|
||||
DEFAULT_IMPERSONATE = "chrome120"
|
||||
|
||||
|
||||
def _get_max_sessions() -> int:
|
||||
raw = os.getenv("CURL_CFFI_MAX_SESSIONS", "20")
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
logger.warning("环境变量 CURL_CFFI_MAX_SESSIONS 非法: {}, 使用默认值 20", raw)
|
||||
return 20
|
||||
return max(1, value)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Session pool (module-level, async-safe)
|
||||
# ---------------------------------------------------------------------------
|
||||
_session_pool: dict[str, AsyncSession] = {}
|
||||
_MAX_SESSIONS = _get_max_sessions()
|
||||
_session_pool: OrderedDict[str, AsyncSession] = OrderedDict()
|
||||
_pool_lock = asyncio.Lock()
|
||||
|
||||
|
||||
@@ -63,9 +76,11 @@ async def _get_or_create_session(
|
||||
) -> AsyncSession:
|
||||
"""Get or create a cached curl_cffi AsyncSession."""
|
||||
key = _session_key(impersonate, proxy)
|
||||
evicted_sessions: list[tuple[str, AsyncSession]] = []
|
||||
async with _pool_lock:
|
||||
session = _session_pool.get(key)
|
||||
if session is not None:
|
||||
_session_pool.move_to_end(key)
|
||||
return session
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
@@ -77,12 +92,22 @@ async def _get_or_create_session(
|
||||
|
||||
session = AsyncSession(**kwargs)
|
||||
_session_pool[key] = session
|
||||
_session_pool.move_to_end(key)
|
||||
while len(_session_pool) > _MAX_SESSIONS:
|
||||
old_key, old_session = _session_pool.popitem(last=False)
|
||||
evicted_sessions.append((old_key, old_session))
|
||||
logger.info(
|
||||
"curl_cffi session created: impersonate={}, proxy={}",
|
||||
impersonate,
|
||||
proxy or "direct",
|
||||
)
|
||||
return session
|
||||
for old_key, old_session in evicted_sessions:
|
||||
try:
|
||||
await old_session.close()
|
||||
logger.debug("curl_cffi session evicted: {}", old_key)
|
||||
except Exception as exc:
|
||||
logger.warning("curl_cffi session close failed during eviction ({}): {}", old_key, exc)
|
||||
return session
|
||||
|
||||
|
||||
async def close_all_sessions() -> None:
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
@@ -19,7 +20,6 @@ import httpx
|
||||
|
||||
from src.config import config
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.fingerprint import KNOWN_IMPERSONATE_PROFILES
|
||||
from src.services.proxy_node.resolver import (
|
||||
build_proxy_url_async,
|
||||
compute_proxy_cache_key,
|
||||
@@ -34,6 +34,19 @@ _proxy_clients_lock = asyncio.Lock()
|
||||
_default_client_lock = asyncio.Lock()
|
||||
|
||||
|
||||
def _get_int_env(name: str, default: int, minimum: int) -> int:
|
||||
"""Read positive integer env value with bounds and fallback."""
|
||||
raw = os.getenv(name)
|
||||
if raw is None:
|
||||
return default
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
logger.warning("环境变量 {} 不是有效整数: {}, 使用默认值 {}", name, raw, default)
|
||||
return default
|
||||
return max(minimum, value)
|
||||
|
||||
|
||||
class HTTPClientPool:
|
||||
"""
|
||||
全局HTTP客户端池单例
|
||||
@@ -268,6 +281,8 @@ class HTTPClientPool:
|
||||
# - direct chrome impersonate profile names (e.g. "chrome124")
|
||||
use_curl_cffi_tls = False
|
||||
if tls_profile_key:
|
||||
from src.services.provider.fingerprint import KNOWN_IMPERSONATE_PROFILES
|
||||
|
||||
use_curl_cffi_tls = (
|
||||
tls_profile_key == "claude_code_nodejs"
|
||||
or tls_profile_key in KNOWN_IMPERSONATE_PROFILES
|
||||
@@ -392,6 +407,78 @@ class HTTPClientPool:
|
||||
|
||||
logger.info("所有HTTP客户端已关闭")
|
||||
|
||||
@classmethod
|
||||
async def cleanup_idle_clients(
|
||||
cls,
|
||||
max_idle_seconds: int | None = None,
|
||||
) -> dict[str, int]:
|
||||
"""清理空闲的代理/Tunnel 客户端并关闭连接池资源。"""
|
||||
idle_seconds = max_idle_seconds
|
||||
if idle_seconds is None:
|
||||
idle_seconds = _get_int_env("HTTP_CLIENT_IDLE_CLEANUP_MAX_SECONDS", 600, minimum=60)
|
||||
|
||||
now = time.time()
|
||||
stale_proxy_clients: list[tuple[str, httpx.AsyncClient]] = []
|
||||
stale_tunnel_clients: list[tuple[str, httpx.AsyncClient]] = []
|
||||
removed_closed_proxy = 0
|
||||
removed_closed_tunnel = 0
|
||||
|
||||
lock = cls._get_proxy_clients_lock()
|
||||
async with lock:
|
||||
for cache_key, (client, last_used) in list(cls._proxy_clients.items()):
|
||||
if client.is_closed:
|
||||
cls._proxy_clients.pop(cache_key, None)
|
||||
removed_closed_proxy += 1
|
||||
continue
|
||||
if now - last_used > idle_seconds:
|
||||
entry = cls._proxy_clients.pop(cache_key, None)
|
||||
if entry is not None:
|
||||
stale_proxy_clients.append((cache_key, entry[0]))
|
||||
|
||||
for node_id, (client, last_used) in list(cls._tunnel_clients.items()):
|
||||
if client.is_closed:
|
||||
cls._tunnel_clients.pop(node_id, None)
|
||||
removed_closed_tunnel += 1
|
||||
continue
|
||||
if now - last_used > idle_seconds:
|
||||
entry = cls._tunnel_clients.pop(node_id, None)
|
||||
if entry is not None:
|
||||
stale_tunnel_clients.append((node_id, entry[0]))
|
||||
|
||||
proxy_closed = 0
|
||||
tunnel_closed = 0
|
||||
for cache_key, client in stale_proxy_clients:
|
||||
try:
|
||||
await client.aclose()
|
||||
proxy_closed += 1
|
||||
except Exception as e:
|
||||
logger.warning("关闭空闲代理客户端失败(key={}): {}", cache_key, e)
|
||||
|
||||
for node_id, client in stale_tunnel_clients:
|
||||
try:
|
||||
await client.aclose()
|
||||
tunnel_closed += 1
|
||||
except Exception as e:
|
||||
logger.warning("关闭空闲 Tunnel 客户端失败(node_id={}): {}", node_id, e)
|
||||
|
||||
if proxy_closed or tunnel_closed or removed_closed_proxy or removed_closed_tunnel:
|
||||
logger.info(
|
||||
"HTTP 客户端空闲清理完成: proxy_closed={}, tunnel_closed={}, "
|
||||
"proxy_already_closed={}, tunnel_already_closed={}, idle_seconds={}",
|
||||
proxy_closed,
|
||||
tunnel_closed,
|
||||
removed_closed_proxy,
|
||||
removed_closed_tunnel,
|
||||
idle_seconds,
|
||||
)
|
||||
|
||||
return {
|
||||
"proxy_closed": proxy_closed,
|
||||
"tunnel_closed": tunnel_closed,
|
||||
"proxy_already_closed": removed_closed_proxy,
|
||||
"tunnel_already_closed": removed_closed_tunnel,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
@asynccontextmanager
|
||||
async def get_temp_client(cls, **kwargs: Any) -> Any:
|
||||
|
||||
+91
-46
@@ -32,26 +32,22 @@ class RedisState(Enum):
|
||||
|
||||
class RedisClientManager:
|
||||
"""
|
||||
Redis客户端管理器(单例)
|
||||
Redis 客户端管理器
|
||||
|
||||
提供 Redis 连接管理、熔断器保护和状态监控。
|
||||
"""
|
||||
|
||||
_instance: RedisClientManager | None = None
|
||||
_redis: aioredis.Redis | None = None
|
||||
|
||||
def __new__(cls) -> "RedisClientManager":
|
||||
"""单例模式"""
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def __init__(self) -> None:
|
||||
# 避免重复初始化
|
||||
if getattr(self, "_initialized", False):
|
||||
return
|
||||
|
||||
self._initialized = True
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
client_name: str,
|
||||
encoding_errors: str = "strict",
|
||||
degraded_warning: str | None = None,
|
||||
) -> None:
|
||||
self._redis: aioredis.Redis | None = None
|
||||
self._client_name = client_name
|
||||
self._encoding_errors = encoding_errors
|
||||
self._degraded_warning = degraded_warning
|
||||
self._circuit_open_until: float | None = None
|
||||
self._consecutive_failures: int = 0
|
||||
self._circuit_threshold = int(os.getenv("REDIS_CIRCUIT_BREAKER_THRESHOLD", "3"))
|
||||
@@ -97,7 +93,7 @@ class RedisClientManager:
|
||||
"""
|
||||
手动重置熔断器(用于管理后台紧急恢复)
|
||||
"""
|
||||
logger.info("Redis 熔断器手动重置")
|
||||
logger.info("{} 熔断器手动重置", self._client_name)
|
||||
self._circuit_open_until = None
|
||||
self._consecutive_failures = 0
|
||||
self._last_error = None
|
||||
@@ -122,7 +118,8 @@ class RedisClientManager:
|
||||
if self._circuit_open_until and time.time() < self._circuit_open_until:
|
||||
remaining = self._circuit_open_until - time.time()
|
||||
logger.warning(
|
||||
"Redis 客户端处于熔断状态,跳过初始化,剩余 {:.1f} 秒 (last_error: {})",
|
||||
"{} 处于熔断状态,跳过初始化,剩余 {:.1f} 秒 (last_error: {})",
|
||||
self._client_name,
|
||||
remaining,
|
||||
self._last_error,
|
||||
)
|
||||
@@ -173,6 +170,7 @@ class RedisClientManager:
|
||||
service_name=sentinel_service,
|
||||
max_connections=redis_max_conn,
|
||||
decode_responses=True,
|
||||
encoding_errors=self._encoding_errors,
|
||||
socket_connect_timeout=5.0,
|
||||
health_check_interval=30, # 每 30 秒检查连接健康状态
|
||||
)
|
||||
@@ -182,6 +180,7 @@ class RedisClientManager:
|
||||
redis_url,
|
||||
encoding="utf-8",
|
||||
decode_responses=True,
|
||||
encoding_errors=self._encoding_errors,
|
||||
socket_timeout=5.0,
|
||||
socket_connect_timeout=5.0,
|
||||
max_connections=redis_max_conn,
|
||||
@@ -191,22 +190,22 @@ class RedisClientManager:
|
||||
|
||||
# 测试连接
|
||||
await self._redis.ping()
|
||||
logger.info(f"[OK] 全局Redis客户端初始化成功: {safe_url}")
|
||||
logger.info("[OK] {} 初始化成功: {}", self._client_name, safe_url)
|
||||
self._consecutive_failures = 0
|
||||
self._circuit_open_until = None
|
||||
return self._redis
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
self._last_error = error_msg
|
||||
logger.error(f"[ERROR] Redis连接失败: {error_msg}")
|
||||
logger.error("[ERROR] {} 连接失败: {}", self._client_name, error_msg)
|
||||
|
||||
self._consecutive_failures += 1
|
||||
if self._consecutive_failures >= self._circuit_threshold:
|
||||
self._circuit_open_until = time.time() + self._circuit_reset_seconds
|
||||
logger.warning(
|
||||
"Redis 初始化连续失败 {} 次,开启熔断 {} 秒。"
|
||||
"熔断期间以下功能将降级: 缓存亲和性、分布式并发控制、RPM限流。"
|
||||
"{} 初始化连续失败 {} 次,开启熔断 {} 秒。"
|
||||
"可通过管理 API /api/admin/system/redis/reset-circuit 手动重置。",
|
||||
self._client_name,
|
||||
self._consecutive_failures,
|
||||
self._circuit_reset_seconds,
|
||||
)
|
||||
@@ -222,12 +221,8 @@ class RedisClientManager:
|
||||
"3. Redis端口(默认6379)是否可访问"
|
||||
) from e
|
||||
|
||||
logger.warning(
|
||||
"[WARN] Redis 不可用,以下功能将降级运行(仅在单实例环境下安全):\n"
|
||||
" - 缓存亲和性: 禁用(每次请求随机选择 Endpoint)\n"
|
||||
" - 分布式并发控制: 降级为本地计数\n"
|
||||
" - RPM 限流: 降级为本地限流"
|
||||
)
|
||||
if self._degraded_warning:
|
||||
logger.warning(self._degraded_warning)
|
||||
self._redis = None
|
||||
return None
|
||||
|
||||
@@ -236,7 +231,7 @@ class RedisClientManager:
|
||||
if self._redis:
|
||||
await self._redis.close()
|
||||
self._redis = None
|
||||
logger.info("全局Redis客户端已关闭")
|
||||
logger.info("{} 已关闭", self._client_name)
|
||||
|
||||
def get_client(self) -> aioredis.Redis | None:
|
||||
"""
|
||||
@@ -250,8 +245,42 @@ class RedisClientManager:
|
||||
return self._redis
|
||||
|
||||
|
||||
# 全局单例
|
||||
_GLOBAL_REDIS_DEGRADED_WARNING = (
|
||||
"[WARN] Redis 不可用,以下功能将降级运行(仅在单实例环境下安全):\n"
|
||||
" - 缓存亲和性: 禁用(每次请求随机选择 Endpoint)\n"
|
||||
" - 分布式并发控制: 降级为本地计数\n"
|
||||
" - RPM 限流: 降级为本地限流"
|
||||
)
|
||||
_USAGE_QUEUE_REDIS_DEGRADED_WARNING = (
|
||||
"[WARN] Usage Queue Redis 不可用,usage queue 写入与消费将暂时不可用"
|
||||
)
|
||||
|
||||
_redis_manager: RedisClientManager | None = None
|
||||
_usage_queue_redis_manager: RedisClientManager | None = None
|
||||
|
||||
|
||||
def _get_global_redis_manager() -> RedisClientManager:
|
||||
global _redis_manager
|
||||
|
||||
if _redis_manager is None:
|
||||
_redis_manager = RedisClientManager(
|
||||
client_name="全局Redis客户端",
|
||||
encoding_errors="strict",
|
||||
degraded_warning=_GLOBAL_REDIS_DEGRADED_WARNING,
|
||||
)
|
||||
return _redis_manager
|
||||
|
||||
|
||||
def _get_usage_queue_redis_manager() -> RedisClientManager:
|
||||
global _usage_queue_redis_manager
|
||||
|
||||
if _usage_queue_redis_manager is None:
|
||||
_usage_queue_redis_manager = RedisClientManager(
|
||||
client_name="Usage Queue Redis客户端",
|
||||
encoding_errors="surrogateescape",
|
||||
degraded_warning=_USAGE_QUEUE_REDIS_DEGRADED_WARNING,
|
||||
)
|
||||
return _usage_queue_redis_manager
|
||||
|
||||
|
||||
async def get_redis_client(require_redis: bool = False) -> aioredis.Redis | None:
|
||||
@@ -267,16 +296,26 @@ async def get_redis_client(require_redis: bool = False) -> aioredis.Redis | None
|
||||
Raises:
|
||||
RuntimeError: 当require_redis=True且连接失败时
|
||||
"""
|
||||
global _redis_manager
|
||||
|
||||
if _redis_manager is None:
|
||||
_redis_manager = RedisClientManager()
|
||||
manager = _get_global_redis_manager()
|
||||
# 如果尚未连接(例如启动时降级、或 close() 后),尝试重新初始化。
|
||||
# initialize() 内部包含熔断器逻辑,避免频繁重试导致抖动。
|
||||
if _redis_manager.get_client() is None:
|
||||
await _redis_manager.initialize(require_redis=require_redis)
|
||||
if manager.get_client() is None:
|
||||
await manager.initialize(require_redis=require_redis)
|
||||
|
||||
return _redis_manager.get_client()
|
||||
return manager.get_client()
|
||||
|
||||
|
||||
async def get_usage_queue_redis_client(require_redis: bool = False) -> aioredis.Redis | None:
|
||||
"""
|
||||
获取 Usage Queue 专用 Redis 客户端。
|
||||
|
||||
与全局 Redis 客户端隔离,专门用于 usage queue 的 msgpack/surrogateescape 编解码链路。
|
||||
"""
|
||||
manager = _get_usage_queue_redis_manager()
|
||||
if manager.get_client() is None:
|
||||
await manager.initialize(require_redis=require_redis)
|
||||
|
||||
return manager.get_client()
|
||||
|
||||
|
||||
def get_redis_client_sync() -> aioredis.Redis | None:
|
||||
@@ -295,11 +334,11 @@ def get_redis_client_sync() -> aioredis.Redis | None:
|
||||
|
||||
|
||||
async def close_redis_client() -> None:
|
||||
"""关闭全局Redis客户端"""
|
||||
global _redis_manager
|
||||
|
||||
"""关闭 Redis 客户端(包含全局客户端和 Usage Queue 专用客户端)"""
|
||||
if _redis_manager:
|
||||
await _redis_manager.close()
|
||||
if _usage_queue_redis_manager:
|
||||
await _usage_queue_redis_manager.close()
|
||||
|
||||
|
||||
def get_redis_state() -> RedisState:
|
||||
@@ -341,13 +380,19 @@ def reset_redis_circuit_breaker() -> bool:
|
||||
"""
|
||||
手动重置 Redis 熔断器(同步方法)
|
||||
|
||||
同时重置全局 Redis 客户端与 Usage Queue 专用客户端,
|
||||
避免其中一方仍停留在熔断状态导致功能未恢复。
|
||||
|
||||
Returns:
|
||||
是否成功重置
|
||||
是否至少重置了一个客户端
|
||||
"""
|
||||
global _redis_manager
|
||||
reset_any = False
|
||||
|
||||
if _redis_manager is None:
|
||||
return False
|
||||
if _redis_manager is not None:
|
||||
_redis_manager.reset_circuit_breaker()
|
||||
reset_any = True
|
||||
if _usage_queue_redis_manager is not None:
|
||||
_usage_queue_redis_manager.reset_circuit_breaker()
|
||||
reset_any = True
|
||||
|
||||
_redis_manager.reset_circuit_breaker()
|
||||
return True
|
||||
return reset_any
|
||||
|
||||
@@ -62,6 +62,14 @@ class StreamDefaults:
|
||||
# 3. 不会占用过多内存
|
||||
MAX_PREFETCH_BYTES = 64 * 1024 # 64KB
|
||||
|
||||
# 单行流式缓冲上限(避免上游长期不换行导致内存无限增长)
|
||||
# 正常 SSE 行通常远小于该值,此值主要作为 OOM 安全护栏。
|
||||
MAX_STREAM_BUFFER_BYTES = 16 * 1024 * 1024 # 16MB
|
||||
|
||||
# 流式总缓冲区硬上限(兜底防护)
|
||||
# 保留单行检查逻辑不变,仅用于防止大量完整行短时间累积占用过高内存。
|
||||
MAX_STREAM_BUFFER_TOTAL_BYTES = 32 * 1024 * 1024 # 32MB
|
||||
|
||||
# 流式转换空产出告警阈值
|
||||
# 连续这么多次空行/非 data 行后记录警告日志
|
||||
# 50 次约等于 50 行非 data SSE 数据,足够覆盖正常事件头
|
||||
|
||||
@@ -324,6 +324,25 @@ class Config:
|
||||
os.getenv("MAINTENANCE_STARTUP_TASKS_ENABLED", "true").lower() == "true"
|
||||
)
|
||||
|
||||
# 启动预热配置(降低懒加载导致的首请求延迟)
|
||||
# STARTUP_WARMUP_ENABLED: 是否启用启动期预热任务(默认 true)
|
||||
# STARTUP_WARMUP_GATE_READINESS: /readyz 是否等待预热完成(默认 true)
|
||||
# STARTUP_WARMUP_PROVIDER_TYPES: 预热时优先 bootstrap 的 provider_type 列表(逗号分隔)
|
||||
self.startup_warmup_enabled = os.getenv("STARTUP_WARMUP_ENABLED", "true").lower() == "true"
|
||||
self.startup_warmup_gate_readiness = (
|
||||
os.getenv("STARTUP_WARMUP_GATE_READINESS", "true").lower() == "true"
|
||||
)
|
||||
warmup_provider_types_env = os.getenv("STARTUP_WARMUP_PROVIDER_TYPES", "").strip()
|
||||
self.startup_warmup_provider_types = (
|
||||
[
|
||||
provider_type.strip()
|
||||
for provider_type in warmup_provider_types_env.split(",")
|
||||
if provider_type.strip()
|
||||
]
|
||||
if warmup_provider_types_env
|
||||
else None
|
||||
)
|
||||
|
||||
# API 文档配置
|
||||
# DOCS_ENABLED: 是否启用 API 文档(/docs, /redoc, /openapi.json)
|
||||
# - 未设置: 开发环境启用,生产环境禁用
|
||||
|
||||
@@ -9,10 +9,14 @@ source -> internal -> target
|
||||
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
|
||||
"""
|
||||
|
||||
import ast
|
||||
import importlib
|
||||
import inspect
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
@@ -43,32 +47,136 @@ def _track_conversion_metrics(
|
||||
)
|
||||
|
||||
|
||||
_MATERIALIZING: tuple[str, str] = ("__materializing__", "")
|
||||
|
||||
|
||||
class FormatConversionRegistry:
|
||||
"""基于 Normalizer 的格式转换注册表"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._normalizers: dict[str, FormatNormalizer] = {}
|
||||
self._lazy_normalizers: dict[str, tuple[str, str]] = {}
|
||||
self._lock = threading.RLock()
|
||||
|
||||
def register(self, normalizer: FormatNormalizer) -> None:
|
||||
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
|
||||
key = str(normalizer.FORMAT_ID).upper()
|
||||
with self._lock:
|
||||
self._normalizers[key] = normalizer
|
||||
self._lazy_normalizers.pop(key, None)
|
||||
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
|
||||
|
||||
def register_lazy(self, format_id: str, module_path: str, class_name: str) -> None:
|
||||
key = str(format_id).upper()
|
||||
with self._lock:
|
||||
if key in self._normalizers:
|
||||
logger.debug(
|
||||
"[FormatConversionRegistry] 跳过 lazy 注册(normalizer 已实例化): {}",
|
||||
key,
|
||||
)
|
||||
return
|
||||
existing = self._lazy_normalizers.get(key)
|
||||
if existing and existing != (module_path, class_name):
|
||||
logger.warning(
|
||||
"[FormatConversionRegistry] FORMAT_ID '{}' 重复 lazy 注册,{}.{}, 将覆盖 {}.{}",
|
||||
key,
|
||||
module_path,
|
||||
class_name,
|
||||
existing[0],
|
||||
existing[1],
|
||||
)
|
||||
self._lazy_normalizers[key] = (module_path, class_name)
|
||||
logger.info(
|
||||
"[FormatConversionRegistry] 注册 lazy normalizer: {} -> {}.{}",
|
||||
key,
|
||||
module_path,
|
||||
class_name,
|
||||
)
|
||||
|
||||
def _materialize_lazy_normalizer(self, key: str) -> FormatNormalizer | None:
|
||||
with self._lock:
|
||||
existing = self._normalizers.get(key)
|
||||
if existing is not None:
|
||||
return existing
|
||||
lazy_spec = self._lazy_normalizers.get(key)
|
||||
if lazy_spec is None or lazy_spec is _MATERIALIZING:
|
||||
return None
|
||||
# 标记为正在加载,防止其他线程重复 materialize
|
||||
self._lazy_normalizers[key] = _MATERIALIZING
|
||||
|
||||
module_path, class_name = lazy_spec
|
||||
try:
|
||||
mod = importlib.import_module(module_path)
|
||||
obj = getattr(mod, class_name, None)
|
||||
if not inspect.isclass(obj) or not issubclass(obj, FormatNormalizer):
|
||||
raise TypeError(f"{module_path}.{class_name} 不是有效的 FormatNormalizer")
|
||||
normalizer = obj()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"[FormatConversionRegistry] lazy 加载 {}.{} 失败: {}",
|
||||
module_path,
|
||||
class_name,
|
||||
e,
|
||||
)
|
||||
# 恢复 lazy_spec 以便后续重试
|
||||
with self._lock:
|
||||
if self._lazy_normalizers.get(key) is _MATERIALIZING:
|
||||
self._lazy_normalizers[key] = lazy_spec
|
||||
return None
|
||||
|
||||
self.register(normalizer)
|
||||
key_upper = str(normalizer.FORMAT_ID).upper()
|
||||
with self._lock:
|
||||
return self._normalizers.get(key) or self._normalizers.get(key_upper)
|
||||
|
||||
def _find_registered_by_data_format_id(self, target_dfid: str) -> FormatNormalizer | None:
|
||||
from src.core.api_format.metadata import get_data_format_id_for_endpoint
|
||||
|
||||
with self._lock:
|
||||
registered_items = list(self._normalizers.items())
|
||||
for reg_key, reg_normalizer in registered_items:
|
||||
if get_data_format_id_for_endpoint(reg_key) == target_dfid:
|
||||
return reg_normalizer
|
||||
return None
|
||||
|
||||
def _find_lazy_key_by_data_format_id(self, target_dfid: str) -> str | None:
|
||||
from src.core.api_format.metadata import get_data_format_id_for_endpoint
|
||||
|
||||
with self._lock:
|
||||
lazy_keys = list(self._lazy_normalizers.keys())
|
||||
for lazy_key in lazy_keys:
|
||||
if get_data_format_id_for_endpoint(lazy_key) == target_dfid:
|
||||
return lazy_key
|
||||
return None
|
||||
|
||||
def get_normalizer(self, format_id: str) -> FormatNormalizer | None:
|
||||
key = str(format_id).upper()
|
||||
# 1. 精确匹配
|
||||
normalizer = self._normalizers.get(key)
|
||||
with self._lock:
|
||||
normalizer = self._normalizers.get(key)
|
||||
if normalizer is not None:
|
||||
return normalizer
|
||||
|
||||
# 2. lazy 精确匹配
|
||||
normalizer = self._materialize_lazy_normalizer(key)
|
||||
if normalizer is not None:
|
||||
return normalizer
|
||||
|
||||
# 2. data_format_id 回退:如 "claude:cli" (dfid="claude") -> ClaudeNormalizer (dfid="claude")
|
||||
from src.core.api_format.metadata import get_data_format_id_for_endpoint
|
||||
|
||||
target_dfid = get_data_format_id_for_endpoint(format_id)
|
||||
if not target_dfid:
|
||||
return None
|
||||
for reg_key, reg_normalizer in self._normalizers.items():
|
||||
reg_dfid = get_data_format_id_for_endpoint(reg_key)
|
||||
if reg_dfid == target_dfid:
|
||||
return reg_normalizer
|
||||
|
||||
# 3. data_format_id 在已实例化 normalizer 中回退
|
||||
normalizer = self._find_registered_by_data_format_id(target_dfid)
|
||||
if normalizer is not None:
|
||||
return normalizer
|
||||
|
||||
# 4. data_format_id 在 lazy normalizer 中回退
|
||||
lazy_key = self._find_lazy_key_by_data_format_id(target_dfid)
|
||||
if lazy_key:
|
||||
return self._materialize_lazy_normalizer(lazy_key)
|
||||
return None
|
||||
|
||||
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
|
||||
@@ -467,13 +575,17 @@ class FormatConversionRegistry:
|
||||
return True
|
||||
|
||||
def list_normalizers(self) -> list[str]:
|
||||
return sorted(self._normalizers.keys())
|
||||
with self._lock:
|
||||
all_keys = set(self._normalizers.keys()) | set(self._lazy_normalizers.keys())
|
||||
return sorted(all_keys)
|
||||
|
||||
def get_supported_targets(self, source_format: str) -> list[str]:
|
||||
src = str(source_format).upper()
|
||||
if src not in self._normalizers:
|
||||
with self._lock:
|
||||
all_keys = set(self._normalizers.keys()) | set(self._lazy_normalizers.keys())
|
||||
if src not in all_keys:
|
||||
return []
|
||||
return [k for k in self._normalizers.keys() if k != src]
|
||||
return [k for k in sorted(all_keys) if k != src]
|
||||
|
||||
|
||||
# 全局注册表(唯一实现)
|
||||
@@ -482,8 +594,91 @@ _DEFAULT_NORMALIZERS_REGISTERED = False
|
||||
_REGISTRATION_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _is_format_normalizer_base(node: ast.expr) -> bool:
|
||||
if isinstance(node, ast.Name):
|
||||
return node.id == "FormatNormalizer"
|
||||
if isinstance(node, ast.Attribute):
|
||||
return node.attr == "FormatNormalizer"
|
||||
return False
|
||||
|
||||
|
||||
def _extract_format_id_literal(class_node: ast.ClassDef) -> str | None:
|
||||
for stmt in class_node.body:
|
||||
if isinstance(stmt, ast.Assign):
|
||||
for target in stmt.targets:
|
||||
if isinstance(target, ast.Name) and target.id == "FORMAT_ID":
|
||||
if isinstance(stmt.value, ast.Constant) and isinstance(stmt.value.value, str):
|
||||
value = stmt.value.value.strip()
|
||||
return value or None
|
||||
elif isinstance(stmt, ast.AnnAssign):
|
||||
target = stmt.target
|
||||
if isinstance(target, ast.Name) and target.id == "FORMAT_ID":
|
||||
value = stmt.value
|
||||
if isinstance(value, ast.Constant) and isinstance(value.value, str):
|
||||
text = value.value.strip()
|
||||
return text or None
|
||||
return None
|
||||
|
||||
|
||||
def _discover_normalizer_specs(normalizers_dir: Path) -> list[tuple[str, str, str]]:
|
||||
specs: list[tuple[str, str, str]] = []
|
||||
|
||||
for py_file in sorted(normalizers_dir.glob("*.py")):
|
||||
if py_file.name.startswith("_"):
|
||||
continue
|
||||
|
||||
module_name = py_file.stem
|
||||
module_path = f"src.core.api_format.conversion.normalizers.{module_name}"
|
||||
module_specs: list[tuple[str, str, str]] = []
|
||||
|
||||
# 优先 AST 发现,避免导入大模块
|
||||
try:
|
||||
source = py_file.read_text(encoding="utf-8")
|
||||
tree = ast.parse(source, filename=str(py_file))
|
||||
for node in tree.body:
|
||||
if not isinstance(node, ast.ClassDef):
|
||||
continue
|
||||
if not any(_is_format_normalizer_base(base) for base in node.bases):
|
||||
continue
|
||||
fmt_id = _extract_format_id_literal(node)
|
||||
if fmt_id:
|
||||
module_specs.append((fmt_id, module_path, node.name))
|
||||
except Exception as e:
|
||||
logger.warning("[FormatConversionRegistry] AST 扫描 {} 失败: {}", module_path, e)
|
||||
|
||||
if module_specs:
|
||||
specs.extend(module_specs)
|
||||
continue
|
||||
|
||||
# AST 无法识别时,回退到反射发现(保持兼容)
|
||||
try:
|
||||
mod = importlib.import_module(module_path)
|
||||
except Exception as e:
|
||||
logger.error("[FormatConversionRegistry] 导入 {} 失败: {}", module_path, e)
|
||||
continue
|
||||
|
||||
for _attr_name, obj in inspect.getmembers(mod, inspect.isclass):
|
||||
if (
|
||||
issubclass(obj, FormatNormalizer)
|
||||
and obj is not FormatNormalizer
|
||||
and hasattr(obj, "FORMAT_ID")
|
||||
and obj.__module__ == mod.__name__
|
||||
):
|
||||
fmt_id = str(getattr(obj, "FORMAT_ID", "")).strip()
|
||||
if fmt_id:
|
||||
module_specs.append((fmt_id, module_path, obj.__name__))
|
||||
|
||||
if not module_specs:
|
||||
logger.warning("[FormatConversionRegistry] 未在 {} 发现可注册 normalizer", module_path)
|
||||
continue
|
||||
|
||||
specs.extend(module_specs)
|
||||
|
||||
return specs
|
||||
|
||||
|
||||
def register_default_normalizers() -> None:
|
||||
"""自动发现并注册 normalizers/ 目录下的所有 FormatNormalizer 实现"""
|
||||
"""自动发现并懒注册 normalizers/ 目录下的所有 FormatNormalizer 实现"""
|
||||
global _DEFAULT_NORMALIZERS_REGISTERED # noqa: PLW0603 - module-level 缓存
|
||||
|
||||
# 快速路径:已注册则直接返回(无锁)
|
||||
@@ -495,43 +690,13 @@ def register_default_normalizers() -> None:
|
||||
if _DEFAULT_NORMALIZERS_REGISTERED:
|
||||
return
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
|
||||
normalizers_dir = Path(__file__).parent / "normalizers"
|
||||
for py_file in sorted(normalizers_dir.glob("*.py")):
|
||||
if py_file.name.startswith("_"):
|
||||
continue
|
||||
module_name = py_file.stem
|
||||
module_path = f"src.core.api_format.conversion.normalizers.{module_name}"
|
||||
try:
|
||||
mod = importlib.import_module(module_path)
|
||||
except Exception as e:
|
||||
logger.error("[FormatConversionRegistry] 导入 {} 失败: {}", module_path, e)
|
||||
continue
|
||||
for _attr_name, obj in inspect.getmembers(mod, inspect.isclass):
|
||||
if (
|
||||
issubclass(obj, FormatNormalizer)
|
||||
and obj is not FormatNormalizer
|
||||
and hasattr(obj, "FORMAT_ID")
|
||||
and obj.__module__ == mod.__name__
|
||||
):
|
||||
fmt_id = str(obj.FORMAT_ID).upper()
|
||||
if format_conversion_registry.get_normalizer(fmt_id) is not None:
|
||||
logger.warning(
|
||||
"[FormatConversionRegistry] FORMAT_ID '{}' 重复注册,{} 将覆盖已有实现",
|
||||
fmt_id,
|
||||
obj.__name__,
|
||||
)
|
||||
try:
|
||||
format_conversion_registry.register(obj())
|
||||
except Exception as e:
|
||||
logger.error("[FormatConversionRegistry] 注册 {} 失败: {}", obj.__name__, e)
|
||||
for fmt_id, module_path, class_name in _discover_normalizer_specs(normalizers_dir):
|
||||
format_conversion_registry.register_lazy(fmt_id, module_path, class_name)
|
||||
|
||||
_DEFAULT_NORMALIZERS_REGISTERED = True
|
||||
logger.info(
|
||||
"[FormatConversionRegistry] 已注册 {} 个 normalizer",
|
||||
"[FormatConversionRegistry] 已懒注册 {} 个 normalizer",
|
||||
len(format_conversion_registry.list_normalizers()),
|
||||
)
|
||||
|
||||
|
||||
+32
-8
@@ -15,10 +15,7 @@ import hashlib
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.utils.perf import PerfRecorder
|
||||
@@ -26,6 +23,9 @@ from src.utils.perf import PerfRecorder
|
||||
from ..config import config
|
||||
from ..core.exceptions import DecryptionException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
|
||||
class CryptoService:
|
||||
"""
|
||||
@@ -36,6 +36,7 @@ class CryptoService:
|
||||
"""
|
||||
|
||||
_instance: CryptoService | None = None
|
||||
_instance_lock = threading.Lock()
|
||||
_cipher: Fernet | None = None
|
||||
_key_source: str = "unknown" # 记录密钥来源,用于调试
|
||||
|
||||
@@ -45,12 +46,17 @@ class CryptoService:
|
||||
|
||||
def __new__(cls) -> CryptoService:
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
cls._instance._initialize()
|
||||
with cls._instance_lock:
|
||||
if cls._instance is None:
|
||||
inst = super().__new__(cls)
|
||||
inst._initialize()
|
||||
cls._instance = inst
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self) -> None:
|
||||
"""初始化加密服务"""
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
logger.info("初始化加密服务")
|
||||
|
||||
encryption_key = config.encryption_key
|
||||
@@ -99,6 +105,10 @@ class CryptoService:
|
||||
Returns:
|
||||
Fernet 兼容的 base64 编码密钥
|
||||
"""
|
||||
from cryptography.fernet import Fernet
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
||||
|
||||
# 首先尝试直接作为 Fernet 密钥使用
|
||||
try:
|
||||
key_bytes = (
|
||||
@@ -235,5 +245,19 @@ class CryptoService:
|
||||
self._decrypt_cache.popitem(last=False)
|
||||
|
||||
|
||||
# 创建全局加密服务实例
|
||||
crypto_service = CryptoService()
|
||||
def get_crypto_service() -> CryptoService:
|
||||
"""获取加密服务单例(首次使用时才会初始化)。"""
|
||||
return CryptoService()
|
||||
|
||||
|
||||
class _LazyCryptoServiceProxy:
|
||||
"""延迟代理,避免 import 阶段触发 cryptography 重载。"""
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(get_crypto_service(), name)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
crypto_service = CryptoService()
|
||||
else:
|
||||
crypto_service = cast(CryptoService, _LazyCryptoServiceProxy())
|
||||
|
||||
@@ -1,106 +0,0 @@
|
||||
"""
|
||||
优化工具类 - 包含Token计数和响应头管理
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import tiktoken
|
||||
|
||||
|
||||
class TokenCounter:
|
||||
"""
|
||||
改进的Token计数器
|
||||
支持多种模型的准确计数
|
||||
"""
|
||||
|
||||
# 模型到编码器的映射
|
||||
MODEL_TO_ENCODING = {
|
||||
"gpt-4": "cl100k_base",
|
||||
"gpt-3.5-turbo": "cl100k_base",
|
||||
"claude-3": "cl100k_base", # Claude使用类似的tokenizer
|
||||
"claude-2": "cl100k_base",
|
||||
}
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._encodings = {}
|
||||
self._default_encoding = None
|
||||
|
||||
def _get_encoding(self, model: str) -> Any:
|
||||
"""获取模型对应的编码器"""
|
||||
# 标准化模型名称
|
||||
model_base = model.lower().split("-")[0]
|
||||
|
||||
if model_base not in self._encodings:
|
||||
encoding_name = self.MODEL_TO_ENCODING.get(model_base, "cl100k_base") # 默认编码器
|
||||
try:
|
||||
self._encodings[model_base] = tiktoken.get_encoding(encoding_name)
|
||||
except Exception:
|
||||
# 如果失败,使用默认编码器
|
||||
if not self._default_encoding:
|
||||
self._default_encoding = tiktoken.get_encoding("cl100k_base")
|
||||
self._encodings[model_base] = self._default_encoding
|
||||
|
||||
return self._encodings[model_base]
|
||||
|
||||
def count_tokens(self, text: str, model: str = "claude-3") -> int:
|
||||
"""
|
||||
精确计算文本的token数量
|
||||
"""
|
||||
if not text:
|
||||
return 0
|
||||
|
||||
try:
|
||||
encoding = self._get_encoding(model)
|
||||
return len(encoding.encode(text))
|
||||
except Exception:
|
||||
# 降级到简单估算
|
||||
return len(text) // 4
|
||||
|
||||
def count_messages_tokens(self, messages: list, model: str = "claude-3") -> int:
|
||||
"""
|
||||
计算消息列表的总token数
|
||||
"""
|
||||
total = 0
|
||||
for message in messages:
|
||||
if isinstance(message, dict):
|
||||
# 计算角色标记
|
||||
total += 4 # 角色和分隔符的开销
|
||||
|
||||
# 计算内容
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, str):
|
||||
total += self.count_tokens(content, model)
|
||||
elif isinstance(content, list):
|
||||
# 处理多模态内容
|
||||
for item in content:
|
||||
if isinstance(item, dict) and "text" in item:
|
||||
total += self.count_tokens(item["text"], model)
|
||||
|
||||
return total
|
||||
|
||||
def estimate_response_tokens(self, response: Any, model: str = "claude-3") -> int:
|
||||
"""
|
||||
估算响应的token数量
|
||||
"""
|
||||
if isinstance(response, dict):
|
||||
# 尝试从响应中提取内容
|
||||
if "content" in response:
|
||||
content = response["content"]
|
||||
if isinstance(content, list):
|
||||
text = " ".join(
|
||||
item.get("text", "") for item in content if isinstance(item, dict)
|
||||
)
|
||||
else:
|
||||
text = str(content)
|
||||
return self.count_tokens(text, model)
|
||||
elif "choices" in response:
|
||||
# OpenAI格式
|
||||
total = 0
|
||||
for choice in response.get("choices", []):
|
||||
message = choice.get("message", {})
|
||||
content = message.get("content", "")
|
||||
total += self.count_tokens(content, model)
|
||||
return total
|
||||
|
||||
# 降级到简单估算
|
||||
return len(str(response)) // 4
|
||||
@@ -1,214 +0,0 @@
|
||||
"""
|
||||
提供商健康度管理
|
||||
基于简单的失败计数和优先级调整
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import Any
|
||||
|
||||
|
||||
class ProviderHealthTracker:
|
||||
"""
|
||||
追踪提供商的健康状态
|
||||
根据失败率动态调整优先级
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
failure_window: int = 300, # 5分钟时间窗口
|
||||
failure_threshold: int = 3, # 3次失败降低优先级
|
||||
recovery_time: int = 600, # 10分钟后重置
|
||||
):
|
||||
self.failure_window = failure_window
|
||||
self.failure_threshold = failure_threshold
|
||||
self.recovery_time = recovery_time
|
||||
|
||||
# 存储每个提供商的失败记录
|
||||
self.failures: dict[str, list] = defaultdict(list)
|
||||
# 存储每个提供商的成功记录
|
||||
self.successes: dict[str, list] = defaultdict(list)
|
||||
# 存储优先级调整
|
||||
self.priority_adjustments: dict[str, int] = {}
|
||||
|
||||
def record_success(self, provider_name: str) -> None:
|
||||
"""记录成功的请求"""
|
||||
current_time = time.time()
|
||||
|
||||
# 记录成功时间
|
||||
self.successes[provider_name].append(current_time)
|
||||
|
||||
# 清理旧记录
|
||||
self._cleanup_old_records(provider_name, current_time)
|
||||
|
||||
# 如果连续成功,可以恢复优先级
|
||||
if len(self.successes[provider_name]) >= 5:
|
||||
if self.priority_adjustments.get(provider_name, 0) < 0:
|
||||
self.priority_adjustments[provider_name] += 1
|
||||
|
||||
def record_failure(self, provider_name: str) -> None:
|
||||
"""记录失败的请求"""
|
||||
current_time = time.time()
|
||||
|
||||
# 记录失败时间
|
||||
self.failures[provider_name].append(current_time)
|
||||
|
||||
# 清理旧记录
|
||||
self._cleanup_old_records(provider_name, current_time)
|
||||
|
||||
# 检查是否需要降低优先级
|
||||
recent_failures = len(self.failures[provider_name])
|
||||
if recent_failures >= self.failure_threshold:
|
||||
# 降低优先级
|
||||
current_adjustment = self.priority_adjustments.get(provider_name, 0)
|
||||
self.priority_adjustments[provider_name] = current_adjustment - 1
|
||||
|
||||
def get_priority_adjustment(self, provider_name: str) -> int:
|
||||
"""
|
||||
获取优先级调整值
|
||||
负数表示降低优先级,正数表示提高优先级
|
||||
"""
|
||||
return self.priority_adjustments.get(provider_name, 0)
|
||||
|
||||
def get_health_status(self, provider_name: str) -> dict:
|
||||
"""
|
||||
获取提供商的健康状态
|
||||
"""
|
||||
current_time = time.time()
|
||||
self._cleanup_old_records(provider_name, current_time)
|
||||
|
||||
recent_failures = len(self.failures[provider_name])
|
||||
recent_successes = len(self.successes[provider_name])
|
||||
total_requests = recent_failures + recent_successes
|
||||
|
||||
failure_rate = recent_failures / total_requests if total_requests > 0 else 0
|
||||
|
||||
return {
|
||||
"provider": provider_name,
|
||||
"recent_failures": recent_failures,
|
||||
"recent_successes": recent_successes,
|
||||
"failure_rate": failure_rate,
|
||||
"priority_adjustment": self.get_priority_adjustment(provider_name),
|
||||
"status": self._get_status_label(failure_rate, recent_failures),
|
||||
}
|
||||
|
||||
def _cleanup_old_records(self, provider_name: str, current_time: float) -> None:
|
||||
"""清理超出时间窗口的记录"""
|
||||
# 清理失败记录
|
||||
self.failures[provider_name] = [
|
||||
t for t in self.failures[provider_name] if current_time - t < self.failure_window
|
||||
]
|
||||
|
||||
# 清理成功记录
|
||||
self.successes[provider_name] = [
|
||||
t for t in self.successes[provider_name] if current_time - t < self.failure_window
|
||||
]
|
||||
|
||||
# 如果很久没有失败,重置优先级调整
|
||||
if not self.failures[provider_name] and self.priority_adjustments.get(provider_name, 0) < 0:
|
||||
# 检查恢复时间
|
||||
if all(current_time - t > self.recovery_time for t in self.successes[provider_name]):
|
||||
self.priority_adjustments[provider_name] = 0
|
||||
|
||||
# 清理已无记录且无优先级调整的 key,防止 dict 无限增长
|
||||
if (
|
||||
not self.failures[provider_name]
|
||||
and not self.successes[provider_name]
|
||||
and self.priority_adjustments.get(provider_name, 0) == 0
|
||||
):
|
||||
self.failures.pop(provider_name, None)
|
||||
self.successes.pop(provider_name, None)
|
||||
self.priority_adjustments.pop(provider_name, None)
|
||||
|
||||
def _get_status_label(self, failure_rate: float, recent_failures: int) -> str:
|
||||
"""根据失败率返回状态标签"""
|
||||
if recent_failures >= self.failure_threshold:
|
||||
return "degraded" # 降级
|
||||
elif failure_rate > 0.5:
|
||||
return "unstable" # 不稳定
|
||||
elif failure_rate > 0.1:
|
||||
return "warning" # 警告
|
||||
else:
|
||||
return "healthy" # 健康
|
||||
|
||||
def should_use_provider(self, provider_name: str) -> bool:
|
||||
"""
|
||||
判断是否应该使用该提供商
|
||||
简单的策略:如果优先级调整低于-3,暂时不使用
|
||||
"""
|
||||
adjustment = self.get_priority_adjustment(provider_name)
|
||||
return adjustment > -3
|
||||
|
||||
def reset_provider_health(self, provider_name: str) -> None:
|
||||
"""重置提供商的健康状态(管理员手动操作)"""
|
||||
self.failures[provider_name] = []
|
||||
self.successes[provider_name] = []
|
||||
self.priority_adjustments[provider_name] = 0
|
||||
|
||||
|
||||
class SimpleProviderSelector:
|
||||
"""
|
||||
简单的提供商选择器
|
||||
基于优先级和健康状态
|
||||
"""
|
||||
|
||||
def __init__(self, health_tracker: ProviderHealthTracker):
|
||||
self.health_tracker = health_tracker
|
||||
|
||||
def select_provider(self, providers: list, specified_provider: str | None = None) -> Any:
|
||||
"""
|
||||
选择提供商
|
||||
|
||||
Args:
|
||||
providers: 可用提供商列表(已按基础优先级排序)
|
||||
specified_provider: 用户指定的提供商
|
||||
|
||||
Returns:
|
||||
选中的提供商
|
||||
"""
|
||||
# 如果用户指定了提供商,直接使用(不管健康状态)
|
||||
if specified_provider:
|
||||
return next((p for p in providers if p.name == specified_provider), None)
|
||||
|
||||
# 否则,根据优先级和健康状态选择
|
||||
# 对提供商列表进行动态排序
|
||||
sorted_providers = sorted(
|
||||
providers,
|
||||
key=lambda p: (
|
||||
p.priority + self.health_tracker.get_priority_adjustment(p.name),
|
||||
-p.id, # 相同优先级时,使用ID作为次要排序
|
||||
),
|
||||
reverse=True, # 优先级高的在前
|
||||
)
|
||||
|
||||
# 选择第一个健康的提供商
|
||||
for provider in sorted_providers:
|
||||
if self.health_tracker.should_use_provider(provider.name):
|
||||
return provider
|
||||
|
||||
# 如果都不健康,还是返回第一个(降级策略)
|
||||
return sorted_providers[0] if sorted_providers else None
|
||||
|
||||
def get_provider_rankings(self, providers: list) -> list:
|
||||
"""
|
||||
获取提供商的当前排名(用于调试和监控)
|
||||
"""
|
||||
rankings = []
|
||||
for provider in providers:
|
||||
health_status = self.health_tracker.get_health_status(provider.name)
|
||||
effective_priority = provider.priority + health_status["priority_adjustment"]
|
||||
|
||||
rankings.append(
|
||||
{
|
||||
"name": provider.name,
|
||||
"base_priority": provider.priority,
|
||||
"adjustment": health_status["priority_adjustment"],
|
||||
"effective_priority": effective_priority,
|
||||
"status": health_status["status"],
|
||||
"failure_rate": health_status["failure_rate"],
|
||||
}
|
||||
)
|
||||
|
||||
# 按有效优先级排序
|
||||
rankings.sort(key=lambda x: x["effective_priority"], reverse=True)
|
||||
return rankings
|
||||
@@ -519,11 +519,13 @@ async def enrich_auth_config(
|
||||
"""Enrich auth_config with non-secret metadata (email/account_id).
|
||||
|
||||
各 provider 的 enrichment 逻辑通过 register_auth_enricher 注册。
|
||||
注意: ensure_providers_bootstrapped() 在应用启动时(main.py lifespan)已显式调用。
|
||||
为支持按需 bootstrap,这里会尝试按 provider_type 触发插件注册。
|
||||
"""
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
pt = normalize_provider_type(provider_type)
|
||||
ensure_providers_bootstrapped(provider_types=[pt] if pt else None)
|
||||
enricher = _auth_enrichers.get(pt)
|
||||
if enricher:
|
||||
return await enricher(auth_config, token_response, access_token, proxy_config)
|
||||
|
||||
+146
-10
@@ -5,12 +5,14 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from src.api.admin import router as admin_router
|
||||
@@ -24,7 +26,7 @@ from src.api.payment import router as payment_router
|
||||
from src.api.public import router as public_router
|
||||
from src.api.user_me import router as me_router
|
||||
from src.api.wallet import router as wallet_router
|
||||
from src.clients.http_client import HTTPClientPool, close_http_clients
|
||||
from src.clients.http_client import close_http_clients
|
||||
|
||||
# 核心模块
|
||||
from src.config import config
|
||||
@@ -106,6 +108,7 @@ class LifecycleState:
|
||||
pool_quota_probe_scheduler: PoolQuotaProbeScheduler | None = None
|
||||
task_poller: TaskPollerService | None = None
|
||||
task_scheduler: TaskScheduler | None = None
|
||||
warmup_task: asyncio.Task[None] | None = None
|
||||
|
||||
|
||||
def _configure_uvicorn_access_log() -> None:
|
||||
@@ -150,9 +153,8 @@ async def _initialize_core_infrastructure(state: LifecycleState) -> None:
|
||||
# 从数据库初始化提供商
|
||||
await initialize_providers()
|
||||
|
||||
# 初始化全局HTTP客户端池
|
||||
logger.info("初始化全局HTTP客户端池...")
|
||||
HTTPClientPool.get_default_client() # 预创建默认客户端
|
||||
# 全局HTTP客户端池按需初始化,避免启动阶段预分配连接池资源
|
||||
logger.info("全局HTTP客户端池采用按需初始化")
|
||||
|
||||
# 初始化全局Redis客户端(可根据配置降级为内存模式)
|
||||
logger.info("初始化全局Redis客户端...")
|
||||
@@ -268,6 +270,96 @@ async def _initialize_plugins_and_modules(app: FastAPI, state: LifecycleState) -
|
||||
register_default_parsers()
|
||||
|
||||
|
||||
def _warmup_lazy_request_dependencies(provider_types: list[str] | None = None) -> int:
|
||||
"""预热懒加载链路,减少首个真实请求的导入与实例化抖动。
|
||||
|
||||
逐个 adapter 独立 try-except,确保单个失败不影响其余组件的预热。
|
||||
"""
|
||||
import importlib
|
||||
|
||||
warmed = 0
|
||||
|
||||
try:
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped(provider_types=provider_types)
|
||||
warmed += 1
|
||||
except Exception as exc:
|
||||
logger.warning("预热 provider bootstrap 失败: {}", exc)
|
||||
|
||||
try:
|
||||
from src.services.health.monitor import get_health_monitor
|
||||
|
||||
get_health_monitor()
|
||||
warmed += 1
|
||||
except Exception as exc:
|
||||
logger.warning("预热 health monitor 失败: {}", exc)
|
||||
|
||||
adapter_specs: list[tuple[str, str]] = [
|
||||
("src.api.handlers.openai", "OpenAIChatAdapter"),
|
||||
("src.api.handlers.openai_cli", "OpenAICliAdapter"),
|
||||
("src.api.handlers.openai_cli", "OpenAICompactAdapter"),
|
||||
("src.api.handlers.claude", "ClaudeChatAdapter"),
|
||||
("src.api.handlers.claude", "ClaudeTokenCountAdapter"),
|
||||
("src.api.handlers.claude_cli", "ClaudeCliAdapter"),
|
||||
("src.api.handlers.gemini", "GeminiChatAdapter"),
|
||||
("src.api.handlers.gemini_cli", "GeminiCliAdapter"),
|
||||
]
|
||||
for module_path, class_name in adapter_specs:
|
||||
try:
|
||||
mod = importlib.import_module(module_path)
|
||||
cls = getattr(mod, class_name)
|
||||
cls()
|
||||
warmed += 1
|
||||
except Exception as exc:
|
||||
logger.warning("预热 {} 失败: {}", class_name, exc)
|
||||
|
||||
return warmed
|
||||
|
||||
|
||||
async def _run_startup_warmup(app: FastAPI) -> None:
|
||||
"""异步执行启动预热,结果写入 app.state 供 /readyz 读取。"""
|
||||
app.state.startup_warmup_started_at = time.monotonic()
|
||||
provider_types = config.startup_warmup_provider_types
|
||||
logger.info(
|
||||
"启动预热任务开始(provider_types={})",
|
||||
provider_types if provider_types else "auto",
|
||||
)
|
||||
|
||||
try:
|
||||
warmed_count = await asyncio.to_thread(
|
||||
_warmup_lazy_request_dependencies,
|
||||
provider_types,
|
||||
)
|
||||
app.state.startup_warmup_status = "ready"
|
||||
elapsed_ms = int((time.monotonic() - app.state.startup_warmup_started_at) * 1000)
|
||||
logger.info("启动预热完成(adapters={}, elapsed_ms={})", warmed_count, elapsed_ms)
|
||||
except Exception as exc:
|
||||
app.state.startup_warmup_status = "failed"
|
||||
app.state.startup_warmup_error = str(exc)
|
||||
logger.exception("启动预热失败,/readyz 将返回 503")
|
||||
finally:
|
||||
app.state.startup_warmup_finished_at = time.monotonic()
|
||||
app.state.startup_warmup_done.set()
|
||||
|
||||
|
||||
def _schedule_startup_warmup(app: FastAPI, state: LifecycleState) -> None:
|
||||
"""初始化预热状态并按配置启动后台预热任务。"""
|
||||
app.state.startup_warmup_done = asyncio.Event()
|
||||
app.state.startup_warmup_status = "pending"
|
||||
app.state.startup_warmup_error = None
|
||||
app.state.startup_warmup_started_at = None
|
||||
app.state.startup_warmup_finished_at = None
|
||||
|
||||
if not config.startup_warmup_enabled:
|
||||
app.state.startup_warmup_status = "disabled"
|
||||
app.state.startup_warmup_done.set()
|
||||
logger.info("启动预热已禁用(STARTUP_WARMUP_ENABLED=false)")
|
||||
return
|
||||
|
||||
state.warmup_task = asyncio.create_task(_run_startup_warmup(app), name="startup_warmup")
|
||||
|
||||
|
||||
async def _start_background_services(state: LifecycleState) -> None:
|
||||
"""启动调度器与后台轮询服务。"""
|
||||
# 启动月卡额度重置调度器(仅一个 worker 执行)
|
||||
@@ -281,16 +373,12 @@ async def _start_background_services(state: LifecycleState) -> None:
|
||||
from src.services.usage.quota_scheduler import get_quota_scheduler
|
||||
from src.utils.task_coordinator import StartupTaskCoordinator
|
||||
|
||||
state.quota_scheduler = get_quota_scheduler()
|
||||
state.maintenance_scheduler = get_maintenance_scheduler()
|
||||
state.model_fetch_scheduler = get_model_fetch_scheduler()
|
||||
state.pool_quota_probe_scheduler = get_pool_quota_probe_scheduler()
|
||||
state.task_poller = get_task_poller()
|
||||
state.task_coordinator = StartupTaskCoordinator(state.redis_client)
|
||||
|
||||
# 启动额度调度器
|
||||
quota_scheduler_active = await state.task_coordinator.acquire("quota_scheduler")
|
||||
if quota_scheduler_active:
|
||||
state.quota_scheduler = get_quota_scheduler()
|
||||
await state.quota_scheduler.start()
|
||||
else:
|
||||
logger.info("检测到其他 worker 已运行额度调度器,本实例跳过")
|
||||
@@ -299,6 +387,7 @@ async def _start_background_services(state: LifecycleState) -> None:
|
||||
# 启动维护调度器
|
||||
maintenance_scheduler_active = await state.task_coordinator.acquire("maintenance_scheduler")
|
||||
if maintenance_scheduler_active:
|
||||
state.maintenance_scheduler = get_maintenance_scheduler()
|
||||
logger.info("启动系统维护调度器...")
|
||||
await state.maintenance_scheduler.start()
|
||||
else:
|
||||
@@ -308,6 +397,7 @@ async def _start_background_services(state: LifecycleState) -> None:
|
||||
# 启动模型自动获取调度器
|
||||
model_fetch_scheduler_active = await state.task_coordinator.acquire("model_fetch_scheduler")
|
||||
if model_fetch_scheduler_active:
|
||||
state.model_fetch_scheduler = get_model_fetch_scheduler()
|
||||
logger.info("启动模型自动获取调度器...")
|
||||
await state.model_fetch_scheduler.start()
|
||||
else:
|
||||
@@ -319,6 +409,7 @@ async def _start_background_services(state: LifecycleState) -> None:
|
||||
"pool_quota_probe_scheduler"
|
||||
)
|
||||
if pool_quota_probe_scheduler_active:
|
||||
state.pool_quota_probe_scheduler = get_pool_quota_probe_scheduler()
|
||||
logger.info("启动号池额度主动探测调度器...")
|
||||
await state.pool_quota_probe_scheduler.start()
|
||||
else:
|
||||
@@ -328,6 +419,7 @@ async def _start_background_services(state: LifecycleState) -> None:
|
||||
# 启动异步任务轮询服务(当前仅视频)
|
||||
task_poller_active = await state.task_coordinator.acquire("task_poller:video")
|
||||
if task_poller_active:
|
||||
state.task_poller = get_task_poller()
|
||||
logger.info("启动 TaskPoller(video)...")
|
||||
await state.task_poller.start()
|
||||
else:
|
||||
@@ -350,6 +442,7 @@ async def _run_startup(app: FastAPI) -> LifecycleState:
|
||||
state = LifecycleState()
|
||||
await _initialize_core_infrastructure(state)
|
||||
await _initialize_plugins_and_modules(app, state)
|
||||
_schedule_startup_warmup(app, state)
|
||||
|
||||
logger.info(f"服务启动成功: http://{config.host}:{config.port}")
|
||||
logger.info("=" * 60)
|
||||
@@ -362,6 +455,15 @@ async def _run_shutdown(state: LifecycleState) -> None:
|
||||
"""执行完整关闭流程。"""
|
||||
logger.info("正在关闭服务...")
|
||||
|
||||
# 停止启动预热任务
|
||||
if state.warmup_task and not state.warmup_task.done():
|
||||
logger.info("停止启动预热任务...")
|
||||
state.warmup_task.cancel()
|
||||
try:
|
||||
await asyncio.wait_for(state.warmup_task, timeout=5.0)
|
||||
except (asyncio.CancelledError, asyncio.TimeoutError):
|
||||
pass
|
||||
|
||||
# 停止 Codex 配额异步同步器(停止前会 flush 待同步事件)
|
||||
logger.info("停止 Codex 配额异步同步器...")
|
||||
from src.services.provider_keys.codex_quota_sync_dispatcher import (
|
||||
@@ -604,6 +706,40 @@ app.include_router(public_router) # 公开API端点(用户可查看提供商
|
||||
app.include_router(monitoring_router) # 监控端点
|
||||
|
||||
|
||||
@app.get("/readyz", include_in_schema=False)
|
||||
async def readiness_check(request: Request) -> Any:
|
||||
"""就绪检查:可选绑定启动预热状态,用于流量门禁。"""
|
||||
warmup_status = getattr(request.app.state, "startup_warmup_status", "unknown")
|
||||
warmup_error = getattr(request.app.state, "startup_warmup_error", None)
|
||||
started_at = getattr(request.app.state, "startup_warmup_started_at", None)
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"status": "ready",
|
||||
"warmup_status": warmup_status,
|
||||
"gate_readiness": config.startup_warmup_gate_readiness,
|
||||
}
|
||||
if isinstance(started_at, (int, float)):
|
||||
payload["warmup_elapsed_ms"] = int((time.monotonic() - started_at) * 1000)
|
||||
|
||||
if not config.startup_warmup_gate_readiness:
|
||||
return payload
|
||||
if warmup_status in {"ready", "disabled"}:
|
||||
return payload
|
||||
|
||||
reason_map = {
|
||||
"failed": "startup_warmup_failed",
|
||||
"pending": "startup_warmup_pending",
|
||||
}
|
||||
detail = {
|
||||
"status": "not_ready",
|
||||
"reason": reason_map.get(warmup_status, "startup_warmup_pending"),
|
||||
"warmup_status": warmup_status,
|
||||
}
|
||||
if warmup_error:
|
||||
detail["warmup_error"] = warmup_error
|
||||
raise HTTPException(status_code=503, detail=detail)
|
||||
|
||||
|
||||
def main() -> Any:
|
||||
# 初始化新日志系统
|
||||
debug_mode = config.environment == "development"
|
||||
|
||||
@@ -41,9 +41,13 @@ class PluginMiddleware:
|
||||
- 纯 ASGI 可以直接透传流式响应,无额外开销
|
||||
"""
|
||||
|
||||
_NOTIFICATION_CACHE_TTL = 30.0 # 通知模块开关缓存 TTL (秒)
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
self.plugin_manager = get_plugin_manager()
|
||||
self._notification_enabled_cache: bool | None = None
|
||||
self._notification_cache_expires: float = 0.0
|
||||
|
||||
# 从配置读取速率限制值
|
||||
self.llm_api_rate_limit = config.llm_api_rate_limit
|
||||
@@ -52,6 +56,7 @@ class PluginMiddleware:
|
||||
# 完全跳过限流的路径(静态资源、文档等)
|
||||
self.skip_rate_limit_paths = [
|
||||
"/health",
|
||||
"/readyz",
|
||||
"/docs",
|
||||
"/redoc",
|
||||
"/openapi.json",
|
||||
@@ -465,6 +470,36 @@ class PluginMiddleware:
|
||||
except Exception as e:
|
||||
logger.error(f"Monitor plugin failed: {e}")
|
||||
|
||||
async def _is_notification_email_module_enabled(self, request: Request) -> bool:
|
||||
"""检查通知邮件模块是否启用(带内存缓存,避免 5xx 雪崩时放大 DB 压力)。"""
|
||||
now = time.time()
|
||||
if self._notification_enabled_cache is not None and now < self._notification_cache_expires:
|
||||
return self._notification_enabled_cache
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.database import create_session
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
config_key = "module.notification_email.enabled"
|
||||
try:
|
||||
request_db = getattr(request.state, "db", None)
|
||||
if isinstance(request_db, Session):
|
||||
result = bool(SystemConfigService.get_config(request_db, config_key, default=False))
|
||||
else:
|
||||
db = create_session()
|
||||
try:
|
||||
result = bool(SystemConfigService.get_config(db, config_key, default=False))
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as e:
|
||||
logger.warning("读取通知邮件模块开关失败: {}", e)
|
||||
return False
|
||||
|
||||
self._notification_enabled_cache = result
|
||||
self._notification_cache_expires = now + self._NOTIFICATION_CACHE_TTL
|
||||
return result
|
||||
|
||||
async def _call_error_plugins(
|
||||
self, request: Request, error: Exception, start_time: float
|
||||
) -> None:
|
||||
@@ -475,6 +510,9 @@ class PluginMiddleware:
|
||||
|
||||
# 通知插件 - 发送严重错误通知
|
||||
if not isinstance(error, HTTPException) or error.status_code >= 500:
|
||||
if not await self._is_notification_email_module_enabled(request):
|
||||
return
|
||||
|
||||
notification_plugin = self.plugin_manager.get_plugin("notification")
|
||||
if notification_plugin and notification_plugin.enabled:
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
"""
|
||||
通知邮件模块
|
||||
|
||||
提供错误通知邮件发送开关,并复用 SMTP 配置页。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from src.core.modules.base import ModuleCategory, ModuleDefinition, ModuleMetadata
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
def _validate_config(db: Session) -> tuple[bool, str]:
|
||||
"""
|
||||
验证通知邮件模块配置。
|
||||
|
||||
启用要求:SMTP 基础配置有效(至少 host + from_email)。
|
||||
"""
|
||||
from src.services.email.email_sender import EmailSenderService
|
||||
|
||||
if not EmailSenderService.is_smtp_configured(db):
|
||||
return False, "请先完成邮件配置(SMTP)"
|
||||
return True, ""
|
||||
|
||||
|
||||
notification_email_module = ModuleDefinition(
|
||||
metadata=ModuleMetadata(
|
||||
name="notification_email",
|
||||
display_name="异常通知",
|
||||
description="为 5xx 异常发送邮件通知,可在模块管理中启用或禁用",
|
||||
category=ModuleCategory.INTEGRATION,
|
||||
env_key="NOTIFICATION_EMAIL_AVAILABLE",
|
||||
default_available=True,
|
||||
required_packages=[],
|
||||
# 通知邮件与 SMTP 配置复用同一页面,不在模块卡片展示独立“配置”入口。
|
||||
admin_route=None,
|
||||
admin_menu_icon="Mail",
|
||||
admin_menu_group="system",
|
||||
admin_menu_order=58,
|
||||
),
|
||||
validate_config=_validate_config,
|
||||
)
|
||||
@@ -16,7 +16,7 @@ WARNING: 多进程环境注意事项
|
||||
3. 使用独立的健康检查服务:所有 worker 共享同一个健康状态源
|
||||
|
||||
目前项目已有 Redis 依赖,建议在高可用场景下将状态迁移到 Redis。
|
||||
参考:src/services/health_monitor.py 中的实现
|
||||
参考:src/services/health/monitor.py 中的实现
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
+68
-6
@@ -45,6 +45,24 @@ class PluginManager:
|
||||
"load_balancer": LoadBalancerStrategy,
|
||||
# 移除 "audit" - 审计功能现在是核心服务
|
||||
}
|
||||
# 默认按需加载集合(未显式配置时使用)
|
||||
# notification 默认不加载,避免未配置插件(如 email)在启动时初始化失败并占用内存。
|
||||
DEFAULT_ENABLED_PLUGIN_MODULES: dict[str, tuple[str, ...]] = {
|
||||
"auth": ("api_key",),
|
||||
"rate_limit": ("sliding_window",),
|
||||
"cache": ("memory",),
|
||||
"monitor": ("prometheus",),
|
||||
"token": ("claude",),
|
||||
"notification": (),
|
||||
"load_balancer": ("sticky_priority",),
|
||||
}
|
||||
# 部分插件“实例名”与“模块名”不同,按需加载时需要映射。
|
||||
MODULE_ALIASES: dict[str, dict[str, str]] = {
|
||||
"token": {
|
||||
"claude": "claude_counter",
|
||||
"tiktoken": "tiktoken_counter",
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self, config: dict[str, Any] | None = None):
|
||||
"""
|
||||
@@ -92,18 +110,62 @@ class PluginManager:
|
||||
if not type_dir.exists():
|
||||
continue
|
||||
|
||||
# 扫描插件目录
|
||||
for file_path in type_dir.glob("*.py"):
|
||||
if file_path.name.startswith("_") or file_path.name == "base.py":
|
||||
continue
|
||||
|
||||
module_name = f"src.plugins.{plugin_type}.{file_path.stem}"
|
||||
enabled_modules = self._resolve_enabled_plugin_modules(plugin_type, type_dir)
|
||||
for module_stem in enabled_modules:
|
||||
module_name = f"src.plugins.{plugin_type}.{module_stem}"
|
||||
try:
|
||||
module = importlib.import_module(module_name)
|
||||
self._load_plugin_from_module(module, plugin_type)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load plugin module {module_name}: {e}")
|
||||
|
||||
def _resolve_enabled_plugin_modules(self, plugin_type: str, type_dir: Path) -> list[str]:
|
||||
"""解析某类型当前应加载的插件模块名列表。"""
|
||||
available_modules = {
|
||||
file_path.stem
|
||||
for file_path in type_dir.glob("*.py")
|
||||
if not file_path.name.startswith("_") and file_path.name != "base.py"
|
||||
}
|
||||
if not available_modules:
|
||||
return []
|
||||
|
||||
type_config_raw = self.config.get(plugin_type, {})
|
||||
type_config = type_config_raw if isinstance(type_config_raw, dict) else {}
|
||||
enabled_modules: set[str] = set()
|
||||
|
||||
# 显式配置优先:仅加载配置中启用的插件 + default 指向插件
|
||||
if type_config:
|
||||
default_name = type_config.get("default")
|
||||
if isinstance(default_name, str) and default_name:
|
||||
enabled_modules.add(default_name)
|
||||
|
||||
for plugin_name, plugin_cfg in type_config.items():
|
||||
if plugin_name == "default":
|
||||
continue
|
||||
if isinstance(plugin_cfg, dict):
|
||||
if plugin_cfg.get("enabled", True):
|
||||
enabled_modules.add(plugin_name)
|
||||
elif plugin_cfg is not False:
|
||||
enabled_modules.add(plugin_name)
|
||||
else:
|
||||
# 无显式配置时使用默认按需集合
|
||||
enabled_modules.update(self.DEFAULT_ENABLED_PLUGIN_MODULES.get(plugin_type, ()))
|
||||
|
||||
aliases = self.MODULE_ALIASES.get(plugin_type, {})
|
||||
resolved_modules = set()
|
||||
for module_name in enabled_modules:
|
||||
resolved_modules.add(aliases.get(module_name, module_name))
|
||||
|
||||
missing_modules = resolved_modules - available_modules
|
||||
if missing_modules:
|
||||
logger.warning(
|
||||
"Plugin modules configured but not found for type {}: {}",
|
||||
plugin_type,
|
||||
sorted(missing_modules),
|
||||
)
|
||||
|
||||
return sorted(resolved_modules & available_modules)
|
||||
|
||||
def _is_api_version_compatible(self, plugin_api_version: str) -> bool:
|
||||
"""
|
||||
检查插件 API 版本是否兼容
|
||||
|
||||
@@ -2,13 +2,13 @@
|
||||
健康监控服务模块
|
||||
|
||||
包含健康监控相关功能:
|
||||
- health_monitor: 健康度监控单例
|
||||
- HealthMonitor: 健康监控类
|
||||
- get_health_monitor: 健康度监控单例获取函数(懒加载)
|
||||
"""
|
||||
|
||||
from .monitor import HealthMonitor, health_monitor
|
||||
from .monitor import HealthMonitor, get_health_monitor
|
||||
|
||||
__all__ = [
|
||||
"health_monitor",
|
||||
"HealthMonitor",
|
||||
"get_health_monitor",
|
||||
]
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import threading
|
||||
from collections import deque
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
@@ -128,7 +129,7 @@ class HealthMonitor:
|
||||
# Key: (key_id, api_format), Value: list of {"ts": float, "ok": bool}
|
||||
# 不再持久化到数据库,进程重启后自然重建
|
||||
_window_cache: dict[tuple[str, str], list[dict[str, Any]]] = {}
|
||||
_WINDOW_CACHE_MAX_ENTRIES = int(os.getenv("HEALTH_WINDOW_CACHE_MAX_ENTRIES", "10000"))
|
||||
_WINDOW_CACHE_MAX_ENTRIES = int(os.getenv("HEALTH_WINDOW_CACHE_MAX_ENTRIES", "5000"))
|
||||
|
||||
# ==================== 数据访问辅助方法 ====================
|
||||
|
||||
@@ -847,7 +848,7 @@ class HealthMonitor:
|
||||
try:
|
||||
endpoint_stats = db.query(
|
||||
func.count(ProviderEndpoint.id).label("total"),
|
||||
func.sum(case((ProviderEndpoint.is_active == True, 1), else_=0)).label("active"),
|
||||
func.sum(case((ProviderEndpoint.is_active.is_(True), 1), else_=0)).label("active"),
|
||||
func.sum(case((ProviderEndpoint.health_score < 0.5, 1), else_=0)).label(
|
||||
"unhealthy"
|
||||
),
|
||||
@@ -978,6 +979,21 @@ class HealthMonitor:
|
||||
return False
|
||||
|
||||
|
||||
# 全局健康监控器实例
|
||||
health_monitor = HealthMonitor()
|
||||
# 全局健康监控器懒加载(避免 import 阶段实例化)
|
||||
_health_monitor_instance: HealthMonitor | None = None
|
||||
_health_monitor_lock = threading.Lock()
|
||||
|
||||
|
||||
def get_health_monitor() -> HealthMonitor:
|
||||
"""获取全局 HealthMonitor 单例(线程安全懒加载)。"""
|
||||
global _health_monitor_instance # noqa: PLW0603
|
||||
|
||||
if _health_monitor_instance is None:
|
||||
with _health_monitor_lock:
|
||||
if _health_monitor_instance is None:
|
||||
_health_monitor_instance = HealthMonitor()
|
||||
|
||||
return _health_monitor_instance
|
||||
|
||||
|
||||
health_open_circuits.set(0)
|
||||
|
||||
@@ -113,7 +113,7 @@ async def fetch_models_for_key(
|
||||
# Ensure provider plugins (including custom model fetchers) are registered.
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
ensure_providers_bootstrapped(provider_types=[ctx.provider_type] if ctx.provider_type else None)
|
||||
|
||||
fetcher = UpstreamModelsFetcherRegistry.get(ctx.provider_type) or _fetch_models_default
|
||||
return await fetcher(ctx, timeout_seconds)
|
||||
|
||||
@@ -24,7 +24,7 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.health.monitor import get_health_monitor
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.provider.pool.config import parse_pool_config
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
@@ -157,7 +157,7 @@ class ErrorHandlerService:
|
||||
)
|
||||
if key:
|
||||
await asyncio.to_thread(
|
||||
health_monitor.record_failure,
|
||||
get_health_monitor().record_failure,
|
||||
db=self.db,
|
||||
key_id=str(key.id),
|
||||
api_format=provider_format_str,
|
||||
@@ -226,7 +226,7 @@ class ErrorHandlerService:
|
||||
# 记录健康失败
|
||||
if key:
|
||||
await asyncio.to_thread(
|
||||
health_monitor.record_failure,
|
||||
get_health_monitor().record_failure,
|
||||
db=self.db,
|
||||
key_id=str(key.id),
|
||||
api_format=provider_format_str,
|
||||
@@ -271,7 +271,7 @@ class ErrorHandlerService:
|
||||
# 记录健康失败
|
||||
if key:
|
||||
await asyncio.to_thread(
|
||||
health_monitor.record_failure,
|
||||
get_health_monitor().record_failure,
|
||||
db=self.db,
|
||||
key_id=str(key.id),
|
||||
api_format=provider_format_str,
|
||||
|
||||
@@ -10,7 +10,9 @@ provider-specific envelopes live in their own service modules.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import threading
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
@@ -117,7 +119,7 @@ def get_provider_envelope(
|
||||
endpoint_sig: str | None,
|
||||
) -> ProviderEnvelope | None:
|
||||
"""Return envelope hooks for the given provider_type + endpoint signature."""
|
||||
ensure_providers_bootstrapped()
|
||||
ensure_providers_bootstrapped(provider_types=[provider_type] if provider_type else None)
|
||||
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
|
||||
@@ -136,37 +138,111 @@ def get_provider_envelope(
|
||||
# ---------------------------------------------------------------------------
|
||||
# 所有 registry 共享同一个 bootstrap,首次访问任何 registry 时自动触发。
|
||||
# 不再依赖模块 import 顺序。
|
||||
_bootstrapped = False
|
||||
_bootstrap_lock = threading.Lock()
|
||||
_bootstrap_condition = threading.Condition(_bootstrap_lock)
|
||||
_bootstrap_in_progress = False
|
||||
_bootstrapped_provider_types: set[str] = set()
|
||||
_auto_detected_provider_types: frozenset[str] | None = None
|
||||
|
||||
_PROVIDER_PLUGIN_MODULES: dict[str, str] = {
|
||||
"antigravity": "src.services.provider.adapters.antigravity.plugin",
|
||||
"claude_code": "src.services.provider.adapters.claude_code.plugin",
|
||||
"codex": "src.services.provider.adapters.codex.plugin",
|
||||
"gemini_cli": "src.services.provider.adapters.gemini_cli.plugin",
|
||||
"kiro": "src.services.provider.adapters.kiro.plugin",
|
||||
"vertex_ai": "src.services.provider.adapters.vertex_ai.plugin",
|
||||
}
|
||||
|
||||
|
||||
def ensure_providers_bootstrapped() -> None:
|
||||
"""确保所有 provider plugin 已注册(幂等,只执行一次)。"""
|
||||
global _bootstrapped # noqa: PLW0603
|
||||
if _bootstrapped:
|
||||
return
|
||||
with _bootstrap_lock:
|
||||
if _bootstrapped:
|
||||
def _normalize_bootstrap_targets(provider_types: Iterable[str] | None) -> set[str]:
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
|
||||
if provider_types is None:
|
||||
return set()
|
||||
if isinstance(provider_types, str):
|
||||
provider_types = [provider_types]
|
||||
|
||||
targets: set[str] = set()
|
||||
for raw in provider_types:
|
||||
pt = normalize_provider_type(raw)
|
||||
if pt in _PROVIDER_PLUGIN_MODULES:
|
||||
targets.add(pt)
|
||||
return targets
|
||||
|
||||
|
||||
def _discover_active_provider_types() -> set[str]:
|
||||
"""从数据库读取活跃 provider_type,用于按需 bootstrap。"""
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
from src.database.database import create_session
|
||||
from src.models.database import Provider
|
||||
|
||||
db = create_session()
|
||||
try:
|
||||
rows = (
|
||||
db.query(Provider.provider_type).filter(Provider.is_active.is_(True)).distinct().all()
|
||||
)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
discovered: set[str] = set()
|
||||
for (raw_provider_type,) in rows:
|
||||
pt = normalize_provider_type(raw_provider_type)
|
||||
if pt in _PROVIDER_PLUGIN_MODULES:
|
||||
discovered.add(pt)
|
||||
return discovered
|
||||
|
||||
|
||||
def _bootstrap_provider_type(provider_type: str) -> None:
|
||||
module_path = _PROVIDER_PLUGIN_MODULES[provider_type]
|
||||
module = importlib.import_module(module_path)
|
||||
register_all = getattr(module, "register_all", None)
|
||||
if callable(register_all):
|
||||
register_all()
|
||||
|
||||
|
||||
def ensure_providers_bootstrapped(provider_types: Iterable[str] | None = None) -> None:
|
||||
"""确保 provider plugins 已注册(幂等,支持按 provider_type 精准注册)。"""
|
||||
global _auto_detected_provider_types, _bootstrap_in_progress # noqa: PLW0603
|
||||
|
||||
targets = _normalize_bootstrap_targets(provider_types)
|
||||
|
||||
# DB 查询在锁外执行,避免慢查询阻塞其他线程的 bootstrap 操作。
|
||||
need_discover = not targets and _auto_detected_provider_types is None
|
||||
if need_discover:
|
||||
try:
|
||||
detected = _discover_active_provider_types()
|
||||
except Exception:
|
||||
detected = set()
|
||||
else:
|
||||
detected = set()
|
||||
|
||||
with _bootstrap_condition:
|
||||
if not targets:
|
||||
if _auto_detected_provider_types is None:
|
||||
# 回退策略:DB 不可用/无记录时,保持原有全量 bootstrap 语义。
|
||||
_auto_detected_provider_types = frozenset(
|
||||
detected if detected else _PROVIDER_PLUGIN_MODULES.keys()
|
||||
)
|
||||
targets = set(_auto_detected_provider_types)
|
||||
|
||||
while _bootstrap_in_progress:
|
||||
_bootstrap_condition.wait()
|
||||
|
||||
missing = targets - _bootstrapped_provider_types
|
||||
if not missing:
|
||||
return
|
||||
_bootstrapped = True
|
||||
_bootstrap_in_progress = True
|
||||
|
||||
from src.services.provider.adapters.antigravity.plugin import (
|
||||
register_all as _reg_antigravity,
|
||||
)
|
||||
from src.services.provider.adapters.claude_code.plugin import (
|
||||
register_all as _reg_claude_code,
|
||||
)
|
||||
from src.services.provider.adapters.codex.plugin import register_all as _reg_codex
|
||||
from src.services.provider.adapters.gemini_cli.plugin import register_all as _reg_gemini_cli
|
||||
from src.services.provider.adapters.kiro.plugin import register_all as _reg_kiro
|
||||
from src.services.provider.adapters.vertex_ai.plugin import register_all as _reg_vertex_ai
|
||||
|
||||
_reg_antigravity()
|
||||
_reg_claude_code()
|
||||
_reg_codex()
|
||||
_reg_gemini_cli()
|
||||
_reg_kiro()
|
||||
_reg_vertex_ai()
|
||||
bootstrapped_now: set[str] = set()
|
||||
try:
|
||||
for pt in sorted(missing):
|
||||
_bootstrap_provider_type(pt)
|
||||
bootstrapped_now.add(pt)
|
||||
finally:
|
||||
with _bootstrap_condition:
|
||||
_bootstrapped_provider_types.update(bootstrapped_now)
|
||||
_bootstrap_in_progress = False
|
||||
_bootstrap_condition.notify_all()
|
||||
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from src.core.provider_types import ProviderType, normalize_provider_type
|
||||
@@ -139,10 +140,25 @@ def _extract_codex_weekly_reset_seconds(metadata: dict[str, Any]) -> float | Non
|
||||
if not isinstance(codex, dict):
|
||||
return None
|
||||
|
||||
parsed = safe_float(codex.get("primary_reset_seconds"))
|
||||
if parsed is None or parsed < 0:
|
||||
now = time.time()
|
||||
|
||||
# 优先绝对时间戳,避免 reset_seconds 快照随时间漂移。
|
||||
reset_at = safe_float(codex.get("primary_reset_at"))
|
||||
if reset_at is not None and reset_at > 0:
|
||||
remaining = reset_at - now
|
||||
return remaining if remaining > 0 else 0.0
|
||||
|
||||
reset_seconds = safe_float(codex.get("primary_reset_seconds"))
|
||||
if reset_seconds is None or reset_seconds < 0:
|
||||
return None
|
||||
return parsed
|
||||
|
||||
updated_at = safe_float(codex.get("updated_at"))
|
||||
if updated_at is not None and updated_at > 0:
|
||||
# 时钟偏移下 updated_at 可能晚于当前时间,elapsed 需要下限钳制到 0。
|
||||
elapsed = max(now - updated_at, 0.0)
|
||||
return max(reset_seconds - elapsed, 0.0)
|
||||
|
||||
return reset_seconds
|
||||
|
||||
|
||||
def extract_reset_seconds(key_obj: Any, provider_type: str | None = None) -> float | None:
|
||||
|
||||
@@ -83,7 +83,7 @@ def get_pool_hook(provider_type: str | None) -> PoolSchedulingHook | None:
|
||||
return None
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
ensure_providers_bootstrapped(provider_types=[provider_type])
|
||||
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
|
||||
|
||||
@@ -27,7 +27,11 @@ def _get_preset_mutex_group(preset_name: str) -> str | None:
|
||||
return _normalize_mutex_group(getattr(dim, "mutex_group", None))
|
||||
|
||||
|
||||
def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None], ...]:
|
||||
def _normalize_presets_from_config(
|
||||
config: Any,
|
||||
*,
|
||||
provider_type: str | None = None,
|
||||
) -> tuple[tuple[str, str | None], ...]:
|
||||
"""Extract enabled (preset_name, mode) tuples from config.scheduling_presets.
|
||||
|
||||
Supports both new SchedulingPreset objects and legacy string lists.
|
||||
@@ -40,6 +44,7 @@ def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None],
|
||||
if not isinstance(raw, (list, tuple)):
|
||||
return ()
|
||||
|
||||
normalized_provider_type = str(provider_type or "").strip().lower()
|
||||
allowed = get_preset_names() | {"lru"}
|
||||
entries: list[tuple[int, str, bool, str | None]] = []
|
||||
seen: set[str] = set()
|
||||
@@ -64,6 +69,15 @@ def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None],
|
||||
seen.add(preset_name)
|
||||
entries.append((idx, preset_name, enabled, mode))
|
||||
|
||||
# Codex 默认启用额度刷新优先维度(除非显式配置了 recent_refresh)。
|
||||
if (
|
||||
normalized_provider_type == "codex"
|
||||
and entries
|
||||
and "recent_refresh" not in {name for _idx, name, _enabled, _mode in entries}
|
||||
and "recent_refresh" in allowed
|
||||
):
|
||||
entries.append((len(entries), "recent_refresh", True, None))
|
||||
|
||||
if not entries:
|
||||
return ()
|
||||
|
||||
@@ -127,7 +141,10 @@ class MultiScoreStrategy:
|
||||
if not isinstance(keys_by_id, dict):
|
||||
keys_by_id = {}
|
||||
|
||||
presets = _normalize_presets_from_config(config)
|
||||
presets = _normalize_presets_from_config(
|
||||
config,
|
||||
provider_type=context.get("provider_type"),
|
||||
)
|
||||
lru_enabled = bool(getattr(config, "lru_enabled", True))
|
||||
if presets:
|
||||
return self._compute_preset_score(
|
||||
|
||||
@@ -197,7 +197,7 @@ def build_provider_url(
|
||||
# Provider transport hook: 如果有注册的 hook 则委托处理
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
ensure_providers_bootstrapped(provider_types=[provider_type] if provider_type else None)
|
||||
if provider_type and endpoint_sig:
|
||||
hook = _transport_hooks.get((provider_type, endpoint_sig))
|
||||
# Codex hook 仅在无 custom_path 时生效
|
||||
|
||||
@@ -93,7 +93,6 @@ class AdaptiveReservationManager:
|
||||
|
||||
def __init__(self, config: ReservationConfig | None = None):
|
||||
self.config = config or ReservationConfig()
|
||||
self._cache: dict[str, ReservationResult] = {} # 简单的内存缓存
|
||||
|
||||
def calculate_reservation(
|
||||
self,
|
||||
|
||||
@@ -16,7 +16,7 @@ from sqlalchemy.orm import Session
|
||||
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.health.monitor import get_health_monitor
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
|
||||
@@ -176,7 +176,7 @@ class RequestExecutor:
|
||||
health_format = provider_format_str or client_format_str
|
||||
|
||||
await asyncio.to_thread(
|
||||
health_monitor.record_success,
|
||||
get_health_monitor().record_success,
|
||||
db=self.db,
|
||||
key_id=key.id,
|
||||
api_format=health_format,
|
||||
|
||||
@@ -34,7 +34,7 @@ from src.models.database import (
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.services.health.monitor import get_health_monitor
|
||||
from src.services.provider.format import normalize_endpoint_signature
|
||||
from src.services.provider.pool.account_state import (
|
||||
resolve_pool_account_state as _resolve_pool_account_state,
|
||||
@@ -365,7 +365,7 @@ class CandidateBuilder:
|
||||
- mapping_matched_model: 通过映射匹配到的模型名(用于实际请求)
|
||||
"""
|
||||
# 检查熔断器状态(使用详细状态方法获取更丰富的跳过原因,按 API 格式)
|
||||
is_available, circuit_reason = health_monitor.get_circuit_breaker_status(
|
||||
is_available, circuit_reason = get_health_monitor().get_circuit_breaker_status(
|
||||
key, api_format=api_format
|
||||
)
|
||||
if not is_available:
|
||||
|
||||
@@ -18,11 +18,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import delete, literal_column, text
|
||||
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.database import create_session
|
||||
@@ -48,6 +50,20 @@ class MaintenanceScheduler:
|
||||
self._stats_aggregation_lock = asyncio.Lock()
|
||||
self._wallet_daily_usage_lock = asyncio.Lock()
|
||||
|
||||
@staticmethod
|
||||
def _get_http_client_idle_cleanup_interval_minutes() -> int:
|
||||
"""获取 HTTP 客户端空闲清理调度间隔(分钟)。"""
|
||||
raw = os.getenv("HTTP_CLIENT_IDLE_CLEANUP_INTERVAL_MINUTES", "5")
|
||||
try:
|
||||
minutes = int(raw)
|
||||
except ValueError:
|
||||
logger.warning(
|
||||
"环境变量 HTTP_CLIENT_IDLE_CLEANUP_INTERVAL_MINUTES 非法: {}, 使用默认值 5",
|
||||
raw,
|
||||
)
|
||||
return 5
|
||||
return max(1, minutes)
|
||||
|
||||
def _get_checkin_time(self) -> tuple[int, int]:
|
||||
"""获取签到任务的执行时间
|
||||
|
||||
@@ -170,6 +186,14 @@ class MaintenanceScheduler:
|
||||
name="连接池监控",
|
||||
)
|
||||
|
||||
# HTTP 代理/Tunnel 客户端空闲清理 - 默认每 5 分钟
|
||||
scheduler.add_interval_job(
|
||||
self._scheduled_http_client_idle_cleanup,
|
||||
minutes=self._get_http_client_idle_cleanup_interval_minutes(),
|
||||
job_id="http_client_idle_cleanup",
|
||||
name="HTTP客户端空闲清理",
|
||||
)
|
||||
|
||||
# Pending 状态清理 - 每 5 分钟
|
||||
scheduler.add_interval_job(
|
||||
self._scheduled_pending_cleanup,
|
||||
@@ -296,7 +320,20 @@ class MaintenanceScheduler:
|
||||
|
||||
log_pool_status()
|
||||
except Exception as e:
|
||||
logger.exception(f"连接池监控任务出错: {e}")
|
||||
logger.exception("连接池监控任务出错: {}", e)
|
||||
|
||||
async def _scheduled_http_client_idle_cleanup(self) -> None:
|
||||
"""HTTP 客户端空闲清理任务(定时调用)。"""
|
||||
try:
|
||||
stats = await HTTPClientPool.cleanup_idle_clients()
|
||||
if stats.get("proxy_closed", 0) or stats.get("tunnel_closed", 0):
|
||||
logger.info(
|
||||
"HTTP 客户端空闲清理释放连接: proxy={}, tunnel={}",
|
||||
stats.get("proxy_closed", 0),
|
||||
stats.get("tunnel_closed", 0),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.exception("HTTP 客户端空闲清理任务出错: {}", e)
|
||||
|
||||
async def _scheduled_pending_cleanup(self) -> None:
|
||||
"""Pending 清理任务(定时调用)"""
|
||||
|
||||
@@ -6,6 +6,7 @@ from src.core.api_format.signature import normalize_signature_key
|
||||
from src.services.billing.token_normalization import normalize_input_tokens_for_billing
|
||||
from src.services.usage._recording_helpers import (
|
||||
build_usage_params,
|
||||
deserialize_body_if_json,
|
||||
sanitize_request_metadata,
|
||||
)
|
||||
from src.services.usage._types import UsageCostInfo, UsageRecordParams
|
||||
@@ -158,6 +159,10 @@ class UsageBillingIntegrationMixin:
|
||||
metadata = sanitize_request_metadata(metadata)
|
||||
|
||||
# 构建 Usage 参数
|
||||
request_body = deserialize_body_if_json(params.request_body)
|
||||
provider_request_body = deserialize_body_if_json(params.provider_request_body)
|
||||
response_body = deserialize_body_if_json(params.response_body)
|
||||
client_response_body = deserialize_body_if_json(params.client_response_body)
|
||||
usage_params = build_usage_params(
|
||||
db=params.db,
|
||||
user=params.user,
|
||||
@@ -183,13 +188,13 @@ class UsageBillingIntegrationMixin:
|
||||
error_message=params.error_message,
|
||||
metadata=metadata,
|
||||
request_headers=params.request_headers,
|
||||
request_body=params.request_body,
|
||||
request_body=request_body,
|
||||
provider_request_headers=params.provider_request_headers,
|
||||
provider_request_body=params.provider_request_body,
|
||||
provider_request_body=provider_request_body,
|
||||
response_headers=params.response_headers,
|
||||
client_response_headers=params.client_response_headers,
|
||||
response_body=params.response_body,
|
||||
client_response_body=params.client_response_body,
|
||||
response_body=response_body,
|
||||
client_response_body=client_response_body,
|
||||
request_id=params.request_id,
|
||||
provider_id=params.provider_id,
|
||||
provider_endpoint_id=params.provider_endpoint_id,
|
||||
|
||||
@@ -46,6 +46,25 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
|
||||
)
|
||||
|
||||
|
||||
def deserialize_body_if_json(value: Any) -> Any:
|
||||
"""写库前按需反序列化 body JSON 字符串。
|
||||
|
||||
仅对 JSON object/array 字符串做 json.loads,其他值保持原样。
|
||||
"""
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
stripped = value.lstrip()
|
||||
if not stripped or stripped[0] not in "{[":
|
||||
return value
|
||||
try:
|
||||
parsed = json.loads(value)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return value
|
||||
if isinstance(parsed, (dict, list)):
|
||||
return parsed
|
||||
return value
|
||||
|
||||
|
||||
def build_usage_params(
|
||||
*,
|
||||
db: Session,
|
||||
|
||||
@@ -7,7 +7,6 @@ Usage Redis Streams consumer.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import time
|
||||
@@ -20,7 +19,7 @@ from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.clients.redis_client import get_usage_queue_redis_client as get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.database.database import create_session
|
||||
@@ -34,19 +33,7 @@ def _consumer_name() -> str:
|
||||
|
||||
|
||||
def _parse_body(value: Any) -> Any:
|
||||
"""将 JSON 字符串 body 反序列化为 dict,否则原样返回。
|
||||
|
||||
QueueTelemetryWriter 会将 body 序列化为 JSON 字符串以便传输,
|
||||
消费者需要将其反序列化回 dict 以正确存入 JSON 列。
|
||||
"""
|
||||
if value is None or isinstance(value, dict):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return json.loads(value)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
# 解析失败,保留原字符串(可能已被截断)
|
||||
return value
|
||||
"""消费者阶段保留原始 body,反序列化延迟到写库阶段。"""
|
||||
return value
|
||||
|
||||
|
||||
@@ -115,7 +102,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("[usage-queue] Created consumer group {}", config.usage_queue_stream_group)
|
||||
except ResponseError as exc:
|
||||
if "BUSYGROUP" in str(exc):
|
||||
return
|
||||
@@ -160,7 +147,7 @@ class UsageQueueConsumer:
|
||||
return
|
||||
self._running = True
|
||||
self._task = asyncio.create_task(self._run(), name="usage-queue-consumer")
|
||||
logger.info(f"[usage-queue] Consumer started: {self._consumer}")
|
||||
logger.info("[usage-queue] Consumer started: {}", self._consumer)
|
||||
|
||||
async def stop(self) -> None:
|
||||
if not self._running:
|
||||
@@ -172,7 +159,7 @@ class UsageQueueConsumer:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
logger.info(f"[usage-queue] Consumer stopped: {self._consumer}")
|
||||
logger.info("[usage-queue] Consumer stopped: {}", self._consumer)
|
||||
|
||||
async def _run(self) -> None:
|
||||
while self._running:
|
||||
@@ -188,10 +175,10 @@ class UsageQueueConsumer:
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except (RedisTimeoutError, RedisConnectionError) as exc:
|
||||
logger.warning(f"[usage-queue] Redis connection issue: {exc}")
|
||||
logger.warning("[usage-queue] Redis connection issue: {}", exc)
|
||||
await asyncio.sleep(1)
|
||||
except Exception as exc:
|
||||
logger.exception(f"[usage-queue] Consumer loop error: {exc}")
|
||||
logger.exception("[usage-queue] Consumer loop error: {}", exc)
|
||||
await asyncio.sleep(1)
|
||||
|
||||
async def _maybe_claim_pending(self, redis_client: Any) -> None:
|
||||
@@ -209,7 +196,7 @@ class UsageQueueConsumer:
|
||||
count=self._batch_size,
|
||||
)
|
||||
except ResponseError as exc:
|
||||
logger.warning(f"[usage-queue] XAUTOCLAIM failed: {exc}")
|
||||
logger.warning("[usage-queue] XAUTOCLAIM failed: {}", exc)
|
||||
return
|
||||
if not result:
|
||||
return
|
||||
@@ -315,12 +302,12 @@ class UsageQueueConsumer:
|
||||
pipe.xack(self._stream_key, self._stream_group, message_id)
|
||||
await pipe.execute()
|
||||
|
||||
logger.debug(f"[usage-queue] Batch processed {len(records)} records")
|
||||
logger.debug("[usage-queue] Batch processed {} records", len(records))
|
||||
|
||||
except Exception as exc:
|
||||
# 批量处理失败,回退到逐条处理(复用已创建的 db session)
|
||||
logger.warning(
|
||||
f"[usage-queue] Batch processing failed, falling back to individual: {exc}"
|
||||
"[usage-queue] Batch processing failed, falling back to individual: {}", exc
|
||||
)
|
||||
try:
|
||||
db.rollback() # 清理批量失败的事务状态
|
||||
@@ -339,7 +326,7 @@ class UsageQueueConsumer:
|
||||
pass
|
||||
if self._is_duplicate_key_error(ie):
|
||||
logger.debug(
|
||||
f"[usage-queue] Duplicate request_id, skipping: {event.request_id}"
|
||||
"[usage-queue] Duplicate request_id, skipping: {}", event.request_id
|
||||
)
|
||||
success_ids.append(message_id)
|
||||
else:
|
||||
@@ -381,13 +368,16 @@ class UsageQueueConsumer:
|
||||
await redis_client.xadd(self._dlq_key, dlq_fields)
|
||||
await redis_client.xack(self._stream_key, self._stream_group, message_id)
|
||||
logger.error(
|
||||
f"[usage-queue] Message moved to DLQ after {retries} attempts: {message_id}"
|
||||
"[usage-queue] Message moved to DLQ after {} attempts: {}", retries, message_id
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(f"[usage-queue] Failed to move message to DLQ: {exc}")
|
||||
logger.error("[usage-queue] Failed to move message to DLQ: {}", exc)
|
||||
else:
|
||||
logger.warning(
|
||||
f"[usage-queue] Processing failed (attempt {retries}): {message_id} error={error}"
|
||||
"[usage-queue] Processing failed (attempt {}): {} error={}",
|
||||
retries,
|
||||
message_id,
|
||||
error,
|
||||
)
|
||||
|
||||
async def _get_delivery_count(self, redis_client: Any, message_id: str) -> int:
|
||||
@@ -531,9 +521,9 @@ class UsageQueueConsumer:
|
||||
break
|
||||
# lag=未读消息数, pending=已读但未ACK的消息数
|
||||
if lag > 0 or pending_count > 0:
|
||||
logger.info(f"[usage-queue] lag={lag} pending={pending_count}")
|
||||
logger.info("[usage-queue] lag={} pending={}", lag, pending_count)
|
||||
except Exception as exc:
|
||||
logger.debug(f"[usage-queue] metrics log failed: {exc}")
|
||||
logger.debug("[usage-queue] metrics log failed: {}", exc)
|
||||
|
||||
|
||||
_consumer_instance: UsageQueueConsumer | None = None
|
||||
|
||||
@@ -10,6 +10,9 @@ from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
import msgpack
|
||||
from msgpack.exceptions import OutOfData
|
||||
|
||||
USAGE_EVENT_VERSION = 1
|
||||
|
||||
|
||||
@@ -38,6 +41,36 @@ def sanitize_payload(data: dict[str, Any]) -> dict[str, Any]:
|
||||
return {str(k): _sanitize_value(v) for k, v in data.items()}
|
||||
|
||||
|
||||
def _decode_payload(raw: Any) -> dict[str, Any]:
|
||||
"""兼容解码:优先 msgpack,回退旧 JSON。
|
||||
|
||||
统一将输入归一化为 bytes 后走单一解码路径:msgpack → JSON fallback。
|
||||
str 输入来自 decode_responses=True + surrogateescape 的 Redis 客户端,
|
||||
通过 surrogateescape 可无损还原回原始 bytes。
|
||||
"""
|
||||
if isinstance(raw, str):
|
||||
raw = raw.encode("utf-8", errors="surrogateescape")
|
||||
|
||||
if not isinstance(raw, (bytes, bytearray, memoryview)):
|
||||
raise ValueError("Invalid payload field in usage event")
|
||||
|
||||
payload_bytes = bytes(raw)
|
||||
|
||||
# 新格式:msgpack
|
||||
try:
|
||||
payload = msgpack.unpackb(payload_bytes, raw=False)
|
||||
except (ValueError, OutOfData, TypeError):
|
||||
# 兼容旧格式:JSON bytes(含 surrogateescape 还原后的纯 UTF-8 JSON)
|
||||
try:
|
||||
payload = json.loads(payload_bytes.decode("utf-8"))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError, TypeError) as exc:
|
||||
raise ValueError("Invalid payload field in usage event") from exc
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("Invalid payload field in usage event")
|
||||
return payload
|
||||
|
||||
|
||||
@dataclass
|
||||
class UsageEvent:
|
||||
event_type: UsageEventType
|
||||
@@ -45,24 +78,29 @@ class UsageEvent:
|
||||
timestamp_ms: int
|
||||
data: dict[str, Any]
|
||||
|
||||
def to_stream_fields(self) -> dict[str, str]:
|
||||
def to_stream_fields(self) -> dict[str, bytes]:
|
||||
"""序列化为 Redis Stream 字段。
|
||||
|
||||
该函数返回 bytes payload,要求读写 usage queue 的 Redis 客户端使用
|
||||
decode_responses=True,并以 surrogateescape 做 UTF-8 编解码,
|
||||
以保证 bytes <-> str 往返无损。
|
||||
"""
|
||||
payload = {
|
||||
"v": USAGE_EVENT_VERSION,
|
||||
"type": self.event_type.value,
|
||||
"request_id": self.request_id,
|
||||
"timestamp_ms": self.timestamp_ms,
|
||||
# 兜底清洗,避免 metadata 中混入非 JSON 类型导致队列写入失败。
|
||||
"data": sanitize_payload(self.data),
|
||||
}
|
||||
return {"payload": json.dumps(payload, ensure_ascii=False)}
|
||||
return {"payload": msgpack.packb(payload, use_bin_type=True)}
|
||||
|
||||
@classmethod
|
||||
def from_stream_fields(cls, fields: dict[str, Any]) -> UsageEvent:
|
||||
raw = fields.get("payload")
|
||||
if not raw:
|
||||
raise ValueError("Missing payload field in usage event")
|
||||
if isinstance(raw, bytes):
|
||||
raw = raw.decode("utf-8", errors="ignore")
|
||||
payload = json.loads(raw)
|
||||
payload = _decode_payload(raw)
|
||||
event_type = UsageEventType(payload["type"])
|
||||
return cls(
|
||||
event_type=event_type,
|
||||
|
||||
@@ -27,6 +27,7 @@ from src.services.usage._recording_helpers import (
|
||||
METADATA_KEEP_KEYS,
|
||||
METADATA_PRUNE_KEYS,
|
||||
build_usage_params,
|
||||
deserialize_body_if_json,
|
||||
sanitize_request_metadata,
|
||||
update_existing_usage,
|
||||
)
|
||||
@@ -700,13 +701,13 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
|
||||
error_message=error_message,
|
||||
metadata=metadata,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
request_body=deserialize_body_if_json(request_body),
|
||||
provider_request_headers=provider_request_headers,
|
||||
provider_request_body=provider_request_body,
|
||||
provider_request_body=deserialize_body_if_json(provider_request_body),
|
||||
response_headers=response_headers,
|
||||
client_response_headers=client_response_headers,
|
||||
response_body=response_body,
|
||||
client_response_body=client_response_body,
|
||||
response_body=deserialize_body_if_json(response_body),
|
||||
client_response_body=deserialize_body_if_json(client_response_body),
|
||||
request_id=request_id,
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=provider_endpoint_id,
|
||||
|
||||
@@ -7,6 +7,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
@@ -21,6 +22,17 @@ from src.models.database import ApiKey, User
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
def _get_response_chunks_max_size() -> int:
|
||||
"""读取响应块存储上限(字节),默认 2MB。"""
|
||||
raw_mb = os.getenv("RESPONSE_CHUNKS_MAX_SIZE_MB", "2")
|
||||
try:
|
||||
mb = int(raw_mb)
|
||||
except ValueError:
|
||||
logger.warning("环境变量 RESPONSE_CHUNKS_MAX_SIZE_MB 非法: {}, 使用默认值 2", raw_mb)
|
||||
mb = 2
|
||||
return max(1, mb) * 1024 * 1024
|
||||
|
||||
|
||||
class StreamUsageTracker:
|
||||
"""流式响应用量跟踪器"""
|
||||
|
||||
@@ -120,7 +132,7 @@ class StreamUsageTracker:
|
||||
self.response_chunks = [] # 保存解析后的响应块
|
||||
self.response_chunks_count = 0 # 响应块总计数(含被丢弃的)
|
||||
self.response_chunks_size = 0 # 响应块累计序列化大小(字节)
|
||||
self._response_chunks_max_size = 4 * 1024 * 1024 # 4MB,留余量给 truncate_body 的 5MB 上限
|
||||
self._response_chunks_max_size = _get_response_chunks_max_size()
|
||||
self.raw_chunks: deque[str | bytes] = deque(
|
||||
maxlen=50
|
||||
) # 仅保留最后50个原始chunk(用于错误诊断)
|
||||
|
||||
@@ -8,7 +8,7 @@ import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.clients.redis_client import get_usage_queue_redis_client as get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.services.usage.events import UsageEventType, build_usage_event
|
||||
@@ -125,7 +125,7 @@ class QueueTelemetryWriter(TelemetryWriter):
|
||||
else:
|
||||
await redis_client.xadd(config.usage_queue_stream_key, event.to_stream_fields())
|
||||
except Exception as exc:
|
||||
logger.error(f"[usage-queue] XADD failed: {exc}")
|
||||
logger.error("[usage-queue] XADD failed: {}", exc)
|
||||
raise
|
||||
|
||||
def _mask_headers(self, headers: Any) -> Any:
|
||||
|
||||
Reference in New Issue
Block a user