Coverage for api/utils/aiohttp.py: 79%
63 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 06:14 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 06:14 +0000
1import asyncio
2import time
3import weakref
5import aiohttp
6import structlog
7from django_asgi_lifespan.signals import asgi_shutdown
10logger = structlog.get_logger(__name__)
13_SESSIONS: weakref.WeakKeyDictionary[
14 asyncio.AbstractEventLoop, aiohttp.ClientSession
15] = weakref.WeakKeyDictionary()
17_LOCKS: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock] = (
18 weakref.WeakKeyDictionary()
19)
22@asgi_shutdown.connect
23async def _close_sessions(sender, **kwargs):
24 logger.debug("Closing aiohttp sessions on application shutdown")
26 closed_sessions = 0
28 while _SESSIONS:
29 loop, session = _SESSIONS.popitem()
30 try:
31 await session.close()
32 closed_sessions += 1
33 except BaseException as exc:
34 logger.error("Error closing sessions", exc=exc, exc_info=True)
36 logger.debug("Successfully closed %s session(s)", closed_sessions)
39async def get_aiohttp_session() -> aiohttp.ClientSession:
40 """
41 Safely retrieve a shared aiohttp session for the current event loop.
43 If the loop already has an aiohttp session associated, it will be reused.
44 If the loop has not yet had an aiohttp session created for it, a new one
45 will be created and returned.
47 While the main application will always run in the same loop, and while
48 that covers 99% of our use cases, it is still possible for `async_to_sync`
49 to cause a new loop to be created if, for example, `force_new_loop` is
50 passed. In order to prevent surprises should that ever be the case, this
51 function assumes that it's possible for multiple loops to be present in
52 the lifetime of the application and therefore we need to verify that each
53 loop gets its own session.
54 """
56 loop = asyncio.get_running_loop()
58 if loop not in _LOCKS:
59 _LOCKS[loop] = asyncio.Lock()
61 async with _LOCKS[loop]:
62 if loop not in _SESSIONS:
63 create_session = True
64 msg = "No session for loop. Creating new session."
65 elif _SESSIONS[loop].closed: 65 ↛ 66line 65 didn't jump to line 66 because the condition on line 65 was never true
66 create_session = True
67 msg = "Loop's previous session closed. Creating new session."
68 else:
69 create_session = False
70 msg = "Reusing existing session for loop."
72 logger.info(msg)
74 if create_session:
75 session = aiohttp.ClientSession(trace_configs=[LogTiming()])
76 _SESSIONS[loop] = session
78 return _SESSIONS[loop]
81class LogTiming(aiohttp.TraceConfig):
82 TIMEOUT_STATUS = -2
83 ERROR_STATUS = -1
85 def __init__(self, *args, **kwargs):
86 super().__init__(*args, **kwargs)
88 self.on_request_start.append(self._start_timing)
89 self.on_request_end.append(self._log_timing)
90 self.on_request_exception.append(self._log_timing)
92 async def _start_timing(
93 self,
94 session: aiohttp.ClientSession,
95 trace_config_ctx,
96 params: aiohttp.TraceRequestStartParams,
97 ):
98 trace_config_ctx.start_time = time.perf_counter()
100 async def _log_timing(
101 self,
102 session: aiohttp.ClientSession,
103 trace_config_ctx,
104 params: aiohttp.TraceRequestEndParams | aiohttp.TraceRequestExceptionParams,
105 ):
106 """Log timing at end of request or when request results in an exception."""
108 if not ( 108 ↛ 112line 108 didn't jump to line 112 because the condition on line 108 was never true
109 trace_config_ctx.trace_request_ctx
110 and "timing_event_name" in trace_config_ctx.trace_request_ctx
111 ):
112 return
114 end_time = time.perf_counter()
115 start_time = trace_config_ctx.start_time
117 request_ctx = trace_config_ctx.trace_request_ctx
119 if hasattr(params, "response"): 119 ↛ 122line 119 didn't jump to line 122 because the condition on line 119 was always true
120 # request end
121 status = params.response.status
122 elif isinstance(params.exception, aiohttp.ClientResponseError):
123 status = params.exception.status
124 elif isinstance(params.exception, asyncio.TimeoutError):
125 status = self.TIMEOUT_STATUS
126 else:
127 status = self.ERROR_STATUS
129 logger.info(
130 request_ctx["timing_event_name"],
131 status=status,
132 time=end_time - start_time,
133 url=str(params.url),
134 **request_ctx.get("timing_event_ctx", {}),
135 )