Source code for api.form

from __future__ import annotations

import humanize

from abc import abstractmethod, ABC
from datetime import date, time
from dateutil.relativedelta import relativedelta
from decimal import Decimal
from onegov.core.utils import binary_to_dictionary
from onegov.form.fields import HoneyPotField
from onegov.form.validators import FileSizeLimit
from onegov.form.validators import If
from onegov.form.validators import Stdnum
from onegov.form.validators import ValidDateRange
from onegov.form.validators import WhitelistedMimeType
from pydantic import create_model, model_validator
from pydantic import AfterValidator, BaseModel, Field
from pydantic import AwareDatetime, Base64Bytes, EmailStr, HttpUrl
from wtforms import HiddenField
from wtforms.validators import DataRequired, InputRequired, Optional
from wtforms.validators import Email, Length, NumberRange, Regexp, URL
from wtforms.validators import HostnameValidation


from typing import Annotated, Any, Literal, TypeVar, TYPE_CHECKING
if TYPE_CHECKING:
    from collections.abc import Callable, Generator, Sequence
    from onegov.core.types import FileDict
    from onegov.form import Form
    from onegov.form.core import FieldDependency
    from onegov.form.types import Validator
    from typing_extensions import TypeForm
    from wtforms import Field as WTField


[docs] type ValidatorAdapter[T: Validator[Any, Any]] = Callable[[T], Any]
[docs] _A = TypeVar('_A', bound='BaseAdapter')
[docs] _V = TypeVar('_V', bound='ValidatorAdapter[Any]')
[docs] class AdapterRegistry: """ Keeps track of all the adapters and WTForms field types they are registered for, making sure each adapter is only instantiated once. """
[docs] adapter_map: dict[str, BaseAdapter]
[docs] validator_map: dict[type[Any], ValidatorAdapter[Any]]
def __init__(self) -> None: self.adapter_map = {} self.validator_map = {}
[docs] def register_for(self, *types: str) -> Callable[[type[_A]], type[_A]]: """ Decorator to register a field adapter. """ def wrapper(adapter: type[_A]) -> type[_A]: instance = adapter() for type in types: assert type not in self.adapter_map self.adapter_map[type] = instance return adapter return wrapper
[docs] def validator(self, validator: type[Any]) -> Callable[[_V], _V]: """ Decorator to register a validator adapter. """ def wrapper(adapter: _V) -> _V: assert validator not in self.validator_map self.validator_map[validator] = adapter return adapter return wrapper
[docs] def adapt(self, field: WTField) -> Generator[Any]: """ Adapts the WTForms field to a pydantic field yielding a sequence of values starting with a type form, followed by any number of pydantic metadata objects. This output will get unpacked into `Annotated`. """ adapter = self.adapter_map[field.type] return adapter(field)
[docs] def adapt_validator(self, validator: Validator[Any, Any]) -> Any: adapter = self.validator_map[validator.__class__] return adapter(validator)
[docs] registry = AdapterRegistry()
@registry.validator(Length)
[docs] def adapt_length(validator: Length) -> Any: return Field( min_length=None if validator.min < 0 else validator.min, max_length=None if validator.max < 0 else validator.max )
@registry.validator(Regexp)
[docs] def adapt_regexp(validator: Regexp) -> Any: return Field(pattern=validator.regex)
@registry.validator(NumberRange)
[docs] def adapt_number_range(validator: NumberRange) -> Any: return Field(ge=validator.min, le=validator.max)
@registry.validator(ValidDateRange)
[docs] def adapt_valid_date_range(validator: ValidDateRange) -> Any: ge = validator.min if isinstance(ge, relativedelta): ge = date.today() + ge # NOTE: In order to get the correct behavior for datetimes # we convert to an exclusive end # FIXME: Will automatic coercion work for datetime, even # though the datetimes will be timezone aware? If # not we may need to emit a custom AfterValidator instead lt = validator.max if isinstance(lt, relativedelta): lt = date.today() + lt if lt is not None: lt += relativedelta(days=1) return Field(ge=ge, lt=lt)
@registry.validator(Stdnum)
[docs] def adapt_stdnum(validator: Stdnum) -> Any: def validate_stdnum(value: str | None) -> str | None: if value is None: return None validator.format.validate(value) return value return AfterValidator(validate_stdnum)
@registry.validator(FileSizeLimit)
[docs] def adapt_file_size_limit(validator: FileSizeLimit) -> Any: def validate_file_size(value: FileDict) -> FileDict: if not value: return value # type: ignore[unreachable] if value.get('size', 0) > validator.max_bytes: raise ValueError(str(validator.message).format( humanize.naturalsize(validator.max_bytes) )) return value return AfterValidator(validate_file_size)
[docs] REQUIRED_DEPENDENT = object()
[docs] class BaseAdapter(ABC): """ Provides utility functions for all adapters. """
[docs] def handle_scalar_field_type( self, t: TypeForm[Any], field: WTField ) -> Generator[Any]: is_dependent = hasattr(field, 'depends_on') if not is_dependent and any( isinstance(validator, (InputRequired, DataRequired)) for validator in field.validators ): yield t # NOTE: This is a bit of a hack to avoid emitting an # Annotated without metadata yield object() return default = field.data if default is None: yield t | None yield Field(default=None) else: yield t if isinstance(default, (list, dict)): yield Field(default_factory=lambda: default.copy()) else: yield Field(default=default) if ( is_dependent and field.validators and isinstance(field.validators[0], If) and any( isinstance(validator, (InputRequired, DataRequired)) for validator in field.validators[0].validators ) ): # we use this marker in the model validator to detect # fields that become required when their dependency # is fulfilled yield REQUIRED_DEPENDENT
[docs] def handle_sequence_field_type( self, t: TypeForm[Any], field: WTField ) -> Generator[Any]: yield list[t] # type: ignore[valid-type] is_dependent = hasattr(field, 'depends_on') upload_required = getattr(field, 'upload_required', False) if not is_dependent and (upload_required or any( isinstance(validator, (InputRequired, DataRequired)) for validator in field.validators )): yield Field(min_length=1) return if default := field.data: yield Field(default_factory=lambda: default.copy()) else: yield Field(default_factory=list) if is_dependent and (upload_required or ( ((validators := field.validators) or ( hasattr(field, 'unbound_field') and (validators := field.unbound_field.kwargs.get( 'validators' )) )) and isinstance(validators[0], If) and any( isinstance(validator, (InputRequired, DataRequired)) for validator in validators[0].validators ) )): # we use this marker in the model validator to detect # fields that become required when their dependency # is fulfilled yield REQUIRED_DEPENDENT
[docs] def adapt_validators( self, validators: Sequence[Validator[Any, Any]] | None ) -> Generator[Any]: if not validators: return if isinstance(validators[0], If): validators = validators[0].validators for validator in validators: if isinstance(validator, ( # already handled via maybe_optional Optional, InputRequired, DataRequired, # already special-cased in Email field adapter Email, # already special-cased in URL field adapter URL, # already special-cased in upload field adapter WhitelistedMimeType, )): continue yield registry.adapt_validator(validator)
@abstractmethod
[docs] def __call__(self, field: WTField) -> Generator[Any]: raise NotImplementedError
@registry.register_for( 'StringField', 'TextAreaField', 'PasswordField' )
[docs] class StringFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(str, field) yield from self.adapt_validators(field.validators)
@registry.register_for('EmailField')
[docs] class EmailFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(EmailStr, field) yield from self.adapt_validators(field.validators)
[docs] validate_hostname = HostnameValidation(require_tld=True, allow_ip=True)
[docs] def validate_and_coerce_url(value: HttpUrl | None) -> str | None: if value is None: return None if value.host is None: # pragma: no cover # NOTE: I don't think it's possible for this to happen with HttpUrl # since the scheme is required, so relative urls don't work # despite the validation error claiming it does... raise ValueError('hostname is required') if not validate_hostname(value.host): raise ValueError('hostname is not valid') # NOTE: Coerce from HttpUrl back to str, since that's what we will # store in our database return str(value)
@registry.register_for( 'URLField', 'VideoURLField' )
[docs] class URLFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: # NOTE: HttpUrl normalizes the URL to some degree, which doesn't # happen in the actual form, maybe we should use `str` # type and rely on the wtforms URL validator instead? yield from self.handle_scalar_field_type(HttpUrl, field) yield AfterValidator(validate_and_coerce_url) yield from self.adapt_validators(field.validators)
@registry.register_for('DateField')
[docs] class DateFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(date, field) yield from self.adapt_validators(field.validators)
@registry.register_for( 'DateTimeLocalField', 'TimezoneDateTimeField' )
[docs] class DateTimeFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(AwareDatetime, field) yield from self.adapt_validators(field.validators)
@registry.register_for('TimeField')
[docs] class TimeFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(time, field) yield from self.adapt_validators(field.validators)
@registry.register_for('DecimalField')
[docs] class DecimalFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(Decimal, field) yield from self.adapt_validators(field.validators)
@registry.register_for('IntegerField')
[docs] class IntegerFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(int, field) yield from self.adapt_validators(field.validators)
@registry.register_for('RadioField')
[docs] class RadioFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(Literal[*( # type: ignore[arg-type] value for value, _ in field.choices # type: ignore[attr-defined] if value )], field) yield from self.adapt_validators(field.validators)
@registry.register_for('MultiCheckboxField')
[docs] class MultiCheckboxFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_sequence_field_type(Literal[*( # type: ignore[arg-type] value for value, _ in field.choices # type: ignore[attr-defined] if value )], field) yield from self.adapt_validators(field.validators)
[docs] class FileUpload(BaseModel):
[docs] filename: str
[docs] data: Base64Bytes
[docs] def mimetypes_validator( mimetypes: set[str] ) -> Callable[[FileDict], FileDict]: def validate_mimetype(value: FileDict) -> FileDict: if not value: return value # type: ignore[unreachable] if value['mimetype'] not in mimetypes: raise ValueError( f'Unsupported mimetype. ' f'Allowed mimetypes are {", ".join(mimetypes)}.' ) return value return validate_mimetype
@registry.register_for('UploadField')
[docs] class UploadFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_scalar_field_type(FileUpload, field) # NOTE: Coerce from FileUpload to FileDict yield AfterValidator( lambda v: {} if v is None # type: ignore[typeddict-item] else binary_to_dictionary(v.data, v.filename) ) yield AfterValidator(mimetypes_validator(field.mimetypes)) # type: ignore[attr-defined] yield from self.adapt_validators(field.validators)
@registry.register_for('UploadMultipleField')
[docs] class UploadMultipleFieldAdapter(BaseAdapter):
[docs] def __call__(self, field: WTField) -> Generator[Any]: yield from self.handle_sequence_field_type( Annotated[ FileUpload, # NOTE: Coerce from FileUpload to FileDict AfterValidator( lambda v: {} if v is None else binary_to_dictionary(v.data, v.filename) ), AfterValidator(mimetypes_validator(field.mimetypes)), *self.adapt_validators(field.unbound_field.kwargs['validators']) ], field )
[docs] def dependency_fulfilled(self: FieldDependency, obj: object) -> bool: result = True for dependency in self.dependencies: data = getattr(obj, dependency['field_id']) choice = dependency['choice'] invert = dependency['invert'] if isinstance(data, bool) and choice in ('y', 'n'): choice = choice == 'y' and True or False if isinstance(data, list): value = choice in data else: value = data == choice result = result and (value ^ invert) return result
[docs] def model_from_form(form: Form) -> type[BaseModel]: validators: dict[str, Any] = {} dependent_fields = { name: field.depends_on for name, field in form._fields.items() if hasattr(field, 'depends_on') } if dependent_fields: @model_validator(mode='after') def validate_required_dependent_fields(self: Any) -> Any: model_fields = type(self).model_fields for name, depends_on in dependent_fields.items(): if REQUIRED_DEPENDENT not in model_fields[name].metadata: continue # if the value is something that will satisfy LaxDataRequired # then we accept it regardless of whether the dependency is # fulfilled value = getattr(self, name) if value is False: # NOTE: we need to special-case this since bool is # an instance of int pass elif isinstance(value, (int, float, Decimal)): continue if isinstance(value, str) and value.strip(): continue if dependency_fulfilled(depends_on, self): raise ValueError( f'{name} is required due to your other submitted data' ) return self validators[ 'validate_required_dependent_fields' ] = validate_required_dependent_fields return create_model( f'{form.__class__.__name__}Model', __base__=None, __module__=form.__class__.__module__, __qualname__=None, __doc__=None, __config__={'frozen': True}, __validators__=validators, __cls_kwargs__=None, **{ name: Annotated[*registry.adapt(field)] for name, field in form._fields.items() if not isinstance(field, (HiddenField, HoneyPotField)) } )