feat: filter image tasks by input count
This commit is contained in:
+27
-1
@@ -1,4 +1,5 @@
|
||||
from django.contrib import admin
|
||||
from django.db.models import Count, Q
|
||||
|
||||
from .models import ImageGenerationTask, ImageGenerationTaskInput
|
||||
|
||||
@@ -11,6 +12,31 @@ class ImageGenerationTaskInputInline(admin.TabularInline):
|
||||
readonly_fields = fields
|
||||
|
||||
|
||||
class ImageInputTypeFilter(admin.SimpleListFilter):
|
||||
title = "输入图片类型"
|
||||
parameter_name = "input_image_type"
|
||||
|
||||
def lookups(self, request, model_admin):
|
||||
return (
|
||||
("single", "单图生图"),
|
||||
("multiple", "多图生图"),
|
||||
)
|
||||
|
||||
def queryset(self, request, queryset):
|
||||
value = self.value()
|
||||
if value not in {"single", "multiple"}:
|
||||
return queryset
|
||||
|
||||
input_counts = queryset.annotate(input_image_count=Count("input_images"))
|
||||
if value == "multiple":
|
||||
task_ids = input_counts.filter(input_image_count__gte=2).values("pk")
|
||||
return queryset.filter(pk__in=task_ids)
|
||||
|
||||
new_single_ids = input_counts.filter(input_image_count=1).values("pk")
|
||||
legacy_single_ids = queryset.filter(input_images__isnull=True).exclude(input_image="").values("pk")
|
||||
return queryset.filter(Q(pk__in=new_single_ids) | Q(pk__in=legacy_single_ids))
|
||||
|
||||
|
||||
@admin.register(ImageGenerationTask)
|
||||
class ImageGenerationTaskAdmin(admin.ModelAdmin):
|
||||
change_form_template = "admin/api/imagegenerationtask/change_form.html"
|
||||
@@ -25,7 +51,7 @@ class ImageGenerationTaskAdmin(admin.ModelAdmin):
|
||||
"created_at",
|
||||
"finished_at",
|
||||
)
|
||||
list_filter = ("status", "created_at", "next_attempt_at", "finished_at")
|
||||
list_filter = (ImageInputTypeFilter, "status", "created_at", "next_attempt_at", "finished_at")
|
||||
search_fields = (
|
||||
"task_id",
|
||||
"user__username",
|
||||
|
||||
@@ -801,6 +801,61 @@ class ImageGenerationTaskAdminTests(TestCase):
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
|
||||
def test_changelist_filters_new_and_legacy_single_image_tasks(self):
|
||||
new_single = self.create_task()
|
||||
self.add_input_image(new_single, ordinal=0, filename="new-single.png")
|
||||
legacy_single = self.create_task()
|
||||
legacy_single.input_image.save("legacy-single.png", ContentFile(b"legacy-image"), save=True)
|
||||
multiple = self.create_task()
|
||||
self.add_input_image(multiple, ordinal=0, filename="multiple-main.png")
|
||||
self.add_input_image(multiple, ordinal=1, filename="multiple-reference.png")
|
||||
no_input = self.create_task()
|
||||
|
||||
response = self.client.get(
|
||||
reverse("admin:api_imagegenerationtask_changelist"),
|
||||
{"input_image_type": "single"},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertContains(response, "输入图片类型")
|
||||
self.assertContains(response, "单图生图")
|
||||
self.assertContains(response, "多图生图")
|
||||
self.assertContains(response, str(new_single.task_id))
|
||||
self.assertContains(response, str(legacy_single.task_id))
|
||||
self.assertNotContains(response, str(multiple.task_id))
|
||||
self.assertNotContains(response, str(no_input.task_id))
|
||||
self.assertCountEqual(
|
||||
response.context["cl"].queryset.values_list("pk", flat=True),
|
||||
[new_single.pk, legacy_single.pk],
|
||||
)
|
||||
|
||||
def test_changelist_filters_multiple_image_tasks_and_composes_with_status_filter(self):
|
||||
succeeded_multiple = self.create_task(status=ImageGenerationTask.Status.SUCCEEDED)
|
||||
self.add_input_image(succeeded_multiple, ordinal=0, filename="succeeded-main.png")
|
||||
self.add_input_image(succeeded_multiple, ordinal=1, filename="succeeded-reference.png")
|
||||
failed_multiple = self.create_task(status=ImageGenerationTask.Status.FAILED)
|
||||
self.add_input_image(failed_multiple, ordinal=0, filename="failed-main.png")
|
||||
self.add_input_image(failed_multiple, ordinal=1, filename="failed-reference.png")
|
||||
single = self.create_task()
|
||||
self.add_input_image(single, ordinal=0, filename="single.png")
|
||||
|
||||
response = self.client.get(
|
||||
reverse("admin:api_imagegenerationtask_changelist"),
|
||||
{
|
||||
"input_image_type": "multiple",
|
||||
"status__exact": ImageGenerationTask.Status.SUCCEEDED,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertContains(response, str(succeeded_multiple.task_id))
|
||||
self.assertNotContains(response, str(failed_multiple.task_id))
|
||||
self.assertNotContains(response, str(single.task_id))
|
||||
self.assertEqual(
|
||||
list(response.context["cl"].queryset.values_list("pk", flat=True)),
|
||||
[succeeded_multiple.pk],
|
||||
)
|
||||
|
||||
|
||||
@override_settings(
|
||||
PAYMENT_CALLBACK_MODE="mock",
|
||||
|
||||
Reference in New Issue
Block a user