# 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.bean.olc_control_request import OlcControlRequest
from olc.bean.tag_group import TagGroup
from olc.control.context.context import OlcContext
from olc.control.handler.group_matcher_handler import GroupMatcherHandler
from olc.statistic.statistic import OlcStatistic


class TestGroupMatcherHandler(unittest.TestCase):

    def setUp(self):
        self.handler = GroupMatcherHandler()

    def test_in_coming_matches_groups(self):
        context = OlcContext(OlcControlRequest(tags={"env": "prod"}))
        groups = []
        tag_groups = [TagGroup(domain="test", name="group1")]

        mock_matcher = MagicMock()
        mock_matcher.match_tag_groups.return_value = tag_groups

        mock_factory = MagicMock()
        mock_factory.get_and_create.return_value = OlcStatistic()

        with patch('olc.control.handler.group_matcher_handler.OlcMatcherProvider.get_group_matcher', return_value=mock_matcher):
            with patch('olc.control.handler.group_matcher_handler.OlcStatisticFactory.get_instance', return_value=mock_factory):
                with patch.object(self.handler, '_next_in_coming') as mock_next:
                    self.handler.in_coming(context, groups)
                    mock_matcher.match_tag_groups.assert_called_once_with(context.request)
                    mock_next.assert_called_once_with(context, tag_groups)

    def test_in_coming_creates_statistics(self):
        context = OlcContext(OlcControlRequest(tags={"env": "prod"}))
        groups = []
        tag_groups = [
            TagGroup(domain="test", name="group1"),
            TagGroup(domain="test", name="group2"),
        ]

        mock_matcher = MagicMock()
        mock_matcher.match_tag_groups.return_value = tag_groups

        statistic1 = OlcStatistic()
        statistic2 = OlcStatistic()

        mock_factory = MagicMock()
        mock_factory.get_and_create.side_effect = [statistic1, statistic2]

        with patch('olc.control.handler.group_matcher_handler.OlcMatcherProvider.get_group_matcher', return_value=mock_matcher):
            with patch('olc.control.handler.group_matcher_handler.OlcStatisticFactory.get_instance', return_value=mock_factory):
                with patch.object(self.handler, '_next_in_coming'):
                    self.handler.in_coming(context, groups)
                    self.assertEqual(len(context.statistics), 2)
                    self.assertIn(tag_groups[0], context.statistics)
                    self.assertIn(tag_groups[1], context.statistics)

    def test_in_coming_with_empty_match(self):
        context = OlcContext(OlcControlRequest(tags={"env": "prod"}))
        groups = []

        mock_matcher = MagicMock()
        mock_matcher.match_tag_groups.return_value = []

        mock_factory = MagicMock()

        with patch('olc.control.handler.group_matcher_handler.OlcMatcherProvider.get_group_matcher', return_value=mock_matcher):
            with patch('olc.control.handler.group_matcher_handler.OlcStatisticFactory.get_instance', return_value=mock_factory):
                with patch.object(self.handler, '_next_in_coming') as mock_next:
                    self.handler.in_coming(context, groups)
                    self.assertEqual(len(context.statistics), 0)
                    mock_next.assert_called_once_with(context, [])

    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)

    def test_in_coming_with_none_request_tags(self):
        context = OlcContext(OlcControlRequest())
        groups = []
        tag_groups = [TagGroup(domain="test", name="group1")]

        mock_matcher = MagicMock()
        mock_matcher.match_tag_groups.return_value = tag_groups

        mock_factory = MagicMock()
        mock_factory.get_and_create.return_value = OlcStatistic()

        with patch('olc.control.handler.group_matcher_handler.OlcMatcherProvider.get_group_matcher', return_value=mock_matcher):
            with patch('olc.control.handler.group_matcher_handler.OlcStatisticFactory.get_instance', return_value=mock_factory):
                with patch.object(self.handler, '_next_in_coming') as mock_next:
                    self.handler.in_coming(context, groups)
                    mock_matcher.match_tag_groups.assert_called_once()

    def test_statistics_persisted_in_context(self):
        context = OlcContext(OlcControlRequest(tags={"env": "prod"}))
        groups = []
        tag_group = TagGroup(domain="test", name="group1")
        tag_groups = [tag_group]

        mock_matcher = MagicMock()
        mock_matcher.match_tag_groups.return_value = tag_groups

        statistic = OlcStatistic()
        mock_factory = MagicMock()
        mock_factory.get_and_create.return_value = statistic

        with patch('olc.control.handler.group_matcher_handler.OlcMatcherProvider.get_group_matcher', return_value=mock_matcher):
            with patch('olc.control.handler.group_matcher_handler.OlcStatisticFactory.get_instance', return_value=mock_factory):
                with patch.object(self.handler, '_next_in_coming'):
                    self.handler.in_coming(context, groups)
                    self.assertIs(context.statistics[tag_group], statistic)


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