hypothesis_mongoengine.py 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207
  1. # -*- coding: utf-8 -*-
  2. # !/usr/bin/env python
  3. """该代码从 hypothesis.extra.django.models 演变而来
  4. 目的为了直接方便的生成一系列需要的数据模型
  5. 当测试需要数据的时候可直接操作,生成的数据保持与模型的对应
  6. 例
  7. models(Device).example()
  8. """
  9. from __future__ import division, print_function, absolute_import
  10. import string
  11. from decimal import Decimal
  12. from typing import Union
  13. import django.db.models as dm
  14. from django.db import IntegrityError
  15. from django.conf import settings as django_settings
  16. from django.core.exceptions import ValidationError
  17. import mongoengine as mg
  18. import mongoengine.fields as mg_fields
  19. import hypothesis.strategies as st
  20. from hypothesis.errors import InvalidArgument
  21. from hypothesis.extra.pytz import timezones
  22. from hypothesis.utils.conventions import UniqueIdentifier
  23. from hypothesis.searchstrategy.strategies import SearchStrategy
  24. class ModelNotSupported(Exception):
  25. pass
  26. def referenced_models(model, seen=None):
  27. if seen is None:
  28. seen = set()
  29. for f in model._meta.concrete_fields:
  30. if isinstance(f, dm.ForeignKey):
  31. t = f.rel.to
  32. if t not in seen:
  33. seen.add(t)
  34. referenced_models(t, seen)
  35. return seen
  36. def get_datetime_strat():
  37. if getattr(django_settings, 'USE_TZ', False):
  38. return st.datetimes(timezones=timezones())
  39. return st.datetimes()
  40. __default_field_mappings = None
  41. def field_mappings():
  42. global __default_field_mappings
  43. if __default_field_mappings is None:
  44. __default_field_mappings = {
  45. mg.fields.IntField: st.integers(-2147483648, 2147483647),
  46. mg.fields.LongField:
  47. st.integers(-9223372036854775808, 9223372036854775807),
  48. mg.fields.BinaryField: st.binary(),
  49. mg.fields.BooleanField: st.booleans(),
  50. mg.fields.DateTimeField: get_datetime_strat(),
  51. mg.fields.FloatField: st.floats(),
  52. mg.fields.ListField: st.lists(),
  53. mg.fields.PointField: st.tuples(st.floats(), st.floats()),
  54. }
  55. return __default_field_mappings
  56. def add_default_field_mapping(field_type, strategy):
  57. field_mappings()[field_type] = strategy
  58. default_value = UniqueIdentifier(u'default_value')
  59. class UnmappedFieldError(Exception):
  60. pass
  61. def validator_to_filter(f):
  62. """Converts the field run_validators method to something suitable for use
  63. in filter."""
  64. def validate(value):
  65. try:
  66. f.run_validators(value)
  67. return True
  68. except ValidationError:
  69. return False
  70. return validate
  71. safe_letters = string.ascii_letters + string.digits + '_-'
  72. domains = st.builds(
  73. lambda x, y: '.'.join(x + [y]),
  74. st.lists(st.text(safe_letters, min_size=1), min_size=1), st.sampled_from([
  75. 'com', 'net', 'org', 'biz', 'info',
  76. ])
  77. )
  78. email_domains = st.one_of(
  79. domains,
  80. st.sampled_from(['gmail.com', 'yahoo.com', 'hotmail.com', 'qq.com'])
  81. )
  82. base_emails = st.text(safe_letters, min_size=1)
  83. emails_with_plus = st.builds(
  84. lambda x, y: '%s+%s' % (x, y), base_emails, base_emails
  85. )
  86. emails = st.builds(
  87. lambda x, y: '%s@%s' % (x, y),
  88. st.one_of(base_emails, emails_with_plus), email_domains
  89. )
  90. def _get_strategy_for_field(f):
  91. # type: () -> Union[SearchStrategy, None, UniqueIdentifier]
  92. #: TODO to replace with mongoengine fields
  93. if f.choices:
  94. choices = [value for (value, name) in f.choices]
  95. if isinstance(f, (mg_fields.StringField, mg_fields.URLField)):
  96. choices.append(u'')
  97. strategy = st.sampled_from(choices)
  98. elif isinstance(f, mg_fields.EmailField):
  99. return emails
  100. elif type(f) in (mg_fields.StringField, ):
  101. strategy = st.text(min_size=f.min_length,
  102. max_size=f.max_length)
  103. elif type(f) == mg.DecimalField:
  104. m = 10 ** f.max_value - 1
  105. div = 10 ** f.precision
  106. q = Decimal('1.' + ('0' * f.decimal_places))
  107. strategy = (
  108. st.integers(min_value=-m, max_value=m)
  109. .map(lambda n: (Decimal(n) / div).quantize(q)))
  110. else:
  111. try:
  112. strategy = field_mappings()[type(f)]
  113. except KeyError:
  114. if f.null:
  115. return None
  116. else:
  117. raise UnmappedFieldError(f)
  118. #if f.validators:
  119. # strategy = strategy.filter(validator_to_filter(f))
  120. if f.null:
  121. strategy = st.one_of(st.none(), strategy)
  122. return strategy
  123. def models(model, **extra):
  124. result = {}
  125. mandatory = set()
  126. for f in model._meta.concrete_fields:
  127. try:
  128. strategy = _get_strategy_for_field(f)
  129. except UnmappedFieldError:
  130. mandatory.add(f.name)
  131. continue
  132. if strategy is not None:
  133. result[f.name] = strategy
  134. missed = {x for x in mandatory if x not in extra}
  135. if missed:
  136. raise InvalidArgument((
  137. u'Missing arguments for mandatory field%s %s for model %s' % (
  138. u's' if len(missed) > 1 else u'',
  139. u', '.join(missed),
  140. model.__name__,
  141. )))
  142. result.update(extra)
  143. # Remove default_values so we don't try to generate anything for those.
  144. result = {k: v for k, v in result.items() if v is not default_value}
  145. return ModelStrategy(model, result)
  146. class ModelStrategy(SearchStrategy):
  147. def __init__(self, model, mappings):
  148. super(ModelStrategy, self).__init__()
  149. self.model = model
  150. self.arg_strategy = st.fixed_dictionaries(mappings)
  151. def __repr__(self):
  152. return u'ModelStrategy(%s)' % (self.model.__name__,)
  153. def do_draw(self, data):
  154. try:
  155. result, _ = self.model.objects.get_or_create(
  156. **self.arg_strategy.do_draw(data)
  157. )
  158. return result
  159. except IntegrityError:
  160. data.mark_invalid()