fix: security fix

This commit is contained in:
Furqan-14
2026-02-16 15:32:05 +05:30
parent 9b41301b09
commit c3017b81da
19 changed files with 382 additions and 253 deletions
+1 -1
View File
@@ -37,7 +37,7 @@ EMAIL_FROM=info@maskantech.in
# JWT Configuration # JWT Configuration
ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg" ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg"
ACCESS_TOKEN_EXPIRES=86400 ACCESS_TOKEN_EXPIRES=900
REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ" REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ"
REFRESH_TOKEN_EXPIRES=864000 REFRESH_TOKEN_EXPIRES=864000
JWT_ALGORITHM=HS256 JWT_ALGORITHM=HS256
+1 -1
View File
@@ -37,7 +37,7 @@ EMAIL_FROM=info@maskantech.in
# JWT Configuration # JWT Configuration
ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg" ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg"
ACCESS_TOKEN_EXPIRES=86400 ACCESS_TOKEN_EXPIRES=900
REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ" REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ"
REFRESH_TOKEN_EXPIRES=864000 REFRESH_TOKEN_EXPIRES=864000
JWT_ALGORITHM=HS256 JWT_ALGORITHM=HS256
+1 -1
View File
@@ -38,7 +38,7 @@ EMAIL_FROM=info@maskantech.in
# JWT Configuration # JWT Configuration
ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg" ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg"
ACCESS_TOKEN_EXPIRES=86400 ACCESS_TOKEN_EXPIRES=900
REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ" REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ"
REFRESH_TOKEN_EXPIRES=864000 REFRESH_TOKEN_EXPIRES=864000
JWT_ALGORITHM=HS256 JWT_ALGORITHM=HS256
+1 -1
View File
@@ -37,7 +37,7 @@ EMAIL_FROM=info@maskantech.in
# JWT Configuration # JWT Configuration
ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg" ACCESS_TOKEN_SECRET="L_ByN0_FIuwsQnDo4sdrOEdJvqlPjKfhVJmqhf76D13v3IWu3mbvzb8hQRnPxHMlr9Y8A9IcOHZZWSs7Kfofpg"
ACCESS_TOKEN_EXPIRES=86400 ACCESS_TOKEN_EXPIRES=900
REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ" REFRESH_TOKEN_SECRET="6Z0yOfkhPjfLH77WTnBh3Iv0JU_gWIfIqpGuGU41GFRV4fnLZMfKN3gAsPTfsqFKv1rRc6szUJRngW8Py0UYUQ"
REFRESH_TOKEN_EXPIRES=864000 REFRESH_TOKEN_EXPIRES=864000
JWT_ALGORITHM=HS256 JWT_ALGORITHM=HS256
+70 -66
View File
@@ -1,3 +1,4 @@
from contextlib import asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.security import HTTPBearer from fastapi.security import HTTPBearer
@@ -31,6 +32,72 @@ logging.basicConfig(
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@asynccontextmanager
async def lifespan(app: FastAPI):
logger.info("Testing database connection...")
try:
with engine.connect() as conn:
conn.execute(text("SELECT 1"))
conn.commit()
logger.info("Database connection successful!")
except Exception as e:
logger.error(f"Database connection failed: {e}")
raise
await redis_client.connect()
await run_in_threadpool(sync_redis_client.connect)
logger.info(
f"{settings.PROJECT_NAME} v{settings.VERSION} started ({settings.APP_ENV})"
)
async def redis_event_consumer():
logger.info("Redis Event Consumer STARTED")
while True:
try:
if not redis_client.client:
await asyncio.sleep(5)
continue
result = await redis_client.client.blpop("saas:events:queue", timeout=5)
if result:
_, event_id = result
try:
with SessionLocal() as db:
await run_in_threadpool(EventService.process_queue_item, db, event_id)
except Exception as e:
logger.error(f"Error processing event {event_id}: {e}")
except Exception as e:
await asyncio.sleep(1)
async def fallback_poller():
logger.info("Fallback Event Poller STARTED")
while True:
try:
with SessionLocal() as db:
await run_in_threadpool(EventService.process_outbox, db)
except Exception as e:
logger.error(f"Fallback poller error: {e}")
await asyncio.sleep(60)
background_tasks = [
asyncio.create_task(redis_event_consumer()),
asyncio.create_task(fallback_poller()),
]
yield
logger.info("Shutting down...")
for task in background_tasks:
task.cancel()
await asyncio.gather(*background_tasks, return_exceptions=True)
await redis_client.close()
await run_in_threadpool(sync_redis_client.close)
def create_app() -> FastAPI: def create_app() -> FastAPI:
app = FastAPI( app = FastAPI(
title=settings.PROJECT_NAME, title=settings.PROJECT_NAME,
@@ -39,6 +106,7 @@ def create_app() -> FastAPI:
docs_url="/docs", docs_url="/docs",
redoc_url="/redoc", redoc_url="/redoc",
openapi_url="/openapi.json", openapi_url="/openapi.json",
lifespan=lifespan,
) )
# === OpenAPI Security Scheme === # === OpenAPI Security Scheme ===
@@ -120,69 +188,6 @@ def create_app() -> FastAPI:
app.include_router(admin_module_env_router, prefix="/api/admin/modules", tags=["Admin - Module Environments"]) app.include_router(admin_module_env_router, prefix="/api/admin/modules", tags=["Admin - Module Environments"])
app.include_router(admin_tenant_modules_router, prefix="/api/admin/tenants", tags=["Admin - Tenant Modules"]) app.include_router(admin_tenant_modules_router, prefix="/api/admin/tenants", tags=["Admin - Tenant Modules"])
# === Startup: Test DB Connection (Sync + SQLAlchemy 2.0 compatible) ===
@app.on_event("startup")
async def startup_event():
logger.info("Testing database connection...")
try:
with engine.connect() as conn:
conn.execute(text("SELECT 1"))
conn.commit()
logger.info("Database connection successful!")
except Exception as e:
logger.error(f"Database connection failed: {e}")
raise
from app.core.redis import redis_client
await redis_client.connect()
logger.info(
f"{settings.PROJECT_NAME} v{settings.VERSION} started ({settings.APP_ENV})"
)
async def redis_event_consumer():
logger.info("Redis Event Consumer STARTED")
while True:
try:
if not redis_client.client:
await asyncio.sleep(5)
continue
result = await redis_client.client.blpop("saas:events:queue", timeout=5)
if result:
_, event_id = result
try:
with SessionLocal() as db:
await run_in_threadpool(EventService.process_queue_item, db, event_id)
except Exception as e:
logger.error(f"Error processing event {event_id}: {e}")
except Exception as e:
await asyncio.sleep(1)
async def fallback_poller():
logger.info("Fallback Event Poller STARTED")
while True:
try:
with SessionLocal() as db:
await run_in_threadpool(EventService.process_outbox, db)
except Exception as e:
logger.error(f"Fallback poller error: {e}")
await asyncio.sleep(60)
asyncio.create_task(redis_event_consumer())
asyncio.create_task(fallback_poller())
@app.on_event("shutdown")
async def shutdown_event():
logger.info("Shutting down...")
await redis_client.close()
# === Basic Routes === # === Basic Routes ===
@app.get("/", tags=["Root"]) @app.get("/", tags=["Root"])
def root(): def root():
@@ -212,7 +217,8 @@ def create_app() -> FastAPI:
conn.execute(text("SELECT 1")) conn.execute(text("SELECT 1"))
db_status = "healthy" db_status = "healthy"
except Exception as e: except Exception as e:
db_status = f"unhealthy: {str(e)}" logger.error(f"Health check DB error: {e}")
db_status = "unhealthy"
return { return {
"status": "healthy" if db_status == "healthy" else "degraded", "status": "healthy" if db_status == "healthy" else "degraded",
@@ -222,5 +228,3 @@ def create_app() -> FastAPI:
} }
return app return app
app = create_app()
+20 -2
View File
@@ -2,11 +2,16 @@ from datetime import datetime, timedelta, timezone
from typing import Optional, Dict, Any from typing import Optional, Dict, Any
import bcrypt import bcrypt
import jwt import jwt
import uuid
import logging
from fastapi import HTTPException, status from fastapi import HTTPException, status
import re import re
import secrets import secrets
import string import string
from app.config.settings import settings from app.config.settings import settings
from app.core.redis import sync_redis_client
logger = logging.getLogger(__name__)
class SecurityUtils: class SecurityUtils:
"""Security utility class for authentication and authorization.""" """Security utility class for authentication and authorization."""
@@ -27,7 +32,7 @@ class SecurityUtils:
"""Generate JWT access token.""" """Generate JWT access token."""
to_encode = data.copy() to_encode = data.copy()
expire = datetime.now(timezone.utc) + timedelta(seconds=settings.ACCESS_TOKEN_EXPIRES) expire = datetime.now(timezone.utc) + timedelta(seconds=settings.ACCESS_TOKEN_EXPIRES)
to_encode.update({"exp": expire, "type": "access"}) to_encode.update({"exp": expire, "type": "access", "jti": str(uuid.uuid4())})
# Include tenant_id if provided # Include tenant_id if provided
if tenant_id: if tenant_id:
@@ -44,7 +49,7 @@ class SecurityUtils:
"""Generate JWT refresh token.""" """Generate JWT refresh token."""
to_encode = data.copy() to_encode = data.copy()
expire = datetime.now(timezone.utc) + timedelta(seconds=settings.REFRESH_TOKEN_EXPIRES) expire = datetime.now(timezone.utc) + timedelta(seconds=settings.REFRESH_TOKEN_EXPIRES)
to_encode.update({"exp": expire, "type": "refresh"}) to_encode.update({"exp": expire, "type": "refresh", "jti": str(uuid.uuid4())})
# Include tenant_id if provided # Include tenant_id if provided
if tenant_id: if tenant_id:
@@ -72,6 +77,19 @@ class SecurityUtils:
detail="Invalid token type" detail="Invalid token type"
) )
jti = payload.get("jti")
if jti and sync_redis_client.client:
try:
if sync_redis_client.client.get(f"blacklist:{jti}"):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Token has been revoked"
)
except HTTPException:
raise
except Exception as e:
logger.warning(f"Redis blacklist check failed: {e}")
return payload return payload
except jwt.ExpiredSignatureError: except jwt.ExpiredSignatureError:
raise HTTPException( raise HTTPException(
+1 -1
View File
@@ -68,7 +68,7 @@ class Settings(BaseSettings):
# JWT settings # JWT settings
ACCESS_TOKEN_SECRET: str ACCESS_TOKEN_SECRET: str
ACCESS_TOKEN_EXPIRES: int = 86400 ACCESS_TOKEN_EXPIRES: int = 900
REFRESH_TOKEN_SECRET: str REFRESH_TOKEN_SECRET: str
REFRESH_TOKEN_EXPIRES: int = 864000 REFRESH_TOKEN_EXPIRES: int = 864000
JWT_ALGORITHM: str = "HS256" JWT_ALGORITHM: str = "HS256"
+2 -2
View File
@@ -40,8 +40,8 @@ class AuthController:
) )
@staticmethod @staticmethod
def logout(current_user: User): def logout(current_user: User, token: str):
return AuthService.logout(current_user) return AuthService.logout(current_user, token)
@staticmethod @staticmethod
def me(db: Session, current_user: User): def me(db: Session, current_user: User):
+1 -1
View File
@@ -115,7 +115,7 @@ class SyncRedisClient:
if not self._redis: if not self._redis:
try: try:
self.connect() self.connect()
except: except Exception:
pass pass
if not self._redis: if not self._redis:
return 0 return 0
+1 -3
View File
@@ -5,6 +5,7 @@ from typing import List
from app.config.database import get_db from app.config.database import get_db
from app.config.security import security from app.config.security import security
from app.models.auth.user_model import User from app.models.auth.user_model import User
from app.models.auth.access_model import Access
security_scheme = HTTPBearer(auto_error=False) security_scheme = HTTPBearer(auto_error=False)
@@ -74,9 +75,6 @@ def can_access(user: User, access_code: str, db: Session) -> bool:
if access_code in user_access_codes: if access_code in user_access_codes:
return True return True
from app.models.auth.access_models import Access
requested_access = db.query(Access).filter( requested_access = db.query(Access).filter(
Access.access_code == access_code Access.access_code == access_code
).first() ).first()
+13 -16
View File
@@ -1,4 +1,5 @@
from fastapi import APIRouter, Depends, status from fastapi import APIRouter, Depends, Request, status, HTTPException
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.config.database import get_db from app.config.database import get_db
from app.controllers.auth.auth_controller import AuthController from app.controllers.auth.auth_controller import AuthController
@@ -14,7 +15,8 @@ from app.schemas.auth.auth_schema import (
VerifyOTPRequest, VerifyOTPRequest,
ResetPasswordWithOTP, ResetPasswordWithOTP,
) )
from app.middleware.auth_middleware import get_current_user from app.schemas.auth.language_schema import UpdateLanguageRequest
from app.middleware.auth_middleware import get_current_user, security_scheme
from app.middleware.tenant_middleware import get_tenant_from_header from app.middleware.tenant_middleware import get_tenant_from_header
from app.models.auth.user_model import User from app.models.auth.user_model import User
import uuid import uuid
@@ -56,21 +58,11 @@ def update_user(
@router.patch("/update/{user_id}/language", response_model=UserResponse) @router.patch("/update/{user_id}/language", response_model=UserResponse)
def update_language( def update_language(
user_id: uuid.UUID, user_id: uuid.UUID,
language_data: dict, language_data: UpdateLanguageRequest,
db: Session = Depends(get_db), db: Session = Depends(get_db),
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
): ):
supported_languages = ["en", "ar"] user_update = UserUpdate(preferred_language=language_data.preferred_language.value)
preferred_language = language_data.get("preferred_language")
if preferred_language not in supported_languages:
from fastapi import HTTPException
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported language. Supported languages: {', '.join(supported_languages)}"
)
user_update = UserUpdate(preferred_language=preferred_language)
return AuthController.update_user(db, user_id, user_update, current_user) return AuthController.update_user(db, user_id, user_update, current_user)
@router.post("/reset-password") @router.post("/reset-password")
@@ -82,8 +74,13 @@ def reset_password(
return AuthController.reset_password(db, current_user, password_data) return AuthController.reset_password(db, current_user, password_data)
@router.post("/logout") @router.post("/logout")
def logout(current_user: User = Depends(get_current_user)): def logout(
return AuthController.logout(current_user) request: Request,
credentials: HTTPAuthorizationCredentials = Depends(security_scheme),
current_user: User = Depends(get_current_user),
):
token = credentials.credentials if credentials else request.cookies.get("access_token")
return AuthController.logout(current_user, token)
@router.post("/forgot-password") @router.post("/forgot-password")
def forgot_password(request: ForgotPasswordRequest, db: Session = Depends(get_db)): def forgot_password(request: ForgotPasswordRequest, db: Session = Depends(get_db)):
+2 -2
View File
@@ -67,10 +67,10 @@ class ForgotPasswordRequest(BaseModel):
class VerifyOTPRequest(BaseModel): class VerifyOTPRequest(BaseModel):
email: EmailStr email: EmailStr
otp: str otp: str = Field(..., min_length=6, max_length=6, pattern=r'^\d{6}$')
class ResetPasswordWithOTP(BaseModel): class ResetPasswordWithOTP(BaseModel):
email: EmailStr email: EmailStr
otp: str otp: str = Field(..., min_length=6, max_length=6, pattern=r'^\d{6}$')
new_password: str = Field(..., min_length=8) new_password: str = Field(..., min_length=8)
+7 -1
View File
@@ -1,4 +1,10 @@
from pydantic import BaseModel from pydantic import BaseModel
from enum import Enum
class LanguageEnum(str, Enum):
"""Supported languages for the application."""
EN = "en"
AR = "ar"
class UpdateLanguageRequest(BaseModel): class UpdateLanguageRequest(BaseModel):
preferred_language: str preferred_language: LanguageEnum
+7 -6
View File
@@ -1,13 +1,15 @@
from sqlalchemy.orm import Session from sqlalchemy.orm import Session, joinedload
from app.models.auth.access_model import Access from app.models.auth.access_model import Access
from typing import List, Any from typing import List, Any
from app.core.redis import sync_redis_client from app.core.redis import sync_redis_client
import json import json
import logging
from app.models.auth.module_access_model import ModuleAccess from app.models.auth.module_access_model import ModuleAccess
from app.models.auth.module_model import Module from app.models.auth.module_model import Module
from sqlalchemy.orm import joinedload
from datetime import datetime from datetime import datetime
logger = logging.getLogger(__name__)
class AccessService: class AccessService:
@staticmethod @staticmethod
@@ -35,7 +37,7 @@ class AccessService:
return deserialized_list return deserialized_list
except Exception as e: except Exception as e:
pass logger.warning(f"Access cache read error: {e}")
query = db.query(Access) query = db.query(Access)
@@ -73,10 +75,9 @@ class AccessService:
"created_at": item.created_at.isoformat() if item.created_at else None, "created_at": item.created_at.isoformat() if item.created_at else None,
}) })
sync_redis_client.client.set(cache_key, json.dumps(serialized), ex=3600) # 1 hour cache sync_redis_client.client.set(cache_key, json.dumps(serialized), ex=3600)
except Exception as e: except Exception as e:
pass logger.warning(f"Access cache write error: {e}")
return result return result
@staticmethod @staticmethod
+168 -70
View File
@@ -3,10 +3,16 @@ from fastapi import HTTPException, status
from app.models.auth.user_model import User from app.models.auth.user_model import User
from app.schemas.auth.auth_schema import UserSignup, UserSignin, UserUpdate from app.schemas.auth.auth_schema import UserSignup, UserSignin, UserUpdate
from app.config.security import security from app.config.security import security
from datetime import datetime, timedelta from datetime import datetime, timedelta, timezone
import random import secrets
import uuid import uuid
import logging
import jwt
from app.config.settings import settings
from app.services.auth.email_service import EmailService from app.services.auth.email_service import EmailService
from app.core.redis import sync_redis_client
logger = logging.getLogger(__name__)
class AuthService: class AuthService:
@@ -159,7 +165,17 @@ class AuthService:
return {"message": "Password updated successfully"} return {"message": "Password updated successfully"}
@staticmethod @staticmethod
def logout(current_user: User): def logout(current_user: User, token: str):
try:
payload = jwt.decode(token, settings.ACCESS_TOKEN_SECRET, algorithms=["HS256"])
jti = payload.get("jti")
exp = payload.get("exp")
if jti and exp and sync_redis_client.client:
remaining_ttl = int(exp - datetime.now(timezone.utc).timestamp())
if remaining_ttl > 0:
sync_redis_client.client.setex(f"blacklist:{jti}", remaining_ttl, "1")
except Exception as e:
logger.warning(f"Failed to blacklist token on logout: {e}")
return {"message": "Logged out successfully"} return {"message": "Logged out successfully"}
@staticmethod @staticmethod
@@ -196,92 +212,174 @@ class AuthService:
@staticmethod @staticmethod
def forgot_password(db: Session, email: str): def forgot_password(db: Session, email: str):
print(f"DEBUG: Processing forgot_password for email: {email}") """Generate and send OTP for password reset via Redis storage."""
user = db.query(User).filter(User.email == email).first() user = db.query(User).filter(User.email == email).first()
if not user: if not user:
print(f"DEBUG: User not found for email: {email}")
return {"message": "If the email is registered, an OTP has been sent."} return {"message": "If the email is registered, an OTP has been sent."}
print(f"DEBUG: User found: {user.id}") otp_code = "".join([secrets.choice("0123456789") for _ in range(6)])
otp_code = "".join([str(random.randint(0, 9)) for _ in range(6)]) redis_key = f"otp:{email}"
expires_at = datetime.utcnow() + timedelta(minutes=10) try:
print(f"DEBUG: Generated OTP: {otp_code}, Expires: {expires_at}") if sync_redis_client.client:
sync_redis_client.client.setex(redis_key, 600, otp_code) # 600s = 10 minutes
db.query(PasswordResetOTP).filter( logger.info(f"OTP generated for email: {email}")
PasswordResetOTP.email == email, PasswordResetOTP.is_used == False else:
).update({"is_used": True}) logger.error("Redis client unavailable for OTP storage")
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Password reset service temporarily unavailable"
)
except Exception as e:
logger.error(f"Redis error storing OTP: {e}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to generate OTP"
)
try: try:
otp_entry = PasswordResetOTP( EmailService.send_otp(email, otp_code)
email=email, otp=otp_code, expires_at=expires_at
)
db.add(otp_entry)
db.commit()
print("DEBUG: OTP stored in database successfully")
except Exception as e: except Exception as e:
print(f"DEBUG: Database error saving OTP: {e}") logger.error(f"Failed to send OTP email to {email}: {e}")
db.rollback()
raise e
# Send Email
print("DEBUG: Attempting to send email...")
email_sent = EmailService.send_otp(email, otp_code)
print(f"DEBUG: Email sending result: {email_sent}")
return {"message": "If the email is registered, an OTP has been sent."} return {"message": "If the email is registered, an OTP has been sent."}
@staticmethod
def _check_otp_attempts(email: str):
"""Check if OTP attempts are exceeded. Raises 429 if locked out."""
attempts_key = f"otp_attempts:{email}"
try:
attempts = sync_redis_client.client.get(attempts_key)
if attempts and int(attempts) >= 5:
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Too many attempts. Try again in 15 minutes."
)
except HTTPException:
raise
except Exception as e:
logger.warning(f"Redis error checking OTP attempts: {e}")
@staticmethod
def _increment_otp_attempts(email: str):
"""Increment failed OTP attempt counter with 15-minute TTL."""
attempts_key = f"otp_attempts:{email}"
try:
pipe = sync_redis_client.client.pipeline()
pipe.incr(attempts_key)
pipe.expire(attempts_key, 900)
pipe.execute()
except Exception as e:
logger.warning(f"Redis error incrementing OTP attempts: {e}")
@staticmethod @staticmethod
def verify_otp(db: Session, email: str, otp: str): def verify_otp(db: Session, email: str, otp: str):
otp_entry = ( """Verify OTP from Redis storage using constant-time comparison."""
db.query(PasswordResetOTP) redis_key = f"otp:{email}"
.filter(
PasswordResetOTP.email == email,
PasswordResetOTP.otp == otp,
PasswordResetOTP.is_used == False,
PasswordResetOTP.expires_at > datetime.utcnow(),
)
.first()
)
if not otp_entry: try:
if not sync_redis_client.client:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Verification service temporarily unavailable"
)
AuthService._check_otp_attempts(email)
stored_otp = sync_redis_client.client.get(redis_key)
if not stored_otp:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or expired OTP"
)
if not secrets.compare_digest(otp, stored_otp):
AuthService._increment_otp_attempts(email)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or expired OTP"
)
sync_redis_client.client.delete(redis_key)
sync_redis_client.client.delete(f"otp_attempts:{email}")
sync_redis_client.client.setex(f"otp_verified:{email}", 300, "1")
logger.info(f"OTP verified successfully for email: {email}")
return {"message": "OTP verified successfully"}
except HTTPException:
raise
except Exception as e:
logger.error(f"Redis error verifying OTP: {e}")
raise HTTPException( raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid or expired OTP" status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to verify OTP"
) )
return {"message": "OTP verified successfully"}
@staticmethod @staticmethod
def reset_password_with_otp(db: Session, email: str, otp: str, new_password: str): def reset_password_with_otp(db: Session, email: str, otp: str, new_password: str):
otp_entry = ( """Reset password after verifying OTP and delete OTP from Redis."""
db.query(PasswordResetOTP) redis_key = f"otp:{email}"
.filter( verified_key = f"otp_verified:{email}"
PasswordResetOTP.email == email,
PasswordResetOTP.otp == otp,
PasswordResetOTP.is_used == False,
PasswordResetOTP.expires_at > datetime.utcnow(),
)
.first()
)
if not otp_entry: try:
if not sync_redis_client.client:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Password reset service temporarily unavailable"
)
AuthService._check_otp_attempts(email)
is_pre_verified = sync_redis_client.client.get(verified_key)
if not is_pre_verified:
stored_otp = sync_redis_client.client.get(redis_key)
if not stored_otp:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or expired OTP"
)
if not secrets.compare_digest(otp, stored_otp):
AuthService._increment_otp_attempts(email)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid or expired OTP"
)
user = db.query(User).filter(User.email == email).first()
if not user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="User not found"
)
if not security.validate_password_strength(new_password):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Password too weak"
)
user.password = security.hash_password(new_password)
user.password_updated_at = datetime.now(timezone.utc)
db.commit()
sync_redis_client.client.delete(redis_key)
sync_redis_client.client.delete(verified_key)
sync_redis_client.client.delete(f"otp_attempts:{email}")
logger.info(f"Password reset successfully for email: {email}")
return {"message": "Password updated successfully"}
except HTTPException:
raise
except Exception as e:
logger.error(f"Error resetting password: {e}")
db.rollback()
raise HTTPException( raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid or expired OTP" status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to reset password"
) )
user = db.query(User).filter(User.email == email).first()
if not user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
)
if not security.validate_password_strength(new_password):
raise HTTPException(status_code=400, detail="Password too weak")
user.password = security.hash_password(new_password)
user.password_updated_at = datetime.utcnow()
otp_entry.is_used = True
db.commit()
return {"message": "Password updated successfully"}
+10 -16
View File
@@ -1,41 +1,35 @@
import smtplib import smtplib
from email.mime.text import MIMEText from email.mime.text import MIMEText
from email.mime.multipart import MIMEMultipart from email.mime.multipart import MIMEMultipart
import os import logging
from dotenv import load_dotenv from app.config.settings import settings
load_dotenv() logger = logging.getLogger(__name__)
class EmailService: class EmailService:
SMTP_HOST = os.getenv("SMTP_HOST")
SMTP_PORT = int(os.getenv("SMTP_PORT", 465))
SMTP_USER = os.getenv("SMTP_USER")
SMTP_PASSWORD = os.getenv("SMTP_PASSWORD")
EMAIL_FROM = os.getenv("EMAIL_FROM")
SMTP_SECURE = os.getenv("SMTP_SECURE", "true").lower() == "true"
@staticmethod @staticmethod
def send_otp(to_email: str, otp: str): def send_otp(to_email: str, otp: str):
try: try:
msg = MIMEMultipart() msg = MIMEMultipart()
msg['From'] = EmailService.EMAIL_FROM msg['From'] = settings.EMAIL_FROM
msg['To'] = to_email msg['To'] = to_email
msg['Subject'] = "Password Reset OTP" msg['Subject'] = "Password Reset OTP"
body = f"Your OTP for password reset is: {otp}. It expires in 10 minutes." body = f"Your OTP for password reset is: {otp}. It expires in 10 minutes."
msg.attach(MIMEText(body, 'plain')) msg.attach(MIMEText(body, 'plain'))
if EmailService.SMTP_SECURE: if settings.SMTP_SECURE:
server = smtplib.SMTP_SSL(EmailService.SMTP_HOST, EmailService.SMTP_PORT) server = smtplib.SMTP_SSL(settings.SMTP_HOST, settings.SMTP_PORT)
else: else:
server = smtplib.SMTP(EmailService.SMTP_HOST, EmailService.SMTP_PORT) server = smtplib.SMTP(settings.SMTP_HOST, settings.SMTP_PORT)
server.starttls() server.starttls()
server.login(EmailService.SMTP_USER, EmailService.SMTP_PASSWORD) server.login(settings.SMTP_USER, settings.SMTP_PASSWORD)
text = msg.as_string() text = msg.as_string()
server.sendmail(EmailService.EMAIL_FROM, to_email, text) server.sendmail(settings.EMAIL_FROM, to_email, text)
server.quit() server.quit()
return True return True
except Exception as e: except Exception as e:
print(f"Failed to send email: {e}") logger.error(f"Failed to send email: {e}")
return False return False
+3 -3
View File
@@ -1,5 +1,5 @@
import uuid import uuid
import requests import httpx
import json import json
import logging import logging
from datetime import datetime, timezone, timedelta from datetime import datetime, timezone, timedelta
@@ -163,7 +163,7 @@ class EventService:
} }
logger.info(f"Sending event {log.event_type} to {log.target_url}") logger.info(f"Sending event {log.event_type} to {log.target_url}")
response = requests.post(log.target_url, data=payload_json, headers=headers, timeout=5) response = httpx.post(log.target_url, content=payload_json, headers=headers, timeout=5)
if response.status_code in range(200, 300): if response.status_code in range(200, 300):
log.status = EventStatus.COMPLETED log.status = EventStatus.COMPLETED
@@ -233,7 +233,7 @@ class EventService:
} }
logger.info(f"Sending event {log.event_type} to {log.target_url}. Payload: {payload_json}") logger.info(f"Sending event {log.event_type} to {log.target_url}. Payload: {payload_json}")
response = requests.post(log.target_url, data=payload_json, headers=headers, timeout=5) response = httpx.post(log.target_url, content=payload_json, headers=headers, timeout=5)
if response.status_code in range(200, 300): if response.status_code in range(200, 300):
log.status = EventStatus.COMPLETED log.status = EventStatus.COMPLETED
+48 -36
View File
@@ -1,4 +1,5 @@
import uuid import uuid
import logging
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import Dict, Any, Optional from typing import Dict, Any, Optional
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -10,9 +11,12 @@ from app.models.auth.tenant_module_model import TenantModule
from app.models.auth.user_model import User from app.models.auth.user_model import User
from app.config.security import security from app.config.security import security
from app.services.auth.trust_service import TrustService from app.services.auth.trust_service import TrustService
from app.core.redis import sync_redis_client
import json import json
import time import time
logger = logging.getLogger(__name__)
class SSOService: class SSOService:
@staticmethod @staticmethod
def generate_grant( def generate_grant(
@@ -62,24 +66,26 @@ class SSOService:
grant_code = str(uuid.uuid4().hex) grant_code = str(uuid.uuid4().hex)
expires_at = datetime.now(timezone.utc) + timedelta(seconds=60)
redirect_url = f"{env.frontend_base_url}{env.sso_entry_path}?grant={grant_code}" redirect_url = f"{env.frontend_base_url}{env.sso_entry_path}?grant={grant_code}"
grant = SSOGrant( grant_data = json.dumps({
grant_code=grant_code, "user_id": str(user_id),
user_id=user_id, "module_id": str(module.id),
module_id=module.id, "tenant_id": str(tenant_id) if tenant_id else None,
tenant_id=tenant_id, "environment_slug": env.slug,
environment_slug=env.slug, "created_at": datetime.now(timezone.utc).isoformat(),
expires_at=expires_at })
)
db.add(grant) if not sync_redis_client.client:
db.commit() raise HTTPException(
db.refresh(grant) status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="SSO service temporarily unavailable"
)
sync_redis_client.client.setex(f"sso_grant:{grant_code}", 60, grant_data)
return { return {
"grant_code": grant.grant_code, "grant_code": grant_code,
"redirect_url": redirect_url "redirect_url": redirect_url
} }
@@ -176,38 +182,44 @@ class SSOService:
""" """
Validates grant and returns a short-lived module-scoped token. Validates grant and returns a short-lived module-scoped token.
This is called by the Module Backend. This is called by the Module Backend.
Grants are stored in Redis — atomically deleted on exchange (one-time use).
""" """
grant = db.query(SSOGrant).filter(SSOGrant.grant_code == grant_code).first() if not sync_redis_client.client:
if not grant: raise HTTPException(
raise HTTPException(status_code=401, detail="Invalid grant code") status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="SSO service temporarily unavailable"
)
if grant.is_used: redis_key = f"sso_grant:{grant_code}"
raise HTTPException(status_code=401, detail="Grant code already used")
if grant.expires_at < datetime.now(timezone.utc): pipe = sync_redis_client.client.pipeline()
raise HTTPException(status_code=401, detail="Grant code expired") pipe.get(redis_key)
pipe.delete(redis_key)
grant_json, _ = pipe.execute()
if not grant_json:
raise HTTPException(status_code=401, detail="Invalid or expired grant code")
grant_data = json.loads(grant_json)
module = db.query(Module).filter(Module.module_id == module_id).first() module = db.query(Module).filter(Module.module_id == module_id).first()
if not module or module.id != grant.module_id: if not module or str(module.id) != grant_data["module_id"]:
raise HTTPException(status_code=401, detail="Grant invalid for this module") raise HTTPException(status_code=401, detail="Grant invalid for this module")
if grant.environment_slug != environment_slug: if grant_data["environment_slug"] != environment_slug:
raise HTTPException(status_code=401, detail="Grant invalid for this environment") raise HTTPException(status_code=401, detail="Grant invalid for this environment")
user = db.query(User).filter(User.id == grant.user_id).first() user = db.query(User).filter(User.id == grant_data["user_id"]).first()
if not user: if not user:
raise HTTPException(status_code=401, detail="User not found") raise HTTPException(status_code=401, detail="User not found")
if grant.tenant_id: grant_tenant_id = grant_data.get("tenant_id")
if user.tenant_id != grant.tenant_id: if grant_tenant_id:
raise HTTPException( if str(user.tenant_id) != grant_tenant_id:
status_code=401, raise HTTPException(
detail="Tenant mismatch for SSO grant" status_code=401,
) detail="Tenant mismatch for SSO grant"
)
grant.is_used = True
grant.used_at = datetime.now(timezone.utc)
db.commit()
permissions = [] permissions = []
if user.role: if user.role:
@@ -219,7 +231,7 @@ class SSOService:
token_payload = { token_payload = {
"sub": str(user.id), "sub": str(user.id),
"email": user.email, "email": user.email,
"tenant_id": str(grant.tenant_id) if grant.tenant_id else None, "tenant_id": grant_tenant_id,
"module_id": module_id, "module_id": module_id,
"environment": environment_slug, "environment": environment_slug,
"permissions": permissions, "permissions": permissions,
+1
View File
@@ -11,3 +11,4 @@ email-validator>=2.1.0
redis==7.1.0 redis==7.1.0
requests==2.32.5 requests==2.32.5
typer==0.21.1 typer==0.21.1
httpx==0.28.1