from collections import namedtuple import json import typing import sqlalchemy from wtforms import fields from wtforms.validators import ValidationError from wtforms_alchemy.fields import ( QuerySelectField, QuerySelectMultipleField, ) from marshmallow.fields import Integer, String from atlas_um import pgdb class ClaimValueSelectField(QuerySelectField): def __init__(self, claim_name: pgdb.ClaimName, **kwargs): self.claim_name = claim_name if claim_name.claim_values_source_id: values = ( pgdb.ClaimValue.query.active() .filter( pgdb.ClaimValue.claim_name == claim_name.claim_values_source ) .order_by(pgdb.ClaimValue.friendly) .all ) else: values = ( pgdb.ClaimValue.query.active() .filter(pgdb.ClaimValue.claim_name == claim_name) .order_by(pgdb.ClaimValue.friendly) .all ) super().__init__( label=claim_name.friendly, query_factory=values, validators=None, get_pk=None, get_label=None, allow_blank=True, blank_text="N/A", **kwargs, ) def pre_validate(self, form): super().pre_validate(form) if not self.claim_name.optional and self.data is None: raise ValidationError("This field is required") class ClaimValueSelectMultipleField(QuerySelectMultipleField): def __init__(self, claim_name: pgdb.ClaimName, **kwargs): self.claim_name = claim_name if claim_name.claim_values_source_id: values = ( pgdb.ClaimValue.query.active() .filter( pgdb.ClaimValue.claim_name == claim_name.claim_values_source ) .order_by(pgdb.ClaimValue.friendly) .all ) else: values = ( pgdb.ClaimValue.query.active() .filter(pgdb.ClaimValue.claim_name == claim_name) .order_by(pgdb.ClaimValue.friendly) .all ) self._tree = pgdb.ClaimValue.query.tree_by_claim_name(self.claim_name) super().__init__( label=claim_name.friendly, query_factory=values, validators=None, default=None, allow_blank=True, **kwargs, ) def pre_validate(self, form): super().pre_validate(form) if not self.claim_name.optional and len(self.data) == 0: raise ValidationError("This field is required") @property def choices_tree(self) -> dict: Choice = namedtuple("Choice", "pk label is_selected") tree = {} for parent, children in self._tree.items(): values = [] for obj in children: values.append( Choice(obj.id, self.get_label(obj), obj in self.data) ) parent_is_selected = False if all(v.is_selected for v in values): parent_is_selected = True elif any(v.is_selected for v in values): parent_is_selected = None if parent: key = Choice( parent.id, self.get_label(parent), parent_is_selected ) else: key = Choice(None, "N/A", parent_is_selected) tree[key] = values return tree class JSONField(fields.StringField): def _value(self): return json.dumps(self.data) if self.data else "" def process_formdata(self, valuelist): if valuelist: try: self.data = json.loads(valuelist[0]) except ValueError: raise ValidationError("Field contains invalid JSON") else: self.data = None def pre_validate(self, form): super().pre_validate(form) if self.data: try: json.dumps(self.data) except TypeError: raise ValidationError("Field contains invalid JSON") class IntegerNoneField(Integer): """ Allows to pass validation with empty value. Similar to WTForms validation logic as used in other WTForms selects in app. """ def _validated(self, value): if value in ("", "__None"): return None return super()._validated(value) class KeyValueField(String): """Expecting key:value string, resulting in key and value tuple.""" default_error_messages = { **String.default_error_messages, "invalid_content": "Not a valid content", } def _serialize(self, value, attr, obj, **kwargs) -> typing.Optional[str]: if value is None: return None if not isinstance(value, tuple) or len(value) != 2: return None return ":".join(value) def _deserialize(self, value, attr, data, **kwargs) -> typing.Any: if not isinstance(value, (str, bytes)): raise self.make_error("invalid") try: key, value = value.split(":") except Exception: raise self.make_error("invalid_content") return key, value class TimezoneField(String): default_error_messages = {"invalid": "Not a valid timezone."} def _deserialize(self, value, attr, data, **kwargs) -> typing.Any: if not value: raise self.make_error("invalid") if value and value.lower == "utc": return value if not pgdb.pgdb.session.execute( sqlalchemy.text( "SELECT name FROM pg_timezone_names WHERE name = :tz_name " ), {"tz_name": value}, ).first(): raise self.make_error("invalid") return value