diff --git a/pokemon_v2/serializers.py b/pokemon_v2/serializers.py index 6da9afb1c..92e442388 100644 --- a/pokemon_v2/serializers.py +++ b/pokemon_v2/serializers.py @@ -3,6 +3,7 @@ from __future__ import annotations import itertools +from collections import defaultdict from typing import TYPE_CHECKING, Any, ClassVar, Protocol, cast from django.db.models import Q @@ -3487,8 +3488,17 @@ def build_chain(self, obj: EvolutionChain) -> dict[str, Any]: PokemonSpeciesEvolutionSerializer(pokemon_objects, many=True, context=self.context).data, ) + evolutions_by_species: dict[int, list[PokemonEvolution]] = defaultdict(list) + if any(species["evolves_from_species"] for species in ref_data): + for evolution in ( + PokemonEvolution.objects.filter(evolved_species__evolution_chain=obj) + .select_related(*self.POKEMON_EVOLUTION_FK_FIELDS) + .order_by("pk") + ): + evolutions_by_species[evolution.evolved_species_id].append(evolution) # pyright: ignore[reportAttributeAccessIssue] + evolution_tree = self.build_evolution_tree(ref_data) - return self.build_chain_link_entry(evolution_tree, summary_data) + return self.build_chain_link_entry(evolution_tree, summary_data, evolutions_by_species) # converts a list of Pokemon species evolution data into a tree representing the evolution chain def build_evolution_tree(self, species_evolution_data: ReturnList[ReturnDict[str, Any]]) -> dict[str, Any]: @@ -3525,15 +3535,16 @@ def build_evolution_tree(self, species_evolution_data: ReturnList[ReturnDict[str # serializes an evolution chain link recursively # chain_link is a tree representing an evolution chain def build_chain_link_entry( - self, chain_link: dict[str, Any], summary_data: ReturnList[ReturnDict[str, Any]] + self, + chain_link: dict[str, Any], + summary_data: ReturnList[ReturnDict[str, Any]], + evolutions_by_species: dict[int, list[PokemonEvolution]], ) -> dict[str, Any]: species = chain_link["species"] evolution_data = None if species["evolves_from_species"]: - evolution_objects = PokemonEvolution.objects.filter(evolved_species=species["id"]).select_related( - *self.POKEMON_EVOLUTION_FK_FIELDS - ) + evolution_objects = evolutions_by_species.get(species["id"], []) evolution_data = cast( "ReturnList[ReturnDict[str, Any]]", PokemonEvolutionSerializer(evolution_objects, many=True, context=self.context).data, @@ -3543,7 +3554,9 @@ def build_chain_link_entry( "is_baby": species["is_baby"], "species": next(x for x in summary_data if x["name"] == species["name"]), "evolution_details": evolution_data or [], - "evolves_to": [self.build_chain_link_entry(c, summary_data) for c in chain_link["children"]], + "evolves_to": [ + self.build_chain_link_entry(c, summary_data, evolutions_by_species) for c in chain_link["children"] + ], } diff --git a/pokemon_v2/tests.py b/pokemon_v2/tests.py index a6d8eac62..a4b43e8e2 100644 --- a/pokemon_v2/tests.py +++ b/pokemon_v2/tests.py @@ -1,6 +1,8 @@ import json from datetime import datetime, timezone +from django.db import connection +from django.test.utils import CaptureQueriesContext from rest_framework import status from rest_framework.test import APITestCase @@ -4976,6 +4978,81 @@ def test_evolution_chain_api_wurmple_bugfix(self): stage_one_second_data = basic_data["evolves_to"][1] self.assertEqual(len(stage_one_second_data["evolves_to"]), 1) + # verifies that building the evolution chain tree issues a constant number of + # queries instead of one PokemonEvolution query per non-root species in the chain + def test_evolution_chain_api_query_count_does_not_scale_with_chain_size(self): + def build_branching_chain(branch_count): + evolution_chain = self.setup_evolution_chain_data() + basic = self.setup_pokemon_species_data( + name=f"bsc for evo chn qc {branch_count}", + evolution_chain=evolution_chain, + ) + for i in range(branch_count): + branch_species = self.setup_pokemon_species_data( + name=f"brnch {i} for evo chn qc {branch_count}", + evolves_from_species=basic, + evolution_chain=evolution_chain, + ) + self.setup_pokemon_evolution_data(evolved_species=branch_species, min_level=7) + return evolution_chain + + small_chain = build_branching_chain(branch_count=1) + large_chain = build_branching_chain(branch_count=6) + + with CaptureQueriesContext(connection) as small_queries: + small_response = self.client.get("{}/evolution-chain/{}/".format(API_V2, small_chain.pk)) + with CaptureQueriesContext(connection) as large_queries: + large_response = self.client.get("{}/evolution-chain/{}/".format(API_V2, large_chain.pk)) + + self.assertEqual(small_response.status_code, status.HTTP_200_OK) + self.assertEqual(large_response.status_code, status.HTTP_200_OK) + self.assertEqual(len(large_response.data["chain"]["evolves_to"]), 6) + + # before the fix, each additional branch added its own PokemonEvolution query + # (one per non-root species), so 5 extra branches meant 5 extra queries here + self.assertEqual( + len(large_queries.captured_queries), + len(small_queries.captured_queries), + ) + + def test_evolution_chain_api_single_species_chain_skips_evolution_query(self): + evolution_chain = self.setup_evolution_chain_data() + self.setup_pokemon_species_data(name="sngl for evo chn", evolution_chain=evolution_chain) + + with CaptureQueriesContext(connection) as queries: + response = self.client.get("{}/evolution-chain/{}/".format(API_V2, evolution_chain.pk)) + + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(response.data["chain"]["evolution_details"], []) + self.assertFalse(any("pokemon_v2_pokemonevolution" in query["sql"] for query in queries.captured_queries)) + + # evolution_details must keep creation (pk) order when a species has many PokemonEvolution + # rows, e.g. Milcery -> Alcremie has one row per flavor/decoration combination + def test_evolution_chain_api_evolution_details_order_with_many_rows_for_same_species(self): + evolution_chain = self.setup_evolution_chain_data() + basic = self.setup_pokemon_species_data( + name="bsc for evo chn ordr", + evolution_chain=evolution_chain, + ) + target = self.setup_pokemon_species_data( + name="trgt for evo chn ordr", + evolves_from_species=basic, + evolution_chain=evolution_chain, + ) + + expected_min_levels = [30, 10, 50, 20, 40] + for min_level in expected_min_levels: + self.setup_pokemon_evolution_data(evolved_species=target, min_level=min_level) + + response = self.client.get("{}/evolution-chain/{}/".format(API_V2, evolution_chain.pk)) + + self.assertEqual(response.status_code, status.HTTP_200_OK) + target_data = response.data["chain"]["evolves_to"][0] + self.assertEqual( + [detail["min_level"] for detail in target_data["evolution_details"]], + expected_min_levels, + ) + # Encounter Tests def test_encounter_method_api(self): encounter_method = self.setup_encounter_method_data(name="base encntr mthd")