KillerKing93/Transformers-InferenceServer-OpenAPI
0
1#!/usr/bin/env python2# -*- coding: utf-8 -*-3"""4Authentication and authorization for AI Marketplace Platform5 6Features:7- JWT token generation and validation8- Password hashing with bcrypt9- Role-based access control (user, supplier, admin)10- Token refresh mechanism11- FastAPI dependencies for protected endpoints12 13Usage:14 from auth import get_current_user, create_access_token15 16 @app.post("/protected")17 def protected_route(current_user: dict = Depends(get_current_user)):18 return {"user": current_user}19"""20 21import os22from datetime import datetime, timedelta23from typing import Optional, Dict, Any24from passlib.context import CryptContext25from jose import JWTError, jwt26from fastapi import Depends, HTTPException, status27from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials28from pydantic import BaseModel, EmailStr29 30# JWT Configuration31SECRET_KEY = os.getenv("JWT_SECRET_KEY", "your-secret-key-change-in-production")32ALGORITHM = "HS256"33ACCESS_TOKEN_EXPIRE_MINUTES = int(os.getenv("ACCESS_TOKEN_EXPIRE_MINUTES", "30"))34REFRESH_TOKEN_EXPIRE_DAYS = int(os.getenv("REFRESH_TOKEN_EXPIRE_DAYS", "7"))35 36# Password hashing37pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")38 39# Security scheme40security = HTTPBearer()41 42 43# Pydantic models44class Token(BaseModel):45 access_token: str46 refresh_token: str47 token_type: str = "bearer"48 49 50class TokenData(BaseModel):51 email: Optional[str] = None52 role: Optional[str] = None53 54 55class UserLogin(BaseModel):56 email: EmailStr57 password: str58 59 60class UserRegisterAuth(BaseModel):61 email: EmailStr62 password: str63 name: str64 role: str = "user" # user, supplier, admin65 66 67# Password utilities68def hash_password(password: str) -> str:69 """Hash a password for storing."""70 return pwd_context.hash(password)71 72 73def verify_password(plain_password: str, hashed_password: str) -> bool:74 """Verify a stored password against one provided by user."""75 return pwd_context.verify(plain_password, hashed_password)76 77 78# Token utilities79def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str:80 """81 Create JWT access token.82 83 Args:84 data: Dictionary with user data (email, role, etc.)85 expires_delta: Optional token expiration time86 87 Returns:88 Encoded JWT token string89 """90 to_encode = data.copy()91 if expires_delta:92 expire = datetime.utcnow() + expires_delta93 else:94 expire = datetime.utcnow() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)95 96 to_encode.update({"exp": expire})97 encoded_jwt = jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)98 return encoded_jwt99 100 101def create_refresh_token(data: dict) -> str:102 """103 Create JWT refresh token with longer expiration.104 105 Args:106 data: Dictionary with user data107 108 Returns:109 Encoded JWT refresh token string110 """111 to_encode = data.copy()112 expire = datetime.utcnow() + timedelta(days=REFRESH_TOKEN_EXPIRE_DAYS)113 to_encode.update({"exp": expire, "type": "refresh"})114 encoded_jwt = jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)115 return encoded_jwt116 117 118def verify_token(token: str) -> Dict[str, Any]:119 """120 Verify and decode JWT token.121 122 Args:123 token: JWT token string124 125 Returns:126 Decoded token payload127 128 Raises:129 HTTPException: If token is invalid or expired130 """131 try:132 payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])133 email: str = payload.get("email")134 if email is None:135 raise HTTPException(136 status_code=status.HTTP_401_UNAUTHORIZED,137 detail="Could not validate credentials",138 headers={"WWW-Authenticate": "Bearer"},139 )140 return payload141 except JWTError:142 raise HTTPException(143 status_code=status.HTTP_401_UNAUTHORIZED,144 detail="Could not validate credentials",145 headers={"WWW-Authenticate": "Bearer"},146 )147 148 149# FastAPI dependencies150async def get_current_user(151 credentials: HTTPAuthorizationCredentials = Depends(security)152) -> Dict[str, Any]:153 """154 FastAPI dependency to get current authenticated user.155 156 Usage:157 @app.get("/protected")158 def protected_route(current_user: dict = Depends(get_current_user)):159 return {"user": current_user}160 161 Returns:162 User data from token payload163 164 Raises:165 HTTPException: If token is invalid166 """167 token = credentials.credentials168 payload = verify_token(token)169 return payload170 171 172async def get_current_active_user(173 current_user: dict = Depends(get_current_user)174) -> Dict[str, Any]:175 """176 Get current active user (additional checks can be added here).177 178 Args:179 current_user: User from token180 181 Returns:182 User data if active183 184 Raises:185 HTTPException: If user is inactive186 """187 # Add additional checks here (e.g., is_active flag from database)188 return current_user189 190 191async def require_role(required_role: str):192 """193 Dependency factory for role-based access control.194 195 Usage:196 @app.post("/admin/users", dependencies=[Depends(require_role("admin"))])197 def admin_only_route():198 return {"message": "Admin access"}199 200 Args:201 required_role: Required role (user, supplier, admin)202 203 Returns:204 Dependency function205 """206 async def role_checker(current_user: dict = Depends(get_current_user)):207 user_role = current_user.get("role", "user")208 if user_role != required_role and user_role != "admin":209 raise HTTPException(210 status_code=status.HTTP_403_FORBIDDEN,211 detail=f"Insufficient permissions. Required role: {required_role}"212 )213 return current_user214 215 return role_checker216 217 218# Helper for protected endpoints219def require_admin(current_user: dict = Depends(get_current_user)):220 """Require admin role for endpoint."""221 if current_user.get("role") != "admin":222 raise HTTPException(223 status_code=status.HTTP_403_FORBIDDEN,224 detail="Admin access required"225 )226 return current_user227 228 229def require_supplier(current_user: dict = Depends(get_current_user)):230 """Require supplier role for endpoint."""231 user_role = current_user.get("role")232 if user_role not in ["supplier", "admin"]:233 raise HTTPException(234 status_code=status.HTTP_403_FORBIDDEN,235 detail="Supplier access required"236 )237 return current_user238 239 240# Utility functions for user authentication241def authenticate_user(email: str, password: str, db_user: Dict[str, Any]) -> bool:242 """243 Authenticate user with email and password.244 245 Args:246 email: User email247 password: Plain password248 db_user: User data from database (must include 'hashed_password' key)249 250 Returns:251 True if authentication successful, False otherwise252 """253 if not db_user:254 return False255 if not verify_password(password, db_user.get("hashed_password", "")):256 return False257 return True258 259 260def create_tokens_for_user(email: str, role: str = "user", **extra_data) -> Token:261 """262 Create access and refresh tokens for user.263 264 Args:265 email: User email266 role: User role (user, supplier, admin)267 **extra_data: Additional data to include in token268 269 Returns:270 Token object with access_token, refresh_token, and token_type271 """272 token_data = {"email": email, "role": role, **extra_data}273 274 access_token = create_access_token(token_data)275 refresh_token = create_refresh_token(token_data)276 277 return Token(278 access_token=access_token,279 refresh_token=refresh_token,280 token_type="bearer"281 )282 283 284# Example: Refresh token endpoint logic285def refresh_access_token(refresh_token: str) -> str:286 """287 Generate new access token from refresh token.288 289 Args:290 refresh_token: Valid refresh token291 292 Returns:293 New access token294 295 Raises:296 HTTPException: If refresh token is invalid or not a refresh type297 """298 payload = verify_token(refresh_token)299 300 # Check if it's a refresh token301 if payload.get("type") != "refresh":302 raise HTTPException(303 status_code=status.HTTP_401_UNAUTHORIZED,304 detail="Invalid token type"305 )306 307 # Create new access token308 token_data = {309 "email": payload.get("email"),310 "role": payload.get("role")311 }312 313 return create_access_token(token_data)314 