From 3ea10d6eb76c30f8220cb472cf9460179de0816c Mon Sep 17 00:00:00 2001 From: Lukas Holecek Date: Aug 03 2018 11:08:14 +0000 Subject: Add option to redirect HTTP to HTTPS --- diff --git a/greenwave/app_factory.py b/greenwave/app_factory.py index e5c9084..24955ca 100644 --- a/greenwave/app_factory.py +++ b/greenwave/app_factory.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: GPL-2.0+ -from flask import Flask +from flask import Flask, redirect, request from greenwave.api_v1 import api from greenwave.utils import json_error, load_config, sha1_mangle_key @@ -9,6 +9,15 @@ from requests import ConnectionError, Timeout from werkzeug.exceptions import default_exceptions +def _setup_http_to_https_redirects(app): + def redirect_http_to_https(): + if request.url.startswith('http://'): + url = request.url.replace('http://', 'https://', 1) + return redirect(url, code=301) + + app.before_request(redirect_http_to_https) + + # applicaiton factory http://flask.pocoo.org/docs/0.12/patterns/appfactories/ def create_app(config_obj=None): app = Flask(__name__) @@ -31,6 +40,9 @@ def create_app(config_obj=None): app.cache = make_region(key_mangler=sha1_mangle_key) app.cache.configure(**app.config['CACHE']) + if app.config.get('REDIRECT_HTTP_TO_HTTPS'): + _setup_http_to_https_redirects(app) + return app diff --git a/greenwave/config.py b/greenwave/config.py index bb693e3..a11d072 100644 --- a/greenwave/config.py +++ b/greenwave/config.py @@ -42,6 +42,7 @@ class Config(object): class ProductionConfig(Config): DEBUG = False PRODUCTION = True + REDIRECT_HTTP_TO_HTTPS = True class DevelopmentConfig(Config): diff --git a/greenwave/tests/test_redirect.py b/greenwave/tests/test_redirect.py new file mode 100644 index 0000000..6dad35e --- /dev/null +++ b/greenwave/tests/test_redirect.py @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: GPL-2.0+ + +from greenwave.app_factory import create_app +from greenwave.config import TestingConfig + + +def test_redirect_http_to_https(): + class Config(TestingConfig): + REDIRECT_HTTP_TO_HTTPS = True + + app = create_app(Config) + client = app.test_client() + r = client.get('/api/v1.0/about') + + assert r.status_code == 301 + assert r.location == 'https://localhost/api/v1.0/about'