"""A local N+1 regression example. Requires Django 5.2.

Run: python django_n_plus_one.py
Uses a fresh in-memory SQLite database, no website settings or network calls.
"""

import unittest

import django
from django.conf import settings

settings.configure(
    INSTALLED_APPS=[],
    DATABASES={"default": {"ENGINE": "django.db.backends.sqlite3", "NAME": ":memory:"}},
    DEFAULT_AUTO_FIELD="django.db.models.AutoField",
)
django.setup()

from django.db import connection, models
from django.test.utils import CaptureQueriesContext


class Customer(models.Model):
    name = models.CharField(max_length=80)

    class Meta:
        app_label = "n_plus_one_example"


class Order(models.Model):
    customer = models.ForeignKey(Customer, on_delete=models.CASCADE)

    class Meta:
        app_label = "n_plus_one_example"


def customer_names(orders):
    return [order.customer.name for order in orders]


class QueryGrowthTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        with connection.schema_editor() as editor:
            editor.create_model(Customer)
            editor.create_model(Order)

    @classmethod
    def tearDownClass(cls):
        with connection.schema_editor() as editor:
            editor.delete_model(Order)
            editor.delete_model(Customer)

    def test_same_output_without_query_growth(self):
        for size in (1, 5, 20):
            with self.subTest(orders=size):
                Order.objects.all().delete()
                Customer.objects.all().delete()
                for index in range(size):
                    customer = Customer.objects.create(name=f"Customer {index}")
                    Order.objects.create(customer=customer)

                with CaptureQueriesContext(connection) as before:
                    original = customer_names(Order.objects.order_by("id"))
                with CaptureQueriesContext(connection) as after:
                    fixed = customer_names(
                        Order.objects.select_related("customer").order_by("id")
                    )

                self.assertEqual(original, [f"Customer {i}" for i in range(size)])
                self.assertEqual(fixed, original)
                self.assertEqual(len(before), size + 1)
                self.assertEqual(len(after), 1)
                print(f"{size} orders: {len(before)} -> {len(after)} queries; same output")


if __name__ == "__main__":
    unittest.main()
