Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 19 additions & 6 deletions pokemon_v2/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]:
Expand Down Expand Up @@ -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,
Expand All @@ -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"]
],
}


Expand Down
77 changes: 77 additions & 0 deletions pokemon_v2/tests.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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")
Expand Down
Loading