"""Tests for handlers.""" import json import logging from typing import Any from unittest.mock import MagicMock, call, patch import pytest from flask import request from flask.ctx import AppContext from flask.testing import FlaskClient from flexmock import flexmock from owsresponse import response from pytest_mock import MockerFixture from requests.exceptions import HTTPError from werkzeug.exceptions import BadRequest, Conflict, HTTPException, NotFound from assets import config, handlers, rate_limiting from assets.constants import authorization, error, field_const from assets.exceptions import ( ArtistNotFound, AssetDeleteFailure, AssetEncodingInProgress, AssetFinalExists, AssetFinalNotFound, AssetStatusNotFound, AssetUploadError, AssetUploadNotFound, AssetUploadTypeNotFound, BitsPerSampleNotFound, DesiredBitsPerSampleNotFound, DuplicateAssetUploadFilename, ErrorDiscardingCorrections, HiveSegmentExists, ImageAssetNotFound, InvalidEntityType, InvalidProductStatus, MissingContext, NoAssetsFound, ProductNotFound, ReleaseCorrectionNotFound, S3MultipartUploadFailed, TrackDoesNotBelongToProduct, TrackNotFound, VendorNotFound, WrongAssetStatus, ) from assets.logic import ( ownership, track, validate_bit_depth_by_upc, ) from assets.rate_limiting import Allowed, Error OA_HEADERS = {"Orchard-User-Id": "oa:12345", "X-Forwarded-For": "0.0.0.0"} ALW_HEADERS = { "Orchard-User-Id": "alw:12345", "Grass-Account-Id": "234234", "Grass-Account-Type": "vendor", "X-Forwarded-For": "0.0.0.0", } TRACK_COPYING_INPUT_DATA = { "from_tuid": "12345_67891", "to_tuid": "21435_65879", } @pytest.mark.parametrize( "headers", [ {**OA_HEADERS, "Orchard-User-Id": "invalid_id"}, {**ALW_HEADERS, "Orchard-User-Id": "invalid_id"}, ], ) def test_get_product_link_failure( fixture_client: FlaskClient, headers: dict[str, str] ) -> None: """Test route fails when there is invalid Orchard-User-Id in header.""" product_id = "1" request_url = f"/asset/product/{product_id}" result = fixture_client.get(request_url, headers=headers) assert result.status_code == 400 response_message = json.loads(result.data.decode()) assert response_message["code"] == error.ERROR_CODE_HEADER_VALIDATION def test_check_for_desired_bit_depth(fixture_client: FlaskClient) -> None: """Test route to check desired bit-depth for provided UPC.""" product_id = 2011621 desired_bit_depth = 64 ( flexmock(validate_bit_depth_by_upc) .should_receive("check_for_desired_bit_depth") .and_return(None) ) request_url = ( "/validators/product/{product_id}" "/desired-bit-depth/{bit_depth}?" "physical_location_id=1&asset_type_id=1".format( product_id=product_id, bit_depth=desired_bit_depth ) ) result = fixture_client.get(request_url) result_message = result.data.decode("utf8") assert result.status_code == 200 assert result.headers.get("Correlation-Id") assert result_message == json.dumps({"status": error.SUCCESS_CODE}) def test_check_for_desired_bit_depth_with_invalid_upc( fixture_client: FlaskClient, ) -> None: """Test route to check for deisered bit-depth with invalid UPC.""" product_id = 000 desired_bit_depth = 64 ( flexmock(validate_bit_depth_by_upc) .should_receive("check_for_desired_bit_depth") .and_raise(DesiredBitsPerSampleNotFound(error.DESIRED_BIT_DEPTH_NOT_FOUND)) ) request_url = ( "/validators/product/{product_id}" "/desired-bit-depth/{bit_depth}?" "physical_location_id=1&asset_type_id=1".format( product_id=product_id, bit_depth=desired_bit_depth ) ) result = fixture_client.get(request_url) result_message = result.data.decode("utf8") assert result.status_code == 400 assert result.headers.get("Correlation-Id") assert error.DESIRED_BIT_DEPTH_NOT_FOUND in result_message @pytest.mark.parametrize( "headers, data, input_data, ownership_mock_called", [ ( {**OA_HEADERS, "Content-Type": "application/json"}, {}, {**TRACK_COPYING_INPUT_DATA, "full_user_id": OA_HEADERS["Orchard-User-Id"]}, False, ), ( {**ALW_HEADERS, "Content-Type": "application/json"}, {}, { **TRACK_COPYING_INPUT_DATA, "full_user_id": ALW_HEADERS["Orchard-User-Id"], }, True, ), ( {"Content-Type": "application/json"}, {"Orchard-User-Id": "alw:321"}, {**TRACK_COPYING_INPUT_DATA, "full_user_id": "alw:321"}, False, ), ], ) def test_copy_track_success( mocker: MockerFixture, fixture_client: FlaskClient, headers: dict[str, str], data: dict[str, Any], input_data: dict[str, str], ownership_mock_called: bool, ) -> None: """Test route that copy track asset.""" request_url_pattern = "/track/{0}/copy/{1}" request_url = request_url_pattern.format( input_data["from_tuid"], input_data["to_tuid"] ) check_track_ownership_mock = mocker.patch.object( ownership, "check_track_ownership", return_value=response.Response() ) ( flexmock(track) .should_receive("copy_track_asset") .with_args(**input_data) .and_return(None) ) result = fixture_client.post(request_url, data=json.dumps(data), headers=headers) assert result.status_code == 200 assert result.headers.get("Correlation-Id") assert check_track_ownership_mock.called == ownership_mock_called @pytest.mark.parametrize( "headers, status, error_code", [ ( {**OA_HEADERS, "Orchard-User-Id": "invalid_id"}, 400, error.ERROR_CODE_HEADER_VALIDATION, ), ( {**ALW_HEADERS, "Orchard-User-Id": "invalid_id"}, 400, error.ERROR_CODE_HEADER_VALIDATION, ), ({"Content-Type": "application/json"}, 401, error.ERROR_CODE_AUTHORIZATION), ], ) def test_copy_track_failure( fixture_client: FlaskClient, headers: dict[str, str], status: int, error_code: int, ) -> None: """Test route that copy track asset failure.""" request_url_pattern = "/track/{0}/copy/{1}" request_url = request_url_pattern.format( TRACK_COPYING_INPUT_DATA["from_tuid"], TRACK_COPYING_INPUT_DATA["to_tuid"] ) data = json.dumps({}) result = fixture_client.post(request_url, data=data, headers=headers) assert result.status_code == status response_message = json.loads(result.data.decode()) assert response_message["code"] == error_code def test_copy_track_to_itself(fixture_client: FlaskClient) -> None: """Test for failure copy_track_asset call.""" headers = ALW_HEADERS request_url_pattern = "/track/{0}/copy/{1}" request_url = request_url_pattern.format(123, 123) result = fixture_client.post(request_url, data=json.dumps({}), headers=headers) response_message = json.loads(result.data.decode()) assert result.status_code == 400 assert response_message["code"] == error.ERROR_CODE_SELF_COPY assert response_message["message"] == error.ERROR_ASSET_SELF_COPY @pytest.mark.parametrize("failure_mock_call_order", [0, 1]) def test_copy_track_with_invalid_ownership( failure_mock_call_order: int, mocker: MockerFixture, fixture_client: FlaskClient, ) -> None: """Test route that copy track asset failure.""" request_url_pattern = "/track/{0}/copy/{1}" request_url = request_url_pattern.format( TRACK_COPYING_INPUT_DATA["from_tuid"], TRACK_COPYING_INPUT_DATA["to_tuid"] ) headers = {**ALW_HEADERS, "Content-Type": "application/json"} ownership_responses = [response.Response(), response.Response()] ownership_responses[failure_mock_call_order] = response.create_error_response( code="not_owned", message="not_owned", status=403 ) mocker.patch.object( ownership, "check_track_ownership", side_effect=ownership_responses, ) result = fixture_client.post(request_url, headers=headers) assert result.status_code == 403 assert result.headers.get("Correlation-Id") @pytest.mark.parametrize( "headers, status, error_code", [ ( {**OA_HEADERS, "Orchard-User-Id": "invalid_id"}, 400, error.ERROR_CODE_HEADER_VALIDATION, ), ( {**ALW_HEADERS, "Orchard-User-Id": "invalid_id"}, 400, error.ERROR_CODE_HEADER_VALIDATION, ), ({"Content-Type": "application/json"}, 401, error.ERROR_CODE_AUTHORIZATION), ], ) def test_copy_product_failure( fixture_client: FlaskClient, headers: dict[str, str], status: int, error_code: int ) -> None: """Test copy product assets failure.""" request_url_pattern = "/track/{0}/copy/{1}" request_url = request_url_pattern.format(10, 20) data = json.dumps({}) result = fixture_client.post(request_url, data=data, headers=headers) assert result.status_code == status response_message = json.loads(result.data.decode()) assert response_message["code"] == error_code def test_hello(fixture_client: FlaskClient) -> None: """Test hello handler success.""" result = fixture_client.get(config.HEALTH_CHECK) assert result.status_code == 200 @patch("assets.handlers.jsonify", side_effect=Exception()) def test_hello_500_error(_mock_jsonify: MagicMock, fixture_client: FlaskClient) -> None: """Test an uncaught Exception results in a 500 status code.""" result = fixture_client.get("/hello/") result_json = json.loads(result.data.decode()) assert result.status_code == 500 assert ( result_json["message"] == "The server encountered an internal error and was unable to complete your request." ) @pytest.mark.parametrize( "raised_exception, expected_error_message", [ ( AssetUploadError("asset upload error"), "asset upload error", ), ( AssetDeleteFailure({"success": None, "failure": True}), {"success": None, "failure": True}, ), ( ErrorDiscardingCorrections("This is a custom exception"), "This is a custom exception", ), ], ) @patch("assets.handlers.g", spec=["log"]) def test_custom_exception_handler( mock_g: MagicMock, raised_exception: AssetDeleteFailure | AssetUploadError | ErrorDiscardingCorrections, expected_error_message: str, ) -> None: """Test an uncaught Exception results in a 500 status code.""" server_response = handlers.custom_exception_handler(raised_exception) mock_g.log.warning.assert_called_with(raised_exception) response_json = json.loads(server_response.data.decode()) assert server_response.status_code == 500 assert response_json["message"] == expected_error_message @patch("assets.handlers.g", spec=["log"]) def test_exception_handler(mock_g: MagicMock) -> None: """Verify exception_Handler returns 500 status code and json payload.""" mock_error = MagicMock() server_response = handlers.exception_handler(mock_error) mock_g.log.exception.assert_called_with(mock_error) # assert status code is 500 assert server_response.status_code == 500 # assert error message matches server_response_json = json.loads(server_response.data.decode()) assert ( server_response_json["message"] == "The server encountered an internal error and was unable to complete your request." ) assert server_response_json["code"] == "internal_error" @pytest.mark.parametrize( "mock_exception, status_code, message, error_code", [ (NotFound("not found error"), 404, "not found error", "not_found_error"), (BadRequest("bad request"), 400, "bad request", "bad_request"), (Conflict("conflict"), 409, "conflict", "conflict"), ( DuplicateAssetUploadFilename("duplicate asset upload filename"), 400, "duplicate asset upload filename", "bad_request", ), (NoAssetsFound("no assets found"), 404, "no assets found", "not_found_error"), (MissingContext("missing context"), 400, "missing context", "bad_request"), (AssetFinalExists("asset final exists"), 409, "asset final exists", "conflict"), ( AssetFinalNotFound("asset final not found"), 404, "asset final not found", "not_found_error", ), ( BitsPerSampleNotFound("bits per sample not found"), 404, "bits per sample not found", "not_found_error", ), ( DesiredBitsPerSampleNotFound("desired bits per sample not found"), 400, "desired bits per sample not found", "bad_request", ), ( AssetStatusNotFound("asset status not found"), 404, "asset status not found", "not_found_error", ), ( HiveSegmentExists("hive segment exists"), 409, "hive segment exists", "conflict", ), ( VendorNotFound("vendor not found"), 404, "vendor not found", "not_found_error", ), ( ArtistNotFound("artist not found"), 404, "artist not found", "not_found_error", ), ( ReleaseCorrectionNotFound("release correction not found"), 404, "release correction not found", "not_found_error", ), ( ProductNotFound("product not found"), 404, "product not found", "not_found_error", ), ( ImageAssetNotFound("image asset not found"), 404, "image asset not found", "not_found_error", ), ( InvalidEntityType("invalid entity type"), 400, "invalid entity type", "bad_request", ), ( AssetUploadNotFound("asset upload not found"), 404, "asset upload not found", "not_found_error", ), ( AssetUploadTypeNotFound("asset upload type not found"), 404, "asset upload type not found", "not_found_error", ), (TrackNotFound("track not found"), 404, "track not found", "not_found_error"), ( InvalidProductStatus("invalid product status"), 400, "invalid product status", "bad_request", ), ( WrongAssetStatus("wrong asset status"), 404, "wrong asset status", "not_found_error", ), ( TrackDoesNotBelongToProduct("track does not belong to product"), 400, "track does not belong to product", "bad_request", ), ( S3MultipartUploadFailed("s3 multipart upload failed"), 400, "s3 multipart upload failed", "bad_request", ), ], ) def test_http_exception_handler( mock_exception: HTTPException, status_code: int, message: str, error_code: str ) -> None: """Verify http_exception_handler returns correct status code and message.""" server_response = handlers.http_exception_handler(mock_exception) server_response_json = json.loads(server_response.data.decode()) assert server_response.status_code == status_code assert server_response_json["message"] == message assert server_response_json["code"] == error_code @pytest.mark.parametrize( "mock_exception, message", [ ( HTTPError("Not Found", response=MagicMock(status_code=404)), "Not Found", ), ( HTTPError("Bad Request", response=MagicMock(status_code=400)), "Bad Request", ), ], ) def test_http_error_handler(mock_exception: HTTPError, message: str) -> None: """Verify http_error_handler returns correct status code and message.""" server_response = handlers.http_error_handler(mock_exception) server_response_json = json.loads(server_response.data.decode()) assert server_response.status_code == 502 assert server_response_json["message"] == message assert server_response_json["code"] == "bad_gateway" def test_asset_encoding_in_progress_handler() -> None: """Verify asset_encoding_Handler returns 204 status code and json payload.""" message = "asset encoding still in progress" server_response = handlers.asset_encoding_in_progress_handler( AssetEncodingInProgress(message) ) assert server_response.status_code == 204 assert server_response.data.decode() == message @pytest.mark.parametrize( ( "rate_limit_result", "expected_logged_errors", "headers", ), [ ( rate_limiting.Error("some error"), {"Rate limiter error for upload:oa:179 - some error - failing open"}, {**OA_HEADERS, "Orchard-User-Id": "oa:179"}, ), ( rate_limiting.Allowed(), {}, OA_HEADERS, ), ( rate_limiting.Allowed(), {}, ALW_HEADERS, ), ], ) @patch("assets.handlers.pdp_auth.shadow_authorization") @patch( "assets.logic.generate_upload_data.create_asset_upload", return_value="create_asset_upload_logic_result", ) @patch( "owsrequest.flask_request.verify_grass_ownership", return_value=response.Response() ) @patch("assets.rate_limiting.get_rate_limiter") def test_create_asset_upload_success( mocked_get_rate_limiter: MagicMock, mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, mocked_shadow_authorization: MagicMock, fixture_client: FlaskClient, headers: dict[str, str], rate_limit_result: Allowed | Error, expected_logged_errors: list[str], caplog: pytest.LogCaptureFixture, ) -> None: """Test create_asset_upload success with oa or alw headers.""" caplog.set_level(logging.ERROR) mocked_get_rate_limiter.return_value.hit.return_value = rate_limit_result request_body = { "product_id": 123, "track_unique_id": 0, "original_filename": "test.tif", "asset_upload_type": "static_artwork", } result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=headers ) assert mocked_get_rate_limiter.return_value.hit.mock_calls == [ call(rate_limiting.RateLimitedResource.UPLOAD, headers["Orchard-User-Id"]) ] for expected_logged_error in expected_logged_errors: assert expected_logged_error in caplog.text mocked_verify_grass_ownership.assert_called_once_with( request, ownership.check_ownership, 123 ) mocked_create_asset_upload.assert_called_with( **request_body, user_id=headers["Orchard-User-Id"], ) assert result.status_code == 200 assert result.text == json.dumps({"filename": "create_asset_upload_logic_result"}) @pytest.mark.parametrize( ("asset_upload_type_value", "expected_missing_log"), [ ("static_artwork", False), (None, True), ("__omit__", True), ], ids=["valid_value", "explicit_null", "absent"], ) @patch("assets.handlers.pdp_auth.shadow_authorization") @patch( "assets.logic.generate_upload_data.create_asset_upload", return_value="create_asset_upload_logic_result", ) @patch( "owsrequest.flask_request.verify_grass_ownership", return_value=response.Response() ) @patch("assets.rate_limiting.get_rate_limiter") def test_create_asset_upload_missing_asset_upload_type_logged( mocked_get_rate_limiter: MagicMock, mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, mocked_shadow_authorization: MagicMock, fixture_client: FlaskClient, asset_upload_type_value: Any, expected_missing_log: bool, caplog: pytest.LogCaptureFixture, ) -> None: """A missing or null asset_upload_type is logged so unmigrated callers are visible.""" caplog.set_level(logging.INFO) mocked_get_rate_limiter.return_value.hit.return_value = rate_limiting.Allowed() request_body: dict[str, Any] = { "product_id": 123, "track_unique_id": 0, "original_filename": "test.tif", } if asset_upload_type_value != "__omit__": request_body["asset_upload_type"] = asset_upload_type_value result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=OA_HEADERS ) expected_log_msg = ( "Missing field 'asset_upload_type' in create_asset_upload request body" ) assert (expected_log_msg in caplog.text) == expected_missing_log mocked_create_asset_upload.assert_called_once() assert result.status_code == 200 @pytest.mark.parametrize( ("asset_upload_type", "track_unique_id"), [ ("static_artwork", 0), ("stereo", 5), ("atmos", 5), ], ids=["image_no_track", "stereo_with_track", "atmos_with_track"], ) @patch("assets.handlers.pdp_auth.shadow_authorization") @patch( "assets.logic.generate_upload_data.create_asset_upload", return_value="create_asset_upload_logic_result", ) @patch( "owsrequest.flask_request.verify_grass_ownership", return_value=response.Response() ) @patch("assets.rate_limiting.get_rate_limiter") def test_create_asset_upload_valid_asset_upload_type_track_combinations( mocked_get_rate_limiter: MagicMock, mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, mocked_shadow_authorization: MagicMock, fixture_client: FlaskClient, asset_upload_type: str, track_unique_id: int, ) -> None: """Valid asset_upload_type + track_unique_id combinations pass the schema allOf checks.""" mocked_get_rate_limiter.return_value.hit.return_value = rate_limiting.Allowed() request_body: dict[str, Any] = { "product_id": 123, "track_unique_id": track_unique_id, "original_filename": "test.file", "asset_upload_type": asset_upload_type, } result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=OA_HEADERS ) mocked_create_asset_upload.assert_called_once() assert result.status_code == 200 @patch("assets.logic.generate_upload_data.create_asset_upload") @patch("owsrequest.flask_request.verify_grass_ownership") @patch("assets.rate_limiting.get_rate_limiter") def test_create_asset_upload_over_limit( mocked_get_rate_limiter: MagicMock, mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, fixture_client: FlaskClient, ) -> None: mocked_get_rate_limiter.return_value.hit.return_value = rate_limiting.Exceeded( "some rate limit message" ) headers = { **OA_HEADERS, "Orchard-User-Id": "oa:179", } request_body = { "product_id": 123, "track_unique_id": 0, "original_filename": "test.tif", } result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=headers ) assert mocked_get_rate_limiter.return_value.hit.mock_calls == [ call(rate_limiting.RateLimitedResource.UPLOAD, headers["Orchard-User-Id"]) ] mocked_verify_grass_ownership.assert_not_called() mocked_create_asset_upload.assert_not_called() assert result.status_code == 429 assert result.json == { "code": error.ERROR_CODE_RATE_LIMIT_EXCEEDED, "message": "some rate limit message", } @patch("assets.logic.generate_upload_data.create_asset_upload") @patch("owsrequest.flask_request.verify_grass_ownership") def test_create_asset_upload_fail_headers( mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, fixture_client: FlaskClient, ) -> None: """Test create_asset_upload failure with no headers.""" request_body = { "product_id": 123, "track_unique_id": 0, "original_filename": "test.tif", } result = fixture_client.post( "/v2/assets/upload", json=request_body, ) result_json = json.loads(result.data.decode()) mocked_verify_grass_ownership.assert_not_called() mocked_create_asset_upload.assert_not_called() assert result.status_code == 400 assert result_json["code"] == "error_header_validation" @patch("assets.logic.generate_upload_data.create_asset_upload") @patch("owsrequest.flask_request.verify_grass_ownership") def test_create_asset_upload_body_schema_failure( mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, fixture_client: FlaskClient, ) -> None: """Test create_asset_upload failure with invalid body.""" request_body = { "product_id": 0, "track_unique_id": 0, "original_filename": "test.tif", } result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=ALW_HEADERS ) result_json = json.loads(result.data.decode()) mocked_verify_grass_ownership.assert_not_called() mocked_create_asset_upload.assert_not_called() assert result.status_code == 400 assert result_json["code"] == "error_body_validation" @patch("assets.handlers.pdp_auth.shadow_authorization") @patch("assets.logic.generate_upload_data.create_asset_upload") @patch( "owsrequest.flask_request.verify_grass_ownership", return_value=response.Response(status=403, message="failed"), ) def test_create_asset_upload_ownership_check_fail( mocked_verify_grass_ownership: MagicMock, mocked_create_asset_upload: MagicMock, mocked_shadow_authorization: MagicMock, fixture_client: FlaskClient, ) -> None: """Test create_asset_upload ownership check fail.""" request_body = { "product_id": 123, "track_unique_id": 0, "original_filename": "test.tif", } result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=ALW_HEADERS ) mocked_verify_grass_ownership.assert_called_once_with( request, ownership.check_ownership, 123 ) mocked_create_asset_upload.assert_not_called() assert result.status_code == 403 assert result.data == b"failed" @pytest.mark.parametrize( ( "test_description", "request_body_keys_to_remove", "request_body_overrides", "expected_message", ), [ ( "product_id is missing", {"product_id"}, {}, "'product_id' is a required property", ), ( "track_unique_id is missing", {"track_unique_id"}, {}, "'track_unique_id' is a required property", ), ( "original_filename is missing", {"original_filename"}, {}, "'original_filename' is a required property", ), ( "image asset_upload_type with track", {}, {"asset_upload_type": "static_artwork", "track_unique_id": 456}, "0 was expected", ), ( "audio asset_upload_type without track", {}, {"asset_upload_type": "stereo", "track_unique_id": 0}, "0 is less than the minimum of 1", ), ( "content_type is not allowed", {}, {"content_type": "image/tiff"}, "Additional properties are not allowed", ), ], ) def test_create_asset_upload_body_validation_errors( fixture_client: FlaskClient, test_description: str, request_body_keys_to_remove: list[str], request_body_overrides: dict[str, Any], expected_message: str, ) -> None: """Test create_asset_upload body schema validation.""" request_body = { "product_id": 123, "track_unique_id": 456, "original_filename": "test.wav", } for key in request_body_keys_to_remove: request_body.pop(key) request_body.update(request_body_overrides) result = fixture_client.post( "/v2/assets/upload", json=request_body, headers=ALW_HEADERS ) result_json = json.loads(result.data.decode()) assert result.status_code == 400 assert result_json["code"] == "error_body_validation" assert expected_message in result_json["message"] @patch( "assets.logic.generate_upload_data.get_presigned_urls_for_asset_upload", return_value="nice", ) def test_get_presigned_urls_for_asset_upload_success( mocked_get_presigned_urls_for_asset_upload: MagicMock, fixture_client: FlaskClient ) -> None: """Test get_presigned_urls_for_asset_upload success.""" result = fixture_client.get( "/v2/assets/upload/some_filename", query_string={"part_numbers": "1,2,3"}, headers=ALW_HEADERS, ) mocked_get_presigned_urls_for_asset_upload.assert_called_once_with( part_numbers=[1, 2, 3], filename="some_filename", user_id="alw:12345" ) assert result.status_code == 200 assert result.data == b"nice" @patch("assets.logic.generate_upload_data.get_presigned_urls_for_asset_upload") def test_get_presigned_urls_for_asset_upload_fail_headers( mocked_get_presigned_urls_for_asset_upload: MagicMock, fixture_client: FlaskClient ) -> None: """Test get_presigned_urls_for_asset_upload failure with no headers.""" result = fixture_client.get( "/v2/assets/upload/some_filename", query_string={"part_numbers": "1,2,3"}, ) result_json = json.loads(result.data.decode()) mocked_get_presigned_urls_for_asset_upload.assert_not_called() assert result.status_code == 400 assert result_json["code"] == "error_header_validation" @patch( "assets.logic.generate_upload_data.get_presigned_urls_for_asset_upload", ) def test_get_presigned_urls_for_asset_upload_fail_query_string( mocked_get_presigned_urls_for_asset_upload: MagicMock, fixture_client: FlaskClient ) -> None: """Test get_presigned_urls_for_asset_upload failure with invalid query string.""" result = fixture_client.get( "/v2/assets/upload/some_filename", query_string={"part_numbers": "3,asdf,234"}, headers=ALW_HEADERS, ) result_json = json.loads(result.data.decode()) mocked_get_presigned_urls_for_asset_upload.assert_not_called() assert result.status_code == 400 assert result_json["code"] == "error_query_validation" assert "3,asdf,234' does not match '^[0-9]+(?:,[0-9]+)*$" in result_json["message"] @pytest.mark.parametrize( ( "part_numbers", "expected_message", ), [ pytest.param( "1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21", "[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21] is too long", id="part_numbers has too many items", ), pytest.param( "0,2,3", "0 is less than the minimum of 1", id="part_numbers has a too small item", ), pytest.param( "1,2,3,50000,6", "50000 is greater than the maximum of 10000", id="part_numbers has a too large item", ), pytest.param( "1,3,2,3,5", "[1, 3, 2, 3, 5] has non-unique elements", id="part_numbers has a duplicate", ), ], ) def test_get_presigned_urls_for_asset_upload_fail_query_string_part_numbers( part_numbers: str, expected_message: str, fixture_client: FlaskClient, ) -> None: """Test get_presigned_urls_for_asset_upload failure with invalid part_numbers.""" result = fixture_client.get( "/v2/assets/upload/some_filename", query_string={"part_numbers": part_numbers}, headers=ALW_HEADERS, ) result_json = json.loads(result.data.decode()) assert result.status_code == 400 assert result_json["code"] == "error_query_validation" assert expected_message in result_json["message"] def test_complete_multipart_upload_request_header_validation_failure( fixture_client: FlaskClient, ) -> None: """Test complete_multipart_upload request header validation failure.""" result = fixture_client.patch("/v2/assets/upload/some_filename") result_json = json.loads(result.data.decode()) assert result.status_code == 400 assert result_json["code"] == "error_header_validation" @pytest.mark.parametrize( ( "request_body", "expected_message", ), [ pytest.param({}, "'parts' is a required property", id="no parts"), pytest.param( {"parts": ["asdf", "asdf"]}, "has non-unique elements", id="non-array parts" ), pytest.param({"parts": list(range(10001))}, "is too long", id="too many parts"), pytest.param({"parts": []}, "[] should be non-empty", id="empty parts"), pytest.param( {"parts": [{"etag": "asdf"}]}, "'part_number' is a required property", id="missing part_number", ), pytest.param( {"parts": [{"part_number": 1}]}, "'etag' is a required property", id="missing etag", ), pytest.param( {"parts": [{"part_number": 1, "etag": ""}]}, "'' should be non-empty", id="empty etag", ), pytest.param( {"parts": [{"part_number": 0, "etag": "asdf"}]}, "0 is less than the minimum of 1", id="too small part_number", ), pytest.param( {"parts": [{"part_number": 10001, "etag": "asdf"}]}, "10001 is greater than the maximum of 10000", id="too large part_number", ), ], ) def test_complete_multipart_upload_request_body_validation_failures( request_body: dict[str, Any], expected_message: str, fixture_client: FlaskClient, ) -> None: """Test complete_multipart_upload request validation failures.""" result = fixture_client.patch( "/v2/assets/upload/some_filename", json=request_body, headers=ALW_HEADERS ) result_json = json.loads(result.data.decode()) assert result.status_code == 400 assert result_json["code"] == "error_body_validation" assert expected_message in result_json["message"] @patch( "assets.logic.generate_upload_data.complete_multipart_upload", return_value=None, ) def test_complete_multipart_upload_success( mocked_complete_multipart_upload: MagicMock, fixture_client: FlaskClient ) -> None: """Test complete_multipart_upload success.""" request_body = { "parts": [ {"part_number": 1, "etag": "asdf"}, {"part_number": 2, "etag": "asdf"}, {"part_number": 3, "etag": "asdf"}, ] } result = fixture_client.patch( "/v2/assets/upload/some_filename", json=request_body, headers=ALW_HEADERS ) mocked_complete_multipart_upload.assert_called_once_with( filename="some_filename", parts=request_body["parts"], user_id="alw:12345" ) assert result.status_code == 200 assert result.data.decode() == json.dumps({"status": error.SUCCESS_CODE}) @pytest.fixture def mock_response_data() -> list[dict[str, Any]]: return [ { "tuid": "123", "product_id": 333, "asset_type": "TIF", "s3_bucket": "my-bucket", "s3_key": "path/to/tif", "updated_date": "2024-01-01T12:00:00Z", "duration": 200, } ] @patch("assets.logic.product.get_assets_info_by_product_id") def test_get_assets_info_by_single_product_id( mock_get_assets_info_by_product_id: MagicMock, fixture_client: FlaskClient, mock_response_data: list[dict[str, Any]], ) -> None: """Test the get_assets_info_by_single_product_id endpoint.""" mock_get_assets_info_by_product_id.return_value = mock_response_data result = fixture_client.get("/product/333/assets") response_data = json.loads(result.data.decode()) assert result.status_code == 200 assert response_data == mock_response_data mock_get_assets_info_by_product_id.assert_called_with(333) @patch("assets.utils.handlers.g") @patch("assets.logic.product.get_assets_info_by_product_ids") def test_get_assets_info_by_many_product_ids( mock_get_assets_info_by_product_ids: MagicMock, mock_g: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, mock_response_data: list[dict[str, Any]], ) -> None: """Test the get_assets_info_by_many_product_ids endpoint.""" mock_get_assets_info_by_product_ids.return_value = mock_response_data mock_g.request_context.jwt_identity_id = ( authorization.BULK_ASSET_DOWNLOAD_IDENTITY_UUID ) app_response = fixture_client.get( "/v2/assets-bulk", query_string={"product_ids": "1,2,3", "asset_types": "FLAC,wav"}, ) assert app_response.status_code == 200 app_response_data = json.loads(app_response.data.decode()) assert app_response_data == mock_response_data mock_get_assets_info_by_product_ids.assert_called_with([1, 2, 3], ["FLAC", "WAV"]) @pytest.mark.parametrize( "query_params", [ ({}), ({"product_ids": "a,b,c"}), ({"product_ids": "-1,2,3"}), ({"product_ids": ",".join([str(x) for x in range(0, 1001)])}), ({"product_ids": "1,2,3", "asset_types": "MP4"}), ({"product_ids": "1,2,3,1", "asset_types": "FLAC,WAV"}), ({"product_ids": "1,2,3", "asset_types": "FLAC,WAV,FLAC"}), ({"product_ids": "[1,2,3]", "asset_types": "FLAC,WAV"}), ({"asset_types": "FLAC,WAV"}), ], ) @patch("assets.handlers.g") @patch("assets.logic.product.get_assets_info_by_product_ids") def test_get_assets_info_by_many_product_ids_invalid( mock_get_assets_info_by_product_ids: MagicMock, mock_g: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, query_params: dict[str, Any], mock_response_data: list[dict[str, Any]], ) -> None: """Test the get_assets_info_by_many_product_ids endpoint with invalid params.""" mock_g.request_context.jwt_identity_id = "dummy_token" app_response = fixture_client.get("/v2/assets-bulk", query_string=query_params) assert app_response.status_code == 401 mock_get_assets_info_by_product_ids.assert_not_called() @patch("assets.utils.handlers.g") @patch("assets.logic.product.get_assets_info_by_product_ids") def test_get_assets_info_by_many_product_ids_with_invalid_token( mock_get_assets_info_by_product_ids: MagicMock, mock_g: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, mock_response_data: list[dict[str, Any]], ) -> None: """Test the get_assets_info_by_many_product_ids endpoint.""" mock_get_assets_info_by_product_ids.return_value = mock_response_data mock_g.request_context.jwt_identity_id = ( authorization.BULK_SESSION_INGEST_COPY_ASSET_LAMBDA_IDENTITY_UUID ) app_response = fixture_client.get( "/v2/assets-bulk", query_string={"product_ids": "1,2,3", "asset_types": "FLAC,wav"}, ) assert app_response.status_code == 403 mock_hive_segment_request_data = { "asset_final_id": 12345, "task": { "task_id": "a22414c0-0cba-11f1-bc98-c7fc13f309f3", "model": "ai_music_classifier_DORIAN_2025_04_02_v00", "model_version": 1, }, "segments": [ { "time": 5, "ai_generated_music": 0.9, "ai_generated_music_vocal": 0.77, "mubert": 0.67, "musicgen": 0.55, "riffusion": 0.33, "stable_audio": 0.89, "suno": 0.78, "udio": 0.66, "yue": 0.11, "minimax": 0.22, "mureka": 0.33, "ace_step": 0.44, "duobao": 0.55, "google": 0.66, "heartmula": 0.77, "loudly": 0.88, } ], } @patch("assets.logic.hive_segment.save_hive_segment_data") @patch("assets.utils.handlers.g") def test_save_hive_segment_data_success( mock_g: MagicMock, mocked_save_hive_segment: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, ) -> None: """Test save_hive_segment_data success.""" mock_g.request_context.jwt_identity_id = authorization.HIVE_AI_DETECTION_UUID request_body = mock_hive_segment_request_data mocked_save_hive_segment.return_value = {} result = fixture_client.post( "/hive_segments", json=request_body, ) assert result.status_code == 200 mocked_save_hive_segment.assert_called_with(request_body) assert json.loads(result.data.decode()) == {"status": error.SUCCESS_CODE} @patch("assets.logic.hive_segment.save_hive_segment_data") @patch("assets.utils.handlers.g") def test_save_hive_segment_data_with_empty_fields( mock_g: MagicMock, mocked_save_hive_segment: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, ) -> None: mock_g.request_context.jwt_identity_id = authorization.HIVE_AI_DETECTION_UUID request_body = { "asset_final_id": 12345, "task": { "task_id": "a22414c0-0cba-11f1-bc98-c7fc13f309f3", "model": "ai_music_classifier_DORIAN_2025_04_02_v00", "model_version": 1, }, "segments": [ { "time": 5, "ai_generated_music": 0.9, "ai_generated_music_vocal": 0.77, } ], } expected_call = { "asset_final_id": 12345, "task": { "task_id": "a22414c0-0cba-11f1-bc98-c7fc13f309f3", "model": "ai_music_classifier_DORIAN_2025_04_02_v00", "model_version": 1, }, "segments": [ { "time": 5, "ai_generated_music": 0.9, "ai_generated_music_vocal": 0.77, "mubert": None, "musicgen": None, "riffusion": None, "stable_audio": None, "suno": None, "udio": None, "yue": None, "minimax": None, "mureka": None, "ace_step": None, "duobao": None, "google": None, "heartmula": None, "loudly": None, } ], } mocked_save_hive_segment.return_value = {} result = fixture_client.post( "/hive_segments", json=request_body, ) assert result.status_code == 200 mocked_save_hive_segment.assert_called_with(expected_call) response_data = json.loads(result.data.decode()) assert response_data == {"status": error.SUCCESS_CODE} @patch("assets.logic.hive_segment.save_hive_segment_data") @patch("assets.utils.handlers.g") def test_save_hive_segment_data_with_invalid_data( mock_g: MagicMock, mocked_save_hive_segment: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, ) -> None: """Test save_hive_segment_data with invalid request data.""" mock_g.request_context.jwt_identity_id = authorization.HIVE_AI_DETECTION_UUID request_body = { "asset_final_id": 12345, "time": "test1", "ai_generated_music": "test", } result = fixture_client.post( "/hive_segments", json=request_body, ) mocked_save_hive_segment.assert_not_called() assert result.status_code == 400 @patch("assets.logic.hive_segment.save_hive_segment_data") @patch("assets.utils.handlers.g") def test_save_hive_segment_data_with_missing_jwt( mock_g: MagicMock, mocked_save_hive_segment: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, ) -> None: """Test save_hive_segment_data with missing jwt.""" mock_g.request_context.jwt_identity_id = "" request_body = mock_hive_segment_request_data result = fixture_client.post( "/hive_segments", json=request_body, ) mocked_save_hive_segment.assert_not_called() assert result.status_code == 401 @patch("assets.logic.hive_segment.save_hive_segment_data") @patch("assets.utils.handlers.g") def test_save_hive_segment_data_with_invalid_jwt( mock_g: MagicMock, mocked_save_hive_segment: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, ) -> None: """Test save_hive_segment_data with invalid jwt.""" mock_g.request_context.jwt_identity_id = "dummy" request_body = mock_hive_segment_request_data result = fixture_client.post( "/hive_segments", json=request_body, ) mocked_save_hive_segment.assert_not_called() assert result.status_code == 403 def test_get_upload_permission_no_user(fixture_client: FlaskClient) -> None: """Test upload-token endpoint fails without user.""" new_headers = ALW_HEADERS.copy() del new_headers[field_const.ORCHARD_USER_ID] result = fixture_client.post( "/upload-token", headers=new_headers, data={field_const.TOKEN: "a_token", field_const.BULK_UPLOAD_COUNT: 10}, ) assert result.status_code == 401 @patch("assets.logic.hive_ai_image_task.save_hive_ai_image_task_data") @patch("assets.utils.handlers.g") def test_post_hive_ai_image_task_data_valid( mock_g: MagicMock, mock_save: MagicMock, fixture_client: FlaskClient, fixture_app: AppContext, ) -> None: mock_g.request_context.jwt_identity_id = authorization.HIVE_AI_DETECTION_UUID payload = { "asset_final_id": 1, "task_id": "task_one", "class_name": "test_class", "score_value": 0.95, } response = fixture_client.post("/hive_ai_image_task", json=payload) assert response.status_code == 200 mock_save.assert_called_once_with( asset_final_id=1, task_id="task_one", class_name="test_class", score_value=0.95, ) @patch("assets.logic.hive_ai_image_task.save_hive_ai_image_task_data") def test_post_hive_ai_image_task_data_invalid( mock_save: MagicMock, fixture_client: FlaskClient ) -> None: payload = { "asset_final_id": 1, } response = fixture_client.post("/hive_ai_image_task", json=payload) assert response.status_code == 400 mock_save.assert_not_called()