#!/usr/bin/env python3
#
# Copyright 2025, The Android Open Source Project
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Unittests for smart_test_finder."""

# pylint: disable=invalid-name

import pathlib
import unittest
from unittest import mock
from atest import atest_utils
from atest import constants
from atest import unittest_constants
from atest.proto import decision_graph_pb2
from atest.test_finders import test_finder_utils
from atest.test_finders.smart_test_finder import atp_test_selector
from atest.test_finders.smart_test_finder import local_info_collector
from atest.test_finders.smart_test_finder import smart_test_filter
from atest.test_finders.smart_test_finder import smart_test_finder
from atest.test_finders.smart_test_finder import test_relevance_client
from google.protobuf import json_format
from pyfakefs import fake_filesystem_unittest

_FAKE_BLOCKLIST_CONTENT = """module,reason
BlockedModule, test_reaon"""

_FAKE_LOOKUP_TABLE_CONTENT = """branch,target,test_name,test_id,postsubmit_pass_rate,test_run_duration_ms_past7days
some_branch,some_target,TestA,a_id,0.99,10000
some_branch2,some_target2,TestFlakyTest,b_id,0.90,500
some_branch3,some_target3,TestRunTimeNotFound,c_id,0.99,-1
some_branch4,some_target4,TestD,d_id,0.99,40000
some_branch6,some_target6,TestF,f_id,0.99,500
some_branch7,some_target7,TestG,g_id,0.98,5
some_branch5,some_target5,TestNotSelectedDueToTimeLimit,e_id,0.99,20000
some_branch8,some_target8,TestNotSelectedDueToTimeLimit2,h_id,0.98,5
some_branch2,some_target2,TestWithoutModule,j_id,0.99,500
some_branch2,some_target2,TestWithoutTestClass,k_id,0.99,500"""

_FAKE_CHANGED_FILE_DETAILS = frozenset([
    atest_utils.ChangedFileDetails(
        filename='/a/b/c',
        number_of_lines_inserted=14,
        number_of_lines_deleted=25,
    ),
    atest_utils.ChangedFileDetails(
        filename='/d/e/f',
        number_of_lines_inserted=36,
        number_of_lines_deleted=47,
    ),
])
_FAKE_CHANGE_INFO = local_info_collector.ChangeInfo(
    project='fake_project',
    branch='fake_branch',
    remote_hostname='stuff-to-be-selected',
    changed_files=_FAKE_CHANGED_FILE_DETAILS,
    user_key='fake_user',
)
_FAKE_SELECTED_TESTS = [
    atp_test_selector.AtpTestInfo(
        name='v2/android-virtual-infra/test_mapping/presubmit-avd',
        target='aosp_cf_x86_64_phone-trunk_staging-userdebug',
        branch='some_aosp-branch2',
    ),
    atp_test_selector.AtpTestInfo(
        name='v2/android-test-harness-team/tradefed/host_unit_tests_zip_validation',
        target='aosp_cf_x86_64_phone-trunk_staging-userdebug',
        branch='some_aosp-branch2',
    ),
]


def _get_decision_graph_check(
    info: smart_test_filter.TestClassInfo,
) -> decision_graph_pb2.Check:
  ants_test = decision_graph_pb2.AnTSTest(
      test_identifier=decision_graph_pb2.TestIdentifier(
          module=info.module,
          test_class=info.test_class,
      ),
      test_identifier_id=info.test_id,
  )
  check_reason = decision_graph_pb2.Check.Reason(relevance_score=info.score)
  return decision_graph_pb2.Check(
      identifier=decision_graph_pb2.Check.Identifier(
          ants_test=ants_test,
      ),
      reason=check_reason,
  )


# pylint: disable=protected-access
class SmartTestFinderFilmsystemUnittests(fake_filesystem_unittest.TestCase):
  """Unit tests for smart_test_finder.py with filesystem access."""

  def setUp(self):
    super().setUp()
    self.setUpPyfakefs()

    self.fake_lookup_table_path = str(
        pathlib.Path(constants.SMART_TEST_SELECTION_ROOT_PATH)
        / 'lookup_tables/tests_with_runtime_and_pass_rate.csv'
    )
    self.fs.create_file(
        self.fake_lookup_table_path,
        contents=_FAKE_LOOKUP_TABLE_CONTENT,
    )
    self.fake_blocklist_path = str(
        pathlib.Path(constants.SMART_TEST_SELECTION_ROOT_PATH)
        / 'lookup_tables/blocklist.csv'
    )
    self.fs.create_file(
        self.fake_blocklist_path,
        contents=_FAKE_BLOCKLIST_CONTENT,
    )

  # TODO(b/410945183): Change this test once the bug is fixed.
  @mock.patch('uuid.uuid4', side_effect=['001-002-003', '002-003-004'])
  @mock.patch.object(
      atp_test_selector,
      'get_selected_atp_tests',
      return_value=_FAKE_SELECTED_TESTS,
  )
  @mock.patch.object(
      local_info_collector,
      'get_local_change_info',
      return_value=_FAKE_CHANGE_INFO,
  )
  @mock.patch.object(test_relevance_client, 'TestRelevanceClient')
  def test_get_smartly_selected_tests(self, mock_client_class, _, __, ___):
    checks = [
        # TestA is selected and ranked first, because it has the highest
        # relevance score.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='a_id',
                atp_test_name='TestA',
                branch='some_branch',
                target='some_target',
                module='TestAModule',
                test_class='testAClass',
                score=1,
            )
        ),
        # This test is not selected because it is blocked.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='blocked_id',
                atp_test_name='SomeTest',
                branch='some_branch',
                target='some_target',
                module='BlockedModule',
                test_class='testClass',
                score=1,
            )
        ),
        # This test is not selected because of no history in the lookup table.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='id_not_found',
                atp_test_name='SomeTest',
                branch='some_branch',
                target='some_target',
                module='SomeTestModule',
                test_class='testSomeClass',
                score=0.98,
            )
        ),
        # This test is not selected because the module name is missing.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='g_id',
                atp_test_name='TestWithoutModule',
                branch='some_branch2',
                target='some_target2',
                test_class='testClass',
                score=1.0,
            )
        ),
        # This test is not selected because the test class name is missing.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='h_id',
                atp_test_name='TestWithoutTestClass',
                branch='some_branch2',
                target='some_target2',
                module='TestWithoutTestClass',
                score=0.99,
            )
        ),
        # This test is not selected because the passing rate of this test class
        # was only 0.90, less than the threshold 0.95.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='b_id',
                atp_test_name='TestFlakyTest',
                branch='some_branch2',
                target='some_target2',
                module='SomeTestModule',
                test_class='testSomeClass',
                score=0.99,
            )
        ),
        # This test is not selected because the execution time history of this
        # test is missing.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='c_id',
                atp_test_name='TestRunTimeNotFound',
                branch='some_branch3',
                target='some_target3',
                module='SomeTestModule',
                test_class='testSomeClass',
                score=0.995,
            )
        ),
        # TestD is selected and ranked right after TestA, because it has the
        # second highest relevance score.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='d_id',
                atp_test_name='TestD',
                branch='some_branch4',
                target='some_target4',
                module='TestDModule',
                test_class='testDClass',
                score=0.98,
            )
        ),
        # After TestA, TestD, TestF and TestG are selected, selecting this test
        # would result in exceeding the estimated execution time (in this test,
        # the user specified the time limit to be one minute), so this test is
        # not selected.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='e_id',
                atp_test_name='TestNotSelectedDueToTimeLimit',
                branch='some_branch5',
                target='some_target5',
                module='TestEModule',
                test_class='testEClass',
                score=0.95,
            )
        ),
        # TestF is selected and ranked right after TestG, because it has the
        # fourth highest relevance score.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='f_id',
                atp_test_name='TestF',
                branch='some_branch6',
                target='some_target6',
                module='TestFModule',
                test_class='TestFModule.TestFClass',
                score=0.96,
            )
        ),
        # TestG is selected and ranked right after TestD, because it has the
        # third highest relevance score of all valid tests.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='g_id',
                atp_test_name='TestG',
                branch='some_branch7',
                target='some_target7',
                module='TestGModule',
                test_class='testGClass',
                score=0.97,
            )
        ),
        # When the decision of not selecting TestNotSelectedDueToTimeLimit was
        # made, we no longer look into the details of other tests with even
        # lower relevant scores, so this test is not selected.
        _get_decision_graph_check(
            smart_test_filter.TestClassInfo(
                test_id='h_id',
                atp_test_name='TestNotSelectedDueToTimeLimit2',
                branch='some_branch8',
                target='some_target8',
                module='SomeTestModule',
                test_class='testSomeClass',
                score=0.949,
            )
        ),
    ]
    dg_outputs = []
    for check in checks:
      dg_output = json_format.MessageToDict(
          decision_graph_pb2.DecisionGraphOutput(
              outputs=[decision_graph_pb2.StageOutput(checks=[check])],
          )
      )
      dg_outputs.append(dg_output)
    mock_client = mock_client_class.return_value
    mock_client.get_tests_with_relevance_score_query_by_query.return_value = (
        dg_outputs
    )

    final_selected_tests = smart_test_finder.get_smartly_selected_tests(
        time_limit_in_minutes=1,
    )

    self.assertEqual(
        final_selected_tests,
        [
            'TestAModule:testAClass',
            'TestDModule:testDClass',
            'TestGModule:testGClass',
            'TestFModule:TestFClass',
        ],
    )

  @mock.patch.object(atp_test_selector, 'get_selected_atp_tests')
  @mock.patch.object(local_info_collector, 'get_local_change_info')
  @mock.patch.object(test_relevance_client, 'TestRelevanceClient')
  def test_get_smartly_selected_tests_return_no_tests_if_api_broken(
      self, mock_client_class, _, __
  ):
    mock_client = mock_client_class.return_value
    mock_client.get_tests_with_relevance_score_query_by_query.side_effect = (
        TimeoutError()
    )
    with self.assertRaises(TimeoutError):
      smart_test_finder.get_smartly_selected_tests()

  @mock.patch.object(local_info_collector, 'get_local_change_info')
  def test_get_smartly_selected_tests_return_no_tests_with_no_changes(
      self, mock_local_info_collector
  ):
    CHANGE_INFO_WITH_NO_CHANGED_FILES = local_info_collector.ChangeInfo(
        project='fake_project',
        branch='fake_branch',
        remote_hostname='stuff-to-be-selected',
        changed_files=[],
        user_key='fake_user',
    )
    mock_local_info_collector.return_value = CHANGE_INFO_WITH_NO_CHANGED_FILES
    results = smart_test_finder.get_smartly_selected_tests()
    self.assertEqual(results, [])

  @mock.patch.object(
      test_finder_utils,
      'find_host_unit_tests',
      return_value=[
          unittest_constants.CLASS_NAME,
          unittest_constants.MODULE2_NAME,
      ],
  )
  @mock.patch('os.getcwd', return_value='/my/main/some/project')
  @mock.patch.object(local_info_collector, 'get_local_change_info')
  def test_get_smartly_selected_tests_return_host_unit_tests_but_no_relevance_score_based_tests(
      self, mock_local_info_collector, _, __
  ):
    CHANGE_INFO_WITH_NO_CHANGED_FILES = local_info_collector.ChangeInfo(
        project='fake_project',
        branch='fake_branch',
        remote_hostname='stuff-to-be-selected',
        changed_files=[],
        user_key='fake_user',
    )
    mock_local_info_collector.return_value = CHANGE_INFO_WITH_NO_CHANGED_FILES

    results = smart_test_finder.get_smartly_selected_tests(
        mod_info=unittest_constants.MODULE_INFO, root_dir='/my/main'
    )

    self.assertCountEqual(
        results,
        [unittest_constants.CLASS_NAME, unittest_constants.MODULE2_NAME],
    )

  @mock.patch.object(local_info_collector, 'get_local_change_info')
  def test_get_smartly_selected_tests_return_no_host_unit_tests_per_user_specification(
      self,
      mock_local_info_collector,
  ):
    CHANGE_INFO_WITH_NO_CHANGED_FILES = local_info_collector.ChangeInfo(
        project='fake_project',
        branch='fake_branch',
        remote_hostname='stuff-to-be-selected',
        changed_files=[],
        user_key='fake_user',
    )
    mock_local_info_collector.return_value = CHANGE_INFO_WITH_NO_CHANGED_FILES

    results = smart_test_finder.get_smartly_selected_tests(
        include_host_unit_tests=False,
        mod_info=unittest_constants.MODULE_INFO,
        root_dir='/my/main',
    )

    self.assertEqual(results, [])


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