# 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 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

    # ---- sync ----
    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:显式覆盖,证明 FlowProcessor 真的会调�?async 版本 ----
    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:
            # 必须�?async_*,不能走 sync_*
            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)
        # l1, l2 必须�?async_roll_back;l3 不应�?roll_back
        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:
            # match 返回一个非�?wrapper 列表
            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)
        # 没有�?sync
        self.assertFalse(any(c[0].startswith("sync_") for c in limiter.calls))


if __name__ == "__main__":
    unittest.main(verbosity=2)