Fix and add test for get_vhosts

This commit is contained in:
yan
2015-04-17 22:24:18 -07:00
parent e5a027ce30
commit d9c8c13f9a
7 changed files with 91 additions and 38 deletions
@@ -163,7 +163,7 @@ class NginxConfigurator(object):
matches = self._get_ranked_matches(target_name) matches = self._get_ranked_matches(target_name)
if len(matches) == 0: if len(matches) == 0:
# No matches at all :'( # No matches at all :'(
break pass
elif matches[0]['rank'] in range(2, 6): elif matches[0]['rank'] in range(2, 6):
# Wildcard match - need to find the longest one # Wildcard match - need to find the longest one
rank = matches[0]['rank'] rank = matches[0]['rank']
+4 -2
View File
@@ -71,7 +71,9 @@ class Addr(object):
def __eq__(self, other): def __eq__(self, other):
if isinstance(other, self.__class__): if isinstance(other, self.__class__):
return self.tup == other.tup return (self.tup == other.tup and
self.ssl == other.ssl and
self.default == other.default)
return False return False
def __hash__(self): def __hash__(self):
@@ -124,7 +126,7 @@ class VirtualHost(object): # pylint: disable=too-few-public-methods
def __eq__(self, other): def __eq__(self, other):
if isinstance(other, self.__class__): if isinstance(other, self.__class__):
return (self.filep == other.filep and return (self.filep == other.filep and
self.addrs == other.addrs and list(self.addrs) == list(other.addrs) and
self.names == other.names and self.names == other.names and
self.ssl == other.ssl and self.enabled == other.enabled) self.ssl == other.ssl and self.enabled == other.enabled)
+35 -26
View File
@@ -105,8 +105,8 @@ class NginxParser(object):
servers[filename] = [] servers[filename] = []
# Find all the server blocks # Find all the server blocks
do_for_subarray(tree, lambda x: x[0] == ['server'], _do_for_subarray(tree, lambda x: x[0] == ['server'],
lambda x: servers[filename].append(x[1])) lambda x: servers[filename].append(x[1]))
# Find 'include' statements in server blocks and append their trees # Find 'include' statements in server blocks and append their trees
for server in servers[filename]: for server in servers[filename]:
@@ -116,10 +116,7 @@ class NginxParser(object):
self.abs_path(directive[1])) self.abs_path(directive[1]))
for f in included_files: for f in included_files:
try: try:
# Assign instead of append because servers[f] server.extend(self.parsed[f])
# should be empty since server blocks cannot
# contain other server blocks.
servers[f] = self.parsed[f]
except: except:
pass pass
@@ -128,10 +125,10 @@ class NginxParser(object):
# Parse the server block into a VirtualHost object # Parse the server block into a VirtualHost object
parsed_server = self._parse_server(server) parsed_server = self._parse_server(server)
vhost = obj.VirtualHost(filename, vhost = obj.VirtualHost(filename,
parsed_server.addrs, parsed_server['addrs'],
parsed_server.ssl, parsed_server['ssl'],
enabled, enabled,
parsed_server.names) parsed_server['names'])
vhosts.append(vhost) vhosts.append(vhost)
return vhosts return vhosts
@@ -144,21 +141,33 @@ class NginxParser(object):
""" """
parsed_server = {} parsed_server = {}
parsed_server.addrs = set() parsed_server['addrs'] = set()
parsed_server.ssl = False parsed_server['ssl'] = False
parsed_server.names = set() parsed_server['names'] = set()
for directive in server: for directive in server:
if directive[0] == 'listen': if directive[0] == 'listen':
addr = obj.Addr.fromstring(directive[1]) addr = obj.Addr.fromstring(directive[1])
parsed_server.addrs.add(addr) parsed_server['addrs'].add(addr)
if not parsed_server.ssl and addr.ssl: if not parsed_server['ssl'] and addr.ssl:
parsed_server.ssl = True parsed_server['ssl'] = True
elif directive[0] == 'server_name': elif directive[0] == 'server_name':
parsed_server.names.update(' '.split(directive[1])) parsed_server['names'].update(
self._get_servernames(directive[1]))
return parsed_server return parsed_server
def _get_servernames(self, names):
"""Turns a server_name string into a list of server names
:param str names: server names
:rtype: list
"""
whitespace_re = re.compile(r'\s+')
names = re.sub(whitespace_re, ' ', names)
return names.split(' ')
def _parse_files(self, filepath): def _parse_files(self, filepath):
"""Parse files from a glob """Parse files from a glob
@@ -260,7 +269,7 @@ class NginxParser(object):
return False return False
if item[0] == 'server_name': if item[0] == 'server_name':
server_names.update((' ').split(item[1])) server_names.update(self._get_servernames(item[1]))
return server_names == names return server_names == names
@@ -302,16 +311,16 @@ class NginxParser(object):
""" """
if replace: if replace:
do_for_subarray(self.parsed[filename], _do_for_subarray(self.parsed[filename],
lambda x: self._has_server_names(x, names), lambda x: self._has_server_names(x, names),
lambda x: self._replace_directives(x, directives)) lambda x: self._replace_directives(x, directives))
else: else:
do_for_subarray(self.parsed[filename], _do_for_subarray(self.parsed[filename],
lambda x: self._has_server_names(x, names), lambda x: self._has_server_names(x, names),
lambda x: x.extend(directives)) lambda x: x.extend(directives))
def do_for_subarray(entry, condition, func): def _do_for_subarray(entry, condition, func):
"""Executes a function for a subarray of a nested array if it matches """Executes a function for a subarray of a nested array if it matches
the given condition. the given condition.
@@ -326,9 +335,9 @@ def do_for_subarray(entry, condition, func):
try: try:
func(item) func(item)
except: except:
logging.warn("Error in do_for_subarray for %s" % item) logging.warn("Error in _do_for_subarray for %s" % item)
else: else:
do_for_subarray(item, condition, func) _do_for_subarray(item, condition, func)
def get_best_match(target_name, names): def get_best_match(target_name, names):
@@ -5,12 +5,11 @@ import shutil
import sys import sys
import unittest import unittest
import mock
import zope.component import zope.component
from letsencrypt.client import errors
from letsencrypt.client.display import util as display_util from letsencrypt.client.display import util as display_util
from letsencrypt.client.plugins.nginx.obj import Addr, VirtualHost
from letsencrypt.client.plugins.nginx.parser import NginxParser from letsencrypt.client.plugins.nginx.parser import NginxParser
from letsencrypt.client.plugins.nginx.tests import util from letsencrypt.client.plugins.nginx.tests import util
@@ -56,7 +55,8 @@ class NginxParserTest(util.NginxTest):
self.assertEqual([['server_name', 'somename alias another.alias']], self.assertEqual([['server_name', 'somename alias another.alias']],
parser.parsed[parser.abs_path('server.conf')]) parser.parsed[parser.abs_path('server.conf')])
self.assertEqual([[['server'], [['listen', '9000'], self.assertEqual([[['server'], [['listen', '9000'],
['server_name', 'example.com']]]], ['server_name', '.example.com'],
['server_name', 'example.*']]]],
parser.parsed[parser.abs_path( parser.parsed[parser.abs_path(
'sites-enabled/example.com')]) 'sites-enabled/example.com')])
@@ -76,9 +76,50 @@ class NginxParserTest(util.NginxTest):
self.assertEqual(2, len( self.assertEqual(2, len(
glob.glob(parser.abs_path('sites-enabled/*.test')))) glob.glob(parser.abs_path('sites-enabled/*.test'))))
self.assertEqual([[['server'], [['listen', '9000'], self.assertEqual([[['server'], [['listen', '9000'],
['server_name', 'example.com']]]], ['server_name', '.example.com'],
['server_name', 'example.*']]]],
parsed[0]) parsed[0])
def test_get_vhosts(self):
parser = NginxParser(self.config_path, self.ssl_options)
vhosts = parser.get_vhosts()
vhost1 = VirtualHost(parser.abs_path('nginx.conf'),
[Addr('', '8080', False, False)],
False, True, set(['localhost']))
vhost2 = VirtualHost(parser.abs_path('nginx.conf'),
[Addr('somename', '8080', False, False),
Addr('', '8000', False, False)],
False, True, set(['somename',
'another.alias', 'alias']))
vhost3 = VirtualHost(parser.abs_path('sites-enabled/example.com'),
[Addr('', '9000', False, False)],
False, True, set(['.example.com', 'example.*']))
vhost4 = VirtualHost(parser.abs_path('sites-enabled/default'),
[Addr('myhost', '', False, True)],
False, True, set(['www.example.org']))
vhost5 = VirtualHost(parser.abs_path('foo.conf'),
[Addr('*', '80', True, True)],
True, True, set(['*.www.foo.com']))
self.assertEqual(5, len(vhosts))
example_com = filter(lambda x: 'example.com' in x.filep, vhosts)[0]
self.assertEqual(vhost3, example_com)
default = filter(lambda x: 'default' in x.filep, vhosts)[0]
self.assertEqual(vhost4, default)
foo = filter(lambda x: 'foo.conf' in x.filep, vhosts)[0]
self.assertEqual(vhost5, foo)
localhost = filter(lambda x: 'localhost' in x.names, vhosts)[0]
self.assertEquals(vhost1, localhost)
somename = filter(lambda x: 'somename' in x.names, vhosts)[0]
self.assertEquals(vhost2, somename)
def test_add_server_directives(self):
pass
def test_get_best_match(self):
pass
# def test_find_dir(self): # def test_find_dir(self):
# from letsencrypt.client.plugins.nginx.parser import case_i # from letsencrypt.client.plugins.nginx.parser import case_i
# test = self.parser.find_dir(case_i("Listen"), "443") # test = self.parser.find_dir(case_i("Listen"), "443")
+2 -2
View File
@@ -2,8 +2,8 @@
user www-data; user www-data;
server { server {
listen 80; listen *:80 default_server ssl;
server_name foo.com; server_name *.www.foo.com;
root /home/ubuntu/sites/foo/; root /home/ubuntu/sites/foo/;
location /status { location /status {
@@ -1,6 +1,6 @@
server { server {
listen 1234; listen myhost default_server;
server_name example.org; server_name www.example.org;
location / { location / {
root html; root html;
@@ -1,4 +1,5 @@
server { server {
listen 9000; listen 9000;
server_name example.com; server_name .example.com;
server_name example.*;
} }