# Copyright (c) Huawei Technologies Co., Ltd. 2026-2026. All rights reserved.
# OpenOLC is licensed under Mulan PSL v2.
# You can use this software according to the terms and conditions of the Mulan PSL v2.
# You may obtain a copy of Mulan PSL v2 at:
#         `http://license.coscl.org.cn/MulanPSL2`
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND,
# EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
# MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.
# See the Mulan PSL v2 for more details.

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()