Add fields

This commit is contained in:
Raphael Michel
2025-01-04 17:18:37 +01:00
parent 3b664f8b76
commit 8f13c03245
8 changed files with 429 additions and 101 deletions
+26 -1
View File
@@ -55,7 +55,7 @@ from django.conf import settings
from django.core.exceptions import ValidationError
from django.db import models, transaction
from django.db.models import (
Case, Exists, F, Max, OuterRef, Q, Subquery, Sum, Value, When,
Case, Exists, F, Max, OuterRef, Prefetch, Q, Subquery, Sum, Value, When,
)
from django.db.models.functions import Coalesce, Greatest
from django.db.models.signals import post_delete
@@ -3095,6 +3095,31 @@ class CheckoutSession(models.Model):
testmode = models.BooleanField(default=False)
session_data = models.JSONField(default=dict)
def get_cart_positions(self, prefetch_questions=False):
qs = CartPosition.objects.filter(event=self.event, cart_id=self.cart_id).select_related(
"item", "variation", "subevent",
)
if prefetch_questions:
qqs = self.event.questions.filter(ask_during_checkin=False, hidden=False)
qs = qs.prefetch_related(
Prefetch("answers",
QuestionAnswer.objects.prefetch_related("options"),
to_attr="answerlist"),
Prefetch("item__questions",
qqs.prefetch_related(
Prefetch("options", QuestionOption.objects.prefetch_related(Prefetch(
# This prefetch statement is utter bullshit, but it actually prevents Django from doing
# a lot of queries since ModelChoiceIterator stops trying to be clever once we have
# a prefetch lookup on this query...
"question",
Question.objects.none(),
to_attr="dummy"
)))
).select_related("dependency_question"),
to_attr="questions_to_ask")
)
return qs
class CartPosition(AbstractPosition):
"""
+2
View File
@@ -0,0 +1,2 @@
class IncompleteError(Exception):
pass
+271
View File
@@ -0,0 +1,271 @@
from django.core.exceptions import ValidationError
from django.core.validators import EmailValidator
from django.utils.translation import gettext_lazy as _
from pretix.base.models import CartPosition, Question
from pretix.base.services.checkin import _save_answers
from pretix.base.storelogic import IncompleteError
from pretix.presale.signals import question_form_fields
class Field:
@property
def identifier(self):
raise NotImplementedError()
@property
def label(self):
raise NotImplementedError()
@property
def help_text(self):
raise NotImplementedError()
@property
def type(self):
raise NotImplementedError()
@property
def required(self):
return True
@property
def validation_hints(self):
raise {}
def validate_input(self, value):
return value
class PositionField(Field):
def save_input(self, position, value):
raise NotImplementedError()
def current_value(self, position):
raise NotImplementedError()
class SessionField(Field):
def save_input(self, session_data, value):
raise NotImplementedError()
def current_value(self, session_data):
raise NotImplementedError()
class QuestionField(PositionField):
def __init__(self, question: Question):
self.question = question
@property
def label(self):
return self.question.question
@property
def help_text(self):
return self.question.help_text
@property
def type(self):
return self.question.type
@property
def identifier(self):
return f"question_{self.question.identifier}"
def validate_input(self, value):
return self.question.clean_answer(value)
def required(self, value):
return self.question.required
def validation_hints(self):
d = {
"valid_number_min": self.question.valid_number_min,
"valid_number_max": self.question.valid_number_max,
"valid_date_min": self.question.valid_date_min,
"valid_date_max": self.question.valid_date_max,
"valid_datetime_min": self.question.valid_datetime_min,
"valid_datetime_max": self.question.valid_datetime_max,
"valid_string_length_max": self.question.valid_string_length_max,
"dependency_on": f"question_{self.question.dependency_question.identifier}" if self.question.dependency_question_id else None,
"dependency_values": self.question.dependency_values,
}
if self.question.type in (Question.TYPE_CHOICE, Question.TYPE_CHOICE_MULTIPLE):
d["choices"] = [
{
"identifier": opt.identifier,
"label": str(opt.answer)
}
for opt in self.question.options.all()
]
return d
def save_input(self, position, value):
answers = [a for a in position.answerlist if a.question_id == self.question.id]
if answers:
answers = {self.question: answers[0]}
else:
answers = {}
_save_answers(position, answers, {self.question: value})
def current_value(self, position):
answers = [a for a in position.answerlist if a.question_id == self.question.id]
if answers:
if self.question.type in (Question.TYPE_CHOICE, Question.TYPE_CHOICE_MULTIPLE):
return ",".join([a.idenitifer for a in answers[0].options.all()])
else:
return answers[0].answer
class SyntheticSessionField(SessionField):
def __init__(self, label, help_text, type, identifier, required, save_func, get_func, validate_func):
self._label = label
self._help_text = help_text
self._type = type
self._identifier = identifier
self._required = required
self._save_func = save_func
self._get_func = get_func
self._validate_func = validate_func
super().__init__()
@property
def label(self):
return self._label
@property
def help_text(self):
return self._help_text
@property
def type(self):
return self._type
@property
def required(self):
return self._required
@property
def identifier(self):
return self._identifier
def validation_hints(self):
return {}
def save_input(self, session_data, value):
self._save_func(session_data, value)
def current_value(self, session_data):
return self._get_func(session_data)
def validate_input(self, value):
return self._validate_func(value)
def get_checkout_fields(event):
fields = []
# TODO: support contact_form_fields
# TODO: support contact_form_fields_override
# email
fields.append(SyntheticSessionField(
label=_("Email"),
help_text=None,
type=Question.TYPE_STRING, # TODO: Add a type?
identifier="email",
required=True,
get_func=lambda session_data: session_data.get("email"),
save_func=lambda session_data, value: session_data.update({"email": value}),
validate_func=lambda value: EmailValidator()(value) or value,
))
# TODO: phone
# TODO: invoice address
return fields
def get_position_fields(event, pos: CartPosition):
# TODO: support override sets
fields = []
for q in pos.item.questions_to_ask:
fields.append(QuestionField(q))
return fields
def ensure_fields_are_completed(event, positions, cart_session, invoice_address, all_optional, cart_is_free):
try:
emailval = EmailValidator()
if not cart_session.get('email') and not all_optional:
raise IncompleteError(_('Please enter a valid email address.'))
if cart_session.get('email'):
emailval(cart_session.get('email'))
except ValidationError:
raise IncompleteError(_('Please enter a valid email address.'))
address_asked = (
event.settings.invoice_address_asked and (not event.settings.invoice_address_not_asked_free or not cart_is_free)
)
if not all_optional:
if address_asked:
if event.settings.invoice_address_required and (not invoice_address or not invoice_address.street):
raise IncompleteError(_('Please enter your invoicing address.'))
if event.settings.invoice_name_required and (not invoice_address or not invoice_address.name):
raise IncompleteError(_('Please enter your name.'))
for cp in positions:
answ = {
aw.question_id: aw for aw in cp.answerlist
}
question_cache = {
q.pk: q for q in cp.item.questions_to_ask
}
def question_is_visible(parentid, qvals):
if parentid not in question_cache:
return False
parentq = question_cache[parentid]
if parentq.dependency_question_id and not question_is_visible(parentq.dependency_question_id,
parentq.dependency_values):
return False
if parentid not in answ:
return False
return (
('True' in qvals and answ[parentid].answer == 'True')
or ('False' in qvals and answ[parentid].answer == 'False')
or (any(qval in [o.identifier for o in answ[parentid].options.all()] for qval in qvals))
)
def question_is_required(q):
return (
q.required and
(not q.dependency_question_id or question_is_visible(q.dependency_question_id, q.dependency_values))
)
if not all_optional:
for q in cp.item.questions_to_ask:
if question_is_required(q) and q.id not in answ:
raise IncompleteError(_('Please fill in answers to all required questions.'))
if cp.item.ask_attendee_data and event.settings.get('attendee_names_required', as_type=bool) \
and not cp.attendee_name_parts:
raise IncompleteError(_('Please fill in answers to all required questions.'))
if cp.item.ask_attendee_data and event.settings.get('attendee_emails_required', as_type=bool) \
and cp.attendee_email is None:
raise IncompleteError(_('Please fill in answers to all required questions.'))
if cp.item.ask_attendee_data and event.settings.get('attendee_company_required', as_type=bool) \
and cp.company is None:
raise IncompleteError(_('Please fill in answers to all required questions.'))
if cp.item.ask_attendee_data and event.settings.get('attendee_addresses_required', as_type=bool) \
and (cp.street is None and cp.city is None and cp.country is None):
raise IncompleteError(_('Please fill in answers to all required questions.'))
responses = question_form_fields.send(sender=event, position=cp)
form_data = cp.meta_info_data.get('question_form_data', {})
for r, response in sorted(responses, key=lambda r: str(r[0])):
for key, value in response.items():
if value.required and not form_data.get(key):
raise IncompleteError(_('Please fill in answers to all required questions.'))
+15 -86
View File
@@ -41,7 +41,6 @@ from django.contrib import messages
from django.core.cache import caches
from django.core.exceptions import ImproperlyConfigured, ValidationError
from django.core.signing import BadSignature, loads
from django.core.validators import EmailValidator
from django.db import models
from django.db.models import Count, F, Q, Sum
from django.db.models.functions import Cast
@@ -71,9 +70,11 @@ from pretix.base.services.orders import perform_order
from pretix.base.services.tasks import EventTask
from pretix.base.settings import PERSON_NAME_SCHEMES
from pretix.base.signals import validate_cart_addons
from pretix.base.storelogic import IncompleteError
from pretix.base.storelogic.addons import (
addons_is_applicable, addons_is_completed, get_addon_groups,
)
from pretix.base.storelogic.fields import ensure_fields_are_completed
from pretix.base.templatetags.money import money_filter
from pretix.base.templatetags.phone_format import phone_format
from pretix.base.templatetags.rich_text import rich_text_snippet
@@ -89,7 +90,7 @@ from pretix.presale.forms.customer import AuthenticationForm, RegistrationForm
from pretix.presale.signals import (
checkout_all_optional, checkout_confirm_messages, checkout_flow_steps,
contact_form_fields, contact_form_fields_overrides,
order_api_meta_from_request, order_meta_from_request, question_form_fields,
order_api_meta_from_request, order_meta_from_request,
question_form_fields_overrides,
)
from pretix.presale.utils import customer_login
@@ -907,93 +908,21 @@ class QuestionsStep(QuestionsViewMixin, CartMixin, TemplateFlowStep):
return redirect_to_url(self.get_next_url(request))
def is_completed(self, request, warn=False):
self.request = request
try:
emailval = EmailValidator()
if not self.cart_session.get('email') and not self.all_optional:
if warn:
messages.warning(request, _('Please enter a valid email address.'))
return False
if self.cart_session.get('email'):
emailval(self.cart_session.get('email'))
except ValidationError:
ensure_fields_are_completed(
self.request.event,
self._positions_for_questions,
self.cart_session,
self.invoice_address,
self.all_optional,
get_cart_is_free(self.request),
)
except IncompleteError as e:
if warn:
messages.warning(request, _('Please enter a valid email address.'))
messages.warning(request, e)
return False
if not self.all_optional:
if self.address_asked:
if request.event.settings.invoice_address_required and (not self.invoice_address or not self.invoice_address.street):
messages.warning(request, _('Please enter your invoicing address.'))
return False
if request.event.settings.invoice_name_required and (not self.invoice_address or not self.invoice_address.name):
messages.warning(request, _('Please enter your name.'))
return False
for cp in self._positions_for_questions:
answ = {
aw.question_id: aw for aw in cp.answerlist
}
question_cache = {
q.pk: q for q in cp.item.questions_to_ask
}
def question_is_visible(parentid, qvals):
if parentid not in question_cache:
return False
parentq = question_cache[parentid]
if parentq.dependency_question_id and not question_is_visible(parentq.dependency_question_id, parentq.dependency_values):
return False
if parentid not in answ:
return False
return (
('True' in qvals and answ[parentid].answer == 'True')
or ('False' in qvals and answ[parentid].answer == 'False')
or (any(qval in [o.identifier for o in answ[parentid].options.all()] for qval in qvals))
)
def question_is_required(q):
return (
q.required and
(not q.dependency_question_id or question_is_visible(q.dependency_question_id, q.dependency_values))
)
if not self.all_optional:
for q in cp.item.questions_to_ask:
if question_is_required(q) and q.id not in answ:
if warn:
messages.warning(request, _('Please fill in answers to all required questions.'))
return False
if cp.item.ask_attendee_data and self.request.event.settings.get('attendee_names_required', as_type=bool) \
and not cp.attendee_name_parts:
if warn:
messages.warning(request, _('Please fill in answers to all required questions.'))
return False
if cp.item.ask_attendee_data and self.request.event.settings.get('attendee_emails_required', as_type=bool) \
and cp.attendee_email is None:
if warn:
messages.warning(request, _('Please fill in answers to all required questions.'))
return False
if cp.item.ask_attendee_data and self.request.event.settings.get('attendee_company_required', as_type=bool) \
and cp.company is None:
if warn:
messages.warning(request, _('Please fill in answers to all required questions.'))
return False
if cp.item.ask_attendee_data and self.request.event.settings.get('attendee_addresses_required', as_type=bool) \
and (cp.street is None and cp.city is None and cp.country is None):
if warn:
messages.warning(request, _('Please fill in answers to all required questions.'))
return False
responses = question_form_fields.send(sender=self.request.event, position=cp)
form_data = cp.meta_info_data.get('question_form_data', {})
for r, response in sorted(responses, key=lambda r: str(r[0])):
for key, value in response.items():
if value.required and not form_data.get(key):
return False
return True
else:
return True
def get_context_data(self, **kwargs):
ctx = super().get_context_data(**kwargs)
+1 -2
View File
@@ -339,8 +339,7 @@ def get_cart(request):
'item__category__position', 'item__category_id', 'item__position', 'item__name', 'variation__value'
).select_related(
'item', 'variation', 'subevent', 'subevent__event', 'subevent__event__organizer',
'item__tax_rule', 'item__category', 'used_membership', 'used_membership__membership_type'
).select_related(
'item__tax_rule', 'item__category', 'used_membership', 'used_membership__membership_type',
'addon_to'
).prefetch_related(
'addons', 'addons__item', 'addons__variation',
+80 -8
View File
@@ -16,6 +16,9 @@ from pretix.base.services.cart import (
add_items_to_cart, error_messages, get_fees, set_cart_addons,
)
from pretix.base.storelogic.addons import get_addon_groups
from pretix.base.storelogic.fields import (
get_checkout_fields, get_position_fields,
)
from pretix.base.timemachine import time_machine_now
from pretix.presale.views.cart import generate_cart_id
from pretix.storefrontapi.endpoints.event import (
@@ -103,7 +106,15 @@ class CartFeeSerializer(serializers.ModelSerializer):
]
class CartPositionSerializer(serializers.ModelSerializer):
class FieldSerializer(serializers.Serializer):
identifier = serializers.CharField()
label = serializers.CharField(allow_null=True)
required = serializers.BooleanField()
type = serializers.CharField()
validation_hints = serializers.DictField()
class MinimalCartPositionSerializer(serializers.ModelSerializer):
# todo: prefetch related items
item = InlineItemSerializer(read_only=True)
variation = InlineItemVariationSerializer(read_only=True)
@@ -124,6 +135,21 @@ class CartPositionSerializer(serializers.ModelSerializer):
]
class CartPositionSerializer(MinimalCartPositionSerializer):
def to_representation(self, instance):
d = super().to_representation(instance)
fields = get_position_fields(self.context["event"], instance)
d["fields"] = FieldSerializer(
fields,
many=True,
context={**self.context, "position": instance}
).data
d["fields_data"] = {
f.identifier: f.current_value(instance) for f in fields
}
return d
class CheckoutSessionSerializer(serializers.ModelSerializer):
class Meta:
@@ -137,9 +163,7 @@ class CheckoutSessionSerializer(serializers.ModelSerializer):
def to_representation(self, checkout):
d = super().to_representation(checkout)
cartpos = CartPosition.objects.filter(
event_id=self.context["event"], cart_id=checkout.cart_id
)
cartpos = checkout.get_cart_positions(prefetch_questions=True)
total = sum(p.price for p in cartpos)
try:
@@ -161,12 +185,28 @@ class CheckoutSessionSerializer(serializers.ModelSerializer):
total += sum([f.value for f in fees])
d["cart_positions"] = CartPositionSerializer(
sorted(cartpos, key=lambda c: c.sort_key), many=True
sorted(cartpos, key=lambda c: c.sort_key), many=True, context=self.context
).data
d["cart_fees"] = CartFeeSerializer(fees, many=True).data
d["cart_fees"] = CartFeeSerializer(fees, many=True, context=self.context).data
d["total"] = str(total)
steps = get_steps(self.context["event"], cartpos)
fields = get_checkout_fields(self.context["event"])
d["fields"] = FieldSerializer(
fields,
many=True,
context={**self.context, "checkout": checkout}
).data
d["fields_data"] = {
f.identifier: f.current_value(checkout.session_data) for f in fields
}
steps = get_steps(
self.context["event"],
cartpos,
getattr(checkout, "invoice_address", None),
checkout.session_data,
total,
)
d["steps"] = {}
for step in steps:
applicable = step.is_applicable()
@@ -267,7 +307,7 @@ class CheckoutViewSet(viewsets.ViewSet):
elif request.method == "GET":
data = [
{
"parent": CartPositionSerializer(grp["pos"], context=ctx).data,
"parent": MinimalCartPositionSerializer(grp["pos"], context=ctx).data,
"categories": [
{
"category": CategorySerializer(
@@ -300,6 +340,38 @@ class CheckoutViewSet(viewsets.ViewSet):
status=200,
)
@action(detail=True, methods=["PATCH"])
def fields(self, request, *args, **kwargs):
cs = get_object_or_404(
self.request.event.checkout_sessions, cart_id=kwargs["cart_id"]
)
server_pos = {
p.pk: p for p in cs.get_cart_positions(prefetch_questions=True)
}
for req_pos in request.data.get("cart_positions", []):
pos = server_pos[req_pos["id"]]
fields = get_position_fields(self.request.event, pos)
fields_data = req_pos["fields_data"]
for f in fields:
if f.identifier in fields_data:
# todo: validation error handing
value = f.validate_input(fields_data[f.identifier])
f.save_input(pos, value)
fields = get_checkout_fields(self.request.event)
fields_data = request.data.get("fields_data", {})
session_data = cs.session_data
for f in fields:
if f.identifier in fields_data:
# todo: validation error handing
value = f.validate_input(fields_data[f.identifier])
f.save_input(session_data, value)
cs.session_data = session_data
cs.save(update_fields=["session_data"])
cs.refresh_from_db()
return self._return_checkout_status(cs, 200)
@action(detail=True, methods=["POST"])
def add_to_cart(self, request, *args, **kwargs):
cs = get_object_or_404(
+1
View File
@@ -137,6 +137,7 @@ class ApiMiddleware:
"OPTIONS",
"PUT",
"DELETE",
"PATCH",
]
)
r["Access-Control-Allow-Headers"] = ", ".join(
+33 -4
View File
@@ -1,12 +1,19 @@
from decimal import Decimal
from pretix.base.storelogic import IncompleteError
from pretix.base.storelogic.addons import (
addons_is_applicable, addons_is_completed,
)
from pretix.base.storelogic.fields import ensure_fields_are_completed
class CheckoutStep:
def __init__(self, event, cart_positions):
def __init__(self, event, cart_positions, invoice_address, cart_session, total):
self.event = event
self.cart_positions = cart_positions
self.cart_session = cart_session
self.invoice_address = invoice_address
self.total = total
@property
def identifier(self):
@@ -29,13 +36,35 @@ class AddonStep(CheckoutStep):
return addons_is_completed(self.cart_positions)
def get_steps(event, cart_positions):
class FieldsStep(CheckoutStep):
identifier = "fields"
def is_applicable(self):
return True
def is_valid(self):
try:
ensure_fields_are_completed(
self.event,
self.cart_positions,
self.cart_session,
self.invoice_address,
False,
cart_is_free=self.total == Decimal("0.00"),
)
except IncompleteError:
return False
else:
return True
def get_steps(event, cart_positions, invoice_address, cart_session, total):
return [
AddonStep(event, cart_positions),
AddonStep(event, cart_positions, invoice_address, cart_session, total),
FieldsStep(event, cart_positions, invoice_address, cart_session, total),
# todo: cross-selling
# todo: customers
# todo: memberships
# todo: questions
# todo: plugin signals
# todo: payment
# todo: confirmations