"""Test Logic Layer.""" from unittest.mock import Mock from unittest.mock import patch from unittest import TestCase import pytest from dbdeploy.base import config from dbdeploy.base import constants from dbdeploy.dtos import ExecutionJob from dbdeploy.logic import monitor_dlq from dbdeploy.logic import check_kafka_topic_exists from dbdeploy.logic import create_kafka_topics from dbdeploy.logic import delete_kafka_topics from dbdeploy.logic import prepare_job_list from dbdeploy.logic import parse_db_pr_file from dbdeploy.logic import push_to_kafka from dbdeploy.logic import start_fargate_task from dbdeploy.logic import start_kafka_connector from dbdeploy.logic import stop_fargate_task from dbdeploy.logic import stop_kafka_connector from dbdeploy.logic import track_kafka_consumer_offsets class TestXMLParsing: def test_parse_db_pr_file(self, mock_xml): """Test XML file parsed.""" config.FLOW_NAME = 'snowflake-neo4j' test_result = parse_db_pr_file(xml_contents=mock_xml) test_query = test_result[0].cypher_query if test_result else None mock_query = 'MERGE (t:TestNode {test: event.test})' assert isinstance(test_result, list) assert len(test_result) == 1 assert test_query == mock_query assert test_result[0].neo4j_server == config.Neo4jServerNames.MUSIC_GRAPH assert test_result[0].sql_query == 'SELECT test_col FROM test_table' assert test_result[0].cypher_query == 'MERGE (t:TestNode {test: event.test})' # noqa: E501 assert not test_result[0].topic_name assert not test_result[0].kafka_cluster_name def test_parse_db_pr_file_several_jobs(self, mock_multi_changeset_xml): """Test XML file with two changesets parsed.""" config.FLOW_NAME = 'snowflake-neo4j' test_result = parse_db_pr_file(xml_contents=mock_multi_changeset_xml) assert isinstance(test_result, list) assert len(test_result) == 2 assert test_result[0].neo4j_server == config.Neo4jServerNames.MUSIC_GRAPH assert test_result[0].sql_query == 'SELECT test_col FROM test_table1' assert test_result[0].cypher_query == 'MERGE (t:TestNode1 {test: event.test})' # noqa: E501 assert not test_result[0].topic_name assert not test_result[0].kafka_cluster_name assert test_result[0].neo4j_server == config.Neo4jServerNames.MUSIC_GRAPH assert test_result[1].sql_query == 'SELECT test_col FROM test_table2' assert test_result[1].cypher_query == 'MERGE (t:TestNode2 {test: event.test})' # noqa: E501 assert not test_result[0].topic_name assert not test_result[0].kafka_cluster_name def test_parse_db_pr_file_default_snowflake_account(self, mock_xml): """Test default snowflake-account param""" config.FLOW_NAME = 'snowflake-neo4j' test_result = parse_db_pr_file(xml_contents=mock_xml) assert test_result[0].snowflake_account == config.SnowflakeAccount.ORCHARD def test_parse_db_pr_file_default_neo4j_server(self, mock_default_neo4j_server_xml): """Test default neo4j-server param.""" config.FLOW_NAME = 'snowflake-neo4j' test_result = parse_db_pr_file(xml_contents=mock_default_neo4j_server_xml) assert test_result[0].neo4j_server == config.Neo4jServerNames.MUSIC_GRAPH def test_parse_db_pr_file_fails(self, mock_bad_xml): """Test bad XML file isn't parsed.""" config.FLOW_NAME = 'snowflake-neo4j' with pytest.raises(ValueError): parse_db_pr_file(xml_contents=mock_bad_xml) def test_parse_db_pr_file_unknown_neo4j_server(self, mock_unknown_neo4j_server_xml): """Test changeset with unknown Neo4j server.""" config.FLOW_NAME = 'snowflake-neo4j' with pytest.raises(ValueError): parse_db_pr_file(xml_contents=mock_unknown_neo4j_server_xml) def test_parse_db_pr_file_skip_cypher(self, mock_skip_cypher): """Test XML file parsed.""" config.FLOW_NAME = 'snowflake-kafka' test_result = parse_db_pr_file(xml_contents=mock_skip_cypher) test_query = test_result[0].cypher_query if test_result else None assert isinstance(test_result, list) assert len(test_result) == 1 assert test_query is None assert test_result[0].topic_name == 'testTopic.name' assert test_result[0].kafka_cluster_name == 'cdc-destination' def test_parse_db_pr_file_skip_cypher_kafka_cluster_not_set( self, mock_skip_cypher_kafka_cluster_not_set): """Test XML file parsed.""" config.FLOW_NAME = 'snowflake-kafka' test_result = parse_db_pr_file(xml_contents=mock_skip_cypher_kafka_cluster_not_set) # noqa: E501 test_query = test_result[0].cypher_query if test_result else None assert isinstance(test_result, list) assert len(test_result) == 1 assert test_query is None assert test_result[0].topic_name == 'testTopic.name' assert test_result[0].kafka_cluster_name == 'managed-kafka-cdc-destination' # noqa: E501 def test_parse_db_pr_file_snowflake_to_mysql( self, mock_skip_snowflake_to_mysql): """Test XML file parsed.""" config.FLOW_NAME = 'snowflake-mysql' test_result = parse_db_pr_file(xml_contents=mock_skip_snowflake_to_mysql) # noqa: E501 assert isinstance(test_result, list) assert len(test_result) == 1 assert test_result[0].table_schema.dict()['fields'] == [ {'field': 'test_col1', 'type': 'string', 'optional': True}, {'field': 'test_col2', 'type': 'string', 'optional': True}] assert test_result[0].topic_name is None assert test_result[0].sql_query == 'SELECT test_col FROM test_table' assert test_result[0].cypher_query is None def test_parse_db_pr_file_skip_cypher_missing_topic_fails( self, mock_skip_cypher_missing_topic_name): """Test bad XML file isn't parsed.""" config.FLOW_NAME = 'snowflake-kafka' with pytest.raises(ValueError): parse_db_pr_file(xml_contents=mock_skip_cypher_missing_topic_name) class TestNameGeneration: def test_generate_names(self, mock_joblist): """Test correct topics and connector names generated.""" test_uuid = 'f007ba11-1111-0000-0000-000000000000' expected_topic_name = f'{config.KAFKA_TOPIC_PREFIX}{test_uuid}' expected_dlq_name = f'{config.KAFKA_DLQ_TOPIC_PREFIX}{test_uuid}' expected_connector_name = constants.CONNECTOR_NAME_TEMPLATE.format(job_id=test_uuid.replace('-', '_')) # noqa: E501 test_result = prepare_job_list(mock_joblist) test_topic_names = {job.topic.name for job in test_result} test_dlq_names = {job.dlq_topic.name for job in test_result} test_connector_names = {job.connector_name for job in test_result} assert expected_topic_name in test_topic_names assert expected_dlq_name in test_dlq_names assert expected_connector_name in test_connector_names def test_generate_names_topic_name_set(self, mock_job_nocypher): """Test correct topics and connector names generated.""" expected_topic_name = 'test.TopicName' expected_dlq_name = None expected_connector_name = None test_result = prepare_job_list([mock_job_nocypher]) test_topic_names = {job.topic.name for job in test_result} test_dlq_names = {job.dlq_topic.name for job in test_result} test_connector_names = {job.connector_name for job in test_result} assert expected_topic_name in test_topic_names assert expected_dlq_name in test_dlq_names assert expected_connector_name in test_connector_names class TestFargateRelatedTasks(TestCase): """Test fargate task starts and stops.""" def setUp(self): self.mock_client = patch( 'dbdeploy.logic.dbdeploy.get_ecs_client').start() self.mock_get_arn = patch( 'dbdeploy.logic.dbdeploy.get_ecs_cluster_arn').start() self.mock_change = patch( 'dbdeploy.logic.dbdeploy.change_fargate_task_count').start() self.mock_check = patch( 'dbdeploy.logic.dbdeploy.check_running_task_count').start() def test_start_fargate_task(self): """Test fargate task started.""" self.mock_change.return_value = True self.mock_check.return_value = True self.mock_get_arn.return_value = 'test-arn-value' test_result = start_fargate_task(desired_count=1) assert test_result is True def test_start_fargate_task_get_arn_fails(self): """Test False returned if cannot get cluster ARN.""" self.mock_get_arn.return_value = None with pytest.raises(RuntimeError): start_fargate_task(desired_count=1) def test_start_fargate_task_check_count_fails(self): """Test False returned check task count fails.""" self.mock_get_arn.return_value = 'test-arn-value' self.mock_check.return_value = False with pytest.raises(RuntimeError): start_fargate_task(desired_count=1) def test_stop_fargate_task(self): """Test fargate task stopped.""" self.mock_change.return_value = True self.mock_check.return_value = True self.mock_get_arn.return_value = 'test-arn-value' test_result = stop_fargate_task() assert test_result is True def test_stop_fargate_task_get_arn_fails(self): """Test False returned if cannot get cluster ARN.""" self.mock_get_arn.return_value = None test_result = stop_fargate_task() assert test_result is False def test_stop_fargate_task_check_count_fails(self): """Test False returned check task count fails.""" self.mock_get_arn.return_value = 'test-arn-value' self.mock_check.return_value = False test_result = stop_fargate_task() assert test_result is False class TestKafkaRelatedTasks: """Test kafka related tasks.""" @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_create_kafka_topics(self, mock_executor, mock_job): """Test kafka topics created.""" mock_executor().admin_client.return_value = (Mock(), None) test_result = create_kafka_topics(job=mock_job) assert test_result.topics_created is True @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_create_kafka_topics_fails(self, mock_executor, mock_job): """Test kafka topics not created.""" config.RETRY_BACKOFF = config.RETRY_BACKOFF / 10000 config.KAFKA_CONNECTION_RETRIES = 1 mock_executor().admin_client.return_value = (None, 'error') with pytest.raises(TimeoutError): create_kafka_topics(job=mock_job) @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_delete_kafka_topics(self, mock_executor, mock_job): """Test kafka topics created.""" response = Mock() response.topic_error_codes = [('topic_name', 0)] mock_executor().admin_client.return_value = (Mock(), None) mock_executor().admin_client()[0].delete_topics.return_value = response test_result = delete_kafka_topics(job=mock_job) assert test_result is True @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_list_kafka_topics(self, mock_executor, mock_job_nocypher): """Test kafka topics created.""" response = ['test.TopicName'] mock_executor().admin_client.return_value = (Mock(), None) mock_executor().admin_client()[0].list_topics.return_value = response test_result = check_kafka_topic_exists(job=mock_job_nocypher) assert test_result is mock_job_nocypher @patch('dbdeploy.logic.dbdeploy.sfcursor') @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_push_to_kafka( self, mock_executor, mock_sfcursor, mock_job): """Test data pushed to the topic.""" mock_sfcursor().__enter__().fetchmany.return_value = [] mock_executor().producer.return_value = (Mock(), None) test_result = push_to_kafka(job=mock_job) assert isinstance(test_result, ExecutionJob) @patch('dbdeploy.logic.dbdeploy.sfcursor') @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_push_to_kafka_no_precondition( self, mock_executor, mock_sfcursor, mock_job_nosql): """Test data pushed to the topic.""" mock_sfcursor().__enter__().fetchmany.return_value = [] mock_executor().producer.return_value = (Mock(), None) test_result = push_to_kafka(job=mock_job_nosql) assert isinstance(test_result, ExecutionJob) @patch('dbdeploy.logic.dbdeploy.compare_offsets') @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_track_kafka_consumer_offsets( self, mock_executor, _, mock_job): """Test offsets tracked in Kafka.""" mock_executor().admin_client.return_value = (Mock(), None) test_result = track_kafka_consumer_offsets(mock_job) assert isinstance(test_result, ExecutionJob) @patch('dbdeploy.logic.dbdeploy.KafkaExecutor') def test_monitor_dlq( self, mock_executor, mock_job): """Test function returns job instance.""" mock_executor().consumer.return_value = (Mock(), None) mock_executor().consumer( )[0].partitions_for_topic.return_value = ['p1', 'p2'] mock_job.is_done = True test_result = monitor_dlq.__wrapped__(mock_job) assert isinstance(test_result, ExecutionJob) class TestKafkaConnectRelatedTasks: """Test kafka-connect related tasks.""" @patch('dbdeploy.logic.dbdeploy.create_connector') def test_start_kafka_connector(self, mock_create, mock_job): """Test connector started.""" mock_create.return_value = (True, None) test_result = start_kafka_connector(mock_job) assert test_result.connector_started is True @patch('dbdeploy.logic.dbdeploy.create_connector') def test_start_kafka_connector_fails(self, mock_create, mock_job): """Test connector not started.""" mock_create.return_value = (False, 'start failed') with pytest.raises(Exception): start_kafka_connector(mock_job) @patch('dbdeploy.logic.dbdeploy.delete_connector') def test_stop_kafka_connector(self, mock_delete, mock_job): """Test connector deleted.""" mock_delete.return_value = (True, None) test_result = stop_kafka_connector(mock_job) assert test_result is True @patch('dbdeploy.logic.dbdeploy.delete_connector') def test_stop_kafka_connector_fails(self, mock_delete, mock_job): """Test connector delete fails.""" mock_delete.return_value = (False, 'start failed') with pytest.raises(Exception): stop_kafka_connector(mock_job) class TestInsertModeXMLParsing: """Test class for insert_mode XML parsing functionality.""" def test_parse_snowflake_to_mysql_with_upsert_mode(self, mock_snowflake_to_mysql_with_upsert): """Test XML parsing with insert-mode="upsert" specified.""" config.FLOW_NAME = 'snowflake-mysql' test_result = parse_db_pr_file(xml_contents=mock_snowflake_to_mysql_with_upsert) assert isinstance(test_result, list) assert len(test_result) == 1 assert test_result[0].insert_mode == 'upsert' assert test_result[0].table_schema.table_name == 'test_table' assert test_result[0].table_schema.primary_key == 'test_col1' def test_parse_snowflake_to_mysql_with_insert_mode(self, mock_snowflake_to_mysql_with_insert): """Test XML parsing with insert-mode="insert" specified.""" config.FLOW_NAME = 'snowflake-mysql' test_result = parse_db_pr_file(xml_contents=mock_snowflake_to_mysql_with_insert) assert isinstance(test_result, list) assert len(test_result) == 1 assert test_result[0].insert_mode == 'insert' assert test_result[0].table_schema.table_name == 'test_table' def test_parse_snowflake_to_mysql_with_delete_mode(self, mock_snowflake_to_mysql_with_delete): """Test XML parsing with insert-mode="delete" specified.""" config.FLOW_NAME = 'snowflake-mysql' test_result = parse_db_pr_file(xml_contents=mock_snowflake_to_mysql_with_delete) assert isinstance(test_result, list) assert len(test_result) == 1 assert test_result[0].insert_mode == 'delete' assert test_result[0].table_schema.table_name == 'test_table' def test_parse_snowflake_to_mysql_default_insert_mode(self, mock_snowflake_to_mysql_default_insert_mode): """Test XML parsing with no insert-mode specified (should default to 'update').""" config.FLOW_NAME = 'snowflake-mysql' test_result = parse_db_pr_file(xml_contents=mock_snowflake_to_mysql_default_insert_mode) assert isinstance(test_result, list) assert len(test_result) == 1 assert test_result[0].insert_mode == 'update' # Should default to 'update' assert test_result[0].table_schema.table_name == 'test_table' def test_parse_snowflake_to_mysql_invalid_insert_mode_fails(self, mock_snowflake_to_mysql_with_invalid_insert_mode): """Test XML parsing fails with invalid insert-mode.""" config.FLOW_NAME = 'snowflake-mysql' with pytest.raises( ValueError, match="The value 'invalid_mode' is not an element of the set {'insert', 'upsert', 'update', 'delete'}" ): parse_db_pr_file(xml_contents=mock_snowflake_to_mysql_with_invalid_insert_mode) def test_existing_snowflake_to_mysql_still_works(self, mock_skip_snowflake_to_mysql): """Test that existing XML without insert-mode still works (backward compatibility).""" config.FLOW_NAME = 'snowflake-mysql' test_result = parse_db_pr_file(xml_contents=mock_skip_snowflake_to_mysql) assert isinstance(test_result, list) assert len(test_result) == 1 assert test_result[0].insert_mode == 'update' # Should default to 'update' assert test_result[0].table_schema.dict()['fields'] == [ {'field': 'test_col1', 'type': 'string', 'optional': True}, {'field': 'test_col2', 'type': 'string', 'optional': True}]