From 35adfadf31611fd8af40d7b596fa7daca56f64f9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Sebasti=C3=A1n=20Ram=C3=ADrez?= Date: Sat, 1 Nov 2025 17:24:33 +0100 Subject: [PATCH] =?UTF-8?q?=E2=9C=85=20Add=20contextvars=20black=20magic?= =?UTF-8?q?=20state=20to=20handle=20multiple=20tests=20changing=20the=20st?= =?UTF-8?q?ate?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../test_dependency_yield_scope_websockets.py | 26 ++++++++++++------- 1 file changed, 17 insertions(+), 9 deletions(-) diff --git a/tests/test_dependency_yield_scope_websockets.py b/tests/test_dependency_yield_scope_websockets.py index c85648574..2bf9b19d4 100644 --- a/tests/test_dependency_yield_scope_websockets.py +++ b/tests/test_dependency_yield_scope_websockets.py @@ -1,3 +1,4 @@ +from contextvars import ContextVar from typing import Any, Dict, Tuple import pytest @@ -6,7 +7,7 @@ from fastapi.exceptions import FastAPIError from fastapi.testclient import TestClient from typing_extensions import Annotated -global_state: Dict[str, Any] = {} +global_context: ContextVar[Dict[str, Any]] = ContextVar("global_context", default={}) # noqa: B039 class Session: @@ -18,6 +19,7 @@ def dep_session() -> Any: s = Session() yield s s.open = False + global_state = global_context.get() global_state["session_closed"] = True @@ -37,6 +39,7 @@ def get_named_session(session: SessionRequestDep, session_b: SessionDefaultDep) named_session = NamedSession(name="named") yield named_session, session_b named_session.open = False + global_state = global_context.get() global_state["named_session_closed"] = True @@ -47,6 +50,7 @@ def get_named_func_session(session: SessionFuncDep) -> Any: named_session = NamedSession(name="named") yield named_session, session named_session.open = False + global_state = global_context.get() global_state["named_func_session_closed"] = True @@ -125,7 +129,8 @@ client = TestClient(app) def test_function_scope() -> None: - global_state["session_closed"] = False + global_context.set({}) + global_state = global_context.get() with client.websocket_connect("/function-scope") as websocket: data = websocket.receive_json() assert data["is_open"] is True @@ -133,7 +138,8 @@ def test_function_scope() -> None: def test_request_scope() -> None: - global_state["session_closed"] = False + global_context.set({}) + global_state = global_context.get() with client.websocket_connect("/request-scope") as websocket: data = websocket.receive_json() assert data["is_open"] is True @@ -141,7 +147,8 @@ def test_request_scope() -> None: def test_two_scopes() -> None: - global_state["session_closed"] = False + global_context.set({}) + global_state = global_context.get() with client.websocket_connect("/two-scopes") as websocket: data = websocket.receive_json() assert data["func_is_open"] is True @@ -150,8 +157,8 @@ def test_two_scopes() -> None: def test_sub() -> None: - global_state["session_closed"] = False - global_state["named_session_closed"] = False + global_context.set({}) + global_state = global_context.get() with client.websocket_connect("/sub") as websocket: data = websocket.receive_json() assert data["named_session_open"] is True @@ -174,8 +181,8 @@ def test_broken_scope() -> None: def test_named_function_scope() -> None: - global_state["session_closed"] = False - global_state["named_func_session_closed"] = False + global_context.set({}) + global_state = global_context.get() with client.websocket_connect("/named-function-scope") as websocket: data = websocket.receive_json() assert data["named_session_open"] is True @@ -185,7 +192,8 @@ def test_named_function_scope() -> None: def test_regular_function_scope() -> None: - global_state["session_closed"] = False + global_context.set({}) + global_state = global_context.get() with client.websocket_connect("/regular-function-scope") as websocket: data = websocket.receive_json() assert data["named_session_open"] is True