124 lines
3.9 KiB
Python
124 lines
3.9 KiB
Python
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.dependencies import get_current_user, get_db
|
|
from app.models.user import User
|
|
from app.schemas.user_oauth import RequestOAuthRequest, RequestOAuthResponse, UserOAuthOut
|
|
from app.services.user_oauth_service import (
|
|
build_oauth_url,
|
|
get_account_info_by_type,
|
|
get_token_by_type,
|
|
save_oauth_token,
|
|
)
|
|
|
|
router = APIRouter(prefix="/user-oauth", tags=["user-oauth"])
|
|
|
|
|
|
@router.post(
|
|
"/request_oauth",
|
|
summary="获取授权链接",
|
|
description="用户提交oauth_type,返回对应的第三方授权链接",
|
|
response_model=RequestOAuthResponse,
|
|
)
|
|
async def request_oauth(
|
|
req: RequestOAuthRequest,
|
|
current_user: User = Depends(get_current_user),
|
|
):
|
|
try:
|
|
auth_url = await build_oauth_url(req.oauth_type, current_user.id)
|
|
return {"auth_url": auth_url}
|
|
except ValueError as e:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=str(e),
|
|
)
|
|
|
|
|
|
@router.get(
|
|
"/juliang_callback",
|
|
summary="巨量授权回调",
|
|
description="巨量引擎授权回调地址,接收code和state参数,获取token并保存",
|
|
)
|
|
async def juliang_callback(
|
|
auth_code: str = Query(..., description="第三方返回的授权码"),
|
|
state: str = Query(..., description="请求时传递的自定义参数"),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
try:
|
|
parts = state.split(":")
|
|
if len(parts) != 4:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail="无效的state参数",
|
|
)
|
|
|
|
oauth_type = int(parts[0])
|
|
user_id = parts[1]
|
|
app_id = parts[2]
|
|
app_type = parts[3]
|
|
|
|
token = await get_token_by_type(auth_code, oauth_type, app_id, app_type)
|
|
account_info = await get_account_info_by_type(token, oauth_type, app_type)
|
|
user_oauth = await save_oauth_token(db, user_id, oauth_type, token, account_info, app_id)
|
|
|
|
return {
|
|
"message": "授权成功",
|
|
"code": 0,
|
|
"data": UserOAuthOut.model_validate(user_oauth),
|
|
}
|
|
|
|
except ValueError as e:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=str(e),
|
|
)
|
|
except Exception as e:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
detail=f"授权失败: {str(e)}",
|
|
)
|
|
|
|
|
|
@router.get(
|
|
"/callback",
|
|
summary="通用授权回调",
|
|
description="其他平台授权回调地址,接收code和state参数",
|
|
)
|
|
async def oauth_callback(
|
|
code: str = Query(..., description="第三方返回的授权码"),
|
|
state: str = Query(..., description="请求时传递的自定义参数"),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
try:
|
|
parts = state.split(":")
|
|
if len(parts) != 4:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail="无效的state参数",
|
|
)
|
|
|
|
oauth_type = int(parts[0])
|
|
user_id = parts[1]
|
|
app_id = parts[2]
|
|
app_type = parts[3]
|
|
|
|
token = await get_token_by_type(code, oauth_type, app_id, app_type)
|
|
account_info = await get_account_info_by_type(token, oauth_type, app_type)
|
|
user_oauth = await save_oauth_token(db, user_id, oauth_type, token, account_info, app_id)
|
|
|
|
return {
|
|
"message": "授权成功",
|
|
"code": 0,
|
|
"data": UserOAuthOut.model_validate(user_oauth),
|
|
}
|
|
|
|
except ValueError as e:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=str(e),
|
|
)
|
|
except Exception as e:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
detail=f"授权失败: {str(e)}",
|
|
) |