from typing import Any import pytest from flexmock import flexmock from sqlalchemy.exc import SQLAlchemyError from assets.connectors import mysql from assets.models import hive_ai_image_task from tests.unit.models import au_operations @pytest.fixture def db_fixture( asset_final_data: list[dict[str, Any]], valid_hive_ai_image_task_data: dict[str, Any], ) -> None: """Set up the asset_final and hive_ai_image_task tables.""" au_operations.truncate_tables() au_operations.seed_asset_final_table(asset_final_data) au_operations.seed_hive_ai_image_task_table(valid_hive_ai_image_task_data) def test_create_hive_ai_image_task( db_fixture: None, valid_hive_ai_image_task_data: dict[str, Any], ) -> None: """Test creating hive_ai_image_task successfully.""" result = hive_ai_image_task.create_hive_ai_image_task( **valid_hive_ai_image_task_data ) assert result["id"] is not None assert result["asset_final_id"] == valid_hive_ai_image_task_data["asset_final_id"] assert result["task_id"] == valid_hive_ai_image_task_data["task_id"] assert result["class_name"] == valid_hive_ai_image_task_data["class_name"] assert result["score_value"] == valid_hive_ai_image_task_data["score_value"] def test_create_hive_ai_image_task_error( valid_hive_ai_image_task_data: dict[str, Any], ) -> None: """Test SQL error handling.""" error_message = "Query error" ( flexmock(mysql) .should_receive("au_db_session") .and_raise(SQLAlchemyError(error_message)) ) with pytest.raises(SQLAlchemyError) as exc: hive_ai_image_task.create_hive_ai_image_task(**valid_hive_ai_image_task_data) assert error_message in str(exc.value)