Fix CI failures: add backend/conftest.py for module resolution and run black formatting

- Add backend/conftest.py that inserts the backend directory into sys.path,
  fixing ModuleNotFoundError when pytest runs from the backend/ directory
  (as CI does with `cd backend && pytest tests/`)
- Run black formatter on all 28 backend files that needed reformatting
- All 53 tests pass with both `pytest tests/` and `python -m pytest tests/`

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