Browse Source

🐛 Refactor router route building to make it thread-safe, mainly relevant for tests running in parallel threads (uncommon) (#16013)

pull/16014/head
Sebastián Ramírez 5 days ago
committed by GitHub
parent
commit
7fe315c21a
No known key found for this signature in database GPG Key ID: B5690EEEBB952194
  1. 79
      fastapi/routing.py
  2. 62
      tests/test_router_include_context.py

79
fastapi/routing.py

@ -7,6 +7,7 @@ import inspect
import json import json
import os import os
import stat import stat
import threading
import types import types
from collections.abc import ( from collections.abc import (
AsyncIterator, AsyncIterator,
@ -1571,6 +1572,9 @@ class RouteContext:
class _IncludedRouter(BaseRoute): class _IncludedRouter(BaseRoute):
original_router: "APIRouter" original_router: "APIRouter"
include_context: _RouterIncludeContext include_context: _RouterIncludeContext
_effective_routes_lock: Any = field(
default_factory=threading.Lock, repr=False, compare=False
)
_effective_candidates: list["_EffectiveRouteContext | _IncludedRouter"] = field( _effective_candidates: list["_EffectiveRouteContext | _IncludedRouter"] = field(
default_factory=list default_factory=list
) )
@ -1584,44 +1588,53 @@ class _IncludedRouter(BaseRoute):
routes_version = self.original_router._get_routes_version() routes_version = self.original_router._get_routes_version()
if routes_version == self._effective_candidates_version: if routes_version == self._effective_candidates_version:
return self._effective_candidates return self._effective_candidates
self._effective_candidates = [] with self._effective_routes_lock:
candidates = self.original_router.routes routes_version = self.original_router._get_routes_version()
for route in candidates: if routes_version == self._effective_candidates_version:
if isinstance(route, _IncludedRouter): return self._effective_candidates
child_context = self.include_context.combine(route.include_context) effective_candidates: list[_EffectiveRouteContext | _IncludedRouter] = []
child_branch = _IncludedRouter( for route in self.original_router.routes:
original_router=route.original_router, if isinstance(route, _IncludedRouter):
include_context=child_context, child_context = self.include_context.combine(route.include_context)
) child_branch = _IncludedRouter(
self._effective_candidates.append(child_branch) original_router=route.original_router,
continue include_context=child_context,
route_context = self._build_effective_context(route) )
if route_context is not None: effective_candidates.append(child_branch)
self._effective_candidates.append(route_context) continue
self._effective_candidates_version = routes_version route_context = self._build_effective_context(route)
return self._effective_candidates if route_context is not None:
effective_candidates.append(route_context)
self._effective_candidates = effective_candidates
self._effective_candidates_version = routes_version
return effective_candidates
def effective_low_priority_routes(self) -> list["_EffectiveRouteContext"]: def effective_low_priority_routes(self) -> list["_EffectiveRouteContext"]:
routes_version = self.original_router._get_routes_version() routes_version = self.original_router._get_routes_version()
if routes_version == self._effective_low_priority_routes_version: if routes_version == self._effective_low_priority_routes_version:
return self._effective_low_priority_routes return self._effective_low_priority_routes
self._effective_low_priority_routes = [] with self._effective_routes_lock:
for route in self.original_router._low_priority_routes: routes_version = self.original_router._get_routes_version()
route_context = self._build_effective_context(route) if routes_version == self._effective_low_priority_routes_version:
if route_context is not None: return self._effective_low_priority_routes
self._effective_low_priority_routes.append(route_context) effective_low_priority_routes: list[_EffectiveRouteContext] = []
for route in self.original_router.routes: for route in self.original_router._low_priority_routes:
if isinstance(route, _IncludedRouter): route_context = self._build_effective_context(route)
child_context = self.include_context.combine(route.include_context) if route_context is not None:
child_branch = _IncludedRouter( effective_low_priority_routes.append(route_context)
original_router=route.original_router, for route in self.original_router.routes:
include_context=child_context, if isinstance(route, _IncludedRouter):
) child_context = self.include_context.combine(route.include_context)
self._effective_low_priority_routes.extend( child_branch = _IncludedRouter(
child_branch.effective_low_priority_routes() original_router=route.original_router,
) include_context=child_context,
self._effective_low_priority_routes_version = routes_version )
return self._effective_low_priority_routes effective_low_priority_routes.extend(
child_branch.effective_low_priority_routes()
)
self._effective_low_priority_routes = effective_low_priority_routes
self._effective_low_priority_routes_version = routes_version
return effective_low_priority_routes
def _build_effective_context( def _build_effective_context(
self, route: BaseRoute self, route: BaseRoute

62
tests/test_router_include_context.py

@ -1,3 +1,5 @@
import threading
from collections.abc import Sequence
from typing import Annotated, cast from typing import Annotated, cast
import pytest import pytest
@ -898,6 +900,66 @@ async def test_included_unknown_route_is_ignored_and_can_return_default_404():
assert messages[0]["status"] == 404 assert messages[0]["status"] == 404
def test_included_router_candidate_cache_is_thread_safe():
router = APIRouter()
route_count = 120
thread_count = 6
for index in range(route_count):
@router.get(f"/items/{index}")
def read_item(index: int = index):
return {"index": index}
app = FastAPI()
app.include_router(router)
included_router = cast(_IncludedRouter, app.router.routes[-1])
barrier = threading.Barrier(thread_count)
results: list[Sequence[object]] = []
def build_candidates() -> None:
barrier.wait()
results.append(included_router.effective_candidates())
threads = [threading.Thread(target=build_candidates) for _ in range(thread_count)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
assert len(results) == thread_count
assert all(result is results[0] for result in results)
assert len(results[0]) == route_count
assert TestClient(app).get("/items/0").json() == {"index": 0}
def test_included_router_low_priority_cache_rechecks_version_after_lock(monkeypatch):
router = APIRouter()
app = FastAPI()
app.include_router(router)
included_router = cast(_IncludedRouter, app.router.routes[-1])
routes_version = router._get_routes_version()
version_checked = threading.Event()
def get_routes_version() -> int:
version_checked.set()
return routes_version
monkeypatch.setattr(router, "_get_routes_version", get_routes_version)
result: list[Sequence[object]] = []
def build_low_priority_routes() -> None:
result.append(included_router.effective_low_priority_routes())
with included_router._effective_routes_lock:
thread = threading.Thread(target=build_low_priority_routes)
thread.start()
assert version_checked.wait(timeout=1)
included_router._effective_low_priority_routes_version = routes_version
thread.join()
assert result == [[]]
def test_no_prefix_include_validation_sees_effective_starlette_route_candidates(): def test_no_prefix_include_validation_sees_effective_starlette_route_candidates():
def endpoint(request): # pragma: no cover def endpoint(request): # pragma: no cover
return PlainTextResponse("ok") return PlainTextResponse("ok")

Loading…
Cancel
Save