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
55from taskiq import AsyncResultBackend
66from taskiq .abc .serializer import TaskiqSerializer
77from taskiq .compat import model_dump , model_validate
1212from 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
0 commit comments