# 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

from olc.admission.admission_controller import (
    AbsAdmissionController,
    AllowController,
    DenyController,
    DynamicController,
)
from olc.admission.admission_factory import AdmissionFactory
from olc.bean.match_wrapper import MatchWrapper
from olc.bean.olc_config_rule import AdmissionPolicy
from olc.bean.olc_control_request import OlcControlRequest
from olc.bean.policy_category import PolicyCategory
from olc.bean.tag_group import TagGroup
from olc.control.context.context import OlcContext
from olc.control.handler.admission_handler import AdmissionHandler, _process
from olc.exception.exception import BlackListException, WhiteListException


def _make_tag_group(name: str = "test_group", domain: str = "test_domain") -> TagGroup:
    return TagGroup(domain=domain, name=name, enabled=True, priority=500)


def _make_admission_policy(
    name: str = "test_policy",
    policy_type: str = "allow",
    category: str = PolicyCategory.ADMISSION.value,
    enabled: bool = True,
    block_msg: str = "",
) -> AdmissionPolicy:
    return AdmissionPolicy(
        name=name,
        category=category,
        enabled=enabled,
        block_msg=block_msg,
        type=policy_type,
        calculate_alg=None,
    )


def _make_wrapper(tag_group: TagGroup, policy: AdmissionPolicy) -> MatchWrapper[AdmissionPolicy]:
    return MatchWrapper(tag_group, policy)


class MockController(AbsAdmissionController):
    def __init__(self, get_result_value: bool = True):
        super().__init__()
        self._get_result_value = get_result_value
        self.get_result_called = 0

    def refresh_status(self):
        pass

    def get_result(self, context: OlcContext) -> bool:
        self.get_result_called += 1
        return self._get_result_value


class TestProcess(unittest.TestCase):

    def test_process_no_matched_policies_returns_true(self):
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group()]

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = []
            result = _process(context, groups)
            self.assertFalse(result)

    def test_process_allow_controller_allows(self):
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_allow")]
        wrapper = _make_wrapper(
            _make_tag_group("group_allow"),
            _make_admission_policy("allow_policy", "allow")
        )

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = [wrapper]
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_controller = AllowController()
                mock_factory.return_value.get_controller.return_value = mock_controller
                with self.assertRaises(WhiteListException):
                    _process(context, groups)

    def test_process_deny_controller_blocks(self):
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_deny")]
        wrapper = _make_wrapper(
            _make_tag_group("group_deny"),
            _make_admission_policy("deny_policy", "deny")
        )

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = [wrapper]
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_controller = DenyController()
                mock_factory.return_value.get_controller.return_value = mock_controller
                with self.assertRaises(BlackListException):
                    _process(context, groups)

    def test_process_controller_is_none_skipped(self):
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_none")]
        wrapper = _make_wrapper(
            _make_tag_group("group_none"),
            _make_admission_policy("none_policy", "dynamic")
        )

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = [wrapper]
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_factory.return_value.get_controller.return_value = None
                result = _process(context, groups)
                self.assertFalse(result)

    def test_process_multiple_policies_all_allow_returns_false(self):
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_multi")]
        wrappers = [
            _make_wrapper(_make_tag_group("group_multi"), _make_admission_policy("allow1", "allow")),
            _make_wrapper(_make_tag_group("group_multi"), _make_admission_policy("allow2", "allow")),
        ]

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = wrappers
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_controller = AllowController()
                mock_factory.return_value.get_controller.return_value = mock_controller
                with self.assertRaises(WhiteListException):
                    _process(context, groups)

    def test_process_logs_debug_for_match_result(self):
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_log")]
        wrappers = [_make_wrapper(_make_tag_group("group_log"), _make_admission_policy("log_policy", "allow"))]

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = wrappers
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_controller = AllowController()
                mock_factory.return_value.get_controller.return_value = mock_controller
                with self.assertLogs('olc.control.handler.admission_handler', level='DEBUG') as cm:
                    with self.assertRaises(WhiteListException):
                        _process(context, groups)
                self.assertTrue(any("match admission policy" in msg for msg in cm.output))

    def test_process_non_admission_category_skipped(self):
        """测试非ADMISSION类别的策略被跳过"""
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_non_admission")]
        wrapper = _make_wrapper(
            _make_tag_group("group_non_admission"),
            _make_admission_policy("flow_policy", "allow", PolicyCategory.FLOW.value)
        )

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = [wrapper]
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_factory.return_value.get_controller.return_value = AllowController()
                result = _process(context, groups)
                self.assertFalse(result)

    def test_process_mixed_admission_and_non_admission_policies(self):
        """测试混合ADMISSION和非ADMISSION策略时的过滤行为"""
        context = OlcContext(OlcControlRequest())
        groups = [_make_tag_group("group_mixed")]

        wrapper_flow = _make_wrapper(
            _make_tag_group("group_mixed"),
            _make_admission_policy("flow_policy", "allow", PolicyCategory.FLOW.value)
        )
        wrapper_admission = _make_wrapper(
            _make_tag_group("group_mixed"),
            _make_admission_policy("admission_policy", "allow", PolicyCategory.ADMISSION.value)
        )

        with patch('olc.control.handler.admission_handler.OlcMatcherProvider.get_admission_matcher') as mock_matcher:
            mock_matcher.return_value.match.return_value = [wrapper_flow, wrapper_admission]
            with patch('olc.control.handler.admission_handler.AdmissionFactory.get_instance') as mock_factory:
                mock_controller = AllowController()
                mock_factory.return_value.get_controller.return_value = mock_controller
                with self.assertRaises(WhiteListException):
                    _process(context, groups)


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