from fastapi import APIRouter, Depends, HTTPException, Query, status from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.dependencies import get_current_user, get_db from app.models.user import User from app.models.user_oauth import UserOAuth from app.models.user_oauth_app import UserOAuthApp from app.schemas.user_oauth import RequestOAuthRequest, RequestOAuthResponse, UserOAuthOut, OAuthListResponse from app.services.user_oauth_service import ( build_oauth_url, get_token, get_oauth_list, ) router = APIRouter(prefix="/user-oauth", tags=["oauth"]) @router.post( "/request_oauth", summary="获取授权链接", description="用户提交开户方式open_type,返回对应的第三方授权链接", response_model=RequestOAuthResponse, ) async def request_oauth( req: RequestOAuthRequest, current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): try: if req.open_type not in [1,2,3,4,5,6,7,8,9,10]: raise ValueError("open_type must be in [1,2,3,4,5,6,7,8,9,10]") auth_url = await build_oauth_url(req.open_type, current_user.id, db) return {"auth_url": auth_url} except ValueError as e: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, 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) != 2 or not parts[0] or not parts[1]: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail="无效的state参数", ) user_id = parts[0] app_id = parts[1] user = await db.execute( select(User).where(User.id == user_id).limit(1) ) user = user.scalar_one_or_none() if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在", ) app = await db.execute( select(UserOAuthApp).where( UserOAuthApp.app_id == app_id, UserOAuthApp.status == 1, UserOAuthApp.deleted_at.is_(None), ).limit(1) ) app = app.scalar_one_or_none() if not app: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail="应用不存在", ) await get_token(auth_code, user_id, app_id, db) return { "message": "授权成功,这里需要跳转页面路径到 /user-oauth/oauth_list", "code": 0, } except ValueError as e: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=str(e), ) except HTTPException: raise except Exception as e: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"授权失败: {str(e)}", ) @router.get( "/oauth_list", summary="获取账户下所有授权列表", description="获取当前用户下所有授权账户列表,支持按授权登录账号、开户方式、授权账户id筛选", response_model=OAuthListResponse, ) async def oauth_list( account_userid: str | None = Query(None, description="授权登录账号id"), open_type: int | None = Query(None, description="开户方式open_type"), account_id: str | None = Query(None, description="授权账户id"), page: int = Query(1, description="页码"), page_size: int = Query(10, description="每页数量"), current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): try: result = await get_oauth_list( user_id=current_user.id, db=db, account_userid=account_userid, open_type=open_type, account_id=account_id, page=page, page_size=page_size, ) return { "code": 0, "message": "查询成功", "data": [ { "id": oauth.id, "account_id": oauth.account_id, "account_name": oauth.account_name, "account_role": oauth.account_role, "account_username": oauth.account_username, "account_userid": oauth.account_userid, "user_id": oauth.user_id, "open_type": oauth.open_type, "port_type": oauth.port_type, "appid": oauth.appid, "material_auth_status": oauth.material_auth_status, "created_at": oauth.created_at, "updated_at": oauth.updated_at, } for oauth in result["data"] ], "pagination": { "page": result["page"], "page_size": result["page_size"], "total": result["total"], "total_pages": (result["total"] + result["page_size"] - 1) // result["page_size"], }, } 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=str(e), )