From d854ff9798bb1081bff293b83a3a358a75f6313d Mon Sep 17 00:00:00 2001 From: Martin Date: Sun, 13 Nov 2016 19:38:45 +0100 Subject: [PATCH] validators are now decorators instead --- api/resources/designers.py | 15 +++++- api/resources/prices.py | 21 ++++---- api/validators.py | 88 +++++++++++++++++++++++++------ requirements.txt | 3 +- tests/resources/test_designers.py | 2 +- tests/test_validators.py | 23 -------- tests/test_web.py | 75 ++++++++++++++++++++++++++ 7 files changed, 172 insertions(+), 55 deletions(-) delete mode 100644 tests/test_validators.py diff --git a/api/resources/designers.py b/api/resources/designers.py index 9b43450..4c93c0d 100644 --- a/api/resources/designers.py +++ b/api/resources/designers.py @@ -1,12 +1,23 @@ from flask import Blueprint, request, jsonify from api import db, tasks from api.models import Designer -from api.validators import validate_keys_exists from api.helpers import pagination_to_dict +from api.validators import validate_jsonschema mod = Blueprint('designers', __name__, url_prefix='/designers') +create_designer_schema = { + "type": "object", + "properties": { + "username": {"type": "string"}, + "name": {"type": "string"}, + "territory": {"type": "string"}, + "currency": {"type": "string"}, + }, + "required": ["username", "name", "territory", "currency"] +} + @mod.route("/", methods=["GET"]) def list(): @@ -22,9 +33,9 @@ def list(): @mod.route("/", methods=["POST"]) +@validate_jsonschema(create_designer_schema) def create(): data = request.json - validate_keys_exists(data, 'username', 'name', 'territory', 'currency') designer = Designer(**data) db.session.add(designer) db.session.commit() diff --git a/api/resources/prices.py b/api/resources/prices.py index bee0366..17e64bd 100644 --- a/api/resources/prices.py +++ b/api/resources/prices.py @@ -1,10 +1,10 @@ from flask import Blueprint, request, jsonify, abort +from api.validators import validate_territory, validate_number, validate_exists from api.models.product import Material from api.models.contract_customer import ContractCustomer from api.models import market as market_model from api.lib.prices import wallpaper_price, old_canvas_price, old_canvas_diy_frame_price, material_price from api.lib.limits import calculate_canvas_limits -from api.validators import validate_territory, validate_number, validate_keys_exists mod = Blueprint('prices', __name__, url_prefix='/prices') @@ -19,11 +19,10 @@ def get_reseller(): @mod.route("/wallpaper", methods=['GET']) +@validate_exists('width', 'height', 'territory') +@validate_number('width', 'height') +@validate_territory('territory') def wallpaper(): - validate_keys_exists(request.args, 'width', 'height', 'territory') - validate_number(request.args.get('width'), request.args.get('height')) - validate_territory(request.args.get('territory')) - width = request.args.get('width', type=int) height = request.args.get('height', type=int) market = market_model.from_territory(request.args.get('territory')) @@ -44,10 +43,10 @@ def wallpaper(): @mod.route('/canvas', methods=['GET']) +@validate_exists('width', 'height', 'territory') +@validate_number('width', 'height') +@validate_territory('territory') def canvas(): - validate_keys_exists(request.args, 'width', 'height', 'territory') - validate_number(request.args.get('width'), request.args.get('height')) - validate_territory(request.args.get('territory')) market = market_model.from_territory(request.args.get('territory')) reseller = get_reseller() material = Material.query.canvas().first() @@ -71,10 +70,10 @@ def canvas(): @mod.route('/diy-frame', methods=['GET']) +@validate_exists('width', 'height', 'territory') +@validate_number('width', 'height') +@validate_territory('territory') def diy_frame(): - validate_keys_exists(request.args, 'width', 'height', 'territory') - validate_number(request.args.get('width'), request.args.get('height')) - validate_territory(request.args.get('territory')) market = market_model.from_territory(request.args.get('territory')) reseller = get_reseller() width, height = calculate_canvas_limits( diff --git a/api/validators.py b/api/validators.py index 126569d..81cb5d1 100644 --- a/api/validators.py +++ b/api/validators.py @@ -1,4 +1,6 @@ -# coding=UTF-8 +from functools import wraps +from flask import request +import jsonschema from api.models.market import iter_markets @@ -9,14 +11,12 @@ class ValidationError(Exception): self.message = message -def validate_keys_exists(obj, *keys): - for key in keys: - if not key in obj: - raise ValidationError('No {} is specified'.format(key)) - return True +def _check_parameter_exists(param): + if not param in request.args: + raise ValidationError('No {} is specified'.format(param)) -def is_number(s): +def _is_number(s): try: float(s) return True @@ -24,15 +24,69 @@ def is_number(s): 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_jsonschema(schema): + def decorator(f): + @wraps(f) + def wrapper(*args, **kwargs): + if not request.json: + raise ValidationError('Payload is not valid json') + try: + j = request.json + jsonschema.validate(j, schema) + except jsonschema.ValidationError as e: + raise ValidationError(e.message) + except jsonschema.SchemaError as e: + raise ValidationError(e.message) + return f(*args, **kwargs) + return wrapper + return decorator -def validate_territory(value): - for market in iter_markets(): - if market.territory == value: - return True - raise ValidationError('{} is not a valid territory'.format(value)) +def validate_exists(*keys): + def decorator(f): + @wraps(f) + def wrapper(*args, **kwargs): + for key in keys: + _check_parameter_exists(key) + if not key in request.args: + raise ValidationError('No {} is specified'.format(key)) + return f(*args, **kwargs) + return wrapper + return decorator + + +def validate_number(*keys): + def decorator(f): + @wraps(f) + def wrapper(*args, **kwargs): + for key in keys: + _check_parameter_exists(key) + value = request.args.get(key) + if not _is_number(value): + raise ValidationError('{} is not a number'.format(key)) + return f(*args, **kwargs) + return wrapper + return decorator + + +def validate_territory(*keys): + def decorator(f): + @wraps(f) + def wrapper(*args, **kwargs): + for key in keys: + if not key in request.args: + raise ValidationError('No {} is specified'.format(key)) + value = request.args.get(key) + valid = False + for market in iter_markets(): + if market.territory == value: + valid = True + break + + if not valid: + raise ValidationError( + '{} is not a valid territory'.format(value)) + + return f(*args, **kwargs) + return wrapper + return decorator diff --git a/requirements.txt b/requirements.txt index 537f400..de5227f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -10,4 +10,5 @@ itsdangerous==0.24 psycopg2==2.6.1 Flask-Testing==0.5.0 unittest2==1.1.0 -factory-boy==2.7.0 \ No newline at end of file +factory-boy==2.7.0 +jsonschema==2.5.1 \ No newline at end of file diff --git a/tests/resources/test_designers.py b/tests/resources/test_designers.py index dc88f5a..c7e102d 100644 --- a/tests/resources/test_designers.py +++ b/tests/resources/test_designers.py @@ -47,7 +47,7 @@ class TestEndpoints(flask_testing.TestCase): self.assertEqual('spider_man', designer['path']) def test_create_designer_missing_data(self): - data = {} + data = {'name': 'Only name'} response = self.client.post( '/designers/', data=json.dumps(data), content_type='application/json') self.assert400(response) diff --git a/tests/test_validators.py b/tests/test_validators.py deleted file mode 100644 index df61fd9..0000000 --- a/tests/test_validators.py +++ /dev/null @@ -1,23 +0,0 @@ -import unittest2 -from api.validators import ValidationError, validate_keys_exists, validate_territory, validate_number - - -class TestValidation(unittest2.TestCase): - - def test_validate_keys_exists(self): - obj = {'a': 1, 'b': 2} - self.assertRaises(ValidationError, validate_keys_exists, obj, 'c') - self.assertRaises( - ValidationError, validate_keys_exists, obj, 'a', 'b', 'c') - self.assertTrue(validate_keys_exists(obj, 'a')) - self.assertTrue(validate_keys_exists(obj, 'a', 'b')) - - 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)) diff --git a/tests/test_web.py b/tests/test_web.py index 1fbe60b..0766b41 100644 --- a/tests/test_web.py +++ b/tests/test_web.py @@ -1,8 +1,11 @@ # coding=UTF-8 import flask_testing +import json import base64 from api import create_app +from flask import Flask +from api.validators import ValidationError, validate_exists, validate_number, validate_territory, validate_jsonschema class TestHttpBasicAuth(flask_testing.TestCase): @@ -35,3 +38,75 @@ class TestHttpBasicAuth(flask_testing.TestCase): def test_basic_auth_disabled(self): response = self.client.get('/') self.assert200(response) + + +class TestValidators(flask_testing.TestCase): + + def create_app(self): + app = Flask(__name__) + + @app.errorhandler(ValidationError) + def on_error(error): + return error.message, 400 + return app + + def test_validate_exists(self): + app = self.app + + @app.route('/', methods=['GET']) + @validate_exists('a', 'b') + def index(): + return 'ok' + + self.assertEqual(b'No b is specified', self.client.get('/?a=1').data) + self.assertEqual(b'ok', self.client.get('/?a=1&b=2').data) + + def test_validate_number(self): + app = self.app + + @app.route('/', methods=['GET']) + @validate_number('age') + def index(): + return 'ok' + + self.assertEqual(b'No age is specified', self.client.get('/').data) + self.assertEqual(b'age is not a number', + self.client.get('/?age=aaa').data) + self.assertEqual(b'ok', self.client.get('/?age=1').data) + + def test_validate_territory(self): + app = self.app + + @app.route('/', methods=['GET']) + @validate_territory('territory') + def index(): + return 'ok' + + self.assertEqual(b'No territory is specified', + self.client.get('/').data) + self.assertEqual(b'abc is not a valid territory', + self.client.get('/?territory=abc').data) + self.assertEqual(b'ok', self.client.get('/?territory=SE').data) + + def test_validate_jsonschema(self): + app = self.app + schema = { + "type": "object", + "properties": { + "name": {"type": "string"}, + "age": {"type": "number"} + }, + "required": ["name", "age"] + } + + @app.route('/', methods=['POST']) + @validate_jsonschema(schema) + def index(): + return 'ok' + + self.assertEqual(b'Payload is not valid json', + self.client.post('/').data) + self.assertEqual(b"'name' is a required property", self.client.post( + '/', data=json.dumps({'foo': 'bar'}), content_type='application/json').data) + self.assertEqual(b"'abc' is not of type 'number'", self.client.post( + '/', data=json.dumps({'name': 'bar', 'age': 'abc'}), content_type='application/json').data)