import pytest from http import HTTPStatus as http_status from apollo_utils.core.constants.dsp import DSP from server.db.models import Message from server.db.session import db_session as session_scope from tests.utils import check_object, get_country_code, create_accounts def get_dsp(i): return (DSP.APPLE if i % 2 else DSP.SPOTIFY).value def get_message_type(i): return {i: v for i, v in enumerate( ( "starred_track_top_chart_entry", "starred_track_top_chart_exit", "starred_track_top_chart_move" ) )}[i % 3] def create_message_data(i, **kwargs) -> dict: data = { "ttl": i * 10, "event_id": None, "account_id": None, "message_id": None, "external_id": None, "meta": { "dsp": get_dsp(i), "country_code": get_country_code(i), "views": ["feed_message", "push_message"], "type": get_message_type(i) }, "data": { "i": i, "track": { "id": f"id_{i}", "isrc": f"isrc_{i}", "name": f"track_name_{i}", "artists": [{"id": f"artist_id_{i}", "name": f"artist_name_{i}"}], "image_url": f"https://image-{i}" }, "current_position": i * 3 if i % 3 != 1 else None, "previous_position": 1 * 7 if i % 3 != 0 else None } } for k, v in kwargs.items(): data[k] = v return data def check_bulk_create_response(response, ok_indexes=None, failed_indexes=None): ok_set, failed_set = set(), set() for item in response.get("ok", []): assert "id" in item ok_set.add(item["index"]) for item in response.get("failed", []): assert item.get("error") failed_set.add(item["index"]) assert ok_set == set(ok_indexes or []) assert failed_set == set(failed_indexes or []) @pytest.mark.parametrize( "data,ok_indexes,failed_indexes,headers,post_status,get_status", ( ({}, None, None, {}, http_status.BAD_REQUEST, http_status.BAD_REQUEST), ( { "public": False, "data": [ create_message_data(1) ], }, None, None, {"X-App-Slug": "app7"}, http_status.NOT_FOUND, http_status.OK ), ( { "public": False, "data": [create_message_data(i, ttl=-100) for i in range(1, 3)] }, None, list(range(2)), {"X-App-Slug": "app1"}, http_status.OK, http_status.OK ), ( { "public": False, "data": [create_message_data(i) for i in range(1, 4)] }, list(range(3)), None, {"X-App-Slug": "app1"}, http_status.OK, http_status.OK ), ( { "public": False, "data": [create_message_data(1, account_id=2, data=None), create_message_data(2, ttl=-100)] }, (0,), (1,), {"X-App-Slug": "app1"}, http_status.OK, http_status.OK ), ) ) async def test_messages( data, ok_indexes, failed_indexes, headers, post_status, get_status, db_session, client, auth): async with session_scope() as session: await create_accounts(session, 2) headers.update(auth) app_slug = headers.get("X-App-Slug") ok_input_items = [item for i, item in enumerate(data.get("data", [])) if ok_indexes and i in ok_indexes] default = { "public": data.get("public", True), "data": None, "meta": None, "event_id": None, "account_id": None, "message_id": None, } id_list = [] post_response = await client.post( "/api/service/messages/", headers=headers, json=data) assert post_response.status == post_status if post_status == http_status.OK: result = await post_response.json() check_bulk_create_response(result, ok_indexes=ok_indexes, failed_indexes=failed_indexes) async with session_scope(): messages, _ = await Message.list(order_by=Message.id.asc()) assert len(messages) == len(ok_indexes or []) for i, message in enumerate(messages): input_item = ok_input_items[i] unset = {k: default[k] for k in (default.keys() - input_item.keys())} unset["app_slug"] = app_slug check_object(message, input_item, **unset) id_list.append(message.id) list_data = { "id": id_list or [1, 2], "include": "all" } get_response = await client.post( "/api/service/messages/list/", headers=headers, json=list_data) assert get_response.status == get_status if get_status == http_status.OK: result = await get_response.json() list_response_items = result.get("data", []) assert len(list_response_items) == len(id_list) for i, response_item in enumerate(list_response_items): input_data = ok_input_items[i] input_item = {k: default[k] for k in (default.keys() - input_data.keys())} input_item.update(input_data) input_item.update({"id": id_list[i], "context": None}) assert response_item.pop("created_at") assert response_item.pop("updated_at") assert response_item == input_item