diff --git a/backend/app/auth/users.py b/backend/app/auth/users.py index 5f731df..08a88b6 100644 --- a/backend/app/auth/users.py +++ b/backend/app/auth/users.py @@ -2,7 +2,7 @@ from collections.abc import AsyncGenerator from typing import Annotated -from fastapi import Depends +from fastapi import Depends, Request from fastapi_users import BaseUserManager, FastAPIUsers, UUIDIDMixin from fastapi_users.db import BaseUserDatabase from fastapi_users_db_sqlalchemy import SQLAlchemyUserDatabase @@ -11,6 +11,7 @@ from app.auth.backend import auth_backend from app.config import get_settings from app.database import get_session +from app.models.household import Household from app.models.user import User @@ -23,16 +24,43 @@ async def get_user_db( class UserManager(UUIDIDMixin, BaseUserManager[User, uuid.UUID]): - # These secrets sign password-reset and verification tokens. - # Changing them invalidates all outstanding tokens. reset_password_token_secret = get_settings().secret_key.get_secret_value() verification_token_secret = get_settings().secret_key.get_secret_value() + def __init__( + self, + user_db: BaseUserDatabase[User, uuid.UUID], + session: AsyncSession, + ) -> None: + super().__init__(user_db) + self.session = session + + async def on_after_register( + self, user: User, request: Request | None = None + ) -> None: + # FastAPI Users commits the user row before calling this hook, so true + # atomicity isn't possible. We use a compensating transaction: if + # household creation fails for any reason, we delete the orphaned user + # so the DB is left in a consistent state. + try: + prefix = user.email.split("@", 1)[0].strip() or "User" + household = Household(name=f"{prefix}'s household") + self.session.add(household) + await self.session.flush() # writes household row and populates household.id + user.household_id = household.id # session tracks this change automatically + await self.session.commit() + except Exception: + await self.user_db.delete(user) + raise + async def get_user_manager( user_db: Annotated[BaseUserDatabase[User, uuid.UUID], Depends(get_user_db)], + session: Annotated[AsyncSession, Depends(get_session)], ) -> AsyncGenerator[UserManager]: - yield UserManager(user_db) + # FastAPI caches dependencies per request, so session here is the same + # instance already open for get_user_db — no extra connection is opened. + yield UserManager(user_db, session) fastapi_users: FastAPIUsers[User, uuid.UUID] = FastAPIUsers( diff --git a/backend/app/database.py b/backend/app/database.py index 135aa72..f432fb3 100644 --- a/backend/app/database.py +++ b/backend/app/database.py @@ -9,7 +9,9 @@ DATABASE_URL = settings.database_url.get_secret_value() -engine = create_async_engine(DATABASE_URL, echo=settings.environment == "development") +# pool_pre_ping=True tests each connection before use and reconnects if Neon +# closed it due to inactivity (Neon drops idle connections after ~5 minutes). +engine = create_async_engine(DATABASE_URL, echo=settings.environment == "development", pool_pre_ping=True) # keeps objects usable after commit (relevant for async sessions) _session_factory = async_sessionmaker(engine, expire_on_commit=False) diff --git a/backend/app/models/household.py b/backend/app/models/household.py index 4c95a19..05a118a 100644 --- a/backend/app/models/household.py +++ b/backend/app/models/household.py @@ -1,10 +1,15 @@ from datetime import UTC, datetime +import sqlalchemy as sa from sqlmodel import Field, SQLModel class Household(SQLModel, table=True): __tablename__: str = "households" + id: int | None = Field(default=None, primary_key=True) name: str - created_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) + created_at: datetime = Field( + default_factory=lambda: datetime.now(UTC), + sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False), + ) diff --git a/backend/app/models/user.py b/backend/app/models/user.py index 247466a..e8123ab 100644 --- a/backend/app/models/user.py +++ b/backend/app/models/user.py @@ -1,6 +1,7 @@ import uuid from datetime import UTC, datetime +import sqlalchemy as sa from fastapi_users import schemas from sqlmodel import Field, SQLModel @@ -18,11 +19,13 @@ class User(SQLModel, table=True): is_active: bool = Field(default=True) is_superuser: bool = Field(default=False) is_verified: bool = Field(default=False) - # Nullable until the on_after_register hook creates and assigns a household. household_id: int | None = Field( - default=None, foreign_key="household.id", index=True + default=None, foreign_key="households.id", index=True + ) + created_at: datetime = Field( + default_factory=lambda: datetime.now(UTC), + sa_column=sa.Column(sa.DateTime(timezone=True), nullable=False), ) - created_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) class UserRead(schemas.BaseUser[uuid.UUID]):