129 lines
4.0 KiB
Python
129 lines
4.0 KiB
Python
"""Disk/Drive API routes for VoIdea."""
|
|
|
|
import logging
|
|
from typing import Annotated
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
|
from pydantic import BaseModel
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.dependencies import get_current_user, get_db
|
|
from app.models.user import User
|
|
from app.services.disk_service import DISK_PROVIDERS, DiskService
|
|
|
|
logger = logging.getLogger(__name__)
|
|
router = APIRouter()
|
|
|
|
|
|
class ConnectRequest(BaseModel):
|
|
provider: str
|
|
code: str
|
|
|
|
|
|
class DisconnectRequest(BaseModel):
|
|
provider: str
|
|
|
|
|
|
@router.get("/disk/providers")
|
|
async def list_disk_providers():
|
|
"""List available disk providers and user's connected ones."""
|
|
return {"providers": list(DISK_PROVIDERS)}
|
|
|
|
|
|
@router.get("/disk/status")
|
|
async def disk_status(
|
|
user: Annotated[User, Depends(get_current_user)],
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
"""Get user's connected disk providers."""
|
|
svc = DiskService(db)
|
|
providers = await svc.get_connected_providers(user)
|
|
return {"connected": providers, "count": len(providers)}
|
|
|
|
|
|
@router.post("/disk/connect")
|
|
async def connect_disk(
|
|
body: ConnectRequest,
|
|
user: Annotated[User, Depends(get_current_user)],
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
"""Exchange OAuth code for tokens and store provider connection."""
|
|
if body.provider not in DISK_PROVIDERS:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Unknown provider")
|
|
|
|
tokens = None
|
|
email = ""
|
|
|
|
if body.provider == "google":
|
|
from app.integrations.oauth.google import exchange_code, get_user_info
|
|
tokens = await exchange_code(body.code)
|
|
if tokens:
|
|
info = await get_user_info(tokens.get("access_token", ""))
|
|
if info:
|
|
email = info.email
|
|
|
|
elif body.provider == "yandex":
|
|
from app.integrations.oauth.yandex import exchange_code, get_user_info
|
|
token_result = await exchange_code(body.code)
|
|
if token_result:
|
|
tokens = {
|
|
"access_token": token_result.access_token,
|
|
"refresh_token": token_result.refresh_token,
|
|
"expires_at": (
|
|
token_result.expires_at.isoformat()
|
|
if token_result.expires_at else None
|
|
),
|
|
}
|
|
info = await get_user_info(token_result.access_token)
|
|
if info:
|
|
email = info.email
|
|
|
|
elif body.provider == "apple":
|
|
from app.integrations.oauth.apple import exchange_code
|
|
result = await exchange_code(body.code)
|
|
if result:
|
|
tokens = result
|
|
|
|
if not tokens:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail="Failed to exchange authorization code",
|
|
)
|
|
|
|
svc = DiskService(db)
|
|
await svc.add_provider(user, body.provider, tokens, email=email)
|
|
return {"success": True, "provider": body.provider}
|
|
|
|
|
|
@router.post("/disk/disconnect")
|
|
async def disconnect_disk(
|
|
body: DisconnectRequest,
|
|
user: Annotated[User, Depends(get_current_user)],
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
"""Disconnect a disk provider."""
|
|
svc = DiskService(db)
|
|
await svc.remove_provider(user, body.provider)
|
|
return {"success": True}
|
|
|
|
|
|
@router.get("/disk/oauth-url/{provider}")
|
|
async def disk_oauth_url(provider: str):
|
|
"""Get OAuth authorize URL for a disk provider."""
|
|
if provider not in DISK_PROVIDERS:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Unknown provider")
|
|
|
|
if provider == "google":
|
|
from app.integrations.oauth.google import get_authorize_url
|
|
url = await get_authorize_url()
|
|
elif provider == "yandex":
|
|
from app.integrations.oauth.yandex import get_authorize_url
|
|
url = await get_authorize_url()
|
|
elif provider == "apple":
|
|
from app.integrations.oauth.apple import get_authorize_url
|
|
url = await get_authorize_url()
|
|
else:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Unknown provider")
|
|
|
|
return {"url": url}
|