Support custom origin configuration
This commit is contained in:
+16
-3
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user