Support custom origin configuration

This commit is contained in:
Sheng
2019-01-19 16:46:25 +08:00
parent 8a8d741230
commit c35f801235
6 changed files with 157 additions and 20 deletions
+16 -3
View File
@@ -215,7 +215,7 @@ class TestWsockHandler(unittest.TestCase):
request = HTTPServerRequest(uri='/')
obj = Mock(spec=WsockHandler, request=request)
options.cows = 0
obj.origin_policy = 'same'
request.headers['Host'] = 'www.example.com:4433'
origin = 'https://www.example.com:4433'
self.assertTrue(WsockHandler.check_origin(obj, origin))
@@ -223,7 +223,7 @@ class TestWsockHandler(unittest.TestCase):
origin = 'https://www.example.com'
self.assertFalse(WsockHandler.check_origin(obj, origin))
options.cows = 1
obj.origin_policy = 'primary'
self.assertTrue(WsockHandler.check_origin(obj, origin))
origin = 'https://blog.example.com'
@@ -232,5 +232,18 @@ class TestWsockHandler(unittest.TestCase):
origin = 'https://blog.example.org'
self.assertFalse(WsockHandler.check_origin(obj, origin))
options.cows = 2
origin = 'https://blog.example.org'
obj.origin_policy = {'https://blog.example.org'}
self.assertTrue(WsockHandler.check_origin(obj, origin))
origin = 'http://blog.example.org'
obj.origin_policy = {'http://blog.example.org'}
self.assertTrue(WsockHandler.check_origin(obj, origin))
origin = 'http://blog.example.org'
obj.origin_policy = {'https://blog.example.org'}
self.assertFalse(WsockHandler.check_origin(obj, origin))
obj.origin_policy = '*'
origin = 'https://blog.example.org'
self.assertTrue(WsockHandler.check_origin(obj, origin))
+30 -1
View File
@@ -1,4 +1,5 @@
import io
import random
import ssl
import sys
import os.path
@@ -10,7 +11,7 @@ from tests.utils import make_tests_data_path
from webssh.policy import load_host_keys
from webssh.settings import (
get_host_keys_settings, get_policy_setting, base_dir, print_version,
get_ssl_context, get_trusted_downstream
get_ssl_context, get_trusted_downstream, get_origin_setting
)
from webssh.utils import UnicodeType
from webssh._version import __version__
@@ -137,3 +138,31 @@ class TestSettings(unittest.TestCase):
tdstream = '1.1.1.1, 2.2.2.'
with self.assertRaises(ValueError):
get_trusted_downstream(tdstream)
def test_get_origin_setting(self):
options.debug = False
options.origin = '*'
with self.assertRaises(ValueError):
get_origin_setting(options)
options.debug = True
self.assertEqual(get_origin_setting(options), '*')
options.origin = random.choice(['Same', 'Primary'])
self.assertEqual(get_origin_setting(options), options.origin.lower())
options.origin = ''
with self.assertRaises(ValueError):
get_origin_setting(options)
options.origin = ','
with self.assertRaises(ValueError):
get_origin_setting(options)
options.origin = 'www.example.com, https://www.example.org'
result = {'http://www.example.com', 'https://www.example.org'}
self.assertEqual(get_origin_setting(options), result)
options.origin = 'www.example.com:80, www.example.org:443'
result = {'http://www.example.com', 'https://www.example.org'}
self.assertEqual(get_origin_setting(options), result)
+32 -1
View File
@@ -2,7 +2,7 @@ import unittest
from webssh.utils import (
is_valid_ip_address, is_valid_port, is_valid_hostname, to_str, to_bytes,
to_int, is_ip_hostname, is_same_primary_domain
to_int, is_ip_hostname, is_same_primary_domain, parse_origin_from_url
)
@@ -90,3 +90,34 @@ class TestUitls(unittest.TestCase):
domain1 = 'xxx.www.example.com'
domain2 = 'xxx.www2.example.com'
self.assertTrue(is_same_primary_domain(domain1, domain2))
def test_parse_origin_from_url(self):
url = ''
self.assertIsNone(parse_origin_from_url(url))
url = 'www.example.com'
self.assertEqual(parse_origin_from_url(url), 'http://www.example.com')
url = 'http://www.example.com'
self.assertEqual(parse_origin_from_url(url), 'http://www.example.com')
url = 'www.example.com:80'
self.assertEqual(parse_origin_from_url(url), 'http://www.example.com')
url = 'http://www.example.com:80'
self.assertEqual(parse_origin_from_url(url), 'http://www.example.com')
url = 'www.example.com:443'
self.assertEqual(parse_origin_from_url(url), 'https://www.example.com')
url = 'https://www.example.com'
self.assertEqual(parse_origin_from_url(url), 'https://www.example.com')
url = 'https://www.example.com:443'
self.assertEqual(parse_origin_from_url(url), 'https://www.example.com')
url = 'https://www.example.com:80'
self.assertEqual(parse_origin_from_url(url), url)
url = 'http://www.example.com:443'
self.assertEqual(parse_origin_from_url(url), url)