From 5ea77c3991e5f541cd84a074cff8031f50c72115 Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Thu, 12 Mar 2026 10:51:45 +0100 Subject: [PATCH 1/2] feat(auth): add multi-user authentication foundation. Related to #151 - Add User domain model and SQLAlchemy UserModel with soft-delete support - Add SqlAlchemyUserStore with email/GitHub user CRUD, password auth, reset tokens - Add JWT signing/verification (python-jose), bcrypt password hashing - Add FastAPI auth dependencies: get_user_id (optional fallback) and get_current_user (strict) - Add /api/auth routes: register, login, github/exchange, me, forgot/reset-password - Add Alembic migrations for users and password_reset_tokens tables - Add AUTH_OPTIONAL env var for gradual migration from legacy 'default' user - Fix account lifecycle bugs: reactivate on OAuth re-login, reject inactive on password login - Add auth API tests --- alembic/versions/0024_users.py | 53 +++++ .../versions/0025_password_reset_tokens.py | 40 ++++ pyproject.toml | 4 +- requirements.txt | 9 +- src/paperbot/api/auth/dependencies.py | 68 ++++++ src/paperbot/api/auth/email.py | 47 ++++ src/paperbot/api/auth/jwt.py | 19 ++ src/paperbot/api/auth/password.py | 11 + src/paperbot/api/main.py | 2 + src/paperbot/api/routes/auth.py | 221 ++++++++++++++++++ src/paperbot/domain/user.py | 16 ++ src/paperbot/infrastructure/stores/models.py | 62 ++++- .../infrastructure/stores/user_store.py | 179 ++++++++++++++ tests/test_auth_api.py | 90 +++++++ 14 files changed, 817 insertions(+), 4 deletions(-) create mode 100644 alembic/versions/0024_users.py create mode 100644 alembic/versions/0025_password_reset_tokens.py create mode 100644 src/paperbot/api/auth/dependencies.py create mode 100644 src/paperbot/api/auth/email.py create mode 100644 src/paperbot/api/auth/jwt.py create mode 100644 src/paperbot/api/auth/password.py create mode 100644 src/paperbot/api/routes/auth.py create mode 100644 src/paperbot/domain/user.py create mode 100644 src/paperbot/infrastructure/stores/user_store.py create mode 100644 tests/test_auth_api.py diff --git a/alembic/versions/0024_users.py b/alembic/versions/0024_users.py new file mode 100644 index 00000000..90d4e51f --- /dev/null +++ b/alembic/versions/0024_users.py @@ -0,0 +1,53 @@ +"""Add users table + +Revision ID: 0024_users +Revises: 0023_intelligence_events +Create Date: 2026-03-09 +""" +from __future__ import annotations + +from alembic import op +import sqlalchemy as sa + + +revision = "0024_users" +down_revision = "0023_intelligence_events" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + + if "users" not in inspector.get_table_names(): + op.create_table( + "users", + sa.Column("id", sa.Integer, primary_key=True, autoincrement=True), + sa.Column("email", sa.String(255), nullable=True), + sa.Column("hashed_password", sa.String(255), nullable=True), + sa.Column("github_id", sa.String(64), nullable=True), + sa.Column("github_username", sa.String(128), nullable=True), + sa.Column("display_name", sa.String(128), nullable=True), + sa.Column("avatar_url", sa.String(512), nullable=True), + sa.Column("is_active", sa.Boolean, nullable=False, server_default=sa.true()), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("last_login_at", sa.DateTime(timezone=True), nullable=True), + sa.CheckConstraint( + "email IS NOT NULL OR github_id IS NOT NULL", + name="ck_users_identity", + ), + ) + + existing_indexes = {i["name"] for i in inspector.get_indexes("users")} + if "uq_users_email" not in existing_indexes: + op.create_index("uq_users_email", "users", ["email"], unique=True) + if "uq_users_github_id" not in existing_indexes: + op.create_index("uq_users_github_id", "users", ["github_id"], unique=True) + + +def downgrade() -> None: + op.drop_index("uq_users_github_id", table_name="users") + op.drop_index("uq_users_email", table_name="users") + op.drop_table("users") + diff --git a/alembic/versions/0025_password_reset_tokens.py b/alembic/versions/0025_password_reset_tokens.py new file mode 100644 index 00000000..0f8826b7 --- /dev/null +++ b/alembic/versions/0025_password_reset_tokens.py @@ -0,0 +1,40 @@ +"""Add password_reset_tokens table + +Revision ID: 0025_password_reset_tokens +Revises: 0024_users +Create Date: 2026-03-10 +""" +from __future__ import annotations + +from alembic import op +import sqlalchemy as sa + + +revision = "0025_password_reset_tokens" +down_revision = "0024_users" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + + if "password_reset_tokens" not in inspector.get_table_names(): + op.create_table( + "password_reset_tokens", + sa.Column("id", sa.Integer, primary_key=True, autoincrement=True), + sa.Column("user_id", sa.Integer, nullable=False), + sa.Column("token", sa.String(64), nullable=False), + sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("used", sa.Boolean, nullable=False, server_default=sa.false()), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + ) + op.create_index("ix_prt_token", "password_reset_tokens", ["token"], unique=True) + op.create_index("ix_prt_user_id", "password_reset_tokens", ["user_id"]) + + +def downgrade() -> None: + op.drop_index("ix_prt_user_id", table_name="password_reset_tokens") + op.drop_index("ix_prt_token", table_name="password_reset_tokens") + op.drop_table("password_reset_tokens") diff --git a/pyproject.toml b/pyproject.toml index 4273577b..54ed88e4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,7 +33,7 @@ dependencies = [ "beautifulsoup4>=4.11.0", "cryptography>=41.0.0", "lxml>=4.9.0", - "pydantic>=2.0", + "pydantic[email]>=2.0", "loguru>=0.7.0", "openai>=1.0.0", "anthropic>=0.3.0", @@ -56,6 +56,8 @@ dependencies = [ "json-repair>=0.22.0", "arq>=0.25.0,<0.26.0", "redis>=5.0.0", + "python-jose[cryptography]>=3.3.0", + "bcrypt>=4.0.0", ] [project.optional-dependencies] diff --git a/requirements.txt b/requirements.txt index 61978405..76c99639 100644 --- a/requirements.txt +++ b/requirements.txt @@ -25,8 +25,8 @@ python-dotenv>=0.19.0 cryptography>=41.0.0 keyring>=25.0.0 -# 配置模型 -pydantic>=2.0.0 +# 配置模型(需要 EmailStr 支持) +pydantic[email]>=2.0.0 # 日志和调试 colorlog>=6.7.0 @@ -101,3 +101,8 @@ sqlite-vec>=0.1.6 # Fuzzy string matching for paper deduplication rapidfuzz>=3.0.0 + +# Authentication (JWT + password hashing) +python-jose[cryptography]>=3.3.0 +bcrypt>=4.0.0 +resend>=0.7.0 diff --git a/src/paperbot/api/auth/dependencies.py b/src/paperbot/api/auth/dependencies.py new file mode 100644 index 00000000..2dcb1b2a --- /dev/null +++ b/src/paperbot/api/auth/dependencies.py @@ -0,0 +1,68 @@ +import logging +import os + +from fastapi import Depends, HTTPException, status +from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials +from jose import JWTError + +from paperbot.api.auth.jwt import decode_token +from paperbot.infrastructure.stores.user_store import SqlAlchemyUserStore + +logger = logging.getLogger(__name__) + +bearer = HTTPBearer(auto_error=False) + +AUTH_OPTIONAL = os.getenv("AUTH_OPTIONAL", "false").lower() in {"1", "true", "yes"} + +_user_store = SqlAlchemyUserStore() + + +def _resolve_user(credentials: HTTPAuthorizationCredentials | None): + if not credentials or not credentials.credentials: + logger.warning("[auth] Missing token — no credentials provided") + if AUTH_OPTIONAL: + return None + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Missing token") + token = credentials.credentials + logger.debug("[auth] Received token (first 20 chars): %s...", token[:20]) + try: + user_id = decode_token(token) + logger.debug("[auth] Token valid, user_id=%s", user_id) + except JWTError as e: + logger.warning("[auth] Invalid token: %s | token prefix: %s...", e, token[:20]) + if AUTH_OPTIONAL: + return None + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid token") + + user = _user_store.get_by_id(user_id) + if not user or not user.is_active: + if AUTH_OPTIONAL: + return None + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found") + return user + + +def get_current_user(credentials: HTTPAuthorizationCredentials | None = Depends(bearer)): + """Strict user dependency: always requires a valid user. + + Use this for endpoints that must not fall back to the legacy "default" namespace. + """ + + user = _resolve_user(credentials) + if user is None: + # AUTH_OPTIONAL only affects get_user_id; this dependency always enforces auth. + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Missing or invalid token") + return user + + +def get_user_id(credentials: HTTPAuthorizationCredentials | None = Depends(bearer)) -> str: + """Return the authenticated user id as a string. + + When AUTH_OPTIONAL=true, missing/invalid tokens fall back to "default" so + legacy callers keep functioning while we migrate to multi-user auth. + """ + + user = _resolve_user(credentials) + if user is None: + return "default" + return str(user.id) diff --git a/src/paperbot/api/auth/email.py b/src/paperbot/api/auth/email.py new file mode 100644 index 00000000..61748df5 --- /dev/null +++ b/src/paperbot/api/auth/email.py @@ -0,0 +1,47 @@ +"""Email sending for auth flows. + +Currently in log mode — no real email is sent. +To switch to Resend, set RESEND_API_KEY and flip _SEND_MODE = "resend". +""" +from __future__ import annotations + +import logging +import os + +logger = logging.getLogger(__name__) + +_SEND_MODE = os.getenv("EMAIL_SEND_MODE", "log") # "log" | "resend" + + +def send_password_reset_email(to_email: str, reset_url: str) -> None: + if _SEND_MODE == "resend": + _send_via_resend(to_email, reset_url) + else: + _log_email(to_email, reset_url) + + +def _log_email(to_email: str, reset_url: str) -> None: + logger.warning( + "[DEV] Password reset link for %s → %s", + to_email, + reset_url, + ) + + +def _send_via_resend(to_email: str, reset_url: str) -> None: # pragma: no cover + import resend # type: ignore[import] + + resend.api_key = os.environ["RESEND_API_KEY"] + from_addr = os.getenv("EMAIL_FROM", "PaperBot ") + + resend.Emails.send({ + "from": from_addr, + "to": [to_email], + "subject": "Reset your PaperBot password", + "html": f""" +

Hi,

+

Click the link below to reset your password. This link expires in 1 hour.

+

{reset_url}

+

If you didn't request a password reset, you can ignore this email.

+ """, + }) diff --git a/src/paperbot/api/auth/jwt.py b/src/paperbot/api/auth/jwt.py new file mode 100644 index 00000000..16b0509c --- /dev/null +++ b/src/paperbot/api/auth/jwt.py @@ -0,0 +1,19 @@ +import os +from datetime import datetime, timedelta, timezone +from jose import jwt + +SECRET_KEY = os.environ.get("PAPERBOT_JWT_SECRET", "change-me-in-production") +ALGORITHM = "HS256" +ACCESS_TOKEN_EXPIRE_MINUTES = 60 * 24 * 7 # 7 days + + +def create_access_token(user_id: int) -> str: + expire = datetime.now(timezone.utc) + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES) + payload = {"sub": str(user_id), "exp": expire} + return jwt.encode(payload, SECRET_KEY, algorithm=ALGORITHM) + + +def decode_token(token: str) -> int: + payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) + return int(payload["sub"]) # raises if missing/invalid + diff --git a/src/paperbot/api/auth/password.py b/src/paperbot/api/auth/password.py new file mode 100644 index 00000000..32729533 --- /dev/null +++ b/src/paperbot/api/auth/password.py @@ -0,0 +1,11 @@ +import bcrypt + + +def hash_password(password: str) -> str: + return bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode() + + +def verify_password(plain: str, hashed: str) -> bool: + if not hashed: + return False + return bcrypt.checkpw(plain.encode(), hashed.encode()) diff --git a/src/paperbot/api/main.py b/src/paperbot/api/main.py index 401dfe7b..d81222a6 100644 --- a/src/paperbot/api/main.py +++ b/src/paperbot/api/main.py @@ -33,6 +33,7 @@ intelligence, push_commands, agent_board, + auth, ) from paperbot.api.error_handling import install_api_error_handling from paperbot.infrastructure.event_log.logging_event_log import LoggingEventLog @@ -90,6 +91,7 @@ async def health_check(): app.include_router(intelligence.router, prefix="/api", tags=["Intelligence"]) app.include_router(push_commands.router, prefix="/api", tags=["Push"]) app.include_router(agent_board.router, tags=["Agent Board"]) +app.include_router(auth.router) @app.on_event("startup") diff --git a/src/paperbot/api/routes/auth.py b/src/paperbot/api/routes/auth.py new file mode 100644 index 00000000..ba6159fb --- /dev/null +++ b/src/paperbot/api/routes/auth.py @@ -0,0 +1,221 @@ +from __future__ import annotations + +import os +import httpx +from fastapi import APIRouter, HTTPException, Depends +from pydantic import BaseModel, EmailStr + +from paperbot.api.auth.password import hash_password +from paperbot.api.auth.jwt import create_access_token +from paperbot.api.auth.dependencies import get_current_user +from paperbot.api.auth.email import send_password_reset_email +from paperbot.infrastructure.stores.user_store import SqlAlchemyUserStore +from paperbot.domain.user import User + + +router = APIRouter(prefix="/api/auth", tags=["auth"]) + +_user_store = SqlAlchemyUserStore() + + +class RegisterRequest(BaseModel): + email: EmailStr + password: str + display_name: str | None = None + + +class TokenResponse(BaseModel): + access_token: str + token_type: str = "bearer" + user_id: int + display_name: str | None + + +@router.post("/register", response_model=TokenResponse, status_code=201) +def register(req: RegisterRequest): + if len(req.password) < 8: + raise HTTPException(status_code=400, detail="Password too short (min 8 chars)") + + existing = _user_store.get_by_email(req.email) + if existing: + if existing.is_active: + raise HTTPException(status_code=400, detail="Email already registered") + # Deactivated account: reactivate and reset password + _user_store.reactivate(existing.id) + _user_store.update_password(existing.id, hash_password(req.password)) + if req.display_name: + _user_store.update_profile(existing.id, req.display_name) + return TokenResponse( + access_token=create_access_token(existing.id), + user_id=existing.id, + display_name=req.display_name or existing.display_name, + ) + + user = _user_store.create_email_user( + email=req.email, + hashed_password=hash_password(req.password), + display_name=req.display_name, + ) + return TokenResponse( + access_token=create_access_token(user.id), + user_id=user.id, + display_name=user.display_name, + ) + + +class LoginRequest(BaseModel): + email: EmailStr + password: str + + +@router.post("/login", response_model=TokenResponse) +def login(req: LoginRequest): + existing = _user_store.get_by_email(req.email) + if not existing: + raise HTTPException(status_code=401, detail="Email not registered.") + if not existing.is_active: + raise HTTPException(status_code=401, detail="Account has been deleted. Please register again.") + user = _user_store.authenticate(req.email, req.password) + if not user: + raise HTTPException(status_code=401, detail="Incorrect password.") + return TokenResponse( + access_token=create_access_token(user.id), + user_id=user.id, + display_name=user.display_name, + ) + + +class MeResponse(BaseModel): + id: int + email: str | None + github_username: str | None + display_name: str | None + avatar_url: str | None + + +@router.get("/me", response_model=MeResponse) +def me(current_user: User = Depends(get_current_user)): + return MeResponse( + id=current_user.id, + email=current_user.email, + github_username=current_user.github_username, + display_name=current_user.display_name, + avatar_url=current_user.avatar_url, + ) + + +class GithubExchangeRequest(BaseModel): + github_id: str + login: str | None = None + name: str | None = None + avatar_url: str | None = None + email: str | None = None + access_token: str + + +@router.post("/github/exchange", response_model=TokenResponse) +async def github_exchange(req: GithubExchangeRequest): + # Verify the GitHub token against GitHub API to prevent forged requests + async with httpx.AsyncClient() as client: + resp = await client.get( + "https://api.github.com/user", + headers={"Authorization": f"Bearer {req.access_token}", "Accept": "application/json"}, + timeout=15.0, + ) + if resp.status_code != 200: + raise HTTPException(status_code=400, detail="Invalid GitHub token") + gh = resp.json() + if str(gh.get("id")) != str(req.github_id): + raise HTTPException(status_code=400, detail="GitHub id mismatch") + + user = _user_store.get_by_github_id(req.github_id) + if user: + if not user.is_active: + _user_store.reactivate(user.id) + _user_store.update_last_login(user.id) + else: + user = _user_store.create_github_user( + github_id=str(req.github_id), + username=req.login or (gh.get("login") or ""), + display_name=(req.name or gh.get("name") or gh.get("login") or ""), + avatar_url=(req.avatar_url or gh.get("avatar_url") or ""), + ) + + return TokenResponse( + access_token=create_access_token(user.id), + user_id=user.id, + display_name=user.display_name, + ) + + +# ── Account management ─────────────────────────────────────────────────────── + +class UpdateMeRequest(BaseModel): + display_name: str | None = None + + +@router.patch("/me", response_model=MeResponse) +def update_me(req: UpdateMeRequest, current_user: User = Depends(get_current_user)): + _user_store.update_profile(current_user.id, req.display_name) + return MeResponse( + id=current_user.id, + email=current_user.email, + github_username=current_user.github_username, + display_name=req.display_name, + avatar_url=current_user.avatar_url, + ) + + +class ChangePasswordRequest(BaseModel): + current_password: str + new_password: str + + +@router.post("/me/change-password", status_code=200) +def change_password(req: ChangePasswordRequest, current_user: User = Depends(get_current_user)): + if not current_user.email: + raise HTTPException(status_code=400, detail="Password change not available for OAuth accounts") + if len(req.new_password) < 8: + raise HTTPException(status_code=400, detail="Password too short (min 8 chars)") + if not _user_store.authenticate(current_user.email, req.current_password): + raise HTTPException(status_code=400, detail="Current password is incorrect") + _user_store.update_password(current_user.id, hash_password(req.new_password)) + return {"detail": "Password updated."} + + +@router.delete("/me", status_code=204) +def delete_me(current_user: User = Depends(get_current_user)): + _user_store.deactivate(current_user.id) + + +# ── Forgot / Reset password ────────────────────────────────────────────────── + +class ForgotPasswordRequest(BaseModel): + email: EmailStr + + +@router.post("/forgot-password", status_code=202) +def forgot_password(req: ForgotPasswordRequest): + """Send a password-reset link. Always returns 202 to avoid email enumeration.""" + user = _user_store.get_by_email(req.email) + if user and user.is_active and user.email: + token = _user_store.create_reset_token(user.id) + frontend_url = os.getenv("FRONTEND_URL", "http://localhost:3000") + reset_url = f"{frontend_url}/reset-password?token={token}" + send_password_reset_email(user.email, reset_url) + return {"detail": "If that email is registered, a reset link has been sent."} + + +class ResetPasswordRequest(BaseModel): + token: str + new_password: str + + +@router.post("/reset-password", status_code=200) +def reset_password(req: ResetPasswordRequest): + if len(req.new_password) < 8: + raise HTTPException(status_code=400, detail="Password too short (min 8 chars)") + ok = _user_store.consume_reset_token(req.token, hash_password(req.new_password)) + if not ok: + raise HTTPException(status_code=400, detail="Invalid or expired reset token") + return {"detail": "Password updated successfully."} diff --git a/src/paperbot/domain/user.py b/src/paperbot/domain/user.py new file mode 100644 index 00000000..17d8f59a --- /dev/null +++ b/src/paperbot/domain/user.py @@ -0,0 +1,16 @@ +from dataclasses import dataclass +from datetime import datetime +from typing import Optional + + +@dataclass +class User: + id: int + email: Optional[str] + github_id: Optional[str] + github_username: Optional[str] + display_name: Optional[str] + avatar_url: Optional[str] + is_active: bool + created_at: datetime + diff --git a/src/paperbot/infrastructure/stores/models.py b/src/paperbot/infrastructure/stores/models.py index a7ab5a44..6988ef10 100644 --- a/src/paperbot/infrastructure/stores/models.py +++ b/src/paperbot/infrastructure/stores/models.py @@ -4,7 +4,19 @@ from datetime import datetime from typing import Any, Dict, Optional -from sqlalchemy import Boolean, DateTime, Float, ForeignKey, Index, Integer, LargeBinary, String, Text, UniqueConstraint +from sqlalchemy import ( + Boolean, + DateTime, + Float, + ForeignKey, + Index, + Integer, + LargeBinary, + String, + Text, + UniqueConstraint, + CheckConstraint, +) from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship @@ -1335,3 +1347,51 @@ class IntelligenceEventModel(Base): updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) payload_json: Mapped[str] = mapped_column(Text, default="{}") + + +# ============================================================================ +# Authentication: Users +# ============================================================================ + + +class UserModel(Base): + """User accounts for authentication and personalization. + + Either email (with hashed_password) or github_id must be provided. + """ + + __tablename__ = "users" + __table_args__ = ( + UniqueConstraint("email", name="uq_users_email"), + UniqueConstraint("github_id", name="uq_users_github_id"), + CheckConstraint("email IS NOT NULL OR github_id IS NOT NULL", name="ck_users_identity"), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + email: Mapped[Optional[str]] = mapped_column(String(255), nullable=True, index=True) + hashed_password: Mapped[Optional[str]] = mapped_column(String(255), nullable=True) + github_id: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, index=True) + github_username: Mapped[Optional[str]] = mapped_column(String(128), nullable=True) + display_name: Mapped[Optional[str]] = mapped_column(String(128), nullable=True) + avatar_url: Mapped[Optional[str]] = mapped_column(String(512), nullable=True) + is_active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + last_login_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True) + + +# ============================================================================ +# Authentication: Password Reset Tokens +# ============================================================================ + + +class PasswordResetTokenModel(Base): + """One-time tokens for password reset. Expires after 1 hour.""" + + __tablename__ = "password_reset_tokens" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + user_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True) + token: Mapped[str] = mapped_column(String(64), nullable=False, unique=True, index=True) + expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + used: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) diff --git a/src/paperbot/infrastructure/stores/user_store.py b/src/paperbot/infrastructure/stores/user_store.py new file mode 100644 index 00000000..6be7311e --- /dev/null +++ b/src/paperbot/infrastructure/stores/user_store.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +import secrets +from datetime import datetime, timedelta, timezone +from typing import Optional + +from paperbot.api.auth.password import verify_password +from paperbot.infrastructure.stores.models import PasswordResetTokenModel, UserModel +from paperbot.infrastructure.stores.sqlalchemy_db import SessionProvider, get_db_url +from paperbot.domain.user import User + + +class SqlAlchemyUserStore: + def __init__(self, db_url: Optional[str] = None): + self.db_url = db_url or get_db_url() + self._provider = SessionProvider(self.db_url) + + def _to_domain(self, row: UserModel) -> User: + return User( + id=row.id, + email=row.email, + github_id=row.github_id, + github_username=row.github_username, + display_name=row.display_name, + avatar_url=row.avatar_url, + is_active=row.is_active, + created_at=row.created_at, + ) + + def get_by_id(self, user_id: int) -> Optional[User]: + with self._provider.session() as session: + row = session.get(UserModel, user_id) + return self._to_domain(row) if row else None + + def get_by_email(self, email: str) -> Optional[User]: + with self._provider.session() as session: + row = session.query(UserModel).filter_by(email=email).first() + return self._to_domain(row) if row else None + + def get_by_github_id(self, github_id: str) -> Optional[User]: + with self._provider.session() as session: + row = session.query(UserModel).filter_by(github_id=github_id).first() + return self._to_domain(row) if row else None + + def authenticate(self, email: str, password: str) -> Optional[User]: + """Verify credentials and update last_login in one session. Returns User or None.""" + now = datetime.now(timezone.utc) + with self._provider.session() as session: + row = session.query(UserModel).filter_by(email=email).first() + if not row or not row.is_active or not verify_password(password, row.hashed_password or ""): + return None + row.last_login_at = now + session.commit() + return self._to_domain(row) + + def create_email_user(self, *, email: str, hashed_password: str, display_name: Optional[str] = None) -> User: + now = datetime.now(timezone.utc) + with self._provider.session() as session: + row = UserModel( + email=email, + hashed_password=hashed_password, + display_name=display_name, + is_active=True, + created_at=now, + ) + session.add(row) + session.commit() + session.refresh(row) + return self._to_domain(row) + + def create_github_user(self, *, github_id: str, username: str, display_name: str, avatar_url: str) -> User: + now = datetime.now(timezone.utc) + with self._provider.session() as session: + row = UserModel( + github_id=github_id, + github_username=username, + display_name=display_name, + avatar_url=avatar_url, + is_active=True, + created_at=now, + last_login_at=now, + ) + session.add(row) + session.commit() + session.refresh(row) + return self._to_domain(row) + + def update_last_login(self, user_id: int) -> None: + now = datetime.now(timezone.utc) + with self._provider.session() as session: + session.query(UserModel).filter_by(id=user_id).update({"last_login_at": now}) + session.commit() + + def update_password(self, user_id: int, hashed_password: str) -> None: + with self._provider.session() as session: + session.query(UserModel).filter_by(id=user_id).update({"hashed_password": hashed_password}) + session.commit() + + def update_profile(self, user_id: int, display_name: Optional[str]) -> None: + with self._provider.session() as session: + session.query(UserModel).filter_by(id=user_id).update({"display_name": display_name}) + session.commit() + + def deactivate(self, user_id: int) -> None: + """Soft-delete: mark is_active=False so the user cannot log in.""" + with self._provider.session() as session: + session.query(UserModel).filter_by(id=user_id).update({"is_active": False}) + session.commit() + + def reactivate(self, user_id: int) -> None: + """Re-enable a previously deactivated account.""" + with self._provider.session() as session: + session.query(UserModel).filter_by(id=user_id).update({"is_active": True}) + session.commit() + + # ── Password reset tokens ──────────────────────────────────────────────── + + def create_reset_token(self, user_id: int) -> str: + """Generate a secure token, store it, and return the raw token string.""" + now = datetime.now(timezone.utc) + token = secrets.token_urlsafe(32) + with self._provider.session() as session: + row = PasswordResetTokenModel( + user_id=user_id, + token=token, + expires_at=now + timedelta(hours=1), + used=False, + created_at=now, + ) + session.add(row) + session.commit() + return token + + def get_valid_reset_token(self, token: str) -> Optional[PasswordResetTokenModel]: + """Return the token row if it exists, is unused, and has not expired.""" + now = datetime.now(timezone.utc) + with self._provider.session() as session: + row = ( + session.query(PasswordResetTokenModel) + .filter_by(token=token, used=False) + .first() + ) + if not row: + return None + # Ensure expires_at is timezone-aware for comparison + expires = row.expires_at + if expires.tzinfo is None: + expires = expires.replace(tzinfo=timezone.utc) + if expires < now: + return None + # Detach from session so caller can read attributes + session.expunge(row) + return row + + def consume_reset_token(self, token: str, new_hashed_password: str) -> bool: + """Mark the token as used and update the user's password atomically. + + Returns True on success, False if the token is invalid/expired. + """ + now = datetime.now(timezone.utc) + with self._provider.session() as session: + row = ( + session.query(PasswordResetTokenModel) + .filter_by(token=token, used=False) + .first() + ) + if not row: + return False + expires = row.expires_at + if expires.tzinfo is None: + expires = expires.replace(tzinfo=timezone.utc) + if expires < now: + return False + row.used = True + session.query(UserModel).filter_by(id=row.user_id).update( + {"hashed_password": new_hashed_password} + ) + session.commit() + return True diff --git a/tests/test_auth_api.py b/tests/test_auth_api.py new file mode 100644 index 00000000..269b6731 --- /dev/null +++ b/tests/test_auth_api.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +import os + +from fastapi.testclient import TestClient + +from paperbot.api.main import app + + +client = TestClient(app) + + +def _set_jwt_secret() -> None: + if not os.getenv("PAPERBOT_JWT_SECRET"): + os.environ["PAPERBOT_JWT_SECRET"] = "test-secret-key" + + +def test_register_login_me_roundtrip(tmp_path, monkeypatch): + """Happy path: register -> login -> /me returns user profile.""" + + _set_jwt_secret() + + email = "test_user@example.com" + password = "s3cretP@ss" + + # Register + r = client.post( + "/api/auth/register", + json={"email": email, "password": password, "display_name": "Tester"}, + ) + assert r.status_code == 201, r.text + data = r.json() + assert data["access_token"] + assert data["user_id"] > 0 + + # Login + r2 = client.post( + "/api/auth/login", + json={"email": email, "password": password}, + ) + assert r2.status_code == 200, r2.text + login_data = r2.json() + token = login_data["access_token"] + + # /me + r3 = client.get( + "/api/auth/me", + headers={"Authorization": f"Bearer {token}"}, + ) + assert r3.status_code == 200, r3.text + me = r3.json() + assert me["email"] == email + + +def test_login_wrong_password_rejected(): + _set_jwt_secret() + email = "wrong_pw@example.com" + password = "correctpass123" + + # First register + client.post( + "/api/auth/register", + json={"email": email, "password": password}, + ) + + # Then attempt login with wrong password + r = client.post( + "/api/auth/login", + json={"email": email, "password": "badpass"}, + ) + assert r.status_code == 401 + + +def test_register_duplicate_email_rejected(): + _set_jwt_secret() + email = "duplicate@example.com" + password = "password123" + + r1 = client.post( + "/api/auth/register", + json={"email": email, "password": password}, + ) + assert r1.status_code == 201, r1.text + + r2 = client.post( + "/api/auth/register", + json={"email": email, "password": password}, + ) + assert r2.status_code == 400 + From 9ba739a474443acf22dd916d72e3bbd751567fff Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Thu, 12 Mar 2026 11:03:46 +0100 Subject: [PATCH 2/2] fix(auth): address code review feedback on PR #365 --- src/paperbot/api/auth/email.py | 6 ++- src/paperbot/api/routes/auth.py | 7 +++- tests/test_auth_api.py | 74 ++++++++++++++++----------------- 3 files changed, 46 insertions(+), 41 deletions(-) diff --git a/src/paperbot/api/auth/email.py b/src/paperbot/api/auth/email.py index 61748df5..af751e8b 100644 --- a/src/paperbot/api/auth/email.py +++ b/src/paperbot/api/auth/email.py @@ -31,7 +31,11 @@ def _log_email(to_email: str, reset_url: str) -> None: def _send_via_resend(to_email: str, reset_url: str) -> None: # pragma: no cover import resend # type: ignore[import] - resend.api_key = os.environ["RESEND_API_KEY"] + api_key = os.getenv("RESEND_API_KEY") + if not api_key: + logger.error("Cannot send email via Resend: RESEND_API_KEY is not set.") + return + resend.api_key = api_key from_addr = os.getenv("EMAIL_FROM", "PaperBot ") resend.Emails.send({ diff --git a/src/paperbot/api/routes/auth.py b/src/paperbot/api/routes/auth.py index ba6159fb..ed04baa9 100644 --- a/src/paperbot/api/routes/auth.py +++ b/src/paperbot/api/routes/auth.py @@ -156,12 +156,15 @@ class UpdateMeRequest(BaseModel): @router.patch("/me", response_model=MeResponse) def update_me(req: UpdateMeRequest, current_user: User = Depends(get_current_user)): - _user_store.update_profile(current_user.id, req.display_name) + update_data = req.model_dump(exclude_unset=True) + if "display_name" in update_data: + _user_store.update_profile(current_user.id, update_data["display_name"]) + updated_user = _user_store.get_by_id(current_user.id) return MeResponse( id=current_user.id, email=current_user.email, github_username=current_user.github_username, - display_name=req.display_name, + display_name=updated_user.display_name if updated_user else current_user.display_name, avatar_url=current_user.avatar_url, ) diff --git a/tests/test_auth_api.py b/tests/test_auth_api.py index 269b6731..33dc5969 100644 --- a/tests/test_auth_api.py +++ b/tests/test_auth_api.py @@ -2,28 +2,45 @@ import os +import pytest from fastapi.testclient import TestClient -from paperbot.api.main import app +@pytest.fixture(autouse=True) +def isolated_db(tmp_path, monkeypatch): + """Give each test its own SQLite database so tests don't share state.""" + db_path = tmp_path / "test_auth.db" + monkeypatch.setenv("PAPERBOT_DB_URL", f"sqlite:///{db_path}") + monkeypatch.setenv("PAPERBOT_JWT_SECRET", "test-secret-key") -client = TestClient(app) + # Re-import app inside each test so stores pick up the new env vars + import importlib + import paperbot.infrastructure.stores.sqlalchemy_db as db_mod + import paperbot.infrastructure.stores.user_store as us_mod + import paperbot.api.routes.auth as auth_mod + importlib.reload(db_mod) + importlib.reload(us_mod) + importlib.reload(auth_mod) -def _set_jwt_secret() -> None: - if not os.getenv("PAPERBOT_JWT_SECRET"): - os.environ["PAPERBOT_JWT_SECRET"] = "test-secret-key" + # Run migrations on the fresh DB + from alembic.config import Config + from alembic import command + alembic_cfg = Config("alembic.ini") + alembic_cfg.set_main_option("sqlalchemy.url", f"sqlite:///{db_path}") + command.upgrade(alembic_cfg, "head") -def test_register_login_me_roundtrip(tmp_path, monkeypatch): - """Happy path: register -> login -> /me returns user profile.""" + from paperbot.api.main import app + yield app - _set_jwt_secret() +def test_register_login_me_roundtrip(isolated_db): + """Happy path: register -> login -> /me returns user profile.""" + client = TestClient(isolated_db) email = "test_user@example.com" password = "s3cretP@ss" - # Register r = client.post( "/api/auth/register", json={"email": email, "password": password, "display_name": "Tester"}, @@ -33,58 +50,39 @@ def test_register_login_me_roundtrip(tmp_path, monkeypatch): assert data["access_token"] assert data["user_id"] > 0 - # Login r2 = client.post( "/api/auth/login", json={"email": email, "password": password}, ) assert r2.status_code == 200, r2.text - login_data = r2.json() - token = login_data["access_token"] + token = r2.json()["access_token"] - # /me r3 = client.get( "/api/auth/me", headers={"Authorization": f"Bearer {token}"}, ) assert r3.status_code == 200, r3.text - me = r3.json() - assert me["email"] == email + assert r3.json()["email"] == email -def test_login_wrong_password_rejected(): - _set_jwt_secret() +def test_login_wrong_password_rejected(isolated_db): + client = TestClient(isolated_db) email = "wrong_pw@example.com" password = "correctpass123" - # First register - client.post( - "/api/auth/register", - json={"email": email, "password": password}, - ) + client.post("/api/auth/register", json={"email": email, "password": password}) - # Then attempt login with wrong password - r = client.post( - "/api/auth/login", - json={"email": email, "password": "badpass"}, - ) + r = client.post("/api/auth/login", json={"email": email, "password": "badpass"}) assert r.status_code == 401 -def test_register_duplicate_email_rejected(): - _set_jwt_secret() +def test_register_duplicate_email_rejected(isolated_db): + client = TestClient(isolated_db) email = "duplicate@example.com" password = "password123" - r1 = client.post( - "/api/auth/register", - json={"email": email, "password": password}, - ) + r1 = client.post("/api/auth/register", json={"email": email, "password": password}) assert r1.status_code == 201, r1.text - r2 = client.post( - "/api/auth/register", - json={"email": email, "password": password}, - ) + r2 = client.post("/api/auth/register", json={"email": email, "password": password}) assert r2.status_code == 400 -