Fixed SANs problem

This commit is contained in:
Brad Warren
2015-12-17 17:28:36 -08:00
parent 20b3188c65
commit 7efdac6c66
4 changed files with 26 additions and 27 deletions
+3 -1
View File
@@ -264,7 +264,9 @@ class TLSSNI01ResponseTest(unittest.TestCase):
def test_verify_bad_cert(self):
self.assertFalse(self.response.verify_cert(
test_util.load_cert('cert.pem')))
OpenSSL.crypto.load_certificate(
OpenSSL.crypto.FILETYPE_PEM,
test_util.load_vector('cert.pem'))))
def test_simple_verify_bad_key_authorization(self):
key2 = jose.JWKRSA.load(test_util.load_vector('rsa256_key.pem'))
+10 -20
View File
@@ -4,11 +4,10 @@ import logging
import socket
import sys
from six.moves import range # pylint: disable=import-error,redefined-builtin
import OpenSSL
from acme import errors
from acme import jose
logger = logging.getLogger(__name__)
@@ -161,31 +160,22 @@ def _pyopenssl_cert_or_req_san(cert_or_req):
:rtype: `list` of `unicode`
"""
# constants based on implementation of
# OpenSSL.crypto.X509Error._subjectAltNameString
# constants based on PyOpenSSL certificate/CSR text dump
label = "DNS"
parts_separator = ", "
part_separator = ":"
extension_short_name = b"subjectAltName"
if hasattr(cert_or_req, 'get_extensions'): # X509Req
extensions = cert_or_req.get_extensions()
else: # X509
extensions = [cert_or_req.get_extension(i)
for i in range(cert_or_req.get_extension_count())]
# pylint: disable=protected-access,no-member
label = OpenSSL.crypto.X509Extension._prefixes[OpenSSL.crypto._lib.GEN_DNS]
assert parts_separator not in label
prefix = label + part_separator
title = "X509v3 Subject Alternative Name:"
san_extensions = [
ext._subjectAltNameString().split(parts_separator)
for ext in extensions if ext.get_short_name() == extension_short_name]
text = jose.ComparableX509(cert_or_req).dump(OpenSSL.crypto.FILETYPE_TEXT)
lines = iter(text.decode("utf-8").splitlines())
sans = [next(lines).split(parts_separator)
for line in lines if title in line]
# WARNING: this function assumes that no SAN can include
# parts_separator, hence the split!
return [part.split(part_separator)[1] for parts in san_extensions
for part in parts if part.startswith(prefix)]
return [part.split(part_separator)[1] for parts in sans
for part in parts if part.lstrip().startswith(prefix)]
def gen_ss_cert(key, domains, not_before=None,
+10 -6
View File
@@ -6,6 +6,8 @@ import unittest
from six.moves import socketserver # pylint: disable=import-error
import OpenSSL
from acme import errors
from acme import jose
from acme import test_util
@@ -64,16 +66,18 @@ class PyOpenSSLCertOrReqSANTest(unittest.TestCase):
"""Test for acme.crypto_util._pyopenssl_cert_or_req_san."""
@classmethod
def _call(cls, loader, name):
def _call(cls, cert_or_req):
# pylint: disable=protected-access
from acme.crypto_util import _pyopenssl_cert_or_req_san
return _pyopenssl_cert_or_req_san(loader(name))
return _pyopenssl_cert_or_req_san(cert_or_req)
def _call_cert(self, name):
return self._call(test_util.load_cert, name)
def _call_cert(self, name, filetype=OpenSSL.crypto.FILETYPE_PEM):
return self._call(OpenSSL.crypto.load_certificate(
filetype, test_util.load_vector(name)))
def _call_csr(self, name):
return self._call(test_util.load_csr, name)
def _call_csr(self, name, filetype=OpenSSL.crypto.FILETYPE_PEM):
return self._call(OpenSSL.crypto.load_certificate_request(
filetype, test_util.load_vector(name)))
def test_cert_no_sans(self):
self.assertEqual(self._call_cert('cert.pem'), [])
+3
View File
@@ -20,6 +20,9 @@ class ComparableX509Test(unittest.TestCase):
self.cert2 = test_util.load_cert('cert.pem')
self.cert_other = test_util.load_cert('cert-san.pem')
def test_getattr_proxy(self):
self.assertTrue(self.cert1.has_expired())
def test_eq(self):
self.assertEqual(self.req1, self.req2)
self.assertEqual(self.cert1, self.cert2)