From 6a6048ff4e3cd8ebfb260bebf97c9598201ba8b6 Mon Sep 17 00:00:00 2001 From: Johannes Maron Date: Wed, 7 Oct 2026 14:34:07 +0200 Subject: [PATCH] Hardcode decode_responses=True on Redis clients The backend now loads both the sync and async Redis clients with decode_responses=True, so task IDs, lease fields, and pub/sub payloads arrive as strings. Drop the bytes decoding at every call site and the bytes handling in the tests. Document the inverted decision in CONTRIBUTING.md. --- CONTRIBUTING.md | 7 ++--- tests/backends/test_redis.py | 20 ++++++------- threadmill/backends/redis.py | 57 ++++++++++++++++-------------------- 3 files changed, 39 insertions(+), 45 deletions(-) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 33e86f8..9bdbdb7 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -15,10 +15,9 @@ curl -sSL https://raw.githubusercontent.com/codingjoe/naming-things/refs/heads/m We require a persistent Redis without eviction. -Redis connections use the redis-py default `decode_responses=False`, so all -values read from Redis are bytes. We do not guard against misconfiguration. -We fail loudly instead. The same goes for Redis data altered mid-flight. -These are deliberate design decisions. +Redis connections explicitly load with `decode_responses=True`, so all values +read from Redis are strings. We do not guard against Redis data altered +mid-flight. We fail loudly instead. These are deliberate design decisions. ## Testing diff --git a/tests/backends/test_redis.py b/tests/backends/test_redis.py index a113c84..a5eea12 100644 --- a/tests/backends/test_redis.py +++ b/tests/backends/test_redis.py @@ -135,7 +135,7 @@ def _claim_expired( f"{backend.key_prefix}:task:", ], ) - return [item.decode() for item in claimed] + return list(claimed) class TestRedisBroker: @@ -852,7 +852,7 @@ def _ack(status: TaskResultStatus) -> str: backend.close() @staticmethod - def _await_message(pubsub, expected: bytes, *, timeout: float = 2.0): + def _await_message(pubsub, expected: str, *, timeout: float = 2.0): """Drain pubsub until a user message with the expected payload arrives.""" deadline = time.monotonic() + timeout while time.monotonic() < deadline: @@ -889,8 +889,8 @@ def test_enqueue__publishes_ingress_telemetry(self): pubsub.subscribe(backend.telemetry_channel) self._drain_subscription(pubsub) backend.enqueue(echo, args=[1]) - message = self._await_message(pubsub, b"ingress:default") - assert message["channel"] == backend.telemetry_channel.encode() + message = self._await_message(pubsub, "ingress:default") + assert message["channel"] == backend.telemetry_channel finally: pubsub.unsubscribe(backend.telemetry_channel) pubsub.close() @@ -924,8 +924,8 @@ def test_acknowledge__publishes_egress_telemetry(self): finished_at=timezone.now(), ) ) - message = self._await_message(pubsub, b"egress:default") - assert message["channel"] == backend.telemetry_channel.encode() + message = self._await_message(pubsub, "egress:default") + assert message["channel"] == backend.telemetry_channel finally: pubsub.unsubscribe(backend.telemetry_channel) pubsub.close() @@ -959,15 +959,15 @@ def test_requeue__publishes_ingress_telemetry(self): pubsub.subscribe(backend.telemetry_channel) self._drain_subscription(pubsub) backend.requeue(failed, timezone.now() + datetime.timedelta(seconds=10)) - message = self._await_message(pubsub, b"ingress:default") - assert message["channel"] == backend.telemetry_channel.encode() + message = self._await_message(pubsub, "ingress:default") + assert message["channel"] == backend.telemetry_channel finally: pubsub.unsubscribe(backend.telemetry_channel) pubsub.close() backend.close() - async def test_worker_telemetry__yields_bytes_reply_as_event(self): - """worker_telemetry() decodes the bytes pub/sub reply and yields the event.""" + async def test_worker_telemetry__yields_reply_as_event(self): + """worker_telemetry() yields the pub/sub reply as an event.""" backend = _make_backend("worker_telemetry_test") try: stream = backend.worker_telemetry() diff --git a/threadmill/backends/redis.py b/threadmill/backends/redis.py index 442c4db..3078031 100644 --- a/threadmill/backends/redis.py +++ b/threadmill/backends/redis.py @@ -42,11 +42,11 @@ def _load_lua(name: str) -> str: return (_LUA_DIR / f"{name}.lua").read_text() -def _parse_lease_started_at(value: bytes | None) -> datetime.datetime | None: +def _parse_lease_started_at(value: str | None) -> datetime.datetime | None: """Return the lease start stamped beside a task, or None when the hash holds no timestamp.""" try: - return datetime.datetime.fromisoformat(value.decode()) - except AttributeError, ValueError: + return datetime.datetime.fromisoformat(value) + except TypeError, ValueError: return None @@ -95,8 +95,7 @@ def _reap_running_queue(self, queue_name: str) -> None: f"{self.backend.key_prefix}:task:", ], ) - for member in claimed_ids: - task_id = member.decode() + for task_id in claimed_ids: try: self._reap_task(task_id) except ImportError as read_error: @@ -118,11 +117,9 @@ def _reap_running_queue(self, queue_name: str) -> None: ) task_result = self.backend._apply_lease( self.backend.deserialize_task_result(json.dumps(payload)), - worker=lease_worker.decode() if lease_worker else None, + worker=lease_worker or None, lease_started_at=_parse_lease_started_at(lease_started_at), - lease_token=( - lease_token.decode() if lease_token is not None else None - ), + lease_token=lease_token, ) self.backend.acknowledge( dataclasses.replace( @@ -226,10 +223,10 @@ def get_leased_task(self, task_id: str) -> ThreadmillTaskResult | None: if data is None: return None return self._apply_lease( - self.deserialize_task_result(data.decode()), - worker=lease_worker.decode() if lease_worker else None, + self.deserialize_task_result(data), + worker=lease_worker or None, lease_started_at=_parse_lease_started_at(lease_started_at), - lease_token=lease_token.decode() if lease_token is not None else None, + lease_token=lease_token, ) def __init__(self, alias: str, params: dict) -> None: @@ -241,7 +238,7 @@ def __init__(self, alias: str, params: dict) -> None: raise ValueError( f"REDIS_URL must be specified in your settings for the {type(self).__name__}." ) from e - self.client = redis.from_url(redis_url) + self.client = redis.from_url(redis_url, decode_responses=True) self._async_client: redis.asyncio.Redis | None = None self.redis_url = redis_url self.key_prefix = f"threadmill:{{{alias}}}" @@ -264,7 +261,9 @@ def __init__(self, alias: str, params: dict) -> None: def async_client(self) -> redis.asyncio.Redis: """Lazily-created async Redis client, reused across calls.""" if self._async_client is None: - self._async_client = redis.asyncio.Redis.from_url(self.redis_url) + self._async_client = redis.asyncio.Redis.from_url( + self.redis_url, decode_responses=True + ) return self._async_client def _compute_score(self, priority: int, enqueued_at: datetime.datetime) -> float: @@ -376,10 +375,10 @@ def acquire( self._miss_count = 0 self._rotation_offset = (self._rotation_offset + 1) % len(queue_names) return self._apply_lease( - self.deserialize_task_result(data.decode()), + self.deserialize_task_result(data), worker=worker, lease_started_at=now, - lease_token=lease_token.decode(), + lease_token=lease_token, ) try: @@ -554,10 +553,8 @@ def _peek_tasks( Leased tasks are yielded as their lease reports them. """ pipe = self.client.pipeline() - for member in self.client.zrange(zset_key, 0, count - 1): - task_key = self.TASK_KEY.format( - prefix=self.key_prefix, task_id=member.decode() - ) + for task_id in self.client.zrange(zset_key, 0, count - 1): + task_key = self.TASK_KEY.format(prefix=self.key_prefix, task_id=task_id) if leased: pipe.hmget(task_key, "data", *self.LEASE_FIELDS) else: @@ -570,33 +567,31 @@ def _peek_tasks( if not data: continue yield self._apply_lease( - self.deserialize_task_result(data.decode()), - worker=lease_worker.decode() if lease_worker else None, + self.deserialize_task_result(data), + worker=lease_worker or None, lease_started_at=_parse_lease_started_at(lease_started_at), - lease_token=( - lease_token.decode() if lease_token is not None else None - ), + lease_token=lease_token, ) else: - yield self.deserialize_task_result(stored.decode()) + yield self.deserialize_task_result(stored) def _peek_results(self, zset_key: str, count: int) -> Generator[TaskResult]: """Yield up to `count` finished results in finish order.""" pipe = self.client.pipeline() - for member in self.client.zrange(zset_key, 0, count - 1): + for result_id in self.client.zrange(zset_key, 0, count - 1): result_key = self.RESULT_KEY.format( - prefix=self.key_prefix, result_id=member.decode() + prefix=self.key_prefix, result_id=result_id ) pipe.get(result_key) for stored in pipe.execute(): if stored: - yield self.deserialize_task_result(stored.decode()) + yield self.deserialize_task_result(stored) def get_result(self, result_id: str) -> TaskResult: if data := self.client.get( self.RESULT_KEY.format(prefix=self.key_prefix, result_id=result_id) ): - return self.deserialize_task_result(data.decode()) + return self.deserialize_task_result(data) raise TaskResultDoesNotExist(f"Task result {result_id!r} does not exist.") async def queue_stats( @@ -639,7 +634,7 @@ async def worker_telemetry( try: async for message in pubsub.listen(): if (data := message.get("data")) is not None: - direction, _, queue_name = data.decode().partition(":") + direction, _, queue_name = data.partition(":") try: event = TelemetryEvent( direction=TelemetryDirection(direction),