157 lines
4.5 KiB
Python
157 lines
4.5 KiB
Python
# coding=UTF-8
|
|
|
|
import flask_testing
|
|
import json
|
|
import base64
|
|
from api import create_app
|
|
from flask import Flask
|
|
from api.helpers import check_api_key
|
|
from api.validators import (
|
|
ValidationError,
|
|
validate_exists,
|
|
validate_number,
|
|
validate_territory,
|
|
validate_jsonschema,
|
|
validate_any_of,
|
|
)
|
|
|
|
|
|
class TestHttpBasicAuth(flask_testing.TestCase):
|
|
|
|
api_key = "secret"
|
|
|
|
def create_app(self):
|
|
app = Flask(__name__)
|
|
app.before_request(check_api_key)
|
|
|
|
@app.route("/", methods=["GET"])
|
|
def index():
|
|
return "ok"
|
|
|
|
return app
|
|
|
|
def _enable_basic_auth(self):
|
|
self.app.config["API_KEYS"] = [self.api_key]
|
|
|
|
def _create_auth_headers(self, username, password=""):
|
|
data = base64.b64encode(bytes(":".join([username, password]), "ascii")).decode(
|
|
"ascii"
|
|
)
|
|
headers = [("Authorization", "Basic %s" % data)]
|
|
return headers
|
|
|
|
def test_missing_credentials(self):
|
|
self._enable_basic_auth()
|
|
response = self.client.get("/")
|
|
self.assert401(response)
|
|
|
|
def test_correct_credentials(self):
|
|
self._enable_basic_auth()
|
|
headers = self._create_auth_headers(self.api_key)
|
|
response = self.client.get("/", headers=headers)
|
|
self.assert200(response)
|
|
|
|
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,
|
|
)
|
|
self.assertEqual(
|
|
b"ok",
|
|
self.client.post(
|
|
"/",
|
|
data=json.dumps({"name": "bar", "age": 55}),
|
|
content_type="application/json",
|
|
).data,
|
|
)
|
|
|
|
def test_validate_any_of(self):
|
|
app = self.app
|
|
|
|
@app.route("/", methods=["GET"])
|
|
@validate_any_of("color", ["red", "green", "blue"])
|
|
def index():
|
|
return "ok"
|
|
|
|
self.assertEqual(b"No color is specified", self.client.get("/").data)
|
|
self.assertEqual(
|
|
b"color contains an invalid value, must be one of: ['red', 'green', 'blue']",
|
|
self.client.get("/?color=black").data,
|
|
)
|
|
self.assertEqual(b"ok", self.client.get("/?color=blue").data)
|