Last active
April 25, 2026 12:10
-
-
Save synodriver/dcf1fac7e553b7998523c7073f83b5a7 to your computer and use it in GitHub Desktop.
like singleflight in go
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| # -*- coding: utf-8 -*- | |
| """ | |
| Copyright (c) 2008-2026 synodriver <diguohuangjiajinweijun@gmail.com> | |
| """ | |
| import asyncio | |
| from enum import IntEnum | |
| from typing import Literal | |
| class NotAvailable(Exception): | |
| pass | |
| class LockState(IntEnum): | |
| empty = 0 # 空 | |
| reading = 1 # 只有读锁 | |
| writing = 2 # 有写锁 | |
| waiting_write = ( | |
| 3 # 有读锁,没有写锁,但是有写锁在等待队列,因此此时不能继续获取读锁 | |
| ) | |
| class RWLock: | |
| def __init__( | |
| self, blocking: bool = True, *, loop: asyncio.AbstractEventLoop | None = None | |
| ): | |
| """ | |
| 读写锁 | |
| :param blocking: 阻塞获取 | |
| :param loop: 事件循环 | |
| """ | |
| self.blocking = blocking | |
| self._read_waiters = [] # type: list[asyncio.Future] | |
| self._write_waiters = [] # type: list[asyncio.Future] | |
| self._pending_reads = 0 # type: int | |
| self._pending_writes = 0 # type: Literal[0, 1] | |
| self._loop = loop or asyncio.get_running_loop() | |
| @staticmethod | |
| def _discard_waiter(waiters: list[asyncio.Future], waiter: asyncio.Future) -> None: | |
| """ | |
| 从等待队列中移除 waiter | |
| :param waiters: 等待队列 | |
| :param waiter: 要移除的 future | |
| :return: | |
| """ | |
| try: | |
| waiters.remove(waiter) | |
| except ValueError: | |
| pass | |
| def _wake_next_writer(self) -> bool: | |
| """ | |
| 唤醒下一个可用的写锁等待者 | |
| :return: 是否成功唤醒写锁等待者 | |
| """ | |
| while self._write_waiters: | |
| waiter = self._write_waiters.pop(0) | |
| if waiter.cancelled() or waiter.done(): | |
| continue | |
| waiter.set_result(None) | |
| self._pending_writes = 1 | |
| return True | |
| return False | |
| def _wake_all_readers(self) -> None: | |
| """ | |
| 唤醒所有可用的读锁等待者 | |
| :return: | |
| """ | |
| waiters = self._read_waiters[:] | |
| self._read_waiters.clear() | |
| for waiter in waiters: | |
| if waiter.cancelled() or waiter.done(): | |
| continue | |
| waiter.set_result(None) | |
| @property | |
| def state(self) -> LockState: | |
| """ | |
| 当前锁状态 | |
| :return: 锁当前的状态 | |
| """ | |
| if self._pending_writes: | |
| return LockState.writing | |
| if not self._pending_reads: | |
| return LockState.empty | |
| if self._write_waiters: | |
| return LockState.waiting_write | |
| return LockState.reading | |
| async def acquire(self, mode: Literal["r", "w"] = "r"): | |
| """ | |
| 获取锁 | |
| :param mode: r or w, r是读锁, w是写锁 | |
| :return: | |
| """ | |
| if mode == "r": | |
| if self.state == LockState.empty or self.state == LockState.reading: | |
| self._pending_reads += 1 | |
| return # 空的 或者 只有读锁 可以立刻获取 | |
| elif self.blocking: | |
| waiter = self._loop.create_future() | |
| self._read_waiters.append(waiter) | |
| try: | |
| await waiter # 被cancel也不要紧,finally总会删除的 | |
| finally: | |
| self._discard_waiter( | |
| self._read_waiters, waiter | |
| ) # 写锁释放的时候才会唤醒获取读锁的协程 | |
| return await self.acquire(mode) # 没出问题才会到这里 | |
| else: | |
| raise NotAvailable | |
| elif mode == "w": | |
| if self.state == LockState.empty: | |
| self._pending_writes = 1 | |
| return | |
| elif self.blocking: | |
| waiter = self._loop.create_future() | |
| self._write_waiters.append(waiter) | |
| try: | |
| await waiter | |
| finally: | |
| self._discard_waiter(self._write_waiters, waiter) | |
| assert self._pending_writes == 1 | |
| else: | |
| raise NotAvailable | |
| else: | |
| raise ValueError("mode must be 'r' or 'w'") | |
| async def release(self, mode: Literal["r", "w"] = "r"): | |
| """ | |
| 释放锁 | |
| :param mode: r or w, r是读锁, w是写锁 | |
| :return: | |
| """ | |
| if mode == "r": | |
| if self.state in (LockState.reading, LockState.waiting_write): | |
| self._pending_reads -= 1 | |
| if self._pending_reads == 0: | |
| self._wake_next_writer() # 轮一下写锁等待队列 | |
| else: | |
| raise ValueError("can not release more than acquire") | |
| elif mode == "w": | |
| if self.state == LockState.writing: | |
| self._pending_writes = 0 | |
| if not self._wake_next_writer(): | |
| assert self.state == LockState.empty | |
| self._wake_all_readers() | |
| else: | |
| raise ValueError("can not release write lock without acquire it") | |
| else: | |
| raise ValueError("mode must be 'r' or 'w'") |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| # -*- coding: utf-8 -*- | |
| """ | |
| Copyright (c) 2008-2026 synodriver <diguohuangjiajinweijun@gmail.com> | |
| """ | |
| from typing import Callable, ParamSpec, TypeVar, Awaitable, NoReturn | |
| from threading import Lock | |
| import asyncio | |
| P = ParamSpec('P') | |
| R = TypeVar('R') | |
| class Caller: | |
| def __init__(self): | |
| self._val = None | |
| self._err = None | |
| self._done = asyncio.Event() | |
| self._done.clear() | |
| def set_result(self, val): | |
| self._val = val | |
| self._done.set() | |
| def set_exception(self, err): | |
| self._err = err | |
| self._done.set() | |
| async def result(self): | |
| await self._done.wait() | |
| if self._err is not None: | |
| raise self._err | |
| return self._val | |
| class SingleFlight: | |
| def __init__(self): | |
| self._cached = {} | |
| self._lock = Lock() | |
| async def do(self, key: str, func: Callable[P, Awaitable[R | NoReturn]], *args: P.args, **kwargs: P.kwargs) -> R | NoReturn: | |
| self._lock.acquire() | |
| if key in self._cached: | |
| caller = self._cached[key] | |
| self._lock.release() | |
| return await caller.result() | |
| caller = Caller() | |
| self._cached[key] = caller | |
| self._lock.release() | |
| try: | |
| val = await func(*args, **kwargs) | |
| caller.set_result(val) | |
| return val | |
| except BaseException as e: | |
| caller.set_exception(e) | |
| raise | |
| finally: | |
| del self._cached[key] |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| # -*- coding: utf-8 -*- | |
| """ | |
| Copyright (c) 2008-2026 synodriver <diguohuangjiajinweijun@gmail.com> | |
| """ | |
| import asyncio | |
| from unittest import IsolatedAsyncioTestCase | |
| import unittest | |
| from singleflight import SingleFlight | |
| class SingleFlightTestCase(IsolatedAsyncioTestCase): | |
| def setUp(self) -> None: | |
| self.sf = SingleFlight() | |
| async def test_singleflight(self): | |
| data = 0 | |
| async def func(): | |
| nonlocal data | |
| await asyncio.sleep(0.1) | |
| data += 1 | |
| return 32 | |
| t1 = asyncio.create_task(self.sf.do("key", func)) | |
| t2 = asyncio.create_task(self.sf.do("key", func)) | |
| t3 = asyncio.create_task(self.sf.do("key", func)) | |
| await asyncio.gather(t1, t2, t3) | |
| self.assertEqual(data, 1) | |
| self.assertEqual(t1.result(), 32) | |
| self.assertEqual(t2.result(), 32) | |
| self.assertEqual(t3.result(), 32) | |
| data = 0 | |
| t1 = asyncio.create_task(self.sf.do("key", func)) | |
| t2 = asyncio.create_task(self.sf.do("key", func)) | |
| t3 = asyncio.create_task(self.sf.do("key2", func)) | |
| t4 = asyncio.create_task(self.sf.do("key2", func)) | |
| await asyncio.gather(t1, t2, t3, t4) | |
| self.assertEqual(data, 2) | |
| async def test_singleflight_err(self): | |
| data = 0 | |
| async def func(): | |
| nonlocal data | |
| data += 1 | |
| await asyncio.sleep(0.1) | |
| raise ValueError | |
| t1 = asyncio.create_task(self.sf.do("key", func)) | |
| t2 = asyncio.create_task(self.sf.do("key", func)) | |
| t3 = asyncio.create_task(self.sf.do("key", func)) | |
| with self.assertRaises(ValueError): | |
| await asyncio.gather(t1, t2, t3) | |
| self.assertEqual(data, 1) | |
| async def test_cancel(self): | |
| data = 0 | |
| async def func(): | |
| nonlocal data | |
| await asyncio.sleep(5) | |
| data += 1 | |
| t1 = asyncio.create_task(self.sf.do("key", func)) | |
| t2 = asyncio.create_task(self.sf.do("key", func)) | |
| t3 = asyncio.create_task(self.sf.do("key", func)) | |
| await asyncio.sleep(0) # 必须有,没有这个,三个task还没调度,t1直接结束,没机会给t2和t3传递cancel | |
| t1.cancel() | |
| try: | |
| await t1 | |
| except asyncio.CancelledError: | |
| pass | |
| # with self.assertRaises(asyncio.CancelledError): | |
| # await t1 | |
| with self.assertRaises(asyncio.CancelledError): | |
| await t2 | |
| with self.assertRaises(asyncio.CancelledError): | |
| await t3 | |
| self.assertEqual(data, 0) | |
| if __name__ == '__main__': | |
| unittest.main() |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment