Skip to content

Testing a FastAPI endpoint

Seed the graph in the test, hand the same session to the application through dependency_overrides, and call the endpoint: the application reads rows that were never committed, and the test rolls them back at the end.

The application

from typing import Annotated

from fastapi import Depends, FastAPI
from sqlalchemy import ForeignKey, create_engine, select
from sqlalchemy.orm import DeclarativeBase, Mapped, Session, mapped_column, relationship


class Base(DeclarativeBase):
    pass


class User(Base):
    __tablename__ = "users"
    id: Mapped[int] = mapped_column(primary_key=True)
    name: Mapped[str]
    posts: Mapped[list["Post"]] = relationship(back_populates="author")


class Post(Base):
    __tablename__ = "posts"
    id: Mapped[int] = mapped_column(primary_key=True)
    title: Mapped[str]
    author_id: Mapped[int] = mapped_column(ForeignKey("users.id"))
    author: Mapped[User] = relationship(back_populates="posts")


production_engine = create_engine("sqlite:///app.db")


def get_session():
    with Session(production_engine) as session:
        yield session


app = FastAPI()


@app.get("/users/{user_id}/posts")
def list_posts(user_id: int, session: Annotated[Session, Depends(get_session)]) -> list[dict]:
    posts = session.scalars(select(Post).where(Post.author_id == user_id).order_by(Post.id))
    return [{"id": post.id, "title": post.title} for post in posts]

The tests

TestClient runs synchronous endpoints in a worker thread. An in-memory SQLite gives each thread its own empty database unless the engine keeps a single connection, hence StaticPool:

import pytest
from fastapi.testclient import TestClient
from sqlalchemy.pool import StaticPool

from seedgraph import seed


@pytest.fixture()
def session():
    engine = create_engine("sqlite://", poolclass=StaticPool, connect_args={"check_same_thread": False})
    Base.metadata.create_all(engine)
    with Session(engine) as session:
        yield session
        session.rollback()
    engine.dispose()


@pytest.fixture()
def client(session):
    app.dependency_overrides[get_session] = lambda: session
    yield TestClient(app)
    app.dependency_overrides.clear()


def test_a_user_sees_only_their_own_posts(session, client):
    graph = seed(session, User, user=2, post=3)
    alice, bob = graph.users

    response = client.get(f"/users/{alice.id}/posts")

    assert response.status_code == 200
    assert [post["id"] for post in response.json()] == [post.id for post in alice.posts]
    assert not {post.id for post in bob.posts} & {post["id"] for post in response.json()}


def test_a_user_without_posts_gets_an_empty_list(session, client):
    graph = seed(session, User)

    assert client.get(f"/users/{graph.users[0].id}/posts").json() == []

Against PostgreSQL, drop StaticPool and check_same_thread: every thread goes through the same session object, so it sees the same transaction.