import unittest
from unittest.mock import MagicMock, patch, PropertyMock
from olc.bean.match_wrapper import MatchWrapper
from olc.bean.olc_config_rule import FlowPolicy
from olc.bean.olc_control_request import OlcControlRequest
from olc.bean.tag_group import TagGroup
from olc.control.context.context import OlcContext
from olc.control.handler.flow_handler import FlowHandler, FlowProcessor
from olc.limit.limiter import Limit
class MockLimiter(Limit):
def __init__(self, wrapper, pre_check_result=True, try_acquire_result=True):
self._wrapper = wrapper
self._pre_check_result = pre_check_result
self._try_acquire_result = try_acquire_result
self.roll_back_called = 0
self.pre_check_called = 0
self.try_acquire_called = 0
def get_match_wrapper(self) -> MatchWrapper[FlowPolicy]:
return self._wrapper
def try_acquire(self, token_number: int) -> bool:
self.try_acquire_called += 1
return self._try_acquire_result
def pre_check(self, token_number: int) -> bool:
self.pre_check_called += 1
return self._pre_check_result
def roll_back(self, token_number: int):
self.roll_back_called += 1
def dec_token(self, token_number: int):
pass
def create_flow_policy(name="test_policy", rate_limit=100):
return FlowPolicy(
name=name,
category="flow",
enabled=True,
block_msg="blocked",
policy_type="NODE",
time_unit="second",
time_interval=1,
rate_limit=rate_limit,
burst_limit=0,
flow_control_mode="qps",
max_wait_time_ms=0,
calculate_alg=None,
assign_alg=None,
)
def create_match_wrapper(group_name="group1", policy_name="policy1"):
tag_group = TagGroup(domain="test", name=group_name)
policy = create_flow_policy(name=policy_name)
return MatchWrapper(tag_group, policy)
class TestFlowHandler(unittest.TestCase):
def setUp(self):
self.handler = FlowHandler()
def test_in_coming_with_stop_flag(self):
context = OlcContext(OlcControlRequest())
context.stop = True
groups = [TagGroup(domain="test", name="group1")]
with patch.object(self.handler, '_next_in_coming') as mock_next:
self.handler.in_coming(context, groups)
mock_next.assert_called_once_with(context, groups)
def test_in_coming_with_empty_match(self):
context = OlcContext(OlcControlRequest())
groups = [TagGroup(domain="test", name="group1")]
with patch('olc.control.handler.flow_handler.OlcMatcherProvider.get_flow_matcher') as mock_matcher:
mock_matcher.return_value.match.return_value = []
with patch.object(self.handler, '_next_in_coming') as mock_next:
self.handler.in_coming(context, groups)
self.assertFalse(context.result.block)
mock_next.assert_called_once()
def test_in_coming_allowed(self):
context = OlcContext(OlcControlRequest())
groups = [TagGroup(domain="test", name="group1")]
wrapper = create_match_wrapper()
with patch('olc.control.handler.flow_handler.OlcMatcherProvider.get_flow_matcher') as mock_matcher:
mock_matcher.return_value.match.return_value = [wrapper]
with patch('olc.control.handler.flow_handler.FlowProcessor.process', return_value=True):
with patch.object(self.handler, '_next_in_coming') as mock_next:
self.handler.in_coming(context, groups)
self.assertFalse(context.result.block)
self.assertFalse(context.stop)
def test_in_coming_blocked(self):
context = OlcContext(OlcControlRequest())
groups = [TagGroup(domain="test", name="group1")]
wrapper = create_match_wrapper()
with patch('olc.control.handler.flow_handler.OlcMatcherProvider.get_flow_matcher') as mock_matcher:
mock_matcher.return_value.match.return_value = [wrapper]
with patch('olc.control.handler.flow_handler.FlowProcessor.process', return_value=False):
with patch.object(self.handler, '_next_in_coming') as mock_next:
self.handler.in_coming(context, groups)
self.assertTrue(context.result.block)
self.assertTrue(context.stop)
def test_out_coming_calls_next(self):
context = OlcContext(OlcControlRequest())
with patch.object(self.handler, '_next_out_coming') as mock_next:
self.handler.out_coming(context)
mock_next.assert_called_once_with(context)
class TestFlowProcessor(unittest.TestCase):
def test_process_empty_limiters(self):
context = OlcContext(OlcControlRequest())
wrappers = [create_match_wrapper()]
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = []
result = FlowProcessor.process(context, wrappers)
self.assertTrue(result)
def test_process_all_limiters_pass(self):
context = OlcContext(OlcControlRequest(token_number=1))
wrapper = create_match_wrapper()
limiter = MockLimiter(wrapper, pre_check_result=True, try_acquire_result=True)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter]
result = FlowProcessor.process(context, [wrapper])
self.assertTrue(result)
self.assertEqual(limiter.pre_check_called, 1)
self.assertEqual(limiter.try_acquire_called, 1)
self.assertEqual(limiter.roll_back_called, 0)
def test_process_pre_check_fails(self):
context = OlcContext(OlcControlRequest(token_number=1))
wrapper = create_match_wrapper()
limiter = MockLimiter(wrapper, pre_check_result=False, try_acquire_result=True)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter]
result = FlowProcessor.process(context, [wrapper])
self.assertFalse(result)
self.assertEqual(context.result.block_rule, wrapper)
def test_process_try_acquire_fails(self):
context = OlcContext(OlcControlRequest(token_number=1))
wrapper = create_match_wrapper()
limiter = MockLimiter(wrapper, pre_check_result=True, try_acquire_result=False)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter]
result = FlowProcessor.process(context, [wrapper])
self.assertFalse(result)
self.assertEqual(context.result.block_rule, wrapper)
def test_process_multiple_limiters_first_fails(self):
context = OlcContext(OlcControlRequest(token_number=1))
wrapper1 = create_match_wrapper(group_name="group1", policy_name="policy1")
wrapper2 = create_match_wrapper(group_name="group2", policy_name="policy2")
limiter1 = MockLimiter(wrapper1, pre_check_result=False)
limiter2 = MockLimiter(wrapper2, pre_check_result=True, try_acquire_result=True)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter1, limiter2]
result = FlowProcessor.process(context, [wrapper1, wrapper2])
self.assertFalse(result)
self.assertEqual(limiter1.pre_check_called, 1)
self.assertEqual(limiter2.pre_check_called, 0)
def test_process_multiple_limiters_second_fails_rollback_first(self):
context = OlcContext(OlcControlRequest(token_number=1))
wrapper1 = create_match_wrapper(group_name="group1", policy_name="policy1")
wrapper2 = create_match_wrapper(group_name="group2", policy_name="policy2")
limiter1 = MockLimiter(wrapper1, pre_check_result=True, try_acquire_result=True)
limiter2 = MockLimiter(wrapper2, pre_check_result=False)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter1, limiter2]
result = FlowProcessor.process(context, [wrapper1, wrapper2])
self.assertFalse(result)
self.assertEqual(limiter1.roll_back_called, 1)
self.assertEqual(limiter2.roll_back_called, 0)
def test_process_with_token_number(self):
context = OlcContext(OlcControlRequest(token_number=5))
wrapper = create_match_wrapper()
limiter = MockLimiter(wrapper)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter]
result = FlowProcessor.process(context, [wrapper])
self.assertTrue(result)
class TestFlowProcessorEdgeCases(unittest.TestCase):
def test_process_with_none_wrappers(self):
context = OlcContext(OlcControlRequest())
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = []
result = FlowProcessor.process(context, None)
self.assertTrue(result)
def test_process_with_zero_token_number(self):
context = OlcContext(OlcControlRequest(token_number=0))
wrapper = create_match_wrapper()
limiter = MockLimiter(wrapper)
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = [limiter]
result = FlowProcessor.process(context, [wrapper])
self.assertTrue(result)
def test_process_multiple_limiters_all_pass(self):
context = OlcContext(OlcControlRequest(token_number=1))
wrappers = [
create_match_wrapper(group_name=f"group{i}", policy_name=f"policy{i}")
for i in range(3)
]
limiters = [MockLimiter(w) for w in wrappers]
with patch('olc.control.handler.flow_handler.OlcRateLimiterFactory.get_instance') as mock_factory:
mock_factory.return_value.get_rate_limiters.return_value = limiters
result = FlowProcessor.process(context, wrappers)
self.assertTrue(result)
for limiter in limiters:
self.assertEqual(limiter.roll_back_called, 0)
if __name__ == '__main__':
unittest.main()