Files

311 lines
10 KiB
Python

"""
OAuth2 Authentication module for BDFR Web Interface
This module handles OAuth2 authentication flow for the web interface,
integrating with BDFR's existing OAuth2 system.
"""
import asyncio
import json
import logging
import secrets
import time
from datetime import datetime, timedelta
from typing import Dict, Optional, Any
from urllib.parse import urlencode
import httpx
from fastapi import HTTPException, status
# Try to import BDFR modules, but handle gracefully if not available
try:
from bdfr.oauth2 import OAuth2Authenticator, OAuth2TokenManager
from bdfr.exceptions import RedditAuthenticationError
BDFR_AVAILABLE = True
except ImportError:
BDFR_AVAILABLE = False
# Create mock classes for when BDFR is not available
class OAuth2Authenticator:
pass
class OAuth2TokenManager:
pass
class RedditAuthenticationError(Exception):
pass
logger = logging.getLogger(__name__)
class WebOAuth2Manager:
"""OAuth2 manager for web interface authentication"""
def __init__(self, client_id: str, client_secret: str, scopes: list = None):
self.client_id = client_id
self.client_secret = client_secret
self.scopes = scopes or ["identity", "history", "read", "save", "mysubreddits"]
# In-memory storage for OAuth2 states and tokens
# In production, this should be replaced with a proper database
self.oauth_states = {}
self.refresh_tokens = {}
self.access_tokens = {}
# Store Reddit usernames per session state
self.usernames = {}
# Reddit OAuth2 endpoints
self.reddit_auth_url = "https://www.reddit.com/api/v1/authorize"
self.reddit_token_url = "https://www.reddit.com/api/v1/access_token"
self.reddit_user_info_url = "https://oauth.reddit.com/api/v1/me"
# Token expiration tracking
self.token_expiry = {}
def generate_state(self) -> str:
"""Generate a secure random state for OAuth2"""
state = secrets.token_urlsafe(32)
self.oauth_states[state] = {
"created_at": time.time(),
"used": False
}
return state
def validate_state(self, state: str) -> bool:
"""Validate OAuth2 state parameter"""
if state not in self.oauth_states:
return False
state_data = self.oauth_states[state]
if state_data["used"]:
return False
# States expire after 10 minutes
if time.time() - state_data["created_at"] > 600:
del self.oauth_states[state]
return False
return True
def mark_state_used(self, state: str):
"""Mark OAuth2 state as used"""
if state in self.oauth_states:
self.oauth_states[state]["used"] = True
def get_authorization_url(self, redirect_uri: str) -> Dict[str, str]:
"""Generate OAuth2 authorization URL"""
state = self.generate_state()
params = {
"client_id": self.client_id,
"response_type": "code",
"state": state,
"redirect_uri": redirect_uri,
"scope": " ".join(self.scopes),
"duration": "permanent"
}
auth_url = f"{self.reddit_auth_url}?{urlencode(params)}"
return {
"authorization_url": auth_url,
"state": state
}
async def exchange_code_for_token(self, code: str, state: str, redirect_uri: str) -> Dict[str, Any]:
"""Exchange authorization code for access token"""
if not self.validate_state(state):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or expired state parameter"
)
self.mark_state_used(state)
data = {
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri
}
headers = {
"User-Agent": "BDFR-Web-Interface/1.0"
}
# Use HTTP Basic Auth for client credentials
auth = (self.client_id, self.client_secret)
async with httpx.AsyncClient() as client:
try:
response = await client.post(
self.reddit_token_url,
data=data,
auth=auth,
headers=headers,
timeout=30.0
)
if response.status_code != 200:
error_detail = response.json() if response.headers.get("content-type", "").startswith("application/json") else response.text
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Token exchange failed: {error_detail}"
)
token_data = response.json()
# Store tokens
access_token = token_data["access_token"]
refresh_token = token_data.get("refresh_token")
if refresh_token:
self.refresh_tokens[state] = refresh_token
self.access_tokens[state] = access_token
# Set expiry (Reddit tokens typically last 1 hour)
self.token_expiry[state] = time.time() + token_data.get("expires_in", 3600)
# Attempt to fetch and store the Reddit username for this session
username = None
try:
user_info = await self.get_user_info(access_token)
username = user_info.get("name")
except Exception as e:
logger.warning(f"Failed to fetch user info during token exchange: {e}")
if username:
self.usernames[state] = username
return {
"access_token": access_token,
"refresh_token": refresh_token,
"expires_in": token_data.get("expires_in", 3600),
"token_type": token_data.get("token_type", "bearer"),
"state": state,
"username": username
}
else:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="No refresh token received"
)
except httpx.TimeoutException:
raise HTTPException(
status_code=status.HTTP_408_REQUEST_TIMEOUT,
detail="Token exchange timed out"
)
except Exception as e:
logger.error(f"Token exchange error: {str(e)}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Internal server error during token exchange"
)
async def get_user_info(self, access_token: str) -> Dict[str, Any]:
"""Get user information using access token"""
headers = {
"Authorization": f"Bearer {access_token}",
"User-Agent": "BDFR-Web-Interface/1.0"
}
async with httpx.AsyncClient() as client:
try:
response = await client.get(
self.reddit_user_info_url,
headers=headers,
timeout=30.0
)
if response.status_code != 200:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid access token"
)
return response.json()
except httpx.TimeoutException:
raise HTTPException(
status_code=status.HTTP_408_REQUEST_TIMEOUT,
detail="User info request timed out"
)
except Exception as e:
logger.error(f"User info error: {str(e)}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Error retrieving user information"
)
def is_token_expired(self, state: str) -> bool:
"""Check if access token is expired"""
if state not in self.token_expiry:
return True
return time.time() > self.token_expiry[state]
def get_valid_token(self, state: str) -> Optional[str]:
"""Get valid access token, refreshing if necessary"""
if state not in self.access_tokens:
return None
if self.is_token_expired(state):
# Token expired, would need refresh logic here
# For now, just return None to indicate re-auth needed
return None
return self.access_tokens[state]
def revoke_session(self, state: str):
"""Revoke OAuth2 session"""
if state in self.oauth_states:
del self.oauth_states[state]
if state in self.refresh_tokens:
del self.refresh_tokens[state]
if state in self.access_tokens:
del self.access_tokens[state]
if state in self.token_expiry:
del self.token_expiry[state]
if state in self.usernames:
del self.usernames[state]
def get_auth_status(self, state: str = None) -> Dict[str, Any]:
"""Get authentication status"""
if not state:
return {
"authenticated": False,
"message": "No active session"
}
if state not in self.access_tokens:
return {
"authenticated": False,
"message": "No tokens found for session"
}
access_token = self.get_valid_token(state)
if not access_token:
return {
"authenticated": False,
"message": "Token expired or invalid"
}
return {
"authenticated": True,
"expires_at": self.token_expiry.get(state, 0),
"scopes": self.scopes,
"username": self.usernames.get(state)
}
# Global OAuth2 manager instance
oauth_manager = None
def init_oauth_manager(client_id: str, client_secret: str, scopes: list = None):
"""Initialize the global OAuth2 manager"""
global oauth_manager
oauth_manager = WebOAuth2Manager(client_id, client_secret, scopes)
def get_oauth_manager() -> WebOAuth2Manager:
"""Get the global OAuth2 manager instance"""
if oauth_manager is None:
raise RuntimeError("OAuth2 manager not initialized")
return oauth_manager