diff --git a/app/api/__init__.py b/app/api/__init__.py index 36f9345a..046d8f1c 100644 --- a/app/api/__init__.py +++ b/app/api/__init__.py @@ -8,6 +8,7 @@ from fastapi import APIRouter from app.api.admin_users import router as admin_users_router from app.api.azure import router as azure_router +from app.api.billing import router as billing_router from app.api.database import router as database_router from app.api.diagnostic import router as diagnostic_router from app.api.dropbox import router as dropbox_router @@ -62,3 +63,4 @@ router.include_router(database_router) router.include_router(subscriptions_router) router.include_router(plans_router) router.include_router(onboarding_router) +router.include_router(billing_router) diff --git a/app/api/billing.py b/app/api/billing.py new file mode 100644 index 00000000..c44dd074 --- /dev/null +++ b/app/api/billing.py @@ -0,0 +1,413 @@ +"""Stripe billing integration for DocuElevate. + +Provides three endpoints: +- POST /api/billing/create-checkout-session — starts Stripe Checkout for a plan upgrade +- POST /api/billing/create-portal-session — opens Stripe Customer Portal (manage/cancel) +- POST /api/billing/webhook — handles Stripe webhook events +- GET /api/billing/success — success landing page after checkout + +Stripe Python SDK license: MIT (compatible with this project's Apache 2.0 license). + +GDPR: Stripe acts as a data processor under a Data Processing Agreement (DPA). + Stripe is SOC 2 Type II certified and supports EU data residency. +SOC2: Stripe is SOC 2 Type II certified. +EU VAT: Configure Stripe Tax in the Stripe Dashboard for automatic VAT collection. +""" + +import json +import logging +import pathlib +from datetime import datetime, timezone +from typing import Any + +import stripe +from fastapi import APIRouter, Depends, HTTPException, Request, status +from fastapi.templating import Jinja2Templates +from pydantic import BaseModel +from sqlalchemy.orm import Session + +from app.auth import require_login +from app.config import settings +from app.database import get_db +from app.models import SubscriptionPlan, UserProfile +from app.utils.user_scope import get_current_owner_id + +logger = logging.getLogger(__name__) +router = APIRouter(prefix="/billing", tags=["billing"]) + +_templates_dir = pathlib.Path(__file__).parents[2] / "frontend" / "templates" +_templates = Jinja2Templates(directory=str(_templates_dir)) + + +def _get_stripe() -> stripe.StripeClient | None: + """Return a configured Stripe client, or None when not configured.""" + if not settings.stripe_secret_key: + return None + return stripe.StripeClient(settings.stripe_secret_key) + + +def _get_or_create_stripe_customer( + client: stripe.StripeClient, + db: Session, + owner_id: str, + email: str | None, + name: str | None, +) -> str: + """Return the Stripe customer_id for *owner_id*, creating one if needed. + + Args: + client: Configured Stripe client. + db: Database session. + owner_id: Stable user identifier. + email: User's email for the Stripe customer record. + name: User's display name for the Stripe customer record. + + Returns: + The Stripe customer ID string. + """ + profile = db.query(UserProfile).filter(UserProfile.user_id == owner_id).first() + if profile and profile.stripe_customer_id: + return profile.stripe_customer_id + + customer = client.customers.create( + params={ + "email": email or "", + "name": name or "", + "metadata": {"docuelevate_user_id": owner_id}, + } + ) + if profile: + profile.stripe_customer_id = customer.id + db.commit() + return customer.id + + +class CheckoutSessionBody(BaseModel): + """Request body for creating a Stripe Checkout session.""" + + plan_id: str + billing_cycle: str = "monthly" # "monthly" | "yearly" + + +class PortalSessionBody(BaseModel): + """Request body for creating a Stripe Customer Portal session.""" + + return_url: str | None = None + + +@router.post("/create-checkout-session", summary="Create a Stripe Checkout session for a plan upgrade") +@require_login +async def create_checkout_session( + request: Request, + body: CheckoutSessionBody, + db: Session = Depends(get_db), +) -> dict[str, Any]: + """Create a Stripe Checkout session. + + The client should redirect the user to the returned ``checkout_url``. + + Raises: + 503: Stripe is not configured. + 404: Plan not found or has no Stripe price configured. + """ + client = _get_stripe() + if not client: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing is not configured.") + + plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == body.plan_id).first() + if plan is None: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Plan {body.plan_id!r} not found.") + + price_id = plan.stripe_price_id_yearly if body.billing_cycle == "yearly" else plan.stripe_price_id_monthly + if not price_id: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=( + f"Stripe price ID not configured for plan {body.plan_id!r} ({body.billing_cycle}). " + "Please set it in the Admin Plan Designer." + ), + ) + + user = request.session.get("user") or {} + owner_id = get_current_owner_id(request) or user.get("email") or "" + email = user.get("email") + name = user.get("name") + + customer_id = _get_or_create_stripe_customer(client, db, owner_id, email, name) + + base = str(request.base_url).rstrip("/") + success_url = settings.stripe_success_url or f"{base}/api/billing/success" + cancel_url = settings.stripe_cancel_url or f"{base}/pricing" + + trial_days = plan.trial_days if plan.trial_days > 0 else None + + session_params: dict[str, Any] = { + "customer": customer_id, + "mode": "subscription", + "line_items": [{"price": price_id, "quantity": 1}], + "success_url": success_url + "?session_id={CHECKOUT_SESSION_ID}", + "cancel_url": cancel_url, + "subscription_data": { + "metadata": { + "docuelevate_user_id": owner_id, + "plan_id": body.plan_id, + "billing_cycle": body.billing_cycle, + }, + }, + "metadata": {"docuelevate_user_id": owner_id, "plan_id": body.plan_id}, + "allow_promotion_codes": True, + "billing_address_collection": "auto", + "tax_id_collection": {"enabled": True}, + "automatic_tax": {"enabled": True}, + } + if trial_days: + session_params["subscription_data"]["trial_period_days"] = trial_days + + checkout_session = client.checkout.sessions.create(params=session_params) + + logger.info( + "Created Stripe checkout session %s for user %s plan %s", + checkout_session.id, + owner_id, + body.plan_id, + ) + return {"checkout_url": checkout_session.url, "session_id": checkout_session.id} + + +@router.post("/create-portal-session", summary="Create a Stripe Customer Portal session") +@require_login +async def create_portal_session( + request: Request, + body: PortalSessionBody, + db: Session = Depends(get_db), +) -> dict[str, Any]: + """Create a Stripe Customer Portal session for subscription self-management. + + Raises: + 503: Stripe not configured. + 404: No Stripe customer found for this user. + """ + client = _get_stripe() + if not client: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing is not configured.") + + user = request.session.get("user") or {} + owner_id = get_current_owner_id(request) or user.get("email") or "" + + profile = db.query(UserProfile).filter(UserProfile.user_id == owner_id).first() + if not profile or not profile.stripe_customer_id: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="No billing account found. Please subscribe to a plan first.", + ) + + base = str(request.base_url).rstrip("/") + return_url = body.return_url or f"{base}/subscription" + + portal = client.billing_portal.sessions.create( + params={ + "customer": profile.stripe_customer_id, + "return_url": return_url, + } + ) + + logger.info("Created Stripe portal session for user %s", owner_id) + return {"portal_url": portal.url} + + +@router.post("/webhook", include_in_schema=False) +async def stripe_webhook(request: Request, db: Session = Depends(get_db)) -> dict[str, str]: + """Handle Stripe webhook events. + + Syncs subscription status to UserProfile.subscription_tier. + + Events handled: + + - ``checkout.session.completed`` — activate subscription after payment + - ``customer.subscription.updated`` — sync tier change + - ``customer.subscription.deleted`` — downgrade to free on cancellation + - ``invoice.payment_failed`` — log failed payment + """ + if not settings.stripe_secret_key: + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing not configured.") + + payload = await request.body() + sig_header = request.headers.get("stripe-signature", "") + + try: + if settings.stripe_webhook_secret: + event = stripe.Webhook.construct_event(payload, sig_header, settings.stripe_webhook_secret) + else: + event = stripe.Event.construct_from(json.loads(payload), stripe.api_key) + except stripe.SignatureVerificationError: + logger.warning("[SECURITY] Stripe webhook signature verification failed") + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid webhook signature.") + except Exception as exc: + logger.warning("Failed to parse Stripe webhook: %s", exc) + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid webhook payload.") + + _handle_stripe_event(db, event) + return {"status": "ok"} + + +@router.get("/success", include_in_schema=False) +@require_login +async def billing_success(request: Request) -> Any: + """Show a success page after a completed Stripe Checkout.""" + return _templates.TemplateResponse("billing_success.html", {"request": request}) + + +def _handle_stripe_event(db: Session, event: Any) -> None: + """Dispatch Stripe event to the appropriate handler. + + Args: + db: Database session. + event: Parsed Stripe event object. + """ + etype = event.get("type", "") if isinstance(event, dict) else getattr(event, "type", "") + data_obj = ( + event.get("data", {}).get("object", {}) + if isinstance(event, dict) + else getattr(getattr(event, "data", None), "object", {}) + ) + + if etype == "checkout.session.completed": + _on_checkout_completed(db, data_obj) + elif etype == "customer.subscription.updated": + _on_subscription_updated(db, data_obj) + elif etype == "customer.subscription.deleted": + _on_subscription_deleted(db, data_obj) + elif etype == "invoice.payment_failed": + customer_id = data_obj.get("customer", "") if isinstance(data_obj, dict) else getattr(data_obj, "customer", "") + logger.warning("Stripe invoice payment failed for customer %s", customer_id) + else: + logger.debug("Unhandled Stripe event type: %s", etype) + + +def _resolve_user_id_from_customer(db: Session, customer_id: str) -> str | None: + """Look up the DocuElevate user_id for a Stripe customer_id. + + Args: + db: Database session. + customer_id: Stripe customer ID. + + Returns: + The matching ``UserProfile.user_id``, or ``None`` if not found. + """ + profile = db.query(UserProfile).filter(UserProfile.stripe_customer_id == customer_id).first() + return profile.user_id if profile else None + + +def _resolve_plan_id_from_price(db: Session, price_id: str) -> str | None: + """Map a Stripe price_id to a DocuElevate plan_id via SubscriptionPlan. + + Args: + db: Database session. + price_id: Stripe price ID. + + Returns: + The matching ``SubscriptionPlan.plan_id``, or ``None`` if not found. + """ + plan = ( + db.query(SubscriptionPlan) + .filter( + (SubscriptionPlan.stripe_price_id_monthly == price_id) + | (SubscriptionPlan.stripe_price_id_yearly == price_id) + ) + .first() + ) + return plan.plan_id if plan else None + + +def _on_checkout_completed(db: Session, data: Any) -> None: + """Activate a subscription after a successful checkout. + + Args: + db: Database session. + data: Stripe ``checkout.session`` object. + """ + meta = data.get("metadata") or {} if isinstance(data, dict) else getattr(data, "metadata", {}) or {} + user_id = meta.get("docuelevate_user_id") if isinstance(meta, dict) else getattr(meta, "docuelevate_user_id", None) + plan_id = meta.get("plan_id") if isinstance(meta, dict) else getattr(meta, "plan_id", None) + billing_cycle = ( + meta.get("billing_cycle", "monthly") if isinstance(meta, dict) else getattr(meta, "billing_cycle", "monthly") + ) + if not user_id: + return + + profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first() + if profile and plan_id: + profile.subscription_tier = plan_id + profile.subscription_billing_cycle = billing_cycle + profile.subscription_period_start = datetime.now(tz=timezone.utc) + customer_id = data.get("customer", "") if isinstance(data, dict) else getattr(data, "customer", "") + if customer_id: + profile.stripe_customer_id = customer_id + db.commit() + logger.info("Activated plan %s/%s for user %s after checkout", plan_id, billing_cycle, user_id) + + +def _on_subscription_updated(db: Session, data: Any) -> None: + """Sync tier change when a subscription is updated. + + Args: + db: Database session. + data: Stripe ``customer.subscription`` object. + """ + customer_id = data.get("customer", "") if isinstance(data, dict) else getattr(data, "customer", "") + user_id = _resolve_user_id_from_customer(db, customer_id) + if not user_id: + return + + items_data = data.get("items") or {} if isinstance(data, dict) else getattr(data, "items", None) or {} + items = items_data.get("data") or [] if isinstance(items_data, dict) else getattr(items_data, "data", []) or [] + if not items: + return + + first_item = items[0] + price_obj = ( + first_item.get("price") or {} if isinstance(first_item, dict) else getattr(first_item, "price", {}) or {} + ) + price_id = price_obj.get("id") if isinstance(price_obj, dict) else getattr(price_obj, "id", None) + if not price_id: + return + + plan_id = _resolve_plan_id_from_price(db, price_id) + if not plan_id: + logger.warning("Unknown Stripe price_id %s on subscription.updated", price_id) + return + + recurring = ( + price_obj.get("recurring", {}) if isinstance(price_obj, dict) else getattr(price_obj, "recurring", {}) or {} + ) + interval = ( + recurring.get("interval", "month") if isinstance(recurring, dict) else getattr(recurring, "interval", "month") + ) + billing_cycle = "yearly" if interval == "year" else "monthly" + + profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first() + if profile: + profile.subscription_tier = plan_id + profile.subscription_billing_cycle = billing_cycle + db.commit() + logger.info("Updated subscription to %s/%s for user %s", plan_id, billing_cycle, user_id) + + +def _on_subscription_deleted(db: Session, data: Any) -> None: + """Downgrade user to free tier after subscription cancellation. + + Args: + db: Database session. + data: Stripe ``customer.subscription`` object. + """ + customer_id = data.get("customer", "") if isinstance(data, dict) else getattr(data, "customer", "") + user_id = _resolve_user_id_from_customer(db, customer_id) + if not user_id: + return + + profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first() + if profile: + profile.subscription_tier = "free" + profile.subscription_billing_cycle = "monthly" + db.commit() + logger.info("Downgraded user %s to free tier after subscription cancellation", user_id) diff --git a/app/api/local_auth.py b/app/api/local_auth.py new file mode 100644 index 00000000..a74d0a33 --- /dev/null +++ b/app/api/local_auth.py @@ -0,0 +1,315 @@ +"""Local user authentication API — signup, email verification, password reset. + +Provides the REST endpoints and page routes for the self-registration flow: + +- GET /signup — signup page (HTML) +- POST /api/auth/signup — create account + send verification email +- GET /verify-email — activate account from email link (redirect) +- GET /verify-email-sent — confirmation landing page (HTML) +- POST /api/auth/resend-verification — re-send verification email +- POST /api/auth/request-password-reset — start password reset +- POST /api/auth/reset-password — set new password using token +- GET /reset-password — password reset form page (HTML) +""" + +import logging +import pathlib +from datetime import datetime, timezone +from typing import Annotated, Any + +from fastapi import APIRouter, Depends, HTTPException, Request, status +from fastapi.templating import Jinja2Templates +from pydantic import BaseModel, Field +from sqlalchemy.orm import Session +from starlette.responses import RedirectResponse + +from app.config import settings +from app.database import get_db +from app.models import LocalUser, UserProfile +from app.utils.local_auth import ( + build_session_user, + generate_token, + hash_password, + is_token_expired, + send_password_reset_email, + send_verification_email, +) + +logger = logging.getLogger(__name__) +router = APIRouter(tags=["local-auth"]) + +_templates_dir = pathlib.Path(__file__).parents[2] / "frontend" / "templates" +templates = Jinja2Templates(directory=str(_templates_dir)) + +DbSession = Annotated[Session, Depends(get_db)] + + +# --------------------------------------------------------------------------- +# Pydantic schemas +# --------------------------------------------------------------------------- + + +class SignupBody(BaseModel): + """Body for the signup endpoint.""" + + email: str = Field(..., max_length=255) + username: str = Field(..., min_length=3, max_length=64, pattern=r"^[a-zA-Z0-9_-]+$") + display_name: str | None = Field(default=None, max_length=255) + password: str = Field(..., min_length=8, max_length=128) + password_confirm: str + + +class ResendVerificationBody(BaseModel): + """Body for the resend-verification endpoint.""" + + email: str + + +class PasswordResetRequestBody(BaseModel): + """Body for the request-password-reset endpoint.""" + + email: str + + +class PasswordResetBody(BaseModel): + """Body for the reset-password endpoint.""" + + token: str + new_password: str = Field(..., min_length=8, max_length=128) + new_password_confirm: str + + +# --------------------------------------------------------------------------- +# Page routes (return HTML) +# --------------------------------------------------------------------------- + + +@router.get("/signup", include_in_schema=False) +async def signup_page(request: Request) -> Any: + """Render the signup page, or redirect to login when signup is disabled.""" + if not settings.allow_local_signup: + return RedirectResponse(url="/login?error=Registration+is+not+enabled", status_code=302) + return templates.TemplateResponse( + "signup.html", + { + "request": request, + "csrf_token": getattr(request.state, "csrf_token", ""), + "app_version": settings.version, + }, + ) + + +@router.get("/verify-email-sent", include_in_schema=False) +async def verify_email_sent_page(request: Request) -> Any: + """Render the verify-email-sent confirmation page.""" + return templates.TemplateResponse("verify_email_sent.html", {"request": request}) + + +@router.get("/reset-password", include_in_schema=False) +async def reset_password_page(request: Request) -> Any: + """Render the password reset form page.""" + token = request.query_params.get("token", "") + return templates.TemplateResponse( + "password_reset_form.html", + { + "request": request, + "token": token, + "csrf_token": getattr(request.state, "csrf_token", ""), + "app_version": settings.version, + }, + ) + + +# --------------------------------------------------------------------------- +# API endpoints (return JSON or redirect) +# --------------------------------------------------------------------------- + + +@router.post("/api/auth/signup", status_code=status.HTTP_201_CREATED) +async def signup(request: Request, body: SignupBody, db: DbSession) -> dict[str, str]: + """Create a new local user account and send a verification email. + + The account is inactive until the user clicks the email link. + + Raises: + 403: Local signup is disabled. + 503: SMTP is not configured. + 422: Passwords do not match. + 409: Email or username already registered. + """ + if not settings.allow_local_signup: + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Registration is not enabled.") + if not settings.email_host: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Email (SMTP) must be configured before local signup can be enabled.", + ) + if body.password != body.password_confirm: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail="Passwords do not match.") + + if db.query(LocalUser).filter(LocalUser.email == body.email).first(): + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Email already registered.") + if db.query(LocalUser).filter(LocalUser.username == body.username).first(): + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Username already taken.") + + token = generate_token() + user = LocalUser( + email=body.email, + username=body.username, + display_name=body.display_name, + hashed_password=hash_password(body.password), + is_active=False, + email_verification_token=token, + email_verification_sent_at=datetime.now(tz=timezone.utc), + ) + db.add(user) + + profile = UserProfile( + user_id=body.email, + display_name=body.display_name or body.username, + ) + db.add(profile) + + try: + db.commit() + except Exception: + db.rollback() + raise + + base_url = str(request.base_url).rstrip("/") + try: + send_verification_email(body.email, body.username, token, base_url) + except Exception as exc: + # Clean up orphan records — don't leave an unverifiable account + try: + db.delete(user) + db.delete(profile) + db.commit() + except Exception: + db.rollback() + logger.warning("Signup email failed for %s: %s", body.email, exc) + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=("Failed to send verification email. Please check that SMTP is correctly configured and try again."), + ) from exc + + logger.info("New local user registered: %s", body.email) + return {"message": "Verification email sent. Please check your inbox."} + + +@router.get("/verify-email", include_in_schema=False) +async def verify_email(request: Request, db: DbSession) -> Any: + """Activate a local user account from the email verification link. + + Redirects to the login page on failure, or to onboarding/upload on success. + """ + token = request.query_params.get("token", "") + user = db.query(LocalUser).filter(LocalUser.email_verification_token == token).first() + + if not user: + return RedirectResponse( + url="/login?error=Invalid+or+expired+verification+link", + status_code=302, + ) + if is_token_expired(user.email_verification_sent_at): + return RedirectResponse( + url="/login?error=Verification+link+has+expired.+Please+request+a+new+one", + status_code=302, + ) + + user.is_active = True + user.email_verification_token = None + user.email_verification_sent_at = None + + # Ensure profile exists + if not db.query(UserProfile).filter(UserProfile.user_id == user.email).first(): + db.add(UserProfile(user_id=user.email, display_name=user.display_name or user.username)) + + db.commit() + + request.session["user"] = build_session_user(user) + logger.info("[SECURITY] EMAIL_VERIFIED user=%s", user.email) + + profile = db.query(UserProfile).filter(UserProfile.user_id == user.email).first() + if profile and not profile.onboarding_completed: + post_onboarding = request.session.pop("redirect_after_login", "/upload") + request.session["post_onboarding_redirect"] = post_onboarding + return RedirectResponse(url="/onboarding", status_code=302) + return RedirectResponse(url="/upload", status_code=302) + + +@router.post("/api/auth/resend-verification") +async def resend_verification(request: Request, body: ResendVerificationBody, db: DbSession) -> dict[str, str]: + """Re-send the verification email for a pending account. + + Always returns 200 to avoid leaking whether an email is registered. + """ + user = db.query(LocalUser).filter(LocalUser.email == body.email).first() + if not user or user.is_active: + return {"message": "Verification email resent if account exists."} + + token = generate_token() + user.email_verification_token = token + user.email_verification_sent_at = datetime.now(tz=timezone.utc) + db.commit() + + base_url = str(request.base_url).rstrip("/") + try: + send_verification_email(user.email, user.username, token, base_url) + except Exception as exc: + logger.warning("Failed to resend verification email to %s: %s", user.email, exc) + + return {"message": "Verification email resent if account exists."} + + +@router.post("/api/auth/request-password-reset") +async def request_password_reset(request: Request, body: PasswordResetRequestBody, db: DbSession) -> dict[str, str]: + """Send a password reset email. + + Always returns 200 to avoid leaking whether an email is registered. + """ + user = db.query(LocalUser).filter(LocalUser.email == body.email).first() + if not user: + return {"message": "Password reset email sent if account exists."} + + token = generate_token() + user.password_reset_token = token + user.password_reset_sent_at = datetime.now(tz=timezone.utc) + db.commit() + + base_url = str(request.base_url).rstrip("/") + try: + send_password_reset_email(user.email, user.username, token, base_url) + except Exception as exc: + logger.warning("Failed to send password reset email to %s: %s", user.email, exc) + + return {"message": "Password reset email sent if account exists."} + + +@router.post("/api/auth/reset-password") +async def reset_password(body: PasswordResetBody, db: DbSession) -> dict[str, str]: + """Set a new password using a valid reset token. + + Raises: + 400: Token is invalid or expired. + 422: Passwords do not match. + """ + user = db.query(LocalUser).filter(LocalUser.password_reset_token == body.token).first() + if not user or is_token_expired(user.password_reset_sent_at): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Invalid or expired reset token.", + ) + if body.new_password != body.new_password_confirm: + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail="Passwords do not match.", + ) + + user.hashed_password = hash_password(body.new_password) + user.password_reset_token = None + user.password_reset_sent_at = None + db.commit() + + logger.info("[SECURITY] PASSWORD_RESET_SUCCESS user=%s", user.email) + return {"message": "Password updated successfully."} diff --git a/app/auth.py b/app/auth.py index b5d7d6dd..5ecbe584 100644 --- a/app/auth.py +++ b/app/auth.py @@ -72,7 +72,7 @@ def get_gravatar_url(email): async def login(request: Request): - """Show login page with appropriate authentication options""" + """Show login page with appropriate authentication options.""" return templates.TemplateResponse( "login.html", { @@ -81,8 +81,9 @@ async def login(request: Request): "message": request.query_params.get("message"), "show_oauth": OAUTH_CONFIGURED, "oauth_provider_name": OAUTH_PROVIDER_NAME, - "app_version": settings.version, # Changed from app_version to version + "app_version": settings.version, "csrf_token": getattr(request.state, "csrf_token", ""), + "allow_signup": settings.allow_local_signup, }, ) @@ -188,14 +189,45 @@ async def oauth_callback(request: Request, db: Session = Depends(get_db)): return RedirectResponse(url=f"/login?error=Authentication+failed:+{str(e)}", status_code=status.HTTP_302_FOUND) -async def auth(request: Request): - """Handle local username/password authentication""" +async def auth(request: Request, db: Session = Depends(get_db)): + """Handle local username/password authentication. + + Checks LocalUser accounts first, then falls back to admin credentials. + """ form_data = await request.form() username = form_data.get("username") password = form_data.get("password") + # --- LocalUser check --- + from app.models import LocalUser as _LocalUser + from app.models import UserProfile as _UserProfile + from app.utils.local_auth import build_session_user as _build_session_user + from app.utils.local_auth import verify_password as _verify_password + + local_user = db.query(_LocalUser).filter((_LocalUser.username == username) | (_LocalUser.email == username)).first() + if local_user is not None: + if not _verify_password(password or "", local_user.hashed_password): + logger.warning("[SECURITY] LOCAL_LOGIN_FAILURE user=%s", username) + return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302) + if not local_user.is_active: + logger.warning("[SECURITY] LOCAL_LOGIN_UNVERIFIED user=%s", username) + return RedirectResponse( + url="/login?error=Please+verify+your+email+address+before+logging+in", + status_code=302, + ) + user_data = _build_session_user(local_user) + request.session["user"] = user_data + logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", local_user.email) + profile = db.query(_UserProfile).filter(_UserProfile.user_id == local_user.email).first() + if profile and not profile.onboarding_completed: + post_onboarding = request.session.pop("redirect_after_login", "/upload") + request.session["post_onboarding_redirect"] = post_onboarding + return RedirectResponse(url="/onboarding", status_code=302) + redirect_url = request.session.pop("redirect_after_login", "/upload") + return RedirectResponse(url=redirect_url, status_code=302) + + # --- Admin credentials fallback --- if username == settings.admin_username and password == settings.admin_password: - # Create user session request.session["user"] = { "id": "admin", "name": "Administrator", @@ -204,12 +236,11 @@ async def auth(request: Request): "picture": "/static/images/default-avatar.svg", "is_admin": True, } - logger.info(f"[SECURITY] LOCAL_LOGIN_SUCCESS user={username}") - # Redirect to original destination or default + logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", username) redirect_url = request.session.pop("redirect_after_login", "/upload") return RedirectResponse(url=redirect_url, status_code=302) else: - logger.warning(f"[SECURITY] LOCAL_LOGIN_FAILURE user={username}") + logger.warning("[SECURITY] LOCAL_LOGIN_FAILURE user=%s", username) return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302) diff --git a/app/config.py b/app/config.py index d0322e82..c59c31c4 100644 --- a/app/config.py +++ b/app/config.py @@ -172,6 +172,23 @@ class Settings(BaseSettings): authentik_config_url: Optional[str] = None oauth_provider_name: Optional[str] = None # Name to display for the OAuth provider + # Local user signup + allow_local_signup: bool = Field( + default=False, + description=( + "Allow users to self-register with email and password. " + "Requires email (SMTP) to be configured for verification emails. " + "Default: False (registration disabled, admin creates users)." + ), + ) + + # Stripe billing + stripe_secret_key: Optional[str] = None + stripe_publishable_key: Optional[str] = None + stripe_webhook_secret: Optional[str] = None + stripe_success_url: Optional[str] = None # e.g. https://app.example.com/billing/success + stripe_cancel_url: Optional[str] = None # e.g. https://app.example.com/pricing + # IMAP 1 imap1_host: Optional[str] = None imap1_port: Optional[int] = 993 diff --git a/app/main.py b/app/main.py index 5f39ae2c..4ed9c081 100644 --- a/app/main.py +++ b/app/main.py @@ -16,6 +16,7 @@ from starlette.middleware.trustedhost import TrustedHostMiddleware from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware from app.api import router as api_router +from app.api.local_auth import router as local_auth_router from app.auth import router as auth_router from app.config import settings from app.database import init_db @@ -246,4 +247,5 @@ def test_500(): app.include_router(frontend_router) app.include_router(files_router) # Explicitly include the files router app.include_router(auth_router) +app.include_router(local_auth_router) app.include_router(api_router, prefix="/api") diff --git a/app/models.py b/app/models.py index 1c38ea84..a81d93ba 100644 --- a/app/models.py +++ b/app/models.py @@ -172,6 +172,31 @@ class WebhookConfig(Base): updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now()) +class LocalUser(Base): + """A locally-registered user authenticated by email and bcrypt password. + + Created during the self-registration flow when ``allow_local_signup`` is + enabled. The account is inactive (``is_active=False``) until the user + clicks the verification link sent to their email address. + """ + + __tablename__ = "local_users" + + id = Column(Integer, primary_key=True, index=True) + email = Column(String(255), unique=True, nullable=False, index=True) + username = Column(String(64), unique=True, nullable=False, index=True) + display_name = Column(String(255), nullable=True) + hashed_password = Column(String(255), nullable=False) + is_active = Column(Boolean, nullable=False, default=False, server_default="0") + is_admin = Column(Boolean, nullable=False, default=False, server_default="0") + email_verification_token = Column(String(128), nullable=True) + email_verification_sent_at = Column(DateTime(timezone=True), nullable=True) + password_reset_token = Column(String(128), nullable=True) + password_reset_sent_at = Column(DateTime(timezone=True), nullable=True) + created_at = Column(DateTime(timezone=True), server_default=func.now()) + updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now()) + + class UserProfile(Base): """Per-user profile for admin-managed settings in multi-user mode. @@ -214,6 +239,7 @@ class UserProfile(Base): onboarding_completed_at = Column(DateTime(timezone=True), nullable=True) contact_email = Column(String(255), nullable=True) preferred_destination = Column(String(50), nullable=True) + stripe_customer_id = Column(String(64), nullable=True) created_at = Column(DateTime(timezone=True), server_default=func.now()) updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now()) @@ -262,6 +288,8 @@ class SubscriptionPlan(Base): sort_order = Column(Integer, nullable=False, default=0) features = Column(Text, nullable=True) # JSON-encoded list[str] api_access = Column(Boolean, nullable=False, default=False) + stripe_price_id_monthly = Column(String(64), nullable=True) + stripe_price_id_yearly = Column(String(64), nullable=True) created_at = Column(DateTime(timezone=True), server_default=func.now()) updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now()) diff --git a/app/utils/local_auth.py b/app/utils/local_auth.py new file mode 100644 index 00000000..ce52de2f --- /dev/null +++ b/app/utils/local_auth.py @@ -0,0 +1,187 @@ +"""Utilities for local (email/password) user authentication. + +Provides password hashing (bcrypt), secure token generation, and +synchronous SMTP email helpers for account verification and password +reset flows. No external dependencies beyond bcrypt (already in +requirements.txt) and Python stdlib. +""" + +import logging +import secrets +import smtplib +import socket +from datetime import datetime, timedelta, timezone +from email.mime.multipart import MIMEMultipart +from email.mime.text import MIMEText + +import bcrypt + +from app.config import settings + +logger = logging.getLogger(__name__) + +TOKEN_BYTES = 32 # 256 bits of entropy +TOKEN_EXPIRY_HOURS = 24 # verification + reset tokens expire after 24 h + + +def hash_password(plain: str) -> str: + """Return a bcrypt hash of *plain*. Stores result as a UTF-8 string.""" + return bcrypt.hashpw(plain.encode("utf-8"), bcrypt.gensalt(rounds=12)).decode("utf-8") + + +def verify_password(plain: str, hashed: str) -> bool: + """Return True when *plain* matches the stored bcrypt *hashed* string.""" + try: + return bcrypt.checkpw(plain.encode("utf-8"), hashed.encode("utf-8")) + except Exception: + return False + + +def generate_token() -> str: + """Return a 256-bit URL-safe random token string.""" + return secrets.token_urlsafe(TOKEN_BYTES) + + +def is_token_expired(sent_at: datetime | None) -> bool: + """Return True when *sent_at* is None or older than TOKEN_EXPIRY_HOURS.""" + if sent_at is None: + return True + return datetime.now(tz=timezone.utc) > sent_at.replace(tzinfo=timezone.utc) + timedelta(hours=TOKEN_EXPIRY_HOURS) + + +def _smtp_send(subject: str, html_body: str, plain_body: str, recipient: str) -> None: + """Send an HTML email via the configured SMTP server. + + Args: + subject: Email subject line. + html_body: HTML version of the email body. + plain_body: Plain-text version of the email body. + recipient: Recipient email address. + + Raises: + RuntimeError: When SMTP is not configured or sending fails. + """ + if not settings.email_host: + raise RuntimeError("SMTP is not configured (EMAIL_HOST missing). Cannot send email.") + + sender = settings.email_sender or settings.email_username or "noreply@docuelevate.local" + + msg = MIMEMultipart("alternative") + msg["Subject"] = subject + msg["From"] = sender + msg["To"] = recipient + msg.attach(MIMEText(plain_body, "plain", "utf-8")) + msg.attach(MIMEText(html_body, "html", "utf-8")) + + try: + socket.gethostbyname(settings.email_host) + except socket.gaierror as exc: + raise RuntimeError(f"Cannot resolve SMTP host {settings.email_host!r}: {exc}") from exc + + with smtplib.SMTP(settings.email_host, settings.email_port or 587, timeout=30) as server: + if settings.email_use_tls: + server.starttls() + if settings.email_username and settings.email_password: + server.login(settings.email_username, settings.email_password) + server.send_message(msg) + + logger.info("Sent %r to %s", subject, recipient) + + +def send_verification_email(email: str, username: str, token: str, base_url: str) -> None: + """Send a double opt-in verification email to *email*. + + Args: + email: Recipient email address. + username: The user's chosen username (used in greeting). + token: The verification token to embed in the link. + base_url: The base URL of the application (e.g. https://app.example.com). + """ + verify_url = f"{base_url}/verify-email?token={token}" + subject = "Verify your DocuElevate account" + html_body = f""" + + + +
+

Welcome to DocuElevate, {username}!

+

Thanks for signing up. Please confirm your email address to activate your account.

+
+ + Confirm my email address + +
+

This link expires in 24 hours. If you did not create an account, you can safely ignore this email.

+
+

DocuElevate · Intelligent Document Processing

+
+ +""" + plain_body = ( + f"Welcome to DocuElevate, {username}!\n\n" + f"Please verify your email address by visiting:\n{verify_url}\n\n" + "This link expires in 24 hours." + ) + _smtp_send(subject, html_body, plain_body, email) + + +def send_password_reset_email(email: str, username: str, token: str, base_url: str) -> None: + """Send a password reset email to *email*. + + Args: + email: Recipient email address. + username: The user's username (used in greeting). + token: The password reset token to embed in the link. + base_url: The base URL of the application. + """ + reset_url = f"{base_url}/reset-password?token={token}" + subject = "Reset your DocuElevate password" + html_body = f""" + + + +
+

Password Reset

+

Hi {username}, you requested a password reset for your DocuElevate account.

+
+ + Reset my password + +
+

This link expires in 24 hours. If you did not request a password reset, you can safely ignore this email.

+
+

DocuElevate · Intelligent Document Processing

+
+ +""" + plain_body = ( + f"Hi {username},\n\n" + f"You requested a password reset. Visit the link below:\n{reset_url}\n\n" + "This link expires in 24 hours. If you did not request this, ignore this email." + ) + _smtp_send(subject, html_body, plain_body, email) + + +def build_session_user(user: object) -> dict: + """Build the session user dict for a LocalUser, matching the OAuth session format. + + Args: + user: A ``LocalUser`` ORM instance. + + Returns: + Dict suitable for storing in ``request.session["user"]``. + """ + from app.auth import get_gravatar_url + + return { + "sub": user.email, # type: ignore[attr-defined] + "id": user.email, # type: ignore[attr-defined] + "email": user.email, # type: ignore[attr-defined] + "preferred_username": user.username, # type: ignore[attr-defined] + "name": user.display_name or user.username, # type: ignore[attr-defined] + "picture": get_gravatar_url(user.email), # type: ignore[attr-defined] + "is_admin": bool(user.is_admin), # type: ignore[attr-defined] + "auth_method": "local", + } diff --git a/docs/BillingSetup.md b/docs/BillingSetup.md new file mode 100644 index 00000000..4a4462a6 --- /dev/null +++ b/docs/BillingSetup.md @@ -0,0 +1,172 @@ +# Billing Setup Guide + +This guide covers how to configure Stripe billing and local user sign-up in DocuElevate. + +## Table of Contents + +- [Local User Sign-up](#local-user-sign-up) +- [Stripe Billing Integration](#stripe-billing-integration) + - [Prerequisites](#prerequisites) + - [Configuration](#configuration) + - [Setting Up Plans](#setting-up-plans) + - [Webhook Configuration](#webhook-configuration) + - [Billing Flows](#billing-flows) +- [Compliance Notes](#compliance-notes) + +--- + +## Local User Sign-up + +By default, user accounts are created by an administrator. To allow users to self-register with an email address and password, set `ALLOW_LOCAL_SIGNUP=true`. + +> **Note:** SMTP must be configured before enabling local sign-up. New accounts require email verification before they can log in. + +### Configuration + +```bash +ALLOW_LOCAL_SIGNUP=true + +# SMTP (required for verification emails) +EMAIL_HOST=smtp.example.com +EMAIL_PORT=587 +EMAIL_USERNAME=noreply@example.com +EMAIL_PASSWORD=yourpassword +EMAIL_USE_TLS=true +EMAIL_SENDER=DocuElevate +``` + +### Sign-up Flow + +1. User visits `/signup` and fills out the registration form. +2. DocuElevate sends a verification email with a 24-hour token link. +3. User clicks the link — their account is activated and they are signed in. +4. First-time users are redirected to the onboarding wizard. + +### Password Reset Flow + +1. User clicks "Forgot password?" on the login page. +2. User enters their email address. +3. DocuElevate sends a password reset email with a 24-hour token link. +4. User clicks the link, enters a new password, and is redirected to sign in. + +### Security + +- Passwords are hashed with bcrypt (12 rounds). +- Verification and reset tokens are 256-bit URL-safe random strings. +- All tokens expire after 24 hours. +- Sign-up and login endpoints return generic error messages to prevent user enumeration. + +--- + +## Stripe Billing Integration + +DocuElevate integrates with [Stripe](https://stripe.com) to handle subscription payments. Stripe acts as a data processor under a Data Processing Agreement (DPA) and is SOC 2 Type II certified. + +### Prerequisites + +- A Stripe account (sign up at [stripe.com](https://stripe.com)) +- Products and prices created in the Stripe Dashboard for each paid plan +- A publicly reachable webhook endpoint (or use [Stripe CLI](https://stripe.com/docs/stripe-cli) for local testing) + +### Configuration + +```bash +STRIPE_SECRET_KEY=sk_live_... # Your Stripe secret key +STRIPE_PUBLISHABLE_KEY=pk_live_... # Your Stripe publishable key (for frontend) +STRIPE_WEBHOOK_SECRET=whsec_... # Webhook signing secret +STRIPE_SUCCESS_URL=https://app.example.com/api/billing/success # Optional override +STRIPE_CANCEL_URL=https://app.example.com/pricing # Optional override +``` + +> **Security:** Never commit your Stripe secret key. Store it in your environment or secrets manager. + +### Setting Up Plans + +After starting DocuElevate, go to **Admin → Plans** to configure each plan: + +1. Open the **Plan Designer** for a paid tier (e.g. Starter, Professional). +2. Enter the **Stripe Price ID (monthly)** from your Stripe Dashboard (e.g. `price_1OtAbc...`). +3. Optionally enter the **Stripe Price ID (yearly)** for annual billing. +4. Save the plan. + +Stripe Price IDs look like `price_1OtAbcDefGhIjKlMnOpQrSt`. Find them in **Products** in your Stripe Dashboard. + +### Webhook Configuration + +Stripe webhooks allow DocuElevate to sync subscription status in real time. + +#### Stripe Dashboard setup + +1. Go to **Developers → Webhooks** in the Stripe Dashboard. +2. Click **Add endpoint**. +3. Set the endpoint URL to: `https://your-app-domain.com/api/billing/webhook` +4. Select the following events: + - `checkout.session.completed` + - `customer.subscription.updated` + - `customer.subscription.deleted` + - `invoice.payment_failed` +5. Copy the **Signing secret** and set `STRIPE_WEBHOOK_SECRET` in your environment. + +#### Local testing with Stripe CLI + +```bash +# Install Stripe CLI and log in +stripe login + +# Forward webhooks to your local server +stripe listen --forward-to http://localhost:8000/api/billing/webhook + +# Trigger a test event +stripe trigger checkout.session.completed +``` + +### Billing Flows + +#### Subscribe to a plan + +1. User visits `/pricing`. +2. User clicks the **CTA button** on a paid plan. +3. DocuElevate calls `POST /api/billing/create-checkout-session`. +4. User is redirected to Stripe Checkout. +5. After payment, Stripe fires `checkout.session.completed`. +6. DocuElevate webhook handler activates the subscription tier. +7. User is redirected to `/api/billing/success`. + +#### Manage or cancel subscription + +1. User visits their account settings. +2. DocuElevate calls `POST /api/billing/create-portal-session`. +3. User is redirected to the Stripe Customer Portal. +4. User can update payment method, upgrade, downgrade, or cancel. +5. Stripe fires `customer.subscription.updated` or `customer.subscription.deleted`. +6. DocuElevate webhook handler syncs the change. + +#### Cancellation + +When a subscription is cancelled, Stripe fires `customer.subscription.deleted` and DocuElevate automatically downgrades the user to the free tier. + +--- + +## Compliance Notes + +| Topic | Details | +|-------|---------| +| **GDPR** | Stripe acts as a data processor. A Data Processing Agreement (DPA) is available in the Stripe Dashboard. Stripe supports EU data residency. | +| **SOC 2** | Stripe is SOC 2 Type II certified. | +| **EU VAT** | Configure [Stripe Tax](https://stripe.com/tax) in the Stripe Dashboard for automatic VAT collection. | +| **PCI DSS** | Card data is handled entirely by Stripe. DocuElevate never sees or stores card details. | + +--- + +## Environment Variable Reference + +| Variable | Type | Default | Description | +|----------|------|---------|-------------| +| `ALLOW_LOCAL_SIGNUP` | bool | `false` | Allow users to self-register with email/password | +| `STRIPE_SECRET_KEY` | string | — | Stripe API secret key | +| `STRIPE_PUBLISHABLE_KEY` | string | — | Stripe API publishable key | +| `STRIPE_WEBHOOK_SECRET` | string | — | Webhook signing secret from Stripe Dashboard | +| `STRIPE_SUCCESS_URL` | string | — | Override redirect URL after successful checkout | +| `STRIPE_CANCEL_URL` | string | — | Override redirect URL when checkout is cancelled | + +See [ConfigurationGuide.md](./ConfigurationGuide.md) for the full environment variable reference. diff --git a/frontend/templates/billing_success.html b/frontend/templates/billing_success.html new file mode 100644 index 00000000..4d3e0862 --- /dev/null +++ b/frontend/templates/billing_success.html @@ -0,0 +1,45 @@ + + + + + + DocuElevate - Subscription Activated + + + + +
+
+ DocuElevate Logo +
+ +
+
+ +
+
+ +

You're all set!

+

+ Your subscription has been activated. Thank you for choosing DocuElevate! +

+ + +
+ + diff --git a/frontend/templates/login.html b/frontend/templates/login.html index a2361d4f..5ead49b9 100644 --- a/frontend/templates/login.html +++ b/frontend/templates/login.html @@ -78,6 +78,15 @@ Return to Home + + {% if allow_signup %} +
+ Don't have an account? + + Create account + +
+ {% endif %}
DocuElevate {{ app_version|default('', true) }} diff --git a/frontend/templates/password_reset_form.html b/frontend/templates/password_reset_form.html new file mode 100644 index 00000000..217ec2b6 --- /dev/null +++ b/frontend/templates/password_reset_form.html @@ -0,0 +1,131 @@ + + + + + + DocuElevate - Reset Password + + + + + +
+
+ DocuElevate Logo +
+ +

Set a new password

+

Enter your new password below.

+ +
+
+
+
+ +
+
+

Password updated successfully!

+ Sign in +
+ +
+ + +
+ + +

Minimum 8 characters.

+
+ +
+ + +
+ + +
+
+ + +
+
+ DocuElevate {{ app_version|default('', true) }} +
+ + diff --git a/frontend/templates/pricing.html b/frontend/templates/pricing.html index f0468883..4764e64f 100644 --- a/frontend/templates/pricing.html +++ b/frontend/templates/pricing.html @@ -117,6 +117,13 @@ {{ tier.cta }} + {% elif tier.stripe_price_id_monthly or tier.stripe_price_id_yearly %} + {% elif tier.highlight %}
{% endblock %} + +{% block extra_scripts %} + +{% endblock %} diff --git a/frontend/templates/signup.html b/frontend/templates/signup.html new file mode 100644 index 00000000..ff2ee9bd --- /dev/null +++ b/frontend/templates/signup.html @@ -0,0 +1,156 @@ + + + + + + DocuElevate - Create Account + + + + + +
+
+ DocuElevate Logo +
+ +

Create your account

+

Already have an account? + Sign in +

+ +
+ + +
+
+ + +
+ +
+ + +

3–64 characters. Letters, numbers, hyphens and underscores only.

+
+ +
+ + +
+ +
+ + +

Minimum 8 characters.

+
+ +
+ + +
+ + +
+
+ + +
+
+ DocuElevate {{ app_version|default('', true) }} +
+ + diff --git a/frontend/templates/verify_email_sent.html b/frontend/templates/verify_email_sent.html new file mode 100644 index 00000000..857824d6 --- /dev/null +++ b/frontend/templates/verify_email_sent.html @@ -0,0 +1,108 @@ + + + + + + DocuElevate - Verify Your Email + + + + + +
+
+ DocuElevate Logo +
+ +
+
+ +
+
+ +

Check your inbox

+

+ We've sent you a verification email. Please click the link in the email to activate your account. +

+

+ The link expires in 24 hours. If you don't see the email, check your spam folder. +

+ +
+

Didn't receive it?

+ +
+

+
+ + +
+
+ + +
+ +
+
+ + +
+ + diff --git a/migrations/versions/018_add_local_users_and_billing.py b/migrations/versions/018_add_local_users_and_billing.py new file mode 100644 index 00000000..b226a0d3 --- /dev/null +++ b/migrations/versions/018_add_local_users_and_billing.py @@ -0,0 +1,49 @@ +"""Add local_users table and billing columns + +Revision ID: 018_add_local_users_and_billing +Revises: 017_add_onboarding_fields +Create Date: 2026-03-09 +""" + +from typing import Union + +import sqlalchemy as sa +from alembic import op + +revision: str = "018_add_local_users_and_billing" +down_revision: Union[str, None] = "017_add_onboarding_fields" +depends_on: Union[str, None] = None + + +def upgrade() -> None: + """Create local_users table and add billing columns.""" + op.create_table( + "local_users", + sa.Column("id", sa.Integer(), nullable=False), + sa.Column("email", sa.String(255), nullable=False), + sa.Column("username", sa.String(64), nullable=False), + sa.Column("display_name", sa.String(255), nullable=True), + sa.Column("hashed_password", sa.String(255), nullable=False), + sa.Column("is_active", sa.Boolean(), nullable=False, server_default="0"), + sa.Column("is_admin", sa.Boolean(), nullable=False, server_default="0"), + sa.Column("email_verification_token", sa.String(128), nullable=True), + sa.Column("email_verification_sent_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("password_reset_token", sa.String(128), nullable=True), + sa.Column("password_reset_sent_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now()), + sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), onupdate=sa.func.now()), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("email"), + sa.UniqueConstraint("username"), + ) + op.add_column("user_profiles", sa.Column("stripe_customer_id", sa.String(64), nullable=True)) + op.add_column("subscription_plans", sa.Column("stripe_price_id_monthly", sa.String(64), nullable=True)) + op.add_column("subscription_plans", sa.Column("stripe_price_id_yearly", sa.String(64), nullable=True)) + + +def downgrade() -> None: + """Reverse the migration.""" + op.drop_column("subscription_plans", "stripe_price_id_yearly") + op.drop_column("subscription_plans", "stripe_price_id_monthly") + op.drop_column("user_profiles", "stripe_customer_id") + op.drop_table("local_users") diff --git a/requirements.txt b/requirements.txt index 3d92dda5..3eb8e52d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -44,3 +44,4 @@ pytesseract>=0.3.10 # Python wrapper for Tesseract OCR pdf2image>=1.17.0 # Convert PDF pages to images (used by Tesseract and EasyOCR providers) ocrmypdf>=16.0.0,<18.0.0 # Post-processing: embeds searchable text layers into PDFs via Tesseract meilisearch>=0.31.0 # Full-text search engine client +stripe>=7.0.0,<15.0.0 # Stripe billing SDK (MIT license) diff --git a/tests/test_auth.py b/tests/test_auth.py index bb833303..f781cbec 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -729,9 +729,16 @@ class TestEnsureUserProfile: class TestAuthFunction: """Tests for auth() function (local authentication).""" + def _make_mock_db(self): + """Create a mock DB that returns None for LocalUser queries (no local users).""" + mock_db = MagicMock() + # query().filter().first() returns None → no LocalUser found + mock_db.query.return_value.filter.return_value.first.return_value = None + return mock_db + @pytest.mark.asyncio async def test_auth_success(self): - """Test successful local authentication.""" + """Test successful local authentication (admin fallback).""" from app.auth import auth mock_request = MagicMock(spec=Request) @@ -743,7 +750,7 @@ class TestAuthFunction: mock_settings.admin_username = "testadmin" mock_settings.admin_password = "testpass" - result = await auth(mock_request) + result = await auth(mock_request, db=self._make_mock_db()) assert isinstance(result, RedirectResponse) assert result.status_code == 302 @@ -766,7 +773,7 @@ class TestAuthFunction: mock_settings.admin_username = "testadmin" mock_settings.admin_password = "testpass" - result = await auth(mock_request) + result = await auth(mock_request, db=self._make_mock_db()) assert isinstance(result, RedirectResponse) assert "/login?error=Invalid+username+or+password" in result.headers["location"] @@ -786,7 +793,7 @@ class TestAuthFunction: mock_settings.admin_username = "testadmin" mock_settings.admin_password = "testpass" - result = await auth(mock_request) + result = await auth(mock_request, db=self._make_mock_db()) assert isinstance(result, RedirectResponse) assert "/login?error=Invalid+username+or+password" in result.headers["location"] @@ -805,7 +812,7 @@ class TestAuthFunction: mock_settings.admin_username = "testadmin" mock_settings.admin_password = "testpass" - result = await auth(mock_request) + result = await auth(mock_request, db=self._make_mock_db()) assert isinstance(result, RedirectResponse) assert result.headers["location"] == "/settings" diff --git a/tests/test_billing.py b/tests/test_billing.py new file mode 100644 index 00000000..a669173a --- /dev/null +++ b/tests/test_billing.py @@ -0,0 +1,486 @@ +"""Tests for the Stripe billing API endpoints. + +Covers: +- POST /api/billing/create-checkout-session +- POST /api/billing/create-portal-session +- POST /api/billing/webhook (all event types) +- GET /api/billing/success +- Internal helpers: _handle_stripe_event, _on_checkout_completed, + _on_subscription_updated, _on_subscription_deleted +""" + +import json +from unittest.mock import MagicMock, patch + +import pytest +from fastapi.testclient import TestClient +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker +from sqlalchemy.pool import StaticPool + +from app.api.billing import ( + _handle_stripe_event, + _on_checkout_completed, + _on_subscription_deleted, + _on_subscription_updated, + _resolve_plan_id_from_price, + _resolve_user_id_from_customer, +) +from app.database import Base, get_db +from app.models import SubscriptionPlan, UserProfile + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture() +def bill_engine(): + """In-memory SQLite engine for billing tests.""" + engine = create_engine( + "sqlite:///:memory:", + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) + Base.metadata.create_all(bind=engine) + yield engine + Base.metadata.drop_all(bind=engine) + + +@pytest.fixture() +def bill_session(bill_engine): + """DB session for one test.""" + Session = sessionmaker(bind=bill_engine) + session = Session() + yield session + session.close() + + +@pytest.fixture() +def bill_client(bill_engine): + """TestClient with DB dependency overridden and a logged-in session.""" + from app.main import app + + Session = sessionmaker(bind=bill_engine) + + def override_get_db(): + db = Session() + try: + yield db + finally: + db.close() + + app.dependency_overrides[get_db] = override_get_db + with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client: + yield client + app.dependency_overrides.pop(get_db, None) + + +@pytest.fixture() +def starter_plan(bill_session): + """A SubscriptionPlan with Stripe price IDs in the DB.""" + plan = SubscriptionPlan( + plan_id="starter", + name="Starter", + price_monthly=9.0, + price_yearly=90.0, + trial_days=0, + stripe_price_id_monthly="price_monthly_starter", + stripe_price_id_yearly="price_yearly_starter", + ) + bill_session.add(plan) + bill_session.commit() + return plan + + +@pytest.fixture() +def user_profile(bill_session): + """A UserProfile for user1@example.com.""" + profile = UserProfile( + user_id="user1@example.com", + display_name="Test User", + stripe_customer_id=None, + ) + bill_session.add(profile) + bill_session.commit() + return profile + + +# --------------------------------------------------------------------------- +# Tests: _get_stripe returns None when not configured +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_get_stripe_returns_none_when_not_configured(): + """_get_stripe returns None when stripe_secret_key is not set.""" + from app.api.billing import _get_stripe + + with patch("app.api.billing.settings") as mock_settings: + mock_settings.stripe_secret_key = None + result = _get_stripe() + assert result is None + + +@pytest.mark.unit +def test_get_stripe_returns_client_when_configured(): + """_get_stripe returns a StripeClient when key is configured.""" + from app.api.billing import _get_stripe + + with patch("app.api.billing.settings") as mock_settings: + mock_settings.stripe_secret_key = "sk_test_fake" + result = _get_stripe() + assert result is not None + + +# --------------------------------------------------------------------------- +# Tests: create-checkout-session +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_create_checkout_session_stripe_not_configured(bill_client): + """POST /api/billing/create-checkout-session returns 503 when Stripe not set.""" + with patch("app.api.billing._get_stripe", return_value=None): + resp = bill_client.post( + "/api/billing/create-checkout-session", + json={"plan_id": "starter", "billing_cycle": "monthly"}, + ) + assert resp.status_code == 503 + + +@pytest.mark.integration +def test_create_checkout_session_plan_not_found(bill_client): + """POST /api/billing/create-checkout-session returns 404 for unknown plan.""" + mock_client = MagicMock() + with patch("app.api.billing._get_stripe", return_value=mock_client): + resp = bill_client.post( + "/api/billing/create-checkout-session", + json={"plan_id": "nonexistent", "billing_cycle": "monthly"}, + ) + assert resp.status_code == 404 + + +@pytest.mark.integration +def test_create_checkout_session_no_price_id(bill_client, bill_session): + """POST /api/billing/create-checkout-session returns 404 when price ID not set.""" + plan = SubscriptionPlan( + plan_id="noprice", + name="No Price", + price_monthly=5.0, + price_yearly=50.0, + trial_days=0, + stripe_price_id_monthly=None, + stripe_price_id_yearly=None, + ) + bill_session.add(plan) + bill_session.commit() + + mock_client = MagicMock() + with patch("app.api.billing._get_stripe", return_value=mock_client): + resp = bill_client.post( + "/api/billing/create-checkout-session", + json={"plan_id": "noprice", "billing_cycle": "monthly"}, + ) + assert resp.status_code == 404 + + +@pytest.mark.integration +def test_create_checkout_session_success(bill_client, starter_plan, user_profile): + """POST /api/billing/create-checkout-session returns checkout_url on success.""" + mock_client = MagicMock() + mock_customer = MagicMock() + mock_customer.id = "cus_test123" + mock_session = MagicMock() + mock_session.id = "cs_test456" + mock_session.url = "https://checkout.stripe.com/test" + + mock_client.customers.create.return_value = mock_customer + mock_client.checkout.sessions.create.return_value = mock_session + + with ( + patch("app.api.billing._get_stripe", return_value=mock_client), + patch("app.api.billing.get_current_owner_id", return_value="user1@example.com"), + ): + resp = bill_client.post( + "/api/billing/create-checkout-session", + json={"plan_id": "starter", "billing_cycle": "monthly"}, + ) + assert resp.status_code == 200 + data = resp.json() + assert "checkout_url" in data + assert data["checkout_url"] == "https://checkout.stripe.com/test" + + +@pytest.mark.integration +def test_create_checkout_session_yearly(bill_client, starter_plan, user_profile): + """POST /api/billing/create-checkout-session uses yearly price ID for yearly cycle.""" + mock_client = MagicMock() + mock_customer = MagicMock() + mock_customer.id = "cus_test123" + mock_session = MagicMock() + mock_session.id = "cs_test456" + mock_session.url = "https://checkout.stripe.com/yearly" + + mock_client.customers.create.return_value = mock_customer + mock_client.checkout.sessions.create.return_value = mock_session + + with ( + patch("app.api.billing._get_stripe", return_value=mock_client), + patch("app.api.billing.get_current_owner_id", return_value="user1@example.com"), + ): + resp = bill_client.post( + "/api/billing/create-checkout-session", + json={"plan_id": "starter", "billing_cycle": "yearly"}, + ) + assert resp.status_code == 200 + # Verify yearly price ID was used + call_params = mock_client.checkout.sessions.create.call_args[1]["params"] + assert call_params["line_items"][0]["price"] == "price_yearly_starter" + + +# --------------------------------------------------------------------------- +# Tests: create-portal-session +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_create_portal_session_stripe_not_configured(bill_client): + """POST /api/billing/create-portal-session returns 503 when not configured.""" + with patch("app.api.billing._get_stripe", return_value=None): + resp = bill_client.post("/api/billing/create-portal-session", json={}) + assert resp.status_code == 503 + + +@pytest.mark.integration +def test_create_portal_session_no_customer(bill_client, user_profile): + """POST /api/billing/create-portal-session returns 404 when no Stripe customer.""" + mock_client = MagicMock() + with ( + patch("app.api.billing._get_stripe", return_value=mock_client), + patch("app.api.billing.get_current_owner_id", return_value="user1@example.com"), + ): + resp = bill_client.post("/api/billing/create-portal-session", json={}) + assert resp.status_code == 404 + + +@pytest.mark.integration +def test_create_portal_session_success(bill_client, bill_session, user_profile): + """POST /api/billing/create-portal-session returns portal_url on success.""" + user_profile.stripe_customer_id = "cus_existing" + bill_session.commit() + + mock_client = MagicMock() + mock_portal = MagicMock() + mock_portal.url = "https://billing.stripe.com/portal/test" + mock_client.billing_portal.sessions.create.return_value = mock_portal + + with ( + patch("app.api.billing._get_stripe", return_value=mock_client), + patch("app.api.billing.get_current_owner_id", return_value="user1@example.com"), + ): + resp = bill_client.post("/api/billing/create-portal-session", json={}) + assert resp.status_code == 200 + assert resp.json()["portal_url"] == "https://billing.stripe.com/portal/test" + + +# --------------------------------------------------------------------------- +# Tests: webhook +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_webhook_not_configured(bill_client): + """POST /api/billing/webhook returns 503 when billing not configured.""" + with patch("app.api.billing.settings") as mock_settings: + mock_settings.stripe_secret_key = None + mock_settings.stripe_webhook_secret = None + resp = bill_client.post( + "/api/billing/webhook", + content=b"{}", + headers={"content-type": "application/json"}, + ) + assert resp.status_code == 503 + + +@pytest.mark.integration +def test_webhook_invalid_signature(bill_client): + """POST /api/billing/webhook returns 400 on invalid Stripe signature.""" + import stripe + + with ( + patch("app.api.billing.settings") as mock_settings, + patch("stripe.Webhook.construct_event", side_effect=stripe.SignatureVerificationError("bad", "sig")), + ): + mock_settings.stripe_secret_key = "sk_test_fake" + mock_settings.stripe_webhook_secret = "whsec_test" + resp = bill_client.post( + "/api/billing/webhook", + content=b'{"type":"test"}', + headers={"stripe-signature": "bad_sig", "content-type": "application/json"}, + ) + assert resp.status_code == 400 + + +@pytest.mark.integration +def test_webhook_checkout_completed(bill_client, bill_session, starter_plan, user_profile): + """POST /api/billing/webhook activates plan on checkout.session.completed.""" + payload = json.dumps( + { + "type": "checkout.session.completed", + "data": { + "object": { + "customer": "cus_new", + "metadata": { + "docuelevate_user_id": "user1@example.com", + "plan_id": "starter", + "billing_cycle": "monthly", + }, + } + }, + } + ).encode() + + with ( + patch("app.api.billing.settings") as mock_settings, + patch("stripe.Event.construct_from", return_value=json.loads(payload)), + ): + mock_settings.stripe_secret_key = "sk_test_fake" + mock_settings.stripe_webhook_secret = None + resp = bill_client.post( + "/api/billing/webhook", + content=payload, + headers={"content-type": "application/json"}, + ) + assert resp.status_code == 200 + + bill_session.expire_all() + profile = bill_session.query(UserProfile).filter(UserProfile.user_id == "user1@example.com").first() + assert profile.subscription_tier == "starter" + + +@pytest.mark.integration +def test_webhook_subscription_deleted(bill_client, bill_session, user_profile): + """POST /api/billing/webhook downgrades to free on subscription deleted.""" + user_profile.stripe_customer_id = "cus_del" + user_profile.subscription_tier = "starter" + bill_session.commit() + + payload = json.dumps( + { + "type": "customer.subscription.deleted", + "data": {"object": {"customer": "cus_del"}}, + } + ).encode() + + with ( + patch("app.api.billing.settings") as mock_settings, + patch("stripe.Event.construct_from", return_value=json.loads(payload)), + ): + mock_settings.stripe_secret_key = "sk_test_fake" + mock_settings.stripe_webhook_secret = None + resp = bill_client.post( + "/api/billing/webhook", + content=payload, + headers={"content-type": "application/json"}, + ) + assert resp.status_code == 200 + + bill_session.expire_all() + profile = bill_session.query(UserProfile).filter(UserProfile.user_id == "user1@example.com").first() + assert profile.subscription_tier == "free" + + +# --------------------------------------------------------------------------- +# Unit tests: internal helpers +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_resolve_user_id_from_customer(bill_session, user_profile): + """_resolve_user_id_from_customer returns user_id for known Stripe customer.""" + user_profile.stripe_customer_id = "cus_known" + bill_session.commit() + result = _resolve_user_id_from_customer(bill_session, "cus_known") + assert result == "user1@example.com" + + +@pytest.mark.unit +def test_resolve_user_id_from_customer_unknown(bill_session): + """_resolve_user_id_from_customer returns None for unknown customer.""" + result = _resolve_user_id_from_customer(bill_session, "cus_unknown") + assert result is None + + +@pytest.mark.unit +def test_resolve_plan_id_from_price(bill_session, starter_plan): + """_resolve_plan_id_from_price finds plan by monthly price ID.""" + result = _resolve_plan_id_from_price(bill_session, "price_monthly_starter") + assert result == "starter" + + +@pytest.mark.unit +def test_resolve_plan_id_from_price_yearly(bill_session, starter_plan): + """_resolve_plan_id_from_price finds plan by yearly price ID.""" + result = _resolve_plan_id_from_price(bill_session, "price_yearly_starter") + assert result == "starter" + + +@pytest.mark.unit +def test_resolve_plan_id_from_price_unknown(bill_session): + """_resolve_plan_id_from_price returns None for unknown price.""" + result = _resolve_plan_id_from_price(bill_session, "price_unknown") + assert result is None + + +@pytest.mark.unit +def test_on_checkout_completed_missing_user_id(bill_session): + """_on_checkout_completed does nothing when user_id is absent.""" + data = {"metadata": {}, "customer": "cus_test"} + _on_checkout_completed(bill_session, data) # Should not raise + + +@pytest.mark.unit +def test_on_subscription_updated_no_items(bill_session, user_profile): + """_on_subscription_updated does nothing when items list is empty.""" + user_profile.stripe_customer_id = "cus_upd" + bill_session.commit() + data = {"customer": "cus_upd", "items": {"data": []}} + _on_subscription_updated(bill_session, data) # Should not raise + + +@pytest.mark.unit +def test_on_subscription_deleted_unknown_customer(bill_session): + """_on_subscription_deleted does nothing for unknown customer.""" + data = {"customer": "cus_nobody"} + _on_subscription_deleted(bill_session, data) # Should not raise + + +@pytest.mark.unit +def test_handle_stripe_event_unhandled_type(bill_session): + """_handle_stripe_event logs but does not raise for unknown event types.""" + event = {"type": "unknown.event.type", "data": {"object": {}}} + _handle_stripe_event(bill_session, event) # Should not raise + + +@pytest.mark.unit +def test_handle_stripe_event_payment_failed(bill_session): + """_handle_stripe_event handles invoice.payment_failed without raising.""" + event = { + "type": "invoice.payment_failed", + "data": {"object": {"customer": "cus_fail"}}, + } + _handle_stripe_event(bill_session, event) # Should not raise + + +# --------------------------------------------------------------------------- +# Tests: billing success page +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_billing_success_page(bill_client): + """GET /api/billing/success returns 200 for logged-in user.""" + resp = bill_client.get("/api/billing/success") + assert resp.status_code == 200 + assert b"subscription" in resp.content.lower() or b"success" in resp.content.lower() diff --git a/tests/test_local_auth.py b/tests/test_local_auth.py new file mode 100644 index 00000000..e2bcf518 --- /dev/null +++ b/tests/test_local_auth.py @@ -0,0 +1,627 @@ +"""Tests for local user authentication: signup, email verification, and password reset. + +Covers: +- POST /api/auth/signup (success, disabled, SMTP missing, password mismatch, conflicts) +- GET /verify-email (valid token, invalid token, expired token) +- POST /api/auth/resend-verification +- POST /api/auth/request-password-reset +- POST /api/auth/reset-password +- GET /signup (page route) +- GET /verify-email-sent (page route) +- GET /reset-password (page route) +- app/utils/local_auth utility functions +- auth() login flow with LocalUser +""" + +from datetime import datetime, timedelta, timezone +from unittest.mock import MagicMock, patch + +import pytest +from fastapi.testclient import TestClient +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker +from sqlalchemy.pool import StaticPool + +from app.database import Base, get_db +from app.models import LocalUser, UserProfile +from app.utils.local_auth import ( + build_session_user, + generate_token, + hash_password, + is_token_expired, + verify_password, +) + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +_TEST_DB_URL = "sqlite:///:memory:" + + +@pytest.fixture() +def la_engine(): + """In-memory SQLite engine for local auth tests.""" + engine = create_engine( + _TEST_DB_URL, + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) + Base.metadata.create_all(bind=engine) + yield engine + Base.metadata.drop_all(bind=engine) + + +@pytest.fixture() +def la_session(la_engine): + """DB session for one test.""" + Session = sessionmaker(bind=la_engine) + session = Session() + yield session + session.close() + + +@pytest.fixture() +def la_client(la_engine): + """TestClient with DB dependency overridden.""" + from app.main import app + + Session = sessionmaker(bind=la_engine) + + def override_get_db(): + db = Session() + try: + yield db + finally: + db.close() + + app.dependency_overrides[get_db] = override_get_db + with TestClient(app, base_url="http://localhost", raise_server_exceptions=True) as client: + yield client + app.dependency_overrides.pop(get_db, None) + + +@pytest.fixture() +def active_user(la_session): + """A fully active LocalUser in the DB.""" + user = LocalUser( + email="active@example.com", + username="activeuser", + display_name="Active User", + hashed_password=hash_password("password123"), + is_active=True, + ) + la_session.add(user) + la_session.add(UserProfile(user_id="active@example.com", display_name="Active User", onboarding_completed=True)) + la_session.commit() + return user + + +@pytest.fixture() +def pending_user(la_session): + """A LocalUser with a pending email verification token.""" + token = "validtoken123" + user = LocalUser( + email="pending@example.com", + username="pendinguser", + hashed_password=hash_password("password123"), + is_active=False, + email_verification_token=token, + email_verification_sent_at=datetime.now(tz=timezone.utc), + ) + la_session.add(user) + la_session.commit() + return user + + +# --------------------------------------------------------------------------- +# Unit tests: local_auth utilities +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_hash_and_verify_password(): + """hash_password produces a bcrypt hash that verify_password validates.""" + plain = "super$ecret99" + hashed = hash_password(plain) + assert hashed != plain + assert verify_password(plain, hashed) is True + assert verify_password("wrong", hashed) is False + + +@pytest.mark.unit +def test_verify_password_bad_hash_returns_false(): + """verify_password returns False for a non-bcrypt string.""" + assert verify_password("any", "notahash") is False + + +@pytest.mark.unit +def test_generate_token_unique(): + """generate_token returns distinct non-empty strings.""" + tokens = {generate_token() for _ in range(10)} + assert len(tokens) == 10 + for t in tokens: + assert len(t) > 20 + + +@pytest.mark.unit +def test_is_token_expired_none(): + """None sent_at is treated as expired.""" + assert is_token_expired(None) is True + + +@pytest.mark.unit +def test_is_token_expired_old(): + """Token sent more than 24 h ago is expired.""" + old = datetime.now(tz=timezone.utc) - timedelta(hours=25) + assert is_token_expired(old) is True + + +@pytest.mark.unit +def test_is_token_expired_fresh(): + """Token sent recently is not expired.""" + fresh = datetime.now(tz=timezone.utc) - timedelta(hours=1) + assert is_token_expired(fresh) is False + + +@pytest.mark.unit +def test_build_session_user(): + """build_session_user returns the expected dict structure.""" + user = MagicMock() + user.email = "u@example.com" + user.username = "uname" + user.display_name = "Display Name" + user.is_admin = False + with patch("app.auth.get_gravatar_url", return_value="https://gravatar.com/test"): + result = build_session_user(user) + assert result["email"] == "u@example.com" + assert result["preferred_username"] == "uname" + assert result["name"] == "Display Name" + assert result["is_admin"] is False + assert result["auth_method"] == "local" + assert "picture" in result + + +# --------------------------------------------------------------------------- +# Integration tests: signup +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_signup_disabled(la_client): + """POST /api/auth/signup returns 403 when allow_local_signup is False.""" + with patch("app.api.local_auth.settings") as mock_settings: + mock_settings.allow_local_signup = False + mock_settings.email_host = "smtp.example.com" + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "a@example.com", + "username": "auser", + "password": "password1", + "password_confirm": "password1", + }, + ) + assert resp.status_code == 403 + + +@pytest.mark.integration +def test_signup_smtp_not_configured(la_client): + """POST /api/auth/signup returns 503 when SMTP is not configured.""" + with patch("app.api.local_auth.settings") as mock_settings: + mock_settings.allow_local_signup = True + mock_settings.email_host = None + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "a@example.com", + "username": "auser", + "password": "password1", + "password_confirm": "password1", + }, + ) + assert resp.status_code == 503 + + +@pytest.mark.integration +def test_signup_password_mismatch(la_client): + """POST /api/auth/signup returns 422 when passwords do not match.""" + with patch("app.api.local_auth.settings") as mock_settings: + mock_settings.allow_local_signup = True + mock_settings.email_host = "smtp.example.com" + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "a@example.com", + "username": "auser", + "password": "password1", + "password_confirm": "different1", + }, + ) + assert resp.status_code == 422 + + +@pytest.mark.integration +def test_signup_success(la_client): + """POST /api/auth/signup creates user and returns 201.""" + with ( + patch("app.api.local_auth.settings") as mock_settings, + patch("app.api.local_auth.send_verification_email") as mock_send, + ): + mock_settings.allow_local_signup = True + mock_settings.email_host = "smtp.example.com" + mock_settings.version = "test" + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "new@example.com", + "username": "newuser", + "password": "password1", + "password_confirm": "password1", + }, + ) + assert resp.status_code == 201 + assert "Verification email sent" in resp.json()["message"] + mock_send.assert_called_once() + + +@pytest.mark.integration +def test_signup_duplicate_email(la_client, active_user): + """POST /api/auth/signup returns 409 when email already registered.""" + with ( + patch("app.api.local_auth.settings") as mock_settings, + patch("app.api.local_auth.send_verification_email"), + ): + mock_settings.allow_local_signup = True + mock_settings.email_host = "smtp.example.com" + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "active@example.com", + "username": "otheruser", + "password": "password1", + "password_confirm": "password1", + }, + ) + assert resp.status_code == 409 + assert "Email" in resp.json()["detail"] + + +@pytest.mark.integration +def test_signup_duplicate_username(la_client, active_user): + """POST /api/auth/signup returns 409 when username already taken.""" + with ( + patch("app.api.local_auth.settings") as mock_settings, + patch("app.api.local_auth.send_verification_email"), + ): + mock_settings.allow_local_signup = True + mock_settings.email_host = "smtp.example.com" + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "different@example.com", + "username": "activeuser", + "password": "password1", + "password_confirm": "password1", + }, + ) + assert resp.status_code == 409 + assert "Username" in resp.json()["detail"] + + +@pytest.mark.integration +def test_signup_smtp_failure_cleans_up(la_client, la_session): + """POST /api/auth/signup cleans up user records if email send fails.""" + with ( + patch("app.api.local_auth.settings") as mock_settings, + patch("app.api.local_auth.send_verification_email", side_effect=RuntimeError("SMTP down")), + ): + mock_settings.allow_local_signup = True + mock_settings.email_host = "smtp.example.com" + resp = la_client.post( + "/api/auth/signup", + json={ + "email": "fail@example.com", + "username": "failuser", + "password": "password1", + "password_confirm": "password1", + }, + ) + assert resp.status_code == 503 + # User should NOT exist in the DB + user = la_session.query(LocalUser).filter(LocalUser.email == "fail@example.com").first() + assert user is None + + +# --------------------------------------------------------------------------- +# Integration tests: email verification +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_verify_email_valid_token(la_client, pending_user): + """GET /verify-email with valid token activates account and redirects.""" + resp = la_client.get( + f"/verify-email?token={pending_user.email_verification_token}", + follow_redirects=False, + ) + assert resp.status_code == 302 + + +@pytest.mark.integration +def test_verify_email_invalid_token(la_client): + """GET /verify-email with unknown token redirects to login with error.""" + resp = la_client.get("/verify-email?token=doesnotexist", follow_redirects=False) + assert resp.status_code == 302 + assert "/login" in resp.headers["location"] + + +@pytest.mark.integration +def test_verify_email_expired_token(la_client, la_session): + """GET /verify-email with expired token redirects to login with error.""" + old_time = datetime.now(tz=timezone.utc) - timedelta(hours=25) + user = LocalUser( + email="expired@example.com", + username="expireduser", + hashed_password=hash_password("password123"), + is_active=False, + email_verification_token="expiredtoken", + email_verification_sent_at=old_time, + ) + la_session.add(user) + la_session.commit() + + resp = la_client.get("/verify-email?token=expiredtoken", follow_redirects=False) + assert resp.status_code == 302 + assert "/login" in resp.headers["location"] + + +# --------------------------------------------------------------------------- +# Integration tests: resend verification +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_resend_verification_always_200(la_client): + """POST /api/auth/resend-verification returns 200 for unknown email.""" + with patch("app.api.local_auth.send_verification_email"): + resp = la_client.post( + "/api/auth/resend-verification", + json={"email": "nobody@example.com"}, + ) + assert resp.status_code == 200 + + +@pytest.mark.integration +def test_resend_verification_sends_email(la_client, pending_user): + """POST /api/auth/resend-verification sends email for pending user.""" + with patch("app.api.local_auth.send_verification_email") as mock_send: + resp = la_client.post( + "/api/auth/resend-verification", + json={"email": pending_user.email}, + ) + assert resp.status_code == 200 + mock_send.assert_called_once() + + +# --------------------------------------------------------------------------- +# Integration tests: password reset +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_request_password_reset_always_200(la_client): + """POST /api/auth/request-password-reset returns 200 for unknown email.""" + with patch("app.api.local_auth.send_password_reset_email"): + resp = la_client.post( + "/api/auth/request-password-reset", + json={"email": "nobody@example.com"}, + ) + assert resp.status_code == 200 + + +@pytest.mark.integration +def test_request_password_reset_sends_email(la_client, active_user): + """POST /api/auth/request-password-reset sends email for known user.""" + with patch("app.api.local_auth.send_password_reset_email") as mock_send: + resp = la_client.post( + "/api/auth/request-password-reset", + json={"email": active_user.email}, + ) + assert resp.status_code == 200 + mock_send.assert_called_once() + + +@pytest.mark.integration +def test_reset_password_success(la_client, la_session): + """POST /api/auth/reset-password updates password with valid token.""" + token = "resettoken123" + user = LocalUser( + email="reset@example.com", + username="resetuser", + hashed_password=hash_password("oldpassword"), + is_active=True, + password_reset_token=token, + password_reset_sent_at=datetime.now(tz=timezone.utc), + ) + la_session.add(user) + la_session.commit() + + resp = la_client.post( + "/api/auth/reset-password", + json={ + "token": token, + "new_password": "newpassword1", + "new_password_confirm": "newpassword1", + }, + ) + assert resp.status_code == 200 + la_session.refresh(user) + assert verify_password("newpassword1", user.hashed_password) + assert user.password_reset_token is None + + +@pytest.mark.integration +def test_reset_password_invalid_token(la_client): + """POST /api/auth/reset-password returns 400 for invalid token.""" + resp = la_client.post( + "/api/auth/reset-password", + json={ + "token": "badtoken", + "new_password": "newpassword1", + "new_password_confirm": "newpassword1", + }, + ) + assert resp.status_code == 400 + + +@pytest.mark.integration +def test_reset_password_mismatch(la_client, la_session): + """POST /api/auth/reset-password returns 422 when passwords do not match.""" + token = "mismatchtoken" + user = LocalUser( + email="mismatch@example.com", + username="mismatchuser", + hashed_password=hash_password("old"), + is_active=True, + password_reset_token=token, + password_reset_sent_at=datetime.now(tz=timezone.utc), + ) + la_session.add(user) + la_session.commit() + + resp = la_client.post( + "/api/auth/reset-password", + json={ + "token": token, + "new_password": "newpassword1", + "new_password_confirm": "different_pw", + }, + ) + assert resp.status_code == 422 + + +# --------------------------------------------------------------------------- +# Integration tests: page routes +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +def test_signup_page_disabled_redirects(la_client): + """GET /signup redirects when allow_local_signup is False.""" + with patch("app.api.local_auth.settings") as mock_settings: + mock_settings.allow_local_signup = False + resp = la_client.get("/signup", follow_redirects=False) + assert resp.status_code == 302 + assert "/login" in resp.headers["location"] + + +@pytest.mark.integration +def test_signup_page_enabled(la_client): + """GET /signup returns 200 when allow_local_signup is True.""" + with patch("app.api.local_auth.settings") as mock_settings: + mock_settings.allow_local_signup = True + mock_settings.version = "test" + resp = la_client.get("/signup") + assert resp.status_code == 200 + assert b"Create" in resp.content + + +@pytest.mark.integration +def test_verify_email_sent_page(la_client): + """GET /verify-email-sent returns 200.""" + resp = la_client.get("/verify-email-sent") + assert resp.status_code == 200 + + +@pytest.mark.integration +def test_reset_password_page(la_client): + """GET /reset-password returns 200.""" + with patch("app.api.local_auth.settings") as mock_settings: + mock_settings.version = "test" + resp = la_client.get("/reset-password?token=abc123") + assert resp.status_code == 200 + assert b"password" in resp.content.lower() + + +# --------------------------------------------------------------------------- +# Integration tests: auth() login flow with LocalUser +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +def test_local_login_success(la_session, active_user): + """auth() with valid LocalUser credentials sets session and redirects.""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + + from app.auth import auth + + mock_request = MagicMock(spec=Request) + mock_request.form = AsyncMock(return_value={"username": "activeuser", "password": "password123"}) + mock_request.session = {} + + result = asyncio.get_event_loop().run_until_complete(auth(mock_request, db=la_session)) + assert result.status_code == 302 + assert "user" in mock_request.session + assert mock_request.session["user"]["email"] == "active@example.com" + + +@pytest.mark.unit +def test_local_login_by_email(la_session, active_user): + """auth() accepts email as username for LocalUser lookup.""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + + from app.auth import auth + + mock_request = MagicMock(spec=Request) + mock_request.form = AsyncMock(return_value={"username": "active@example.com", "password": "password123"}) + mock_request.session = {} + + result = asyncio.get_event_loop().run_until_complete(auth(mock_request, db=la_session)) + assert result.status_code == 302 + assert "user" in mock_request.session + + +@pytest.mark.unit +def test_local_login_wrong_password(la_session, active_user): + """auth() with wrong password redirects to login with error.""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + + from app.auth import auth + + mock_request = MagicMock(spec=Request) + mock_request.form = AsyncMock(return_value={"username": "activeuser", "password": "wrongpassword"}) + mock_request.session = {} + + result = asyncio.get_event_loop().run_until_complete(auth(mock_request, db=la_session)) + assert result.status_code == 302 + assert "/login" in result.headers["location"] + assert "user" not in mock_request.session + + +@pytest.mark.unit +def test_local_login_unverified(la_session, pending_user): + """auth() for unverified user redirects with verification message.""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + + from app.auth import auth + + mock_request = MagicMock(spec=Request) + mock_request.form = AsyncMock(return_value={"username": "pendinguser", "password": "password123"}) + mock_request.session = {} + + result = asyncio.get_event_loop().run_until_complete(auth(mock_request, db=la_session)) + assert result.status_code == 302 + assert "verify" in result.headers["location"].lower()