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"])
+3 -2
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,7 +14,7 @@ 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)"""
@@ -32,5 +33,5 @@ async def get_admin_stats(
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,
} }
+24 -35
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,41 +12,37 @@ 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)
@@ -59,15 +56,12 @@ async def register(
@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:
@@ -88,8 +82,7 @@ async def login(
# 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
@@ -106,8 +99,7 @@ async def login(
@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.
@@ -116,24 +108,21 @@ async def google_oauth(
# 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()
@@ -151,12 +140,12 @@ async def google_oauth(
# 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)
+30 -31
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,11 +24,13 @@ 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"""
@@ -45,7 +52,7 @@ async def create_mail_account(
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
@@ -68,7 +75,7 @@ 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)
@@ -81,7 +88,7 @@ async def create_mail_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,21 +104,19 @@ 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,28 +127,28 @@ 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)
@@ -158,21 +163,19 @@ 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)
@@ -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)
@@ -220,6 +220,5 @@ async def auto_detect_mail_settings(
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,
} }
+3 -2
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,7 +23,7 @@ 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:
+5 -5
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
@@ -11,10 +12,7 @@ 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
@@ -28,7 +26,9 @@ class Settings(BaseSettings):
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
+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
+12 -15
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.
@@ -75,8 +80,7 @@ async def get_current_user(
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,12 +113,7 @@ 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)
@@ -125,7 +122,7 @@ def check_subscription_tier(required_tier: str):
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
+4 -1
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
@@ -25,7 +26,9 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
# 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 = (
+23 -13
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
@@ -84,10 +94,10 @@ class CredentialEncryption:
# 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(
@@ -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
+4 -3
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__)
@@ -29,7 +30,7 @@ def create_application() -> FastAPI:
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)
@@ -54,7 +55,7 @@ def create_application() -> FastAPI:
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")
+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",
] ]
+99 -37
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,11 +69,14 @@ 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)
@@ -69,28 +87,41 @@ class User(Base):
# 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
@@ -134,25 +165,32 @@ class MailAccount(Base):
# 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)
@@ -172,19 +210,24 @@ class ProcessingRun(Base):
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)
@@ -205,17 +248,20 @@ class ProcessingLog(Base):
# 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)
@@ -235,19 +281,20 @@ class NotificationConfig(Base):
# 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)
@@ -270,11 +317,14 @@ class MailServerPreset(Base):
# 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)
@@ -296,7 +346,9 @@ class SubscriptionPlan(Base):
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
@@ -304,17 +356,22 @@ class SubscriptionPlan(Base):
# 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
@@ -332,17 +389,20 @@ class AuditLog(Base):
# 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)
@@ -361,7 +421,9 @@ class GmailCredential(Base):
# 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")
+3
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
@@ -156,6 +157,7 @@ class MailAccountResponse(MailAccountBase):
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
+30 -25
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
@@ -26,14 +27,16 @@ class OAuthService:
"""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.
@@ -48,53 +51,55 @@ 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:
@@ -103,7 +108,7 @@ class OAuthService:
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
@@ -123,7 +128,7 @@ class OAuthService:
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:
+113 -100
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,21 +24,25 @@ 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
@@ -75,13 +80,11 @@ 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
@@ -112,15 +115,11 @@ class MailProcessor:
# 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()
@@ -128,14 +127,14 @@ class MailProcessor:
# 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)
@@ -173,13 +172,11 @@ 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
@@ -188,7 +185,9 @@ class MailProcessor:
# 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 = []
@@ -197,10 +196,12 @@ class MailProcessor:
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}")
@@ -231,45 +232,43 @@ class MailProcessor:
# 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
@@ -278,7 +277,7 @@ class MailProcessor:
# 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}")
@@ -300,7 +299,7 @@ class MailProcessor:
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.
@@ -327,15 +326,17 @@ class MailProcessor:
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"
@@ -349,26 +350,32 @@ class MailProcessor:
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
@@ -543,7 +550,7 @@ class MailServerAutoDetect:
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 = []
@@ -553,60 +560,66 @@ class MailServerAutoDetect:
# 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
+25 -17
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
@@ -57,7 +62,7 @@ async def process_mail_account(account_id: int):
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()
@@ -79,9 +84,7 @@ async def process_mail_account(account_id: int):
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
@@ -123,11 +126,13 @@ 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()
@@ -146,10 +151,7 @@ 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
@@ -191,14 +193,16 @@ async def process_mail_account(account_id: int):
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)
@@ -219,7 +223,9 @@ 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]
),
) )
) )
) )
@@ -232,7 +238,9 @@ async def process_all_enabled_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
+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))
+8 -1
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
@@ -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,6 +137,7 @@ 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(
@@ -166,6 +172,7 @@ 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)
+9 -2
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
@@ -56,5 +57,11 @@ class TestConfigValidation:
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
+2 -1
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,
@@ -48,7 +49,7 @@ class TestJWT:
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: