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:
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": "解绑成功"}
|
||||
Reference in New Issue
Block a user