validators are now decorators instead

This commit is contained in:
Martin
2016-11-13 19:38:45 +01:00
parent 66b1787722
commit d854ff9798
7 changed files with 172 additions and 55 deletions
+13 -2
View File
@@ -1,12 +1,23 @@
from flask import Blueprint, request, jsonify from flask import Blueprint, request, jsonify
from api import db, tasks from api import db, tasks
from api.models import Designer from api.models import Designer
from api.validators import validate_keys_exists
from api.helpers import pagination_to_dict from api.helpers import pagination_to_dict
from api.validators import validate_jsonschema
mod = Blueprint('designers', __name__, url_prefix='/designers') 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"]) @mod.route("/", methods=["GET"])
def list(): def list():
@@ -22,9 +33,9 @@ def list():
@mod.route("/", methods=["POST"]) @mod.route("/", methods=["POST"])
@validate_jsonschema(create_designer_schema)
def create(): def create():
data = request.json data = request.json
validate_keys_exists(data, 'username', 'name', 'territory', 'currency')
designer = Designer(**data) designer = Designer(**data)
db.session.add(designer) db.session.add(designer)
db.session.commit() db.session.commit()
+10 -11
View File
@@ -1,10 +1,10 @@
from flask import Blueprint, request, jsonify, abort 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.product import Material
from api.models.contract_customer import ContractCustomer from api.models.contract_customer import ContractCustomer
from api.models import market as market_model 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.prices import wallpaper_price, old_canvas_price, old_canvas_diy_frame_price, material_price
from api.lib.limits import calculate_canvas_limits 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') mod = Blueprint('prices', __name__, url_prefix='/prices')
@@ -19,11 +19,10 @@ def get_reseller():
@mod.route("/wallpaper", methods=['GET']) @mod.route("/wallpaper", methods=['GET'])
@validate_exists('width', 'height', 'territory')
@validate_number('width', 'height')
@validate_territory('territory')
def wallpaper(): 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) width = request.args.get('width', type=int)
height = request.args.get('height', type=int) height = request.args.get('height', type=int)
market = market_model.from_territory(request.args.get('territory')) market = market_model.from_territory(request.args.get('territory'))
@@ -44,10 +43,10 @@ def wallpaper():
@mod.route('/canvas', methods=['GET']) @mod.route('/canvas', methods=['GET'])
@validate_exists('width', 'height', 'territory')
@validate_number('width', 'height')
@validate_territory('territory')
def canvas(): 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')) market = market_model.from_territory(request.args.get('territory'))
reseller = get_reseller() reseller = get_reseller()
material = Material.query.canvas().first() material = Material.query.canvas().first()
@@ -71,10 +70,10 @@ def canvas():
@mod.route('/diy-frame', methods=['GET']) @mod.route('/diy-frame', methods=['GET'])
@validate_exists('width', 'height', 'territory')
@validate_number('width', 'height')
@validate_territory('territory')
def diy_frame(): 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')) market = market_model.from_territory(request.args.get('territory'))
reseller = get_reseller() reseller = get_reseller()
width, height = calculate_canvas_limits( width, height = calculate_canvas_limits(
+71 -17
View File
@@ -1,4 +1,6 @@
# coding=UTF-8 from functools import wraps
from flask import request
import jsonschema
from api.models.market import iter_markets from api.models.market import iter_markets
@@ -9,14 +11,12 @@ class ValidationError(Exception):
self.message = message self.message = message
def validate_keys_exists(obj, *keys): def _check_parameter_exists(param):
for key in keys: if not param in request.args:
if not key in obj: raise ValidationError('No {} is specified'.format(param))
raise ValidationError('No {} is specified'.format(key))
return True
def is_number(s): def _is_number(s):
try: try:
float(s) float(s)
return True return True
@@ -24,15 +24,69 @@ def is_number(s):
return False return False
def validate_number(*values): def validate_jsonschema(schema):
for val in values: def decorator(f):
if not is_number(val): @wraps(f)
raise ValidationError('{} is not a number'.format(val)) def wrapper(*args, **kwargs):
return True 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): def validate_exists(*keys):
for market in iter_markets(): def decorator(f):
if market.territory == value: @wraps(f)
return True def wrapper(*args, **kwargs):
raise ValidationError('{} is not a valid territory'.format(value)) 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
+2 -1
View File
@@ -10,4 +10,5 @@ itsdangerous==0.24
psycopg2==2.6.1 psycopg2==2.6.1
Flask-Testing==0.5.0 Flask-Testing==0.5.0
unittest2==1.1.0 unittest2==1.1.0
factory-boy==2.7.0 factory-boy==2.7.0
jsonschema==2.5.1
+1 -1
View File
@@ -47,7 +47,7 @@ class TestEndpoints(flask_testing.TestCase):
self.assertEqual('spider_man', designer['path']) self.assertEqual('spider_man', designer['path'])
def test_create_designer_missing_data(self): def test_create_designer_missing_data(self):
data = {} data = {'name': 'Only name'}
response = self.client.post( response = self.client.post(
'/designers/', data=json.dumps(data), content_type='application/json') '/designers/', data=json.dumps(data), content_type='application/json')
self.assert400(response) self.assert400(response)
-23
View File
@@ -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))
+75
View File
@@ -1,8 +1,11 @@
# coding=UTF-8 # coding=UTF-8
import flask_testing import flask_testing
import json
import base64 import base64
from api import create_app 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): class TestHttpBasicAuth(flask_testing.TestCase):
@@ -35,3 +38,75 @@ class TestHttpBasicAuth(flask_testing.TestCase):
def test_basic_auth_disabled(self): def test_basic_auth_disabled(self):
response = self.client.get('/') response = self.client.get('/')
self.assert200(response) 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)