mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
Merge branch 'fix/python314-upgrade'
# Conflicts: # src/api/handlers/base/base_handler.py # src/api/handlers/base/request_builder.py # src/models/endpoint_models.py # src/services/orchestration/candidate_resolver.py # src/services/orchestration/fallback_orchestrator.py
This commit is contained in:
@@ -10,7 +10,6 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
@@ -34,7 +33,7 @@ class EnableAdaptiveRequest(BaseModel):
|
||||
"""启用自适应模式请求"""
|
||||
|
||||
enabled: bool = Field(..., description="是否启用自适应模式(true=自适应,false=固定限制)")
|
||||
fixed_limit: Optional[int] = Field(
|
||||
fixed_limit: int | None = Field(
|
||||
None, ge=1, le=100, description="固定 RPM 限制(仅当 enabled=false 时生效,1-100)"
|
||||
)
|
||||
|
||||
@@ -43,30 +42,30 @@ class AdaptiveStatsResponse(BaseModel):
|
||||
"""自适应统计响应"""
|
||||
|
||||
adaptive_mode: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||
rpm_limit: Optional[int] = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||
effective_limit: Optional[int] = Field(
|
||||
rpm_limit: int | None = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||
effective_limit: int | None = Field(
|
||||
None, description="当前有效限制(自适应使用学习值,固定使用配置值)"
|
||||
)
|
||||
learned_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
||||
learned_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||
concurrent_429_count: int
|
||||
rpm_429_count: int
|
||||
last_429_at: Optional[str]
|
||||
last_429_type: Optional[str]
|
||||
last_429_at: str | None
|
||||
last_429_type: str | None
|
||||
adjustment_count: int
|
||||
recent_adjustments: List[dict]
|
||||
recent_adjustments: list[dict]
|
||||
|
||||
|
||||
class KeyListItem(BaseModel):
|
||||
"""Key 列表项"""
|
||||
|
||||
id: str
|
||||
name: Optional[str]
|
||||
name: str | None
|
||||
provider_id: str
|
||||
api_formats: List[str] = Field(default_factory=list)
|
||||
api_formats: list[str] = Field(default_factory=list)
|
||||
is_adaptive: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||
rpm_limit: Optional[int] = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||
effective_limit: Optional[int] = Field(None, description="当前有效限制")
|
||||
learned_rpm_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
||||
rpm_limit: int | None = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||
effective_limit: int | None = Field(None, description="当前有效限制")
|
||||
learned_rpm_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||
concurrent_429_count: int
|
||||
rpm_429_count: int
|
||||
|
||||
@@ -76,12 +75,12 @@ class KeyListItem(BaseModel):
|
||||
|
||||
@router.get(
|
||||
"/keys",
|
||||
response_model=List[KeyListItem],
|
||||
response_model=list[KeyListItem],
|
||||
summary="获取所有启用自适应模式的Key",
|
||||
)
|
||||
async def list_adaptive_keys(
|
||||
request: Request,
|
||||
provider_id: Optional[str] = Query(None, description="按 Provider 过滤"),
|
||||
provider_id: str | None = Query(None, description="按 Provider 过滤"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -207,7 +206,7 @@ async def get_adaptive_summary(
|
||||
|
||||
@dataclass
|
||||
class ListAdaptiveKeysAdapter(AdminApiAdapter):
|
||||
provider_id: Optional[str] = None
|
||||
provider_id: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
# 自适应模式:rpm_limit = NULL
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
@@ -25,7 +24,7 @@ from src.services.user.apikey import ApiKeyService
|
||||
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
|
||||
|
||||
|
||||
def parse_expiry_date(date_str: Optional[str]) -> Optional[datetime]:
|
||||
def parse_expiry_date(date_str: str | None) -> datetime | None:
|
||||
"""解析过期日期字符串为 datetime 对象。
|
||||
|
||||
Args:
|
||||
@@ -70,7 +69,7 @@ async def list_standalone_api_keys(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -330,7 +329,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
|
||||
self,
|
||||
skip: int,
|
||||
limit: int,
|
||||
is_active: Optional[bool],
|
||||
is_active: bool | None,
|
||||
):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
|
||||
@@ -5,7 +5,6 @@ Endpoint 健康监控 API
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -128,7 +127,7 @@ async def get_api_format_health_monitor(
|
||||
async def get_key_health(
|
||||
key_id: str,
|
||||
request: Request,
|
||||
api_format: Optional[str] = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||
api_format: str | None = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> HealthStatusResponse:
|
||||
"""
|
||||
@@ -161,7 +160,7 @@ async def get_key_health(
|
||||
async def recover_key_health(
|
||||
key_id: str,
|
||||
request: Request,
|
||||
api_format: Optional[str] = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||
api_format: str | None = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""
|
||||
@@ -278,7 +277,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建所有格式的 provider_count 映射
|
||||
all_formats: Dict[str, int] = {}
|
||||
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)
|
||||
@@ -295,7 +294,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
||||
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||
active_provider_formats: set[tuple[str, str]] = set()
|
||||
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
|
||||
api_format = (
|
||||
@@ -305,7 +304,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
active_provider_formats.add((str(provider_id), api_format))
|
||||
|
||||
# 1.2 统计每个 API 格式可用的活跃 Key 数量(Key 属于 Provider,通过 api_formats 关联格式)
|
||||
key_counts: Dict[str, int] = {}
|
||||
key_counts: dict[str, int] = {}
|
||||
if active_provider_formats:
|
||||
active_provider_keys = (
|
||||
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
|
||||
@@ -342,7 +341,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建每个格式的状态统计
|
||||
status_counts: Dict[str, Dict[str, int]] = {}
|
||||
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)
|
||||
@@ -370,7 +369,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped_attempts: Dict[str, List[RequestCandidate]] = {}
|
||||
grouped_attempts: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for attempt, api_format_enum, provider_id in rows:
|
||||
api_format = (
|
||||
@@ -384,7 +383,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
grouped_attempts[api_format].append(attempt)
|
||||
|
||||
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||
monitors: List[ApiFormatHealthMonitor] = []
|
||||
monitors: list[ApiFormatHealthMonitor] = []
|
||||
for api_format in all_formats:
|
||||
attempts = grouped_attempts.get(api_format, [])
|
||||
# 获取窗口内的真实统计数据
|
||||
@@ -399,7 +398,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
|
||||
# 时间线按时间正序
|
||||
attempts_sorted = list(reversed(attempts))
|
||||
events: List[EndpointHealthEvent] = []
|
||||
events: list[EndpointHealthEvent] = []
|
||||
for attempt in attempts_sorted:
|
||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||
events.append(
|
||||
@@ -462,7 +461,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
key_id: str
|
||||
api_format: Optional[str] = None
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
|
||||
@@ -500,7 +499,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
|
||||
key_id: str
|
||||
api_format: Optional[str] = None
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
@@ -6,7 +6,6 @@ import json
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -146,14 +145,14 @@ async def delete_endpoint_key(
|
||||
# ========== Provider Keys API ==========
|
||||
|
||||
|
||||
@router.get("/providers/{provider_id}/keys", response_model=List[EndpointAPIKeyResponse])
|
||||
@router.get("/providers/{provider_id}/keys", response_model=list[EndpointAPIKeyResponse])
|
||||
async def list_provider_keys(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[EndpointAPIKeyResponse]:
|
||||
) -> list[EndpointAPIKeyResponse]:
|
||||
"""
|
||||
获取 Provider 的所有 Keys
|
||||
|
||||
@@ -503,12 +502,12 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_base_url_map: Dict[tuple[str, str], str] = {}
|
||||
endpoint_base_url_map: dict[tuple[str, str], str] = {}
|
||||
for provider_id, api_format, base_url in endpoints:
|
||||
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
|
||||
|
||||
grouped: Dict[str, List[dict]] = {}
|
||||
grouped: dict[str, list[dict]] = {}
|
||||
for key, provider in keys:
|
||||
api_formats = key.api_formats or []
|
||||
|
||||
|
||||
@@ -5,10 +5,9 @@ ProviderEndpoint CRUD 管理 API
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import and_, func
|
||||
from sqlalchemy import and_
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
@@ -29,7 +28,7 @@ router = APIRouter(tags=["Endpoint Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
||||
def mask_proxy_password(proxy_config: dict | None) -> dict | None:
|
||||
"""对代理配置中的密码进行脱敏处理"""
|
||||
if not proxy_config:
|
||||
return None
|
||||
@@ -39,14 +38,14 @@ def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
||||
return masked
|
||||
|
||||
|
||||
@router.get("/providers/{provider_id}/endpoints", response_model=List[ProviderEndpointResponse])
|
||||
@router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse])
|
||||
async def list_provider_endpoints(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ProviderEndpointResponse]:
|
||||
) -> list[ProviderEndpointResponse]:
|
||||
"""
|
||||
获取指定 Provider 的所有 Endpoints
|
||||
|
||||
@@ -245,7 +244,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
|
||||
if is_active:
|
||||
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
|
||||
|
||||
result: List[ProviderEndpointResponse] = []
|
||||
result: list[ProviderEndpointResponse] = []
|
||||
for endpoint in endpoints:
|
||||
endpoint_format = (
|
||||
endpoint.api_format
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""LDAP配置管理API端点。"""
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, Field, ValidationError, field_validator
|
||||
@@ -30,9 +30,9 @@ BCRYPT_HASH_PATTERN = re.compile(r"^\$2[aby]\$\d{2}\$.{53}$")
|
||||
class LDAPConfigResponse(BaseModel):
|
||||
"""LDAP配置响应(不返回密码)"""
|
||||
|
||||
server_url: Optional[str] = None
|
||||
bind_dn: Optional[str] = None
|
||||
base_dn: Optional[str] = None
|
||||
server_url: str | None = None
|
||||
bind_dn: str | None = None
|
||||
base_dn: str | None = None
|
||||
has_bind_password: bool = False
|
||||
user_search_filter: str
|
||||
username_attr: str
|
||||
@@ -50,7 +50,7 @@ class LDAPConfigUpdate(BaseModel):
|
||||
server_url: str = Field(..., min_length=1, max_length=255)
|
||||
bind_dn: str = Field(..., min_length=1, max_length=255)
|
||||
# 允许空字符串表示"清除密码";非空时自动 strip 并校验不能为空
|
||||
bind_password: Optional[str] = Field(None, max_length=1024)
|
||||
bind_password: str | None = Field(None, max_length=1024)
|
||||
base_dn: str = Field(..., min_length=1, max_length=255)
|
||||
user_search_filter: str = Field(default="(uid={username})", max_length=500)
|
||||
username_attr: str = Field(default="uid", max_length=50)
|
||||
@@ -63,7 +63,7 @@ class LDAPConfigUpdate(BaseModel):
|
||||
|
||||
@field_validator("bind_password")
|
||||
@classmethod
|
||||
def validate_bind_password(cls, v: Optional[str]) -> Optional[str]:
|
||||
def validate_bind_password(cls, v: str | None) -> str | None:
|
||||
if v is None or v == "":
|
||||
return v
|
||||
v = v.strip()
|
||||
@@ -114,22 +114,22 @@ class LDAPTestResponse(BaseModel):
|
||||
class LDAPConfigTest(BaseModel):
|
||||
"""LDAP配置测试请求(全部可选,用于临时覆盖)"""
|
||||
|
||||
server_url: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
bind_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
bind_password: Optional[str] = Field(None, min_length=1)
|
||||
base_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
user_search_filter: Optional[str] = Field(None, max_length=500)
|
||||
username_attr: Optional[str] = Field(None, max_length=50)
|
||||
email_attr: Optional[str] = Field(None, max_length=50)
|
||||
display_name_attr: Optional[str] = Field(None, max_length=50)
|
||||
is_enabled: Optional[bool] = None
|
||||
is_exclusive: Optional[bool] = None
|
||||
use_starttls: Optional[bool] = None
|
||||
connect_timeout: Optional[int] = Field(None, ge=1, le=60)
|
||||
server_url: str | None = Field(None, min_length=1, max_length=255)
|
||||
bind_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||
bind_password: str | None = Field(None, min_length=1)
|
||||
base_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||
user_search_filter: str | None = Field(None, max_length=500)
|
||||
username_attr: str | None = Field(None, max_length=50)
|
||||
email_attr: str | None = Field(None, max_length=50)
|
||||
display_name_attr: str | None = Field(None, max_length=50)
|
||||
is_enabled: bool | None = None
|
||||
is_exclusive: bool | None = None
|
||||
use_starttls: bool | None = None
|
||||
connect_timeout: int | None = Field(None, ge=1, le=60)
|
||||
|
||||
@field_validator("user_search_filter")
|
||||
@classmethod
|
||||
def validate_search_filter(cls, v: Optional[str]) -> Optional[str]:
|
||||
def validate_search_filter(cls, v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
if "{username}" not in v:
|
||||
@@ -263,7 +263,7 @@ async def test_ldap_connection(request: Request, db: Session = Depends(get_db))
|
||||
|
||||
|
||||
class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
config = db.query(LDAPConfig).first()
|
||||
|
||||
@@ -300,7 +300,7 @@ class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, str]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, str]: # type: ignore[override]
|
||||
db = context.db
|
||||
payload = context.ensure_json_body()
|
||||
|
||||
@@ -421,7 +421,7 @@ class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.services.auth.ldap import LDAPService
|
||||
|
||||
db = context.db
|
||||
@@ -442,7 +442,7 @@ class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
config_data: Dict[str, Any] = {}
|
||||
config_data: dict[str, Any] = {}
|
||||
|
||||
if saved_config:
|
||||
config_data = {
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""管理员 Management Token 管理端点"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -46,8 +45,8 @@ class AdminManagementTokenApiAdapter(AdminApiAdapter):
|
||||
@router.get("")
|
||||
async def list_all_management_tokens(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="筛选用户 ID"),
|
||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
||||
user_id: str | None = Query(None, description="筛选用户 ID"),
|
||||
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -174,8 +173,8 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
||||
"""列出所有 Management Tokens"""
|
||||
|
||||
name: str = "admin_list_management_tokens"
|
||||
user_id: Optional[str] = None
|
||||
is_active: Optional[bool] = None
|
||||
user_id: str | None = None
|
||||
is_active: bool | None = None
|
||||
skip: int = 0
|
||||
limit: int = 50
|
||||
|
||||
@@ -197,7 +196,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
||||
)
|
||||
|
||||
# 预加载用户信息
|
||||
user_ids = list(set(t.user_id for t in tokens))
|
||||
user_ids = list({t.user_id for t in tokens})
|
||||
users = {u.id: u for u in context.db.query(User).filter(User.id.in_(user_ids)).all()}
|
||||
for token in tokens:
|
||||
token.user = users.get(token.user_id)
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
@@ -64,12 +63,12 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
db: Session = context.db
|
||||
|
||||
# 1. 获取所有活跃的 GlobalModel
|
||||
global_models: List[GlobalModel] = (
|
||||
global_models: list[GlobalModel] = (
|
||||
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
||||
)
|
||||
|
||||
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
|
||||
models: List[Model] = (
|
||||
models: list[Model] = (
|
||||
db.query(Model)
|
||||
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
||||
.filter(Model.is_active == True)
|
||||
@@ -77,17 +76,17 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 按 GlobalModel ID 组织关联提供商
|
||||
models_by_global_model: Dict[str, List[Model]] = {}
|
||||
models_by_global_model: dict[str, list[Model]] = {}
|
||||
for model in models:
|
||||
if model.global_model_id:
|
||||
models_by_global_model.setdefault(model.global_model_id, []).append(model)
|
||||
|
||||
# 3. 为每个 GlobalModel 构建 catalog item
|
||||
catalog_items: List[ModelCatalogItem] = []
|
||||
catalog_items: list[ModelCatalogItem] = []
|
||||
|
||||
for gm in global_models:
|
||||
gm_id = gm.id
|
||||
provider_entries: List[ModelCatalogProviderDetail] = []
|
||||
provider_entries: list[ModelCatalogProviderDetail] = []
|
||||
# 从 config JSON 读取能力标志
|
||||
gm_config = gm.config or {}
|
||||
capability_flags = {
|
||||
|
||||
@@ -3,8 +3,7 @@ models.dev 外部模型数据代理
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -43,7 +42,7 @@ OFFICIAL_PROVIDERS = {
|
||||
}
|
||||
|
||||
|
||||
async def _get_cached_data() -> Optional[dict[str, Any]]:
|
||||
async def _get_cached_data() -> dict[str, Any] | None:
|
||||
"""从 Redis 获取缓存数据"""
|
||||
redis = await get_redis_client()
|
||||
if redis is None:
|
||||
|
||||
@@ -5,7 +5,6 @@ GlobalModel Admin API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -37,8 +36,8 @@ async def list_global_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=1000),
|
||||
is_active: Optional[bool] = Query(None),
|
||||
search: Optional[str] = Query(None),
|
||||
is_active: bool | None = Query(None),
|
||||
search: str | None = Query(None),
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelListResponse:
|
||||
"""
|
||||
@@ -254,8 +253,8 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
is_active: Optional[bool]
|
||||
search: Optional[str]
|
||||
is_active: bool | None
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy import func
|
||||
|
||||
@@ -9,7 +9,6 @@ GlobalModel 请求链路预览 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
@@ -47,22 +46,22 @@ class RoutingKeyInfo(BaseModel):
|
||||
name: str
|
||||
masked_key: str = Field("", description="脱敏的 API Key")
|
||||
internal_priority: int = Field(..., description="Key 内部优先级")
|
||||
global_priority_by_format: Optional[Dict[str, int]] = Field(None, description="按 API 格式的全局优先级")
|
||||
rpm_limit: Optional[int] = Field(None, description="RPM 限制,null 表示自适应")
|
||||
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: Optional[int] = Field(None, description="有效 RPM 限制")
|
||||
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
|
||||
cache_ttl_minutes: int = Field(0, description="缓存 TTL(分钟)")
|
||||
health_score: float = Field(1.0, description="健康度分数(0-1 小数格式)")
|
||||
is_active: bool
|
||||
api_formats: List[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
api_formats: list[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
# 模型白名单
|
||||
allowed_models: Optional[List[str]] = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
allowed_models: list[str] | None = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
# 熔断状态
|
||||
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
|
||||
circuit_breaker_formats: List[str] = Field(
|
||||
circuit_breaker_formats: list[str] = Field(
|
||||
default_factory=list, description="熔断的 API 格式列表"
|
||||
)
|
||||
next_probe_at: Optional[str] = Field(None, description="下次探测时间(ISO格式)")
|
||||
next_probe_at: str | None = Field(None, description="下次探测时间(ISO格式)")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -73,9 +72,9 @@ class RoutingEndpointInfo(BaseModel):
|
||||
id: str
|
||||
api_format: str
|
||||
base_url: str
|
||||
custom_path: Optional[str] = None
|
||||
custom_path: str | None = None
|
||||
is_active: bool
|
||||
keys: List[RoutingKeyInfo] = Field(default_factory=list)
|
||||
keys: list[RoutingKeyInfo] = Field(default_factory=list)
|
||||
total_keys: int = 0
|
||||
active_keys: int = 0
|
||||
|
||||
@@ -87,7 +86,7 @@ class RoutingModelMapping(BaseModel):
|
||||
|
||||
name: str = Field(..., description="映射名称")
|
||||
priority: int = Field(..., description="优先级(数字越小优先级越高)")
|
||||
api_formats: Optional[List[str]] = Field(None, description="作用域(适用的 API 格式)")
|
||||
api_formats: list[str] | None = Field(None, description="作用域(适用的 API 格式)")
|
||||
|
||||
|
||||
class RoutingProviderInfo(BaseModel):
|
||||
@@ -97,18 +96,18 @@ class RoutingProviderInfo(BaseModel):
|
||||
name: str
|
||||
model_id: str = Field(..., description="Model ID(GlobalModel 与 Provider 的关联记录 ID)")
|
||||
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
|
||||
billing_type: Optional[str] = Field(None, description="计费类型")
|
||||
monthly_quota_usd: Optional[float] = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: Optional[float] = Field(None, description="已用额度(美元)")
|
||||
billing_type: str | None = Field(None, description="计费类型")
|
||||
monthly_quota_usd: float | None = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: float | None = Field(None, description="已用额度(美元)")
|
||||
is_active: bool
|
||||
# 模型映射信息
|
||||
provider_model_name: str = Field(..., description="提供商侧的模型名称")
|
||||
model_mappings: List[RoutingModelMapping] = Field(
|
||||
model_mappings: list[RoutingModelMapping] = Field(
|
||||
default_factory=list, description="模型名称映射列表"
|
||||
)
|
||||
model_is_active: bool = Field(True, description="Model 是否活跃")
|
||||
# Endpoint 和 Key 信息
|
||||
endpoints: List[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
endpoints: list[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
total_endpoints: int = 0
|
||||
active_endpoints: int = 0
|
||||
|
||||
@@ -123,7 +122,7 @@ class GlobalKeyWhitelistItem(BaseModel):
|
||||
masked_key: str = Field(..., description="脱敏的 API Key")
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
provider_name: str = Field(..., description="Provider 名称")
|
||||
allowed_models: List[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -136,11 +135,11 @@ class ModelRoutingPreviewResponse(BaseModel):
|
||||
display_name: str
|
||||
is_active: bool
|
||||
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
|
||||
global_model_mappings: List[str] = Field(
|
||||
global_model_mappings: list[str] = Field(
|
||||
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
|
||||
)
|
||||
# 链路信息
|
||||
providers: List[RoutingProviderInfo] = Field(
|
||||
providers: list[RoutingProviderInfo] = Field(
|
||||
default_factory=list, description="按优先级排序的提供商列表"
|
||||
)
|
||||
total_providers: int = 0
|
||||
@@ -149,7 +148,7 @@ class ModelRoutingPreviewResponse(BaseModel):
|
||||
scheduling_mode: str = Field("cache_affinity", description="调度模式")
|
||||
priority_mode: str = Field("provider", description="优先级模式")
|
||||
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key)
|
||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = Field(
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = Field(
|
||||
default_factory=list, description="所有 Provider 的 Key 白名单数据"
|
||||
)
|
||||
|
||||
@@ -227,7 +226,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
provider_ids = [m.provider_id for m in models if m.provider_id]
|
||||
|
||||
# 批量获取 Provider 的 Endpoints
|
||||
endpoints_by_provider: Dict[str, List[ProviderEndpoint]] = {}
|
||||
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
|
||||
if provider_ids:
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
@@ -240,7 +239,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
endpoints_by_provider[ep.provider_id].append(ep)
|
||||
|
||||
# 批量获取 Provider 的 Keys
|
||||
keys_by_provider: Dict[str, List[ProviderAPIKey]] = {}
|
||||
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
|
||||
if provider_ids:
|
||||
keys = (
|
||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
|
||||
@@ -251,14 +250,14 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
keys_by_provider[key.provider_id].append(key)
|
||||
|
||||
# 提取 GlobalModel 的 model_mappings(用于 Key 白名单匹配)
|
||||
global_model_mappings: List[str] = []
|
||||
global_model_mappings: list[str] = []
|
||||
if global_model.config and isinstance(global_model.config, dict):
|
||||
mappings = global_model.config.get("model_mappings")
|
||||
if isinstance(mappings, list):
|
||||
global_model_mappings = [m for m in mappings if isinstance(m, str)]
|
||||
|
||||
# 构建 Provider 路由信息
|
||||
provider_infos: List[RoutingProviderInfo] = []
|
||||
provider_infos: list[RoutingProviderInfo] = []
|
||||
for model in models:
|
||||
provider = model.provider
|
||||
if not provider:
|
||||
@@ -281,7 +280,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
provider_keys = keys_by_provider.get(provider.id, [])
|
||||
|
||||
# 按 api_format 组织 Keys
|
||||
keys_by_endpoint: Dict[str, List[ProviderAPIKey]] = {}
|
||||
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {}
|
||||
for key in provider_keys:
|
||||
# 每个 Key 可能支持多个 api_formats
|
||||
for fmt in key.api_formats or []:
|
||||
@@ -355,8 +354,8 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
|
||||
# 检查熔断状态
|
||||
circuit_breaker_open = False
|
||||
circuit_breaker_formats: List[str] = []
|
||||
next_probe_at: Optional[str] = None
|
||||
circuit_breaker_formats: list[str] = []
|
||||
next_probe_at: str | None = None
|
||||
if key.circuit_breaker_by_format:
|
||||
for fmt, cb_state in key.circuit_breaker_by_format.items():
|
||||
if isinstance(cb_state, dict) and cb_state.get("open"):
|
||||
@@ -462,7 +461,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
|
||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = []
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = []
|
||||
crypto = CryptoService()
|
||||
|
||||
# 获取所有活跃的 Key(带白名单),使用 selectinload 避免 N+1 查询
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""模块管理 API 端点"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel
|
||||
@@ -28,18 +29,18 @@ class ModuleStatusResponse(BaseModel):
|
||||
enabled: bool
|
||||
active: bool
|
||||
config_validated: bool
|
||||
config_error: Optional[str]
|
||||
config_error: str | None
|
||||
display_name: str
|
||||
description: str
|
||||
category: str
|
||||
admin_route: Optional[str]
|
||||
admin_menu_icon: Optional[str]
|
||||
admin_menu_group: Optional[str]
|
||||
admin_route: str | None
|
||||
admin_menu_icon: str | None
|
||||
admin_menu_group: str | None
|
||||
admin_menu_order: int
|
||||
health: str
|
||||
|
||||
@classmethod
|
||||
def from_status(cls, status: ModuleStatus) -> "ModuleStatusResponse":
|
||||
def from_status(cls, status: ModuleStatus) -> ModuleStatusResponse:
|
||||
return cls(
|
||||
name=status.name,
|
||||
available=status.available,
|
||||
@@ -130,7 +131,7 @@ async def set_module_enabled(
|
||||
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
|
||||
"""获取所有模块状态"""
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
all_status = await registry.get_all_status_async(context.db)
|
||||
|
||||
@@ -146,7 +147,7 @@ class AdminGetModuleStatusAdapter(AdminApiAdapter):
|
||||
|
||||
module_name: str
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
status = await registry.get_module_status_async(self.module_name, context.db)
|
||||
|
||||
@@ -162,7 +163,7 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
|
||||
|
||||
module_name: str
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
|
||||
# 检查模块是否存在
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -33,8 +32,8 @@ pipeline = ApiRequestPipeline()
|
||||
@router.get("/audit-logs")
|
||||
async def get_audit_logs(
|
||||
request: Request,
|
||||
username: Optional[str] = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
||||
username: str | None = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||
days: int = Query(7, description="查询天数"),
|
||||
limit: int = Query(100, description="返回数量限制"),
|
||||
offset: int = Query(0, description="偏移量"),
|
||||
@@ -212,8 +211,8 @@ async def get_circuit_history(
|
||||
|
||||
@dataclass
|
||||
class AdminGetAuditLogsAdapter(AdminApiAdapter):
|
||||
username: Optional[str]
|
||||
event_type: Optional[str]
|
||||
username: str | None
|
||||
event_type: str | None
|
||||
days: int
|
||||
limit: int
|
||||
offset: int
|
||||
@@ -497,8 +496,8 @@ class AdminCircuitHistoryAdapter(AdminApiAdapter):
|
||||
return {"items": history, "count": len(history)}
|
||||
|
||||
|
||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> List[str]:
|
||||
recommendations: List[str] = []
|
||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> list[str]:
|
||||
recommendations: list[str] = []
|
||||
if health_score < 50:
|
||||
recommendations.append("系统健康状况严重,请立即检查错误日志")
|
||||
if error_stats.get("total_errors", 0) > 100:
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import PlainTextResponse
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pagination import build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.crypto import crypto_service
|
||||
@@ -28,7 +28,7 @@ router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitori
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int = 4) -> Optional[str]:
|
||||
def mask_api_key(api_key: str | None, prefix_len: int = 8, suffix_len: int = 4) -> str | None:
|
||||
"""
|
||||
脱敏 API Key,显示前缀 + 星号 + 后缀
|
||||
例如: sk-jhiId-xxxxxxxxxxxAABB -> sk-jhiId-********AABB
|
||||
@@ -47,7 +47,7 @@ def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int =
|
||||
return f"{api_key[:prefix_len]}********{api_key[-suffix_len:]}"
|
||||
|
||||
|
||||
def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optional[str]:
|
||||
def decrypt_and_mask(encrypted_key: str | None, prefix_len: int = 8) -> str | None:
|
||||
"""
|
||||
解密 API Key 后脱敏显示
|
||||
|
||||
@@ -65,7 +65,7 @@ def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optio
|
||||
return None
|
||||
|
||||
|
||||
def resolve_user_identifier(db: Session, identifier: str) -> Optional[str]:
|
||||
def resolve_user_identifier(db: Session, identifier: str) -> str | None:
|
||||
"""
|
||||
将用户标识符(username/email/user_id/api_key_id)解析为 user_id
|
||||
|
||||
@@ -181,7 +181,7 @@ async def get_user_affinity(
|
||||
@router.get("/affinities")
|
||||
async def list_affinities(
|
||||
request: Request,
|
||||
keyword: Optional[str] = None,
|
||||
keyword: str | None = None,
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回数量限制"),
|
||||
offset: int = Query(0, ge=0, description="偏移量"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -421,7 +421,7 @@ async def get_cache_metrics(
|
||||
|
||||
|
||||
class AdminCacheStatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
# 读取系统配置,确保监控接口与编排器使用一致的模式
|
||||
@@ -487,14 +487,14 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
logger.exception(f"导出缓存指标失败: {exc}")
|
||||
raise HTTPException(status_code=500, detail=f"导出缓存指标失败: {exc}")
|
||||
|
||||
def _format_prometheus(self, stats: Dict[str, Any]) -> str:
|
||||
def _format_prometheus(self, stats: dict[str, Any]) -> str:
|
||||
"""
|
||||
将 scheduler/affinity 指标转换为 Prometheus 文本格式。
|
||||
"""
|
||||
scheduler_metrics = stats.get("scheduler_metrics", {})
|
||||
affinity_stats = stats.get("affinity_stats", {})
|
||||
|
||||
metric_map: List[Tuple[str, str, float]] = [
|
||||
metric_map: list[tuple[str, str, float]] = [
|
||||
(
|
||||
"cache_scheduler_total_batches",
|
||||
"Total batches pulled from provider list",
|
||||
@@ -542,7 +542,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
),
|
||||
]
|
||||
|
||||
affinity_map: List[Tuple[str, str, float]] = [
|
||||
affinity_map: list[tuple[str, str, float]] = [
|
||||
(
|
||||
"cache_affinity_total",
|
||||
"Total cache affinities stored",
|
||||
@@ -596,7 +596,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||
user_identifier: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
user_id = resolve_user_identifier(db, self.user_identifier)
|
||||
@@ -673,11 +673,11 @@ class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
keyword: Optional[str]
|
||||
keyword: str | None
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
redis_client = get_redis_client_sync()
|
||||
if not redis_client:
|
||||
@@ -686,7 +686,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
matched_user_id = None
|
||||
matched_api_key_id = None
|
||||
raw_affinities: List[Dict[str, Any]] = []
|
||||
raw_affinities: list[dict[str, Any]] = []
|
||||
|
||||
if self.keyword:
|
||||
# 首先检查是否是 API Key ID(affinity_key)
|
||||
@@ -724,14 +724,14 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
}
|
||||
|
||||
# 批量查询用户 API Key 信息
|
||||
user_api_key_map: Dict[str, ApiKey] = {}
|
||||
user_api_key_map: dict[str, ApiKey] = {}
|
||||
if affinity_keys:
|
||||
user_api_keys = db.query(ApiKey).filter(ApiKey.id.in_(list(affinity_keys))).all()
|
||||
user_api_key_map = {str(k.id): k for k in user_api_keys}
|
||||
|
||||
# 收集所有 user_id
|
||||
user_ids = {str(k.user_id) for k in user_api_key_map.values()}
|
||||
user_map: Dict[str, User] = {}
|
||||
user_map: dict[str, User] = {}
|
||||
if user_ids:
|
||||
users = db.query(User).filter(User.id.in_(list(user_ids))).all()
|
||||
user_map = {str(user.id): user for user in users}
|
||||
@@ -771,7 +771,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
global_model_ids = {
|
||||
item.get("model_name") for item in raw_affinities if item.get("model_name")
|
||||
}
|
||||
global_model_map: Dict[str, GlobalModel] = {}
|
||||
global_model_map: dict[str, GlobalModel] = {}
|
||||
if global_model_ids:
|
||||
# model_name 可能是 UUID 格式的 global_model_id,也可能是原始模型名称
|
||||
global_models = db.query(GlobalModel).filter(
|
||||
@@ -885,7 +885,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
class AdminClearUserCacheAdapter(AdminApiAdapter):
|
||||
user_identifier: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
@@ -995,7 +995,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
||||
model_id: str
|
||||
api_format: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
@@ -1048,7 +1048,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
@@ -1068,7 +1068,7 @@ class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||
class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
@@ -1091,7 +1091,7 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminCacheConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
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.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
@@ -1260,7 +1260,7 @@ async def clear_provider_model_mapping_cache(
|
||||
|
||||
|
||||
class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
import json
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
@@ -1510,7 +1510,7 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
@@ -1552,7 +1552,7 @@ class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||
class AdminClearModelMappingCacheByNameAdapter(AdminApiAdapter):
|
||||
model_name: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
@@ -1599,7 +1599,7 @@ class AdminClearProviderModelMappingCacheAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
@@ -28,29 +27,29 @@ class CandidateResponse(BaseModel):
|
||||
request_id: str
|
||||
candidate_index: int
|
||||
retry_index: int = 0 # 重试序号(从0开始)
|
||||
provider_id: Optional[str] = None
|
||||
provider_name: Optional[str] = None
|
||||
provider_website: Optional[str] = None # Provider 官网
|
||||
endpoint_id: Optional[str] = None
|
||||
endpoint_name: Optional[str] = None # 端点显示名称(api_format)
|
||||
key_id: Optional[str] = None
|
||||
key_name: Optional[str] = None # 密钥名称
|
||||
key_preview: Optional[str] = None # 密钥脱敏预览(如 sk-***abc)
|
||||
key_capabilities: Optional[dict] = None # Key 支持的能力
|
||||
required_capabilities: Optional[dict] = None # 请求实际需要的能力标签
|
||||
provider_id: str | None = None
|
||||
provider_name: str | None = None
|
||||
provider_website: str | None = None # Provider 官网
|
||||
endpoint_id: str | None = None
|
||||
endpoint_name: str | None = None # 端点显示名称(api_format)
|
||||
key_id: str | None = None
|
||||
key_name: str | None = None # 密钥名称
|
||||
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abc)
|
||||
key_capabilities: dict | None = None # Key 支持的能力
|
||||
required_capabilities: dict | None = None # 请求实际需要的能力标签
|
||||
status: str # 'pending', 'success', 'failed', 'skipped'
|
||||
skip_reason: Optional[str] = None
|
||||
skip_reason: str | None = None
|
||||
is_cached: bool = False
|
||||
# 执行结果字段
|
||||
status_code: Optional[int] = None
|
||||
error_type: Optional[str] = None
|
||||
error_message: Optional[str] = None
|
||||
latency_ms: Optional[int] = None
|
||||
concurrent_requests: Optional[int] = None
|
||||
extra_data: Optional[dict] = None
|
||||
status_code: int | None = None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
latency_ms: int | None = None
|
||||
concurrent_requests: int | None = None
|
||||
extra_data: dict | None = None
|
||||
created_at: datetime
|
||||
started_at: Optional[datetime] = None
|
||||
finished_at: Optional[datetime] = None
|
||||
started_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -62,7 +61,7 @@ class RequestTraceResponse(BaseModel):
|
||||
total_candidates: int
|
||||
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
|
||||
total_latency_ms: int
|
||||
candidates: List[CandidateResponse]
|
||||
candidates: list[CandidateResponse]
|
||||
|
||||
|
||||
@router.get("/{request_id}", response_model=RequestTraceResponse)
|
||||
@@ -253,7 +252,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
key_preview_map[k.id] = "***"
|
||||
|
||||
# 构建 candidate 响应列表
|
||||
candidate_responses: List[CandidateResponse] = []
|
||||
candidate_responses: list[CandidateResponse] = []
|
||||
for candidate in candidates:
|
||||
provider_name = (
|
||||
provider_map.get(candidate.provider_id) if candidate.provider_id else None
|
||||
|
||||
@@ -9,7 +9,7 @@ Provider 操作 API 路由
|
||||
"""
|
||||
|
||||
from dataclasses import asdict, is_dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
@@ -18,9 +18,7 @@ from sqlalchemy.orm import Session
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, User
|
||||
from src.services.provider_ops import (
|
||||
ActionStatus,
|
||||
ConnectorAuthType,
|
||||
ConnectorStatus,
|
||||
ProviderActionType,
|
||||
ProviderOpsConfig,
|
||||
ProviderOpsService,
|
||||
@@ -40,46 +38,46 @@ class ArchitectureInfo(BaseModel):
|
||||
architecture_id: str
|
||||
display_name: str
|
||||
description: str
|
||||
supported_auth_types: List[Dict[str, str]]
|
||||
supported_actions: List[Dict[str, Any]]
|
||||
default_connector: Optional[str]
|
||||
supported_auth_types: list[dict[str, str]]
|
||||
supported_actions: list[dict[str, Any]]
|
||||
default_connector: str | None
|
||||
|
||||
|
||||
class ConnectorConfigRequest(BaseModel):
|
||||
"""连接器配置请求"""
|
||||
|
||||
auth_type: str = Field(..., description="认证类型")
|
||||
config: Dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||
credentials: Dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||
credentials: dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||
|
||||
|
||||
class ActionConfigRequest(BaseModel):
|
||||
"""操作配置请求"""
|
||||
|
||||
enabled: bool = Field(True, description="是否启用")
|
||||
config: Dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||
|
||||
|
||||
class SaveConfigRequest(BaseModel):
|
||||
"""保存配置请求"""
|
||||
|
||||
architecture_id: str = Field("generic_api", description="架构 ID")
|
||||
base_url: Optional[str] = Field(None, description="API 基础地址")
|
||||
base_url: str | None = Field(None, description="API 基础地址")
|
||||
connector: ConnectorConfigRequest
|
||||
actions: Dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||
schedule: Dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||
actions: dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||
schedule: dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||
|
||||
|
||||
class ConnectRequest(BaseModel):
|
||||
"""连接请求"""
|
||||
|
||||
credentials: Optional[Dict[str, Any]] = Field(None, description="凭据(可选,使用已保存的)")
|
||||
credentials: dict[str, Any] | None = Field(None, description="凭据(可选,使用已保存的)")
|
||||
|
||||
|
||||
class ExecuteActionRequest(BaseModel):
|
||||
"""执行操作请求"""
|
||||
|
||||
config: Optional[Dict[str, Any]] = Field(None, description="操作配置(覆盖默认)")
|
||||
config: dict[str, Any] | None = Field(None, description="操作配置(覆盖默认)")
|
||||
|
||||
|
||||
class ConnectionStatusResponse(BaseModel):
|
||||
@@ -87,9 +85,9 @@ class ConnectionStatusResponse(BaseModel):
|
||||
|
||||
status: str
|
||||
auth_type: str
|
||||
connected_at: Optional[str]
|
||||
expires_at: Optional[str]
|
||||
last_error: Optional[str]
|
||||
connected_at: str | None
|
||||
expires_at: str | None
|
||||
last_error: str | None
|
||||
|
||||
|
||||
class ActionResultResponse(BaseModel):
|
||||
@@ -97,10 +95,10 @@ class ActionResultResponse(BaseModel):
|
||||
|
||||
status: str
|
||||
action_type: str
|
||||
data: Optional[Any]
|
||||
message: Optional[str]
|
||||
data: Any | None
|
||||
message: str | None
|
||||
executed_at: str
|
||||
response_time_ms: Optional[int]
|
||||
response_time_ms: int | None
|
||||
cache_ttl_seconds: int
|
||||
|
||||
|
||||
@@ -109,9 +107,9 @@ class ProviderOpsStatusResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
is_configured: bool
|
||||
architecture_id: Optional[str]
|
||||
architecture_id: str | None
|
||||
connection_status: ConnectionStatusResponse
|
||||
enabled_actions: List[str]
|
||||
enabled_actions: list[str]
|
||||
|
||||
|
||||
class ProviderOpsConfigResponse(BaseModel):
|
||||
@@ -119,17 +117,17 @@ class ProviderOpsConfigResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
is_configured: bool
|
||||
architecture_id: Optional[str] = None
|
||||
base_url: Optional[str] = None
|
||||
connector: Optional[Dict[str, Any]] = None # 脱敏后的连接器配置
|
||||
architecture_id: str | None = None
|
||||
base_url: str | None = None
|
||||
connector: dict[str, Any] | None = None # 脱敏后的连接器配置
|
||||
|
||||
|
||||
class VerifyAuthResponse(BaseModel):
|
||||
"""验证认证响应"""
|
||||
|
||||
success: bool
|
||||
message: Optional[str] = None
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
message: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
|
||||
# ==================== Helper Functions ====================
|
||||
@@ -147,7 +145,7 @@ def _serialize_data(data: Any) -> Any:
|
||||
# ==================== Routes ====================
|
||||
|
||||
|
||||
@router.get("/architectures", response_model=List[ArchitectureInfo])
|
||||
@router.get("/architectures", response_model=list[ArchitectureInfo])
|
||||
async def list_architectures(_: User = Depends(require_admin)):
|
||||
"""获取所有可用的架构"""
|
||||
registry = get_registry()
|
||||
@@ -490,7 +488,7 @@ async def checkin(
|
||||
|
||||
@router.post("/batch/balance")
|
||||
async def batch_query_balance(
|
||||
provider_ids: Optional[List[str]] = None,
|
||||
provider_ids: list[str] | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
|
||||
@@ -4,7 +4,6 @@ Provider Query API 端点
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -40,7 +39,7 @@ class ModelsQueryRequest(BaseModel):
|
||||
"""模型列表查询请求"""
|
||||
|
||||
provider_id: str
|
||||
api_key_id: Optional[str] = None
|
||||
api_key_id: str | None = None
|
||||
force_refresh: bool = False # 强制刷新,跳过缓存
|
||||
|
||||
|
||||
@@ -49,11 +48,11 @@ class TestModelRequest(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
model_name: str
|
||||
api_key_id: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None # 指定使用的端点ID
|
||||
api_key_id: str | None = None
|
||||
endpoint_id: str | None = None # 指定使用的端点ID
|
||||
stream: bool = False
|
||||
message: Optional[str] = "你好"
|
||||
api_format: Optional[str] = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||
message: str | None = "你好"
|
||||
api_format: str | None = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||
|
||||
|
||||
# ============ API Endpoints ============
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -25,11 +24,11 @@ pipeline = ApiRequestPipeline()
|
||||
|
||||
class ProviderBillingUpdate(BaseModel):
|
||||
billing_type: ProviderBillingType
|
||||
monthly_quota_usd: Optional[float] = None
|
||||
monthly_quota_usd: float | None = None
|
||||
quota_reset_day: int = Field(default=30, ge=1, le=365) # 重置周期(天数)
|
||||
quota_last_reset_at: Optional[str] = None # 当前周期开始时间
|
||||
quota_expires_at: Optional[str] = None
|
||||
rpm_limit: Optional[int] = Field(default=None, ge=0)
|
||||
quota_last_reset_at: str | None = None # 当前周期开始时间
|
||||
quota_expires_at: str | None = None
|
||||
rpm_limit: int | None = Field(default=None, ge=0)
|
||||
provider_priority: int = Field(default=100, ge=0, le=200)
|
||||
|
||||
|
||||
@@ -163,13 +162,12 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
provider.quota_reset_day = config.quota_reset_day
|
||||
provider.provider_priority = config.provider_priority
|
||||
|
||||
from dateutil import parser
|
||||
from sqlalchemy import func
|
||||
|
||||
from src.models.database import Usage
|
||||
|
||||
if config.quota_last_reset_at:
|
||||
new_reset_at = parser.parse(config.quota_last_reset_at)
|
||||
new_reset_at = datetime.fromisoformat(config.quota_last_reset_at)
|
||||
# 确保有时区信息,如果没有则假设为 UTC
|
||||
if new_reset_at.tzinfo is None:
|
||||
new_reset_at = new_reset_at.replace(tzinfo=timezone.utc)
|
||||
@@ -188,7 +186,7 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
|
||||
|
||||
if config.quota_expires_at:
|
||||
expires_at = parser.parse(config.quota_expires_at)
|
||||
expires_at = datetime.fromisoformat(config.quota_expires_at)
|
||||
# 确保有时区信息,如果没有则假设为 UTC
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||
|
||||
@@ -3,7 +3,7 @@ Provider 模型管理 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
@@ -40,15 +40,15 @@ router = APIRouter(tags=["Model Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/{provider_id}/models", response_model=List[ModelResponse])
|
||||
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
|
||||
async def list_provider_models(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ModelResponse]:
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
获取提供商的所有模型
|
||||
|
||||
@@ -222,13 +222,13 @@ async def delete_provider_model(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{provider_id}/models/batch", response_model=List[ModelResponse])
|
||||
@router.post("/{provider_id}/models/batch", response_model=list[ModelResponse])
|
||||
async def batch_create_provider_models(
|
||||
provider_id: str,
|
||||
models_data: List[ModelCreate],
|
||||
models_data: list[ModelCreate],
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ModelResponse]:
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
批量创建模型
|
||||
|
||||
@@ -375,7 +375,7 @@ async def import_models_from_upstream(
|
||||
@dataclass
|
||||
class AdminListProviderModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
is_active: Optional[bool]
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -482,7 +482,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
models_data: List[ModelCreate]
|
||||
models_data: list[ModelCreate]
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
@@ -525,7 +525,7 @@ class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 2. 构建以 GlobalModel 为主键的字典
|
||||
global_models_dict: Dict[str, Dict[str, Any]] = {}
|
||||
global_models_dict: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for model in models:
|
||||
global_model = model.global_model
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
@@ -48,7 +47,7 @@ class MappingMatchingGlobalModel(BaseModel):
|
||||
global_model_name: str
|
||||
display_name: str
|
||||
is_active: bool
|
||||
matched_models: List[MappingMatchedModel] = Field(
|
||||
matched_models: list[MappingMatchedModel] = Field(
|
||||
default_factory=list, description="匹配到的模型列表"
|
||||
)
|
||||
|
||||
@@ -62,8 +61,8 @@ class MappingMatchingKey(BaseModel):
|
||||
key_name: str
|
||||
masked_key: str
|
||||
is_active: bool
|
||||
allowed_models: List[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||
matching_global_models: List[MappingMatchingGlobalModel] = Field(
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||
matching_global_models: list[MappingMatchingGlobalModel] = Field(
|
||||
default_factory=list, description="匹配到的 GlobalModel 列表"
|
||||
)
|
||||
|
||||
@@ -75,7 +74,7 @@ class ProviderMappingPreviewResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
provider_name: str
|
||||
keys: List[MappingMatchingKey] = Field(
|
||||
keys: list[MappingMatchingKey] = Field(
|
||||
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
|
||||
)
|
||||
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
|
||||
@@ -95,7 +94,7 @@ async def list_providers(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -209,7 +208,7 @@ async def delete_provider(provider_id: str, request: Request, db: Session = Depe
|
||||
|
||||
|
||||
class AdminListProvidersAdapter(AdminApiAdapter):
|
||||
def __init__(self, skip: int, limit: int, is_active: Optional[bool]):
|
||||
def __init__(self, skip: int, limit: int, is_active: bool | None):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
self.is_active = is_active
|
||||
@@ -473,7 +472,7 @@ async def get_provider_mapping_preview(
|
||||
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
|
||||
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
logger.warning(f"映射预览超时: provider_id={provider_id}")
|
||||
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
|
||||
|
||||
@@ -565,7 +564,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
|
||||
|
||||
# 构建有映射配置的 GlobalModel 映射
|
||||
models_with_mappings: Dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||
models_with_mappings: dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||
for gm in global_models:
|
||||
config = gm.config or {}
|
||||
mappings = config.get("model_mappings", [])
|
||||
@@ -585,7 +584,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
truncated_models=0,
|
||||
)
|
||||
|
||||
key_infos: List[MappingMatchingKey] = []
|
||||
key_infos: list[MappingMatchingKey] = []
|
||||
total_matches = 0
|
||||
|
||||
# 创建 CryptoService 实例
|
||||
@@ -611,10 +610,10 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
pass
|
||||
|
||||
# 查找匹配的 GlobalModel
|
||||
matching_global_models: List[MappingMatchingGlobalModel] = []
|
||||
matching_global_models: list[MappingMatchingGlobalModel] = []
|
||||
|
||||
for gm_id, (gm, mappings) in models_with_mappings.items():
|
||||
matched_models: List[MappingMatchedModel] = []
|
||||
matched_models: list[MappingMatchedModel] = []
|
||||
|
||||
for allowed_model in allowed_models_list:
|
||||
for mapping_pattern in mappings:
|
||||
|
||||
@@ -4,7 +4,6 @@ Provider 摘要与健康监控 API
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import case, func
|
||||
@@ -35,11 +34,11 @@ router = APIRouter(tags=["Provider Summary"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/summary", response_model=List[ProviderWithEndpointsSummary])
|
||||
@router.get("/summary", response_model=list[ProviderWithEndpointsSummary])
|
||||
async def get_providers_summary(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ProviderWithEndpointsSummary]:
|
||||
) -> list[ProviderWithEndpointsSummary]:
|
||||
"""
|
||||
获取所有提供商摘要信息
|
||||
|
||||
@@ -381,8 +380,8 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
attempts = attempts_query.limit(limit_rows).all()
|
||||
|
||||
buffered_attempts: Dict[str, List[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||
counters: Dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||
buffered_attempts: dict[str, list[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||
counters: dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||
|
||||
for attempt in attempts:
|
||||
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
|
||||
@@ -392,10 +391,10 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
buffered_attempts[attempt.endpoint_id].append(attempt)
|
||||
counters[attempt.endpoint_id] += 1
|
||||
|
||||
endpoint_monitors: List[EndpointHealthMonitor] = []
|
||||
endpoint_monitors: list[EndpointHealthMonitor] = []
|
||||
for endpoint in endpoints:
|
||||
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
|
||||
events: List[EndpointHealthEvent] = []
|
||||
events: list[EndpointHealthEvent] = []
|
||||
for attempt in attempt_list:
|
||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||
events.append(
|
||||
|
||||
@@ -4,8 +4,6 @@ IP 安全管理接口
|
||||
提供 IP 黑白名单管理和速率限制统计
|
||||
"""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -14,7 +12,6 @@ from src.api.base.adapter import ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
|
||||
@@ -30,7 +27,7 @@ class AddIPToBlacklistRequest(BaseModel):
|
||||
|
||||
ip_address: str = Field(..., description="IP 地址")
|
||||
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因")
|
||||
ttl: Optional[int] = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||
ttl: int | None = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||
|
||||
|
||||
class RemoveIPFromBlacklistRequest(BaseModel):
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
"""系统设置API端点。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import ValidationError
|
||||
@@ -649,9 +647,8 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
|
||||
class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
"""手动触发清理任务"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import func
|
||||
|
||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -34,8 +33,8 @@ pipeline = ApiRequestPipeline()
|
||||
async def get_usage_aggregation(
|
||||
request: Request,
|
||||
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
limit: int = Query(20, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -75,8 +74,8 @@ async def get_usage_aggregation(
|
||||
@router.get("/stats")
|
||||
async def get_usage_stats(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -122,14 +121,14 @@ async def get_activity_heatmap(
|
||||
@router.get("/records")
|
||||
async def get_usage_records(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
search: Optional[str] = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||
user_id: Optional[str] = None,
|
||||
username: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
provider: Optional[str] = None,
|
||||
status: Optional[str] = None, # stream, standard, error
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
search: str | None = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||
user_id: str | None = None,
|
||||
username: str | None = None,
|
||||
model: str | None = None,
|
||||
provider: str | None = None,
|
||||
status: str | None = None, # stream, standard, error
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
offset: int = Query(0, ge=0),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -179,7 +178,7 @@ async def get_usage_records(
|
||||
@router.get("/active")
|
||||
async def get_active_requests(
|
||||
request: Request,
|
||||
ids: Optional[str] = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||
ids: str | None = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -259,7 +258,7 @@ async def get_usage_detail(
|
||||
|
||||
|
||||
class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime]):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
|
||||
@@ -339,7 +338,7 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -386,7 +385,7 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -436,7 +435,7 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -446,7 +445,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
|
||||
# 这样可以正确统计 Fallback 场景(一个请求可能尝试多个 Provider)
|
||||
from sqlalchemy import case, Integer
|
||||
from sqlalchemy import case
|
||||
|
||||
attempt_query = db.query(
|
||||
RequestCandidate.provider_id,
|
||||
@@ -550,7 +549,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -608,14 +607,14 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
def __init__(
|
||||
self,
|
||||
start_date: Optional[datetime],
|
||||
end_date: Optional[datetime],
|
||||
search: Optional[str],
|
||||
user_id: Optional[str],
|
||||
username: Optional[str],
|
||||
model: Optional[str],
|
||||
provider: Optional[str],
|
||||
status: Optional[str],
|
||||
start_date: datetime | None,
|
||||
end_date: datetime | None,
|
||||
search: str | None,
|
||||
user_id: str | None,
|
||||
username: str | None,
|
||||
model: str | None,
|
||||
provider: str | None,
|
||||
status: str | None,
|
||||
limit: int,
|
||||
offset: int,
|
||||
):
|
||||
@@ -744,7 +743,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
|
||||
for req_id, candidates in request_candidates.items():
|
||||
# 提取所有不同的 candidate_index
|
||||
unique_candidates = set(c[0] for c in candidates)
|
||||
unique_candidates = {c[0] for c in candidates}
|
||||
# 如果有多个不同的 candidate_index,说明发生了 Fallback(Provider 切换)
|
||||
fallback_map[req_id] = len(unique_candidates) > 1
|
||||
|
||||
@@ -877,7 +876,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
class AdminActiveRequestsAdapter(AdminApiAdapter):
|
||||
"""轻量级活跃请求状态查询适配器"""
|
||||
|
||||
def __init__(self, ids: Optional[str]):
|
||||
def __init__(self, ids: str | None):
|
||||
self.ids = ids
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
@@ -1033,8 +1032,8 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
@router.get("/cache-affinity/ttl-analysis")
|
||||
async def analyze_cache_affinity_ttl(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1057,8 +1056,8 @@ async def analyze_cache_affinity_ttl(
|
||||
@router.get("/cache-affinity/hit-analysis")
|
||||
async def analyze_cache_hit(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1080,8 +1079,8 @@ class CacheAffinityTTLAnalysisAdapter(AdminApiAdapter):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
api_key_id: Optional[str],
|
||||
user_id: str | None,
|
||||
api_key_id: str | None,
|
||||
hours: int,
|
||||
):
|
||||
self.user_id = user_id
|
||||
@@ -1114,8 +1113,8 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
api_key_id: Optional[str],
|
||||
user_id: str | None,
|
||||
api_key_id: str | None,
|
||||
hours: int,
|
||||
):
|
||||
self.user_id = user_id
|
||||
@@ -1147,7 +1146,7 @@ async def get_interval_timeline(
|
||||
request: Request,
|
||||
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1177,7 +1176,7 @@ class IntervalTimelineAdapter(AdminApiAdapter):
|
||||
self,
|
||||
hours: int,
|
||||
limit: int,
|
||||
user_id: Optional[str] = None,
|
||||
user_id: str | None = None,
|
||||
include_user_info: bool = False,
|
||||
):
|
||||
self.hours = hours
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""用户管理 API 端点。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
@@ -48,8 +47,8 @@ async def list_users(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
||||
role: Optional[str] = Query(None, description="按角色筛选(user/admin)"),
|
||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
||||
role: str | None = Query(None, description="按角色筛选(user/admin)"),
|
||||
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -136,7 +135,7 @@ async def reset_user_quota(user_id: str, request: Request, db: Session = Depends
|
||||
async def get_user_api_keys(
|
||||
user_id: str,
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
||||
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -274,7 +273,7 @@ class AdminCreateUserAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminListUsersAdapter(AdminApiAdapter):
|
||||
def __init__(self, skip: int, limit: int, role: Optional[str], is_active: Optional[bool]):
|
||||
def __init__(self, skip: int, limit: int, role: str | None, is_active: bool | None):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
self.role = role
|
||||
@@ -467,7 +466,7 @@ class AdminResetUserQuotaAdapter(AdminApiAdapter):
|
||||
class AdminGetUserKeysAdapter(AdminApiAdapter):
|
||||
"""获取用户的API Keys"""
|
||||
|
||||
def __init__(self, user_id: str, is_active: Optional[bool]):
|
||||
def __init__(self, user_id: str, is_active: bool | None):
|
||||
self.user_id = user_id
|
||||
self.is_active = is_active
|
||||
|
||||
|
||||
Reference in New Issue
Block a user