Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 3 additions & 4 deletions CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
20 changes: 10 additions & 10 deletions tests/backends/test_redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,7 +135,7 @@ def _claim_expired(
f"{backend.key_prefix}:task:",
],
)
return [item.decode() for item in claimed]
return list(claimed)


class TestRedisBroker:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand Down
57 changes: 26 additions & 31 deletions threadmill/backends/redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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:
Expand All @@ -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(
Expand Down Expand Up @@ -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:
Expand All @@ -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}}}"
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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(
Expand Down Expand Up @@ -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),
Expand Down
Loading