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

1import asyncio 

2import time 

3import weakref 

4 

5import aiohttp 

6import structlog 

7from django_asgi_lifespan.signals import asgi_shutdown 

8 

9 

10logger = structlog.get_logger(__name__) 

11 

12 

13_SESSIONS: weakref.WeakKeyDictionary[ 

14 asyncio.AbstractEventLoop, aiohttp.ClientSession 

15] = weakref.WeakKeyDictionary() 

16 

17_LOCKS: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, asyncio.Lock] = ( 

18 weakref.WeakKeyDictionary() 

19) 

20 

21 

22@asgi_shutdown.connect 

23async def _close_sessions(sender, **kwargs): 

24 logger.debug("Closing aiohttp sessions on application shutdown") 

25 

26 closed_sessions = 0 

27 

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) 

35 

36 logger.debug("Successfully closed %s session(s)", closed_sessions) 

37 

38 

39async def get_aiohttp_session() -> aiohttp.ClientSession: 

40 """ 

41 Safely retrieve a shared aiohttp session for the current event loop. 

42 

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. 

46 

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 """ 

55 

56 loop = asyncio.get_running_loop() 

57 

58 if loop not in _LOCKS: 

59 _LOCKS[loop] = asyncio.Lock() 

60 

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." 

71 

72 logger.info(msg) 

73 

74 if create_session: 

75 session = aiohttp.ClientSession(trace_configs=[LogTiming()]) 

76 _SESSIONS[loop] = session 

77 

78 return _SESSIONS[loop] 

79 

80 

81class LogTiming(aiohttp.TraceConfig): 

82 TIMEOUT_STATUS = -2 

83 ERROR_STATUS = -1 

84 

85 def __init__(self, *args, **kwargs): 

86 super().__init__(*args, **kwargs) 

87 

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) 

91 

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() 

99 

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.""" 

107 

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 

113 

114 end_time = time.perf_counter() 

115 start_time = trace_config_ctx.start_time 

116 

117 request_ctx = trace_config_ctx.trace_request_ctx 

118 

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 

128 

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 )