diff --git a/src/pretix/base/models/orders.py b/src/pretix/base/models/orders.py index fa33104ec8..53c1cb494c 100644 --- a/src/pretix/base/models/orders.py +++ b/src/pretix/base/models/orders.py @@ -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): """ diff --git a/src/pretix/base/storelogic/__init__.py b/src/pretix/base/storelogic/__init__.py index e69de29bb2..9b91f3eeca 100644 --- a/src/pretix/base/storelogic/__init__.py +++ b/src/pretix/base/storelogic/__init__.py @@ -0,0 +1,2 @@ +class IncompleteError(Exception): + pass diff --git a/src/pretix/base/storelogic/fields.py b/src/pretix/base/storelogic/fields.py new file mode 100644 index 0000000000..48a3908137 --- /dev/null +++ b/src/pretix/base/storelogic/fields.py @@ -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.')) diff --git a/src/pretix/presale/checkoutflow.py b/src/pretix/presale/checkoutflow.py index b3e47d10bb..8af745fa41 100644 --- a/src/pretix/presale/checkoutflow.py +++ b/src/pretix/presale/checkoutflow.py @@ -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) diff --git a/src/pretix/presale/views/__init__.py b/src/pretix/presale/views/__init__.py index e5dd3de071..8d067ebf82 100644 --- a/src/pretix/presale/views/__init__.py +++ b/src/pretix/presale/views/__init__.py @@ -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', diff --git a/src/pretix/storefrontapi/endpoints/checkout.py b/src/pretix/storefrontapi/endpoints/checkout.py index 5a0a36812c..5d1efffabb 100644 --- a/src/pretix/storefrontapi/endpoints/checkout.py +++ b/src/pretix/storefrontapi/endpoints/checkout.py @@ -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( diff --git a/src/pretix/storefrontapi/middleware.py b/src/pretix/storefrontapi/middleware.py index c6a8cf9b4e..0f70f492f3 100644 --- a/src/pretix/storefrontapi/middleware.py +++ b/src/pretix/storefrontapi/middleware.py @@ -137,6 +137,7 @@ class ApiMiddleware: "OPTIONS", "PUT", "DELETE", + "PATCH", ] ) r["Access-Control-Allow-Headers"] = ", ".join( diff --git a/src/pretix/storefrontapi/steps.py b/src/pretix/storefrontapi/steps.py index dfa679b350..4c5c362292 100644 --- a/src/pretix/storefrontapi/steps.py +++ b/src/pretix/storefrontapi/steps.py @@ -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