Skip to content

Commit 71bb85f

Browse files
author
zach
committed
cleanup: add locking
1 parent c4eae6d commit 71bb85f

2 files changed

Lines changed: 120 additions & 30 deletions

File tree

‎extism/pool.py‎

Lines changed: 47 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -30,50 +30,67 @@ class PoolError(Exception):
3030
pass
3131

3232
class Pool:
33-
plugins: Dict[str, Callable[[], Plugin]]
34-
instances: Dict[str, List[PoolPlugin]]
33+
_plugins: Dict[str, Callable[[], Plugin]]
34+
_instances: Dict[str, List[PoolPlugin]]
35+
_locks: Dict[str, Lock]
3536

3637
def __init__(self, max_instances=1):
3738
self.max_instances = max_instances
38-
self.instances = {}
39-
self.plugins = {}
40-
self.count = {}
39+
self._instances = {}
40+
self._locks = {}
41+
self._plugins = {}
4142

4243
def add(self, name, source: Callable[[], Plugin]):
43-
self.plugins[name] = source
44-
self.instances[name] = []
45-
46-
def find_available(self, name):
47-
entry = self.instances[name]
44+
self._plugins[name] = source
45+
self._instances[name] = []
46+
self._locks[name] = Lock()
47+
48+
def count(self, name):
49+
self._locks[name].acquire()
50+
n = len(self._instances[name])
51+
self._locks[name].release()
52+
return n
53+
54+
def _find_available(self, name):
55+
entry = self._instances[name]
4856
for instance in entry:
4957
if not instance.active:
5058
return instance.make_active()
5159
return None
5260

5361
async def async_get(self, name, timeout=None):
5462
start = time.time()
55-
entry = self.instances[name]
56-
57-
p = self.find_available(name)
58-
if p is not None:
59-
return p
60-
61-
if len(entry) < self.max_instances:
62-
p = PoolPlugin(self.plugins[name](), active=True)
63-
entry.append(p)
64-
self.instances[name] = entry
65-
return p
66-
67-
while True:
68-
p = self.find_available(name)
63+
entry = self._instances[name]
64+
lock = self._locks[name]
65+
lock.acquire()
66+
try:
67+
p = self._find_available(name)
6968
if p is not None:
69+
lock.release()
7070
return p
71-
else:
72-
if timeout is None:
73-
await asyncio.sleep(0)
74-
continue
75-
elif (time.time() - start) >= timeout:
76-
raise PoolError("Timed out getting instance for key " + name)
71+
72+
if len(entry) < self.max_instances:
73+
p = PoolPlugin(self._plugins[name](), active=True)
74+
entry.append(p)
75+
self._instances[name] = entry
76+
lock.release()
77+
return p
78+
79+
while True:
80+
p = self._find_available(name)
81+
if p is not None:
82+
lock.release()
83+
return p
84+
else:
85+
if timeout is None:
86+
await asyncio.sleep(0)
87+
continue
88+
elif (time.time() - start) >= timeout:
89+
lock.release()
90+
raise PoolError("Timed out getting instance for key " + name)
91+
except Exception as exc:
92+
lock.release()
93+
raise exc
7794

7895
def get(self, name, timeout=None):
7996
fut = self.async_get(name, timeout)

‎tests/test_extism.py‎

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -183,3 +183,76 @@ def read_test_wasm(p):
183183
path = join(dirname(__file__), "..", "wasm", p)
184184
with open(path, "rb") as wasm_file:
185185
return wasm_file.read()
186+
187+
188+
# pool = Pool(max_instances=5)
189+
# manifest = {'wasm': [{'path': 'code.wasm'}]}
190+
# pool.add('test', lambda: Plugin(manifest, wasi=True))
191+
192+
# def run_test(sleep, input):
193+
# with pool.get('test') as plugin:
194+
# time.sleep(sleep)
195+
# print(plugin.call('count_vowels', input))
196+
197+
# def test_thread(sleep, input):
198+
# t = Thread(target=run_test, args=[sleep, input])
199+
# t.start()
200+
# return t
201+
202+
# threads = [
203+
# test_thread(1, 'aaa'),
204+
# test_thread(1, 'aaa'),
205+
# test_thread(1, 'aaa'),
206+
# test_thread(1, 'aaa'),
207+
# test_thread(1, 'aaa'),
208+
# test_thread(1, 'aaa'),
209+
# test_thread(2, ''),
210+
# test_thread(2, ''),
211+
# test_thread(2, ''),
212+
# test_thread(2, ''),
213+
# test_thread(2, ''),
214+
# test_thread(2, ''),
215+
# test_thread(0, 'abc'),
216+
# test_thread(0, 'abc'),
217+
# test_thread(0, 'abc'),
218+
# test_thread(0, 'abc'),
219+
# test_thread(0, 'abc'),
220+
# test_thread(0, 'abc'),
221+
# ]
222+
223+
# for t in threads:
224+
# t.join()
225+
226+
# async def test_async_inner(sleep, input):
227+
# with await pool.get('test') as plugin:
228+
# await asyncio.sleep(sleep)
229+
# print(plugin.call('count_vowels', input))
230+
231+
# async def test_async(*args):
232+
# await asyncio.create_task(test_async_inner(*args))
233+
234+
# futures = [
235+
# test_async(1, 'aaa'),
236+
# test_async(1, 'aaa'),
237+
# test_async(1, 'aaa'),
238+
# test_async(1, 'aaa'),
239+
# test_async(1, 'aaa'),
240+
# test_async(1, 'aaa'),
241+
# test_async(2, ''),
242+
# test_async(2, ''),
243+
# test_async(2, ''),
244+
# test_async(2, ''),
245+
# test_async(2, ''),
246+
# test_async(2, ''),
247+
# test_async(0, 'abc'),
248+
# test_async(0, 'abc'),
249+
# test_async(0, 'abc'),
250+
# test_async(0, 'abc'),
251+
# test_async(0, 'abc'),
252+
# test_async(0, 'abc'),
253+
# ]
254+
255+
# async def main():
256+
# await asyncio.gather(*futures)
257+
258+
# asyncio.get_event_loop().run_until_complete(main())

0 commit comments

Comments
 (0)