Files
kite/backend/api/cloud.py
T

659 lines
22 KiB
Python

"""Cloud storage connection management endpoints."""
from __future__ import annotations
import asyncio
import secrets
import uuid
import urllib.parse
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.responses import JSONResponse, RedirectResponse
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from api.schemas import CloudConnectionOut
from config import settings
from db.models import CloudConnection, User
from deps.auth import get_regular_user
from deps.db import get_db
from deps.utils import get_client_ip
from services.audit import write_audit_log
from services.rate_limiting import account_limiter
from storage.cloud_backend_factory import build_cloud_backend
from storage.cloud_utils import decrypt_credentials, encrypt_credentials, validate_cloud_url
router = APIRouter(prefix="/api/cloud", tags=["cloud"])
users_router = APIRouter(prefix="/api/users", tags=["users"])
VALID_OAUTH_PROVIDERS = {"google_drive", "onedrive"}
VALID_WEBDAV_PROVIDERS = {"nextcloud", "webdav"}
VALID_CLOUD_PROVIDERS = VALID_OAUTH_PROVIDERS | VALID_WEBDAV_PROVIDERS
_VALID_BACKENDS = frozenset({"minio", *VALID_CLOUD_PROVIDERS})
_DISPLAY_NAMES = {
"google_drive": "Google Drive",
"onedrive": "OneDrive",
"nextcloud": "Nextcloud",
"webdav": "WebDAV server",
}
class WebDAVConnectRequest(BaseModel):
server_url: str
username: str
password: str
provider: str # "nextcloud" or "webdav"
class DefaultStorageRequest(BaseModel):
backend: str
def _master_key() -> bytes:
return settings.cloud_creds_key.encode()
def _oauth_redirect_uri(provider: str) -> str:
return f"{settings.backend_url}/api/cloud/oauth/callback/{provider}"
def _oauth_state_key(state_token: str) -> str:
return f"oauth_state:{state_token}"
def _settings_redirect(**params: str) -> RedirectResponse:
query = urllib.parse.urlencode(params)
return RedirectResponse(url=f"{settings.frontend_url}/settings?{query}", status_code=302)
def _unsupported_provider_error(provider: str, valid_providers: set[str]) -> HTTPException:
return HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported provider: {provider}. Valid providers: {sorted(valid_providers)}",
)
def _ensure_oauth_configured(provider: str) -> None:
if provider == "google_drive" and (not settings.google_client_id or not settings.google_client_secret):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Google Drive OAuth is not configured on this server. Set GOOGLE_CLIENT_ID and GOOGLE_CLIENT_SECRET in your environment.",
)
if provider == "onedrive" and (
not settings.onedrive_client_id
or not settings.onedrive_client_secret
or not settings.onedrive_tenant_id
):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="OneDrive OAuth is not configured on this server. Set ONEDRIVE_CLIENT_ID, ONEDRIVE_CLIENT_SECRET, and ONEDRIVE_TENANT_ID in your environment.",
)
def _webdav_credentials(body: WebDAVConnectRequest) -> dict[str, str]:
return {
"server_url": body.server_url,
"username": body.username,
"password": body.password,
}
async def _upsert_cloud_connection(
session: AsyncSession,
user_id: uuid.UUID,
provider: str,
credentials_enc: str,
) -> CloudConnection:
result = await session.execute(
select(CloudConnection).where(
CloudConnection.user_id == user_id,
CloudConnection.provider == provider,
)
)
conn = result.scalar_one_or_none()
if conn is not None:
conn.credentials_enc = credentials_enc
conn.status = "ACTIVE"
else:
conn = CloudConnection(
id=uuid.uuid4(),
user_id=user_id,
provider=provider,
display_name=_DISPLAY_NAMES.get(provider, provider),
credentials_enc=credentials_enc,
status="ACTIVE",
)
session.add(conn)
return conn
async def _get_owned_connection(
session: AsyncSession,
connection_id: uuid.UUID,
user_id: uuid.UUID,
) -> CloudConnection:
conn = await session.get(CloudConnection, connection_id)
if conn is None or conn.user_id != user_id:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Connection not found")
return conn
async def _get_active_connection(
session: AsyncSession,
user_id: uuid.UUID,
provider: str,
) -> CloudConnection:
result = await session.execute(
select(CloudConnection).where(
CloudConnection.user_id == user_id,
CloudConnection.provider == provider,
CloudConnection.status == "ACTIVE",
)
)
conn = result.scalar_one_or_none()
if conn is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="No active cloud connection found for this provider",
)
return conn
def _decrypt_connection(conn: CloudConnection, user_id: uuid.UUID) -> dict:
return decrypt_credentials(_master_key(), str(user_id), conn.credentials_enc)
async def _oauth_authorization_url(provider: str, redirect_uri: str, state_token: str) -> str:
if provider == "google_drive":
from google_auth_oauthlib.flow import Flow # lazy import
flow = Flow.from_client_config(
{
"web": {
"client_id": settings.google_client_id,
"client_secret": settings.google_client_secret,
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
}
},
scopes=["https://www.googleapis.com/auth/drive.file"],
redirect_uri=redirect_uri,
)
authorization_url, _ = flow.authorization_url(
access_type="offline",
prompt="consent",
state=state_token,
)
return authorization_url
if provider == "onedrive":
import msal # lazy import
app = msal.ConfidentialClientApplication(
settings.onedrive_client_id,
client_credential=settings.onedrive_client_secret,
authority=f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}",
)
return await asyncio.to_thread(
app.get_authorization_request_url,
scopes=["Files.ReadWrite", "offline_access"],
redirect_uri=redirect_uri,
state=state_token,
)
raise ValueError(f"Unsupported OAuth provider: {provider}")
async def _oauth_credentials_from_code(provider: str, redirect_uri: str, code: Optional[str]) -> dict:
if provider == "google_drive":
from google_auth_oauthlib.flow import Flow # lazy import
flow = Flow.from_client_config(
{
"web": {
"client_id": settings.google_client_id,
"client_secret": settings.google_client_secret,
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
"token_uri": "https://oauth2.googleapis.com/token",
}
},
scopes=["https://www.googleapis.com/auth/drive.file"],
redirect_uri=redirect_uri,
)
await asyncio.to_thread(flow.fetch_token, code=code)
token_credentials = flow.credentials
return {
"access_token": token_credentials.token,
"refresh_token": token_credentials.refresh_token,
"token_uri": token_credentials.token_uri,
"client_id": token_credentials.client_id,
"client_secret": token_credentials.client_secret,
"expiry": token_credentials.expiry.isoformat() if token_credentials.expiry else None,
}
if provider == "onedrive":
import msal # lazy import
app = msal.ConfidentialClientApplication(
settings.onedrive_client_id,
client_credential=settings.onedrive_client_secret,
authority=f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}",
)
result = await asyncio.to_thread(
app.acquire_token_by_authorization_code,
code,
scopes=["Files.ReadWrite", "offline_access"],
redirect_uri=redirect_uri,
)
if "error" in result:
raise ValueError(f"Token exchange failed: {result.get('error_description', result['error'])}")
return {
"access_token": result["access_token"],
"refresh_token": result.get("refresh_token", ""),
"token_uri": f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}/oauth2/v2.0/token",
"client_id": settings.onedrive_client_id,
"client_secret": settings.onedrive_client_secret,
}
raise ValueError(f"Unsupported OAuth provider: {provider}")
def _connection_response(conn: CloudConnection) -> dict:
data = CloudConnectionOut.model_validate(conn).model_dump()
if conn.provider not in VALID_WEBDAV_PROVIDERS:
return data
try:
credentials = _decrypt_connection(conn, conn.user_id)
except Exception:
return data
data["server_url"] = credentials.get("server_url")
data["connection_username"] = credentials.get("username")
return data
async def _fetch_google_drive_folders(credentials: dict, folder_id: str) -> list:
from google.oauth2.credentials import Credentials # lazy import
from googleapiclient.discovery import build # lazy import
import datetime as dt
expiry = None
if expiry_str := credentials.get("expiry"):
try:
expiry = dt.datetime.fromisoformat(expiry_str)
except ValueError:
pass
creds = Credentials(
token=credentials.get("access_token"),
refresh_token=credentials.get("refresh_token"),
token_uri=credentials.get("token_uri", "https://oauth2.googleapis.com/token"),
client_id=credentials.get("client_id"),
client_secret=credentials.get("client_secret"),
expiry=expiry,
)
def _list_files() -> list:
service = build("drive", "v3", credentials=creds, cache_discovery=False)
response = service.files().list(
q=f"'{folder_id}' in parents and trashed=false",
fields="files(id,name,mimeType,size)",
pageSize=200,
).execute()
items = []
for item in response.get("files", []):
is_dir = item.get("mimeType") == "application/vnd.google-apps.folder"
items.append({
"id": item["id"],
"name": item["name"],
"is_dir": is_dir,
"size": int(item.get("size", 0)) if not is_dir else 0,
})
return items
return await asyncio.to_thread(_list_files)
async def _fetch_onedrive_folders(credentials: dict, folder_id: str) -> list:
import httpx # lazy import
if folder_id in ("root", ""):
url = "https://graph.microsoft.com/v1.0/me/drive/root/children"
else:
url = f"https://graph.microsoft.com/v1.0/me/drive/items/{folder_id}/children"
async with httpx.AsyncClient() as client:
resp = await client.get(
url,
headers={"Authorization": f"Bearer {credentials.get('access_token', '')}"},
timeout=30,
)
resp.raise_for_status()
data = resp.json()
items = []
for item in data.get("value", []):
is_dir = "folder" in item
items.append({
"id": item["id"],
"name": item["name"],
"is_dir": is_dir,
"size": item.get("size", 0) if not is_dir else 0,
})
return items
async def _fetch_webdav_folders(provider: str, credentials: dict, folder_id: str) -> list:
webdav_path = "" if folder_id == "root" else folder_id
backend = build_cloud_backend(provider, credentials)
return await backend.list_folder(webdav_path)
@router.get("/oauth/initiate/{provider}")
@account_limiter.limit("100/minute")
async def oauth_initiate(
provider: str,
request: Request,
current_user: User = Depends(get_regular_user),
) -> dict:
"""Start an OAuth flow and return the provider authorization URL."""
request.state.current_user = current_user
if provider not in VALID_OAUTH_PROVIDERS:
raise _unsupported_provider_error(provider, VALID_OAUTH_PROVIDERS)
_ensure_oauth_configured(provider)
state_token = secrets.token_urlsafe(32)
await request.app.state.redis.setex(_oauth_state_key(state_token), 1800, str(current_user.id))
authorization_url = await _oauth_authorization_url(
provider,
redirect_uri=_oauth_redirect_uri(provider),
state_token=state_token,
)
return JSONResponse({"url": authorization_url})
@router.get("/oauth/callback/{provider}", response_class=RedirectResponse)
async def oauth_callback(
provider: str,
request: Request,
session: AsyncSession = Depends(get_db),
):
"""Exchange OAuth code, store encrypted credentials, and redirect to settings."""
state = request.query_params.get("state")
code = request.query_params.get("code")
error_param = request.query_params.get("error")
try:
if error_param:
raise ValueError(f"OAuth provider returned error: {error_param}")
if provider not in VALID_OAUTH_PROVIDERS:
raise ValueError(f"Unsupported OAuth provider: {provider}")
if not state:
raise ValueError("Missing OAuth state parameter")
redis_client = request.app.state.redis
state_key = _oauth_state_key(state)
stored_user_id = await redis_client.get(state_key)
if not stored_user_id:
return _settings_redirect(
cloud_error="Invalid or expired OAuth state. Please try connecting again."
)
await redis_client.delete(state_key)
if isinstance(stored_user_id, bytes):
stored_user_id = stored_user_id.decode("utf-8")
user_id = uuid.UUID(stored_user_id)
user = await session.get(User, user_id)
if user is None or not user.is_active:
raise ValueError("User not found or inactive")
credentials = await _oauth_credentials_from_code(provider, _oauth_redirect_uri(provider), code)
credentials_enc = encrypt_credentials(_master_key(), str(user_id), credentials)
conn = await _upsert_cloud_connection(session, user_id, provider, credentials_enc)
await session.flush()
await write_audit_log(
session,
event_type="cloud.connected",
user_id=user.id,
actor_id=user.id,
resource_id=conn.id,
ip_address=None,
metadata_={"provider": provider},
)
await session.commit()
return _settings_redirect(cloud_connected=provider)
except Exception as exc:
return _settings_redirect(cloud_error=str(exc))
@router.post("/connections/webdav", status_code=status.HTTP_201_CREATED)
@account_limiter.limit("100/minute")
async def connect_webdav(
body: WebDAVConnectRequest,
request: Request,
session: AsyncSession = Depends(get_db),
current_user: User = Depends(get_regular_user),
) -> dict:
"""Connect a WebDAV or Nextcloud server."""
request.state.current_user = current_user
if body.provider not in VALID_WEBDAV_PROVIDERS:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Unsupported WebDAV provider: {body.provider}. Valid values: {sorted(VALID_WEBDAV_PROVIDERS)}",
)
try:
validate_cloud_url(body.server_url)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Invalid server URL: {exc}",
) from exc
try:
credentials = _webdav_credentials(body)
backend = build_cloud_backend(body.provider, credentials)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Invalid server URL: {exc}",
) from exc
try:
ok = await backend.health_check()
if not ok:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="Connection test failed — check server URL and credentials",
)
except HTTPException:
raise
except Exception as exc:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Connection test failed — check server URL and credentials: {exc}",
) from exc
credentials_enc = encrypt_credentials(_master_key(), str(current_user.id), credentials)
conn = await _upsert_cloud_connection(session, current_user.id, body.provider, credentials_enc)
await session.flush()
_ip = get_client_ip(request)
await write_audit_log(
session,
event_type="cloud.connected",
user_id=current_user.id,
actor_id=current_user.id,
resource_id=conn.id,
ip_address=_ip,
metadata_={"provider": body.provider},
)
await session.commit()
await session.refresh(conn)
return CloudConnectionOut.model_validate(conn).model_dump()
@router.get("/connections")
@account_limiter.limit("100/minute")
async def list_connections(
request: Request,
session: AsyncSession = Depends(get_db),
current_user: User = Depends(get_regular_user),
) -> dict:
"""List the current user's cloud connections."""
request.state.current_user = current_user
result = await session.execute(
select(CloudConnection).where(CloudConnection.user_id == current_user.id)
)
return {"items": [_connection_response(conn) for conn in result.scalars().all()]}
@router.get("/connections/{connection_id}/config")
@account_limiter.limit("100/minute")
async def get_connection_config(
request: Request,
connection_id: uuid.UUID,
session: AsyncSession = Depends(get_db),
current_user: User = Depends(get_regular_user),
) -> dict:
"""Return non-secret WebDAV/Nextcloud connection fields."""
request.state.current_user = current_user
conn = await _get_owned_connection(session, connection_id, current_user.id)
if conn.provider not in VALID_WEBDAV_PROVIDERS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Connection config is only available for WebDAV/Nextcloud connections",
)
try:
credentials = _decrypt_connection(conn, current_user.id)
except Exception:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Failed to decrypt connection credentials",
)
return {
"id": str(conn.id),
"provider": conn.provider,
"server_url": credentials.get("server_url", ""),
"connection_username": credentials.get("username", ""),
}
@router.delete("/connections/{connection_id}", status_code=status.HTTP_204_NO_CONTENT)
@account_limiter.limit("100/minute")
async def delete_connection(
connection_id: uuid.UUID,
request: Request,
session: AsyncSession = Depends(get_db),
current_user: User = Depends(get_regular_user),
) -> None:
"""Disconnect a cloud connection."""
request.state.current_user = current_user
conn = await _get_owned_connection(session, connection_id, current_user.id)
provider = conn.provider
from services.cloud_cache import invalidate_provider_cache # lazy import
invalidate_provider_cache(str(current_user.id), provider)
_ip = get_client_ip(request)
await write_audit_log(
session,
event_type="cloud.disconnected",
user_id=current_user.id,
actor_id=current_user.id,
resource_id=conn.id,
ip_address=_ip,
metadata_={"provider": provider},
)
await session.delete(conn)
await session.commit()
@router.get("/folders/{provider}/{folder_id:path}")
@account_limiter.limit("100/minute")
async def list_cloud_folders(
request: Request,
provider: str,
folder_id: str,
session: AsyncSession = Depends(get_db),
current_user: User = Depends(get_regular_user),
) -> dict:
"""List folder contents for a connected cloud provider."""
request.state.current_user = current_user
if provider not in VALID_CLOUD_PROVIDERS:
raise _unsupported_provider_error(provider, VALID_CLOUD_PROVIDERS)
conn = await _get_active_connection(session, current_user.id, provider)
credentials = _decrypt_connection(conn, current_user.id)
from services.cloud_cache import get_cloud_folders_cached # lazy import
if provider == "google_drive":
fetcher = lambda: _fetch_google_drive_folders(credentials, folder_id)
elif provider == "onedrive":
fetcher = lambda: _fetch_onedrive_folders(credentials, folder_id)
else:
fetcher = lambda: _fetch_webdav_folders(provider, credentials, folder_id)
items = await get_cloud_folders_cached(str(current_user.id), provider, folder_id, fetcher)
return {"items": items}
@users_router.patch("/me/default-storage")
@account_limiter.limit("100/minute")
async def update_default_storage(
request: Request,
body: DefaultStorageRequest,
session: AsyncSession = Depends(get_db),
current_user: User = Depends(get_regular_user),
) -> dict:
"""Update the current user's default storage backend.
The backend value is validated against the allowlist before storage.
Returns the updated default_storage_backend value.
"""
if body.backend not in _VALID_BACKENDS:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=f"Invalid backend. Valid values: {sorted(_VALID_BACKENDS)}",
)
request.state.current_user = current_user
user = await session.get(User, current_user.id)
if user is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
user.default_storage_backend = body.backend
session.add(user)
await session.commit()
return {"default_storage_backend": user.default_storage_backend}