from marshmallow import Schema, ValidationError, fields, post_load, pre_dump, validate, pre_load from server.core.schemas import BaseSchema from server.artist.constants import COUNTRIES_LIST_SLUG_NAME from server.artist.schemas import ArtistDiscoveryCountRequestSchema, ArtistSearchByNameDataResponseSchema from server.track.constants import GENRES_LIST_SLUG_NAME from server.track.schemas import TrackSearchByNameDataResponseSchema, TrackDiscoveryCountRequestSchema from server.dna.constants import ARTIST, TRACK from server.core.exceptions import BaseSchemaErrorRequest PARAMS_MAPPER = {ARTIST: ArtistDiscoveryCountRequestSchema, TRACK: TrackDiscoveryCountRequestSchema} DATA_MAPPER = {ARTIST: ArtistSearchByNameDataResponseSchema, TRACK: TrackSearchByNameDataResponseSchema} class UserSearchSchema: class RequestSchema(BaseSchema): name = fields.String(required=True) params = fields.Dict(required=True) type = fields.String(validate=validate.OneOf([ARTIST, TRACK]), required=True) @pre_load def load_params_schema(self, data, **kwargs): if all(keys in data for keys in ("params", "type")): self.fields["params"] = fields.Nested( PARAMS_MAPPER[data["type"]]().load(data=data["params"]), required=True ) return data @post_load def checK_params(self, data, **kwargs): if not data["params"]: raise ValidationError(message="Params should not be empty", field_name="error") return data class RequestArchiveSchema(BaseSchema): archived = fields.Boolean(required=True) class ResponseGetSchema(Schema): id = fields.Integer() name = fields.String() user_id = fields.String() params = fields.Dict() created_at = fields.DateTime() updated_at = fields.DateTime() archived = fields.DateTime() type = fields.String() class ResponseSchema(Schema): id = fields.Integer() class ResponseArchiveSchema(Schema): archived = fields.Boolean(required=True) class ResponseDeleteSchema(Schema): message = fields.String(default="Success") class ProfileSchemaMixin(Schema): email = fields.String() name = fields.String() class ProfileInfoResponse(ProfileSchemaMixin): labels = fields.Dict(keys=fields.String(), values=fields.String()) roles = fields.List(fields.String()) user_id = fields.String() class ProfileResetResponse(Schema): user_id = fields.String() message = fields.String(default="Profile was reset") class LabelsWithUsersResponse(Schema): id = fields.String() name = fields.String() class Meta: class UserSchema(Schema): pass include = {"users": fields.List(fields.Nested(ProfileSchemaMixin))} class ErrorResponseSchema(Schema): error = fields.String() class RecentSearch: class CreatedRecentSearchResponseSchema(Schema): id = fields.Integer() class EntityIdResponseSchema(BaseSchema): entity_id = fields.String(required=True, validate=validate.Length(min=1, max=255)) class EntityTypeSchema(EntityIdResponseSchema): entity_type = fields.String(validate=validate.OneOf([ARTIST, TRACK]), required=True) class ResponseSchema(EntityTypeSchema): data = fields.Dict() @pre_dump def check_data_field(self, data, **kwargs): self.fields["data"] = fields.Nested(DATA_MAPPER[data["entity_type"]]().load(data=data["data"])) return data class RequestSchema(EntityTypeSchema): class Meta: fields = ( "entity_id", "entity_type", ) class EntityTypeSchema(BaseSchema): entity_type = fields.String(validate=validate.OneOf(["genres", "countries"]), required=True) class FavoritesRequestShema(BaseSchema): entity_id_mapper = { "genres": GENRES_LIST_SLUG_NAME, "countries": COUNTRIES_LIST_SLUG_NAME, } entity_type = fields.String(validate=validate.OneOf(["genres", "countries"]), required=True) entity_id = fields.String() @post_load def check_entity_id(self, data, **kwargs): if data["entity_id"] not in self.entity_id_mapper[data["entity_type"]]: raise BaseSchemaErrorRequest( error_message=f"irrelevant entity_id {data['entity_id']} for entity_type {data['entity_type']}" ) return data