Skip to content

Commit d8f2d42

Browse files
authored
feat: batch check updates (#41)
2 parents d27eb1b + 095c3dd commit d8f2d42

2 files changed

Lines changed: 19 additions & 14 deletions

File tree

openfga_sdk/client/client.py

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -511,17 +511,20 @@ async def check(self, body: ClientCheckRequest, options: dict[str, str] = None):
511511
)
512512
return api_response
513513

514-
async def _single_batch_check(self, body: ClientCheckRequest, options: dict[str, str] = None): # noqa: E501
514+
async def _single_batch_check(self, body: ClientCheckRequest, semaphore: asyncio.Semaphore, options: dict[str, str] = None): # noqa: E501
515515
"""
516516
Run a single batch request and return body in a SingleBatchCheckResponse
517517
:param body - ClientCheckRequest defining check request
518518
:param authorization_model_id(options) - Overrides the authorization model id in the configuration
519519
"""
520+
await semaphore.acquire()
520521
try:
521522
api_response = await self.check(body, options)
522523
return BatchCheckResponse(allowed=api_response.allowed, request=body, response=api_response, error=None)
523524
except Exception as err:
524525
return BatchCheckResponse(allowed=False, request=body, response=None, error=err)
526+
finally:
527+
semaphore.release()
525528

526529
async def batch_check(self, body: List[ClientCheckRequest], options: dict[str, str] = None): # noqa: E501
527530
"""
@@ -543,13 +546,11 @@ async def batch_check(self, body: List[ClientCheckRequest], options: dict[str, s
543546
max_parallel_requests = 10
544547
if options is not None and "max_parallel_requests" in options:
545548
max_parallel_requests = options["max_parallel_requests"]
546-
# Break the batch into chunks
547-
request_batches = _chuck_array(body, max_parallel_requests)
548-
batch_check_response = []
549-
for request_batch in request_batches:
550-
request = [self._single_batch_check(i, options) for i in request_batch]
551-
response = await asyncio.gather(*request)
552-
batch_check_response.extend(response)
549+
550+
sem = asyncio.Semaphore(max_parallel_requests)
551+
batch_check_coros = [self._single_batch_check(request, sem, options) for request in body]
552+
batch_check_response = await asyncio.gather(*batch_check_coros)
553+
553554
return batch_check_response
554555

555556
async def expand(self, body: ClientExpandRequest, options: dict[str, str] = None): # noqa: E501

openfga_sdk/sync/client/client.py

Lines changed: 10 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -40,9 +40,9 @@
4040
from openfga_sdk.models.write_request import WriteRequest
4141
from openfga_sdk.validation import is_well_formed_ulid_string
4242

43-
import time
4443
import uuid
4544
from typing import List
45+
from concurrent.futures import ThreadPoolExecutor
4646

4747
CLIENT_METHOD_HEADER = "X-OpenFGA-Client-Method"
4848
CLIENT_BULK_REQUEST_ID_HEADER = "X-OpenFGA-Client-Bulk-Request-Id"
@@ -543,12 +543,16 @@ def batch_check(self, body: List[ClientCheckRequest], options: dict[str, str] =
543543
max_parallel_requests = 10
544544
if options is not None and "max_parallel_requests" in options:
545545
max_parallel_requests = options["max_parallel_requests"]
546-
# Break the batch into chunks
547-
request_batches = _chuck_array(body, max_parallel_requests)
546+
548547
batch_check_response = []
549-
for request_batch in request_batches:
550-
response = [self._single_batch_check(i, options) for i in request_batch]
551-
batch_check_response.extend(response)
548+
549+
def single_batch_check(request):
550+
return self._single_batch_check(request, options)
551+
552+
with ThreadPoolExecutor(max_workers=max_parallel_requests) as executor:
553+
for response in executor.map(single_batch_check, body):
554+
batch_check_response.append(response)
555+
552556
return batch_check_response
553557

554558
def expand(self, body: ClientExpandRequest, options: dict[str, str] = None): # noqa: E501

0 commit comments

Comments
 (0)