Skip to content

Instantly share code, notes, and snippets.

@synodriver
Last active April 25, 2026 12:10
Show Gist options
  • Select an option

  • Save synodriver/dcf1fac7e553b7998523c7073f83b5a7 to your computer and use it in GitHub Desktop.

Select an option

Save synodriver/dcf1fac7e553b7998523c7073f83b5a7 to your computer and use it in GitHub Desktop.
like singleflight in go
# -*- 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'")
# -*- 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]
# -*- 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