feat: 添加 OAuth 认证支持及相关改进

- 新增 OAuth 模块,支持 LinuxDo/GitHub/Google 等第三方登录
- 用户邮箱改为可选字段,支持无邮箱注册
- 新增模块配置验证状态 (config_validated/config_error)
- 系统设置界面改为分块独立保存
- 用户设置新增 OAuth 绑定管理和首次密码设置
- 登录界面支持 OAuth 按钮展示
- 邮箱验证设置移至邮件设置页面
This commit is contained in:
fawney19
2026-01-19 03:19:17 +08:00
parent e2e14fd09c
commit 3d88dfd98a
61 changed files with 4548 additions and 950 deletions

View File

@@ -27,6 +27,8 @@ class ModuleStatusResponse(BaseModel):
available: bool
enabled: bool
active: bool
config_validated: bool
config_error: Optional[str]
display_name: str
description: str
category: str
@@ -43,6 +45,8 @@ class ModuleStatusResponse(BaseModel):
available=status.available,
enabled=status.enabled,
active=status.active,
config_validated=status.config_validated,
config_error=status.config_error,
display_name=status.display_name,
description=status.description,
category=status.category.value,
@@ -180,6 +184,12 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
except Exception:
raise InvalidRequestException("请求体格式错误,需要 enabled 字段")
# 如果是启用模块,必须先通过配置验证
if req.enabled:
config_validated, config_error = registry.validate_config(self.module_name, context.db)
if not config_validated:
raise InvalidRequestException(f"模块配置未验证通过: {config_error}")
# 设置启用状态
registry.set_enabled(self.module_name, req.enabled, context.db)

View File

@@ -3,7 +3,6 @@ Provider Query API 端点
用于查询提供商的模型列表等信息
"""
import asyncio
from typing import Optional
import httpx
@@ -11,14 +10,17 @@ from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
from sqlalchemy.orm import Session, joinedload
from src.api.handlers.base.chat_adapter_base import get_adapter_class
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class
from src.config.constants import TimeoutDefaults
from src.core.crypto import crypto_service
from src.core.headers import get_extra_headers_from_endpoint
from src.core.logger import logger
from src.database.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User
from src.models.database import Provider, ProviderEndpoint, User
from src.services.model.upstream_fetcher import (
_get_adapter_for_format,
build_all_format_configs,
fetch_models_from_endpoints,
)
from src.utils.auth_utils import get_current_user
@@ -50,21 +52,6 @@ class TestModelRequest(BaseModel):
# ============ API Endpoints ============
def _get_adapter_for_format(api_format: str):
"""根据 API 格式获取对应的 Adapter 类"""
# 先检查 Chat Adapter 注册表
adapter_class = get_adapter_class(api_format)
if adapter_class:
return adapter_class
# 再检查 CLI Adapter 注册表
cli_adapter_class = get_cli_adapter_class(api_format)
if cli_adapter_class:
return cli_adapter_class
return None
@router.post("/models")
async def query_available_models(
request: ModelsQueryRequest,
@@ -75,11 +62,7 @@ async def query_available_models(
查询提供商可用模型
优先从缓存获取(缓存由定时任务刷新),缓存未命中时实时调用上游 API。
遍历所有活跃端点,根据端点的 API 格式选择正确的 Adapter 进行请求:
- OPENAI/OPENAI_CLI: 使用 OpenAIChatAdapter.fetch_models
- CLAUDE/CLAUDE_CLI: 使用 ClaudeChatAdapter.fetch_models
- GEMINI/GEMINI_CLI: 使用 GeminiChatAdapter.fetch_models
从所有 API 格式尝试获取模型,然后聚合去重。
Args:
request: 查询请求
@@ -107,9 +90,6 @@ async def query_available_models(
raise HTTPException(status_code=404, detail="Provider not found")
# 如果指定了 api_key_id 且不是强制刷新,优先从缓存获取
# 注:不指定 api_key_id 时Provider 级别查询)不使用缓存,因为:
# 1. Provider 级别查询会遍历多个 Key结果不稳定
# 2. 缓存按 Key 粒度存储,与定时任务的刷新逻辑一致
if request.api_key_id and not request.force_refresh:
cached_models = await get_upstream_models_from_cache(
request.provider_id, request.api_key_id
@@ -126,127 +106,42 @@ async def query_available_models(
# 缓存未命中或强制刷新,实时获取
# 收集所有活跃端点的配置
endpoint_configs: list[dict] = []
# 构建 api_format -> endpoint 映射
format_to_endpoint: dict[str, ProviderEndpoint] = {}
for endpoint in provider.endpoints:
if endpoint.is_active:
format_to_endpoint[endpoint.api_format] = endpoint
if not format_to_endpoint:
raise HTTPException(status_code=400, detail="No active endpoints found for this provider")
# 获取 API Key
if request.api_key_id:
# 指定了特定的 API Key(从 provider.api_keys 查找)
# 指定了特定的 API Key
api_key = next(
(key for key in provider.api_keys if key.id == request.api_key_id),
None
)
if not api_key:
raise HTTPException(status_code=404, detail="API Key not found")
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
except Exception as e:
logger.error(f"Failed to decrypt API key: {e}")
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
# 根据 Key 的 api_formats 找对应的 Endpoint
key_formats = api_key.api_formats or []
for fmt in key_formats:
endpoint = format_to_endpoint.get(fmt)
if endpoint:
endpoint_configs.append({
"api_key": api_key_value,
"base_url": endpoint.base_url,
"api_format": fmt,
"extra_headers": get_extra_headers_from_endpoint(endpoint),
})
if not endpoint_configs:
raise HTTPException(
status_code=400,
detail="No matching endpoint found for this API Key's formats"
)
else:
# 遍历所有活跃端点,为每个端点找一个支持该格式的 Key
for endpoint in provider.endpoints:
if not endpoint.is_active:
continue
# 找第一个支持该格式的可用 Key
for api_key in provider.api_keys:
if not api_key.is_active:
continue
key_formats = api_key.api_formats or []
if endpoint.api_format not in key_formats:
continue
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
except Exception as e:
logger.error(f"Failed to decrypt API key: {e}")
continue
endpoint_configs.append({
"api_key": api_key_value,
"base_url": endpoint.base_url,
"api_format": endpoint.api_format,
"extra_headers": get_extra_headers_from_endpoint(endpoint),
})
break # 只取第一个可用的 Key
if not endpoint_configs:
# 使用第一个可用的 Key
api_key = next(
(key for key in provider.api_keys if key.is_active),
None
)
if not api_key:
raise HTTPException(status_code=400, detail="No active API Key found for this provider")
# 并发请求所有端点的模型列表
all_models: list = []
errors: list[str] = []
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
except Exception as e:
logger.error(f"Failed to decrypt API key: {e}")
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
async def fetch_endpoint_models(
client: httpx.AsyncClient, config: dict
) -> tuple[list, Optional[str]]:
base_url = config["base_url"]
if not base_url:
return [], None
base_url = base_url.rstrip("/")
api_format = config["api_format"]
api_key_value = config["api_key"]
extra_headers = config.get("extra_headers")
try:
# 获取对应的 Adapter 类并调用 fetch_models
adapter_class = _get_adapter_for_format(api_format)
if not adapter_class:
return [], f"Unknown API format: {api_format}"
models, error = await adapter_class.fetch_models(
client, base_url, api_key_value, extra_headers
)
# 确保所有模型都有 api_format 字段
for m in models:
if "api_format" not in m:
m["api_format"] = api_format
return models, error
except Exception as e:
logger.error(f"Error fetching models from {api_format} endpoint: {e}")
return [], f"{api_format}: {str(e)}"
# 限制并发请求数量,避免触发上游速率限制
MAX_CONCURRENT_REQUESTS = 5
semaphore = asyncio.Semaphore(MAX_CONCURRENT_REQUESTS)
async def fetch_with_semaphore(
client: httpx.AsyncClient, config: dict
) -> tuple[list, Optional[str]]:
async with semaphore:
return await fetch_endpoint_models(client, config)
async with httpx.AsyncClient(timeout=30.0) as client:
results = await asyncio.gather(
*[fetch_with_semaphore(client, c) for c in endpoint_configs]
)
for models, error in results:
all_models.extend(models)
if error:
errors.append(error)
# 使用公共函数构建所有格式的端点配置并获取模型
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
all_models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
# 按 model id + api_format 去重(保留第一个)
seen_keys: set[str] = set()

View File

@@ -1486,8 +1486,17 @@ class AdminImportUsersAdapter(AdminApiAdapter):
stats["users"]["skipped"] += 1
continue
# 导入必须有邮箱email 是导入的主键)
import_email = user_data.get("email")
if not import_email:
stats["errors"].append(
f"跳过无邮箱用户: {user_data.get('username', '未知')}"
)
stats["users"]["skipped"] += 1
continue
existing_user = (
db.query(User).filter(User.email == user_data["email"]).first()
db.query(User).filter(User.email == import_email).first()
)
if existing_user:
@@ -1496,7 +1505,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
stats["users"]["skipped"] += 1
elif merge_mode == "error":
raise InvalidRequestException(
f"用户 '{user_data['email']}' 已存在"
f"用户 '{import_email}' 已存在"
)
elif merge_mode == "overwrite":
# 更新现有用户
@@ -1527,8 +1536,9 @@ class AdminImportUsersAdapter(AdminApiAdapter):
new_user = User(
id=str(uuid.uuid4()),
email=user_data["email"],
username=user_data.get("username", user_data["email"].split("@")[0]),
email=import_email,
email_verified=user_data.get("email_verified", True),
username=user_data.get("username") or import_email.split("@")[0],
password_hash=user_data.get("password_hash", ""),
role=role,
allowed_providers=user_data.get("allowed_providers"),