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

View File

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

View File

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

View File

@@ -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
View 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
View 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
View 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
View 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": "解绑成功"}

View File

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