diff --git a/care/utils/models/validators.py b/care/utils/models/validators.py index 58624d1143..9b7a3e3581 100644 --- a/care/utils/models/validators.py +++ b/care/utils/models/validators.py @@ -1,5 +1,5 @@ import re -from collections.abc import Iterable +from collections.abc import Collection, Iterable from fractions import Fraction from pathlib import Path @@ -103,9 +103,23 @@ class PhoneNumberValidator(RegexValidator): "support": support_number_regex, } - def __init__(self, types: Iterable[str], *args, **kwargs): - if not isinstance(types, Iterable) or isinstance(types, str) or len(types) == 0: - msg = "The `types` argument must be a non-empty iterable." + def __init__(self, types: Collection[str], *args, **kwargs): + if not isinstance(types, Collection) or isinstance(types, str): + msg = "The `types` argument must be a non-empty collection." + raise ValueError(msg) + + types = tuple(types) + if not types: + msg = "The `types` argument must be a non-empty collection." + raise ValueError(msg) + + unsupported_types = [ + type_ + for type_ in types + if not isinstance(type_, str) or type_ not in self.regex_map + ] + if unsupported_types: + msg = f"Unsupported phone number type(s): {', '.join(str(type_) for type_ in unsupported_types)}." raise ValueError(msg) self.types = types diff --git a/care/utils/tests/test_phone_number_validator.py b/care/utils/tests/test_phone_number_validator.py index 5f378c9954..f29b020e6e 100644 --- a/care/utils/tests/test_phone_number_validator.py +++ b/care/utils/tests/test_phone_number_validator.py @@ -130,3 +130,40 @@ def test_invalid_support_numbers(self): for number in self.invalid_support_numbers: with self.assertRaises(ValidationError, msg=f"Failed for {number}"): self.support_validator(number) + + def test_types_must_be_non_empty_collection(self): + invalid_types = ["mobile", (), (type_ for type_ in ("mobile",))] + + for types in invalid_types: + with self.assertRaisesMessage( + ValueError, + "The `types` argument must be a non-empty collection.", + ): + PhoneNumberValidator(types=types) + + def test_unsupported_types_raise_value_error(self): + with self.assertRaisesMessage( + ValueError, + "Unsupported phone number type(s): pager.", + ): + PhoneNumberValidator(types=("mobile", "pager")) + + def test_unhashable_types_raise_value_error(self): + with self.assertRaisesMessage( + ValueError, + "Unsupported phone number type(s): [].", + ): + PhoneNumberValidator(types=([],)) + + def test_non_string_types_raise_value_error(self): + with self.assertRaisesMessage( + ValueError, + "Unsupported phone number type(s): 1.", + ): + PhoneNumberValidator(types=(1,)) + + def test_types_accepts_reiterable_collection(self): + validator = PhoneNumberValidator(types=["mobile", "landline"]) + + self.assertIsNone(validator("+919876543210")) + self.assertIsNone(validator("+914902626488"))