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:
fawney19
2026-03-14 11:59:07 +08:00
co-authored by AAEE86
parent 45985f1c04
commit e0286aebe3
111 changed files with 2774 additions and 1101 deletions
+2 -2
View File
@@ -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 ====================
+2 -2
View File
@@ -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]:
+2 -2
View File
@@ -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()
+2 -2
View File
@@ -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)
+6 -6
View File
@@ -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} 不存在")
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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}$")
+2 -2
View File
@@ -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()
# ============== 安全基类 ==============
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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 ==========
+2 -2
View File
@@ -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 ==========
+2 -2
View File
@@ -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")
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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
View File
@@ -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)
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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])
+2 -2
View File
@@ -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()
# 映射预览配置(管理后台功能,限制宽松)
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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()
# ---------------------------------------------------------------------------
+2 -2
View File
@@ -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 模型 ==========
+2 -2
View File
@@ -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
View File
@@ -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": "模板已重置为默认值",
+2 -2
View File
@@ -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(
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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("")
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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()
# ============== 公共端点(所有用户可访问) ==============
+2 -2
View File
@@ -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端点
+8
View File
@@ -684,3 +684,11 @@ class ApiRequestPipeline:
except Exception:
return str(value)
return str(value)
_shared_pipeline = ApiRequestPipeline()
def get_pipeline() -> ApiRequestPipeline:
"""返回全局共享的无状态请求管道实例。"""
return _shared_pipeline
+5 -3
View File
@@ -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)
+6 -6
View File
@@ -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)
+61 -2
View File
@@ -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")
+35
View File
@@ -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,
+24 -15
View File
@@ -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
+54 -5
View File
@@ -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",
+25 -17
View File
@@ -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
+23 -10
View File
@@ -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
+24 -13
View File
@@ -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
+23 -10
View File
@@ -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
+2 -2
View File
@@ -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")
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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")
+2 -2
View File
@@ -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")
+6 -6
View File
@@ -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
View File
@@ -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,
+8 -4
View File
@@ -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,
+2 -2
View File
@@ -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 --------------------
+2 -2
View File
@@ -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()
# ============== 安全基类 ==============
+2 -2
View File
@@ -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:
+2 -2
View File
@@ -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]:
+27 -2
View File
@@ -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:
+88 -1
View File
@@ -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
View File
@@ -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
+8
View File
@@ -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 数据,足够覆盖正常事件头
+19
View File
@@ -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)
# - 未设置: 开发环境启用,生产环境禁用
+208 -43
View File
@@ -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
View File
@@ -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())
-106
View File
@@ -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
-214
View File
@@ -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
+3 -1
View File
@@ -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
View File
@@ -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"
+38
View File
@@ -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,
)
+1 -1
View File
@@ -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
View File
@@ -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 版本是否兼容
+3 -3
View File
@@ -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",
]
+20 -4
View File
@@ -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)
+1 -1
View File
@@ -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)
+4 -4
View File
@@ -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,
+103 -27
View File
@@ -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:
+1 -1
View File
@@ -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(
+1 -1
View File
@@ -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,
+2 -2
View File
@@ -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,
+2 -2
View File
@@ -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:
+38 -1
View File
@@ -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 清理任务(定时调用)"""
+9 -4
View File
@@ -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,
+19
View File
@@ -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,
+19 -29
View File
@@ -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
+43 -5
View File
@@ -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,
+5 -4
View File
@@ -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,
+13 -1
View File
@@ -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(用于错误诊断)
+2 -2
View File
@@ -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: