Skip to content

Commit 3a99adf

Browse files
maxim-f1maxim-f1mmzeynalli
authored
[Fix] AJAX fields: persist selected values after validation errors (#1039)
* * Add `format_by_pk` * Fix `form` in context after form validation * Now, even with override Field, the loader field will be used (for example, for custom AJAX widgets) * If the loader field is None, an empty list will be returned instead of an error. * If the model has not been transferred, it will be automatically loaded by the PK field. * fix: enhance AJAX loader with improved error handling and logging and fix composite pk formatting --------- Co-authored-by: maxim-f1 <“maxconal228@gmail.com”> Co-authored-by: Miradil Zeynalli <miradil.zeynalli@gmail.com>
1 parent 3b752c8 commit 3a99adf

7 files changed

Lines changed: 236 additions & 7 deletions

File tree

sqladmin/ajax.py

Lines changed: 38 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,11 @@
44

55
from sqlalchemy import String, cast, inspect, or_, select
66

7-
from sqladmin.helpers import get_object_identifier, get_primary_keys
7+
from sqladmin.helpers import (
8+
get_object_identifier,
9+
get_primary_keys,
10+
object_identifier_values,
11+
)
812

913
if TYPE_CHECKING:
1014
from sqladmin.models import ModelView
@@ -61,6 +65,39 @@ def format(self, model: type) -> dict[str, Any]:
6165

6266
return {"id": str(get_object_identifier(model)), "text": str(model)}
6367

68+
async def format_by_pk(self, pk: Any) -> dict[str, Any]:
69+
if pk is None:
70+
return {}
71+
72+
stmt = select(self.model)
73+
primary_keys = tuple(inspect(self.model).primary_key)
74+
75+
try:
76+
values = object_identifier_values(str(pk), self.model)
77+
except (TypeError, ValueError):
78+
return {}
79+
80+
if len(values) != len(primary_keys):
81+
return {}
82+
83+
conditions = [field == value for field, value in zip(primary_keys, values)]
84+
stmt = stmt.where(*conditions)
85+
86+
if self.order_by:
87+
if isinstance(self.order_by, list):
88+
for o in self.order_by:
89+
stmt = stmt.order_by(o)
90+
else:
91+
stmt = stmt.order_by(self.order_by)
92+
93+
stmt = stmt.limit(1)
94+
95+
result = await self.model_admin._run_query(stmt)
96+
if len(result) < 1:
97+
return {}
98+
99+
return {"id": str(get_object_identifier(result[0])), "text": str(result[0])}
100+
64101
async def get_list(self, term: str) -> list[Any]:
65102
stmt = select(self.model)
66103

sqladmin/application.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -704,8 +704,9 @@ async def edit(self, request: Request) -> Response:
704704

705705
form_data = await self._handle_form_data(request, model)
706706
form = Form(form_data)
707+
context["form"] = form
708+
707709
if not form.validate():
708-
context["form"] = form
709710
return await self.templates.TemplateResponse(
710711
request, model_view.edit_template, context, status_code=400
711712
)

sqladmin/fields.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -299,7 +299,7 @@ def pre_validate(self, form: Form) -> None:
299299

300300

301301
class AjaxSelectField(fields.SelectFieldBase):
302-
widget = sqladmin_widgets.AjaxSelect2Widget()
302+
widget = sqladmin_widgets.AjaxSelect2Widget() # type: ignore[assignment]
303303
separator = ","
304304

305305
def __init__(

sqladmin/forms.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -314,6 +314,9 @@ async def convert(
314314
if not issubclass(override, Field):
315315
raise TypeError("Expected Field, got %s" % type(override))
316316

317+
if loader:
318+
kwargs.setdefault("loader", loader)
319+
317320
return override(**kwargs)
318321

319322
multiple = (

sqladmin/widgets.py

Lines changed: 51 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
# mypy: disable-error-code="override"
22

33
import json
4+
import logging
45
from typing import TYPE_CHECKING, Any
56

67
from markupsafe import Markup
@@ -17,6 +18,8 @@
1718
"Select2TagsWidget",
1819
]
1920

21+
logger = logging.getLogger(__name__)
22+
2023

2124
class DatePickerWidget(widgets.TextInput):
2225
"""
@@ -43,7 +46,7 @@ def __init__(self, multiple: bool = False):
4346
self.multiple = multiple
4447
self.lookup_url = ""
4548

46-
def __call__(self, field: "AjaxSelectField", **kwargs: Any) -> Markup:
49+
async def __call__(self, field: "AjaxSelectField", **kwargs: Any) -> Markup:
4750
kwargs.setdefault("data-role", "select2-ajax")
4851
kwargs.setdefault("data-url", field.loader.model_admin.ajax_lookup_url)
4952

@@ -55,14 +58,59 @@ def __call__(self, field: "AjaxSelectField", **kwargs: Any) -> Markup:
5558
kwargs.setdefault("type", "hidden")
5659

5760
if self.multiple:
58-
result = [field.loader.format(value) for value in field.data]
61+
result = []
62+
for value in field.data:
63+
try:
64+
result.append(field.loader.format(value))
65+
continue
66+
except Exception:
67+
logger.debug(
68+
"Fallback to format_by_pk for ajax value=%r",
69+
value,
70+
exc_info=True,
71+
)
72+
73+
try:
74+
result_value = await field.loader.format_by_pk(value)
75+
if result_value == {}:
76+
continue
77+
else:
78+
result.append(result_value)
79+
except Exception:
80+
logger.debug(
81+
"Unable to resolve ajax value by pk for field=%s value=%r",
82+
field.name,
83+
value,
84+
exc_info=True,
85+
)
86+
5987
kwargs.setdefault("data-json", json.dumps(result))
6088
kwargs.setdefault("multiple", "1")
89+
6190
else:
6291
try:
6392
data = field.loader.format(field.data)
6493
except Exception:
65-
data = None
94+
logger.debug(
95+
"Fallback to format_by_pk for ajax field=%s value=%r",
96+
field.name,
97+
field.data,
98+
exc_info=True,
99+
)
100+
try:
101+
data = await field.loader.format_by_pk(field.data)
102+
if data == {}:
103+
data = None
104+
except Exception:
105+
logger.debug(
106+
"Unable to resolve ajax single value by pk for "
107+
"field=%s value=%r",
108+
field.name,
109+
field.data,
110+
exc_info=True,
111+
)
112+
data = None
113+
66114
if data:
67115
kwargs.setdefault("data-json", json.dumps([data]))
68116

tests/test_ajax.py

Lines changed: 114 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
from starlette.applications import Starlette
99

1010
from sqladmin import Admin, ModelView
11-
from sqladmin.ajax import create_ajax_loader
11+
from sqladmin.ajax import QueryAjaxModelLoader, create_ajax_loader
1212
from tests.common import async_engine as engine
1313

1414
pytestmark = pytest.mark.anyio
@@ -77,6 +77,40 @@ def __str__(self) -> str:
7777
return f"Room {self.id}"
7878

7979

80+
class Team(Base):
81+
__tablename__ = "teams"
82+
83+
id = Column(Integer, primary_key=True)
84+
name = Column(String(length=32), nullable=False)
85+
86+
def __str__(self) -> str:
87+
return f"Team {self.id}"
88+
89+
90+
class Member(Base):
91+
__tablename__ = "members"
92+
93+
id = Column(Integer, primary_key=True)
94+
name = Column(String(length=32), nullable=False)
95+
team_id = Column(Integer, ForeignKey("teams.id"))
96+
97+
team = relationship("Team")
98+
99+
def __str__(self) -> str:
100+
return f"Member {self.id}"
101+
102+
103+
class CompositeTag(Base):
104+
__tablename__ = "composite_tags"
105+
106+
key = Column(String(length=16), primary_key=True)
107+
locale = Column(String(length=8), primary_key=True)
108+
label = Column(String(length=32), nullable=False)
109+
110+
def __str__(self) -> str:
111+
return f"{self.label}:{self.key}:{self.locale}"
112+
113+
80114
class UserAdmin(ModelView, model=User):
81115
form_ajax_refs = {
82116
"addresses": {
@@ -110,9 +144,18 @@ class RoomAdmin(ModelView, model=Room):
110144
}
111145

112146

147+
class MemberAdmin(ModelView, model=Member):
148+
form_ajax_refs = {
149+
"team": {
150+
"fields": ("name",),
151+
}
152+
}
153+
154+
113155
admin.add_view(UserAdmin)
114156
admin.add_view(AddressAdmin)
115157
admin.add_view(RoomAdmin)
158+
admin.add_view(MemberAdmin)
116159

117160

118161
@pytest.fixture(autouse=True)
@@ -330,3 +373,73 @@ async def test_create_and_edit_forms(client: AsyncClient) -> None:
330373

331374
user = result.scalar_one()
332375
assert len(user.addresses) == 2
376+
377+
378+
async def test_edit_validation_error_preserves_selected_ajax_value(
379+
client: AsyncClient,
380+
) -> None:
381+
async with session_maker() as s:
382+
s.add_all([Team(name="A"), Team(name="B")])
383+
await s.commit()
384+
385+
async with session_maker() as s:
386+
member = Member(name="John", team_id=1)
387+
s.add(member)
388+
await s.commit()
389+
390+
response = await client.post(
391+
"/admin/member/edit/1",
392+
data={"name": "", "team": "2"},
393+
)
394+
395+
assert response.status_code == 400
396+
assert (
397+
'data-json="[{&#34;id&#34;: &#34;2&#34;, &#34;text&#34;: &#34;Team 2&#34;}]"'
398+
in response.text
399+
)
400+
401+
402+
async def test_format_by_pk_single_pk() -> None:
403+
async with session_maker() as s:
404+
user = User(name="Arya")
405+
s.add(user)
406+
await s.commit()
407+
408+
loader = QueryAjaxModelLoader(
409+
name="user",
410+
model=User,
411+
model_admin=UserAdmin(),
412+
fields=("name",),
413+
)
414+
415+
assert await loader.format_by_pk(1) == {"id": "1", "text": "User 1"}
416+
417+
418+
async def test_format_by_pk_composite_pk_identifier() -> None:
419+
async with session_maker() as s:
420+
tag = CompositeTag(key="greeting", locale="en", label="Hello")
421+
s.add(tag)
422+
await s.commit()
423+
424+
loader = QueryAjaxModelLoader(
425+
name="composite",
426+
model=CompositeTag,
427+
model_admin=UserAdmin(),
428+
fields=("label",),
429+
)
430+
431+
assert await loader.format_by_pk("greeting;en") == {
432+
"id": "greeting;en",
433+
"text": "Hello:greeting:en",
434+
}
435+
436+
437+
async def test_format_by_pk_returns_empty_for_missing_record() -> None:
438+
loader = QueryAjaxModelLoader(
439+
name="user",
440+
model=User,
441+
model_admin=UserAdmin(),
442+
fields=("name",),
443+
)
444+
445+
assert await loader.format_by_pk("999") == {}

tests/test_forms/test_forms.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@
3232
from wtforms.fields.core import UnboundField
3333

3434
from sqladmin import ModelView
35+
from sqladmin.ajax import create_ajax_loader
3536
from sqladmin.fields import Select2TagsField, SelectField
3637
from sqladmin.forms import ModelConverter, converts, get_model_form
3738
from tests.common import async_engine as engine
@@ -233,6 +234,32 @@ class ExampleField(Field):
233234
assert not isinstance(Form()._fields["email"], ExampleField)
234235

235236

237+
async def test_model_form_override_receives_ajax_loader() -> None:
238+
class LoaderAwareField(Field):
239+
def __init__(self, *args: Any, loader=None, **kwargs: Any) -> None:
240+
self.loader = loader
241+
kwargs.pop("allow_blank", None)
242+
super().__init__(*args, **kwargs)
243+
244+
class AddressAdmin(ModelView, model=Address):
245+
form_ajax_refs = {"user": {"fields": ("name",)}}
246+
247+
loader = create_ajax_loader(
248+
model_admin=AddressAdmin(),
249+
name="user",
250+
options={"fields": ("name",)},
251+
)
252+
253+
Form = await get_model_form(
254+
model=Address,
255+
session_maker=session_maker,
256+
form_overrides={"user": LoaderAwareField},
257+
form_ajax_refs={"user": loader},
258+
)
259+
260+
assert Form()._fields["user"].loader is loader
261+
262+
236263
@pytest.mark.skipif(engine.name != "postgresql", reason="PostgreSQL only")
237264
async def test_model_form_postgresql() -> None:
238265
class PostgresModel(Base):

0 commit comments

Comments
 (0)