import asyncio
import unittest
from unittest.mock import MagicMock
from olc.control.context.context import OlcContext
from olc.control.handler.flow_handler import FlowProcessor
from olc.limit.limiter import Limit
def _arun(coro):
return asyncio.new_event_loop().run_until_complete(coro)
class _RecordingLimit(Limit):
"""记录所有同�?异步方法调用次数�?limiter,便于断言走了哪条路径"""
def __init__(self, name="rec", allow=True):
self.name = name
self.allow = allow
self.calls = []
self._wrapper = MagicMock()
self._wrapper.get_policy.return_value = MagicMock()
self._wrapper.tag_group.name = name
def get_match_wrapper(self):
return self._wrapper
def try_acquire(self, n):
self.calls.append(("sync_try_acquire", n))
return self.allow
def pre_check(self, n):
self.calls.append(("sync_pre_check", n))
return True
def roll_back(self, n):
self.calls.append(("sync_roll_back", n))
def dec_token(self, n):
self.calls.append(("sync_dec_token", n))
async def async_try_acquire(self, n):
self.calls.append(("async_try_acquire", n))
return self.allow
async def async_pre_check(self, n):
self.calls.append(("async_pre_check", n))
return True
async def async_roll_back(self, n):
self.calls.append(("async_roll_back", n))
async def async_dec_token(self, n):
self.calls.append(("async_dec_token", n))
def _make_context(token_number=1):
request = MagicMock()
request.token_number = token_number
ctx = OlcContext(request)
return ctx
class TestFlowProcessorAsyncPath(unittest.TestCase):
"""证明 FlowProcessor.async_process 真的串联 async_* 方法"""
def test_all_limiters_allow_uses_async_methods(self):
limiters = [_RecordingLimit(name=f"l{i}", allow=True) for i in range(3)]
from unittest.mock import patch
with patch(
"olc.control.handler.flow_handler.OlcRateLimiterFactory"
) as factory_mock:
factory_mock.get_instance.return_value.get_rate_limiters.return_value = limiters
ctx = _make_context(token_number=2)
allowed = _arun(FlowProcessor.async_process(ctx, [MagicMock()] * 3))
self.assertTrue(allowed)
for l in limiters:
sync_called = any(c[0].startswith("sync_") for c in l.calls)
self.assertFalse(
sync_called,
f"{l.name} 走了同步路径:{l.calls}",
)
self.assertTrue(("async_pre_check", 2) in l.calls)
self.assertTrue(("async_try_acquire", 2) in l.calls)
def test_fail_short_circuits_and_rolls_back_async(self):
"""�?2 �?limiter 允许,第 3 个拒绝;�?2 个应�?async_roll_back"""
l1 = _RecordingLimit(name="l1", allow=True)
l2 = _RecordingLimit(name="l2", allow=True)
l3 = _RecordingLimit(name="l3", allow=False)
limiters = [l1, l2, l3]
from unittest.mock import patch
with patch(
"olc.control.handler.flow_handler.OlcRateLimiterFactory"
) as factory_mock:
factory_mock.get_instance.return_value.get_rate_limiters.return_value = limiters
ctx = _make_context(token_number=1)
allowed = _arun(FlowProcessor.async_process(ctx, [MagicMock()] * 3))
self.assertFalse(allowed)
self.assertTrue(("async_roll_back", 1) in l1.calls)
self.assertTrue(("async_roll_back", 1) in l2.calls)
self.assertFalse(any("roll_back" in c[0] for c in l3.calls))
class TestFlowHandlerAsyncPath(unittest.TestCase):
"""证明 FlowHandler.async_in_coming / async_out_coming 串联完整"""
def test_async_in_coming_uses_async_processor(self):
from olc.control.handler.flow_handler import FlowHandler
from unittest.mock import patch
limiter = _RecordingLimit(name="L", allow=True)
with patch(
"olc.control.handler.flow_handler.OlcMatcherProvider"
) as matcher_mock, patch(
"olc.control.handler.flow_handler.OlcRateLimiterFactory"
) as factory_mock:
wrapper = MagicMock()
matcher_mock.get_flow_matcher.return_value.match.return_value = [wrapper]
factory_mock.get_instance.return_value.get_rate_limiters.return_value = [limiter]
handler = FlowHandler()
ctx = _make_context(token_number=1)
ctx.stop = False
ctx.add_match_wrapper = MagicMock()
_arun(handler.async_in_coming(ctx, None))
self.assertTrue(("async_try_acquire", 1) in limiter.calls)
self.assertFalse(any(c[0].startswith("sync_") for c in limiter.calls))
if __name__ == "__main__":
unittest.main(verbosity=2)