diff --git a/api/validators.py b/api/validators.py index 98d65a0..33f5083 100644 --- a/api/validators.py +++ b/api/validators.py @@ -10,6 +10,21 @@ def validate_keys_exists(obj, *keys): return True +def is_number(s): + try: + float(s) + return True + except ValueError: + return False + + +def validate_number(*values): + for val in values: + if not is_number(val): + raise ValidationError('{} is not a number'.format(val)) + return True + + def validate_territory(value): for market in iter_markets(): if market.territory == value: diff --git a/tests/test_validators.py b/tests/test_validators.py index c76dbe9..e8fe779 100644 --- a/tests/test_validators.py +++ b/tests/test_validators.py @@ -1,5 +1,6 @@ import unittest2 -from api.validators import ValidationError, validate_keys_exists, validate_territory +from api.validators import validate_keys_exists, validate_territory, validate_number +from api.exceptions import ValidationError class TestValidation(unittest2.TestCase): @@ -15,3 +16,9 @@ class TestValidation(unittest2.TestCase): def test_validate_territory(self): self.assertTrue(validate_territory('SE')) self.assertRaises(ValidationError, validate_territory, 'unknown') + + def test_validate_number(self): + self.assertRaises(ValidationError, validate_number, 'foo') + self.assertRaises(ValidationError, validate_number, 'foo', 12) + self.assertTrue(validate_number(12)) + self.assertTrue(validate_number(12, 45))