mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
refactor: 全局适配 ApiFamily/EndpointKind 结构化标识体系
将新的 (ApiFamily, EndpointKind) / `family:kind` 签名体系应用到整个代码库: - API Handlers: 所有 adapter/handler 使用新的签名格式 - Services: provider, model, usage, cache, auth 等服务层适配 - Database: ProviderEndpoint 新增 api_family/endpoint_kind 字段 - Frontend: Provider 管理、Usage 表格等组件适配 - Tests: 更新所有相关测试用例
This commit is contained in:
@@ -11,20 +11,20 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
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.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
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(prefix="/api/admin/adaptive", tags=["Adaptive RPM"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
@@ -5,15 +5,16 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
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.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
@@ -21,8 +22,6 @@ from src.database import get_db
|
||||
from src.models.api import CreateApiKeyRequest
|
||||
from src.models.database import ApiKey
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
|
||||
# 应用时区配置,默认为 Asia/Shanghai
|
||||
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
|
||||
@@ -432,7 +431,9 @@ class AdminCreateStandaloneKeyAdapter(AdminApiAdapter):
|
||||
auto_delete_on_expiry=self.key_data.auto_delete_on_expiry,
|
||||
)
|
||||
|
||||
logger.info(f"管理员创建独立余额Key: ID {api_key.id}, 初始余额 ${self.key_data.initial_balance_usd}")
|
||||
logger.info(
|
||||
f"管理员创建独立余额Key: ID {api_key.id}, 初始余额 ${self.key_data.initial_balance_usd}"
|
||||
)
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="create_standalone_api_key",
|
||||
@@ -548,7 +549,9 @@ class AdminToggleApiKeyAdapter(AdminApiAdapter):
|
||||
db.commit()
|
||||
db.refresh(api_key)
|
||||
|
||||
logger.info(f"管理员切换API密钥状态: Key ID {self.key_id}, 新状态 {'启用' if api_key.is_active else '禁用'}")
|
||||
logger.info(
|
||||
f"管理员切换API密钥状态: Key ID {self.key_id}, 新状态 {'启用' if api_key.is_active else '禁用'}"
|
||||
)
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="toggle_api_key",
|
||||
@@ -581,7 +584,9 @@ class AdminToggleLockApiKeyAdapter(AdminApiAdapter):
|
||||
db.commit()
|
||||
db.refresh(api_key)
|
||||
|
||||
logger.info(f"管理员切换API密钥锁定状态: Key ID {self.key_id}, 新状态 {'锁定' if api_key.is_locked else '解锁'}")
|
||||
logger.info(
|
||||
f"管理员切换API密钥锁定状态: Key ID {self.key_id}, 新状态 {'锁定' if api_key.is_locked else '解锁'}"
|
||||
)
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="toggle_lock_api_key",
|
||||
@@ -611,7 +616,9 @@ class AdminDeleteApiKeyAdapter(AdminApiAdapter):
|
||||
db.delete(api_key)
|
||||
db.commit()
|
||||
|
||||
logger.info(f"管理员删除API密钥: Key ID {self.key_id}, 用户 {user.email if user else '未知'}")
|
||||
logger.info(
|
||||
f"管理员删除API密钥: Key ID {self.key_id}, 用户 {user.email if user else '未知'}"
|
||||
)
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="delete_api_key",
|
||||
|
||||
@@ -2,20 +2,20 @@
|
||||
Key RPM 限制管理 API
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
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.core.exceptions import NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.models.endpoint_models import KeyRpmStatusResponse
|
||||
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(tags=["RPM Control"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
@@ -4,16 +4,17 @@ Endpoint 健康监控 API
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import func
|
||||
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.core.exceptions import NotFoundException
|
||||
from src.core.logger import logger
|
||||
@@ -27,10 +28,16 @@ from src.models.endpoint_models import (
|
||||
HealthSummaryResponse,
|
||||
)
|
||||
from src.services.health.endpoint import EndpointHealthService
|
||||
from src.services.health.monitor import health_monitor
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.services.health.monitor import HealthMonitor, health_monitor
|
||||
|
||||
router = APIRouter(tags=["Endpoint Health"])
|
||||
|
||||
|
||||
def _format_str(api_format_enum: Any) -> str:
|
||||
"""将 DB 查询返回的 api_format(可能是 enum 或 str)统一转为 str。"""
|
||||
return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
|
||||
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@@ -234,8 +241,6 @@ class AdminEndpointHealthStatusAdapter(AdminApiAdapter):
|
||||
lookback_hours: int
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from src.services.health.endpoint import EndpointHealthService
|
||||
|
||||
db = context.db
|
||||
|
||||
# 使用共享服务获取健康状态(管理员视图)
|
||||
@@ -265,30 +270,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
now = datetime.now(timezone.utc)
|
||||
since = now - timedelta(hours=self.lookback_hours)
|
||||
|
||||
# 1. 获取所有活跃的 API 格式及其 Provider 数量
|
||||
active_formats = (
|
||||
db.query(
|
||||
ProviderEndpoint.api_format,
|
||||
func.count(func.distinct(ProviderEndpoint.provider_id)).label("provider_count"),
|
||||
)
|
||||
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||
.filter(
|
||||
ProviderEndpoint.is_active.is_(True),
|
||||
Provider.is_active.is_(True),
|
||||
)
|
||||
.group_by(ProviderEndpoint.api_format)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 构建所有格式的 provider_count 映射
|
||||
all_formats: dict[str, int] = {}
|
||||
for api_format_enum, provider_count in active_formats:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
)
|
||||
all_formats[api_format] = provider_count
|
||||
|
||||
# 1.1 建立每个 API 格式对应的 Endpoint ID 列表(用于时间线生成),并收集活跃的 provider+format 组合
|
||||
# 1. 单次查询获取所有活跃 endpoint 行,在内存中聚合 provider_count / endpoint_map
|
||||
endpoint_rows = (
|
||||
db.query(ProviderEndpoint.api_format, ProviderEndpoint.id, ProviderEndpoint.provider_id)
|
||||
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||
@@ -298,14 +280,20 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
all_formats: dict[str, int] = {} # api_format -> distinct provider count
|
||||
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||
active_provider_formats: set[tuple[str, str]] = set()
|
||||
_provider_sets: dict[str, set[str]] = defaultdict(set)
|
||||
|
||||
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
)
|
||||
endpoint_map[api_format].append(endpoint_id)
|
||||
active_provider_formats.add((str(provider_id), api_format))
|
||||
fmt = _format_str(api_format_enum)
|
||||
endpoint_map[fmt].append(endpoint_id)
|
||||
_provider_sets[fmt].add(str(provider_id))
|
||||
active_provider_formats.add((str(provider_id), fmt))
|
||||
|
||||
for fmt, pids in _provider_sets.items():
|
||||
all_formats[fmt] = len(pids)
|
||||
|
||||
# 1.2 统计每个 API 格式可用的活跃 Key 数量(Key 属于 Provider,通过 api_formats 关联格式)
|
||||
key_counts: dict[str, int] = {}
|
||||
@@ -321,7 +309,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
for provider_id, api_formats in active_provider_keys:
|
||||
pid = str(provider_id)
|
||||
for fmt in (api_formats or []):
|
||||
for fmt in api_formats or []:
|
||||
if (pid, fmt) not in active_provider_formats:
|
||||
continue
|
||||
key_counts[fmt] = key_counts.get(fmt, 0) + 1
|
||||
@@ -347,12 +335,10 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
# 构建每个格式的状态统计
|
||||
status_counts: dict[str, dict[str, int]] = {}
|
||||
for api_format_enum, status, count in status_counts_query:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
)
|
||||
if api_format not in status_counts:
|
||||
status_counts[api_format] = {"success": 0, "failed": 0, "skipped": 0}
|
||||
status_counts[api_format][status] = count
|
||||
fmt = _format_str(api_format_enum)
|
||||
if fmt not in status_counts:
|
||||
status_counts[fmt] = {"success": 0, "failed": 0, "skipped": 0}
|
||||
status_counts[fmt][status] = count
|
||||
|
||||
# 3. 获取最近一段时间的 RequestCandidate(限制数量)
|
||||
# 使用上面定义的 final_statuses,排除中间状态
|
||||
@@ -376,15 +362,13 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
grouped_attempts: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for attempt, api_format_enum, provider_id in rows:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
)
|
||||
if api_format not in grouped_attempts:
|
||||
grouped_attempts[api_format] = []
|
||||
fmt = _format_str(api_format_enum)
|
||||
if fmt not in grouped_attempts:
|
||||
grouped_attempts[fmt] = []
|
||||
|
||||
# 只保留每个 API 格式最近 per_format_limit 条记录
|
||||
if len(grouped_attempts[api_format]) < self.per_format_limit:
|
||||
grouped_attempts[api_format].append(attempt)
|
||||
if len(grouped_attempts[fmt]) < self.per_format_limit:
|
||||
grouped_attempts[fmt].append(attempt)
|
||||
|
||||
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||
monitors: list[ApiFormatHealthMonitor] = []
|
||||
@@ -551,17 +535,22 @@ class AdminRecoverAllKeysHealthAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
# 查找所有有熔断格式的 Key(检查 circuit_breaker_by_format JSON 字段)
|
||||
all_keys = db.query(ProviderAPIKey).all()
|
||||
# 粗过滤:仅加载 circuit_breaker_by_format 非空的 Key,避免全表扫描
|
||||
candidates = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(
|
||||
ProviderAPIKey.circuit_breaker_by_format.isnot(None),
|
||||
ProviderAPIKey.circuit_breaker_by_format != "{}",
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 筛选出有任何格式熔断的 Key
|
||||
circuit_open_keys = []
|
||||
for key in all_keys:
|
||||
circuit_by_format = key.circuit_breaker_by_format or {}
|
||||
for fmt, circuit_data in circuit_by_format.items():
|
||||
if circuit_data.get("open"):
|
||||
circuit_open_keys.append(key)
|
||||
break
|
||||
# 精确筛选有任何格式熔断的 Key
|
||||
circuit_open_keys = [
|
||||
key
|
||||
for key in candidates
|
||||
if any(cb.get("open") for cb in (key.circuit_breaker_by_format or {}).values())
|
||||
]
|
||||
|
||||
if not circuit_open_keys:
|
||||
return {
|
||||
@@ -586,11 +575,8 @@ class AdminRecoverAllKeysHealthAdapter(AdminApiAdapter):
|
||||
|
||||
db.commit()
|
||||
|
||||
# 重置健康监控器的计数
|
||||
from src.services.health.monitor import HealthMonitor, health_open_circuits
|
||||
|
||||
HealthMonitor._open_circuit_keys = 0
|
||||
health_open_circuits.set(0)
|
||||
# 重置健康监控器的熔断计数
|
||||
HealthMonitor.reset_open_circuit_count()
|
||||
|
||||
logger.info(f"管理员批量恢复 {len(recovered_keys)} 个 Key 的健康状态")
|
||||
|
||||
|
||||
@@ -4,16 +4,17 @@ Provider API Keys 管理
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import json
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
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.core.crypto import crypto_service
|
||||
@@ -28,7 +29,6 @@ from src.models.endpoint_models import (
|
||||
EndpointAPIKeyUpdate,
|
||||
)
|
||||
from src.services.cache.provider_cache import ProviderCacheService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(tags=["Provider Keys"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -261,7 +261,9 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
|
||||
update_data["api_key"] = crypto_service.encrypt(update_data["api_key"])
|
||||
# 加密 auth_config(包含敏感的 Service Account 凭证)
|
||||
if "auth_config" in update_data and update_data["auth_config"]:
|
||||
update_data["auth_config"] = crypto_service.encrypt(json.dumps(update_data["auth_config"]))
|
||||
update_data["auth_config"] = crypto_service.encrypt(
|
||||
json.dumps(update_data["auth_config"])
|
||||
)
|
||||
|
||||
# 特殊处理 rpm_limit:需要区分"未提供"和"显式设置为 null"
|
||||
if "rpm_limit" in self.key_data.model_fields_set:
|
||||
@@ -410,10 +412,10 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
|
||||
# 检查是否是新格式的占位符(表示 auth_config 丢失)
|
||||
if decrypted_key == "__placeholder__":
|
||||
logger.error(f"Vertex AI Key 缺少 auth_config: ID={self.key_id}")
|
||||
raise InvalidRequestException(
|
||||
"认证配置丢失,请重新添加该密钥。"
|
||||
)
|
||||
logger.info(f"[REVEAL] 查看完整 Key (legacy vertex_ai): ID={self.key_id}, Name={key.name}")
|
||||
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
|
||||
logger.info(
|
||||
f"[REVEAL] 查看完整 Key (legacy vertex_ai): ID={self.key_id}, Name={key.name}"
|
||||
)
|
||||
return {"auth_type": "vertex_ai", "auth_config": decrypted_key}
|
||||
except InvalidRequestException:
|
||||
raise
|
||||
@@ -760,10 +762,14 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
|
||||
auto_fetch_models=self.key_data.auto_fetch_models,
|
||||
locked_models=self.key_data.locked_models if self.key_data.locked_models else None,
|
||||
model_include_patterns=(
|
||||
self.key_data.model_include_patterns if self.key_data.model_include_patterns else None
|
||||
self.key_data.model_include_patterns
|
||||
if self.key_data.model_include_patterns
|
||||
else None
|
||||
),
|
||||
model_exclude_patterns=(
|
||||
self.key_data.model_exclude_patterns if self.key_data.model_exclude_patterns else None
|
||||
self.key_data.model_exclude_patterns
|
||||
if self.key_data.model_exclude_patterns
|
||||
else None
|
||||
),
|
||||
request_count=0,
|
||||
success_count=0,
|
||||
|
||||
@@ -4,10 +4,10 @@ ProviderEndpoint CRUD 管理 API
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import and_
|
||||
@@ -15,13 +15,14 @@ from sqlalchemy.orm import Session
|
||||
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.core.api_format.signature import parse_signature_key
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.models.endpoint_models import (
|
||||
ProviderEndpointCreate,
|
||||
ProviderEndpointResponse,
|
||||
@@ -243,7 +244,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
|
||||
total_keys_map: dict[str, int] = {}
|
||||
active_keys_map: dict[str, int] = {}
|
||||
for api_formats, is_active in keys:
|
||||
for fmt in (api_formats or []):
|
||||
for fmt in api_formats or []:
|
||||
total_keys_map[fmt] = total_keys_map.get(fmt, 0) + 1
|
||||
if is_active:
|
||||
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
|
||||
@@ -299,11 +300,18 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
sig = parse_signature_key(self.endpoint_data.api_format)
|
||||
api_family = sig.api_family.value
|
||||
endpoint_kind = sig.endpoint_kind.value
|
||||
# 使用归一化后的 signature key,确保格式一致性
|
||||
normalized_api_format = sig.key
|
||||
|
||||
new_endpoint = ProviderEndpoint(
|
||||
id=str(uuid.uuid4()),
|
||||
provider_id=self.provider_id,
|
||||
api_format=self.endpoint_data.api_format,
|
||||
api_format=normalized_api_format,
|
||||
api_family=api_family,
|
||||
endpoint_kind=endpoint_kind,
|
||||
base_url=self.endpoint_data.base_url,
|
||||
custom_path=self.endpoint_data.custom_path,
|
||||
header_rules=self.endpoint_data.header_rules,
|
||||
@@ -323,7 +331,9 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
|
||||
# 清除 /v1/models 列表缓存
|
||||
await invalidate_models_list_cache()
|
||||
|
||||
logger.info(f"[OK] 创建 Endpoint: Provider={provider.name}, Format={self.endpoint_data.api_format}, ID={new_endpoint.id}")
|
||||
logger.info(
|
||||
f"[OK] 创建 Endpoint: Provider={provider.name}, Format={self.endpoint_data.api_format}, ID={new_endpoint.id}"
|
||||
)
|
||||
|
||||
endpoint_dict = {
|
||||
k: v
|
||||
@@ -417,6 +427,11 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
|
||||
# proxy 为 None 时保留,用于清除代理配置
|
||||
for field, value in update_data.items():
|
||||
setattr(endpoint, field, value)
|
||||
|
||||
# Phase 3/4: 自动维护新架构字段,确保新增/历史数据都能被调度器按 family/kind 查询
|
||||
sig = parse_signature_key(endpoint.api_format)
|
||||
endpoint.api_family = sig.api_family.value
|
||||
endpoint.endpoint_kind = sig.endpoint_kind.value
|
||||
endpoint.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
db.commit()
|
||||
@@ -426,10 +441,14 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
|
||||
await invalidate_models_list_cache()
|
||||
|
||||
provider = db.query(Provider).filter(Provider.id == endpoint.provider_id).first()
|
||||
logger.info(f"[OK] 更新 Endpoint: ID={self.endpoint_id}, Updates={list(update_data.keys())}")
|
||||
logger.info(
|
||||
f"[OK] 更新 Endpoint: ID={self.endpoint_id}, Updates={list(update_data.keys())}"
|
||||
)
|
||||
|
||||
endpoint_format = (
|
||||
endpoint.api_format if isinstance(endpoint.api_format, str) else endpoint.api_format.value
|
||||
endpoint.api_format
|
||||
if isinstance(endpoint.api_format, str)
|
||||
else endpoint.api_format.value
|
||||
)
|
||||
keys = (
|
||||
db.query(ProviderAPIKey.api_formats, ProviderAPIKey.is_active)
|
||||
@@ -472,7 +491,9 @@ class AdminDeleteProviderEndpointAdapter(AdminApiAdapter):
|
||||
raise NotFoundException(f"Endpoint {self.endpoint_id} 不存在")
|
||||
|
||||
endpoint_format = (
|
||||
endpoint.api_format if isinstance(endpoint.api_format, str) else endpoint.api_format.value
|
||||
endpoint.api_format
|
||||
if isinstance(endpoint.api_format, str)
|
||||
else endpoint.api_format.value
|
||||
)
|
||||
|
||||
# 查询包含该格式的所有 Key,并从 api_formats 中移除该格式
|
||||
@@ -488,7 +509,7 @@ class AdminDeleteProviderEndpointAdapter(AdminApiAdapter):
|
||||
# 移除该格式
|
||||
new_formats = [f for f in key.api_formats if f != endpoint_format]
|
||||
key.api_formats = new_formats if new_formats else []
|
||||
flag_modified(key, 'api_formats')
|
||||
flag_modified(key, "api_formats")
|
||||
|
||||
db.delete(endpoint)
|
||||
db.commit()
|
||||
|
||||
@@ -10,6 +10,7 @@ from pydantic import BaseModel, Field, ValidationError, field_validator
|
||||
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.core.crypto import crypto_service
|
||||
from src.core.enums import AuthSource
|
||||
@@ -17,7 +18,6 @@ from src.core.exceptions import InvalidRequestException, translate_pydantic_erro
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import AuditEventType, LDAPConfig, User, UserRole
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.services.system.audit import AuditService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/ldap", tags=["Admin - LDAP"])
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -222,9 +222,7 @@ class AdminGetManagementTokenAdapter(AdminManagementTokenApiAdapter):
|
||||
token_id: str = ""
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
token = ManagementTokenService.get_token_by_id(
|
||||
db=context.db, token_id=self.token_id
|
||||
)
|
||||
token = ManagementTokenService.get_token_by_id(db=context.db, token_id=self.token_id)
|
||||
|
||||
if not token:
|
||||
raise NotFoundException("Management Token 不存在")
|
||||
@@ -245,9 +243,7 @@ class AdminDeleteManagementTokenAdapter(AdminManagementTokenApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
# 先获取 token 信息用于审计
|
||||
token = ManagementTokenService.get_token_by_id(
|
||||
db=context.db, token_id=self.token_id
|
||||
)
|
||||
token = ManagementTokenService.get_token_by_id(db=context.db, token_id=self.token_id)
|
||||
|
||||
if not token:
|
||||
raise NotFoundException("Management Token 不存在")
|
||||
@@ -258,9 +254,7 @@ class AdminDeleteManagementTokenAdapter(AdminManagementTokenApiAdapter):
|
||||
owner_user_id=token.user_id,
|
||||
)
|
||||
|
||||
success = ManagementTokenService.delete_token(
|
||||
db=context.db, token_id=self.token_id
|
||||
)
|
||||
success = ManagementTokenService.delete_token(db=context.db, token_id=self.token_id)
|
||||
|
||||
if not success:
|
||||
raise NotFoundException("Management Token 不存在")
|
||||
@@ -277,9 +271,7 @@ class AdminToggleManagementTokenAdapter(AdminManagementTokenApiAdapter):
|
||||
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_UPDATED
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
token = ManagementTokenService.toggle_status(
|
||||
db=context.db, token_id=self.token_id
|
||||
)
|
||||
token = ManagementTokenService.toggle_status(db=context.db, token_id=self.token_id)
|
||||
|
||||
if not token:
|
||||
raise NotFoundException("Management Token 不存在")
|
||||
|
||||
@@ -4,17 +4,17 @@
|
||||
基于 GlobalModel 的聚合视图
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
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.database import get_db
|
||||
from src.models.database import GlobalModel, Model
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.models.pydantic_models import (
|
||||
ModelCapabilities,
|
||||
ModelCatalogItem,
|
||||
|
||||
@@ -6,13 +6,14 @@ GlobalModel Admin API
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request, Response
|
||||
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.core.logger import logger
|
||||
@@ -29,7 +30,6 @@ from src.models.pydantic_models import (
|
||||
ModelCatalogProviderDetail,
|
||||
)
|
||||
from src.services.model.global_model import GlobalModelService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
@@ -17,6 +17,7 @@ from pydantic import BaseModel, ConfigDict, Field
|
||||
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.core.crypto import CryptoService
|
||||
from src.core.model_permissions import (
|
||||
@@ -32,7 +33,6 @@ from src.models.database import (
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
@@ -49,7 +49,9 @@ class RoutingKeyInfo(BaseModel):
|
||||
name: str
|
||||
masked_key: str = Field("", description="脱敏的 API Key")
|
||||
internal_priority: int = Field(..., description="Key 内部优先级")
|
||||
global_priority_by_format: dict[str, int] | None = Field(None, description="按 API 格式的全局优先级")
|
||||
global_priority_by_format: dict[str, int] | None = Field(
|
||||
None, description="按 API 格式的全局优先级"
|
||||
)
|
||||
rpm_limit: int | None = Field(None, description="RPM 限制,null 表示自适应")
|
||||
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
|
||||
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
|
||||
@@ -320,11 +322,13 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
|
||||
# 按优先级排序(使用当前格式的全局优先级)
|
||||
api_format = ep.api_format or ""
|
||||
|
||||
def get_key_priority(k: ProviderAPIKey) -> tuple[int, int]:
|
||||
format_priority = 999
|
||||
if k.global_priority_by_format and api_format in k.global_priority_by_format:
|
||||
format_priority = k.global_priority_by_format[api_format]
|
||||
return (format_priority, k.internal_priority or 0)
|
||||
|
||||
ep_keys.sort(key=get_key_priority)
|
||||
|
||||
key_infos = []
|
||||
@@ -414,11 +418,21 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
)
|
||||
)
|
||||
|
||||
# 按 APIFormat 枚举定义的顺序排序 Endpoints
|
||||
from src.core.api_format import APIFormat
|
||||
|
||||
format_order = {fmt.value: i for i, fmt in enumerate(APIFormat)}
|
||||
endpoint_infos.sort(key=lambda e: format_order.get(e.api_format, 999))
|
||||
# 按 endpoint signature 的推荐顺序排序 Endpoints(与前端展示保持一致)
|
||||
preferred_order = [
|
||||
"openai:chat",
|
||||
"openai:cli",
|
||||
"openai:video",
|
||||
"claude:chat",
|
||||
"claude:cli",
|
||||
"gemini:chat",
|
||||
"gemini:cli",
|
||||
"gemini:video",
|
||||
]
|
||||
order_map = {key: i for i, key in enumerate(preferred_order)}
|
||||
endpoint_infos.sort(
|
||||
key=lambda e: order_map.get(str(e.api_format or "").strip().lower(), 999)
|
||||
)
|
||||
|
||||
active_endpoints = sum(1 for e in endpoint_infos if e.is_active)
|
||||
provider_infos.append(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""模块管理 API 端点"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
@@ -9,10 +10,10 @@ from pydantic import BaseModel
|
||||
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.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.modules import ModuleStatus, get_module_registry
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/admin/modules", tags=["Admin - Modules"])
|
||||
|
||||
@@ -2,15 +2,16 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
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, paginate_query
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.logger import logger
|
||||
@@ -26,8 +27,6 @@ from src.models.database import User as DBUser
|
||||
from src.services.health.monitor import HealthMonitor
|
||||
from src.services.system.audit import audit_service
|
||||
from src.utils.database_helpers import escape_like_pattern
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
|
||||
router = APIRouter(prefix="/api/admin/monitoring", tags=["Admin - Monitoring"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
@@ -776,9 +776,9 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
global_model_map: dict[str, GlobalModel] = {}
|
||||
if global_model_ids:
|
||||
# model_name 可能是 UUID 格式的 global_model_id,也可能是原始模型名称
|
||||
global_models = db.query(GlobalModel).filter(
|
||||
GlobalModel.id.in_(list(global_model_ids))
|
||||
).all()
|
||||
global_models = (
|
||||
db.query(GlobalModel).filter(GlobalModel.id.in_(list(global_model_ids))).all()
|
||||
)
|
||||
global_model_map = {str(gm.id): gm for gm in global_models}
|
||||
|
||||
keyword_lower = self.keyword.lower() if self.keyword else None
|
||||
@@ -830,12 +830,14 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
"global_model_id": affinity.get("model_name"), # 原始的 global_model_id
|
||||
"model_name": (
|
||||
global_model_map.get(affinity.get("model_name")).name
|
||||
if affinity.get("model_name") and global_model_map.get(affinity.get("model_name"))
|
||||
if affinity.get("model_name")
|
||||
and global_model_map.get(affinity.get("model_name"))
|
||||
else affinity.get("model_name") # 如果找不到 GlobalModel,显示原始值
|
||||
),
|
||||
"model_display_name": (
|
||||
global_model_map.get(affinity.get("model_name")).display_name
|
||||
if affinity.get("model_name") and global_model_map.get(affinity.get("model_name"))
|
||||
if affinity.get("model_name")
|
||||
and global_model_map.get(affinity.get("model_name"))
|
||||
else None
|
||||
),
|
||||
"api_format": affinity.get("api_format"),
|
||||
@@ -916,7 +918,9 @@ class AdminClearUserCacheAdapter(AdminApiAdapter):
|
||||
)
|
||||
count += 1
|
||||
|
||||
logger.info(f"已清除API Key缓存亲和性: api_key_name={api_key.name}, affinity_key={affinity_key[:8]}..., 清除数量={count}")
|
||||
logger.info(
|
||||
f"已清除API Key缓存亲和性: api_key_name={api_key.name}, affinity_key={affinity_key[:8]}..., 清除数量={count}"
|
||||
)
|
||||
|
||||
response = {
|
||||
"status": "ok",
|
||||
@@ -969,7 +973,9 @@ class AdminClearUserCacheAdapter(AdminApiAdapter):
|
||||
)
|
||||
count += 1
|
||||
|
||||
logger.info(f"已清除用户缓存亲和性: username={user.username}, user_id={user_id[:8]}..., 清除数量={count}")
|
||||
logger.info(
|
||||
f"已清除用户缓存亲和性: username={user.username}, user_id={user_id[:8]}..., 清除数量={count}"
|
||||
)
|
||||
|
||||
response = {
|
||||
"status": "ok",
|
||||
@@ -1075,7 +1081,9 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
count = await affinity_mgr.invalidate_all_for_provider(self.provider_id)
|
||||
logger.info(f"已清除Provider缓存亲和性: provider_id={self.provider_id[:8]}..., count={count}")
|
||||
logger.info(
|
||||
f"已清除Provider缓存亲和性: provider_id={self.provider_id[:8]}..., count={count}"
|
||||
)
|
||||
context.add_audit_metadata(
|
||||
action="cache_clear_provider",
|
||||
provider_id=self.provider_id,
|
||||
@@ -1094,8 +1102,8 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
|
||||
class AdminCacheConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||
from src.config.constants import ConcurrencyDefaults
|
||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
|
||||
# 获取动态预留管理器的配置
|
||||
@@ -1334,11 +1342,13 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
if cached_str == "NOT_FOUND":
|
||||
unmapped_entries.append({
|
||||
"mapping_name": mapping_name,
|
||||
"status": "not_found",
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
})
|
||||
unmapped_entries.append(
|
||||
{
|
||||
"mapping_name": mapping_name,
|
||||
"status": "not_found",
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
}
|
||||
)
|
||||
else:
|
||||
try:
|
||||
cached_data = json.loads(cached_str)
|
||||
@@ -1380,27 +1390,33 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
provider_names.append(provider.name)
|
||||
provider_names = sorted(list(set(provider_names)))
|
||||
|
||||
mappings.append({
|
||||
"mapping_name": mapping_name,
|
||||
"global_model_name": global_model_name,
|
||||
"global_model_display_name": global_model_display_name,
|
||||
"providers": provider_names,
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
})
|
||||
mappings.append(
|
||||
{
|
||||
"mapping_name": mapping_name,
|
||||
"global_model_name": global_model_name,
|
||||
"global_model_display_name": global_model_display_name,
|
||||
"providers": provider_names,
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
}
|
||||
)
|
||||
|
||||
except json.JSONDecodeError:
|
||||
unmapped_entries.append({
|
||||
"mapping_name": mapping_name,
|
||||
"status": "invalid",
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
})
|
||||
unmapped_entries.append(
|
||||
{
|
||||
"mapping_name": mapping_name,
|
||||
"status": "invalid",
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"解析缓存键 {key} 失败: {e}")
|
||||
unmapped_entries.append({
|
||||
"mapping_name": mapping_name,
|
||||
"status": "error",
|
||||
"ttl": None,
|
||||
})
|
||||
unmapped_entries.append(
|
||||
{
|
||||
"mapping_name": mapping_name,
|
||||
"status": "error",
|
||||
"ttl": None,
|
||||
}
|
||||
)
|
||||
|
||||
# 按 mapping_name 排序
|
||||
mappings.sort(key=lambda x: x["mapping_name"])
|
||||
@@ -1408,8 +1424,13 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
# 3. 解析 provider_global 缓存(Provider 级别的模型解析缓存)
|
||||
provider_model_mappings = []
|
||||
# 预加载 Provider 和 GlobalModel 数据
|
||||
provider_map = {str(p.id): p for p in db.query(Provider).filter(Provider.is_active.is_(True)).all()}
|
||||
global_model_map = {str(gm.id): gm for gm in db.query(GlobalModel).filter(GlobalModel.is_active.is_(True)).all()}
|
||||
provider_map = {
|
||||
str(p.id): p for p in db.query(Provider).filter(Provider.is_active.is_(True)).all()
|
||||
}
|
||||
global_model_map = {
|
||||
str(gm.id): gm
|
||||
for gm in db.query(GlobalModel).filter(GlobalModel.is_active.is_(True)).all()
|
||||
}
|
||||
|
||||
for key in provider_global_keys[:100]: # 最多处理 100 个
|
||||
# key 格式: model:provider_global:{provider_id}:{global_model_id}
|
||||
@@ -1447,7 +1468,9 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
mapping_names = []
|
||||
if cached_model_mappings:
|
||||
for mapping_entry in cached_model_mappings:
|
||||
if isinstance(mapping_entry, dict) and mapping_entry.get("name"):
|
||||
if isinstance(mapping_entry, dict) and mapping_entry.get(
|
||||
"name"
|
||||
):
|
||||
mapping_names.append(mapping_entry["name"])
|
||||
|
||||
# provider_model_name 为空时跳过
|
||||
@@ -1463,19 +1486,23 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
if has_name_mapping or has_mappings:
|
||||
# 构建用于展示的映射列表
|
||||
# 如果只有名称映射没有额外映射,则用 global_model_name 作为"请求名称"
|
||||
display_mappings = mapping_names if mapping_names else [global_model.name]
|
||||
display_mappings = (
|
||||
mapping_names if mapping_names else [global_model.name]
|
||||
)
|
||||
|
||||
provider_model_mappings.append({
|
||||
"provider_id": provider_id,
|
||||
"provider_name": provider.name,
|
||||
"global_model_id": global_model_id,
|
||||
"global_model_name": global_model.name,
|
||||
"global_model_display_name": global_model.display_name,
|
||||
"provider_model_name": provider_model_name,
|
||||
"aliases": display_mappings,
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
"hit_count": hit_count,
|
||||
})
|
||||
provider_model_mappings.append(
|
||||
{
|
||||
"provider_id": provider_id,
|
||||
"provider_name": provider.name,
|
||||
"global_model_id": global_model_id,
|
||||
"global_model_name": global_model.name,
|
||||
"global_model_display_name": global_model.display_name,
|
||||
"provider_model_name": provider_model_name,
|
||||
"aliases": display_mappings,
|
||||
"ttl": ttl if ttl > 0 else None,
|
||||
"hit_count": hit_count,
|
||||
}
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
except Exception as e:
|
||||
@@ -1496,7 +1523,9 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
"global_model_resolve": len(global_model_resolve_keys),
|
||||
},
|
||||
"mappings": mappings,
|
||||
"provider_model_mappings": provider_model_mappings if provider_model_mappings else None,
|
||||
"provider_model_mappings": (
|
||||
provider_model_mappings if provider_model_mappings else None
|
||||
),
|
||||
"unmapped": unmapped_entries if unmapped_entries else None,
|
||||
}
|
||||
|
||||
|
||||
@@ -4,21 +4,21 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, ProviderEndpoint, ProviderAPIKey
|
||||
from src.core.crypto import crypto_service
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
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()
|
||||
@@ -184,8 +184,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
# 4. status="pending" 表示请求尚未开始执行
|
||||
# 5. status="cancelled" 表示客户端主动断开连接(不算失败)
|
||||
has_success = any(
|
||||
c.status == "success"
|
||||
or (c.status_code is not None and 200 <= c.status_code < 300)
|
||||
c.status == "success" or (c.status_code is not None and 200 <= c.status_code < 300)
|
||||
for c in candidates
|
||||
)
|
||||
has_streaming = any(c.status == "streaming" for c in candidates)
|
||||
@@ -221,7 +220,9 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
endpoint_ids = {c.endpoint_id for c in candidates if c.endpoint_id}
|
||||
endpoint_map = {}
|
||||
if endpoint_ids:
|
||||
endpoints = db.query(ProviderEndpoint).filter(ProviderEndpoint.id.in_(endpoint_ids)).all()
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint).filter(ProviderEndpoint.id.in_(endpoint_ids)).all()
|
||||
)
|
||||
endpoint_map = {e.id: e.api_format for e in endpoints}
|
||||
|
||||
# 批量加载 key 信息
|
||||
@@ -245,7 +246,9 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
prefix_end = len(prefix)
|
||||
break
|
||||
if prefix_end > 0:
|
||||
key_preview_map[k.id] = f"{decrypted_key[:prefix_end]}***{decrypted_key[-4:]}"
|
||||
key_preview_map[k.id] = (
|
||||
f"{decrypted_key[:prefix_end]}***{decrypted_key[-4:]}"
|
||||
)
|
||||
else:
|
||||
key_preview_map[k.id] = f"{decrypted_key[:4]}***{decrypted_key[-4:]}"
|
||||
elif len(decrypted_key) > 4:
|
||||
@@ -267,12 +270,8 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
endpoint_name = (
|
||||
endpoint_map.get(candidate.endpoint_id) if candidate.endpoint_id else None
|
||||
)
|
||||
key_name = (
|
||||
key_map.get(candidate.key_id) if candidate.key_id else None
|
||||
)
|
||||
key_preview = (
|
||||
key_preview_map.get(candidate.key_id) if candidate.key_id else None
|
||||
)
|
||||
key_name = key_map.get(candidate.key_id) if candidate.key_id else None
|
||||
key_preview = key_preview_map.get(candidate.key_id) if candidate.key_id else None
|
||||
key_capabilities = (
|
||||
key_capabilities_map.get(candidate.key_id) if candidate.key_id else None
|
||||
)
|
||||
|
||||
@@ -5,8 +5,8 @@ Provider Query API 端点
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import asyncio
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -14,11 +14,15 @@ from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.config.constants import TimeoutDefaults
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.api_format import get_extra_headers_from_endpoint
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.database.database import get_db
|
||||
from src.models.database import Provider, ProviderEndpoint, User
|
||||
from src.services.model.fetch_scheduler import (
|
||||
get_upstream_models_from_cache,
|
||||
set_upstream_models_to_cache,
|
||||
)
|
||||
from src.services.model.upstream_fetcher import (
|
||||
_get_adapter_for_format,
|
||||
build_all_format_configs,
|
||||
@@ -26,11 +30,6 @@ from src.services.model.upstream_fetcher import (
|
||||
)
|
||||
from src.utils.auth_utils import get_current_user
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
from src.services.model.fetch_scheduler import (
|
||||
get_upstream_models_from_cache,
|
||||
set_upstream_models_to_cache,
|
||||
)
|
||||
|
||||
|
||||
router = APIRouter(prefix="/api/admin/provider-query", tags=["Provider Query"])
|
||||
|
||||
@@ -124,9 +123,7 @@ async def query_available_models(
|
||||
async def fetch_for_key(api_key: Any) -> Any:
|
||||
# 非强制刷新时,先检查缓存
|
||||
if not request.force_refresh:
|
||||
cached_models = await get_upstream_models_from_cache(
|
||||
request.provider_id, api_key.id
|
||||
)
|
||||
cached_models = await get_upstream_models_from_cache(request.provider_id, api_key.id)
|
||||
if cached_models is not None:
|
||||
return cached_models, None, True # models, error, from_cache
|
||||
|
||||
@@ -252,10 +249,7 @@ async def _fetch_models_for_single_key(
|
||||
) -> Any:
|
||||
"""获取单个 Key 的模型列表"""
|
||||
# 查找指定的 Key
|
||||
api_key = next(
|
||||
(key for key in provider.api_keys if key.id == api_key_id),
|
||||
None
|
||||
)
|
||||
api_key = next((key for key in provider.api_keys if key.id == api_key_id), None)
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=404, detail="API Key not found")
|
||||
|
||||
@@ -347,19 +341,22 @@ async def test_model(
|
||||
if not endpoint:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"No active endpoint found for API format: {request.api_format}"
|
||||
detail=f"No active endpoint found for API format: {request.api_format}",
|
||||
)
|
||||
|
||||
if request.api_key_id:
|
||||
# 使用指定的 Key,但需要校验是否支持该格式
|
||||
api_key = next(
|
||||
(key for key in provider.api_keys if key.id == request.api_key_id and key.is_active),
|
||||
None
|
||||
(
|
||||
key
|
||||
for key in provider.api_keys
|
||||
if key.id == request.api_key_id and key.is_active
|
||||
),
|
||||
None,
|
||||
)
|
||||
if api_key and request.api_format not in (api_key.api_formats or []):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"API Key does not support format: {request.api_format}"
|
||||
status_code=400, detail=f"API Key does not support format: {request.api_format}"
|
||||
)
|
||||
else:
|
||||
# 找支持该格式的第一个可用 Key
|
||||
@@ -378,13 +375,17 @@ async def test_model(
|
||||
if request.api_key_id:
|
||||
# 同时指定了 Key,需要校验是否支持该端点格式
|
||||
api_key = next(
|
||||
(key for key in provider.api_keys if key.id == request.api_key_id and key.is_active),
|
||||
None
|
||||
(
|
||||
key
|
||||
for key in provider.api_keys
|
||||
if key.id == request.api_key_id and key.is_active
|
||||
),
|
||||
None,
|
||||
)
|
||||
if api_key and endpoint.api_format not in (api_key.api_formats or []):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"API Key does not support endpoint format: {endpoint.api_format}"
|
||||
detail=f"API Key does not support endpoint format: {endpoint.api_format}",
|
||||
)
|
||||
else:
|
||||
# 找支持该端点格式的第一个可用 Key
|
||||
@@ -398,11 +399,11 @@ async def test_model(
|
||||
# 使用指定的 API Key
|
||||
api_key = next(
|
||||
(key for key in provider.api_keys if key.id == request.api_key_id and key.is_active),
|
||||
None
|
||||
None,
|
||||
)
|
||||
if api_key:
|
||||
# 找到该 Key 支持的第一个活跃 Endpoint
|
||||
for fmt in (api_key.api_formats or []):
|
||||
for fmt in api_key.api_formats or []:
|
||||
if fmt in format_to_endpoint:
|
||||
endpoint = format_to_endpoint[fmt]
|
||||
break
|
||||
@@ -469,7 +470,9 @@ async def test_model(
|
||||
}
|
||||
|
||||
# 发送测试请求
|
||||
async with httpx.AsyncClient(timeout=endpoint_config["timeout"], verify=get_ssl_context()) as client:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=endpoint_config["timeout"], verify=get_ssl_context()
|
||||
) as client:
|
||||
# 非流式测试
|
||||
logger.debug(f"[test-model] 开始非流式测试...")
|
||||
|
||||
@@ -492,46 +495,47 @@ async def test_model(
|
||||
logger.debug(f"[test-model] 非流式测试结果:")
|
||||
logger.debug(f"[test-model] Status Code: {response.get('status_code')}")
|
||||
logger.debug(f"[test-model] Response Headers: {response.get('headers', {})}")
|
||||
response_data = response.get('response', {})
|
||||
response_body = response_data.get('response_body', {})
|
||||
response_data = response.get("response", {})
|
||||
response_body = response_data.get("response_body", {})
|
||||
logger.debug(f"[test-model] Response Data: {response_data}")
|
||||
logger.debug(f"[test-model] Response Body: {response_body}")
|
||||
# 尝试解析 response_body (通常是 JSON 字符串)
|
||||
parsed_body = response_body
|
||||
import json
|
||||
|
||||
if isinstance(response_body, str):
|
||||
try:
|
||||
parsed_body = json.loads(response_body)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
if isinstance(parsed_body, dict) and 'error' in parsed_body:
|
||||
error_obj = parsed_body['error']
|
||||
if isinstance(parsed_body, dict) and "error" in parsed_body:
|
||||
error_obj = parsed_body["error"]
|
||||
# 兼容 error 可能是字典或字符串的情况
|
||||
if isinstance(error_obj, dict):
|
||||
logger.debug(f"[test-model] Error Message: {error_obj.get('message')}")
|
||||
raise HTTPException(status_code=500, detail=error_obj.get('message'))
|
||||
raise HTTPException(status_code=500, detail=error_obj.get("message"))
|
||||
else:
|
||||
logger.debug(f"[test-model] Error: {error_obj}")
|
||||
raise HTTPException(status_code=500, detail=error_obj)
|
||||
elif 'error' in response:
|
||||
elif "error" in response:
|
||||
logger.debug(f"[test-model] Error: {response['error']}")
|
||||
raise HTTPException(status_code=500, detail=response['error'])
|
||||
raise HTTPException(status_code=500, detail=response["error"])
|
||||
else:
|
||||
# 如果有选择或消息,记录内容预览
|
||||
if isinstance(response_data, dict):
|
||||
if 'choices' in response_data and response_data['choices']:
|
||||
choice = response_data['choices'][0]
|
||||
if 'message' in choice:
|
||||
content = choice['message'].get('content', '')
|
||||
if "choices" in response_data and response_data["choices"]:
|
||||
choice = response_data["choices"][0]
|
||||
if "message" in choice:
|
||||
content = choice["message"].get("content", "")
|
||||
logger.debug(f"[test-model] Content Preview: {content[:200]}...")
|
||||
elif 'content' in response_data and response_data['content']:
|
||||
content = str(response_data['content'])
|
||||
elif "content" in response_data and response_data["content"]:
|
||||
content = str(response_data["content"])
|
||||
logger.debug(f"[test-model] Content Preview: {content[:200]}...")
|
||||
|
||||
# 检查测试是否成功(基于HTTP状态码)
|
||||
status_code = response.get('status_code', 0)
|
||||
is_success = status_code == 200 and 'error' not in response
|
||||
status_code = response.get("status_code", 0)
|
||||
is_success = status_code == 200 and "error" not in response
|
||||
|
||||
return {
|
||||
"success": is_success,
|
||||
@@ -561,9 +565,13 @@ async def test_model(
|
||||
"name": provider.name,
|
||||
},
|
||||
"model": request.model_name,
|
||||
"endpoint": {
|
||||
"id": endpoint.id,
|
||||
"api_format": endpoint.api_format,
|
||||
"base_url": endpoint.base_url,
|
||||
} if endpoint else None,
|
||||
"endpoint": (
|
||||
{
|
||||
"id": endpoint.id,
|
||||
"api_format": endpoint.api_format,
|
||||
"base_url": endpoint.base_url,
|
||||
}
|
||||
if endpoint
|
||||
else None
|
||||
),
|
||||
}
|
||||
|
||||
@@ -4,8 +4,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -13,13 +13,13 @@ from pydantic import BaseModel, Field, ValidationError
|
||||
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.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.models.database_extensions import ProviderUsageTracking
|
||||
|
||||
router = APIRouter(prefix="/api/admin/provider-strategy", tags=["Provider Strategy"])
|
||||
@@ -187,7 +187,9 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
.scalar()
|
||||
)
|
||||
provider.monthly_used_usd = float(period_usage or 0)
|
||||
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
|
||||
logger.info(
|
||||
f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}"
|
||||
)
|
||||
|
||||
if config.quota_expires_at:
|
||||
expires_at = datetime.fromisoformat(config.quota_expires_at)
|
||||
|
||||
@@ -11,6 +11,7 @@ from fastapi import APIRouter, Depends, Request
|
||||
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.core.exceptions import InvalidRequestException, NotFoundException
|
||||
@@ -21,23 +22,22 @@ from src.models.api import (
|
||||
ModelResponse,
|
||||
ModelUpdate,
|
||||
)
|
||||
from src.models.pydantic_models import (
|
||||
BatchAssignModelsToProviderRequest,
|
||||
BatchAssignModelsToProviderResponse,
|
||||
ImportFromUpstreamRequest,
|
||||
ImportFromUpstreamResponse,
|
||||
ImportFromUpstreamSuccessItem,
|
||||
ImportFromUpstreamErrorItem,
|
||||
ProviderAvailableSourceModel,
|
||||
ProviderAvailableSourceModelsResponse,
|
||||
)
|
||||
from src.models.database import (
|
||||
GlobalModel,
|
||||
Model,
|
||||
Provider,
|
||||
)
|
||||
from src.models.pydantic_models import (
|
||||
BatchAssignModelsToProviderRequest,
|
||||
BatchAssignModelsToProviderResponse,
|
||||
ImportFromUpstreamErrorItem,
|
||||
ImportFromUpstreamRequest,
|
||||
ImportFromUpstreamResponse,
|
||||
ImportFromUpstreamSuccessItem,
|
||||
ProviderAvailableSourceModel,
|
||||
ProviderAvailableSourceModelsResponse,
|
||||
)
|
||||
from src.services.model.service import ModelService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(tags=["Model Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -322,9 +322,7 @@ async def batch_assign_global_models_to_provider(
|
||||
- `global_model_name`: 全局模型名称(如果可用)
|
||||
- `error`: 错误信息
|
||||
"""
|
||||
adapter = AdminBatchAssignModelsToProviderAdapter(
|
||||
provider_id=provider_id, payload=payload
|
||||
)
|
||||
adapter = AdminBatchAssignModelsToProviderAdapter(provider_id=provider_id, payload=payload)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@@ -407,7 +405,9 @@ class AdminCreateProviderModelAdapter(AdminApiAdapter):
|
||||
|
||||
try:
|
||||
model = ModelService.create_model(db, self.provider_id, self.model_data)
|
||||
logger.info(f"Model created: {model.provider_model_name} for provider {provider.name} by {context.user.username}")
|
||||
logger.info(
|
||||
f"Model created: {model.provider_model_name} for provider {provider.name} by {context.user.username}"
|
||||
)
|
||||
# 缓存失效已在 ModelService.create_model 中处理
|
||||
return ModelService.convert_to_response(model)
|
||||
except Exception as exc:
|
||||
@@ -450,7 +450,9 @@ class AdminUpdateProviderModelAdapter(AdminApiAdapter):
|
||||
|
||||
try:
|
||||
updated_model = ModelService.update_model(db, self.model_id, self.model_data)
|
||||
logger.info(f"Model updated: {updated_model.provider_model_name} by {context.user.username}")
|
||||
logger.info(
|
||||
f"Model updated: {updated_model.provider_model_name} by {context.user.username}"
|
||||
)
|
||||
# 缓存失效已在 ModelService.update_model 中处理
|
||||
return ModelService.convert_to_response(updated_model)
|
||||
except Exception as exc:
|
||||
@@ -495,7 +497,9 @@ class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
||||
|
||||
try:
|
||||
models = ModelService.batch_create_models(db, self.provider_id, self.models_data)
|
||||
logger.info(f"Batch created {len(models)} models for provider {provider.name} by {context.user.username}")
|
||||
logger.info(
|
||||
f"Batch created {len(models)} models for provider {provider.name} by {context.user.username}"
|
||||
)
|
||||
# 缓存失效已在 ModelService.batch_create_models 中处理
|
||||
return [ModelService.convert_to_response(model) for model in models]
|
||||
except Exception as exc:
|
||||
@@ -642,6 +646,7 @@ class AdminBatchAssignModelsToProviderAdapter(AdminApiAdapter):
|
||||
if success:
|
||||
# Provider 新增模型实现后,清除同进程的 ModelMapper 缓存,避免 TTL 内仍返回 None
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
|
||||
cache_service = get_cache_invalidation_service()
|
||||
cache_service.on_model_changed(self.provider_id, success[0].get("global_model_id", ""))
|
||||
|
||||
@@ -669,9 +674,12 @@ class AdminImportFromUpstreamAdapter(AdminApiAdapter):
|
||||
# 获取价格覆盖配置
|
||||
tiered_pricing = None
|
||||
price_per_request = None
|
||||
if hasattr(self.payload, 'tiered_pricing') and self.payload.tiered_pricing:
|
||||
if hasattr(self.payload, "tiered_pricing") and self.payload.tiered_pricing:
|
||||
tiered_pricing = self.payload.tiered_pricing
|
||||
if hasattr(self.payload, 'price_per_request') and self.payload.price_per_request is not None:
|
||||
if (
|
||||
hasattr(self.payload, "price_per_request")
|
||||
and self.payload.price_per_request is not None
|
||||
):
|
||||
price_per_request = self.payload.price_per_request
|
||||
|
||||
for model_id in self.payload.model_ids:
|
||||
@@ -679,7 +687,11 @@ class AdminImportFromUpstreamAdapter(AdminApiAdapter):
|
||||
if not model_id or len(model_id) > 100:
|
||||
errors.append(
|
||||
ImportFromUpstreamErrorItem(
|
||||
model_id=model_id[:50] + "..." if model_id and len(model_id) > 50 else model_id or "<empty>",
|
||||
model_id=(
|
||||
model_id[:50] + "..."
|
||||
if model_id and len(model_id) > 50
|
||||
else model_id or "<empty>"
|
||||
),
|
||||
error="Invalid model_id: must be 1-100 characters",
|
||||
)
|
||||
)
|
||||
@@ -705,7 +717,9 @@ class AdminImportFromUpstreamAdapter(AdminApiAdapter):
|
||||
ImportFromUpstreamSuccessItem(
|
||||
model_id=model_id,
|
||||
global_model_id=existing.global_model_id or "",
|
||||
global_model_name=existing.global_model.name if existing.global_model else "",
|
||||
global_model_name=(
|
||||
existing.global_model.name if existing.global_model else ""
|
||||
),
|
||||
provider_model_id=existing.id,
|
||||
created_global_model=False,
|
||||
)
|
||||
|
||||
@@ -2,17 +2,17 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
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.services.cache.model_cache import ModelCacheService
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
@@ -21,8 +21,8 @@ from src.core.model_permissions import match_model_with_pattern, parse_allowed_m
|
||||
from src.database import get_db
|
||||
from src.models.admin_requests import CreateProviderRequest, UpdateProviderRequest
|
||||
from src.models.database import GlobalModel, Provider, ProviderAPIKey
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
from src.services.cache.provider_cache import ProviderCacheService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(tags=["Provider CRUD"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -159,7 +159,9 @@ async def create_provider(request: Request, db: Session = Depends(get_db)) -> An
|
||||
|
||||
|
||||
@router.put("/{provider_id}")
|
||||
async def update_provider(provider_id: str, request: Request, db: Session = Depends(get_db)) -> None:
|
||||
async def update_provider(
|
||||
provider_id: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> None:
|
||||
"""
|
||||
更新提供商配置
|
||||
|
||||
@@ -195,7 +197,9 @@ async def update_provider(provider_id: str, request: Request, db: Session = Depe
|
||||
|
||||
|
||||
@router.delete("/{provider_id}")
|
||||
async def delete_provider(provider_id: str, request: Request, db: Session = Depends(get_db)) -> None:
|
||||
async def delete_provider(
|
||||
provider_id: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> None:
|
||||
"""
|
||||
删除提供商
|
||||
|
||||
|
||||
@@ -2,15 +2,16 @@
|
||||
Provider 摘要与健康监控 API
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import case, func
|
||||
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.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import NotFoundException
|
||||
@@ -23,7 +24,6 @@ from src.models.database import (
|
||||
ProviderEndpoint,
|
||||
RequestCandidate,
|
||||
)
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.models.endpoint_models import (
|
||||
EndpointHealthEvent,
|
||||
EndpointHealthMonitor,
|
||||
@@ -229,10 +229,14 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
|
||||
active_keys = int(key_stats.active or 0)
|
||||
|
||||
# Model 统计(合并为单个查询)
|
||||
model_stats = db.query(
|
||||
func.count(Model.id).label("total"),
|
||||
func.sum(case((Model.is_active == True, 1), else_=0)).label("active"),
|
||||
).filter(Model.provider_id == provider.id).first()
|
||||
model_stats = (
|
||||
db.query(
|
||||
func.count(Model.id).label("total"),
|
||||
func.sum(case((Model.is_active == True, 1), else_=0)).label("active"),
|
||||
)
|
||||
.filter(Model.provider_id == provider.id)
|
||||
.first()
|
||||
)
|
||||
total_models = model_stats.total or 0
|
||||
active_models = int(model_stats.active or 0)
|
||||
|
||||
@@ -294,7 +298,9 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
|
||||
# 检查是否配置了 Provider Ops(余额监控等)
|
||||
provider_ops_config = (provider.config or {}).get("provider_ops")
|
||||
ops_configured = bool(provider_ops_config)
|
||||
ops_architecture_id = provider_ops_config.get("architecture_id") if provider_ops_config else None
|
||||
ops_architecture_id = (
|
||||
provider_ops_config.get("architecture_id") if provider_ops_config else None
|
||||
)
|
||||
|
||||
return ProviderWithEndpointsSummary(
|
||||
id=provider.id,
|
||||
|
||||
@@ -7,17 +7,18 @@ IP 安全管理接口
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
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.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.database import get_db
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(prefix="/api/admin/security/ip", tags=["Admin - Security"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -78,7 +79,9 @@ async def add_to_blacklist(request: Request, db: Session = Depends(get_db)) -> N
|
||||
|
||||
|
||||
@router.delete("/blacklist/{ip_address}")
|
||||
async def remove_from_blacklist(ip_address: str, request: Request, db: Session = Depends(get_db)) -> None:
|
||||
async def remove_from_blacklist(
|
||||
ip_address: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> None:
|
||||
"""
|
||||
从黑名单移除 IP
|
||||
|
||||
@@ -133,7 +136,9 @@ async def add_to_whitelist(request: Request, db: Session = Depends(get_db)) -> N
|
||||
|
||||
|
||||
@router.delete("/whitelist/{ip_address}")
|
||||
async def remove_from_whitelist(ip_address: str, request: Request, db: Session = Depends(get_db)) -> None:
|
||||
async def remove_from_whitelist(
|
||||
ip_address: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> None:
|
||||
"""
|
||||
从白名单移除 IP
|
||||
|
||||
|
||||
@@ -1,17 +1,17 @@
|
||||
"""系统设置API端点。"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
import copy
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import ValidationError
|
||||
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.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
@@ -20,7 +20,6 @@ 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.system.config import SystemConfigService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(prefix="/api/admin/system", tags=["Admin - System"])
|
||||
|
||||
@@ -179,7 +178,9 @@ async def check_update() -> Any:
|
||||
tag_commit_sha = latest_tag_info.get("commit", {}).get("sha")
|
||||
if tag_commit_sha:
|
||||
# 尝试获取 annotated tag 的信息
|
||||
tag_ref_url = f"https://api.github.com/repos/{github_repo}/git/refs/tags/{latest_tag_name}"
|
||||
tag_ref_url = (
|
||||
f"https://api.github.com/repos/{github_repo}/git/refs/tags/{latest_tag_name}"
|
||||
)
|
||||
ref_response = await client.get(
|
||||
tag_ref_url,
|
||||
headers={
|
||||
@@ -210,7 +211,9 @@ async def check_update() -> Any:
|
||||
|
||||
# 如果没有获取到时间,从 commit 获取
|
||||
if not published_at:
|
||||
commit_url = f"https://api.github.com/repos/{github_repo}/commits/{tag_commit_sha}"
|
||||
commit_url = (
|
||||
f"https://api.github.com/repos/{github_repo}/commits/{tag_commit_sha}"
|
||||
)
|
||||
commit_response = await client.get(
|
||||
commit_url,
|
||||
headers={
|
||||
@@ -220,7 +223,9 @@ async def check_update() -> Any:
|
||||
)
|
||||
if commit_response.status_code == 200:
|
||||
commit_data = commit_response.json()
|
||||
published_at = commit_data.get("commit", {}).get("committer", {}).get("date")
|
||||
published_at = (
|
||||
commit_data.get("commit", {}).get("committer", {}).get("date")
|
||||
)
|
||||
|
||||
return {
|
||||
"current_version": current_version,
|
||||
@@ -599,6 +604,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
|
||||
# 对敏感配置进行加密
|
||||
if self.key in self.ENCRYPTED_KEYS and value:
|
||||
from src.core.crypto import crypto_service
|
||||
|
||||
value = crypto_service.encrypt(value)
|
||||
|
||||
config = SystemConfigService.set_config(
|
||||
@@ -653,7 +659,6 @@ class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
||||
"""手动触发清理任务"""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
|
||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||
|
||||
db = context.db
|
||||
@@ -714,21 +719,41 @@ class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
||||
class AdminGetApiFormatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
"""获取所有可用的API格式"""
|
||||
from src.core.api_format import API_FORMAT_DEFINITIONS, APIFormat
|
||||
from src.core.api_format import list_endpoint_definitions
|
||||
|
||||
_ = context # 参数保留以符合接口规范
|
||||
|
||||
formats = []
|
||||
for api_format in APIFormat:
|
||||
definition = API_FORMAT_DEFINITIONS.get(api_format)
|
||||
formats.append(
|
||||
{
|
||||
"value": api_format.value,
|
||||
"label": api_format.value,
|
||||
"default_path": definition.default_path if definition else "/",
|
||||
"aliases": list(definition.aliases) if definition else [],
|
||||
}
|
||||
)
|
||||
def _label_for(sig: str) -> str:
|
||||
fam, kind = (sig.split(":", 1) + [""])[:2]
|
||||
fam_title = {"claude": "Claude", "openai": "OpenAI", "gemini": "Gemini"}.get(fam, fam)
|
||||
if kind == "chat":
|
||||
return fam_title
|
||||
kind_title = {"cli": "CLI", "video": "Video", "image": "Image"}.get(kind, kind)
|
||||
return f"{fam_title} {kind_title}".strip()
|
||||
|
||||
endpoint_defs = list_endpoint_definitions()
|
||||
preferred_order = [
|
||||
"openai:chat",
|
||||
"openai:cli",
|
||||
"openai:video",
|
||||
"claude:chat",
|
||||
"claude:cli",
|
||||
"gemini:chat",
|
||||
"gemini:cli",
|
||||
"gemini:video",
|
||||
]
|
||||
order_map = {key: i for i, key in enumerate(preferred_order)}
|
||||
endpoint_defs.sort(key=lambda d: order_map.get(d.signature_key, 999))
|
||||
|
||||
formats = [
|
||||
{
|
||||
"value": d.signature_key,
|
||||
"label": _label_for(d.signature_key),
|
||||
"default_path": d.default_path,
|
||||
"aliases": list(d.aliases or []),
|
||||
}
|
||||
for d in endpoint_defs
|
||||
]
|
||||
|
||||
return {"formats": formats}
|
||||
|
||||
@@ -738,8 +763,14 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
# Provider Ops 中需要解密的敏感字段
|
||||
SENSITIVE_CREDENTIALS = {
|
||||
"api_key", "password", "session_token", "session_cookie",
|
||||
"token_cookie", "auth_cookie", "cookie_string", "cookie"
|
||||
"api_key",
|
||||
"password",
|
||||
"session_token",
|
||||
"session_cookie",
|
||||
"token_cookie",
|
||||
"auth_cookie",
|
||||
"cookie_string",
|
||||
"cookie",
|
||||
}
|
||||
|
||||
def _decrypt_provider_config(self, config: dict, crypto_service: Any) -> dict:
|
||||
@@ -797,9 +828,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
for provider in providers:
|
||||
# 导出 Endpoints
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
.filter(ProviderEndpoint.provider_id == provider.id)
|
||||
.all()
|
||||
db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
|
||||
)
|
||||
endpoints_data = []
|
||||
for ep in endpoints:
|
||||
@@ -903,6 +932,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
# 导出 LDAP 配置
|
||||
from src.models.database import LDAPConfig
|
||||
|
||||
ldap_config = db.query(LDAPConfig).first()
|
||||
ldap_data = None
|
||||
if ldap_config:
|
||||
@@ -931,6 +961,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
# 导出 OAuth Providers 配置
|
||||
from src.models.database import OAuthProvider
|
||||
|
||||
oauth_providers = db.query(OAuthProvider).all()
|
||||
oauth_data = []
|
||||
for oauth in oauth_providers:
|
||||
@@ -942,21 +973,23 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
except Exception as e:
|
||||
logger.debug(f"解密 OAuth '{oauth.provider_type}' client_secret 失败: {e}")
|
||||
|
||||
oauth_data.append({
|
||||
"provider_type": oauth.provider_type,
|
||||
"display_name": oauth.display_name,
|
||||
"client_id": oauth.client_id,
|
||||
"client_secret": client_secret,
|
||||
"authorization_url_override": oauth.authorization_url_override,
|
||||
"token_url_override": oauth.token_url_override,
|
||||
"userinfo_url_override": oauth.userinfo_url_override,
|
||||
"scopes": oauth.scopes,
|
||||
"redirect_uri": oauth.redirect_uri,
|
||||
"frontend_callback_url": oauth.frontend_callback_url,
|
||||
"attribute_mapping": oauth.attribute_mapping,
|
||||
"extra_config": oauth.extra_config,
|
||||
"is_enabled": oauth.is_enabled,
|
||||
})
|
||||
oauth_data.append(
|
||||
{
|
||||
"provider_type": oauth.provider_type,
|
||||
"display_name": oauth.display_name,
|
||||
"client_id": oauth.client_id,
|
||||
"client_secret": client_secret,
|
||||
"authorization_url_override": oauth.authorization_url_override,
|
||||
"token_url_override": oauth.token_url_override,
|
||||
"userinfo_url_override": oauth.userinfo_url_override,
|
||||
"scopes": oauth.scopes,
|
||||
"redirect_uri": oauth.redirect_uri,
|
||||
"frontend_callback_url": oauth.frontend_callback_url,
|
||||
"attribute_mapping": oauth.attribute_mapping,
|
||||
"extra_config": oauth.extra_config,
|
||||
"is_enabled": oauth.is_enabled,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"version": "2.1",
|
||||
@@ -976,8 +1009,14 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
# Provider Ops 中需要加密的敏感字段
|
||||
SENSITIVE_CREDENTIALS = {
|
||||
"api_key", "password", "session_token", "session_cookie",
|
||||
"token_cookie", "auth_cookie", "cookie_string", "cookie"
|
||||
"api_key",
|
||||
"password",
|
||||
"session_token",
|
||||
"session_cookie",
|
||||
"token_cookie",
|
||||
"auth_cookie",
|
||||
"cookie_string",
|
||||
"cookie",
|
||||
}
|
||||
|
||||
def _encrypt_provider_config(self, config: dict, crypto_service: Any) -> dict:
|
||||
@@ -1045,9 +1084,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
# 导入 GlobalModels
|
||||
global_model_map = {} # name -> id 映射
|
||||
for gm_data in global_models_data:
|
||||
existing = (
|
||||
db.query(GlobalModel).filter(GlobalModel.name == gm_data["name"]).first()
|
||||
)
|
||||
existing = db.query(GlobalModel).filter(GlobalModel.name == gm_data["name"]).first()
|
||||
|
||||
if existing:
|
||||
global_model_map[gm_data["name"]] = existing.id
|
||||
@@ -1055,23 +1092,17 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
stats["global_models"]["skipped"] += 1
|
||||
continue
|
||||
elif merge_mode == "error":
|
||||
raise InvalidRequestException(
|
||||
f"GlobalModel '{gm_data['name']}' 已存在"
|
||||
)
|
||||
raise InvalidRequestException(f"GlobalModel '{gm_data['name']}' 已存在")
|
||||
elif merge_mode == "overwrite":
|
||||
# 更新现有记录
|
||||
existing.display_name = gm_data.get(
|
||||
"display_name", existing.display_name
|
||||
)
|
||||
existing.display_name = gm_data.get("display_name", existing.display_name)
|
||||
existing.default_price_per_request = gm_data.get(
|
||||
"default_price_per_request"
|
||||
)
|
||||
existing.default_tiered_pricing = gm_data.get(
|
||||
"default_tiered_pricing", existing.default_tiered_pricing
|
||||
)
|
||||
existing.supported_capabilities = gm_data.get(
|
||||
"supported_capabilities"
|
||||
)
|
||||
existing.supported_capabilities = gm_data.get("supported_capabilities")
|
||||
existing.config = gm_data.get("config")
|
||||
existing.is_active = gm_data.get("is_active", True)
|
||||
existing.updated_at = datetime.now(timezone.utc)
|
||||
@@ -1085,7 +1116,15 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
default_price_per_request=gm_data.get("default_price_per_request"),
|
||||
default_tiered_pricing=gm_data.get(
|
||||
"default_tiered_pricing",
|
||||
{"tiers": [{"up_to": None, "input_price_per_1m": 0, "output_price_per_1m": 0}]},
|
||||
{
|
||||
"tiers": [
|
||||
{
|
||||
"up_to": None,
|
||||
"input_price_per_1m": 0,
|
||||
"output_price_per_1m": 0,
|
||||
}
|
||||
]
|
||||
},
|
||||
),
|
||||
supported_capabilities=gm_data.get("supported_capabilities"),
|
||||
config=gm_data.get("config"),
|
||||
@@ -1108,33 +1147,23 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
stats["providers"]["skipped"] += 1
|
||||
# 仍然需要处理 endpoints 和 models(如果存在)
|
||||
elif merge_mode == "error":
|
||||
raise InvalidRequestException(
|
||||
f"Provider '{prov_data['name']}' 已存在"
|
||||
)
|
||||
raise InvalidRequestException(f"Provider '{prov_data['name']}' 已存在")
|
||||
elif merge_mode == "overwrite":
|
||||
# 更新现有记录
|
||||
existing_provider.name = prov_data.get(
|
||||
"name", existing_provider.name
|
||||
)
|
||||
existing_provider.name = prov_data.get("name", existing_provider.name)
|
||||
existing_provider.description = prov_data.get("description")
|
||||
existing_provider.website = prov_data.get("website")
|
||||
if prov_data.get("billing_type"):
|
||||
existing_provider.billing_type = ProviderBillingType(
|
||||
prov_data["billing_type"]
|
||||
)
|
||||
existing_provider.monthly_quota_usd = prov_data.get(
|
||||
"monthly_quota_usd"
|
||||
)
|
||||
existing_provider.quota_reset_day = prov_data.get(
|
||||
"quota_reset_day", 30
|
||||
)
|
||||
existing_provider.monthly_quota_usd = prov_data.get("monthly_quota_usd")
|
||||
existing_provider.quota_reset_day = prov_data.get("quota_reset_day", 30)
|
||||
existing_provider.provider_priority = prov_data.get(
|
||||
"provider_priority", 100
|
||||
)
|
||||
existing_provider.is_active = prov_data.get("is_active", True)
|
||||
existing_provider.concurrent_limit = prov_data.get(
|
||||
"concurrent_limit"
|
||||
)
|
||||
existing_provider.concurrent_limit = prov_data.get("concurrent_limit")
|
||||
existing_provider.max_retries = prov_data.get(
|
||||
"max_retries", existing_provider.max_retries
|
||||
)
|
||||
@@ -1178,11 +1207,17 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
# 导入 Endpoints
|
||||
for ep_data in prov_data.get("endpoints", []):
|
||||
from src.core.api_format.signature import (
|
||||
normalize_signature_key,
|
||||
parse_signature_key,
|
||||
)
|
||||
|
||||
ep_format = normalize_signature_key(ep_data["api_format"])
|
||||
existing_ep = (
|
||||
db.query(ProviderEndpoint)
|
||||
.filter(
|
||||
ProviderEndpoint.provider_id == provider_id,
|
||||
ProviderEndpoint.api_format == ep_data["api_format"],
|
||||
ProviderEndpoint.api_format == ep_format,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
@@ -1192,25 +1227,32 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
stats["endpoints"]["skipped"] += 1
|
||||
elif merge_mode == "error":
|
||||
raise InvalidRequestException(
|
||||
f"Endpoint '{ep_data['api_format']}' 已存在于 Provider '{prov_data['name']}'"
|
||||
f"Endpoint '{ep_format}' 已存在于 Provider '{prov_data['name']}'"
|
||||
)
|
||||
elif merge_mode == "overwrite":
|
||||
existing_ep.base_url = ep_data.get(
|
||||
"base_url", existing_ep.base_url
|
||||
)
|
||||
existing_ep.base_url = ep_data.get("base_url", existing_ep.base_url)
|
||||
existing_ep.header_rules = ep_data.get("header_rules")
|
||||
existing_ep.max_retries = ep_data.get("max_retries", 2)
|
||||
existing_ep.is_active = ep_data.get("is_active", True)
|
||||
existing_ep.custom_path = ep_data.get("custom_path")
|
||||
existing_ep.config = ep_data.get("config")
|
||||
existing_ep.proxy = ep_data.get("proxy")
|
||||
sig = parse_signature_key(ep_format)
|
||||
existing_ep.api_format = sig.key # 使用归一化后的格式
|
||||
existing_ep.api_family = sig.api_family.value
|
||||
existing_ep.endpoint_kind = sig.endpoint_kind.value
|
||||
existing_ep.updated_at = datetime.now(timezone.utc)
|
||||
stats["endpoints"]["updated"] += 1
|
||||
else:
|
||||
sig = parse_signature_key(ep_format)
|
||||
api_family = sig.api_family.value
|
||||
endpoint_kind = sig.endpoint_kind.value
|
||||
new_ep = ProviderEndpoint(
|
||||
id=str(uuid.uuid4()),
|
||||
provider_id=provider_id,
|
||||
api_format=ep_data["api_format"],
|
||||
api_format=sig.key, # 使用归一化后的格式
|
||||
api_family=api_family,
|
||||
endpoint_kind=endpoint_kind,
|
||||
base_url=ep_data["base_url"],
|
||||
header_rules=ep_data.get("header_rules"),
|
||||
max_retries=ep_data.get("max_retries", 2),
|
||||
@@ -1232,11 +1274,11 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
endpoint_formats: set[str] = set()
|
||||
for (api_format,) in endpoint_format_rows:
|
||||
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
endpoint_formats.add(fmt.strip().upper())
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
endpoint_formats.add(normalize_signature_key(fmt))
|
||||
existing_keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.provider_id == provider_id)
|
||||
.all()
|
||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == provider_id).all()
|
||||
)
|
||||
existing_key_values = set()
|
||||
for ek in existing_keys:
|
||||
@@ -1248,9 +1290,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
|
||||
for key_data in prov_data.get("api_keys", []):
|
||||
if not key_data.get("api_key"):
|
||||
stats["errors"].append(
|
||||
f"跳过空 API Key (Provider: {prov_data['name']})"
|
||||
)
|
||||
stats["errors"].append(f"跳过空 API Key (Provider: {prov_data['name']})")
|
||||
continue
|
||||
|
||||
plaintext_key = key_data["api_key"]
|
||||
@@ -1368,21 +1408,13 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
existing_model.provider_model_mappings = model_data.get(
|
||||
"provider_model_mappings"
|
||||
)
|
||||
existing_model.price_per_request = model_data.get(
|
||||
"price_per_request"
|
||||
)
|
||||
existing_model.tiered_pricing = model_data.get(
|
||||
"tiered_pricing"
|
||||
)
|
||||
existing_model.supports_vision = model_data.get(
|
||||
"supports_vision"
|
||||
)
|
||||
existing_model.price_per_request = model_data.get("price_per_request")
|
||||
existing_model.tiered_pricing = model_data.get("tiered_pricing")
|
||||
existing_model.supports_vision = model_data.get("supports_vision")
|
||||
existing_model.supports_function_calling = model_data.get(
|
||||
"supports_function_calling"
|
||||
)
|
||||
existing_model.supports_streaming = model_data.get(
|
||||
"supports_streaming"
|
||||
)
|
||||
existing_model.supports_streaming = model_data.get("supports_streaming")
|
||||
existing_model.supports_extended_thinking = model_data.get(
|
||||
"supports_extended_thinking"
|
||||
)
|
||||
@@ -1399,22 +1431,14 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
provider_id=provider_id,
|
||||
global_model_id=global_model_id,
|
||||
provider_model_name=model_data["provider_model_name"],
|
||||
provider_model_mappings=model_data.get(
|
||||
"provider_model_mappings"
|
||||
),
|
||||
provider_model_mappings=model_data.get("provider_model_mappings"),
|
||||
price_per_request=model_data.get("price_per_request"),
|
||||
tiered_pricing=model_data.get("tiered_pricing"),
|
||||
supports_vision=model_data.get("supports_vision"),
|
||||
supports_function_calling=model_data.get(
|
||||
"supports_function_calling"
|
||||
),
|
||||
supports_function_calling=model_data.get("supports_function_calling"),
|
||||
supports_streaming=model_data.get("supports_streaming"),
|
||||
supports_extended_thinking=model_data.get(
|
||||
"supports_extended_thinking"
|
||||
),
|
||||
supports_image_generation=model_data.get(
|
||||
"supports_image_generation"
|
||||
),
|
||||
supports_extended_thinking=model_data.get("supports_extended_thinking"),
|
||||
supports_image_generation=model_data.get("supports_image_generation"),
|
||||
is_active=model_data.get("is_active", True),
|
||||
config=model_data.get("config"),
|
||||
)
|
||||
@@ -1439,7 +1463,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
elif merge_mode == "error":
|
||||
raise InvalidRequestException("LDAP 配置已存在")
|
||||
elif merge_mode == "overwrite":
|
||||
existing_ldap.server_url = ldap_data.get("server_url", existing_ldap.server_url)
|
||||
existing_ldap.server_url = ldap_data.get(
|
||||
"server_url", existing_ldap.server_url
|
||||
)
|
||||
existing_ldap.bind_dn = ldap_data.get("bind_dn", existing_ldap.bind_dn)
|
||||
# 加密绑定密码
|
||||
if ldap_data.get("bind_password"):
|
||||
@@ -1453,11 +1479,15 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
existing_ldap.username_attr = ldap_data.get(
|
||||
"username_attr", existing_ldap.username_attr
|
||||
)
|
||||
existing_ldap.email_attr = ldap_data.get("email_attr", existing_ldap.email_attr)
|
||||
existing_ldap.email_attr = ldap_data.get(
|
||||
"email_attr", existing_ldap.email_attr
|
||||
)
|
||||
existing_ldap.display_name_attr = ldap_data.get(
|
||||
"display_name_attr", existing_ldap.display_name_attr
|
||||
)
|
||||
existing_ldap.is_enabled = ldap_data.get("is_enabled", existing_ldap.is_enabled)
|
||||
existing_ldap.is_enabled = ldap_data.get(
|
||||
"is_enabled", existing_ldap.is_enabled
|
||||
)
|
||||
existing_ldap.is_exclusive = ldap_data.get(
|
||||
"is_exclusive", existing_ldap.is_exclusive
|
||||
)
|
||||
@@ -1476,7 +1506,8 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
bind_dn=ldap_data["bind_dn"],
|
||||
bind_password_encrypted=(
|
||||
crypto_service.encrypt(ldap_data["bind_password"])
|
||||
if ldap_data.get("bind_password") else None
|
||||
if ldap_data.get("bind_password")
|
||||
else None
|
||||
),
|
||||
base_dn=ldap_data["base_dn"],
|
||||
user_search_filter=ldap_data.get("user_search_filter", "(uid={username})"),
|
||||
@@ -1494,6 +1525,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
# 导入 OAuth Providers(2.1 新增)
|
||||
if oauth_data:
|
||||
from src.models.database import OAuthProvider
|
||||
|
||||
for oauth_item in oauth_data:
|
||||
provider_type = oauth_item.get("provider_type")
|
||||
if not provider_type:
|
||||
@@ -1548,7 +1580,11 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
stats["oauth"]["updated"] += 1
|
||||
else:
|
||||
# 创建新的 OAuth Provider - 校验必填字段
|
||||
required_oauth_fields = ["client_id", "redirect_uri", "frontend_callback_url"]
|
||||
required_oauth_fields = [
|
||||
"client_id",
|
||||
"redirect_uri",
|
||||
"frontend_callback_url",
|
||||
]
|
||||
missing = [f for f in required_oauth_fields if not oauth_item.get(f)]
|
||||
if missing:
|
||||
stats["errors"].append(
|
||||
@@ -1562,7 +1598,8 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
client_id=oauth_item["client_id"],
|
||||
client_secret_encrypted=(
|
||||
crypto_service.encrypt(oauth_item["client_secret"])
|
||||
if oauth_item.get("client_secret") else None
|
||||
if oauth_item.get("client_secret")
|
||||
else None
|
||||
),
|
||||
authorization_url_override=oauth_item.get("authorization_url_override"),
|
||||
token_url_override=oauth_item.get("token_url_override"),
|
||||
@@ -1588,10 +1625,14 @@ class AdminImportConfigAdapter(AdminApiAdapter):
|
||||
# 触发开启了 auto_fetch_models 的 Key 的模型获取
|
||||
keys_to_fetch = stats.get("keys_to_fetch", [])
|
||||
if keys_to_fetch:
|
||||
logger.info(f"[AUTO_FETCH] 导入了 {len(keys_to_fetch)} 个开启自动获取模型的 Key,触发模型获取")
|
||||
logger.info(
|
||||
f"[AUTO_FETCH] 导入了 {len(keys_to_fetch)} 个开启自动获取模型的 Key,触发模型获取"
|
||||
)
|
||||
try:
|
||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||
import asyncio
|
||||
|
||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||
|
||||
scheduler = get_model_fetch_scheduler()
|
||||
for key_id in keys_to_fetch:
|
||||
asyncio.create_task(scheduler._fetch_models_for_key_by_id(key_id))
|
||||
@@ -1649,18 +1690,18 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
return data
|
||||
|
||||
# 导出 Users(排除管理员)
|
||||
users = db.query(User).filter(
|
||||
User.is_deleted.is_(False),
|
||||
User.role != UserRole.ADMIN
|
||||
).all()
|
||||
users = db.query(User).filter(User.is_deleted.is_(False), User.role != UserRole.ADMIN).all()
|
||||
users_data = []
|
||||
for user in users:
|
||||
# 导出用户的 API Keys(排除独立余额Key,独立Key单独导出)
|
||||
api_keys = db.query(ApiKey).filter(
|
||||
ApiKey.user_id == user.id,
|
||||
ApiKey.is_standalone.is_(False)
|
||||
).all()
|
||||
api_keys_data = [_serialize_api_key(key, include_is_standalone=True) for key in api_keys]
|
||||
api_keys = (
|
||||
db.query(ApiKey)
|
||||
.filter(ApiKey.user_id == user.id, ApiKey.is_standalone.is_(False))
|
||||
.all()
|
||||
)
|
||||
api_keys_data = [
|
||||
_serialize_api_key(key, include_is_standalone=True) for key in api_keys
|
||||
]
|
||||
|
||||
users_data.append(
|
||||
{
|
||||
@@ -1751,27 +1792,30 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
f"API Key '{key_data.get('name', key_hash[:8])}' 的 expires_at 格式无效"
|
||||
)
|
||||
|
||||
return ApiKey(
|
||||
id=str(uuid.uuid4()),
|
||||
user_id=owner_id,
|
||||
key_hash=key_hash,
|
||||
key_encrypted=key_data.get("key_encrypted"),
|
||||
name=key_data.get("name"),
|
||||
is_standalone=is_standalone or key_data.get("is_standalone", False),
|
||||
balance_used_usd=key_data.get("balance_used_usd", 0.0),
|
||||
current_balance_usd=key_data.get("current_balance_usd"),
|
||||
allowed_providers=key_data.get("allowed_providers"),
|
||||
allowed_api_formats=key_data.get("allowed_api_formats"),
|
||||
allowed_models=key_data.get("allowed_models"),
|
||||
rate_limit=key_data.get("rate_limit"),
|
||||
concurrent_limit=key_data.get("concurrent_limit", 5),
|
||||
force_capabilities=key_data.get("force_capabilities"),
|
||||
is_active=key_data.get("is_active", True),
|
||||
expires_at=expires_at,
|
||||
auto_delete_on_expiry=key_data.get("auto_delete_on_expiry", False),
|
||||
total_requests=key_data.get("total_requests", 0),
|
||||
total_cost_usd=key_data.get("total_cost_usd", 0.0),
|
||||
), "created"
|
||||
return (
|
||||
ApiKey(
|
||||
id=str(uuid.uuid4()),
|
||||
user_id=owner_id,
|
||||
key_hash=key_hash,
|
||||
key_encrypted=key_data.get("key_encrypted"),
|
||||
name=key_data.get("name"),
|
||||
is_standalone=is_standalone or key_data.get("is_standalone", False),
|
||||
balance_used_usd=key_data.get("balance_used_usd", 0.0),
|
||||
current_balance_usd=key_data.get("current_balance_usd"),
|
||||
allowed_providers=key_data.get("allowed_providers"),
|
||||
allowed_api_formats=key_data.get("allowed_api_formats"),
|
||||
allowed_models=key_data.get("allowed_models"),
|
||||
rate_limit=key_data.get("rate_limit"),
|
||||
concurrent_limit=key_data.get("concurrent_limit", 5),
|
||||
force_capabilities=key_data.get("force_capabilities"),
|
||||
is_active=key_data.get("is_active", True),
|
||||
expires_at=expires_at,
|
||||
auto_delete_on_expiry=key_data.get("auto_delete_on_expiry", False),
|
||||
total_requests=key_data.get("total_requests", 0),
|
||||
total_cost_usd=key_data.get("total_cost_usd", 0.0),
|
||||
),
|
||||
"created",
|
||||
)
|
||||
|
||||
try:
|
||||
for user_data in users_data:
|
||||
@@ -1785,29 +1829,21 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
# 导入必须有邮箱(email 是导入的主键)
|
||||
import_email = user_data.get("email")
|
||||
if not import_email:
|
||||
stats["errors"].append(
|
||||
f"跳过无邮箱用户: {user_data.get('username', '未知')}"
|
||||
)
|
||||
stats["errors"].append(f"跳过无邮箱用户: {user_data.get('username', '未知')}")
|
||||
stats["users"]["skipped"] += 1
|
||||
continue
|
||||
|
||||
existing_user = (
|
||||
db.query(User).filter(User.email == import_email).first()
|
||||
)
|
||||
existing_user = db.query(User).filter(User.email == import_email).first()
|
||||
|
||||
if existing_user:
|
||||
user_id = existing_user.id
|
||||
if merge_mode == "skip":
|
||||
stats["users"]["skipped"] += 1
|
||||
elif merge_mode == "error":
|
||||
raise InvalidRequestException(
|
||||
f"用户 '{import_email}' 已存在"
|
||||
)
|
||||
raise InvalidRequestException(f"用户 '{import_email}' 已存在")
|
||||
elif merge_mode == "overwrite":
|
||||
# 更新现有用户
|
||||
existing_user.username = user_data.get(
|
||||
"username", existing_user.username
|
||||
)
|
||||
existing_user.username = user_data.get("username", existing_user.username)
|
||||
if user_data.get("password_hash"):
|
||||
existing_user.password_hash = user_data["password_hash"]
|
||||
if user_data.get("role"):
|
||||
@@ -1898,8 +1934,8 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
"""测试 SMTP 连接"""
|
||||
from src.core.crypto import crypto_service
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.email.email_sender import EmailSenderService
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
db = context.db
|
||||
payload = context.ensure_json_body() or {}
|
||||
@@ -1917,16 +1953,23 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
|
||||
# 前端可传入未保存的配置,优先使用前端值,否则回退数据库
|
||||
config = {
|
||||
"smtp_host": payload.get("smtp_host") or SystemConfigService.get_config(db, "smtp_host"),
|
||||
"smtp_port": payload.get("smtp_port") or SystemConfigService.get_config(db, "smtp_port", default=587),
|
||||
"smtp_user": payload.get("smtp_user") or SystemConfigService.get_config(db, "smtp_user"),
|
||||
"smtp_host": payload.get("smtp_host")
|
||||
or SystemConfigService.get_config(db, "smtp_host"),
|
||||
"smtp_port": payload.get("smtp_port")
|
||||
or SystemConfigService.get_config(db, "smtp_port", default=587),
|
||||
"smtp_user": payload.get("smtp_user")
|
||||
or SystemConfigService.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),
|
||||
"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),
|
||||
"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)
|
||||
),
|
||||
"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)
|
||||
),
|
||||
"smtp_from_email": payload.get("smtp_from_email")
|
||||
or SystemConfigService.get_config(db, "smtp_from_email"),
|
||||
"smtp_from_name": payload.get("smtp_from_name")
|
||||
@@ -1935,12 +1978,14 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
|
||||
# 验证必要配置
|
||||
missing_fields = [
|
||||
field for field in ["smtp_host", "smtp_user", "smtp_password", "smtp_from_email"] if not config.get(field)
|
||||
field
|
||||
for field in ["smtp_host", "smtp_user", "smtp_password", "smtp_from_email"]
|
||||
if not config.get(field)
|
||||
]
|
||||
if missing_fields:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"SMTP 配置不完整,请检查 {', '.join(missing_fields)}"
|
||||
"message": f"SMTP 配置不完整,请检查 {', '.join(missing_fields)}",
|
||||
}
|
||||
|
||||
# 测试连接
|
||||
@@ -1950,20 +1995,11 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
if success:
|
||||
return {
|
||||
"success": True,
|
||||
"message": "SMTP 连接测试成功"
|
||||
}
|
||||
return {"success": True, "message": "SMTP 连接测试成功"}
|
||||
else:
|
||||
return {
|
||||
"success": False,
|
||||
"message": error_msg
|
||||
}
|
||||
return {"success": False, "message": error_msg}
|
||||
except Exception as e:
|
||||
return {
|
||||
"success": False,
|
||||
"message": str(e)
|
||||
}
|
||||
return {"success": False, "message": str(e)}
|
||||
|
||||
|
||||
# -------- 邮件模板适配器 --------
|
||||
|
||||
@@ -2,16 +2,17 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
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.database import get_db
|
||||
from src.models.database import (
|
||||
@@ -24,7 +25,6 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.usage.service import UsageService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
router = APIRouter(prefix="/api/admin/usage", tags=["Admin - Usage"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
@@ -36,7 +36,9 @@ pipeline = ApiRequestPipeline()
|
||||
@router.get("/aggregation/stats")
|
||||
async def get_usage_aggregation(
|
||||
request: Request,
|
||||
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
|
||||
group_by: str = Query(
|
||||
..., description="Aggregation dimension: model, user, provider, or api_format"
|
||||
),
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
limit: int = Query(20, ge=1, le=100),
|
||||
@@ -66,11 +68,13 @@ async def get_usage_aggregation(
|
||||
elif group_by == "provider":
|
||||
adapter = AdminUsageByProviderAdapter(start_date=start_date, end_date=end_date, limit=limit)
|
||||
elif group_by == "api_format":
|
||||
adapter = AdminUsageByApiFormatAdapter(start_date=start_date, end_date=end_date, limit=limit)
|
||||
adapter = AdminUsageByApiFormatAdapter(
|
||||
start_date=start_date, end_date=end_date, limit=limit
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid group_by value: {group_by}. Must be one of: model, user, provider, api_format"
|
||||
detail=f"Invalid group_by value: {group_by}. Must be one of: model, user, provider, api_format",
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
@@ -454,12 +458,10 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
attempt_query = db.query(
|
||||
RequestCandidate.provider_id,
|
||||
func.count(RequestCandidate.id).label("attempt_count"),
|
||||
func.sum(
|
||||
case((RequestCandidate.status == "success", 1), else_=0)
|
||||
).label("success_count"),
|
||||
func.sum(
|
||||
case((RequestCandidate.status == "failed", 1), else_=0)
|
||||
).label("failed_count"),
|
||||
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
|
||||
"success_count"
|
||||
),
|
||||
func.sum(case((RequestCandidate.status == "failed", 1), else_=0)).label("failed_count"),
|
||||
func.avg(RequestCandidate.latency_ms).label("avg_latency_ms"),
|
||||
).filter(
|
||||
RequestCandidate.provider_id.isnot(None),
|
||||
@@ -537,17 +539,19 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
# 从 usage_map 获取 token 和费用信息
|
||||
usage_stat = usage_map.get(provider_id_str)
|
||||
|
||||
result.append({
|
||||
"provider_id": provider_id_str,
|
||||
"provider": provider_map.get(provider_id_str, "Unknown"),
|
||||
"request_count": attempt_count, # 尝试次数
|
||||
"total_tokens": int(usage_stat.total_tokens or 0) if usage_stat else 0,
|
||||
"total_cost": float(usage_stat.total_cost or 0) if usage_stat else 0,
|
||||
"actual_cost": float(usage_stat.actual_cost or 0) if usage_stat else 0,
|
||||
"avg_response_time_ms": float(stat.avg_latency_ms or 0),
|
||||
"success_rate": round(success_rate, 2),
|
||||
"error_count": failed_count,
|
||||
})
|
||||
result.append(
|
||||
{
|
||||
"provider_id": provider_id_str,
|
||||
"provider": provider_map.get(provider_id_str, "Unknown"),
|
||||
"request_count": attempt_count, # 尝试次数
|
||||
"total_tokens": int(usage_stat.total_tokens or 0) if usage_stat else 0,
|
||||
"total_cost": float(usage_stat.total_cost or 0) if usage_stat else 0,
|
||||
"actual_cost": float(usage_stat.actual_cost or 0) if usage_stat else 0,
|
||||
"avg_response_time_ms": float(stat.avg_latency_ms or 0),
|
||||
"success_rate": round(success_rate, 2),
|
||||
"error_count": failed_count,
|
||||
}
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@@ -581,9 +585,7 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
query = query.filter(Usage.created_at <= self.end_date)
|
||||
|
||||
query = (
|
||||
query.group_by(Usage.api_format)
|
||||
.order_by(func.count(Usage.id).desc())
|
||||
.limit(self.limit)
|
||||
query.group_by(Usage.api_format).order_by(func.count(Usage.id).desc()).limit(self.limit)
|
||||
)
|
||||
stats = query.all()
|
||||
|
||||
@@ -691,9 +693,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
elif self.status == "standard":
|
||||
query = query.filter(Usage.is_stream == False) # noqa: E712
|
||||
elif self.status == "error":
|
||||
query = query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
)
|
||||
query = query.filter((Usage.status_code >= 400) | (Usage.error_message.isnot(None)))
|
||||
elif self.status in ("pending", "streaming", "completed"):
|
||||
# 新的状态筛选:直接按 status 字段过滤
|
||||
query = query.filter(Usage.status == self.status)
|
||||
@@ -702,9 +702,9 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
# 1. 新方式:status = "failed"
|
||||
# 2. 旧方式:status_code >= 400 或 error_message 不为空
|
||||
query = query.filter(
|
||||
(Usage.status == "failed") |
|
||||
(Usage.status_code >= 400) |
|
||||
(Usage.error_message.isnot(None))
|
||||
(Usage.status == "failed")
|
||||
| (Usage.status_code >= 400)
|
||||
| (Usage.error_message.isnot(None))
|
||||
)
|
||||
elif self.status == "active":
|
||||
# 活跃请求:pending 或 streaming 状态
|
||||
@@ -761,9 +761,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
retry_map[req_id] = has_retry
|
||||
|
||||
# 检查是否有整流:任意候选的 extra_data 中有 rectified=True
|
||||
rectified_map[req_id] = any(
|
||||
c[2].get("rectified", False) for c in candidates
|
||||
)
|
||||
rectified_map[req_id] = any(c[2].get("rectified", False) for c in candidates)
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="usage_records",
|
||||
|
||||
@@ -2,14 +2,15 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
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.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
@@ -20,8 +21,6 @@ from src.models.database import ApiKey, User, UserRole
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.services.user.service import UserService
|
||||
from src.api.base.context import ApiRequestContext
|
||||
|
||||
|
||||
router = APIRouter(prefix="/api/admin/users", tags=["Admin - Users"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
Reference in New Issue
Block a user