Skip to content

Commit 6a895a6

Browse files
committed
feat: support starlette 1.0+
1 parent bc4f971 commit 6a895a6

7 files changed

Lines changed: 267 additions & 18 deletions

File tree

.github/workflows/ci.yml

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -56,17 +56,21 @@ jobs:
5656
run: make test_sqlite_regexp
5757
env:
5858
PYTHONDEVMODE: 1
59-
- name: Test FastAPI/Blacksheep/Sanic Examples
59+
- name: Test FastAPI/Blacksheep/Sanic/Starlette Examples
6060
run: |
6161
PYTHONPATH=$DEST_FASTAPI uv run --frozen tortoise -c config.TORTOISE_ORM migrate
6262
PYTHONPATH=$DEST_FASTAPI uv run --frozen pytest $PYTEST_ARGS_SEQ $DEST_FASTAPI/_tests.py
6363
rm -f $DEST_FASTAPI/db.sqlite3
6464
PYTHONPATH=$DEST_BLACKSHEEP uv run --frozen pytest $PYTEST_ARGS $DEST_BLACKSHEEP/_tests.py
6565
PYTHONPATH=$DEST_SANIC uv run --frozen pytest $PYTEST_ARGS $DEST_SANIC/_tests.py
66+
PYTHONPATH=$DEST_STARLETTE uv run --frozen pytest $PYTEST_ARGS $DEST_STARLETTE/_tests.py
67+
uv pip install "starlette<1.0"
68+
PYTHONPATH=$DEST_STARLETTE uv run --no-sync pytest $PYTEST_ARGS $DEST_STARLETTE/_tests.py
6669
env:
6770
DEST_FASTAPI: examples/fastapi
6871
DEST_BLACKSHEEP: examples/blacksheep
6972
DEST_SANIC: examples/sanic
73+
DEST_STARLETTE: examples/starlette
7074
PYTHONDEVMODE: 1
7175
PYTEST_ARGS: "-n auto --cov=tortoise --cov-append --cov-branch --tb=native -q"
7276
PYTEST_ARGS_SEQ: "--cov=tortoise --cov-append --cov-branch --tb=native -q"

examples/starlette/_tests.py

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
from pathlib import Path
2+
3+
import pytest
4+
from asgi_lifespan import LifespanManager
5+
from httpx import ASGITransport, AsyncClient
6+
7+
try:
8+
from main import app
9+
from models import Users
10+
except ImportError:
11+
if (cwd := Path.cwd()) == (parent := Path(__file__).parent):
12+
dirpath = "."
13+
else:
14+
dirpath = str(parent.relative_to(cwd))
15+
print(f"You may need to explicitly declare python path:\n\nexport PYTHONPATH={dirpath}\n")
16+
raise
17+
18+
19+
@pytest.fixture(scope="module")
20+
def anyio_backend() -> str:
21+
return "asyncio"
22+
23+
24+
@pytest.mark.anyio
25+
async def test_app() -> None:
26+
async with LifespanManager(app):
27+
transport = ASGITransport(app=app)
28+
# note: you _must_ set `base_url` for relative urls like "/" to work
29+
async with AsyncClient(transport=transport, base_url="http://testserver") as client:
30+
r = await client.get("/")
31+
assert r.status_code == 200
32+
assert r.json() == {"users": []}
33+
assert await Users.all() == []
34+
35+
r = await client.post("/user/", json={"username": "Iron"})
36+
assert r.status_code == 201
37+
assert r.json() == {"user": "Users(id=1, username='Iron')"}
38+
assert await Users.get(id=1) == await Users.last()
39+
40+
r = await client.get("/")
41+
assert r.status_code == 200
42+
assert r.json() == {"users": ["User 1: Iron"]}
43+
assert await Users.all() == [await Users.first()]

examples/starlette/main.py

Lines changed: 19 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,16 @@
1+
#!/usr/bin/env python
12
# pylint: disable=E0401,E0611
23
import logging
34
from json import JSONDecodeError
5+
from pathlib import Path
46

57
from models import Users
68
from starlette.applications import Starlette
79
from starlette.exceptions import HTTPException
810
from starlette.requests import Request
911
from starlette.responses import JSONResponse
12+
from starlette.routing import Mount, Route
1013
from starlette.status import HTTP_201_CREATED, HTTP_400_BAD_REQUEST
11-
from uvicorn.main import run
1214

1315
from tortoise.contrib.starlette import register_tortoise
1416

@@ -17,29 +19,39 @@
1719
app = Starlette()
1820

1921

20-
@app.route("/", methods=["GET"])
2122
async def list_all(_: Request) -> JSONResponse:
2223
users = await Users.all()
2324
return JSONResponse({"users": [str(user) for user in users]})
2425

2526

26-
@app.route("/user", methods=["POST"])
2727
async def add_user(request: Request) -> JSONResponse:
2828
try:
2929
payload = await request.json()
3030
username = payload["username"]
3131
except JSONDecodeError:
32-
raise HTTPException(status_code=HTTP_400_BAD_REQUEST, detail="cannot parse request body")
32+
raise HTTPException(
33+
status_code=HTTP_400_BAD_REQUEST, detail="cannot parse request body"
34+
) from None
3335
except KeyError:
34-
raise HTTPException(status_code=HTTP_400_BAD_REQUEST, detail="username is required")
36+
raise HTTPException(
37+
status_code=HTTP_400_BAD_REQUEST, detail="username is required"
38+
) from None
3539

3640
user = await Users.create(username=username)
37-
return JSONResponse({"user": str(user)}, status_code=HTTP_201_CREATED)
41+
return JSONResponse({"user": repr(user)}, status_code=HTTP_201_CREATED)
3842

3943

44+
app = Starlette(
45+
routes=[
46+
Route("/", list_all),
47+
Mount("/user", routes=[Route("/", add_user, methods=["POST"])]),
48+
]
49+
)
4050
register_tortoise(
4151
app, db_url="sqlite://:memory:", modules={"models": ["models"]}, generate_schemas=True
4252
)
4353

4454
if __name__ == "__main__":
45-
run(app)
55+
import uvicorn
56+
57+
uvicorn.run("__main__:app", reload=True, reload_dirs=[str(Path(__file__).parent)])

examples/starlette/models.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,3 +7,8 @@ class Users(models.Model):
77

88
def __str__(self) -> str:
99
return f"User {self.id}: {self.username}"
10+
11+
def __repr__(self) -> str:
12+
fields = sorted(self._meta.db_fields) # ['id', 'username']
13+
values = ", ".join(f"{f}={getattr(self, f)!r}" for f in fields)
14+
return f"{self.__class__.__name__}({values})"

tests/contrib/test_starlette.py

Lines changed: 121 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,121 @@
1+
from __future__ import annotations
2+
3+
from contextlib import asynccontextmanager
4+
5+
import pytest
6+
from asgi_lifespan import LifespanManager
7+
from httpx import ASGITransport, AsyncClient
8+
from starlette.applications import Starlette
9+
from starlette.requests import Request
10+
from starlette.responses import JSONResponse
11+
from starlette.routing import Route
12+
13+
from tortoise import fields, models
14+
from tortoise.context import TortoiseContext, get_current_context
15+
from tortoise.contrib.starlette import TortoiseContextMiddleware, register_tortoise
16+
17+
18+
class StarletteMiddlewareModel(models.Model):
19+
id = fields.IntField(primary_key=True)
20+
21+
22+
def _register_tortoise(app: Starlette) -> None:
23+
register_tortoise(
24+
app,
25+
db_url="sqlite://:memory:",
26+
modules={"models": [__name__]},
27+
)
28+
29+
30+
def _register_tortoise_with_lifespan(app: Starlette) -> None:
31+
register_tortoise(
32+
app,
33+
db_url="sqlite://:memory:",
34+
modules={"models": [__name__]},
35+
generate_schemas=True,
36+
)
37+
38+
39+
async def _context_endpoint(_: Request) -> JSONResponse:
40+
ctx = get_current_context()
41+
return JSONResponse({"inited": ctx.inited if ctx is not None else False})
42+
43+
44+
async def _get(client: AsyncClient) -> dict:
45+
response = await client.get("/")
46+
assert response.status_code == 200
47+
return response.json()
48+
49+
50+
@pytest.mark.asyncio
51+
async def test_starlette_register_tortoise_uses_middleware_without_patching_endpoint() -> None:
52+
route = Route("/", _context_endpoint)
53+
app = Starlette(routes=[route])
54+
original_endpoint = route.endpoint
55+
56+
_register_tortoise(app)
57+
58+
assert route.endpoint is original_endpoint
59+
assert any(
60+
middleware.cls is TortoiseContextMiddleware # type:ignore[comparison-overlap]
61+
for middleware in app.user_middleware
62+
)
63+
64+
transport = ASGITransport(app=app)
65+
async with AsyncClient(transport=transport, base_url="http://testserver") as client:
66+
assert await _get(client) == {"inited": True}
67+
68+
69+
@pytest.mark.asyncio
70+
async def test_starlette_middleware_covers_routes_added_after_register_tortoise() -> None:
71+
app = Starlette()
72+
_register_tortoise(app)
73+
app.add_route("/", _context_endpoint)
74+
75+
transport = ASGITransport(app=app)
76+
async with AsyncClient(transport=transport, base_url="http://testserver") as client:
77+
assert await _get(client) == {"inited": True}
78+
79+
80+
@pytest.mark.asyncio
81+
async def test_starlette_middleware_does_not_reuse_unrelated_global_context() -> None:
82+
unrelated_ctx = TortoiseContext()
83+
unrelated_ctx.__enter__()
84+
await unrelated_ctx.init(
85+
db_url="sqlite://:memory:",
86+
modules={"models": [__name__]},
87+
_enable_global_fallback=True,
88+
)
89+
unrelated_ctx.__exit__(None, None, None)
90+
91+
async def endpoint(_: Request) -> JSONResponse:
92+
return JSONResponse({"unrelated": get_current_context() is unrelated_ctx})
93+
94+
try:
95+
app = Starlette(routes=[Route("/", endpoint)])
96+
_register_tortoise(app)
97+
98+
transport = ASGITransport(app=app)
99+
async with AsyncClient(transport=transport, base_url="http://testserver") as client:
100+
assert await _get(client) == {"unrelated": False}
101+
finally:
102+
await unrelated_ctx.close_connections()
103+
104+
105+
@pytest.mark.asyncio
106+
async def test_starlette_lifespan_context_is_stored_and_cleared() -> None:
107+
@asynccontextmanager
108+
async def lifespan(app: Starlette):
109+
yield
110+
111+
app = Starlette(routes=[Route("/", _context_endpoint)], lifespan=lifespan)
112+
_register_tortoise_with_lifespan(app)
113+
114+
async with LifespanManager(app):
115+
assert hasattr(app.state, "_tortoise_context")
116+
117+
transport = ASGITransport(app=app)
118+
async with AsyncClient(transport=transport, base_url="http://testserver") as client:
119+
assert await _get(client) == {"inited": True}
120+
121+
assert not hasattr(app.state, "_tortoise_context")

tortoise/contrib/starlette/__init__.py

Lines changed: 69 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,18 +1,53 @@
11
from __future__ import annotations
22

3-
from collections.abc import Iterable
3+
from collections.abc import AsyncIterator, Iterable, Mapping
4+
from contextlib import asynccontextmanager
45
from types import ModuleType
6+
from typing import Any, cast
57

8+
import starlette
69
from starlette.applications import Starlette # pylint: disable=E0401
10+
from starlette.routing import _DefaultLifespan as StarletteDefaultLifespan
11+
from starlette.types import ASGIApp, Lifespan, Receive, Scope, Send
712

813
from tortoise import Tortoise
14+
from tortoise.config import TortoiseConfig
915
from tortoise.connection import get_connections
16+
from tortoise.context import TortoiseContext, get_current_context
1017
from tortoise.log import logger
1118

19+
_TORTOISE_CONTEXT_STATE = "_tortoise_context"
20+
21+
22+
class TortoiseContextMiddleware:
23+
def __init__(self, app: ASGIApp, config: TortoiseConfig, starlette_app: Starlette) -> None:
24+
self.app = app
25+
self.config = config
26+
self.starlette_app = starlette_app
27+
28+
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
29+
if scope["type"] not in {"http", "websocket"}:
30+
await self.app(scope, receive, send)
31+
return
32+
33+
app_context = getattr(self.starlette_app.state, _TORTOISE_CONTEXT_STATE, None)
34+
if app_context is None:
35+
async with TortoiseContext() as request_context:
36+
await request_context.init(self.config)
37+
await self.app(scope, receive, send)
38+
return
39+
40+
if get_current_context() is not app_context:
41+
with app_context:
42+
await self.app(scope, receive, send)
43+
return
44+
45+
await self.app(scope, receive, send)
46+
1247

1348
def register_tortoise(
1449
app: Starlette,
15-
config: dict | None = None,
50+
config: dict[str, Any] | TortoiseConfig | None = None,
1651
config_file: str | None = None,
1752
db_url: str | None = None,
1853
modules: dict[str, Iterable[str | ModuleType]] | None = None,
@@ -79,16 +114,45 @@ def register_tortoise(
79114
ConfigurationError
80115
For any configuration error
81116
"""
117+
typed_config = TortoiseConfig.resolve_args(config, config_file, db_url, modules)
82118

83-
@app.on_event("startup")
84119
async def init_orm() -> None: # pylint: disable=W0612
85-
await Tortoise.init(config=config, config_file=config_file, db_url=db_url, modules=modules)
120+
ctx = await Tortoise.init(config=typed_config, _enable_global_fallback=True)
121+
setattr(app.state, _TORTOISE_CONTEXT_STATE, ctx)
86122
logger.info("Tortoise-ORM started, %s, %s", get_connections()._get_storage(), Tortoise.apps)
87123
if generate_schemas:
88124
logger.info("Tortoise-ORM generating schema")
89125
await Tortoise.generate_schemas()
90126

91-
@app.on_event("shutdown")
92127
async def close_orm() -> None: # pylint: disable=W0612
93128
await Tortoise.close_connections()
129+
if hasattr(app.state, _TORTOISE_CONTEXT_STATE):
130+
delattr(app.state, _TORTOISE_CONTEXT_STATE)
94131
logger.info("Tortoise-ORM shutdown")
132+
133+
if starlette.__version__ < "1":
134+
if (on_event := getattr(app, "on_event", None)) is not None:
135+
on_event("startup")(init_orm)
136+
on_event("shutdown")(close_orm)
137+
return
138+
139+
original_lifespan = app.router.lifespan_context
140+
141+
if generate_schemas or not isinstance(original_lifespan, StarletteDefaultLifespan):
142+
143+
@asynccontextmanager
144+
async def orm_inited_lifespan(app_: Starlette) -> AsyncIterator[Mapping[str, Any] | None]:
145+
await init_orm()
146+
try:
147+
async with original_lifespan(app_) as maybe_state:
148+
yield maybe_state
149+
finally:
150+
await close_orm()
151+
152+
app.router.lifespan_context = cast("Lifespan[Any]", orm_inited_lifespan)
153+
154+
if not any(
155+
middleware.cls is TortoiseContextMiddleware # type:ignore[comparison-overlap]
156+
for middleware in app.user_middleware
157+
):
158+
app.add_middleware(TortoiseContextMiddleware, config=typed_config, starlette_app=app)

uv.lock

Lines changed: 5 additions & 5 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)