mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
feat: 添加 OAuth 认证支持及相关改进
- 新增 OAuth 模块,支持 LinuxDo/GitHub/Google 等第三方登录 - 用户邮箱改为可选字段,支持无邮箱注册 - 新增模块配置验证状态 (config_validated/config_error) - 系统设置界面改为分块独立保存 - 用户设置新增 OAuth 绑定管理和首次密码设置 - 登录界面支持 OAuth 按钮展示 - 邮箱验证设置移至邮件设置页面
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -266,8 +266,21 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
||||
if not user_id:
|
||||
return None
|
||||
user = (
|
||||
context.db.query(User).filter(User.id == user_id, User.is_active.is_(True)).first()
|
||||
context.db.query(User)
|
||||
.filter(
|
||||
User.id == user_id,
|
||||
User.is_active.is_(True),
|
||||
User.is_deleted.is_(False),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if not user:
|
||||
return None
|
||||
|
||||
if not AuthService.token_identity_matches_user(payload, user):
|
||||
return None
|
||||
|
||||
return user
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@@ -299,13 +299,12 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
||||
access_token = AuthService.create_access_token(
|
||||
data={
|
||||
"user_id": user.id,
|
||||
"email": user.email,
|
||||
"role": user.role.value,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
)
|
||||
refresh_token = AuthService.create_refresh_token(
|
||||
data={"user_id": user.id, "email": user.email}
|
||||
data={"user_id": user.id, "created_at": user.created_at.isoformat() if user.created_at else None}
|
||||
)
|
||||
response = LoginResponse(
|
||||
access_token=access_token,
|
||||
@@ -345,19 +344,23 @@ class AuthRefreshAdapter(AuthPublicAdapter):
|
||||
)
|
||||
if not user.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="用户已禁用")
|
||||
if user.is_deleted:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="用户不存在或已禁用")
|
||||
|
||||
if not AuthService.token_identity_matches_user(token_payload, user):
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无效的刷新令牌")
|
||||
|
||||
new_access_token = AuthService.create_access_token(
|
||||
data={
|
||||
"user_id": user.id,
|
||||
"email": user.email,
|
||||
"role": user.role.value,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
)
|
||||
new_refresh_token = AuthService.create_refresh_token(
|
||||
data={"user_id": user.id, "email": user.email}
|
||||
data={"user_id": user.id, "created_at": user.created_at.isoformat() if user.created_at else None}
|
||||
)
|
||||
logger.info(f"令牌刷新成功: {user.email}")
|
||||
logger.info(f"令牌刷新成功: user_id={user.id}")
|
||||
return RefreshTokenResponse(
|
||||
access_token=new_access_token,
|
||||
refresh_token=new_refresh_token,
|
||||
@@ -378,10 +381,16 @@ class AuthRegistrationSettingsAdapter(AuthPublicAdapter):
|
||||
|
||||
enable_registration = SystemConfigService.get_config(db, "enable_registration", default=False)
|
||||
require_verification = SystemConfigService.get_config(db, "require_email_verification", default=False)
|
||||
email_configured = EmailSenderService.is_smtp_configured(db)
|
||||
|
||||
# 如果邮箱服务未配置,强制 require_email_verification 为 False
|
||||
if not email_configured:
|
||||
require_verification = False
|
||||
|
||||
return RegistrationSettingsResponse(
|
||||
enable_registration=bool(enable_registration),
|
||||
require_email_verification=bool(require_verification),
|
||||
email_configured=email_configured,
|
||||
).model_dump()
|
||||
|
||||
|
||||
@@ -430,61 +439,78 @@ class AuthRegisterAdapter(AuthPublicAdapter):
|
||||
AuditService.log_event(
|
||||
db=db,
|
||||
event_type=AuditEventType.UNAUTHORIZED_ACCESS,
|
||||
description=f"Registration attempt rejected - registration disabled: {register_request.email}",
|
||||
description=f"Registration attempt rejected - registration disabled: {register_request.username}",
|
||||
ip_address=client_ip,
|
||||
user_agent=user_agent,
|
||||
metadata={"email": register_request.email, "reason": "registration_disabled"},
|
||||
metadata={"username": register_request.username, "reason": "registration_disabled"},
|
||||
)
|
||||
db.commit()
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="系统暂不开放注册")
|
||||
|
||||
# 检查邮箱后缀是否允许
|
||||
suffix_allowed, suffix_error = validate_email_suffix(db, register_request.email)
|
||||
if not suffix_allowed:
|
||||
logger.warning(f"注册失败:邮箱后缀不允许: {register_request.email}")
|
||||
AuditService.log_event(
|
||||
db=db,
|
||||
event_type=AuditEventType.UNAUTHORIZED_ACCESS,
|
||||
description=f"Registration attempt rejected - email suffix not allowed: {register_request.email}",
|
||||
ip_address=client_ip,
|
||||
user_agent=user_agent,
|
||||
metadata={"email": register_request.email, "reason": "email_suffix_not_allowed"},
|
||||
)
|
||||
db.commit()
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=suffix_error,
|
||||
)
|
||||
|
||||
# 检查是否需要邮箱验证
|
||||
email = register_request.email
|
||||
email_configured = EmailSenderService.is_smtp_configured(db)
|
||||
require_verification = SystemConfigService.get_config(db, "require_email_verification", default=False)
|
||||
|
||||
# 如果邮箱服务未配置,强制不要求邮箱验证
|
||||
if not email_configured:
|
||||
require_verification = False
|
||||
|
||||
# 如果系统要求邮箱验证,则必须提供邮箱
|
||||
if require_verification:
|
||||
if not email:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="系统要求邮箱验证,请填写邮箱",
|
||||
)
|
||||
# 检查邮箱是否已验证
|
||||
is_verified = await EmailVerificationService.is_email_verified(register_request.email)
|
||||
is_verified = await EmailVerificationService.is_email_verified(email)
|
||||
if not is_verified:
|
||||
logger.warning(f"注册失败:邮箱未验证: {register_request.email}")
|
||||
logger.warning(f"注册失败:邮箱未验证: {email}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="请先完成邮箱验证。请发送验证码并验证后再注册。",
|
||||
)
|
||||
|
||||
# 如果提供了邮箱,进行后缀验证
|
||||
if email:
|
||||
suffix_allowed, suffix_error = validate_email_suffix(db, email)
|
||||
if not suffix_allowed:
|
||||
logger.warning(f"注册失败:邮箱后缀不允许: {email}")
|
||||
AuditService.log_event(
|
||||
db=db,
|
||||
event_type=AuditEventType.UNAUTHORIZED_ACCESS,
|
||||
description=f"Registration attempt rejected - email suffix not allowed: {email}",
|
||||
ip_address=client_ip,
|
||||
user_agent=user_agent,
|
||||
metadata={"email": email, "reason": "email_suffix_not_allowed"},
|
||||
)
|
||||
db.commit()
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=suffix_error,
|
||||
)
|
||||
|
||||
try:
|
||||
# 读取系统配置的默认配额
|
||||
default_quota = SystemConfigService.get_config(db, "default_user_quota_usd", default=10.0)
|
||||
|
||||
# email_verified 逻辑:
|
||||
# - 要求邮箱验证且已通过验证:True
|
||||
# - 提供了邮箱但不要求验证:False(用户可后续自行验证)
|
||||
# - 未提供邮箱:False
|
||||
user = UserService.create_user(
|
||||
db=db,
|
||||
email=register_request.email,
|
||||
email=email, # 可以为 None
|
||||
username=register_request.username,
|
||||
password=register_request.password,
|
||||
role=UserRole.USER,
|
||||
quota_usd=default_quota,
|
||||
email_verified=bool(require_verification and email),
|
||||
)
|
||||
AuditService.log_event(
|
||||
db=db,
|
||||
event_type=AuditEventType.USER_CREATED,
|
||||
description=f"User registered: {user.email}",
|
||||
description=f"User registered: {user.username}" + (f" ({user.email})" if user.email else ""),
|
||||
user_id=user.id,
|
||||
ip_address=client_ip,
|
||||
user_agent=user_agent,
|
||||
@@ -494,9 +520,9 @@ class AuthRegisterAdapter(AuthPublicAdapter):
|
||||
db.commit()
|
||||
|
||||
# 注册成功后清除验证状态(在 commit 后清理,即使清理失败也不影响注册结果)
|
||||
if require_verification:
|
||||
if require_verification and email:
|
||||
try:
|
||||
await EmailVerificationService.clear_verification(register_request.email)
|
||||
await EmailVerificationService.clear_verification(email)
|
||||
except Exception as e:
|
||||
logger.warning(f"清理验证状态失败: {e}")
|
||||
|
||||
@@ -510,10 +536,10 @@ class AuthRegisterAdapter(AuthPublicAdapter):
|
||||
AuditService.log_event(
|
||||
db=db,
|
||||
event_type=AuditEventType.UNAUTHORIZED_ACCESS,
|
||||
description=f"Registration failed: {register_request.email} - {exc}",
|
||||
description=f"Registration failed: {register_request.username} - {exc}",
|
||||
ip_address=client_ip,
|
||||
user_agent=user_agent,
|
||||
metadata={"email": register_request.email, "error": str(exc)},
|
||||
metadata={"username": register_request.username, "error": str(exc)},
|
||||
)
|
||||
db.commit()
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc))
|
||||
|
||||
@@ -238,9 +238,12 @@ class ApiRequestPipeline:
|
||||
|
||||
# 直接查询数据库,确保返回的是当前 Session 绑定的对象
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user or not user.is_active:
|
||||
if not user or not user.is_active or user.is_deleted:
|
||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
||||
|
||||
if not self.auth_service.token_identity_matches_user(payload, user):
|
||||
raise HTTPException(status_code=403, detail="无效的管理员令牌")
|
||||
|
||||
# 检查管理员权限
|
||||
if user.role != UserRole.ADMIN:
|
||||
logger.warning(f"非管理员尝试通过 JWT 访问管理端点: {user.email}")
|
||||
@@ -291,9 +294,12 @@ class ApiRequestPipeline:
|
||||
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
||||
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user or not user.is_active:
|
||||
if not user or not user.is_active or user.is_deleted:
|
||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
||||
|
||||
if not self.auth_service.token_identity_matches_user(payload, user):
|
||||
raise HTTPException(status_code=403, detail="无效的用户令牌")
|
||||
|
||||
request.state.user_id = user.id
|
||||
return user, None
|
||||
|
||||
|
||||
15
src/api/oauth/__init__.py
Normal file
15
src/api/oauth/__init__.py
Normal file
@@ -0,0 +1,15 @@
|
||||
"""OAuth API 路由聚合。"""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from src.api.oauth.admin import router as admin_router
|
||||
from src.api.oauth.public import router as public_router
|
||||
from src.api.oauth.user import router as user_router
|
||||
|
||||
router = APIRouter()
|
||||
router.include_router(public_router)
|
||||
router.include_router(user_router)
|
||||
router.include_router(admin_router)
|
||||
|
||||
__all__ = ["router"]
|
||||
|
||||
264
src/api/oauth/admin.py
Normal file
264
src/api/oauth/admin.py
Normal file
@@ -0,0 +1,264 @@
|
||||
"""OAuth 管理端点(管理员)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, 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 OAuthProvider
|
||||
from src.services.auth.oauth.registry import get_oauth_provider_registry
|
||||
from src.services.auth.oauth.service import OAuthService
|
||||
|
||||
router = APIRouter(prefix="/api/admin/oauth", tags=["Admin - OAuth"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
class SupportedOAuthType(BaseModel):
|
||||
provider_type: str
|
||||
display_name: str
|
||||
default_authorization_url: str
|
||||
default_token_url: str
|
||||
default_userinfo_url: str
|
||||
default_scopes: List[str]
|
||||
|
||||
|
||||
class OAuthProviderUpsertRequest(BaseModel):
|
||||
display_name: str = Field(..., min_length=1, max_length=100)
|
||||
client_id: str = Field(..., min_length=1, max_length=255)
|
||||
client_secret: Optional[str] = Field(None, max_length=2048)
|
||||
|
||||
authorization_url_override: Optional[str] = Field(None, max_length=500)
|
||||
token_url_override: Optional[str] = Field(None, max_length=500)
|
||||
userinfo_url_override: Optional[str] = Field(None, max_length=500)
|
||||
scopes: Optional[List[str]] = None
|
||||
|
||||
redirect_uri: str = Field(..., min_length=1, max_length=500)
|
||||
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
|
||||
|
||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
||||
extra_config: Optional[Dict[str, Any]] = None
|
||||
|
||||
is_enabled: bool = False
|
||||
force: bool = False
|
||||
|
||||
|
||||
class OAuthProviderAdminResponse(BaseModel):
|
||||
provider_type: str
|
||||
display_name: str
|
||||
client_id: str
|
||||
has_secret: bool
|
||||
authorization_url_override: Optional[str] = None
|
||||
token_url_override: Optional[str] = None
|
||||
userinfo_url_override: Optional[str] = None
|
||||
scopes: Optional[List[str]] = None
|
||||
redirect_uri: str
|
||||
frontend_callback_url: str
|
||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
||||
extra_config: Optional[Dict[str, Any]] = None
|
||||
is_enabled: bool
|
||||
|
||||
|
||||
class OAuthProviderTestResponse(BaseModel):
|
||||
authorization_url_reachable: bool
|
||||
token_url_reachable: bool
|
||||
secret_status: str
|
||||
details: str = ""
|
||||
|
||||
|
||||
class OAuthProviderTestRequest(BaseModel):
|
||||
"""测试请求,使用表单数据而非数据库配置"""
|
||||
|
||||
client_id: str = Field(..., min_length=1)
|
||||
client_secret: Optional[str] = None
|
||||
authorization_url_override: Optional[str] = None
|
||||
token_url_override: Optional[str] = None
|
||||
redirect_uri: str = Field(..., min_length=1)
|
||||
|
||||
|
||||
@router.get("/supported-types", response_model=List[SupportedOAuthType])
|
||||
async def get_supported_types(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = GetSupportedTypesAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/providers", response_model=List[OAuthProviderAdminResponse])
|
||||
async def list_provider_configs(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = ListOAuthProviderConfigsAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/providers/{provider_type}", response_model=OAuthProviderAdminResponse)
|
||||
async def get_provider_config(provider_type: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = GetOAuthProviderConfigAdapter(provider_type=provider_type)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.put("/providers/{provider_type}", response_model=OAuthProviderAdminResponse)
|
||||
async def upsert_provider_config(provider_type: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = UpsertOAuthProviderConfigAdapter(provider_type=provider_type)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.delete("/providers/{provider_type}")
|
||||
async def delete_provider_config(provider_type: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = DeleteOAuthProviderConfigAdapter(provider_type=provider_type)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/providers/{provider_type}/test", response_model=OAuthProviderTestResponse)
|
||||
async def test_provider_config(provider_type: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = TestOAuthProviderConfigAdapter(provider_type=provider_type)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
class GetSupportedTypesAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
registry = get_oauth_provider_registry()
|
||||
types = registry.get_supported_types()
|
||||
return [
|
||||
SupportedOAuthType(
|
||||
provider_type=t.provider_type,
|
||||
display_name=t.display_name,
|
||||
default_authorization_url=t.default_authorization_url,
|
||||
default_token_url=t.default_token_url,
|
||||
default_userinfo_url=t.default_userinfo_url,
|
||||
default_scopes=list(t.default_scopes),
|
||||
).model_dump()
|
||||
for t in types
|
||||
]
|
||||
|
||||
|
||||
class ListOAuthProviderConfigsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
rows = context.db.query(OAuthProvider).order_by(OAuthProvider.provider_type.asc()).all()
|
||||
return [
|
||||
OAuthProviderAdminResponse(
|
||||
provider_type=str(row.provider_type or ""),
|
||||
display_name=str(row.display_name or ""),
|
||||
client_id=str(row.client_id or ""),
|
||||
has_secret=bool(row.client_secret_encrypted),
|
||||
authorization_url_override=row.authorization_url_override,
|
||||
token_url_override=row.token_url_override,
|
||||
userinfo_url_override=row.userinfo_url_override,
|
||||
scopes=row.scopes,
|
||||
redirect_uri=str(row.redirect_uri or ""),
|
||||
frontend_callback_url=str(row.frontend_callback_url or ""),
|
||||
attribute_mapping=row.attribute_mapping,
|
||||
extra_config=row.extra_config,
|
||||
is_enabled=bool(row.is_enabled),
|
||||
).model_dump()
|
||||
for row in rows
|
||||
]
|
||||
|
||||
|
||||
class GetOAuthProviderConfigAdapter(AdminApiAdapter):
|
||||
def __init__(self, provider_type: str):
|
||||
self.provider_type = provider_type
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
row = context.db.query(OAuthProvider).filter(OAuthProvider.provider_type == self.provider_type).first()
|
||||
if not row:
|
||||
raise InvalidRequestException("Provider 配置不存在")
|
||||
return OAuthProviderAdminResponse(
|
||||
provider_type=str(row.provider_type or ""),
|
||||
display_name=str(row.display_name or ""),
|
||||
client_id=str(row.client_id or ""),
|
||||
has_secret=bool(row.client_secret_encrypted),
|
||||
authorization_url_override=row.authorization_url_override,
|
||||
token_url_override=row.token_url_override,
|
||||
userinfo_url_override=row.userinfo_url_override,
|
||||
scopes=row.scopes,
|
||||
redirect_uri=str(row.redirect_uri or ""),
|
||||
frontend_callback_url=str(row.frontend_callback_url or ""),
|
||||
attribute_mapping=row.attribute_mapping,
|
||||
extra_config=row.extra_config,
|
||||
is_enabled=bool(row.is_enabled),
|
||||
).model_dump()
|
||||
|
||||
|
||||
class UpsertOAuthProviderConfigAdapter(AdminApiAdapter):
|
||||
def __init__(self, provider_type: str):
|
||||
self.provider_type = provider_type
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = OAuthProviderUpsertRequest.model_validate(payload)
|
||||
except ValidationError as exc:
|
||||
errors = exc.errors()
|
||||
if errors:
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
row = await OAuthService.upsert_provider_config(
|
||||
db=context.db,
|
||||
provider_type=self.provider_type,
|
||||
data=req,
|
||||
)
|
||||
|
||||
return OAuthProviderAdminResponse(
|
||||
provider_type=str(row.provider_type or ""),
|
||||
display_name=str(row.display_name or ""),
|
||||
client_id=str(row.client_id or ""),
|
||||
has_secret=bool(row.client_secret_encrypted),
|
||||
authorization_url_override=row.authorization_url_override,
|
||||
token_url_override=row.token_url_override,
|
||||
userinfo_url_override=row.userinfo_url_override,
|
||||
scopes=row.scopes,
|
||||
redirect_uri=str(row.redirect_uri or ""),
|
||||
frontend_callback_url=str(row.frontend_callback_url or ""),
|
||||
attribute_mapping=row.attribute_mapping,
|
||||
extra_config=row.extra_config,
|
||||
is_enabled=bool(row.is_enabled),
|
||||
).model_dump()
|
||||
|
||||
|
||||
class DeleteOAuthProviderConfigAdapter(AdminApiAdapter):
|
||||
def __init__(self, provider_type: str):
|
||||
self.provider_type = provider_type
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
await OAuthService.delete_provider_config(context.db, self.provider_type)
|
||||
return {"message": "删除成功"}
|
||||
|
||||
|
||||
class TestOAuthProviderConfigAdapter(AdminApiAdapter):
|
||||
def __init__(self, provider_type: str):
|
||||
self.provider_type = provider_type
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
req = OAuthProviderTestRequest.model_validate(payload)
|
||||
except ValidationError as exc:
|
||||
errors = exc.errors()
|
||||
if errors:
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
# 如果没有提供 client_secret,尝试从数据库获取已保存的
|
||||
client_secret = req.client_secret
|
||||
if not client_secret:
|
||||
existing = context.db.query(OAuthProvider).filter(
|
||||
OAuthProvider.provider_type == self.provider_type
|
||||
).first()
|
||||
if existing and existing.client_secret_encrypted:
|
||||
client_secret = existing.get_client_secret()
|
||||
|
||||
result = await OAuthService.test_provider_config_with_data(
|
||||
provider_type=self.provider_type,
|
||||
client_id=req.client_id,
|
||||
client_secret=client_secret,
|
||||
authorization_url_override=req.authorization_url_override,
|
||||
token_url_override=req.token_url_override,
|
||||
redirect_uri=req.redirect_uri,
|
||||
)
|
||||
return OAuthProviderTestResponse(**result).model_dump()
|
||||
59
src/api/oauth/public.py
Normal file
59
src/api/oauth/public.py
Normal file
@@ -0,0 +1,59 @@
|
||||
"""OAuth 公开端点(无需登录)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, status
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.responses import RedirectResponse
|
||||
|
||||
from src.database import get_db
|
||||
from src.services.auth.oauth.service import OAuthService
|
||||
|
||||
router = APIRouter(prefix="/api/oauth", tags=["OAuth"])
|
||||
|
||||
|
||||
@router.get("/providers")
|
||||
async def list_oauth_providers(db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
"""
|
||||
获取可用 OAuth Providers 列表。
|
||||
|
||||
模块未启用时返回空列表(前端友好)。
|
||||
"""
|
||||
providers = await OAuthService.list_public_providers(db)
|
||||
return {"providers": providers}
|
||||
|
||||
|
||||
@router.get("/{provider_type}/authorize")
|
||||
async def oauth_authorize(provider_type: str, db: Session = Depends(get_db)) -> RedirectResponse:
|
||||
"""
|
||||
发起 OAuth 登录(login flow)。
|
||||
"""
|
||||
url = await OAuthService.build_login_authorize_url(db, provider_type)
|
||||
return RedirectResponse(url=url, status_code=status.HTTP_302_FOUND)
|
||||
|
||||
|
||||
@router.get("/{provider_type}/callback")
|
||||
async def oauth_callback(
|
||||
provider_type: str,
|
||||
db: Session = Depends(get_db),
|
||||
code: Optional[str] = Query(None),
|
||||
state: Optional[str] = Query(None),
|
||||
error: Optional[str] = Query(None),
|
||||
error_description: Optional[str] = Query(None),
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
OAuth 回调端点。
|
||||
|
||||
成功/失败都会重定向到前端回调页。
|
||||
"""
|
||||
redirect_url = await OAuthService.handle_callback(
|
||||
db=db,
|
||||
provider_type=provider_type,
|
||||
state=state or "",
|
||||
code=code,
|
||||
error=error,
|
||||
error_description=error_description,
|
||||
)
|
||||
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
|
||||
84
src/api/oauth/user.py
Normal file
84
src/api/oauth/user.py
Normal file
@@ -0,0 +1,84 @@
|
||||
"""OAuth 用户端点(需登录)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.responses import RedirectResponse
|
||||
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.database import get_db
|
||||
from src.services.auth.oauth.service import OAuthService
|
||||
|
||||
router = APIRouter(prefix="/api/user/oauth", tags=["User - OAuth"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/bindable-providers")
|
||||
async def list_bindable_providers(request: Request, db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
adapter = ListBindableProvidersAdapter()
|
||||
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
return cast(dict[str, Any], result)
|
||||
|
||||
|
||||
@router.get("/links")
|
||||
async def list_my_oauth_links(request: Request, db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
adapter = ListMyOAuthLinksAdapter()
|
||||
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
return cast(dict[str, Any], result)
|
||||
|
||||
|
||||
@router.get("/{provider_type}/bind")
|
||||
async def bind_oauth_provider(
|
||||
provider_type: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> RedirectResponse:
|
||||
adapter = BindOAuthProviderAdapter(provider_type=provider_type)
|
||||
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
return cast(RedirectResponse, result)
|
||||
|
||||
|
||||
@router.delete("/{provider_type}")
|
||||
async def unbind_oauth_provider(
|
||||
provider_type: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> dict[str, Any]:
|
||||
adapter = UnbindOAuthProviderAdapter(provider_type=provider_type)
|
||||
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
return cast(dict[str, Any], result)
|
||||
|
||||
|
||||
class ListBindableProvidersAdapter(AuthenticatedApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
assert context.user is not None
|
||||
providers = await OAuthService.list_bindable_providers(context.db, context.user)
|
||||
return {"providers": providers}
|
||||
|
||||
|
||||
class ListMyOAuthLinksAdapter(AuthenticatedApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
assert context.user is not None
|
||||
links = await OAuthService.list_user_links(context.db, context.user)
|
||||
return {"links": links}
|
||||
|
||||
|
||||
class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
def __init__(self, provider_type: str):
|
||||
self.provider_type = provider_type
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
|
||||
assert context.user is not None
|
||||
url = await OAuthService.build_bind_authorize_url(context.db, context.user, self.provider_type)
|
||||
return RedirectResponse(url=url, status_code=status.HTTP_302_FOUND)
|
||||
|
||||
|
||||
class UnbindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
def __init__(self, provider_type: str):
|
||||
self.provider_type = provider_type
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
assert context.user is not None
|
||||
await OAuthService.unbind_provider(context.db, context.user, self.provider_type)
|
||||
return {"message": "解绑成功"}
|
||||
@@ -446,16 +446,31 @@ class ChangePasswordAdapter(AuthenticatedApiAdapter):
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
if not user.verify_password(request.old_password):
|
||||
raise InvalidRequestException("旧密码错误")
|
||||
# LDAP 用户不能修改密码
|
||||
from src.core.enums import AuthSource
|
||||
if user.auth_source == AuthSource.LDAP:
|
||||
raise ForbiddenException("LDAP 用户不能在此修改密码")
|
||||
|
||||
# 判断用户是否已有密码
|
||||
has_password = bool(user.password_hash)
|
||||
|
||||
if has_password:
|
||||
# 已有密码:需要验证旧密码
|
||||
if not request.old_password:
|
||||
raise InvalidRequestException("请输入当前密码")
|
||||
if not user.verify_password(request.old_password):
|
||||
raise InvalidRequestException("旧密码错误")
|
||||
# 无密码(如 OAuth 用户首次设置):无需旧密码
|
||||
|
||||
if len(request.new_password) < 6:
|
||||
raise InvalidRequestException("密码长度至少6位")
|
||||
|
||||
user.set_password(request.new_password)
|
||||
user.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
logger.info(f"用户修改密码: {user.email}")
|
||||
return {"message": "密码修改成功"}
|
||||
action = "修改" if has_password else "设置"
|
||||
logger.info(f"用户{action}密码: {user.email}")
|
||||
return {"message": f"密码{action}成功"}
|
||||
|
||||
|
||||
class ListMyApiKeysAdapter(AuthenticatedApiAdapter):
|
||||
|
||||
Reference in New Issue
Block a user