diff --git a/sparkmeter/tariff/templates/tariff-modal-form.html b/sparkmeter/tariff/templates/tariff-modal-form.html
new file mode 100644
index 0000000..d300e18
--- /dev/null
+++ b/sparkmeter/tariff/templates/tariff-modal-form.html
@@ -0,0 +1,6 @@
+{%- from "_macros.html" import form_errors -%}
+
+
diff --git a/sparkmeter/tariff/templates/tariff/_tariff-form-fields.html b/sparkmeter/tariff/templates/tariff/_tariff-form-fields.html
new file mode 100644
index 0000000..bcea382
--- /dev/null
+++ b/sparkmeter/tariff/templates/tariff/_tariff-form-fields.html
@@ -0,0 +1,217 @@
+{% macro form_tariff_error(form, field, class_=None) %}
+ {% if form.errors[field] %}
+
+
+
+ {% for error in form.errors[field] %}
+ {{ error }}
+ {% endfor %}
+
+
+ {% endif %}
+{% endmacro %}
+
+
+
+
+{{ form.hidden_tag() }}
+
+
+{{ form_tariff_error(form, 'name') }}
+
+
+{{ form_tariff_error(form, 'cycle_start_day_of_month', class_='form-group cycle-start-day-of-month') }}
+
+
+{{ form_tariff_error(form, 'load_limit_type') }}
+
+
+{% set extra = '' %}
+{% if form.data.load_limit_type != 'flat' %}
+{% set extra = ' hide' %}
+{% endif %}
+{{ form_tariff_error(form, 'flat_load_limit', class_='form-group load_limit_type' + extra) }}
+
+
+{% set extra = '' %}
+{% if form.data.load_limit_type != 'scheduled' %}
+{% set extra = ' hide' %}
+{% endif %}
+{{ form_tariff_error(form, 'load_limits', class_='form-group load-limits' + extra) }}
+
+
+{{ form_tariff_error(form, 'low_balance_threshold') }}
+
+
+{{ form_tariff_error(form, 'plan_enabled') }}
+
+{% set extra = '' %}
+{% if not form.data.plan_enabled %}
+{% set extra = ' hide' %}
+{% endif %}
+
+
+{{ form_tariff_error(form, 'plan_fixed_fee', class_='form-group plan-fixed-fee' + extra) }}
+
+
+{{ form_tariff_error(form, 'plan_price', class_='form-group plan-price' + extra) }}
+
+
+{{ form_tariff_error(form, 'tariff_type') }}
+
+
+{% set extra = '' %}
+{% if form.data.tariff_type != 'flat' %}
+{% set extra = ' hide' %}
+{% endif %}
+{{ form_tariff_error(form, 'flat_price', class_='form-group tariff_type' + extra) }}
+
+
+{{ form_tariff_error(form, 'blockrates') }}
+
+
+{{ form_tariff_error(form, 'tou_enabled') }}
+
+
+{% set extra = '' %}
+{% if not form.data.tou_enabled %}
+{% set extra = ' hide' %}
+{% endif %}
+{{ form_tariff_error(form, 'tous', class_='form-group tou' + extra) }}
+
+
+{{ form_tariff_error(form, 'daily_energy_limit_enabled') }}
+
+
+{% if form.data.daily_energy_limit_enabled -%}
+{{ form_tariff_error(form, 'daily_energy_limit_reset_hour', class_='form-group daily-energy-limit-reset-hour') }}
+{%- endif %}
+
+
+{% if form.data.daily_energy_limit_enabled -%}
+{{ form_tariff_error(form, 'daily_energy_limit_value', class_='form-group daily-energy-limit-reset-value') }}
+{%- endif %}
diff --git a/sparkmeter/tariff/tests/test_tariffviews.py b/sparkmeter/tariff/tests/test_tariffviews.py
index 828b9b0..d351070 100644
--- a/sparkmeter/tariff/tests/test_tariffviews.py
+++ b/sparkmeter/tariff/tests/test_tariffviews.py
@@ -13,10 +13,10 @@
from sparkmeter.event.eventdomain import Event
from sparkmeter.meter.meterdomain import MeterConfig
from sparkmeter.misc.htmlutils import build_link
-from sparkmeter.misc.jsonutils import json_dumps
+from sparkmeter.misc.jsonutils import json_dumps, json_loads
from sparkmeter.tariff.tariffdomain import Tariff
from sparkmeter.tests.base import WebViewTestCaseBase
-from sparkmeter.tests.test_data_factory import MeterFactory, TariffFactory
+from sparkmeter.tests.test_data_factory import MeterFactory, TariffFactory, VendorFactory
@pytest.fixture(scope="module", autouse=True)
@@ -55,6 +55,124 @@ def test_add(self, client):
response = client.get(path)
self.verify_response(response)
+ def test_add_modal_get(self, client):
+ path = "/tariff/add-modal"
+
+ response = client.get(path)
+ assert response.status_code == http.client.OK
+ assert "X-Form-Errors" not in response.headers
+ self.verify_response(response)
+
+ def test_add_modal_get_renders_empty_collections(self, client):
+ """A freshly opened modal must match /tariff/add, which posts back cleanly.
+
+ With no formdata the JSON fields keep ``None``, which renders as
+ ``null``/``""`` instead of ``[]`` and cannot be posted back.
+ """
+ response = client.get("/tariff/add-modal")
+ standalone = client.get("/tariff/add")
+
+ for attribute in ('data-blockrates="[]"', 'data-tous="[]"', 'data-load-limits="[]"'):
+ assert attribute in response.text
+ assert attribute in standalone.text
+
+ def test_add_modal_untouched_post_is_parseable(self, client):
+ """Submitting an untouched modal must not post unparseable collections."""
+ data = dict(name="", blockrates="[]", tous="[]", load_limits="[]")
+
+ response = client.post("/tariff/add-modal", data=data)
+
+ assert response.status_code == http.client.BAD_REQUEST
+ errors = json_loads(response.headers["X-Form-Errors"])
+ assert "blockrates" not in errors
+ assert "tous" not in errors
+ assert "load_limits" not in errors
+
+ def test_add_modal_forbidden_without_permission(self, client, vendor_role):
+ """The modal endpoint is behind the same tariff:add permission as /tariff/add."""
+ client.login_as(VendorFactory(roles=[vendor_role]))
+
+ assert client.get("/tariff/add-modal").status_code == http.client.NOT_FOUND
+ assert client.post("/tariff/add-modal", data={}).status_code == http.client.NOT_FOUND
+
+ def test_add_modal_post_valid(self, client, config):
+ path = "/tariff/add-modal"
+ data = dict(
+ name="MODAL TARIFF",
+ flat_load_limit=150,
+ plan_price=0,
+ cycle_start_day_of_month=1,
+ tariff_type="flat",
+ flat_price=4,
+ tous="",
+ )
+
+ config["HEROKU"] = False
+ response = client.post(path, data=data)
+
+ assert response.status_code == http.client.OK
+ body = response.json()
+ assert body["message"] == "Tariff created."
+
+ tariffs = Tariff.get_all()
+ assert len(tariffs) == 1
+ assert body["tariff"]["name"] == "MODAL TARIFF"
+ assert body["tariff"]["id"] == str(tariffs[0].id)
+
+ def test_add_modal_post_invalid(self, client):
+ path = "/tariff/add-modal"
+ data = dict(name="", flat_load_limit=150, flat_price=4)
+
+ response = client.post(path, data=data)
+
+ assert response.status_code == http.client.BAD_REQUEST
+ errors = json_loads(response.headers["X-Form-Errors"])
+ assert errors["name"] == ["Please set a name for this tariff"]
+ assert "Please set a name for this tariff" in response.text
+ assert not Tariff.query.scalar()
+
+ @pytest.mark.parametrize(
+ "field, data, message",
+ [
+ (
+ "blockrates",
+ dict(
+ tariff_type=Tariff.TYPE_BLOCKRATE,
+ blockrates=json_dumps([{"lower": "1", "upper": "20", "value": "1"}]),
+ ),
+ "Block rates contain at least one gap, between 0 and 65535",
+ ),
+ (
+ "tous",
+ dict(
+ tou_enabled=True,
+ tous=json_dumps([{"start": "00:00", "end": "12:00", "value": -100}]),
+ ),
+ "The TOU period modifier must be a positive number.",
+ ),
+ (
+ "load_limits",
+ dict(load_limit_type=Tariff.LOAD_LIMIT_TYPE_SCHEDULED, load_limits=json_dumps([])),
+ "Please add some Load limit periods.",
+ ),
+ ],
+ )
+ def test_add_modal_post_collection_error(self, client, field, data, message):
+ """These validators store the raw exception, which is not JSON serializable.
+
+ Reporting them used to raise a TypeError out of the error header and
+ turn the response into a 500.
+ """
+ path = "/tariff/add-modal"
+ data = dict(data, name="TARIFF", flat_load_limit=150, flat_price=4)
+
+ response = client.post(path, data=data)
+
+ assert response.status_code == http.client.BAD_REQUEST
+ errors = json_loads(response.headers["X-Form-Errors"])
+ assert errors[field] == [message]
+ assert not Tariff.query.scalar()
+
def test_add_form(self, client, config):
path = "/tariff/add"
diff --git a/sparkmeter/web/forms.py b/sparkmeter/web/forms.py
index 84b02cd..17144c3 100644
--- a/sparkmeter/web/forms.py
+++ b/sparkmeter/web/forms.py
@@ -39,6 +39,31 @@ def getall(self, key):
return [self[key]]
+def set_form_errors_header(response, form):
+ """Attach a form's validation errors to a response as ``X-Form-Errors``.
+
+ This is for development, and especially so that unittests can show a nicer
+ error when there is a form error.
+
+ Every error is coerced to ``str``: validators are free to store exception
+ instances rather than messages, and those are not JSON serializable.
+
+ :param response: the response to annotate.
+ :param form: the form whose errors should be reported.
+ :return: the same response, for convenience.
+ :rtype: Response
+ """
+ if not form.errors:
+ return response
+
+ error_dict = {}
+ for name, errors in list(form.errors.items()):
+ error_dict[name] = list(map(str, errors))
+ response.headers["X-Form-Errors"] = json_dumps(error_dict)
+ logger.warning("{} errors: {} {}".format(type(form).__name__, error_dict, form.data))
+ return response
+
+
class BaseForm(FlaskForm):
"""Base form, used by all other forms in the application."""
@@ -102,15 +127,7 @@ def render(self, **context):
:rtype: Response
"""
body = render_template(self.template_filename, form=self, **context)
- response = Response(body)
- if self.errors:
- error_dict = {}
- for name, errors in list(self.errors.items()):
- error_dict[name] = list(map(str, errors))
- response.headers["X-Form-Errors"] = json_dumps(error_dict)
- logger.warning("{} errors: {} {}".format(type(self).__name__, error_dict, self.data))
-
- return response
+ return set_form_errors_header(Response(body), self)
def flatten_json(form, json, parent_key="", separator="-", skip_unknown_keys=True): # pragma: nocoverage
diff --git a/test-data/meter/test_meterviews-MeterViewTest.test_add.page b/test-data/meter/test_meterviews-MeterViewTest.test_add.page
index 4140e36..abb8f61 100644
--- a/test-data/meter/test_meterviews-MeterViewTest.test_add.page
+++ b/test-data/meter/test_meterviews-MeterViewTest.test_add.page
@@ -50,7 +50,7 @@
-
+