codekingpro/portable-devtools
114k
1import _thread2import json3import logging4import random5import time6import typing7 8from redis import client9 10from . import exceptions, utils11 12logger = logging.getLogger(__name__)13 14DEFAULT_UNAVAILABLE_TIMEOUT = 115DEFAULT_THREAD_SLEEP_TIME = 0.116 17 18class PubSubWorkerThread(client.PubSubWorkerThread): # type: ignore19 def run(self):20 try:21 super().run()22 except Exception: # pragma: no cover23 _thread.interrupt_main()24 raise25 26 27class RedisLock(utils.LockBase):28 '''29 An extremely reliable Redis lock based on pubsub with a keep-alive thread30 31 As opposed to most Redis locking systems based on key/value pairs,32 this locking method is based on the pubsub system. The big advantage is33 that if the connection gets killed due to network issues, crashing34 processes or otherwise, it will still immediately unlock instead of35 waiting for a lock timeout.36 37 To make sure both sides of the lock know about the connection state it is38 recommended to set the `health_check_interval` when creating the redis39 connection..40 41 Args:42 channel: the redis channel to use as locking key.43 connection: an optional redis connection if you already have one44 or if you need to specify the redis connection45 timeout: timeout when trying to acquire a lock46 check_interval: check interval while waiting47 fail_when_locked: after the initial lock failed, return an error48 or lock the file. This does not wait for the timeout.49 thread_sleep_time: sleep time between fetching messages from redis to50 prevent a busy/wait loop. In the case of lock conflicts this51 increases the time it takes to resolve the conflict. This should52 be smaller than the `check_interval` to be useful.53 unavailable_timeout: If the conflicting lock is properly connected54 this should never exceed twice your redis latency. Note that this55 will increase the wait time possibly beyond your `timeout` and is56 always executed if a conflict arises.57 redis_kwargs: The redis connection arguments if no connection is58 given. The `DEFAULT_REDIS_KWARGS` are used as default, if you want59 to override these you need to explicitly specify a value (e.g.60 `health_check_interval=0`)61 62 '''63 64 redis_kwargs: typing.Dict[str, typing.Any]65 thread: typing.Optional[PubSubWorkerThread]66 channel: str67 timeout: float68 connection: typing.Optional[client.Redis]69 pubsub: typing.Optional[client.PubSub] = None70 close_connection: bool71 72 DEFAULT_REDIS_KWARGS: typing.ClassVar[typing.Dict[str, typing.Any]] = dict(73 health_check_interval=10,74 )75 76 def __init__(77 self,78 channel: str,79 connection: typing.Optional[client.Redis] = None,80 timeout: typing.Optional[float] = None,81 check_interval: typing.Optional[float] = None,82 fail_when_locked: typing.Optional[bool] = False,83 thread_sleep_time: float = DEFAULT_THREAD_SLEEP_TIME,84 unavailable_timeout: float = DEFAULT_UNAVAILABLE_TIMEOUT,85 redis_kwargs: typing.Optional[typing.Dict] = None,86 ):87 # We don't want to close connections given as an argument88 self.close_connection = not connection89 90 self.thread = None91 self.channel = channel92 self.connection = connection93 self.thread_sleep_time = thread_sleep_time94 self.unavailable_timeout = unavailable_timeout95 self.redis_kwargs = redis_kwargs or dict()96 97 for key, value in self.DEFAULT_REDIS_KWARGS.items():98 self.redis_kwargs.setdefault(key, value)99 100 super().__init__(101 timeout=timeout,102 check_interval=check_interval,103 fail_when_locked=fail_when_locked,104 )105 106 def get_connection(self) -> client.Redis:107 if not self.connection:108 self.connection = client.Redis(**self.redis_kwargs)109 110 return self.connection111 112 def channel_handler(self, message):113 if message.get('type') != 'message': # pragma: no cover114 return115 116 try:117 data = json.loads(message.get('data'))118 except TypeError: # pragma: no cover119 logger.debug('TypeError while parsing: %r', message)120 return121 122 assert self.connection is not None123 self.connection.publish(data['response_channel'], str(time.time()))124 125 @property126 def client_name(self):127 return f'{self.channel}-lock'128 129 def acquire(130 self,131 timeout: typing.Optional[float] = None,132 check_interval: typing.Optional[float] = None,133 fail_when_locked: typing.Optional[bool] = None,134 ):135 timeout = utils.coalesce(timeout, self.timeout, 0.0)136 check_interval = utils.coalesce(137 check_interval,138 self.check_interval,139 0.0,140 )141 fail_when_locked = utils.coalesce(142 fail_when_locked,143 self.fail_when_locked,144 )145 146 assert not self.pubsub, 'This lock is already active'147 connection = self.get_connection()148 149 timeout_generator = self._timeout_generator(timeout, check_interval)150 for _ in timeout_generator: # pragma: no branch151 subscribers = connection.pubsub_numsub(self.channel)[0][1]152 153 if subscribers:154 logger.debug(155 'Found %d lock subscribers for %s',156 subscribers,157 self.channel,158 )159 160 if self.check_or_kill_lock(161 connection,162 self.unavailable_timeout,163 ): # pragma: no branch164 continue165 else: # pragma: no cover166 subscribers = 0167 168 # Note: this should not be changed to an elif because the if169 # above can still end up here170 if not subscribers:171 connection.client_setname(self.client_name)172 self.pubsub = connection.pubsub()173 self.pubsub.subscribe(**{self.channel: self.channel_handler})174 self.thread = PubSubWorkerThread(175 self.pubsub,176 sleep_time=self.thread_sleep_time,177 )178 self.thread.start()179 180 subscribers = connection.pubsub_numsub(self.channel)[0][1]181 if subscribers == 1: # pragma: no branch182 return self183 else: # pragma: no cover184 # Race condition, let's try again185 self.release()186 187 if fail_when_locked: # pragma: no cover188 raise exceptions.AlreadyLocked(exceptions)189 190 raise exceptions.AlreadyLocked(exceptions)191 192 def check_or_kill_lock(self, connection, timeout):193 # Random channel name to get messages back from the lock194 response_channel = f'{self.channel}-{random.random()}'195 196 pubsub = connection.pubsub()197 pubsub.subscribe(response_channel)198 connection.publish(199 self.channel,200 json.dumps(201 dict(202 response_channel=response_channel,203 message='ping',204 ),205 ),206 )207 208 check_interval = min(self.thread_sleep_time, timeout / 10)209 for _ in self._timeout_generator(210 timeout,211 check_interval,212 ): # pragma: no branch213 if pubsub.get_message(timeout=check_interval):214 pubsub.close()215 return True216 217 for client_ in connection.client_list('pubsub'): # pragma: no cover218 if client_.get('name') == self.client_name:219 logger.warning('Killing unavailable redis client: %r', client_)220 connection.client_kill_filter(client_.get('id'))221 return None222 223 def release(self):224 if self.thread: # pragma: no branch225 self.thread.stop()226 self.thread.join()227 self.thread = None228 time.sleep(0.01)229 230 if self.pubsub: # pragma: no branch231 self.pubsub.unsubscribe(self.channel)232 self.pubsub.close()233 self.pubsub = None234 235 def __del__(self):236 self.release()237 