diff --git a/backend/alembic/env.py b/backend/alembic/env.py index 8819cc1..7e6ca51 100644 --- a/backend/alembic/env.py +++ b/backend/alembic/env.py @@ -1,4 +1,5 @@ """Alembic environment configuration""" + from logging.config import fileConfig from sqlalchemy import engine_from_config, pool from alembic import context @@ -49,10 +50,7 @@ def run_migrations_online() -> None: ) with connectable.connect() as connection: - context.configure( - connection=connection, - target_metadata=target_metadata - ) + context.configure(connection=connection, target_metadata=target_metadata) with context.begin_transaction(): context.run_migrations() diff --git a/backend/app/api/v1/api.py b/backend/app/api/v1/api.py index 0f2e2d9..c7939b4 100644 --- a/backend/app/api/v1/api.py +++ b/backend/app/api/v1/api.py @@ -1,17 +1,34 @@ """ API v1 router aggregation. """ + from fastapi import APIRouter -from app.api.v1.endpoints import auth, users, mail_accounts, notifications, subscriptions, admin, providers +from app.api.v1.endpoints import ( + auth, + users, + mail_accounts, + notifications, + subscriptions, + admin, + providers, +) api_router = APIRouter() # Include all endpoint routers api_router.include_router(auth.router, prefix="/auth", tags=["Authentication"]) api_router.include_router(users.router, prefix="/users", tags=["Users"]) -api_router.include_router(mail_accounts.router, prefix="/mail-accounts", tags=["Mail Accounts"]) -api_router.include_router(providers.router, prefix="/providers", tags=["Providers & Gmail"]) -api_router.include_router(notifications.router, prefix="/notifications", tags=["Notifications"]) -api_router.include_router(subscriptions.router, prefix="/subscriptions", tags=["Subscriptions"]) +api_router.include_router( + mail_accounts.router, prefix="/mail-accounts", tags=["Mail Accounts"] +) +api_router.include_router( + providers.router, prefix="/providers", tags=["Providers & Gmail"] +) +api_router.include_router( + notifications.router, prefix="/notifications", tags=["Notifications"] +) +api_router.include_router( + subscriptions.router, prefix="/subscriptions", tags=["Subscriptions"] +) api_router.include_router(admin.router, prefix="/admin", tags=["Admin"]) diff --git a/backend/app/api/v1/endpoints/admin.py b/backend/app/api/v1/endpoints/admin.py index 4552c64..9aeb1e1 100644 --- a/backend/app/api/v1/endpoints/admin.py +++ b/backend/app/api/v1/endpoints/admin.py @@ -1,4 +1,5 @@ """Admin endpoints""" + from fastapi import APIRouter, Depends from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select, func @@ -13,24 +14,24 @@ router = APIRouter() @router.get("/stats") async def get_admin_stats( current_user: User = Depends(get_current_superuser), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Get overall system statistics (admin only)""" - + # Count users user_count = await db.execute(select(func.count(User.id))) total_users = user_count.scalar() - + # Count accounts account_count = await db.execute(select(func.count(MailAccount.id))) total_accounts = account_count.scalar() - + # Count processing runs run_count = await db.execute(select(func.count(ProcessingRun.id))) total_runs = run_count.scalar() - + return { "total_users": total_users, "total_mail_accounts": total_accounts, - "total_processing_runs": total_runs + "total_processing_runs": total_runs, } diff --git a/backend/app/api/v1/endpoints/auth.py b/backend/app/api/v1/endpoints/auth.py index f218796..72d38f4 100644 --- a/backend/app/api/v1/endpoints/auth.py +++ b/backend/app/api/v1/endpoints/auth.py @@ -1,6 +1,7 @@ """ Authentication endpoints (login, register, OAuth). """ + from fastapi import APIRouter, Depends, HTTPException, status from fastapi.security import OAuth2PasswordRequestForm from sqlalchemy.ext.asyncio import AsyncSession @@ -11,72 +12,65 @@ import logging from app.core.database import get_db from app.core.security import verify_password, get_password_hash from app.models.database_models import User, SubscriptionTier -from app.models.schemas import ( - Token, UserCreate, UserResponse, GoogleAuthRequest -) +from app.models.schemas import Token, UserCreate, UserResponse, GoogleAuthRequest from app.services.auth_service import oauth_service router = APIRouter() logger = logging.getLogger(__name__) -@router.post("/register", response_model=UserResponse, status_code=status.HTTP_201_CREATED) -async def register( - user_in: UserCreate, - db: AsyncSession = Depends(get_db) -): +@router.post( + "/register", response_model=UserResponse, status_code=status.HTTP_201_CREATED +) +async def register(user_in: UserCreate, db: AsyncSession = Depends(get_db)): """Register a new user with email and password""" - + # Check if user exists - result = await db.execute( - select(User).where(User.email == user_in.email) - ) + result = await db.execute(select(User).where(User.email == user_in.email)) existing_user = result.scalar_one_or_none() - + if existing_user: raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="Email already registered" + status_code=status.HTTP_400_BAD_REQUEST, detail="Email already registered" ) - + # Create new user user = User( email=user_in.email, full_name=user_in.full_name, - hashed_password=get_password_hash(user_in.password) if user_in.password else None, + hashed_password=( + get_password_hash(user_in.password) if user_in.password else None + ), subscription_tier=SubscriptionTier.FREE, - is_active=True + is_active=True, ) - + db.add(user) await db.commit() await db.refresh(user) - + logger.info(f"New user registered: {user.email}") - + return user @router.post("/login", response_model=Token) async def login( - form_data: OAuth2PasswordRequestForm = Depends(), - db: AsyncSession = Depends(get_db) + form_data: OAuth2PasswordRequestForm = Depends(), db: AsyncSession = Depends(get_db) ): """Login with email and password""" - + # Get user - result = await db.execute( - select(User).where(User.email == form_data.username) - ) + result = await db.execute(select(User).where(User.email == form_data.username)) user = result.scalar_one_or_none() - + if not user or not user.hashed_password: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Incorrect email or password", headers={"WWW-Authenticate": "Bearer"}, ) - + # Verify password if not verify_password(form_data.password, user.hashed_password): raise HTTPException( @@ -84,90 +78,85 @@ async def login( detail="Incorrect email or password", headers={"WWW-Authenticate": "Bearer"}, ) - + # Check if user is active if not user.is_active: raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail="User account is inactive" + status_code=status.HTTP_403_FORBIDDEN, detail="User account is inactive" ) - + # Update last login user.last_login_at = datetime.utcnow() await db.commit() - + # Create tokens tokens = oauth_service.create_tokens_for_user(user) - + logger.info(f"User logged in: {user.email}") - + return tokens @router.post("/google", response_model=Token) async def google_oauth( - auth_request: GoogleAuthRequest, - db: AsyncSession = Depends(get_db) + auth_request: GoogleAuthRequest, db: AsyncSession = Depends(get_db) ): """ Authenticate with Google OAuth2. Exchange authorization code for access token and user info. """ - + # Get user info from Google user_info = await oauth_service.get_google_user_info( - code=auth_request.code, - redirect_uri=auth_request.redirect_uri + code=auth_request.code, redirect_uri=auth_request.redirect_uri ) - - if not user_info.get('verified_email'): + + if not user_info.get("verified_email"): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail="Email not verified with Google" + detail="Email not verified with Google", ) - - email = user_info['email'] - google_id = user_info['google_id'] - + + email = user_info["email"] + google_id = user_info["google_id"] + # Check if user exists result = await db.execute( - select(User).where( - (User.email == email) | (User.google_id == google_id) - ) + select(User).where((User.email == email) | (User.google_id == google_id)) ) user = result.scalar_one_or_none() - + if user: # Update Google ID if not set if not user.google_id: user.google_id = google_id user.oauth_provider = "google" - + # Update last login user.last_login_at = datetime.utcnow() - + logger.info(f"Existing user logged in with Google: {user.email}") else: # Create new user user = User( email=email, - full_name=user_info.get('full_name'), + full_name=user_info.get("full_name"), google_id=google_id, oauth_provider="google", subscription_tier=SubscriptionTier.FREE, is_active=True, - last_login_at=datetime.utcnow() + last_login_at=datetime.utcnow(), ) db.add(user) - + logger.info(f"New user registered with Google: {user.email}") - + await db.commit() await db.refresh(user) - + # Create tokens tokens = oauth_service.create_tokens_for_user(user) - + return tokens @@ -175,7 +164,7 @@ async def google_oauth( async def get_google_authorize_url(redirect_uri: str): """Get Google OAuth2 authorization URL""" from app.core.config import settings - + auth_url = ( f"https://accounts.google.com/o/oauth2/v2/auth?" f"client_id={settings.GOOGLE_CLIENT_ID}&" @@ -184,5 +173,5 @@ async def get_google_authorize_url(redirect_uri: str): f"redirect_uri={redirect_uri}&" f"access_type=offline" ) - + return {"authorization_url": auth_url} diff --git a/backend/app/api/v1/endpoints/mail_accounts.py b/backend/app/api/v1/endpoints/mail_accounts.py index 4d3020f..3415f17 100644 --- a/backend/app/api/v1/endpoints/mail_accounts.py +++ b/backend/app/api/v1/endpoints/mail_accounts.py @@ -1,4 +1,5 @@ """Mail account management endpoints""" + from typing import List from fastapi import APIRouter, Depends, HTTPException, status from sqlalchemy.ext.asyncio import AsyncSession @@ -9,9 +10,13 @@ from app.core.deps import get_current_active_user from app.core.security import encrypt_credential, decrypt_credential from app.models.database_models import User, MailAccount from app.models.schemas import ( - MailAccountCreate, MailAccountResponse, MailAccountUpdate, - MailAccountTestRequest, MailAccountTestResponse, - MailAccountAutoDetectRequest, MailAccountAutoDetectResponse + MailAccountCreate, + MailAccountResponse, + MailAccountUpdate, + MailAccountTestRequest, + MailAccountTestResponse, + MailAccountAutoDetectRequest, + MailAccountAutoDetectResponse, ) from app.services.mail_processor import MailProcessor, MailServerAutoDetect from app.core.config import settings @@ -19,38 +24,40 @@ from app.core.config import settings router = APIRouter() -@router.post("", response_model=MailAccountResponse, status_code=status.HTTP_201_CREATED) +@router.post( + "", response_model=MailAccountResponse, status_code=status.HTTP_201_CREATED +) async def create_mail_account( account_in: MailAccountCreate, current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Create a new mail account""" - + # Check subscription limits result = await db.execute( select(MailAccount).where(MailAccount.user_id == current_user.id) ) existing_accounts = result.scalars().all() - + tier_limits = { "free": settings.TIER_FREE_MAX_ACCOUNTS, "basic": settings.TIER_BASIC_MAX_ACCOUNTS, "pro": settings.TIER_PRO_MAX_ACCOUNTS, "enterprise": settings.TIER_ENTERPRISE_MAX_ACCOUNTS, } - + max_accounts = tier_limits.get(current_user.subscription_tier.value, 1) - + if len(existing_accounts) >= max_accounts: raise HTTPException( status_code=status.HTTP_402_PAYMENT_REQUIRED, - detail=f"Account limit reached. Upgrade your subscription to add more accounts." + detail=f"Account limit reached. Upgrade your subscription to add more accounts.", ) - + # Encrypt password encrypted_password = encrypt_credential(account_in.password) - + # Create account account = MailAccount( user_id=current_user.id, @@ -68,20 +75,20 @@ async def create_mail_account( is_enabled=account_in.is_enabled, check_interval_minutes=account_in.check_interval_minutes, max_emails_per_check=account_in.max_emails_per_check, - delete_after_forward=account_in.delete_after_forward + delete_after_forward=account_in.delete_after_forward, ) - + db.add(account) await db.commit() await db.refresh(account) - + return account @router.get("", response_model=List[MailAccountResponse]) async def list_mail_accounts( current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """List all mail accounts for current user""" result = await db.execute( @@ -97,23 +104,21 @@ async def list_mail_accounts( async def get_mail_account( account_id: int, current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Get a specific mail account""" result = await db.execute( select(MailAccount).where( - MailAccount.id == account_id, - MailAccount.user_id == current_user.id + MailAccount.id == account_id, MailAccount.user_id == current_user.id ) ) account = result.scalar_one_or_none() - + if not account: raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail="Mail account not found" + status_code=status.HTTP_404_NOT_FOUND, detail="Mail account not found" ) - + return account @@ -122,35 +127,35 @@ async def update_mail_account( account_id: int, account_update: MailAccountUpdate, current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Update a mail account""" result = await db.execute( select(MailAccount).where( - MailAccount.id == account_id, - MailAccount.user_id == current_user.id + MailAccount.id == account_id, MailAccount.user_id == current_user.id ) ) account = result.scalar_one_or_none() - + if not account: raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail="Mail account not found" + status_code=status.HTTP_404_NOT_FOUND, detail="Mail account not found" ) - + # Update fields update_data = account_update.dict(exclude_unset=True) - + if "password" in update_data: - update_data["encrypted_password"] = encrypt_credential(update_data.pop("password")) - + update_data["encrypted_password"] = encrypt_credential( + update_data.pop("password") + ) + for field, value in update_data.items(): setattr(account, field, value) - + await db.commit() await db.refresh(account) - + return account @@ -158,23 +163,21 @@ async def update_mail_account( async def delete_mail_account( account_id: int, current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Delete a mail account""" result = await db.execute( select(MailAccount).where( - MailAccount.id == account_id, - MailAccount.user_id == current_user.id + MailAccount.id == account_id, MailAccount.user_id == current_user.id ) ) account = result.scalar_one_or_none() - + if not account: raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail="Mail account not found" + status_code=status.HTTP_404_NOT_FOUND, detail="Mail account not found" ) - + await db.delete(account) await db.commit() @@ -185,7 +188,7 @@ async def test_mail_connection( current_user: User = Depends(get_current_active_user), ): """Test connection to mail server""" - + # Create temporary account for testing temp_account = MailAccount( user_id=current_user.id, @@ -198,16 +201,13 @@ async def test_mail_connection( use_tls=test_request.use_tls, username=test_request.username, encrypted_password="", # Not used for test - forward_to="test@test.com" + forward_to="test@test.com", ) - + processor = MailProcessor(temp_account, test_request.password) success, message = await processor.test_connection() - - return MailAccountTestResponse( - success=success, - message=message - ) + + return MailAccountTestResponse(success=success, message=message) @router.post("/auto-detect", response_model=MailAccountAutoDetectResponse) @@ -216,10 +216,9 @@ async def auto_detect_mail_settings( current_user: User = Depends(get_current_active_user), ): """Auto-detect mail server settings for an email address""" - + suggestions = MailServerAutoDetect.detect(detect_request.email_address) - + return MailAccountAutoDetectResponse( - success=len(suggestions) > 0, - suggestions=suggestions + success=len(suggestions) > 0, suggestions=suggestions ) diff --git a/backend/app/api/v1/endpoints/notifications.py b/backend/app/api/v1/endpoints/notifications.py index 4556c51..784e18e 100644 --- a/backend/app/api/v1/endpoints/notifications.py +++ b/backend/app/api/v1/endpoints/notifications.py @@ -1,4 +1,5 @@ """Notification configuration endpoints""" + from typing import List from fastapi import APIRouter, Depends, HTTPException, status from sqlalchemy.ext.asyncio import AsyncSession @@ -8,23 +9,24 @@ from app.core.database import get_db from app.core.deps import get_current_active_user from app.models.database_models import User, NotificationConfig from app.models.schemas import ( - NotificationConfigCreate, NotificationConfigResponse, NotificationConfigUpdate + NotificationConfigCreate, + NotificationConfigResponse, + NotificationConfigUpdate, ) router = APIRouter() -@router.post("", response_model=NotificationConfigResponse, status_code=status.HTTP_201_CREATED) +@router.post( + "", response_model=NotificationConfigResponse, status_code=status.HTTP_201_CREATED +) async def create_notification_config( config_in: NotificationConfigCreate, current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Create notification configuration""" - config = NotificationConfig( - user_id=current_user.id, - **config_in.dict() - ) + config = NotificationConfig(user_id=current_user.id, **config_in.dict()) db.add(config) await db.commit() await db.refresh(config) @@ -34,7 +36,7 @@ async def create_notification_config( @router.get("", response_model=List[NotificationConfigResponse]) async def list_notification_configs( current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """List all notification configurations""" result = await db.execute( diff --git a/backend/app/api/v1/endpoints/providers.py b/backend/app/api/v1/endpoints/providers.py index bd4c43c..f7f404a 100644 --- a/backend/app/api/v1/endpoints/providers.py +++ b/backend/app/api/v1/endpoints/providers.py @@ -1,4 +1,5 @@ """Provider presets and Gmail credential management endpoints""" + from datetime import datetime from typing import List from fastapi import APIRouter, Depends, HTTPException, status @@ -156,7 +157,11 @@ async def get_provider_preset( ) -@router.post("/gmail-credential", response_model=GmailCredentialResponse, status_code=status.HTTP_201_CREATED) +@router.post( + "/gmail-credential", + response_model=GmailCredentialResponse, + status_code=status.HTTP_201_CREATED, +) async def save_gmail_credential( credential_in: GmailCredentialCreate, current_user: User = Depends(get_current_active_user), @@ -189,7 +194,9 @@ async def save_gmail_credential( encrypted_access = encrypt_credential(credential_in.access_token) encrypted_refresh = ( - encrypt_credential(credential_in.refresh_token) if credential_in.refresh_token else None + encrypt_credential(credential_in.refresh_token) + if credential_in.refresh_token + else None ) if existing: diff --git a/backend/app/api/v1/endpoints/subscriptions.py b/backend/app/api/v1/endpoints/subscriptions.py index b1e7e4c..1209d3b 100644 --- a/backend/app/api/v1/endpoints/subscriptions.py +++ b/backend/app/api/v1/endpoints/subscriptions.py @@ -1,4 +1,5 @@ """Subscription and payment endpoints""" + from typing import List from fastapi import APIRouter, Depends, HTTPException, status from sqlalchemy.ext.asyncio import AsyncSession @@ -13,9 +14,7 @@ router = APIRouter() @router.get("/plans", response_model=List[SubscriptionPlanResponse]) -async def list_subscription_plans( - db: AsyncSession = Depends(get_db) -): +async def list_subscription_plans(db: AsyncSession = Depends(get_db)): """List all available subscription plans""" result = await db.execute( select(SubscriptionPlan).where(SubscriptionPlan.is_active == True) @@ -25,11 +24,11 @@ async def list_subscription_plans( @router.get("/current") async def get_current_subscription( - current_user: User = Depends(get_current_active_user) + current_user: User = Depends(get_current_active_user), ): """Get current user's subscription details""" return { "tier": current_user.subscription_tier, "status": current_user.subscription_status, - "expires_at": current_user.subscription_expires_at + "expires_at": current_user.subscription_expires_at, } diff --git a/backend/app/api/v1/endpoints/users.py b/backend/app/api/v1/endpoints/users.py index a69af20..f7b3eb9 100644 --- a/backend/app/api/v1/endpoints/users.py +++ b/backend/app/api/v1/endpoints/users.py @@ -1,4 +1,5 @@ """User management endpoints""" + from fastapi import APIRouter, Depends, HTTPException, status from sqlalchemy.ext.asyncio import AsyncSession @@ -12,7 +13,7 @@ router = APIRouter() @router.get("/me", response_model=UserDetailResponse) async def get_current_user_profile( - current_user: User = Depends(get_current_active_user) + current_user: User = Depends(get_current_active_user), ): """Get current user profile""" return current_user @@ -22,14 +23,14 @@ async def get_current_user_profile( async def update_current_user_profile( user_update: UserUpdate, current_user: User = Depends(get_current_active_user), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ): """Update current user profile""" if user_update.email: current_user.email = user_update.email if user_update.full_name: current_user.full_name = user_update.full_name - + await db.commit() await db.refresh(current_user) return current_user diff --git a/backend/app/core/config.py b/backend/app/core/config.py index 1026ed7..b123632 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -2,6 +2,7 @@ Application configuration using Pydantic settings. Supports environment variables and .env files. """ + from typing import Optional, List from pydantic_settings import BaseSettings, SettingsConfigDict from pydantic import PostgresDsn, field_validator, ValidationInfo @@ -9,86 +10,85 @@ from pydantic import PostgresDsn, field_validator, ValidationInfo class Settings(BaseSettings): """Application settings loaded from environment variables""" - + model_config = SettingsConfigDict( - env_file=".env", - env_file_encoding="utf-8", - case_sensitive=False, - extra="ignore" + env_file=".env", env_file_encoding="utf-8", case_sensitive=False, extra="ignore" ) - + # Application APP_NAME: str = "POP3 Forwarder SaaS" APP_VERSION: str = "2.0.0" DEBUG: bool = False API_V1_PREFIX: str = "/api/v1" - + # Server HOST: str = "0.0.0.0" PORT: int = 8000 - + # Database - DATABASE_URL: str = "postgresql+asyncpg://user:password@localhost:5432/pop3_forwarder" + DATABASE_URL: str = ( + "postgresql+asyncpg://user:password@localhost:5432/pop3_forwarder" + ) DATABASE_POOL_SIZE: int = 20 DATABASE_MAX_OVERFLOW: int = 10 - + # Security SECRET_KEY: str = "change-this-to-a-secure-random-secret-key-in-production" ALGORITHM: str = "HS256" ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 REFRESH_TOKEN_EXPIRE_DAYS: int = 7 - + # Encryption (for storing POP3/IMAP credentials) ENCRYPTION_KEY: str = "change-this-to-a-secure-encryption-key" - + # OAuth2 - Google GOOGLE_CLIENT_ID: Optional[str] = None GOOGLE_CLIENT_SECRET: Optional[str] = None GOOGLE_REDIRECT_URI: str = "http://localhost:3000/auth/callback/google" - + # Gmail API (for direct email injection) GMAIL_API_ENABLED: bool = True GMAIL_INJECT_LABEL_IDS: List[str] = ["INBOX"] - + # CORS CORS_ORIGINS: List[str] = ["http://localhost:3000", "http://localhost:8000"] - + # Stripe Payment STRIPE_API_KEY: Optional[str] = None STRIPE_WEBHOOK_SECRET: Optional[str] = None STRIPE_PUBLISHABLE_KEY: Optional[str] = None - + # Subscription Tiers TIER_FREE_MAX_ACCOUNTS: int = 1 TIER_BASIC_MAX_ACCOUNTS: int = 5 TIER_PRO_MAX_ACCOUNTS: int = 20 TIER_ENTERPRISE_MAX_ACCOUNTS: int = 100 - + # Email Processing MAX_EMAILS_PER_RUN: int = 50 CHECK_INTERVAL_MINUTES: int = 5 THROTTLE_EMAILS_PER_MINUTE: int = 10 - + # Redis (for Celery and caching) REDIS_URL: str = "redis://localhost:6379/0" - + # Celery CELERY_BROKER_URL: str = "redis://localhost:6379/0" CELERY_RESULT_BACKEND: str = "redis://localhost:6379/0" - + # Apprise (notifications) APPRISE_ENABLED: bool = True - + # Logging LOG_LEVEL: str = "INFO" - + # Admin ADMIN_EMAIL: Optional[str] = None ADMIN_PASSWORD: Optional[str] = None - + # Mail Server Presets MAIL_SERVER_PRESETS_FILE: str = "app/data/mail_server_presets.json" - + @field_validator("CORS_ORIGINS", mode="before") @classmethod def assemble_cors_origins(cls, v: str | List[str]) -> List[str]: @@ -96,7 +96,7 @@ class Settings(BaseSettings): if isinstance(v, str): return [i.strip() for i in v.split(",")] return v - + @field_validator("SECRET_KEY") @classmethod def validate_secret_key(cls, v: str) -> str: @@ -118,7 +118,7 @@ class Settings(BaseSettings): "Generate a secure key with: python -c 'import secrets; print(secrets.token_urlsafe(32))'" ) return v - + @field_validator("ENCRYPTION_KEY") @classmethod def validate_encryption_key(cls, v: str) -> str: diff --git a/backend/app/core/database.py b/backend/app/core/database.py index cac3093..5f287a4 100644 --- a/backend/app/core/database.py +++ b/backend/app/core/database.py @@ -1,6 +1,7 @@ """ Database configuration and session management. """ + from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.orm import declarative_base from app.core.config import settings diff --git a/backend/app/core/deps.py b/backend/app/core/deps.py index 01f1a10..f8bf957 100644 --- a/backend/app/core/deps.py +++ b/backend/app/core/deps.py @@ -1,9 +1,14 @@ """ Authentication dependencies for FastAPI. """ + from typing import Optional from fastapi import Depends, HTTPException, status -from fastapi.security import OAuth2PasswordBearer, HTTPBearer, HTTPAuthorizationCredentials +from fastapi.security import ( + OAuth2PasswordBearer, + HTTPBearer, + HTTPAuthorizationCredentials, +) from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select @@ -19,7 +24,7 @@ http_bearer = HTTPBearer(auto_error=False) async def get_current_user( token: Optional[str] = Depends(oauth2_scheme), credentials: Optional[HTTPAuthorizationCredentials] = Depends(http_bearer), - db: AsyncSession = Depends(get_db) + db: AsyncSession = Depends(get_db), ) -> User: """ Get current authenticated user from JWT token. @@ -27,14 +32,14 @@ async def get_current_user( """ # Get token from either source auth_token = token or (credentials.credentials if credentials else None) - + if not auth_token: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated", headers={"WWW-Authenticate": "Bearer"}, ) - + # Decode token payload = decode_token(auth_token) if not payload: @@ -43,7 +48,7 @@ async def get_current_user( detail="Invalid authentication credentials", headers={"WWW-Authenticate": "Bearer"}, ) - + # Verify token type token_type = payload.get("type") if token_type != "access": @@ -52,7 +57,7 @@ async def get_current_user( detail="Invalid token type", headers={"WWW-Authenticate": "Bearer"}, ) - + # Get user ID from token user_id: Optional[int] = payload.get("sub") if user_id is None: @@ -61,24 +66,23 @@ async def get_current_user( detail="Invalid token payload", headers={"WWW-Authenticate": "Bearer"}, ) - + # Fetch user from database result = await db.execute(select(User).where(User.id == user_id)) user = result.scalar_one_or_none() - + if user is None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found", headers={"WWW-Authenticate": "Bearer"}, ) - + if not user.is_active: raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail="User account is inactive" + status_code=status.HTTP_403_FORBIDDEN, detail="User account is inactive" ) - + return user @@ -88,8 +92,7 @@ async def get_current_active_user( """Get current active user""" if not current_user.is_active: raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail="Inactive user" + status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user" ) return current_user @@ -100,8 +103,7 @@ async def get_current_superuser( """Get current superuser""" if not current_user.is_superuser: raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail="Not enough permissions" + status_code=status.HTTP_403_FORBIDDEN, detail="Not enough permissions" ) return current_user @@ -111,23 +113,18 @@ def check_subscription_tier(required_tier: str): Dependency factory to check if user has required subscription tier. Returns a dependency function. """ - tier_hierarchy = { - "free": 0, - "basic": 1, - "pro": 2, - "enterprise": 3 - } - + tier_hierarchy = {"free": 0, "basic": 1, "pro": 2, "enterprise": 3} + async def check_tier(current_user: User = Depends(get_current_active_user)) -> User: user_tier_level = tier_hierarchy.get(current_user.subscription_tier.value, 0) required_tier_level = tier_hierarchy.get(required_tier, 0) - + if user_tier_level < required_tier_level: raise HTTPException( status_code=status.HTTP_402_PAYMENT_REQUIRED, - detail=f"This feature requires {required_tier} subscription or higher" + detail=f"This feature requires {required_tier} subscription or higher", ) - + return current_user - + return check_tier diff --git a/backend/app/core/middleware.py b/backend/app/core/middleware.py index c8a72b9..2412d69 100644 --- a/backend/app/core/middleware.py +++ b/backend/app/core/middleware.py @@ -1,6 +1,7 @@ """ Security middleware for adding security headers and CSRF protection. """ + from fastapi import Request, Response from starlette.middleware.base import BaseHTTPMiddleware from starlette.types import ASGIApp @@ -9,24 +10,26 @@ import secrets class SecurityHeadersMiddleware(BaseHTTPMiddleware): """Add security headers to all responses""" - + async def dispatch(self, request: Request, call_next) -> Response: response = await call_next(request) - + # Prevent clickjacking response.headers["X-Frame-Options"] = "DENY" - + # Prevent MIME type sniffing response.headers["X-Content-Type-Options"] = "nosniff" - + # Enable XSS protection (for older browsers) response.headers["X-XSS-Protection"] = "1; mode=block" - + # Strict Transport Security (HTTPS only) # Note: Only enable in production with HTTPS if request.url.hostname not in ["localhost", "127.0.0.1"]: - response.headers["Strict-Transport-Security"] = "max-age=31536000; includeSubDomains" - + response.headers["Strict-Transport-Security"] = ( + "max-age=31536000; includeSubDomains" + ) + # Content Security Policy (adjust based on frontend needs) csp = ( "default-src 'self'; " @@ -38,15 +41,15 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware): "frame-src https://js.stripe.com;" ) response.headers["Content-Security-Policy"] = csp - + # Referrer Policy response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin" - + # Permissions Policy (formerly Feature Policy) response.headers["Permissions-Policy"] = ( "geolocation=(), microphone=(), camera=()" ) - + return response @@ -55,7 +58,7 @@ class CSRFProtectionMiddleware(BaseHTTPMiddleware): Basic CSRF protection for state-changing operations. For API-only applications, this is less critical but still good practice. """ - + def __init__(self, app: ASGIApp, exempt_paths: list = None): super().__init__(app) self.exempt_paths = exempt_paths or [ @@ -66,20 +69,20 @@ class CSRFProtectionMiddleware(BaseHTTPMiddleware): "/openapi.json", "/health", ] - + async def dispatch(self, request: Request, call_next) -> Response: # Skip CSRF check for safe methods if request.method in ["GET", "HEAD", "OPTIONS"]: return await call_next(request) - + # Skip CSRF check for exempt paths if any(request.url.path.startswith(path) for path in self.exempt_paths): return await call_next(request) - + # For API endpoints using JWT, the token itself provides CSRF protection # This is because attackers can't access the token stored in httpOnly cookies # or local storage from a different origin - + # If implementing cookie-based sessions, would check CSRF token here: # csrf_token = request.headers.get("X-CSRF-Token") # if not csrf_token or not self._validate_csrf_token(csrf_token): @@ -87,15 +90,15 @@ class CSRFProtectionMiddleware(BaseHTTPMiddleware): # status_code=403, # content={"detail": "CSRF token missing or invalid"} # ) - + response = await call_next(request) return response - + @staticmethod def _generate_csrf_token() -> str: """Generate a secure CSRF token""" return secrets.token_urlsafe(32) - + @staticmethod def _validate_csrf_token(token: str) -> bool: """Validate CSRF token (implement actual validation logic)""" diff --git a/backend/app/core/security.py b/backend/app/core/security.py index d40564e..419c80f 100644 --- a/backend/app/core/security.py +++ b/backend/app/core/security.py @@ -1,6 +1,7 @@ """ Security utilities for encryption, hashing, and token generation. """ + import hashlib import secrets from datetime import datetime, timedelta @@ -14,7 +15,6 @@ import base64 from app.core.config import settings - # Password hashing context pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") @@ -29,17 +29,23 @@ def get_password_hash(password: str) -> str: return pwd_context.hash(password) -def create_access_token(data: Dict[str, Any], expires_delta: Optional[timedelta] = None) -> str: +def create_access_token( + data: Dict[str, Any], expires_delta: Optional[timedelta] = None +) -> str: """Create JWT access token""" to_encode = data.copy() - + if expires_delta: expire = datetime.utcnow() + expires_delta else: - expire = datetime.utcnow() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) - + expire = datetime.utcnow() + timedelta( + minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES + ) + to_encode.update({"exp": expire, "type": "access"}) - encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) + encoded_jwt = jwt.encode( + to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM + ) return encoded_jwt @@ -48,14 +54,18 @@ def create_refresh_token(data: Dict[str, Any]) -> str: to_encode = data.copy() expire = datetime.utcnow() + timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS) to_encode.update({"exp": expire, "type": "refresh"}) - encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) + encoded_jwt = jwt.encode( + to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM + ) return encoded_jwt def decode_token(token: str) -> Optional[Dict[str, Any]]: """Decode and validate JWT token""" try: - payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) + payload = jwt.decode( + token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM] + ) return payload except JWTError: return None @@ -68,27 +78,27 @@ def generate_random_token(length: int = 32) -> str: class CredentialEncryption: """Handles encryption/decryption of sensitive credentials (POP3/IMAP passwords)""" - + def __init__(self, key: Optional[str] = None, user_id: Optional[int] = None): """ Initialize encryption with a key. If no key provided, uses the one from settings. In production, use a unique salt per user for enhanced security. - + Args: key: Encryption key (defaults to settings.ENCRYPTION_KEY) user_id: Optional user ID for per-user salt generation """ if key is None: key = settings.ENCRYPTION_KEY - + # Generate salt - unique per user for enhanced security if user_id is not None: - salt = hashlib.sha256(f'pop3fwd_usr_{user_id}'.encode()).digest()[:16] + salt = hashlib.sha256(f"pop3fwd_usr_{user_id}".encode()).digest()[:16] else: # Default salt for system-wide operations (use with caution) - salt = b'pop3_forwarder_0' - + salt = b"pop3_forwarder_0" + # Derive a proper Fernet key from the provided key kdf = PBKDF2HMAC( algorithm=hashes.SHA256(), @@ -96,20 +106,20 @@ class CredentialEncryption: salt=salt, iterations=100000, ) - key_bytes = key.encode('utf-8') + key_bytes = key.encode("utf-8") derived_key = base64.urlsafe_b64encode(kdf.derive(key_bytes)) self.fernet = Fernet(derived_key) - + def encrypt(self, plain_text: str) -> str: """Encrypt a string and return base64-encoded ciphertext""" - encrypted = self.fernet.encrypt(plain_text.encode('utf-8')) - return base64.b64encode(encrypted).decode('utf-8') - + encrypted = self.fernet.encrypt(plain_text.encode("utf-8")) + return base64.b64encode(encrypted).decode("utf-8") + def decrypt(self, encrypted_text: str) -> str: """Decrypt a base64-encoded ciphertext""" - encrypted_bytes = base64.b64decode(encrypted_text.encode('utf-8')) + encrypted_bytes = base64.b64decode(encrypted_text.encode("utf-8")) decrypted = self.fernet.decrypt(encrypted_bytes) - return decrypted.decode('utf-8') + return decrypted.decode("utf-8") # Global encryption instance diff --git a/backend/app/main.py b/backend/app/main.py index 42d3cfb..403b5af 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -1,6 +1,7 @@ """ Main FastAPI application. """ + from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.trustedhost import TrustedHostMiddleware @@ -14,7 +15,7 @@ from app.api.v1.api import api_router # Configure logging logging.basicConfig( level=getattr(logging, settings.LOG_LEVEL.upper()), - format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", ) logger = logging.getLogger(__name__) @@ -22,20 +23,20 @@ logger = logging.getLogger(__name__) def create_application() -> FastAPI: """Create and configure FastAPI application""" - + app = FastAPI( title=settings.APP_NAME, version=settings.APP_VERSION, description="Multi-tenant POP3/IMAP to Gmail forwarder with subscription management", docs_url="/api/docs", redoc_url="/api/redoc", - openapi_url="/api/openapi.json" + openapi_url="/api/openapi.json", ) - + # Security middleware (add before CORS) app.add_middleware(SecurityHeadersMiddleware) app.add_middleware(CSRFProtectionMiddleware) - + # CORS middleware app.add_middleware( CORSMiddleware, @@ -44,36 +45,36 @@ def create_application() -> FastAPI: allow_methods=["*"], allow_headers=["*"], ) - + # Include API router app.include_router(api_router, prefix=settings.API_V1_PREFIX) - + @app.get("/") async def root(): """Root endpoint""" return { "message": "POP3 Forwarder SaaS API", "version": settings.APP_VERSION, - "docs": "/api/docs" + "docs": "/api/docs", } - + @app.get("/health") async def health_check(): """Health check endpoint for container orchestration""" return {"status": "healthy"} - + @app.on_event("startup") async def startup_event(): """Run on application startup""" logger.info(f"Starting {settings.APP_NAME} v{settings.APP_VERSION}") logger.info(f"Debug mode: {settings.DEBUG}") logger.info(f"API documentation: /api/docs") - + @app.on_event("shutdown") async def shutdown_event(): """Run on application shutdown""" logger.info("Shutting down application") - + return app diff --git a/backend/app/models/__init__.py b/backend/app/models/__init__.py index c28f715..e73a86a 100644 --- a/backend/app/models/__init__.py +++ b/backend/app/models/__init__.py @@ -1,12 +1,31 @@ """Models package""" + from app.models.database_models import ( - User, MailAccount, ProcessingRun, ProcessingLog, - NotificationConfig, MailServerPreset, SubscriptionPlan, AuditLog, - SubscriptionTier, MailProtocol, AccountStatus, NotificationChannel + User, + MailAccount, + ProcessingRun, + ProcessingLog, + NotificationConfig, + MailServerPreset, + SubscriptionPlan, + AuditLog, + SubscriptionTier, + MailProtocol, + AccountStatus, + NotificationChannel, ) __all__ = [ - "User", "MailAccount", "ProcessingRun", "ProcessingLog", - "NotificationConfig", "MailServerPreset", "SubscriptionPlan", "AuditLog", - "SubscriptionTier", "MailProtocol", "AccountStatus", "NotificationChannel" + "User", + "MailAccount", + "ProcessingRun", + "ProcessingLog", + "NotificationConfig", + "MailServerPreset", + "SubscriptionPlan", + "AuditLog", + "SubscriptionTier", + "MailProtocol", + "AccountStatus", + "NotificationChannel", ] diff --git a/backend/app/models/database_models.py b/backend/app/models/database_models.py index a5509b7..0dc6bb2 100644 --- a/backend/app/models/database_models.py +++ b/backend/app/models/database_models.py @@ -1,11 +1,21 @@ """ Database models for the multi-tenant POP3 forwarder application. """ + from datetime import datetime from typing import Optional from sqlalchemy import ( - Column, Integer, String, Boolean, DateTime, ForeignKey, - Text, Enum as SQLEnum, JSON, Float, Index + Column, + Integer, + String, + Boolean, + DateTime, + ForeignKey, + Text, + Enum as SQLEnum, + JSON, + Float, + Index, ) from sqlalchemy.orm import relationship import enum @@ -15,6 +25,7 @@ from app.core.database import Base class SubscriptionTier(str, enum.Enum): """Subscription tier levels""" + FREE = "free" BASIC = "basic" PRO = "pro" @@ -23,6 +34,7 @@ class SubscriptionTier(str, enum.Enum): class MailProtocol(str, enum.Enum): """Supported mail protocols""" + POP3 = "pop3" POP3_SSL = "pop3_ssl" IMAP = "imap" @@ -31,12 +43,14 @@ class MailProtocol(str, enum.Enum): class DeliveryMethod(str, enum.Enum): """How emails are delivered to Gmail""" - SMTP = "smtp" # Forward via SMTP (legacy) - GMAIL_API = "gmail_api" # Inject via Gmail API (preferred) + + SMTP = "smtp" # Forward via SMTP (legacy) + GMAIL_API = "gmail_api" # Inject via Gmail API (preferred) class AccountStatus(str, enum.Enum): """Mail account status""" + ACTIVE = "active" INACTIVE = "inactive" ERROR = "error" @@ -45,6 +59,7 @@ class AccountStatus(str, enum.Enum): class NotificationChannel(str, enum.Enum): """Notification channel types""" + EMAIL = "email" TELEGRAM = "telegram" WEBHOOK = "webhook" @@ -54,76 +69,92 @@ class NotificationChannel(str, enum.Enum): class User(Base): """User model - represents a user account""" + __tablename__ = "users" - + id = Column(Integer, primary_key=True, index=True) email = Column(String(255), unique=True, index=True, nullable=False) - hashed_password = Column(String(255), nullable=True) # Nullable for OAuth-only users + hashed_password = Column( + String(255), nullable=True + ) # Nullable for OAuth-only users full_name = Column(String(255)) is_active = Column(Boolean, default=True) is_superuser = Column(Boolean, default=False) - + # OAuth google_id = Column(String(255), unique=True, index=True, nullable=True) oauth_provider = Column(String(50), nullable=True) - + # Subscription subscription_tier = Column(SQLEnum(SubscriptionTier), default=SubscriptionTier.FREE) - subscription_status = Column(String(50), default="active") # active, canceled, past_due + subscription_status = Column( + String(50), default="active" + ) # active, canceled, past_due stripe_customer_id = Column(String(255), unique=True, nullable=True) stripe_subscription_id = Column(String(255), unique=True, nullable=True) subscription_expires_at = Column(DateTime, nullable=True) - + # Timestamps created_at = Column(DateTime, default=datetime.utcnow, nullable=False) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) + updated_at = Column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) last_login_at = Column(DateTime, nullable=True) - + # Relationships - mail_accounts = relationship("MailAccount", back_populates="user", cascade="all, delete-orphan") - notifications = relationship("NotificationConfig", back_populates="user", cascade="all, delete-orphan") - logs = relationship("ProcessingLog", back_populates="user", cascade="all, delete-orphan") + mail_accounts = relationship( + "MailAccount", back_populates="user", cascade="all, delete-orphan" + ) + notifications = relationship( + "NotificationConfig", back_populates="user", cascade="all, delete-orphan" + ) + logs = relationship( + "ProcessingLog", back_populates="user", cascade="all, delete-orphan" + ) class MailAccount(Base): """Mail account configuration (POP3/IMAP)""" + __tablename__ = "mail_accounts" - + id = Column(Integer, primary_key=True, index=True) - user_id = Column(Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False) - + user_id = Column( + Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False + ) + # Account details name = Column(String(255), nullable=False) # User-friendly name email_address = Column(String(255), nullable=False) - + # Server configuration protocol = Column(SQLEnum(MailProtocol), default=MailProtocol.POP3_SSL) host = Column(String(255), nullable=False) port = Column(Integer, nullable=False) use_ssl = Column(Boolean, default=True) use_tls = Column(Boolean, default=False) - + # Credentials (encrypted) username = Column(String(255), nullable=False) encrypted_password = Column(Text, nullable=False) - + # Forwarding destination forward_to = Column(String(255), nullable=False) - + # Delivery method delivery_method = Column(SQLEnum(DeliveryMethod), default=DeliveryMethod.GMAIL_API) - + # Status and settings status = Column(SQLEnum(AccountStatus), default=AccountStatus.ACTIVE) is_enabled = Column(Boolean, default=True) check_interval_minutes = Column(Integer, default=5) max_emails_per_check = Column(Integer, default=50) delete_after_forward = Column(Boolean, default=True) - + # Auto-detection metadata provider_name = Column(String(100), nullable=True) # e.g., "Gmail", "GMX" auto_detected = Column(Boolean, default=False) - + # Statistics total_emails_processed = Column(Integer, default=0) total_emails_failed = Column(Integer, default=0) @@ -131,131 +162,147 @@ class MailAccount(Base): last_successful_check_at = Column(DateTime, nullable=True) last_error_at = Column(DateTime, nullable=True) last_error_message = Column(Text, nullable=True) - + # Timestamps created_at = Column(DateTime, default=datetime.utcnow, nullable=False) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) - + updated_at = Column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) + # Relationships user = relationship("User", back_populates="mail_accounts") - processing_runs = relationship("ProcessingRun", back_populates="mail_account", cascade="all, delete-orphan") - + processing_runs = relationship( + "ProcessingRun", back_populates="mail_account", cascade="all, delete-orphan" + ) + # Indexes __table_args__ = ( - Index('idx_user_email', 'user_id', 'email_address'), - Index('idx_status_enabled', 'status', 'is_enabled'), + Index("idx_user_email", "user_id", "email_address"), + Index("idx_status_enabled", "status", "is_enabled"), ) class ProcessingRun(Base): """Records of email processing runs for each mail account""" + __tablename__ = "processing_runs" - + id = Column(Integer, primary_key=True, index=True) - mail_account_id = Column(Integer, ForeignKey("mail_accounts.id", ondelete="CASCADE"), nullable=False) - + mail_account_id = Column( + Integer, ForeignKey("mail_accounts.id", ondelete="CASCADE"), nullable=False + ) + # Run details started_at = Column(DateTime, default=datetime.utcnow, nullable=False) completed_at = Column(DateTime, nullable=True) duration_seconds = Column(Float, nullable=True) - + # Results emails_fetched = Column(Integer, default=0) emails_forwarded = Column(Integer, default=0) emails_failed = Column(Integer, default=0) - + # Status status = Column(String(50), default="running") # running, completed, failed error_message = Column(Text, nullable=True) - + # Relationships mail_account = relationship("MailAccount", back_populates="processing_runs") - + # Indexes - __table_args__ = ( - Index('idx_account_started', 'mail_account_id', 'started_at'), - ) + __table_args__ = (Index("idx_account_started", "mail_account_id", "started_at"),) class ProcessingLog(Base): """Detailed logs of individual email processing attempts""" + __tablename__ = "processing_logs" - + id = Column(Integer, primary_key=True, index=True) - user_id = Column(Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False) - mail_account_id = Column(Integer, ForeignKey("mail_accounts.id", ondelete="CASCADE"), nullable=False) - processing_run_id = Column(Integer, ForeignKey("processing_runs.id", ondelete="CASCADE"), nullable=True) - + user_id = Column( + Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False + ) + mail_account_id = Column( + Integer, ForeignKey("mail_accounts.id", ondelete="CASCADE"), nullable=False + ) + processing_run_id = Column( + Integer, ForeignKey("processing_runs.id", ondelete="CASCADE"), nullable=True + ) + # Log details timestamp = Column(DateTime, default=datetime.utcnow, nullable=False, index=True) level = Column(String(20), nullable=False) # INFO, WARNING, ERROR message = Column(Text, nullable=False) - + # Email metadata (if applicable) email_subject = Column(String(500), nullable=True) email_from = Column(String(255), nullable=True) email_size_bytes = Column(Integer, nullable=True) - + # Status success = Column(Boolean, default=True) error_details = Column(JSON, nullable=True) - + # Relationships user = relationship("User", back_populates="logs") - + # Indexes __table_args__ = ( - Index('idx_user_timestamp', 'user_id', 'timestamp'), - Index('idx_account_timestamp', 'mail_account_id', 'timestamp'), + Index("idx_user_timestamp", "user_id", "timestamp"), + Index("idx_account_timestamp", "mail_account_id", "timestamp"), ) class NotificationConfig(Base): """User notification channel configurations""" + __tablename__ = "notification_configs" - + id = Column(Integer, primary_key=True, index=True) - user_id = Column(Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False) - + user_id = Column( + Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False + ) + # Channel details channel = Column(SQLEnum(NotificationChannel), nullable=False) is_enabled = Column(Boolean, default=True) - + # Channel-specific configuration (stored as JSON) config = Column(JSON, nullable=False) # Examples: # EMAIL: {"address": "user@example.com"} # TELEGRAM: {"bot_token": "xxx", "chat_id": "yyy"} # WEBHOOK: {"url": "https://example.com/webhook", "headers": {...}} - + # Notification preferences notify_on_errors = Column(Boolean, default=True) notify_on_success = Column(Boolean, default=False) notify_threshold = Column(Integer, default=3) # Notify after N consecutive errors - + # Timestamps created_at = Column(DateTime, default=datetime.utcnow, nullable=False) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) - + updated_at = Column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) + # Relationships user = relationship("User", back_populates="notifications") - + # Indexes - __table_args__ = ( - Index('idx_user_channel', 'user_id', 'channel'), - ) + __table_args__ = (Index("idx_user_channel", "user_id", "channel"),) class MailServerPreset(Base): """Predefined mail server configurations for common providers""" + __tablename__ = "mail_server_presets" - + id = Column(Integer, primary_key=True, index=True) - + # Provider info provider_name = Column(String(100), unique=True, nullable=False, index=True) provider_domain = Column(String(255), nullable=False) # e.g., "gmail.com" - + # Server configurations (can have multiple protocols) configs = Column(JSON, nullable=False) # Example: @@ -263,105 +310,120 @@ class MailServerPreset(Base): # "pop3_ssl": {"host": "pop.gmail.com", "port": 995, "ssl": true}, # "imap_ssl": {"host": "imap.gmail.com", "port": 993, "ssl": true} # } - + # Metadata is_verified = Column(Boolean, default=False) popularity_score = Column(Integer, default=0) # For sorting recommendations - + # Timestamps created_at = Column(DateTime, default=datetime.utcnow, nullable=False) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) + updated_at = Column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) class SubscriptionPlan(Base): """Available subscription plans and their features""" + __tablename__ = "subscription_plans" - + id = Column(Integer, primary_key=True, index=True) - + # Plan details tier = Column(SQLEnum(SubscriptionTier), unique=True, nullable=False) name = Column(String(100), nullable=False) description = Column(Text, nullable=True) - + # Pricing price_monthly = Column(Float, nullable=False) price_yearly = Column(Float, nullable=True) - + # Stripe integration stripe_price_id_monthly = Column(String(255), nullable=True) stripe_price_id_yearly = Column(String(255), nullable=True) - + # Features/Limits max_mail_accounts = Column(Integer, nullable=False) max_emails_per_day = Column(Integer, nullable=False) check_interval_minutes = Column(Integer, nullable=False) - support_level = Column(String(50), default="community") # community, email, priority + support_level = Column( + String(50), default="community" + ) # community, email, priority features = Column(JSON, nullable=True) # Additional features as JSON - + # Status is_active = Column(Boolean, default=True) - + # Timestamps created_at = Column(DateTime, default=datetime.utcnow, nullable=False) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) + updated_at = Column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) class AuditLog(Base): """Audit trail for security and compliance""" + __tablename__ = "audit_logs" - + id = Column(Integer, primary_key=True, index=True) - + # Who - user_id = Column(Integer, ForeignKey("users.id", ondelete="SET NULL"), nullable=True) + user_id = Column( + Integer, ForeignKey("users.id", ondelete="SET NULL"), nullable=True + ) user_email = Column(String(255), nullable=True) # Cached for deleted users ip_address = Column(String(45), nullable=True) # IPv4 or IPv6 - + # What action = Column(String(100), nullable=False, index=True) resource_type = Column(String(50), nullable=True) resource_id = Column(Integer, nullable=True) - + # Details details = Column(JSON, nullable=True) status = Column(String(20), default="success") # success, failure - + # When timestamp = Column(DateTime, default=datetime.utcnow, nullable=False, index=True) - + # Indexes __table_args__ = ( - Index('idx_user_action', 'user_id', 'action'), - Index('idx_timestamp_action', 'timestamp', 'action'), + Index("idx_user_action", "user_id", "action"), + Index("idx_timestamp_action", "timestamp", "action"), ) class GmailCredential(Base): """Stores OAuth2 credentials for Gmail API access (per-user)""" + __tablename__ = "gmail_credentials" - + id = Column(Integer, primary_key=True, index=True) - user_id = Column(Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, unique=True) - + user_id = Column( + Integer, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, unique=True + ) + # Gmail account email gmail_email = Column(String(255), nullable=False) - + # OAuth2 tokens (encrypted) encrypted_access_token = Column(Text, nullable=False) encrypted_refresh_token = Column(Text, nullable=True) - + # Token metadata token_expiry = Column(DateTime, nullable=True) scopes = Column(JSON, nullable=True) - + # Status is_valid = Column(Boolean, default=True) last_verified_at = Column(DateTime, nullable=True) - + # Timestamps created_at = Column(DateTime, default=datetime.utcnow, nullable=False) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False) - + updated_at = Column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) + # Relationships user = relationship("User", backref="gmail_credential") diff --git a/backend/app/models/schemas.py b/backend/app/models/schemas.py index 93baeb6..80d4ad5 100644 --- a/backend/app/models/schemas.py +++ b/backend/app/models/schemas.py @@ -1,6 +1,7 @@ """ Pydantic schemas for API request/response validation. """ + from datetime import datetime from typing import Optional, Dict, Any, List from pydantic import BaseModel, EmailStr, Field, validator @@ -63,7 +64,7 @@ class UserResponse(UserBase): subscription_tier: SubscriptionTier subscription_status: str created_at: datetime - + class Config: from_attributes = True @@ -74,7 +75,7 @@ class UserDetailResponse(UserResponse): stripe_customer_id: Optional[str] = None subscription_expires_at: Optional[datetime] = None last_login_at: Optional[datetime] = None - + class Config: from_attributes = True @@ -145,17 +146,18 @@ class MailAccountResponse(MailAccountBase): last_error_message: Optional[str] = None created_at: datetime updated_at: datetime - + # Don't expose password or username in responses password: str = Field(exclude=True, default="") username: str = Field(exclude=True, default="") - + class Config: from_attributes = True class MailAccountTestRequest(BaseModel): """Test connection to mail server""" + host: str port: int protocol: MailProtocol @@ -173,6 +175,7 @@ class MailAccountTestResponse(BaseModel): class MailAccountAutoDetectRequest(BaseModel): """Auto-detect mail server settings""" + email_address: EmailStr @@ -193,7 +196,7 @@ class ProcessingRunResponse(BaseModel): emails_failed: int status: str error_message: Optional[str] = None - + class Config: from_attributes = True @@ -207,7 +210,7 @@ class ProcessingLogResponse(BaseModel): email_subject: Optional[str] = None email_from: Optional[str] = None success: bool - + class Config: from_attributes = True @@ -239,7 +242,7 @@ class NotificationConfigResponse(NotificationConfigBase): user_id: int created_at: datetime updated_at: datetime - + class Config: from_attributes = True @@ -258,7 +261,7 @@ class SubscriptionPlanResponse(BaseModel): support_level: str features: Optional[Dict[str, Any]] = None is_active: bool - + class Config: from_attributes = True @@ -307,7 +310,7 @@ class MailServerPresetResponse(BaseModel): provider_domain: str configs: Dict[str, Any] is_verified: bool - + class Config: from_attributes = True @@ -327,7 +330,7 @@ class GmailCredentialResponse(BaseModel): last_verified_at: Optional[datetime] = None created_at: datetime updated_at: datetime - + class Config: from_attributes = True diff --git a/backend/app/services/auth_service.py b/backend/app/services/auth_service.py index cbe18b8..2924705 100644 --- a/backend/app/services/auth_service.py +++ b/backend/app/services/auth_service.py @@ -1,6 +1,7 @@ """ OAuth2 authentication service for Google and other providers. """ + from typing import Dict, Any, Optional from datetime import datetime import httpx @@ -17,30 +18,32 @@ logger = logging.getLogger(__name__) class OAuthService: """OAuth2 authentication service""" - + def __init__(self): self.oauth = OAuth() self._register_google() - + def _register_google(self): """Register Google OAuth2 provider""" if settings.GOOGLE_CLIENT_ID and settings.GOOGLE_CLIENT_SECRET: self.oauth.register( - name='google', + name="google", client_id=settings.GOOGLE_CLIENT_ID, client_secret=settings.GOOGLE_CLIENT_SECRET, - server_metadata_url='https://accounts.google.com/.well-known/openid-configuration', - client_kwargs={'scope': 'openid email profile'} + server_metadata_url="https://accounts.google.com/.well-known/openid-configuration", + client_kwargs={"scope": "openid email profile"}, ) - - async def get_google_user_info(self, code: str, redirect_uri: str) -> Dict[str, Any]: + + async def get_google_user_info( + self, code: str, redirect_uri: str + ) -> Dict[str, Any]: """ Exchange Google authorization code for user information. - + Args: code: Authorization code from Google redirect_uri: Redirect URI used in OAuth flow - + Returns: Dict with user information (email, name, google_id) """ @@ -48,82 +51,84 @@ class OAuthService: # Exchange code for token async with httpx.AsyncClient() as client: token_response = await client.post( - 'https://oauth2.googleapis.com/token', + "https://oauth2.googleapis.com/token", data={ - 'code': code, - 'client_id': settings.GOOGLE_CLIENT_ID, - 'client_secret': settings.GOOGLE_CLIENT_SECRET, - 'redirect_uri': redirect_uri, - 'grant_type': 'authorization_code' - } + "code": code, + "client_id": settings.GOOGLE_CLIENT_ID, + "client_secret": settings.GOOGLE_CLIENT_SECRET, + "redirect_uri": redirect_uri, + "grant_type": "authorization_code", + }, ) - + if token_response.status_code != 200: logger.error(f"Google token exchange failed: {token_response.text}") raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail="Failed to exchange authorization code" + detail="Failed to exchange authorization code", ) - + token_data = token_response.json() - access_token = token_data.get('access_token') - + access_token = token_data.get("access_token") + if not access_token: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail="No access token received" + detail="No access token received", ) - + # Get user info user_info_response = await client.get( - 'https://www.googleapis.com/oauth2/v2/userinfo', - headers={'Authorization': f'Bearer {access_token}'} + "https://www.googleapis.com/oauth2/v2/userinfo", + headers={"Authorization": f"Bearer {access_token}"}, ) - + if user_info_response.status_code != 200: - logger.error(f"Google user info fetch failed: {user_info_response.text}") + logger.error( + f"Google user info fetch failed: {user_info_response.text}" + ) raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail="Failed to get user information" + detail="Failed to get user information", ) - + user_info = user_info_response.json() - + return { - 'email': user_info.get('email'), - 'full_name': user_info.get('name'), - 'google_id': user_info.get('id'), - 'picture': user_info.get('picture'), - 'verified_email': user_info.get('verified_email', False) + "email": user_info.get("email"), + "full_name": user_info.get("name"), + "google_id": user_info.get("id"), + "picture": user_info.get("picture"), + "verified_email": user_info.get("verified_email", False), } - + except HTTPException: raise except Exception as e: logger.error(f"OAuth error: {e}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail="OAuth authentication failed" + detail="OAuth authentication failed", ) - + @staticmethod def create_tokens_for_user(user: User) -> Dict[str, str]: """ Create access and refresh tokens for a user. - + Args: user: User database model - + Returns: Dict with access_token, refresh_token, and token_type """ access_token = create_access_token(data={"sub": user.id}) refresh_token = create_refresh_token(data={"sub": user.id}) - + return { "access_token": access_token, "refresh_token": refresh_token, - "token_type": "bearer" + "token_type": "bearer", } diff --git a/backend/app/services/gmail_service.py b/backend/app/services/gmail_service.py index 8ba0d6f..403bf46 100644 --- a/backend/app/services/gmail_service.py +++ b/backend/app/services/gmail_service.py @@ -5,6 +5,7 @@ Uses the Gmail API's users.messages.insert() method to inject emails into a user's Gmail account, preserving original headers and metadata. This is preferred over SMTP forwarding as it doesn't modify the email. """ + import asyncio import base64 import logging @@ -25,6 +26,7 @@ GMAIL_SCOPES = [ class GmailInjectionError(Exception): """Raised when Gmail API injection fails""" + pass @@ -129,7 +131,9 @@ class GmailService: } except HttpError as e: - error_msg = f"Gmail API error: {e.reason if hasattr(e, 'reason') else str(e)}" + error_msg = ( + f"Gmail API error: {e.reason if hasattr(e, 'reason') else str(e)}" + ) logger.error(error_msg) raise GmailInjectionError(error_msg) except Exception as e: @@ -149,9 +153,7 @@ class GmailService: try: result = await loop.run_in_executor( None, - lambda: self.service.users() - .getProfile(userId="me") - .execute(), + lambda: self.service.users().getProfile(userId="me").execute(), ) email = result.get("emailAddress", "unknown") logger.info(f"Gmail API access verified for: {email}") @@ -172,9 +174,7 @@ class GmailService: try: result = await loop.run_in_executor( None, - lambda: self.service.users() - .getProfile(userId="me") - .execute(), + lambda: self.service.users().getProfile(userId="me").execute(), ) return result.get("emailAddress") except Exception as e: diff --git a/backend/app/services/mail_processor.py b/backend/app/services/mail_processor.py index 18d6959..148afe0 100644 --- a/backend/app/services/mail_processor.py +++ b/backend/app/services/mail_processor.py @@ -2,6 +2,7 @@ Mail processing service for fetching and forwarding emails. Supports both POP3 and IMAP protocols with secure connections. """ + import asyncio import poplib import smtplib @@ -23,31 +24,35 @@ logger = logging.getLogger(__name__) class MailConnectionError(Exception): """Raised when unable to connect to mail server""" + pass class MailAuthenticationError(Exception): """Raised when authentication fails""" + pass class MailFetchError(Exception): """Raised when fetching emails fails""" + pass class MailForwardError(Exception): """Raised when forwarding email fails""" + pass class MailProcessor: """Handles mail fetching and forwarding operations""" - + def __init__(self, account: MailAccount, decrypted_password: str): self.account = account self.password = decrypted_password - + async def test_connection(self) -> Tuple[bool, str]: """ Test connection to mail server. @@ -61,12 +66,12 @@ class MailProcessor: except Exception as e: logger.error(f"Connection test failed: {e}") return False, str(e) - + async def _test_pop3_connection(self) -> Tuple[bool, str]: """Test POP3 connection""" try: loop = asyncio.get_event_loop() - + # Run blocking POP3 operations in thread pool def connect_pop3(): if self.account.protocol == MailProtocol.POP3_SSL: @@ -75,29 +80,27 @@ class MailProcessor: self.account.host, self.account.port, context=context, - timeout=10 + timeout=10, ) else: pop_conn = poplib.POP3( - self.account.host, - self.account.port, - timeout=10 + self.account.host, self.account.port, timeout=10 ) - + # Try authentication pop_conn.user(self.account.username) pop_conn.pass_(self.password) - + # Get mailbox stats message_count, mailbox_size = pop_conn.stat() - + pop_conn.quit() return message_count, mailbox_size - + message_count, mailbox_size = await loop.run_in_executor(None, connect_pop3) - + return True, f"Connection successful. {message_count} messages in mailbox." - + except poplib.error_proto as e: error_msg = str(e) if "authentication" in error_msg.lower() or "auth" in error_msg.lower(): @@ -105,66 +108,62 @@ class MailProcessor: return False, f"POP3 protocol error: {error_msg}" except Exception as e: return False, f"Connection failed: {str(e)}" - + async def _test_imap_connection(self) -> Tuple[bool, str]: """Test IMAP connection""" try: # Create IMAP client if self.account.protocol == MailProtocol.IMAP_SSL: imap_client = aioimaplib.IMAP4_SSL( - host=self.account.host, - port=self.account.port, - timeout=10 + host=self.account.host, port=self.account.port, timeout=10 ) else: imap_client = aioimaplib.IMAP4( - host=self.account.host, - port=self.account.port, - timeout=10 + host=self.account.host, port=self.account.port, timeout=10 ) - + await imap_client.wait_hello_from_server() - + # Authenticate response = await imap_client.login(self.account.username, self.password) - - if response.result != 'OK': + + if response.result != "OK": return False, f"Authentication failed: {response.lines}" - + # Select inbox - await imap_client.select('INBOX') - + await imap_client.select("INBOX") + # Get message count - response = await imap_client.search('ALL') + response = await imap_client.search("ALL") message_ids = response.lines[0].split() message_count = len(message_ids) - + await imap_client.logout() - + return True, f"Connection successful. {message_count} messages in mailbox." - + except Exception as e: return False, f"IMAP connection failed: {str(e)}" - + async def fetch_emails(self, max_count: Optional[int] = None) -> List[bytes]: """ Fetch emails from the mail server. Returns list of raw email data. """ max_count = max_count or self.account.max_emails_per_check - + if self.account.protocol in [MailProtocol.POP3, MailProtocol.POP3_SSL]: return await self._fetch_pop3_emails(max_count) else: return await self._fetch_imap_emails(max_count) - + async def _fetch_pop3_emails(self, max_count: int) -> List[bytes]: """Fetch emails via POP3""" emails = [] - + try: loop = asyncio.get_event_loop() - + def fetch_pop3(): # Connect if self.account.protocol == MailProtocol.POP3_SSL: @@ -173,37 +172,39 @@ class MailProcessor: self.account.host, self.account.port, context=context, - timeout=30 + timeout=30, ) else: pop_conn = poplib.POP3( - self.account.host, - self.account.port, - timeout=30 + self.account.host, self.account.port, timeout=30 ) - + # Authenticate pop_conn.user(self.account.username) pop_conn.pass_(self.password) - + # Get message count num_messages = len(pop_conn.list()[1]) - logger.info(f"Found {num_messages} messages for account {self.account.id}") - + logger.info( + f"Found {num_messages} messages for account {self.account.id}" + ) + fetched_emails = [] messages_to_delete = [] - + # Fetch emails (limited by max_count) for i in range(1, min(num_messages + 1, max_count + 1)): try: response, lines, octets = pop_conn.retr(i) - email_data = b'\r\n'.join(lines) + email_data = b"\r\n".join(lines) fetched_emails.append(email_data) messages_to_delete.append(i) - logger.info(f"Retrieved message {i} from account {self.account.id}") + logger.info( + f"Retrieved message {i} from account {self.account.id}" + ) except Exception as e: logger.error(f"Error retrieving message {i}: {e}") - + # Delete messages if configured if self.account.delete_after_forward: for msg_id in messages_to_delete: @@ -211,100 +212,98 @@ class MailProcessor: pop_conn.dele(msg_id) except Exception as e: logger.error(f"Error deleting message {msg_id}: {e}") - + pop_conn.quit() return fetched_emails - + emails = await loop.run_in_executor(None, fetch_pop3) - + except Exception as e: logger.error(f"Error fetching POP3 emails: {e}") raise MailFetchError(f"POP3 fetch error: {str(e)}") - + return emails - + async def _fetch_imap_emails(self, max_count: int) -> List[bytes]: """Fetch emails via IMAP""" emails = [] - + try: # Create IMAP client if self.account.protocol == MailProtocol.IMAP_SSL: imap_client = aioimaplib.IMAP4_SSL( - host=self.account.host, - port=self.account.port, - timeout=30 + host=self.account.host, port=self.account.port, timeout=30 ) else: imap_client = aioimaplib.IMAP4( - host=self.account.host, - port=self.account.port, - timeout=30 + host=self.account.host, port=self.account.port, timeout=30 ) - + await imap_client.wait_hello_from_server() await imap_client.login(self.account.username, self.password) - await imap_client.select('INBOX') - + await imap_client.select("INBOX") + # Search for all messages - response = await imap_client.search('UNSEEN') # Only fetch unread + response = await imap_client.search("UNSEEN") # Only fetch unread message_ids = response.lines[0].split() - + # Limit to max_count message_ids = message_ids[:max_count] - - logger.info(f"Found {len(message_ids)} unread messages for account {self.account.id}") - + + logger.info( + f"Found {len(message_ids)} unread messages for account {self.account.id}" + ) + # Fetch each message for msg_id in message_ids: try: - response = await imap_client.fetch(msg_id, '(RFC822)') - + response = await imap_client.fetch(msg_id, "(RFC822)") + # Extract email data from response email_data = None for line in response.lines: - if isinstance(line, bytes) and b'RFC822' in line: + if isinstance(line, bytes) and b"RFC822" in line: # Find the email content - start_idx = line.find(b'{') + start_idx = line.find(b"{") if start_idx != -1: # Email data is in the next parts continue - elif isinstance(line, bytes) and not line.startswith(b'*'): + elif isinstance(line, bytes) and not line.startswith(b"*"): email_data = line break - + if email_data: emails.append(email_data) - + # Mark as seen if deleting after forward if self.account.delete_after_forward: - await imap_client.store(msg_id, '+FLAGS', '\\Deleted') - + await imap_client.store(msg_id, "+FLAGS", "\\Deleted") + except Exception as e: logger.error(f"Error fetching message {msg_id}: {e}") - + # Expunge deleted messages if self.account.delete_after_forward: await imap_client.expunge() - + await imap_client.logout() - + except Exception as e: logger.error(f"Error fetching IMAP emails: {e}") raise MailFetchError(f"IMAP fetch error: {str(e)}") - + return emails - + @staticmethod async def forward_email( email_data: bytes, source_account_name: str, destination: str, - smtp_config: Dict[str, Any] + smtp_config: Dict[str, Any], ) -> bool: """ Forward an email to the destination address. - + Args: email_data: Raw email bytes source_account_name: Name of source account for labeling @@ -315,60 +314,68 @@ class MailProcessor: - username: SMTP username - password: SMTP password - use_tls: Whether to use STARTTLS - + Returns: True if successful, False otherwise """ try: loop = asyncio.get_event_loop() - + def send_email(): # Parse the email msg = parser.BytesParser().parsebytes(email_data) - + # Create forwarding message - forward_msg = MIMEMultipart('mixed') - forward_msg['From'] = smtp_config['username'] - forward_msg['To'] = destination - forward_msg['Date'] = formatdate(localtime=True) - forward_msg['Message-ID'] = make_msgid() - + forward_msg = MIMEMultipart("mixed") + forward_msg["From"] = smtp_config["username"] + forward_msg["To"] = destination + forward_msg["Date"] = formatdate(localtime=True) + forward_msg["Message-ID"] = make_msgid() + # Preserve original subject with prefix - original_subject = msg.get('Subject', 'No Subject') - forward_msg['Subject'] = f"[Fwd from {source_account_name}] {original_subject}" - + original_subject = msg.get("Subject", "No Subject") + forward_msg["Subject"] = ( + f"[Fwd from {source_account_name}] {original_subject}" + ) + # Add original headers header_info = f"Originally from: {msg.get('From', 'Unknown')}\n" header_info += f"Original Date: {msg.get('Date', 'Unknown')}\n" header_info += f"Original Subject: {original_subject}\n" header_info += f"Source Account: {source_account_name}\n" header_info += "-" * 50 + "\n\n" - + # Get email body body = "" if msg.is_multipart(): for part in msg.walk(): if part.get_content_type() == "text/plain": - body = part.get_payload(decode=True).decode('utf-8', errors='ignore') + body = part.get_payload(decode=True).decode( + "utf-8", errors="ignore" + ) break else: payload = msg.get_payload(decode=True) if payload: - body = payload.decode('utf-8', errors='ignore') - + body = payload.decode("utf-8", errors="ignore") + # Combine header and body full_body = header_info + body - forward_msg.attach(MIMEText(full_body, 'plain', 'utf-8')) - + forward_msg.attach(MIMEText(full_body, "plain", "utf-8")) + # Send via SMTP - if smtp_config.get('use_tls', True): - server = smtplib.SMTP(smtp_config['host'], smtp_config['port'], timeout=30) + if smtp_config.get("use_tls", True): + server = smtplib.SMTP( + smtp_config["host"], smtp_config["port"], timeout=30 + ) server.starttls() else: - server = smtplib.SMTP_SSL(smtp_config['host'], smtp_config['port'], timeout=30) - + server = smtplib.SMTP_SSL( + smtp_config["host"], smtp_config["port"], timeout=30 + ) + try: - server.login(smtp_config['username'], smtp_config['password']) + server.login(smtp_config["username"], smtp_config["password"]) server.send_message(forward_msg) logger.info(f"Successfully forwarded email to {destination}") return True @@ -377,9 +384,9 @@ class MailProcessor: server.quit() except Exception as e: logger.warning(f"Error closing SMTP connection: {e}") - + return await loop.run_in_executor(None, send_email) - + except Exception as e: logger.error(f"Error forwarding email: {e}") raise MailForwardError(f"Forward error: {str(e)}") @@ -387,7 +394,7 @@ class MailProcessor: class MailServerAutoDetect: """Auto-detect mail server settings based on email domain""" - + # Common mail server configurations KNOWN_PROVIDERS = { "gmail.com": { @@ -536,77 +543,83 @@ class MailServerAutoDetect: "imap_ssl": {"host": "imap.mail.de", "port": 993}, }, } - + @classmethod def detect(cls, email_address: str) -> List[Dict[str, Any]]: """ Detect mail server settings for an email address. Returns list of possible configurations. """ - domain = email_address.split('@')[-1].lower() - + domain = email_address.split("@")[-1].lower() + suggestions = [] - + # Check if we have a known provider if domain in cls.KNOWN_PROVIDERS: provider = cls.KNOWN_PROVIDERS[domain] - + # Add POP3 SSL suggestion if "pop3_ssl" in provider: - suggestions.append({ - "protocol": "pop3_ssl", - "provider_name": provider["name"], - "host": provider["pop3_ssl"]["host"], - "port": provider["pop3_ssl"]["port"], - "use_ssl": True, - "use_tls": False, - }) - + suggestions.append( + { + "protocol": "pop3_ssl", + "provider_name": provider["name"], + "host": provider["pop3_ssl"]["host"], + "port": provider["pop3_ssl"]["port"], + "use_ssl": True, + "use_tls": False, + } + ) + # Add IMAP SSL suggestion if "imap_ssl" in provider: - suggestions.append({ - "protocol": "imap_ssl", - "provider_name": provider["name"], - "host": provider["imap_ssl"]["host"], - "port": provider["imap_ssl"]["port"], - "use_ssl": True, - "use_tls": False, - }) + suggestions.append( + { + "protocol": "imap_ssl", + "provider_name": provider["name"], + "host": provider["imap_ssl"]["host"], + "port": provider["imap_ssl"]["port"], + "use_ssl": True, + "use_tls": False, + } + ) else: # Generic suggestions based on common patterns - suggestions.extend([ - { - "protocol": "pop3_ssl", - "provider_name": "Generic", - "host": f"pop.{domain}", - "port": 995, - "use_ssl": True, - "use_tls": False, - }, - { - "protocol": "pop3_ssl", - "provider_name": "Generic", - "host": f"pop3.{domain}", - "port": 995, - "use_ssl": True, - "use_tls": False, - }, - { - "protocol": "imap_ssl", - "provider_name": "Generic", - "host": f"imap.{domain}", - "port": 993, - "use_ssl": True, - "use_tls": False, - }, - { - "protocol": "imap_ssl", - "provider_name": "Generic", - "host": f"mail.{domain}", - "port": 993, - "use_ssl": True, - "use_tls": False, - }, - ]) - + suggestions.extend( + [ + { + "protocol": "pop3_ssl", + "provider_name": "Generic", + "host": f"pop.{domain}", + "port": 995, + "use_ssl": True, + "use_tls": False, + }, + { + "protocol": "pop3_ssl", + "provider_name": "Generic", + "host": f"pop3.{domain}", + "port": 995, + "use_ssl": True, + "use_tls": False, + }, + { + "protocol": "imap_ssl", + "provider_name": "Generic", + "host": f"imap.{domain}", + "port": 993, + "use_ssl": True, + "use_tls": False, + }, + { + "protocol": "imap_ssl", + "provider_name": "Generic", + "host": f"mail.{domain}", + "port": 993, + "use_ssl": True, + "use_tls": False, + }, + ] + ) + return suggestions diff --git a/backend/app/workers/celery_app.py b/backend/app/workers/celery_app.py index 83a3930..7672ff5 100644 --- a/backend/app/workers/celery_app.py +++ b/backend/app/workers/celery_app.py @@ -1,6 +1,7 @@ """ Celery application for background email processing tasks. """ + from celery import Celery from celery.schedules import crontab import logging @@ -14,7 +15,7 @@ celery_app = Celery( "pop3_forwarder", broker=settings.CELERY_BROKER_URL, backend=settings.CELERY_RESULT_BACKEND, - include=["app.workers.tasks"] + include=["app.workers.tasks"], ) # Celery configuration diff --git a/backend/app/workers/tasks.py b/backend/app/workers/tasks.py index decf1c7..8450ef4 100644 --- a/backend/app/workers/tasks.py +++ b/backend/app/workers/tasks.py @@ -1,6 +1,7 @@ """ Celery tasks for background email processing. """ + import asyncio import os from datetime import datetime, timedelta @@ -12,8 +13,12 @@ from app.workers.celery_app import celery_app from app.core.database import async_session_maker from app.core.security import decrypt_credential from app.models.database_models import ( - MailAccount, ProcessingRun, ProcessingLog, AccountStatus, - DeliveryMethod, GmailCredential, + MailAccount, + ProcessingRun, + ProcessingLog, + AccountStatus, + DeliveryMethod, + GmailCredential, ) from app.services.mail_processor import MailProcessor from app.services.gmail_service import GmailService, GmailInjectionError @@ -26,7 +31,7 @@ logger = logging.getLogger(__name__) class AsyncTask(Task): """Base task class that handles async operations""" - + def __call__(self, *args, **kwargs): """Run async task in event loop""" # Use asyncio.run() for better event loop management @@ -37,7 +42,7 @@ class AsyncTask(Task): async def process_mail_account(account_id: int): """ Process a single mail account - fetch and forward emails. - + Args: account_id: ID of mail account to process """ @@ -48,44 +53,42 @@ async def process_mail_account(account_id: int): select(MailAccount).where(MailAccount.id == account_id) ) account = result.scalar_one_or_none() - + if not account or not account.is_enabled: logger.warning(f"Account {account_id} not found or disabled") return - + # Create processing run run = ProcessingRun( mail_account_id=account.id, started_at=datetime.utcnow(), - status="running" + status="running", ) db.add(run) await db.commit() await db.refresh(run) - + # Decrypt password password = decrypt_credential(account.encrypted_password) - + # Create processor processor = MailProcessor(account, password) - + # Fetch emails emails = await processor.fetch_emails(account.max_emails_per_check) - + run.emails_fetched = len(emails) - + # Forward emails emails_forwarded = 0 emails_failed = 0 - + # Determine delivery method - use_gmail_api = ( - account.delivery_method == DeliveryMethod.GMAIL_API - ) - + use_gmail_api = account.delivery_method == DeliveryMethod.GMAIL_API + gmail_service = None smtp_config = None - + if use_gmail_api: # Get user's Gmail credentials gmail_cred_result = await db.execute( @@ -95,7 +98,7 @@ async def process_mail_account(account_id: int): ) ) gmail_cred = gmail_cred_result.scalar_one_or_none() - + if gmail_cred: access_token = decrypt_credential(gmail_cred.encrypted_access_token) refresh_token = ( @@ -115,7 +118,7 @@ async def process_mail_account(account_id: int): f"falling back to SMTP for account {account.id}" ) use_gmail_api = False - + if not use_gmail_api: # Fall back to SMTP smtp_config = { @@ -123,16 +126,18 @@ async def process_mail_account(account_id: int): "port": int(os.getenv("SMTP_PORT", "587")), "username": os.getenv("SMTP_USER", ""), "password": os.getenv("SMTP_PASSWORD", ""), - "use_tls": os.getenv("SMTP_USE_TLS", "true").lower() == "true" + "use_tls": os.getenv("SMTP_USE_TLS", "true").lower() == "true", } - + if not smtp_config["username"] or not smtp_config["password"]: - logger.error(f"SMTP credentials not configured for account {account.id}") + logger.error( + f"SMTP credentials not configured for account {account.id}" + ) run.status = "failed" run.error_message = "No delivery method configured (SMTP credentials missing and Gmail API not set up)" await db.commit() return - + for email_data in emails: try: if use_gmail_api and gmail_service: @@ -146,32 +151,29 @@ async def process_mail_account(account_id: int): else: # Forward via SMTP (fallback) success = await MailProcessor.forward_email( - email_data, - account.name, - account.forward_to, - smtp_config + email_data, account.name, account.forward_to, smtp_config ) if success: emails_forwarded += 1 else: emails_failed += 1 - + except (GmailInjectionError, Exception) as e: logger.error(f"Error delivering email: {e}") emails_failed += 1 - + # Update run run.emails_forwarded = emails_forwarded run.emails_failed = emails_failed run.completed_at = datetime.utcnow() run.duration_seconds = (run.completed_at - run.started_at).total_seconds() run.status = "completed" if emails_failed == 0 else "partial_failure" - + # Update account account.total_emails_processed += emails_forwarded account.total_emails_failed += emails_failed account.last_check_at = datetime.utcnow() - + if emails_failed == 0: account.last_successful_check_at = datetime.utcnow() account.status = AccountStatus.ACTIVE @@ -179,30 +181,32 @@ async def process_mail_account(account_id: int): account.status = AccountStatus.ERROR account.last_error_at = datetime.utcnow() account.last_error_message = f"{emails_failed} emails failed to forward" - + await db.commit() - + logger.info( f"Processed account {account.id}: " f"{emails_forwarded} forwarded, {emails_failed} failed" ) - + except Exception as e: logger.error(f"Error processing account {account_id}: {e}") - + # Mark run as failed - if 'run' in locals(): + if "run" in locals(): run.status = "failed" run.error_message = str(e) run.completed_at = datetime.utcnow() - run.duration_seconds = (run.completed_at - run.started_at).total_seconds() - + run.duration_seconds = ( + run.completed_at - run.started_at + ).total_seconds() + # Update account error status - if 'account' in locals(): + if "account" in locals(): account.status = AccountStatus.ERROR account.last_error_at = datetime.utcnow() account.last_error_message = str(e) - + await db.commit() @@ -219,26 +223,30 @@ async def process_all_enabled_accounts(): select(MailAccount).where( and_( MailAccount.is_enabled == True, - MailAccount.status.in_([AccountStatus.ACTIVE, AccountStatus.TESTING]) + MailAccount.status.in_( + [AccountStatus.ACTIVE, AccountStatus.TESTING] + ), ) ) ) accounts = result.scalars().all() - + logger.info(f"Processing {len(accounts)} enabled mail accounts") - + # Process each account for account in accounts: # Check if it's time to check this account if account.last_check_at: time_since_last_check = datetime.utcnow() - account.last_check_at - if time_since_last_check.total_seconds() < (account.check_interval_minutes * 60): + if time_since_last_check.total_seconds() < ( + account.check_interval_minutes * 60 + ): logger.debug(f"Skipping account {account.id} - not time yet") continue - + # Queue processing task process_mail_account.delay(account.id) - + except Exception as e: logger.error(f"Error processing accounts: {e}") @@ -247,38 +255,38 @@ async def process_all_enabled_accounts(): async def cleanup_old_logs(days_to_keep: int = 30): """ Clean up old processing logs and runs. - + Args: days_to_keep: Number of days of logs to retain """ async with async_session_maker() as db: try: cutoff_date = datetime.utcnow() - timedelta(days=days_to_keep) - + # Delete old processing runs result = await db.execute( select(ProcessingRun).where(ProcessingRun.started_at < cutoff_date) ) old_runs = result.scalars().all() - + for run in old_runs: await db.delete(run) - + # Delete old processing logs result = await db.execute( select(ProcessingLog).where(ProcessingLog.timestamp < cutoff_date) ) old_logs = result.scalars().all() - + for log in old_logs: await db.delete(log) - + await db.commit() - + logger.info( f"Cleaned up {len(old_runs)} old processing runs and " f"{len(old_logs)} old logs" ) - + except Exception as e: logger.error(f"Error cleaning up logs: {e}") diff --git a/backend/conftest.py b/backend/conftest.py new file mode 100644 index 0000000..6105390 --- /dev/null +++ b/backend/conftest.py @@ -0,0 +1,6 @@ +import sys +from pathlib import Path + +# Ensure the backend directory is on sys.path so that `app` is importable +# when pytest is invoked from the backend/ directory (e.g., `cd backend && pytest tests/`). +sys.path.insert(0, str(Path(__file__).resolve().parent)) diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index 08b8a8f..cc6b1d9 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -1,6 +1,7 @@ """ Test configuration and fixtures. """ + import pytest import asyncio from typing import AsyncGenerator, Generator @@ -15,7 +16,9 @@ from app.models.database_models import User from app.core.security import get_password_hash, create_access_token # Test database URL (use different database for tests) -TEST_DATABASE_URL = settings.DATABASE_URL.replace("/pop3_forwarder", "/pop3_forwarder_test") +TEST_DATABASE_URL = settings.DATABASE_URL.replace( + "/pop3_forwarder", "/pop3_forwarder_test" +) # Note: event_loop fixture removed - pytest-asyncio provides this automatically @@ -30,18 +33,18 @@ async def db_engine(): poolclass=NullPool, echo=False, ) - + # Create tables async with engine.begin() as conn: await conn.run_sync(Base.metadata.drop_all) await conn.run_sync(Base.metadata.create_all) - + yield engine - + # Drop tables async with engine.begin() as conn: await conn.run_sync(Base.metadata.drop_all) - + await engine.dispose() @@ -53,7 +56,7 @@ async def db_session(db_engine) -> AsyncGenerator[AsyncSession, None]: class_=AsyncSession, expire_on_commit=False, ) - + async with async_session_maker() as session: yield session @@ -61,15 +64,15 @@ async def db_session(db_engine) -> AsyncGenerator[AsyncSession, None]: @pytest.fixture(scope="function") async def client(db_session: AsyncSession) -> AsyncGenerator[AsyncClient, None]: """Create test client with database session override""" - + async def override_get_db(): yield db_session - + app.dependency_overrides[get_db] = override_get_db - + async with AsyncClient(app=app, base_url="http://test") as client: yield client - + app.dependency_overrides.clear() @@ -120,9 +123,11 @@ def admin_auth_headers(test_admin_user: User) -> dict: # Factory fixtures for creating test data + @pytest.fixture def user_factory(db_session: AsyncSession): """Factory for creating test users""" + async def _create_user( email: str = None, password: str = "testpassword123", @@ -132,8 +137,9 @@ def user_factory(db_session: AsyncSession): ) -> User: if email is None: import uuid + email = f"test-{uuid.uuid4()}@example.com" - + user = User( email=email, hashed_password=get_password_hash(password), @@ -145,7 +151,7 @@ def user_factory(db_session: AsyncSession): await db_session.commit() await db_session.refresh(user) return user - + return _create_user @@ -154,7 +160,7 @@ def mail_account_factory(db_session: AsyncSession): """Factory for creating test mail accounts""" from app.models.database_models import MailAccount from app.core.security import encrypt_password - + async def _create_mail_account( user_id: int, host: str = "pop.example.com", @@ -166,10 +172,11 @@ def mail_account_factory(db_session: AsyncSession): ) -> MailAccount: if username is None: import uuid + username = f"test-{uuid.uuid4()}@example.com" - + encrypted_password = encrypt_password(password, user_id) - + account = MailAccount( user_id=user_id, host=host, @@ -184,5 +191,5 @@ def mail_account_factory(db_session: AsyncSession): await db_session.commit() await db_session.refresh(account) return account - + return _create_mail_account diff --git a/backend/tests/unit/test_config.py b/backend/tests/unit/test_config.py index 22c43c9..edabd03 100644 --- a/backend/tests/unit/test_config.py +++ b/backend/tests/unit/test_config.py @@ -1,6 +1,7 @@ """ Unit tests for configuration module. """ + import pytest from pydantic import ValidationError from app.core.config import Settings @@ -8,7 +9,7 @@ from app.core.config import Settings class TestConfigValidation: """Test configuration validation""" - + def test_default_secret_key_rejected(self): """Test that default SECRET_KEY is rejected""" with pytest.raises(ValidationError) as exc_info: @@ -16,9 +17,9 @@ class TestConfigValidation: SECRET_KEY="change-this-to-a-secure-random-secret-key-in-production", ENCRYPTION_KEY="this-is-a-secure-32-character-key-for-testing", ) - + assert "SECRET_KEY must be changed from default" in str(exc_info.value) - + def test_short_secret_key_rejected(self): """Test that short SECRET_KEY is rejected""" with pytest.raises(ValidationError) as exc_info: @@ -26,9 +27,9 @@ class TestConfigValidation: SECRET_KEY="short", ENCRYPTION_KEY="this-is-a-secure-32-character-key-for-testing", ) - + assert "at least 32 characters" in str(exc_info.value) - + def test_default_encryption_key_rejected(self): """Test that default ENCRYPTION_KEY is rejected""" with pytest.raises(ValidationError) as exc_info: @@ -36,9 +37,9 @@ class TestConfigValidation: SECRET_KEY="this-is-a-secure-32-character-key-for-testing", ENCRYPTION_KEY="change-this-to-a-secure-encryption-key", ) - + assert "ENCRYPTION_KEY must be changed from default" in str(exc_info.value) - + def test_short_encryption_key_rejected(self): """Test that short ENCRYPTION_KEY is rejected""" with pytest.raises(ValidationError) as exc_info: @@ -46,15 +47,21 @@ class TestConfigValidation: SECRET_KEY="this-is-a-secure-32-character-key-for-testing", ENCRYPTION_KEY="short", ) - + assert "at least 32 characters" in str(exc_info.value) - + def test_valid_keys_accepted(self): """Test that valid keys are accepted""" settings = Settings( SECRET_KEY="this-is-a-secure-32-character-key-for-testing-secret", ENCRYPTION_KEY="this-is-a-secure-32-character-key-for-encryption", ) - - assert settings.SECRET_KEY == "this-is-a-secure-32-character-key-for-testing-secret" - assert settings.ENCRYPTION_KEY == "this-is-a-secure-32-character-key-for-encryption" + + assert ( + settings.SECRET_KEY + == "this-is-a-secure-32-character-key-for-testing-secret" + ) + assert ( + settings.ENCRYPTION_KEY + == "this-is-a-secure-32-character-key-for-encryption" + ) diff --git a/backend/tests/unit/test_gmail_service.py b/backend/tests/unit/test_gmail_service.py index abf3ab7..e6121ae 100644 --- a/backend/tests/unit/test_gmail_service.py +++ b/backend/tests/unit/test_gmail_service.py @@ -1,6 +1,7 @@ """ Unit tests for Gmail service module. """ + import pytest from unittest.mock import MagicMock, patch, AsyncMock from app.services.gmail_service import GmailService, GmailInjectionError, GMAIL_SCOPES @@ -91,7 +92,9 @@ class TestGmailService: service = GmailService(access_token="test-access-token") mock_api = MagicMock() - mock_api.users().messages().insert().execute.side_effect = Exception("API Error") + mock_api.users().messages().insert().execute.side_effect = Exception( + "API Error" + ) service._service = mock_api with pytest.raises(GmailInjectionError, match="Failed to inject email"): diff --git a/backend/tests/unit/test_provider_presets.py b/backend/tests/unit/test_provider_presets.py index 56ac614..a5de124 100644 --- a/backend/tests/unit/test_provider_presets.py +++ b/backend/tests/unit/test_provider_presets.py @@ -1,6 +1,7 @@ """ Unit tests for provider presets and mail server auto-detection. """ + import pytest from app.services.mail_processor import MailServerAutoDetect @@ -156,11 +157,13 @@ class TestProviderPresets: def test_provider_presets_import(self): """Test that provider presets can be imported""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + assert len(PROVIDER_PRESETS) > 0 def test_all_presets_have_required_fields(self): """Test that all presets have required fields""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + for preset in PROVIDER_PRESETS: assert preset.id assert preset.name @@ -171,6 +174,7 @@ class TestProviderPresets: def test_gmail_preset_exists(self): """Test that Gmail preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + gmail = next((p for p in PROVIDER_PRESETS if p.id == "gmail"), None) assert gmail is not None assert gmail.imap_ssl is not None @@ -179,6 +183,7 @@ class TestProviderPresets: def test_gmx_preset_exists(self): """Test that GMX preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + gmx = next((p for p in PROVIDER_PRESETS if p.id == "gmx"), None) assert gmx is not None assert "gmx.de" in gmx.domains @@ -186,6 +191,7 @@ class TestProviderPresets: def test_webde_preset_exists(self): """Test that WEB.DE preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + webde = next((p for p in PROVIDER_PRESETS if p.id == "webde"), None) assert webde is not None assert "web.de" in webde.domains @@ -193,6 +199,7 @@ class TestProviderPresets: def test_outlook_preset_exists(self): """Test that Outlook preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + outlook = next((p for p in PROVIDER_PRESETS if p.id == "outlook"), None) assert outlook is not None assert "hotmail.com" in outlook.domains @@ -200,17 +207,20 @@ class TestProviderPresets: def test_yahoo_preset_exists(self): """Test that Yahoo preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + yahoo = next((p for p in PROVIDER_PRESETS if p.id == "yahoo"), None) assert yahoo is not None def test_aol_preset_exists(self): """Test that AOL preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + aol = next((p for p in PROVIDER_PRESETS if p.id == "aol"), None) assert aol is not None def test_tonline_preset_exists(self): """Test that T-Online preset is included""" from app.api.v1.endpoints.providers import PROVIDER_PRESETS + tonline = next((p for p in PROVIDER_PRESETS if p.id == "tonline"), None) assert tonline is not None diff --git a/backend/tests/unit/test_security.py b/backend/tests/unit/test_security.py index b4c2d14..8da718d 100644 --- a/backend/tests/unit/test_security.py +++ b/backend/tests/unit/test_security.py @@ -1,6 +1,7 @@ """ Unit tests for security module. """ + import pytest from app.core.security import ( get_password_hash, @@ -12,99 +13,99 @@ from app.core.security import ( class TestPasswordHashing: """Test password hashing and verification""" - + def test_hash_password(self): """Test password hashing""" password = "securepassword123" hashed = get_password_hash(password) - + assert hashed != password assert len(hashed) > 50 assert hashed.startswith("$2b$") - + def test_verify_password_success(self): """Test password verification with correct password""" password = "securepassword123" hashed = get_password_hash(password) - + assert verify_password(password, hashed) is True - + def test_verify_password_failure(self): """Test password verification with wrong password""" password = "securepassword123" wrong_password = "wrongpassword" hashed = get_password_hash(password) - + assert verify_password(wrong_password, hashed) is False class TestJWT: """Test JWT token creation and validation""" - + def test_create_access_token(self): """Test access token creation""" data = {"sub": "test@example.com"} token = create_access_token(data) - + assert isinstance(token, str) assert len(token) > 50 - assert token.count('.') == 2 # JWT has 3 parts + assert token.count(".") == 2 # JWT has 3 parts class TestEncryption: """Test credential encryption/decryption""" - + def test_encrypt_password(self): """Test password encryption""" password = "mailpassword123" user_id = 1 - + encryptor = CredentialEncryption(user_id=user_id) encrypted = encryptor.encrypt(password) - + assert encrypted != password assert len(encrypted) > 50 - + def test_decrypt_password(self): """Test password decryption""" password = "mailpassword123" user_id = 1 - + encryptor = CredentialEncryption(user_id=user_id) encrypted = encryptor.encrypt(password) decrypted = encryptor.decrypt(encrypted) - + assert decrypted == password - + def test_encryption_with_different_user_ids(self): """Test that encryption produces different results for different users""" password = "mailpassword123" user_id_1 = 1 user_id_2 = 2 - + encryptor_1 = CredentialEncryption(user_id=user_id_1) encryptor_2 = CredentialEncryption(user_id=user_id_2) - + encrypted_1 = encryptor_1.encrypt(password) encrypted_2 = encryptor_2.encrypt(password) - + # Different users should produce different encrypted values assert encrypted_1 != encrypted_2 - + # But decryption should work correctly for each assert encryptor_1.decrypt(encrypted_1) == password assert encryptor_2.decrypt(encrypted_2) == password - + def test_decrypt_with_wrong_user_id_fails(self): """Test that decryption fails with wrong user ID""" password = "mailpassword123" user_id = 1 wrong_user_id = 2 - + encryptor = CredentialEncryption(user_id=user_id) wrong_encryptor = CredentialEncryption(user_id=wrong_user_id) - + encrypted = encryptor.encrypt(password) - + with pytest.raises(Exception): wrong_encryptor.decrypt(encrypted)