mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/ - 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层 - 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构 - 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations) - 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image - 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
14
_deprecated_py_src/api/oauth/__init__.py
Normal file
14
_deprecated_py_src/api/oauth/__init__.py
Normal file
@@ -0,0 +1,14 @@
|
||||
"""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"]
|
||||
295
_deprecated_py_src/api/oauth/admin.py
Normal file
295
_deprecated_py_src/api/oauth/admin.py
Normal file
@@ -0,0 +1,295 @@
|
||||
"""OAuth 管理端点(管理员)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, 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 get_pipeline
|
||||
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 = get_pipeline()
|
||||
_OAUTH_ADMIN_LEGACY_DETAIL = "OAuth admin routes are retired; use Rust maintenance backend"
|
||||
|
||||
|
||||
def _raise_oauth_admin_legacy_unavailable() -> None:
|
||||
raise HTTPException(status_code=503, detail=_OAUTH_ADMIN_LEGACY_DETAIL)
|
||||
|
||||
|
||||
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: str | None = Field(None, max_length=2048)
|
||||
|
||||
authorization_url_override: str | None = Field(None, max_length=500)
|
||||
token_url_override: str | None = Field(None, max_length=500)
|
||||
userinfo_url_override: str | None = Field(None, max_length=500)
|
||||
scopes: list[str] | None = None
|
||||
|
||||
redirect_uri: str = Field(..., min_length=1, max_length=500)
|
||||
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
|
||||
|
||||
attribute_mapping: dict[str, Any] | None = None
|
||||
extra_config: dict[str, Any] | None = 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: str | None = None
|
||||
token_url_override: str | None = None
|
||||
userinfo_url_override: str | None = None
|
||||
scopes: list[str] | None = None
|
||||
redirect_uri: str
|
||||
frontend_callback_url: str
|
||||
attribute_mapping: dict[str, Any] | None = None
|
||||
extra_config: dict[str, Any] | None = 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: str | None = None
|
||||
authorization_url_override: str | None = None
|
||||
token_url_override: str | None = 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:
|
||||
_ = request, db
|
||||
_raise_oauth_admin_legacy_unavailable()
|
||||
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:
|
||||
_ = request, db
|
||||
_raise_oauth_admin_legacy_unavailable()
|
||||
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:
|
||||
_ = provider_type, request, db
|
||||
_raise_oauth_admin_legacy_unavailable()
|
||||
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:
|
||||
_ = provider_type, request, db
|
||||
_raise_oauth_admin_legacy_unavailable()
|
||||
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:
|
||||
_ = provider_type, request, db
|
||||
_raise_oauth_admin_legacy_unavailable()
|
||||
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:
|
||||
_ = provider_type, request, db
|
||||
_raise_oauth_admin_legacy_unavailable()
|
||||
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
_deprecated_py_src/api/oauth/public.py
Normal file
59
_deprecated_py_src/api/oauth/public.py
Normal file
@@ -0,0 +1,59 @@
|
||||
"""OAuth 公开端点(无需登录)。"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.responses import RedirectResponse
|
||||
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/oauth", tags=["OAuth"])
|
||||
_OAUTH_PUBLIC_LEGACY_DETAIL = "OAuth public routes are retired; use Rust maintenance backend"
|
||||
|
||||
|
||||
def _raise_oauth_public_legacy_unavailable() -> None:
|
||||
raise HTTPException(status_code=503, detail=_OAUTH_PUBLIC_LEGACY_DETAIL)
|
||||
|
||||
|
||||
@router.get("/providers")
|
||||
async def list_oauth_providers(db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
"""
|
||||
获取可用 OAuth Providers 列表。
|
||||
|
||||
模块未启用时返回空列表(前端友好)。
|
||||
"""
|
||||
_ = db
|
||||
_raise_oauth_public_legacy_unavailable()
|
||||
|
||||
|
||||
@router.get("/{provider_type}/authorize")
|
||||
async def oauth_authorize(
|
||||
provider_type: str,
|
||||
client_device_id: str = Query(..., min_length=1, max_length=128),
|
||||
db: Session = Depends(get_db),
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
发起 OAuth 登录(login flow)。
|
||||
"""
|
||||
_ = provider_type, client_device_id, db
|
||||
_raise_oauth_public_legacy_unavailable()
|
||||
|
||||
|
||||
@router.get("/{provider_type}/callback")
|
||||
async def oauth_callback(
|
||||
provider_type: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
code: str | None = Query(None),
|
||||
state: str | None = Query(None),
|
||||
error: str | None = Query(None),
|
||||
error_description: str | None = Query(None),
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
OAuth 回调端点。
|
||||
|
||||
成功/失败都会重定向到前端回调页。
|
||||
"""
|
||||
_ = provider_type, request, db, code, state, error, error_description
|
||||
_raise_oauth_public_legacy_unavailable()
|
||||
196
_deprecated_py_src/api/oauth/user.py
Normal file
196
_deprecated_py_src/api/oauth/user.py
Normal file
@@ -0,0 +1,196 @@
|
||||
"""OAuth 用户端点(需登录)。"""
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.responses import RedirectResponse
|
||||
|
||||
from src.api.base.adapter import ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.database import get_db
|
||||
from src.models.database import User
|
||||
from src.services.auth.oauth.service import OAuthService
|
||||
from src.services.auth.oauth.state import consume_oauth_bind_token, create_oauth_bind_token
|
||||
from src.services.auth.session_service import CLIENT_DEVICE_ID_HEADER, SessionService
|
||||
|
||||
router = APIRouter(prefix="/api/user/oauth", tags=["User - OAuth"])
|
||||
pipeline = get_pipeline()
|
||||
_OAUTH_USER_LEGACY_DETAIL = "OAuth user routes are retired; use Rust maintenance backend"
|
||||
|
||||
|
||||
def _raise_oauth_user_legacy_unavailable() -> None:
|
||||
raise HTTPException(status_code=503, detail=_OAUTH_USER_LEGACY_DETAIL)
|
||||
|
||||
|
||||
@router.get("/bindable-providers")
|
||||
async def list_bindable_providers(
|
||||
request: Request, db: Session = Depends(get_db)
|
||||
) -> dict[str, Any]:
|
||||
_ = request, db
|
||||
_raise_oauth_user_legacy_unavailable()
|
||||
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]:
|
||||
_ = request, db
|
||||
_raise_oauth_user_legacy_unavailable()
|
||||
adapter = ListMyOAuthLinksAdapter()
|
||||
result = await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
return cast(dict[str, Any], result)
|
||||
|
||||
|
||||
@router.post("/{provider_type}/bind-token")
|
||||
async def create_bind_token(
|
||||
provider_type: str, request: Request, db: Session = Depends(get_db)
|
||||
) -> dict[str, Any]:
|
||||
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景的安全认证"""
|
||||
_ = provider_type, request, db
|
||||
_raise_oauth_user_legacy_unavailable()
|
||||
adapter = CreateBindTokenAdapter(provider_type=provider_type)
|
||||
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),
|
||||
bind_token: str | None = None,
|
||||
) -> RedirectResponse:
|
||||
"""发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
|
||||
_ = provider_type, request, db, bind_token
|
||||
_raise_oauth_user_legacy_unavailable()
|
||||
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
|
||||
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]:
|
||||
_ = provider_type, request, db
|
||||
_raise_oauth_user_legacy_unavailable()
|
||||
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 CreateBindTokenAdapter(AuthenticatedApiAdapter):
|
||||
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景"""
|
||||
|
||||
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
|
||||
assert context.user.id is not None
|
||||
|
||||
# 验证 provider 是否存在且可绑定
|
||||
bindable = await OAuthService.list_bindable_providers(context.db, context.user)
|
||||
if not any(p["provider_type"] == self.provider_type for p in bindable):
|
||||
raise HTTPException(status_code=400, detail="无法绑定该 Provider")
|
||||
|
||||
redis = await get_redis_client(require_redis=True)
|
||||
if redis is None:
|
||||
raise HTTPException(status_code=503, detail="Redis 不可用")
|
||||
|
||||
token = await create_oauth_bind_token(
|
||||
redis, user_id=context.user.id, provider_type=self.provider_type
|
||||
)
|
||||
return {"bind_token": token}
|
||||
|
||||
|
||||
class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
"""发起 OAuth 绑定流程,支持两种认证方式:
|
||||
1. Authorization header (标准方式)
|
||||
2. bind_token 参数 (浏览器跳转场景)
|
||||
"""
|
||||
|
||||
def __init__(self, provider_type: str, bind_token: str | None = None):
|
||||
self.provider_type = provider_type
|
||||
self.bind_token = bind_token
|
||||
self._user_from_bind_token: User | None = None
|
||||
|
||||
@property
|
||||
def mode(self) -> ApiMode: # type: ignore[override]
|
||||
# 如果有 bind_token,使用 PUBLIC mode 跳过 header 认证
|
||||
if self.bind_token:
|
||||
return ApiMode.PUBLIC
|
||||
return ApiMode.USER
|
||||
|
||||
def authorize(self, context: ApiRequestContext) -> None:
|
||||
# 如果是 bind_token 模式,不在这里检查(会在 handle 中验证)
|
||||
if self.bind_token:
|
||||
return
|
||||
# 标准模式,检查用户
|
||||
if not context.user:
|
||||
raise HTTPException(status_code=401, detail="未登录")
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
|
||||
user: User | None = context.user
|
||||
|
||||
# 如果使用 bind_token,验证并获取用户
|
||||
if self.bind_token:
|
||||
redis = await get_redis_client(require_redis=True)
|
||||
if redis is None:
|
||||
raise HTTPException(status_code=503, detail="Redis 不可用")
|
||||
|
||||
token_data = await consume_oauth_bind_token(redis, self.bind_token)
|
||||
if not token_data:
|
||||
raise HTTPException(status_code=401, detail="无效或过期的绑定令牌")
|
||||
|
||||
# 验证 provider_type 匹配
|
||||
if token_data.provider_type != self.provider_type:
|
||||
raise HTTPException(status_code=400, detail="绑定令牌与 Provider 不匹配")
|
||||
|
||||
# 从数据库获取用户
|
||||
user = context.db.query(User).filter(User.id == token_data.user_id).first()
|
||||
if not user or not user.is_active or user.is_deleted:
|
||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
||||
|
||||
assert user is not None
|
||||
client_device_id: str | None = None
|
||||
if context.request.headers.get(CLIENT_DEVICE_ID_HEADER) or context.request.query_params.get(
|
||||
"client_device_id"
|
||||
):
|
||||
client_device_id = SessionService.extract_client_device_id(context.request)
|
||||
url = await OAuthService.build_bind_authorize_url(
|
||||
context.db,
|
||||
user,
|
||||
self.provider_type,
|
||||
client_device_id=client_device_id,
|
||||
)
|
||||
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