File size: 6,117 Bytes
ee7d7b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
import os
import bcrypt
from datetime import datetime, timedelta, timezone
from typing import Optional
from jose import jwt, JWTError
from fastapi import Request, HTTPException, Depends
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
import logging
from authlib.integrations.starlette_client import OAuth
from starlette.config import Config

logger = logging.getLogger(__name__)

# Setup Authlib OAuth — reads from .env file or OS environment variables (HuggingFace Spaces secrets)
oauth_config = Config(".env")
oauth = OAuth(oauth_config)

# Register Google OAuth (only works if GOOGLE_CLIENT_ID is set)
try:
    oauth.register(
        name='google',
        server_metadata_url='https://accounts.google.com/.well-known/openid-configuration',
        client_kwargs={
            'scope': 'openid email profile'
        }
    )
except Exception as e:
    logger.warning(f"Google OAuth registration failed: {e}")

# Register GitHub OAuth (only works if GITHUB_CLIENT_ID is set)
try:
    oauth.register(
        name='github',
        api_base_url='https://api.github.com/',
        access_token_url='https://github.com/login/oauth/access_token',
        authorize_url='https://github.com/login/oauth/authorize',
        client_kwargs={
            'scope': 'user:email'
        }
    )
except Exception as e:
    logger.warning(f"GitHub OAuth registration failed: {e}")

SECRET_KEY = os.environ.get("JWT_SECRET") or os.environ.get("JWT_SECRET_KEY") or "datavision-production-jwt-secret-key-32bytes-long!"
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 60 * 24 * 7 # 7 days for ease of use

security = HTTPBearer(auto_error=False)

class AuthenticatedUser:
    def __init__(self, user_id: str, email: Optional[str] = None, role: str = "authenticated", is_guest: bool = False):
        self.id = user_id
        self.user_id = user_id
        self.email = email
        self.role = role
        self.is_guest = is_guest

    def __repr__(self):
        return f"User(id={self.id}, email={self.email}, role={self.role})"

def verify_password(plain_password: str, hashed_password: str) -> bool:
    try:
        return bcrypt.checkpw(plain_password.encode('utf-8'), hashed_password.encode('utf-8'))
    except Exception:
        return False

def get_password_hash(password: str) -> str:
    salt = bcrypt.gensalt()
    return bcrypt.hashpw(password.encode('utf-8'), salt).decode('utf-8')

def create_access_token(data: dict, expires_delta: Optional[timedelta] = None):
    to_encode = data.copy()
    if expires_delta:
        expire = datetime.now(timezone.utc) + expires_delta
    else:
        expire = datetime.now(timezone.utc) + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
    to_encode.update({"exp": expire})
    encoded_jwt = jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
    return encoded_jwt

def decode_jwt_token(token: str) -> dict:
    """Decode and validate a JWT token. Standard HS256 verification."""
    try:
        return jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
    except JWTError as e:
        raise HTTPException(status_code=401, detail=f"Invalid token: {str(e)}")


def decode_jwt(token: str) -> dict:
    return decode_jwt_token(token)

def extract_token_from_request(request: Request) -> Optional[str]:
    auth_header = request.headers.get("Authorization", "")
    if auth_header.startswith("Bearer "):
        return auth_header[7:]
    return request.cookies.get("dv-access-token") or request.query_params.get("token")

def generate_guest_id(request: Request) -> str:
    import hashlib
    ip = request.client.host if request.client else "unknown"
    ua = request.headers.get("User-Agent", "unknown")[:100]
    fingerprint = hashlib.sha256(f"{ip}:{ua}".encode()).hexdigest()[:12]
    return f"guest_{fingerprint}"

async def get_current_user(
    request: Request,
    credentials: Optional[HTTPAuthorizationCredentials] = Depends(security)
) -> AuthenticatedUser:
    token = credentials.credentials if credentials else extract_token_from_request(request)
    
    if token and token not in ["null", "undefined", ""]:
        try:
            payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
            user_id = payload.get("sub")
            if not user_id:
                raise HTTPException(status_code=401, detail="Token missing user ID")
            
            return AuthenticatedUser(
                user_id=user_id,
                email=payload.get("email"),
                role=payload.get("role", "authenticated"),
                is_guest=False
            )
        except JWTError as e:
            logger.warning(f"Token validation failed: {e}")
    
    guest_id = generate_guest_id(request)
    return AuthenticatedUser(user_id=guest_id, email=None, role="anon", is_guest=True)

async def get_current_user_optional(request: Request, credentials: Optional[HTTPAuthorizationCredentials] = Depends(security)):
    try:
        return await get_current_user(request, credentials)
    except HTTPException:
        return None

async def require_authenticated_user(user: AuthenticatedUser = Depends(get_current_user)):
    if user.is_guest:
        raise HTTPException(status_code=401, detail="Authentication required")
    return user

async def get_user_id_from_token(user: AuthenticatedUser = Depends(get_current_user)) -> str:
    return user.id

def get_user_id_from_body_deprecated(body_user_id: Optional[str], user: AuthenticatedUser) -> str:
    if not user.is_guest:
        return user.id
    if body_user_id and body_user_id not in ["default", "guest", "", "null", "undefined"]:
        return body_user_id
    return user.id

async def get_admin_user(user: AuthenticatedUser = Depends(require_authenticated_user)):
    if user.role not in ["admin", "super_admin"]:
        raise HTTPException(status_code=403, detail="Admin access required")
    return user

async def get_super_admin_user(user: AuthenticatedUser = Depends(require_authenticated_user)):
    if user.role != "super_admin":
        raise HTTPException(status_code=403, detail="Super admin access required")
    return user