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:
fawney19
2026-02-01 17:28:00 +08:00
parent c246ccfc91
commit 7b66505634
219 changed files with 4732 additions and 2545 deletions

View File

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

View File

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

View File

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

View File

@@ -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 的健康状态")

View File

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

View File

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

View File

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

View File

@@ -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 不存在")

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -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:
"""
删除提供商

View File

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

View File

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

View File

@@ -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 Providers2.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)}
# -------- 邮件模板适配器 --------

View File

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

View File

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