Skip to content

Commit 1e78039

Browse files
authored
feat: add new result_backend with capo-s3 (#40)
1 parent af34d48 commit 1e78039

4 files changed

Lines changed: 346 additions & 97 deletions

File tree

‎pyproject.toml‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ dependencies = [
3232
"taskiq>=0.12.6",
3333
"aiobotocore>=2.13.3",
3434
"capo-s3>=0.15.0",
35+
"capo-sqs>=0.6.0",
3536
]
3637

3738
[project.urls]
@@ -43,7 +44,8 @@ dependencies = [
4344
dev = [
4445
{include-group = "lint"},
4546
{include-group = "test"},
46-
{include-group = "types"},
47+
{include-group = "examples"},
48+
{include-group = "docs"},
4749
"prek>=0.5.3",
4850
]
4951
test = [
@@ -54,14 +56,15 @@ test = [
5456
lint = [
5557
"ruff>=0.16.7",
5658
"zizmor>=1.30.1",
57-
]
58-
types = [
5959
"mypy>=2.3.1",
6060
"types-aiobotocore[essential]>=3.7.0",
6161
]
6262
examples = [
6363
"python-dotenv>=1.2.3",
6464
]
65+
docs = [
66+
"zensical>=0.0.62",
67+
]
6568

6669

6770
[build-system]

‎src/taskiq_sqs/result_backend.py‎

Lines changed: 49 additions & 76 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
1-
from typing import TYPE_CHECKING, Any, TypeVar
1+
import contextlib
2+
from typing import Any, TypeVar
23

3-
from aiobotocore.session import get_session
4-
from botocore.exceptions import ClientError
4+
import capo_s3
55
from taskiq import AsyncResultBackend
66
from taskiq.abc.serializer import TaskiqSerializer
77
from taskiq.compat import model_dump, model_validate
@@ -12,9 +12,6 @@
1212
from taskiq_sqs.types import S3Bucket
1313

1414

15-
if TYPE_CHECKING:
16-
from types_aiobotocore_s3.client import S3Client
17-
1815
_ReturnType = TypeVar("_ReturnType")
1916

2017

@@ -48,61 +45,57 @@ def __init__(
4845
self._aws_secret_access_key = aws_secret_access_key
4946
self._bucket = bucket
5047
self._base_path = base_path
51-
self._session = get_session()
5248
self._serializer = serializer or JSONSerializer()
5349

54-
async def _get_client(self) -> "S3Client":
55-
"""
56-
Retrieves the S3 client, creating it if necessary.
57-
58-
Returns:
59-
S3Client: The initialized S3 client.
60-
"""
61-
self._client_context_creator = self._session.create_client(
62-
"s3",
63-
region_name=self._aws_region,
64-
endpoint_url=self._aws_endpoint_url,
65-
aws_access_key_id=self._aws_access_key_id,
66-
aws_secret_access_key=self._aws_secret_access_key,
67-
)
68-
return await self._client_context_creator.__aenter__()
69-
7050
async def startup(self) -> None:
71-
"""Initialize the result backend."""
72-
self._s3_client = await self._get_client()
51+
"""Initialize the S3 client and ensure the bucket exists."""
52+
credentials = None
53+
if self._aws_access_key_id and self._aws_secret_access_key:
54+
credentials = capo_s3.Credentials(
55+
access_key=self._aws_access_key_id,
56+
secret_key=self._aws_secret_access_key,
57+
)
58+
self._s3_client = capo_s3.AsyncS3Client(
59+
region=self._aws_region,
60+
endpoint=self._aws_endpoint_url,
61+
credentials=credentials,
62+
force_path_style=True,
63+
)
64+
await self._s3_client.__aenter__()
7365
try:
7466
await self._ensure_bucket_exists()
7567
except Exception:
76-
await self._client_context_creator.__aexit__(None, None, None)
68+
await self._s3_client.__aexit__(None, None, None)
7769
raise
7870
return await super().startup()
7971

8072
async def _ensure_bucket_exists(self) -> None:
8173
try:
82-
await self._s3_client.head_bucket(Bucket=self._bucket["name"])
83-
except ClientError as exc:
84-
code = exc.response.get("Error", {}).get("Code")
85-
if code not in ("404", "NoSuchBucket"):
86-
raise exceptions.ResultBackendError(code=code) from exc
74+
await self._s3_client.head_bucket(bucket=self._bucket["name"])
75+
except capo_s3.errors.NotFound:
8776
if not self._bucket.get("declare", True):
88-
raise exceptions.BucketNotFoundError(bucket_name=self._bucket["name"]) from exc
77+
raise exceptions.BucketNotFoundError(bucket_name=self._bucket["name"]) from None
8978
await self._create_bucket()
79+
except capo_s3.errors.ServiceError as exc:
80+
raise exceptions.ResultBackendError(code=exc.code) from exc
9081

9182
async def _create_bucket(self) -> None:
92-
create_kwargs: dict[str, Any] = {"Bucket": self._bucket["name"]}
83+
create_kwargs: dict[str, Any] = {}
9384
if self._aws_region and self._aws_region != constants.AWS_DEFAULT_REGION:
94-
create_kwargs["CreateBucketConfiguration"] = {"LocationConstraint": self._aws_region}
95-
try:
96-
await self._s3_client.create_bucket(**create_kwargs)
97-
except ClientError as exc:
98-
if exc.response.get("Error", {}).get("Code") != "BucketAlreadyOwnedByYou": # can be raise between workers
99-
raise
85+
create_kwargs["create_bucket_configuration"] = {"location_constraint": self._aws_region}
86+
with contextlib.suppress(capo_s3.errors.BucketAlreadyOwnedByYou):
87+
await self._s3_client.create_bucket(bucket=self._bucket["name"], **create_kwargs)
10088

10189
async def shutdown(self) -> None:
10290
"""Shut down the result backend."""
103-
await self._client_context_creator.__aexit__(None, None, None)
91+
await self._s3_client.__aexit__(None, None, None)
10492
return await super().shutdown()
10593

94+
def _build_key(self, task_id: str) -> str:
95+
if self._base_path:
96+
return f"{self._base_path.rstrip('/')}/{task_id}"
97+
return task_id
98+
10699
async def set_result(
107100
self,
108101
task_id: str,
@@ -114,13 +107,10 @@ async def set_result(
114107
:param task_id: current task id.
115108
:param result: result of execution.
116109
"""
117-
if self._base_path:
118-
task_id = f"{self._base_path.rstrip('/')}/{task_id}"
119-
120110
await self._s3_client.put_object(
121-
Bucket=self._bucket["name"],
122-
Key=task_id,
123-
Body=self._serializer.dumpb(model_dump(result)),
111+
bucket=self._bucket["name"],
112+
key=self._build_key(task_id),
113+
body=self._serializer.dumpb(model_dump(result)),
124114
)
125115

126116
async def get_result(
@@ -138,27 +128,17 @@ async def get_result(
138128
:param with_logs: whether to fetch logs.
139129
:return: result.
140130
"""
141-
result = None
142-
if self._base_path:
143-
task_id = f"{self._base_path.rstrip('/')}/{task_id}"
144131
try:
145-
if response := await self._s3_client.get_object(
146-
Bucket=self._bucket["name"],
147-
Key=task_id,
148-
):
149-
async with response["Body"] as stream:
150-
result = await stream.read()
151-
except ClientError as exc:
152-
code = exc.response.get("Error", {}).get("Code")
153-
if code in ["NoSuchKey", "404"]:
154-
raise exceptions.ResultIsMissingError(task_id=task_id) from exc
155-
raise exceptions.ResultBackendError(code=code) from exc
156-
if result is None:
157-
raise exceptions.ResultIsMissingError(task_id=task_id)
132+
async with self._s3_client.get_object(bucket=self._bucket["name"], key=self._build_key(task_id)) as output:
133+
body = b"".join([chunk async for chunk in output["body"]])
134+
except capo_s3.errors.NoSuchKey as exc:
135+
raise exceptions.ResultIsMissingError(task_id=task_id) from exc
136+
except capo_s3.errors.ServiceError as exc:
137+
raise exceptions.ResultBackendError(code=exc.code) from exc
158138

159139
taskiq_result = model_validate(
160140
TaskiqResult[_ReturnType],
161-
self._serializer.loadb(result),
141+
self._serializer.loadb(body),
162142
)
163143

164144
if not with_logs:
@@ -173,17 +153,10 @@ async def is_result_ready(self, task_id: str) -> bool:
173153
:param task_id: id of a task.
174154
:return: True if result is ready.
175155
"""
176-
if self._base_path:
177-
task_id = f"{self._base_path.rstrip('/')}/{task_id}"
178156
try:
179-
if await self._s3_client.head_object(Bucket=self._bucket["name"], Key=task_id):
180-
return True
181-
except ClientError as exc:
182-
code = exc.response.get("Error", {}).get("Code")
183-
if code in ["NoSuchKey", "404"]:
184-
pass
185-
else:
186-
raise exceptions.ResultBackendError(
187-
code=code,
188-
) from exc
189-
return False
157+
await self._s3_client.head_object(bucket=self._bucket["name"], key=self._build_key(task_id))
158+
except capo_s3.errors.NotFound:
159+
return False
160+
except capo_s3.errors.ServiceError as exc:
161+
raise exceptions.ResultBackendError(code=exc.code) from exc
162+
return True

‎tests/test_result_backend.py‎

Lines changed: 4 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
import uuid
22
from collections.abc import AsyncGenerator
33
from typing import TYPE_CHECKING, Any
4-
from unittest.mock import AsyncMock
54

65
import pytest
76
from taskiq.result import TaskiqResult
@@ -26,12 +25,13 @@ class TestResultBackend:
2625
async def test_when_result_set__then_result_is_actually_saved_to_s3(
2726
self,
2827
s3_backend: S3ResultBackend,
28+
s3_client: "S3Client",
2929
s3_bucket: str,
3030
taskiq_result: TaskiqResult,
3131
) -> None:
3232
await s3_backend.set_result("test_task_id", taskiq_result)
3333

34-
response = await s3_backend._s3_client.get_object(
34+
response = await s3_client.get_object(
3535
Bucket=s3_bucket,
3636
Key="test_task_id",
3737
)
@@ -48,24 +48,17 @@ async def test_when_result_present_in_s3__then_get_result_return_it(
4848
assert retrieved_result.return_value == "test_value"
4949
assert retrieved_result.is_err is True
5050

51-
async def test_when_result_is_missing__then_get_result_raise_exception(
52-
self,
53-
s3_backend: S3ResultBackend,
54-
) -> None:
55-
s3_backend._s3_client.get_object = AsyncMock(return_value={}) # Simulate a response with no Body
56-
with pytest.raises(ResultIsMissingError):
57-
await s3_backend.get_result("test_task_id")
58-
5951
async def test_when_set_result_is_called__then_save_it_to_right_path(
6052
self,
6153
s3_backend: S3ResultBackend,
54+
s3_client: "S3Client",
6255
s3_bucket: str,
6356
taskiq_result: TaskiqResult,
6457
) -> None:
6558
s3_backend._base_path = "results"
6659
await s3_backend.set_result("test_task_id", taskiq_result)
6760

68-
response = await s3_backend._s3_client.head_object(
61+
response = await s3_client.head_object(
6962
Bucket=s3_bucket,
7063
Key="results/test_task_id",
7164
)

0 commit comments

Comments
 (0)