import json import pickle as _pickle import pytest from DBotTrainClustering import ( MESSAGE_CLUSTERING_NOT_VALID, MESSAGE_INCORRECT_FIELD, MESSAGE_INVALID_FIELD, MESSAGE_NO_FIELD_NAME_OR_CLUSTERING, base64, check_list_of_dict, datetime, demisto, get_model_if_not_expired, main, preprocess_incidents_field, ) from freezegun import freeze_time PARAMETERS_DICT = { "fromDate": "", "toDate": "", "limit": "1000", "query": "", "minNumberofIncidentPerCluster": "2", "modelName": "model_name", "storeModel": "False", "minHomogeneityCluster": 0.6, "type": "Phishing", "maxRatioOfMissingValue": 0.5, "modelExpiration": 24, "forceRetrain": "True", "modelHidden": "False", "numberOfFeaturesPerField": 500, "analyzer": "char", } FETCHED_INCIDENT_NOT_EMPTY = [ { "id": "1", "created": "2021-01-30", "name": "name_1", "field_1": "powershell IP=1.1.1.1", "field_2": "powershell.exe", "entityname": "powershell", }, { "id": "2", "created": "2021-01-30", "name": "name_2", "field_1": "nmap port 1", "field_2": "nmap.exe", "entityname": "nmap", }, { "id": "3", "created": "2021-01-30", "name": "name_3", "field_1": "powershell IP=1.1.1.2", "field_2": "powershell", "entityname": "powershell", }, { "id": "4", "created": "2021-01-30", "name": "name_4", "field_1": "nmap port 2", "field_2": "nmap", "entityname": "nmap", }, { "id": "5", "created": "2021-01-30", "name": "name_3", "field_1": "powershell IP=1.1.1.3", "field_2": "powershell", "entityname": "powershell", }, { "id": "6", "created": "2021-01-30", "name": "name_4", "field_1": "nmap port 3", "field_2": "nmap", "entityname": "nmap", }, ] FETCHED_INCIDENT_NOT_EMPTY_MULTIPLE_NAME = [ { "id": "1", "created": "2021-01-30", "name": "name_1", "field_1": "powershell IP=1.1.1.1", "field_2": "powershell.exe", "entityname": ["powershell", "powershell", "nmap"], }, { "id": "2", "created": "2021-01-30", "name": "name_2", "field_1": "nmap port 1", "field_2": "nmap.exe", "entityname": ["powershell", "nmap", "nmap"], }, { "id": "3", "created": "2021-01-30", "name": "name_3", "field_1": "powershell IP=1.1.1.2", "field_2": "powershell", "entityname": ["powershell", "powershell", "nmap"], }, { "id": "4", "created": "2021-01-30", "name": "name_4", "field_1": "nmap port 2", "field_2": "nmap", "entityname": ["powershell", "nmap", "nmap"], }, { "id": "5", "created": "2021-01-30", "name": "name_3", "field_1": "powershell IP=1.1.1.3", "field_2": "powershell", "entityname": ["powershell", "powershell", "nmap"], }, { "id": "6", "created": "2021-01-30", "name": "name_4", "field_1": "nmap port 3", "field_2": "nmap", "entityname": ["powershell", "nmap", "nmap"], }, ] FETCHED_INCIDENT_NOT_EMPTY_WITH_NOT_ENOUGH_VALUES = [ { "id": "1", "created": "2021-01-30", "field_1": "powershell IP=1.1.1.1", "field_2": "", "entityname": "powershell", }, { "id": "2", "created": "2021-01-30", "field_1": "nmap port 1", "field_2": "", "entityname": "nmap", }, { "id": "3", "created": "2021-01-30", "field_1": "powershell IP=1.1.1.2", "field_2": "", "entityname": "powershell", }, { "id": "4", "created": "2021-01-30", "field_1": "nmap port 2", "field_2": "nmap", "entityname": "nmap", }, { "id": "5", "created": "2021-01-30", "field_1": "powershell IP=1.1.1.3", "field_2": "", "entityname": "powershell", }, { "id": "6", "created": "2021-01-30", "field_1": "nmap port 3", "field_2": "nmap", "entityname": "nmap", }, ] FETCHED_INCIDENT_NOT_EMPTY_SAME_CLUSTER_NAME = [ { "id": "1", "created": "2021-01-30", "name": "name_1", "field_1": "powershell IP=1.1.1.1", "field_2": "powershell.exe", "entityname": "powershell", }, { "id": "2", "created": "2021-01-30", "name": "name_2", "field_1": "nmap port 1", "field_2": "nmap.exe", "entityname": "nmap", }, { "id": "3", "created": "2021-01-30", "name": "name_3", "field_1": "powershell IP=1.1.1.2", "field_2": "powershell.exe", "entityname": "powershell", }, { "id": "4", "created": "2021-01-30", "name": "name_4", "field_1": "nmap port 2", "field_2": "nmap.exe", "entityname": "nmap", }, { "id": "5", "created": "2021-01-30", "name": "name_3", "field_1": "powershell IP=1.1.1.3", "field_2": "powershell.exe", "entityname": "powershell", }, { "id": "6", "created": "2021-01-20", "name": "name_4", "field_1": "nmap port 3", "field_2": "nmap.exe", "entityname": "nmap", }, ] FETCHED_INCIDENT_EMPTY = [] sub_dict_0 = { "data": [3], "dataType": "incident", "incidents_ids": ["1", "3", "5"], "name": "powershell", "query": "type:Phishing", } sub_dict_1 = { "data": [3], "dataType": "incident", "incidents_ids": ["2", "4", "6"], "name": "nmap", "query": "type:Phishing", } class PostProcessing: def __init__(self, date_training): self.date_training = date_training self.json = '{"data": [{"name": "name", "incidents_ids": []}]}' def executeCommand(command, args): import DBotTrainClustering match command: case "GetIncidentsByQuery": return [{"Contents": json.dumps(FETCHED_INCIDENT), "Type": "note"}] case "getMLModel": model = PostProcessing(datetime.now().strftime("%m/%d/%Y %H:%M:%S")) # Add test module class to allowlist so safe_pickle_loads can deserialize it DBotTrainClustering._ALLOWED_CLASSES.add(("DBotTrainClustering_test", "PostProcessing")) model_data = base64.b64encode(_pickle.dumps(model)).decode("utf-8") # guardrails-disable-line return [ { "Contents": {"modelData": model_data}, "Type": "note", } ] case _: return None def test_preprocess_incidents_field(): assert preprocess_incidents_field("incident.commandline") == "commandline" assert preprocess_incidents_field("commandline") == "commandline" def test_check_list_of_dict(): assert check_list_of_dict([{"test": "value_test"}, {"test1": "value_test1"}]) is True assert check_list_of_dict({"test": "value_test"}) is False # Test regular training def test_main_regular(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY args = PARAMETERS_DICT | { "fieldsForClustering": "field_1, field_2, wrong_field", "fieldForClusterName": "entityname", } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) _, output_clustering_json, msg = main() output_json = json.loads(output_clustering_json) cluster_0 = output_json["data"][0] cluster_1 = output_json["data"][1] assert MESSAGE_INCORRECT_FIELD % "wrong_field" in msg assert cluster_0["incidents_ids"] == ["1", "3", "5"] assert cluster_1["incidents_ids"] == ["2", "4", "6"] assert all(item in cluster_0.items() for item in sub_dict_0.items()) assert all(item in cluster_1.items() for item in sub_dict_1.items()) assert not all(item in cluster_0.items() for item in sub_dict_1.items()) assert not all(item in cluster_1.items() for item in sub_dict_0.items()) # Test if wrong cluster name def test_wrong_cluster_name(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY args = PARAMETERS_DICT | { "fieldsForClustering": "field_1, field_2", "fieldForClusterName": "wrong_cluster_name_field", } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, output_clustering_json, msg = main() assert MESSAGE_INCORRECT_FIELD % "wrong_cluster_name_field" in msg assert not output_clustering_json assert not model # Test if empty cluster name def test_empty_cluster_name(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY args = PARAMETERS_DICT | {"fieldsForClustering": "field_1, field_2", "fieldForClusterName": ""} mocker.patch.object(demisto, "args", return_value=args) sub_dict_0 = { "data": [3], "dataType": "incident", "incidents_ids": ["1", "3", "5"], "name": "Cluster 0", "query": "type:Phishing", } sub_dict_1 = { "data": [3], "dataType": "incident", "incidents_ids": ["2", "4", "6"], "name": "Cluster 1", "query": "type:Phishing", } mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, output_clustering_json, msg = main() output_json = json.loads(output_clustering_json) cluster_0 = output_json["data"][0] cluster_1 = output_json["data"][1] assert cluster_0["incidents_ids"] == ["1", "3", "5"] assert cluster_1["incidents_ids"] == ["2", "4", "6"] assert all(item in cluster_0.items() for item in sub_dict_0.items()) assert all(item in cluster_1.items() for item in sub_dict_1.items()) assert not all(item in cluster_0.items() for item in sub_dict_1.items()) assert not all(item in cluster_1.items() for item in sub_dict_0.items()) # Test if incorrect all incorrrect field name def test_all_incorrect_fields(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY args = PARAMETERS_DICT | { "fieldsForClustering": "field_1_wrong, field_2_wrong", "fieldForClusterName": "name", } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, output_clustering_json, msg = main() assert MESSAGE_INCORRECT_FIELD % " , ".join(["field_1_wrong", "field_2_wrong"]) in msg assert MESSAGE_NO_FIELD_NAME_OR_CLUSTERING in msg assert not output_clustering_json assert not model # Test if one field has no enough value def test_missing_too_many_values(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY_WITH_NOT_ENOUGH_VALUES args = PARAMETERS_DICT | {"fieldsForClustering": "field_1, field_2", "fieldForClusterName": "entityname"} mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, output_clustering_json, msg = main() assert MESSAGE_INVALID_FIELD % "field_2" in msg assert output_clustering_json assert model # Test for nested fields def test_main_incident_nested(mocker): """ Test if fetched incident truncated - Should return MESSAGE_WARNING_TRUNCATED in the message :param mocker: :return: """ global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY nested_field = "xdralerts.cmd" args = PARAMETERS_DICT | {"fieldsForClustering": nested_field, "fieldForClusterName": nested_field} mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "dt", return_value=["nested_val_1", "nested_val_2"]) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, output_clustering_json, msg = main() assert model is None assert output_clustering_json == {} assert MESSAGE_CLUSTERING_NOT_VALID in msg # Test to validate that if the model is still valid then it won't train again def test_model_exist_and_valid(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY args = PARAMETERS_DICT | { "fieldsForClustering": "field_1, field_2, wrong_field", "fieldForClusterName": "entityname", "forceRetrain": "False", } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) _, output_clustering_json, msg = main() assert not msg assert output_clustering_json == PostProcessing(None).json # Test to validate that if the model has expired then it will train again def test_model_exist_and_expired(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY time = "1e-20" args = PARAMETERS_DICT | { "fieldsForClustering": "field_1, field_2", "fieldForClusterName": "entityname", "forceRetrain": "False", "modelExpiration": time, } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) _, output_clustering_json, _ = main() output_json = json.loads(output_clustering_json) cluster_0 = output_json["data"][0] cluster_1 = output_json["data"][1] assert all(item in cluster_0.items() for item in sub_dict_0.items()) assert all(item in cluster_1.items() for item in sub_dict_1.items()) assert not all(item in cluster_0.items() for item in sub_dict_1.items()) assert not all(item in cluster_1.items() for item in sub_dict_0.items()) # Test if cluster name field has value of type list def test_main_name_cluster_is_list(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY_MULTIPLE_NAME args = PARAMETERS_DICT | { "fieldsForClustering": "field_1, field_2, wrong_field", "fieldForClusterName": "entityname", } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, output_clustering_json, msg = main() output_json = json.loads(output_clustering_json) cluster_0 = output_json["data"][0] cluster_1 = output_json["data"][1] assert MESSAGE_INCORRECT_FIELD % "wrong_field" in msg assert all(item in cluster_0.items() for item in sub_dict_0.items()) assert all(item in cluster_1.items() for item in sub_dict_1.items()) assert not all(item in cluster_0.items() for item in sub_dict_1.items()) assert not all(item in cluster_1.items() for item in sub_dict_0.items()) # Test same cluster name should created prefixes def test_same_cluster_name(mocker): global FETCHED_INCIDENT FETCHED_INCIDENT = FETCHED_INCIDENT_NOT_EMPTY_SAME_CLUSTER_NAME args = PARAMETERS_DICT | { "fieldsForClustering": "field_1, field_2, wrong_field", "fieldForClusterName": "entityname", } mocker.patch.object(demisto, "args", return_value=args) mocker.patch.object(demisto, "executeCommand", side_effect=executeCommand) model, *_ = main() cluster_names = [x["clusterName"] for x in model.selected_clusters.values()] assert cluster_names == ["", "powershell", "nmap"] @pytest.mark.parametrize( "force_retrain, model_expiration, model, expected_result_obj", [ (True, 48, PostProcessing(datetime(2023, 1, 1)), type(None)), (False, 48, None, type(None)), (False, 48, PostProcessing(datetime(2023, 1, 1)), type(None)), (False, 48, PostProcessing(datetime(2023, 2, 1)), PostProcessing), ], ) @freeze_time("2023-02-02") def test_get_model_if_not_expired(mocker, force_retrain, model_expiration, model, expected_result_obj): # Mock get_model function mocker.patch("DBotTrainClustering.get_model", return_value=model) result = get_model_if_not_expired(force_retrain, model_expiration, "name") assert isinstance(result, expected_result_obj)