Add SSL/TLS support for HTTP and gRPC servers (#18973)
Co-authored-by: guys@spotify.com
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -266,5 +267,185 @@ class TestPortArgs(unittest.TestCase):
|
||||
self.assertIn("expected ':' after ']'", str(context.exception))
|
||||
|
||||
|
||||
class TestSSLArgs(unittest.TestCase):
|
||||
def test_default_ssl_fields_are_none(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
self.assertIsNone(server_args.ssl_keyfile)
|
||||
self.assertIsNone(server_args.ssl_certfile)
|
||||
self.assertIsNone(server_args.ssl_ca_certs)
|
||||
self.assertIsNone(server_args.ssl_keyfile_password)
|
||||
|
||||
def test_ssl_keyfile_without_certfile_raises(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(model_path="dummy", ssl_keyfile="key.pem")
|
||||
self.assertIn("--ssl-certfile", str(context.exception))
|
||||
|
||||
def test_ssl_certfile_without_keyfile_raises(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(model_path="dummy", ssl_certfile="cert.pem")
|
||||
self.assertIn("--ssl-keyfile", str(context.exception))
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_ssl_both_keyfile_and_certfile_accepted(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy", ssl_keyfile="key.pem", ssl_certfile="cert.pem"
|
||||
)
|
||||
self.assertEqual(server_args.ssl_keyfile, "key.pem")
|
||||
self.assertEqual(server_args.ssl_certfile, "cert.pem")
|
||||
|
||||
def test_url_returns_http_without_ssl(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
self.assertTrue(server_args.url().startswith("http://"))
|
||||
|
||||
def test_url_rewrites_all_interfaces_to_loopback(self):
|
||||
server_args = ServerArgs(model_path="dummy", host="0.0.0.0")
|
||||
self.assertEqual(server_args.url(), "http://127.0.0.1:30000")
|
||||
|
||||
def test_url_rewrites_empty_host_to_loopback(self):
|
||||
server_args = ServerArgs(model_path="dummy", host="")
|
||||
self.assertEqual(server_args.url(), "http://127.0.0.1:30000")
|
||||
|
||||
def test_url_rewrites_ipv6_all_interfaces_to_loopback(self):
|
||||
server_args = ServerArgs(model_path="dummy", host="::")
|
||||
self.assertEqual(server_args.url(), "http://[::1]:30000")
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_url_returns_https_with_ssl(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy", ssl_keyfile="key.pem", ssl_certfile="cert.pem"
|
||||
)
|
||||
self.assertTrue(server_args.url().startswith("https://"))
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_ssl_cli_args_parsed(self, _mock_isfile):
|
||||
server_args = prepare_server_args(
|
||||
[
|
||||
"--model-path",
|
||||
"dummy",
|
||||
"--ssl-keyfile",
|
||||
"key.pem",
|
||||
"--ssl-certfile",
|
||||
"cert.pem",
|
||||
"--ssl-ca-certs",
|
||||
"ca.pem",
|
||||
"--ssl-keyfile-password",
|
||||
"secret",
|
||||
]
|
||||
)
|
||||
self.assertEqual(server_args.ssl_keyfile, "key.pem")
|
||||
self.assertEqual(server_args.ssl_certfile, "cert.pem")
|
||||
self.assertEqual(server_args.ssl_ca_certs, "ca.pem")
|
||||
self.assertEqual(server_args.ssl_keyfile_password, "secret")
|
||||
|
||||
def test_ssl_verify_without_ssl(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
self.assertIs(server_args.ssl_verify(), True)
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_ssl_verify_with_ssl_no_ca(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy", ssl_keyfile="key.pem", ssl_certfile="cert.pem"
|
||||
)
|
||||
self.assertIs(server_args.ssl_verify(), False)
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_ssl_verify_with_ssl_and_ca(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
ssl_keyfile="key.pem",
|
||||
ssl_certfile="cert.pem",
|
||||
ssl_ca_certs="ca.pem",
|
||||
)
|
||||
self.assertEqual(server_args.ssl_verify(), "ca.pem")
|
||||
|
||||
def test_ssl_ca_certs_without_certfile_raises(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(model_path="dummy", ssl_ca_certs="ca.pem")
|
||||
self.assertIn("--ssl-ca-certs", str(context.exception))
|
||||
|
||||
def test_ssl_keyfile_password_without_certfile_raises(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(model_path="dummy", ssl_keyfile_password="secret")
|
||||
self.assertIn("--ssl-keyfile-password", str(context.exception))
|
||||
|
||||
def test_ssl_keyfile_not_found_raises(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(
|
||||
model_path="dummy",
|
||||
ssl_keyfile="/nonexistent/key.pem",
|
||||
ssl_certfile="/nonexistent/cert.pem",
|
||||
)
|
||||
self.assertIn("not found", str(context.exception))
|
||||
|
||||
def test_ssl_certfile_not_found_raises(self):
|
||||
with tempfile.NamedTemporaryFile(suffix=".pem") as keyfile:
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(
|
||||
model_path="dummy",
|
||||
ssl_keyfile=keyfile.name,
|
||||
ssl_certfile="/nonexistent/cert.pem",
|
||||
)
|
||||
self.assertIn("SSL certificate file not found", str(context.exception))
|
||||
|
||||
def test_ssl_ca_certs_not_found_raises(self):
|
||||
with tempfile.NamedTemporaryFile(suffix=".pem") as keyfile:
|
||||
with tempfile.NamedTemporaryFile(suffix=".pem") as certfile:
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(
|
||||
model_path="dummy",
|
||||
ssl_keyfile=keyfile.name,
|
||||
ssl_certfile=certfile.name,
|
||||
ssl_ca_certs="/nonexistent/ca.pem",
|
||||
)
|
||||
self.assertIn(
|
||||
"SSL CA certificates file not found", str(context.exception)
|
||||
)
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_url_returns_https_with_ssl_and_ipv6(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
host="::1",
|
||||
ssl_keyfile="key.pem",
|
||||
ssl_certfile="cert.pem",
|
||||
)
|
||||
self.assertEqual(server_args.url(), "https://[::1]:30000")
|
||||
|
||||
def test_enable_ssl_refresh_default_false(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
self.assertFalse(server_args.enable_ssl_refresh)
|
||||
|
||||
def test_enable_ssl_refresh_without_ssl_raises(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
ServerArgs(model_path="dummy", enable_ssl_refresh=True)
|
||||
self.assertIn("--enable-ssl-refresh", str(context.exception))
|
||||
self.assertIn("--ssl-certfile", str(context.exception))
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_enable_ssl_refresh_with_ssl_accepted(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
ssl_keyfile="key.pem",
|
||||
ssl_certfile="cert.pem",
|
||||
enable_ssl_refresh=True,
|
||||
)
|
||||
self.assertTrue(server_args.enable_ssl_refresh)
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_enable_ssl_refresh_cli_flag(self, _mock_isfile):
|
||||
server_args = prepare_server_args(
|
||||
[
|
||||
"--model-path",
|
||||
"dummy",
|
||||
"--ssl-keyfile",
|
||||
"key.pem",
|
||||
"--ssl-certfile",
|
||||
"cert.pem",
|
||||
"--enable-ssl-refresh",
|
||||
]
|
||||
)
|
||||
self.assertTrue(server_args.enable_ssl_refresh)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user