Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/redis/asyncio/lock.py: 24%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

131 statements  

1import asyncio 

2import logging 

3import threading 

4import uuid 

5from types import SimpleNamespace 

6from typing import TYPE_CHECKING, Awaitable, Literal, Optional, Union 

7 

8from redis.exceptions import LockError, LockNotOwnedError 

9from redis.typing import Number 

10 

11if TYPE_CHECKING: 

12 from redis.asyncio import Redis, RedisCluster 

13 

14logger = logging.getLogger(__name__) 

15 

16 

17class Lock: 

18 """ 

19 A shared, distributed Lock. Using Redis for locking allows the Lock 

20 to be shared across processes and/or machines. 

21 

22 It's left to the user to resolve deadlock issues and make sure 

23 multiple clients play nicely together. 

24 """ 

25 

26 lua_release = None 

27 lua_extend = None 

28 lua_reacquire = None 

29 

30 # KEYS[1] - lock name 

31 # ARGV[1] - token 

32 # return 1 if the lock was released, otherwise 0 

33 LUA_RELEASE_SCRIPT = """ 

34 local token = redis.call('get', KEYS[1]) 

35 if not token or token ~= ARGV[1] then 

36 return 0 

37 end 

38 redis.call('del', KEYS[1]) 

39 return 1 

40 """ 

41 

42 # KEYS[1] - lock name 

43 # ARGV[1] - token 

44 # ARGV[2] - additional milliseconds 

45 # ARGV[3] - "0" if the additional time should be added to the lock's 

46 # existing ttl or "1" if the existing ttl should be replaced 

47 # return 1 if the locks time was extended, otherwise 0 

48 LUA_EXTEND_SCRIPT = """ 

49 local token = redis.call('get', KEYS[1]) 

50 if not token or token ~= ARGV[1] then 

51 return 0 

52 end 

53 local expiration = redis.call('pttl', KEYS[1]) 

54 if not expiration then 

55 expiration = 0 

56 end 

57 if expiration < 0 then 

58 return 0 

59 end 

60 

61 local newttl = ARGV[2] 

62 if ARGV[3] == "0" then 

63 newttl = ARGV[2] + expiration 

64 end 

65 if tonumber(newttl) <= 0 then 

66 return 1 

67 end 

68 redis.call('pexpire', KEYS[1], newttl) 

69 return 1 

70 """ 

71 

72 # KEYS[1] - lock name 

73 # ARGV[1] - token 

74 # ARGV[2] - milliseconds 

75 # return 1 if the locks time was reacquired, otherwise 0 

76 LUA_REACQUIRE_SCRIPT = """ 

77 local token = redis.call('get', KEYS[1]) 

78 if not token or token ~= ARGV[1] then 

79 return 0 

80 end 

81 if tonumber(ARGV[2]) > 0 then 

82 redis.call('pexpire', KEYS[1], ARGV[2]) 

83 end 

84 return 1 

85 """ 

86 

87 def __init__( 

88 self, 

89 redis: Union["Redis", "RedisCluster"], 

90 name: Union[str, bytes, memoryview], 

91 timeout: Optional[float] = None, 

92 sleep: float = 0.1, 

93 blocking: bool = True, 

94 blocking_timeout: Optional[Number] = None, 

95 thread_local: bool = True, 

96 raise_on_release_error: bool = True, 

97 ): 

98 """ 

99 Create a new Lock instance named ``name`` using the Redis client 

100 supplied by ``redis``. 

101 

102 ``timeout`` indicates a maximum life for the lock in seconds. 

103 By default, it will remain locked until release() is called. 

104 ``timeout`` can be specified as a float or integer, both representing 

105 the number of seconds to wait. 

106 

107 ``sleep`` indicates the amount of time to sleep in seconds per loop 

108 iteration when the lock is in blocking mode and another client is 

109 currently holding the lock. 

110 

111 ``blocking`` indicates whether calling ``acquire`` should block until 

112 the lock has been acquired or to fail immediately, causing ``acquire`` 

113 to return False and the lock not being acquired. Defaults to True. 

114 Note this value can be overridden by passing a ``blocking`` 

115 argument to ``acquire``. 

116 

117 ``blocking_timeout`` indicates the maximum amount of time in seconds to 

118 spend trying to acquire the lock. A value of ``None`` indicates 

119 continue trying forever. ``blocking_timeout`` can be specified as a 

120 float or integer, both representing the number of seconds to wait. 

121 

122 ``thread_local`` indicates whether the lock token is placed in 

123 thread-local storage. By default, the token is placed in thread local 

124 storage so that a thread only sees its token, not a token set by 

125 another thread. Consider the following timeline: 

126 

127 time: 0, thread-1 acquires `my-lock`, with a timeout of 5 seconds. 

128 thread-1 sets the token to "abc" 

129 time: 1, thread-2 blocks trying to acquire `my-lock` using the 

130 Lock instance. 

131 time: 5, thread-1 has not yet completed. redis expires the lock 

132 key. 

133 time: 5, thread-2 acquired `my-lock` now that it's available. 

134 thread-2 sets the token to "xyz" 

135 time: 6, thread-1 finishes its work and calls release(). if the 

136 token is *not* stored in thread local storage, then 

137 thread-1 would see the token value as "xyz" and would be 

138 able to successfully release the thread-2's lock. 

139 

140 ``raise_on_release_error`` indicates whether to raise an exception when 

141 the lock is no longer owned when exiting the context manager. By default, 

142 this is True, meaning an exception will be raised. If False, the warning 

143 will be logged and the exception will be suppressed. 

144 

145 In some use cases it's necessary to disable thread local storage. For 

146 example, if you have code where one thread acquires a lock and passes 

147 that lock instance to a worker thread to release later. If thread 

148 local storage isn't disabled in this case, the worker thread won't see 

149 the token set by the thread that acquired the lock. Our assumption 

150 is that these cases aren't common and as such default to using 

151 thread local storage. 

152 """ 

153 self.redis = redis 

154 self.name = name 

155 self.timeout = timeout 

156 self.sleep = sleep 

157 self.blocking = blocking 

158 self.blocking_timeout = blocking_timeout 

159 self.thread_local = bool(thread_local) 

160 self.local = threading.local() if self.thread_local else SimpleNamespace() 

161 self.raise_on_release_error = raise_on_release_error 

162 self.local.token = None 

163 self.register_scripts() 

164 

165 def register_scripts(self): 

166 cls = self.__class__ 

167 client = self.redis 

168 if cls.lua_release is None: 

169 cls.lua_release = client.register_script(cls.LUA_RELEASE_SCRIPT) 

170 if cls.lua_extend is None: 

171 cls.lua_extend = client.register_script(cls.LUA_EXTEND_SCRIPT) 

172 if cls.lua_reacquire is None: 

173 cls.lua_reacquire = client.register_script(cls.LUA_REACQUIRE_SCRIPT) 

174 

175 async def __aenter__(self): 

176 if await self.acquire(): 

177 return self 

178 raise LockError( 

179 "Unable to acquire lock within the time specified", 

180 lock_name=self.name, 

181 ) 

182 

183 async def __aexit__(self, exc_type, exc_value, traceback): 

184 try: 

185 await self.release() 

186 except LockError: 

187 if self.raise_on_release_error: 

188 raise 

189 logger.warning( 

190 "Lock was unlocked or no longer owned when exiting context manager." 

191 ) 

192 

193 async def acquire( 

194 self, 

195 blocking: Optional[bool] = None, 

196 blocking_timeout: Optional[Number] = None, 

197 token: Optional[Union[str, bytes]] = None, 

198 ): 

199 """ 

200 Use Redis to hold a shared, distributed lock named ``name``. 

201 Returns True once the lock is acquired. 

202 

203 If ``blocking`` is False, always return immediately. If the lock 

204 was acquired, return True, otherwise return False. 

205 

206 ``blocking_timeout`` specifies the maximum number of seconds to 

207 wait trying to acquire the lock. 

208 

209 ``token`` specifies the token value to be used. If provided, token 

210 must be a bytes object or a string that can be encoded to a bytes 

211 object with the default encoding. If a token isn't specified, a UUID 

212 will be generated. 

213 """ 

214 sleep = self.sleep 

215 if token is None: 

216 token = uuid.uuid1().hex.encode() 

217 else: 

218 try: 

219 encoder = self.redis.connection_pool.get_encoder() 

220 except AttributeError: 

221 # Cluster 

222 encoder = self.redis.get_encoder() 

223 token = encoder.encode(token) 

224 if blocking is None: 

225 blocking = self.blocking 

226 if blocking_timeout is None: 

227 blocking_timeout = self.blocking_timeout 

228 stop_trying_at = None 

229 if blocking_timeout is not None: 

230 stop_trying_at = asyncio.get_running_loop().time() + blocking_timeout 

231 while True: 

232 if await self.do_acquire(token): 

233 self.local.token = token 

234 return True 

235 if not blocking: 

236 return False 

237 next_try_at = asyncio.get_running_loop().time() + sleep 

238 if stop_trying_at is not None and next_try_at > stop_trying_at: 

239 return False 

240 await asyncio.sleep(sleep) 

241 

242 async def do_acquire(self, token: Union[str, bytes]) -> bool: 

243 if self.timeout: 

244 # convert to milliseconds 

245 timeout = int(self.timeout * 1000) 

246 else: 

247 timeout = None 

248 if await self.redis.set(self.name, token, nx=True, px=timeout): 

249 return True 

250 return False 

251 

252 async def locked(self) -> bool: 

253 """ 

254 Returns True if this key is locked by any process, otherwise False. 

255 """ 

256 return await self.redis.get(self.name) is not None 

257 

258 async def owned(self) -> bool: 

259 """ 

260 Returns True if this key is locked by this lock, otherwise False. 

261 """ 

262 stored_token = await self.redis.get(self.name) 

263 # need to always compare bytes to bytes 

264 # TODO: this can be simplified when the context manager is finished 

265 if stored_token and not isinstance(stored_token, bytes): 

266 try: 

267 encoder = self.redis.connection_pool.get_encoder() 

268 except AttributeError: 

269 # Cluster 

270 encoder = self.redis.get_encoder() 

271 stored_token = encoder.encode(stored_token) 

272 return self.local.token is not None and stored_token == self.local.token 

273 

274 async def release(self) -> None: 

275 """Releases the already acquired lock. 

276 

277 The token is only cleared after the Redis release operation completes 

278 successfully. This ensures that if the release is cancelled mid-operation, 

279 the lock state remains consistent and can be retried. 

280 """ 

281 expected_token = self.local.token 

282 if expected_token is None: 

283 raise LockError( 

284 "Cannot release a lock that's not owned or is already unlocked.", 

285 lock_name=self.name, 

286 ) 

287 try: 

288 await self.do_release(expected_token) 

289 except LockNotOwnedError: 

290 # Lock doesn't exist in Redis, so clear the token, but only if it is still 

291 # ours. threading.local() is shared by every task on one event loop, so a 

292 # Lock shared between tasks may already hold another acquirer's token; 

293 # overwriting it would leave the new owner unable to release. 

294 if self.local.token == expected_token: 

295 self.local.token = None 

296 raise 

297 # Only clear token after successful release, and only if it is still ours. 

298 if self.local.token == expected_token: 

299 self.local.token = None 

300 

301 async def do_release(self, expected_token: bytes) -> None: 

302 if not bool( 

303 await self.lua_release( 

304 keys=[self.name], args=[expected_token], client=self.redis 

305 ) 

306 ): 

307 raise LockNotOwnedError( 

308 "Cannot release a lock that's no longer owned", 

309 lock_name=self.name, 

310 ) 

311 

312 def extend( 

313 self, additional_time: Number, replace_ttl: bool = False 

314 ) -> Awaitable[Literal[True]]: 

315 """ 

316 Adds more time to an already acquired lock. 

317 

318 ``additional_time`` can be specified as an integer or a float, both 

319 representing the number of seconds to add. 

320 

321 ``replace_ttl`` if False (the default), add `additional_time` to 

322 the lock's existing ttl. If True, replace the lock's ttl with 

323 `additional_time`. 

324 

325 When the resulting TTL is non-positive for a lock with an expiry, the 

326 current expiry is left unchanged and ``True`` is returned. 

327 """ 

328 if self.local.token is None: 

329 raise LockError("Cannot extend an unlocked lock", lock_name=self.name) 

330 if self.timeout is None: 

331 raise LockError("Cannot extend a lock with no timeout", lock_name=self.name) 

332 return self.do_extend(additional_time, replace_ttl) 

333 

334 async def do_extend(self, additional_time, replace_ttl) -> Literal[True]: 

335 additional_time = int(additional_time * 1000) 

336 if not bool( 

337 await self.lua_extend( 

338 keys=[self.name], 

339 args=[self.local.token, additional_time, replace_ttl and "1" or "0"], 

340 client=self.redis, 

341 ) 

342 ): 

343 raise LockNotOwnedError( 

344 "Cannot extend a lock that's no longer owned", 

345 lock_name=self.name, 

346 ) 

347 return True 

348 

349 def reacquire(self) -> Awaitable[Literal[True]]: 

350 """ 

351 Resets a TTL of an already acquired lock back to a timeout value. 

352 

353 When the resulting TTL is non-positive, the current expiry is left 

354 unchanged and ``True`` is returned. 

355 """ 

356 if self.local.token is None: 

357 raise LockError("Cannot reacquire an unlocked lock", lock_name=self.name) 

358 if self.timeout is None: 

359 raise LockError( 

360 "Cannot reacquire a lock with no timeout", lock_name=self.name 

361 ) 

362 return self.do_reacquire() 

363 

364 async def do_reacquire(self) -> Literal[True]: 

365 timeout = int(self.timeout * 1000) 

366 if not bool( 

367 await self.lua_reacquire( 

368 keys=[self.name], args=[self.local.token, timeout], client=self.redis 

369 ) 

370 ): 

371 raise LockNotOwnedError( 

372 "Cannot reacquire a lock that's no longer owned", 

373 lock_name=self.name, 

374 ) 

375 return True