diff --git a/src/client_backend/client_backend.h b/src/client_backend/client_backend.h index 07db7ca1..6e6ea62c 100644 --- a/src/client_backend/client_backend.h +++ b/src/client_backend/client_backend.h @@ -258,6 +258,7 @@ struct SslOptionsBase { std::string ssl_grpc_root_certifications_file = ""; std::string ssl_grpc_private_key_file = ""; std::string ssl_grpc_certificate_chain_file = ""; + std::string ssl_grpc_target_name_override = ""; long ssl_https_verify_peer = 1L; long ssl_https_verify_host = 2L; std::string ssl_https_ca_certificates_file = ""; diff --git a/src/client_backend/tensorflow_serving/tfserve_grpc_client.cc b/src/client_backend/tensorflow_serving/tfserve_grpc_client.cc index f53e4d17..a7f1a8f4 100644 --- a/src/client_backend/tensorflow_serving/tfserve_grpc_client.cc +++ b/src/client_backend/tensorflow_serving/tfserve_grpc_client.cc @@ -112,6 +112,9 @@ GetChannel(const std::string& url, bool use_ssl, const SslOptions& ssl_options) grpc::ChannelArguments arguments; arguments.SetMaxSendMessageSize(tc::MAX_GRPC_MESSAGE_SIZE); arguments.SetMaxReceiveMessageSize(tc::MAX_GRPC_MESSAGE_SIZE); + if (!ssl_options.target_name_override.empty()) { + arguments.SetSslTargetNameOverride(ssl_options.target_name_override); + } std::shared_ptr credentials; if (use_ssl) { std::string root; diff --git a/src/client_backend/tensorflow_serving/tfserve_grpc_client.h b/src/client_backend/tensorflow_serving/tfserve_grpc_client.h index bfa475b8..426610bc 100644 --- a/src/client_backend/tensorflow_serving/tfserve_grpc_client.h +++ b/src/client_backend/tensorflow_serving/tfserve_grpc_client.h @@ -53,6 +53,7 @@ struct SslOptions { // This parameter can be empty if the client does not have a // certificate chain. std::string certificate_chain; + std::string target_name_override; }; class InferResult; diff --git a/src/client_backend/triton/test_triton_client_backend.cc b/src/client_backend/triton/test_triton_client_backend.cc index c32ad17b..30d50967 100644 --- a/src/client_backend/triton/test_triton_client_backend.cc +++ b/src/client_backend/triton/test_triton_client_backend.cc @@ -27,6 +27,7 @@ #include #include #include +#include #include "../../doctest.h" #include "triton_client_backend.h" @@ -46,6 +47,49 @@ class TestTritonClientBackend : public TritonClientBackend { } }; +TEST_CASE("TritonClientBackend::Create with GRPC and SSL target name override") +{ + using triton::perfanalyzer::clientbackend::Error; + using triton::perfanalyzer::clientbackend::SslOptionsBase; + using triton::perfanalyzer::clientbackend::ClientBackend; + using triton::perfanalyzer::clientbackend::tritonremote::TritonClientBackend; + + const std::string url{"localhost:8001"}; + const bool verbose{false}; + const std::string metrics_url{}; + const auto protocol = cb::ProtocolType::GRPC; + const auto input_tensor_format = cb::TensorFormat::BINARY; + const auto output_tensor_format = cb::TensorFormat::BINARY; + const grpc_compression_algorithm compression = GRPC_COMPRESS_NONE; + std::shared_ptr headers = std::make_shared(); + std::map> trace_options{}; + + SUBCASE("override set") + { + SslOptionsBase ssl{}; + ssl.ssl_grpc_target_name_override = "my.host.name"; + + std::unique_ptr backend; + Error err = TritonClientBackend::Create( + url, protocol, ssl, trace_options, compression, headers, verbose, + metrics_url, input_tensor_format, output_tensor_format, &backend); + CHECK(err.IsOk()); + CHECK(backend != nullptr); + } + + SUBCASE("override empty") + { + SslOptionsBase ssl{}; + + std::unique_ptr backend; + Error err = TritonClientBackend::Create( + url, protocol, ssl, trace_options, compression, headers, verbose, + metrics_url, input_tensor_format, output_tensor_format, &backend); + CHECK(err.IsOk()); + CHECK(backend != nullptr); + } +} + TEST_CASE("testing the ParseAndStoreMetric function") { TestTritonClientBackend ttcb{}; diff --git a/src/client_backend/triton/triton_client_backend.cc b/src/client_backend/triton/triton_client_backend.cc index 40ce3489..f7945010 100644 --- a/src/client_backend/triton/triton_client_backend.cc +++ b/src/client_backend/triton/triton_client_backend.cc @@ -27,6 +27,7 @@ #include "triton_client_backend.h" #include +#include #include #include @@ -124,9 +125,17 @@ TritonClientBackend::Create( ParseGrpcSslOptions(ssl_options); bool use_ssl = grpc_ssl_options_pair.first; triton::client::SslOptions grpc_ssl_options = grpc_ssl_options_pair.second; + + grpc::ChannelArguments channel_args; + channel_args.SetMaxSendMessageSize(tc::MAX_GRPC_MESSAGE_SIZE); + channel_args.SetMaxReceiveMessageSize(tc::MAX_GRPC_MESSAGE_SIZE); + if (!ssl_options.ssl_grpc_target_name_override.empty()) { + channel_args.SetSslTargetNameOverride( + ssl_options.ssl_grpc_target_name_override); + } RETURN_IF_TRITON_ERROR(tc::InferenceServerGrpcClient::Create( - &(triton_client_backend->client_.grpc_client_), url, verbose, use_ssl, - grpc_ssl_options)); + &(triton_client_backend->client_.grpc_client_), url, channel_args, + verbose, use_ssl, grpc_ssl_options, true /*use_cached_channel*/)); if (!trace_options.empty()) { inference::TraceSettingResponse response; RETURN_IF_TRITON_ERROR( diff --git a/src/command_line_parser.cc b/src/command_line_parser.cc index 2a12b49b..4819b4dc 100644 --- a/src/command_line_parser.cc +++ b/src/command_line_parser.cc @@ -170,6 +170,7 @@ CLParser::Usage(const std::string& msg) std::cerr << "\t--ssl-grpc-root-certifications-file " << std::endl; std::cerr << "\t--ssl-grpc-private-key-file " << std::endl; std::cerr << "\t--ssl-grpc-certificate-chain-file " << std::endl; + std::cerr << "\t--ssl-grpc-target-name-override " << std::endl; std::cerr << "\t--ssl-https-verify-peer " << std::endl; std::cerr << "\t--ssl-https-verify-host " << std::endl; std::cerr << "\t--ssl-https-ca-certificates-file " << std::endl; @@ -670,6 +671,14 @@ CLParser::Usage(const std::string& msg) "PEM encoding of the client's certificate chain.", 38) << std::endl; + std::cerr << std::setw(38) << std::left + << " --ssl-grpc-target-name-override: " + << FormatMessage( + "Override the target name used for TLS hostname verification " + "on the gRPC channel. Useful when connecting to an IP with " + "a certificate issued for a hostname.", + 38) + << std::endl; std::cerr << std::setw(38) << std::left << " --ssl-https-verify-peer: " << FormatMessage( "Number (0|1) to verify the " @@ -925,6 +934,8 @@ CLParser::ParseCommandLine(int argc, char** argv) long_option_idx_base + 40}, {"ssl-https-private-key-type", required_argument, 0, long_option_idx_base + 41}, + {"ssl-grpc-target-name-override", required_argument, 0, + long_option_idx_base + 67}, {"verbose-csv", no_argument, 0, long_option_idx_base + 42}, {"enable-mpi", no_argument, 0, long_option_idx_base + 43}, {"trace-level", required_argument, 0, long_option_idx_base + 44}, @@ -1420,6 +1431,10 @@ CLParser::ParseCommandLine(int argc, char** argv) } break; } + case long_option_idx_base + 67: { + params_->ssl_options.ssl_grpc_target_name_override = optarg; + break; + } case long_option_idx_base + 35: { if (std::atol(optarg) == 0 || std::atol(optarg) == 1) { params_->ssl_options.ssl_https_verify_peer = std::atol(optarg); diff --git a/src/test_command_line_parser.cc b/src/test_command_line_parser.cc index a7bba7c3..47e3c36a 100644 --- a/src/test_command_line_parser.cc +++ b/src/test_command_line_parser.cc @@ -74,6 +74,7 @@ CHECK_PARAMS(PAParamsPtr act, PAParamsPtr exp) for (size_t i = 0; i < act->user_data.size(); i++) { CHECK_STRING(act->user_data[i], exp->user_data[i]); } + CHECK(act->input_shapes.size() == exp->input_shapes.size()); for (auto act_shape : act->input_shapes) { auto exp_shape = exp->input_shapes.find(act_shape.first); @@ -156,6 +157,9 @@ CHECK_PARAMS(PAParamsPtr act, PAParamsPtr exp) CHECK( act->ssl_options.ssl_https_verify_peer == exp->ssl_options.ssl_https_verify_peer); + CHECK_STRING( + act->ssl_options.ssl_grpc_target_name_override, + exp->ssl_options.ssl_grpc_target_name_override); CHECK(act->verbose_csv == exp->verbose_csv); CHECK(act->enable_mpi == exp->enable_mpi); CHECK(act->trace_options.size() == exp->trace_options.size()); @@ -307,6 +311,9 @@ TEST_CASE("Testing PerfAnalyzerParameters") "ssl_grpc_root_certifications_file", params->ssl_options.ssl_grpc_root_certifications_file, ""); CHECK(params->ssl_options.ssl_grpc_use_ssl == false); + CHECK_STRING( + "ssl_grpc_target_name_override", + params->ssl_options.ssl_grpc_target_name_override, ""); CHECK_STRING( "ssl_https_ca_certificates_file", params->ssl_options.ssl_https_ca_certificates_file, ""); @@ -622,6 +629,12 @@ TEST_CASE("Testing Command Line Parser") exp->max_threads = 1; exp->max_threads_specified = true; + + REQUIRE(act->user_data.size() == exp->user_data.size()); + for (size_t i = 0; i < act->user_data.size(); i++) { + CHECK_STRING(act->user_data[i], exp->user_data[i]); + } + CHECK(act->input_shapes.size() == exp->input_shapes.size()); } SUBCASE("set to max") @@ -651,7 +664,7 @@ TEST_CASE("Testing Command Line Parser") SUBCASE("bad value") { - int argc = 4; + int argc = 5; char* argv[argc] = {app_name, "-m", model_name, "--max-threads", "bad"}; CHECK_THROWS_WITH_AS( @@ -1983,6 +1996,35 @@ TEST_CASE("Testing Command Line Parser") } } + SUBCASE("Option : --ssl-grpc-target-name-override") + { + SUBCASE("set to my.host.name") + { + int argc = 5; + char* argv[argc] = {app_name, "-m", model_name, + "--ssl-grpc-target-name-override", "my.host.name"}; + + REQUIRE_NOTHROW(act = parser.Parse(argc, argv)); + CHECK(!parser.UsageCalled()); + + exp->ssl_options.ssl_grpc_target_name_override = "my.host.name"; + } + + SUBCASE("missing value") + { + int argc = 4; + char* argv[argc] = { + app_name, "-m", model_name, "--ssl-grpc-target-name-override"}; + + CHECK_THROWS_WITH_AS( + act = parser.Parse(argc, argv), + "Error: Missing value for option '--ssl-grpc-target-name-override'", + PerfAnalyzerException); + + check_params = false; + } + } + if (check_params) { if (act == nullptr) { std::cerr