From 58faf0a1f699b6d8c93d170a662535c85e14a4f4 Mon Sep 17 00:00:00 2001 From: ksmaze Date: Tue, 28 Oct 2025 16:54:11 -0700 Subject: [PATCH 1/3] feat: add SSL target name override for gRPC connections - Added ssl_grpc_target_name_override option to override TLS hostname verification in gRPC channels - Implemented target name override support in Triton and TensorFlow Serving gRPC clients - Added command line flag --ssl-grpc-target-name-override with usage documentation - Added unit tests to verify target name override functionality - Updated SSL options structs to include the new target_name_override field --- src/client_backend/client_backend.h | 1 + .../tensorflow_serving/tfserve_grpc_client.cc | 3 + .../tensorflow_serving/tfserve_grpc_client.h | 1 + .../triton/test_triton_client_backend.cc | 44 +++++++++++ .../triton/triton_client_backend.cc | 13 +++- src/command_line_parser.cc | 15 ++++ src/test_command_line_parser.cc | 73 ++++++++++++++++++- 7 files changed, 147 insertions(+), 3 deletions(-) diff --git a/src/client_backend/client_backend.h b/src/client_backend/client_backend.h index 07db7ca18..6e6ea62c8 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 f53e4d179..a7f1a8f40 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 bfa475b8c..426610bc9 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 c32ad17be..5f8c6c887 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" @@ -44,6 +45,49 @@ class TestTritonClientBackend : public TritonClientBackend { TritonClientBackend::ParseAndStoreMetric( metrics_endpoint_text, metric_id, metric_per_gpu); } + +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") diff --git a/src/client_backend/triton/triton_client_backend.cc b/src/client_backend/triton/triton_client_backend.cc index 40ce34891..f79450103 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 2a12b49b6..4819b4dce 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 a7bba7c3e..6d2a20cca 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); @@ -87,6 +88,35 @@ CHECK_PARAMS(PAParamsPtr act, PAParamsPtr exp) "Unexpected shape value for: ", act_shape.first, "[", i, "]"); } } + + 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; + } + } CHECK(act->measurement_window_ms == exp->measurement_window_ms); CHECK(act->inference_load_mode == exp->inference_load_mode); CHECK(act->concurrency_range.start == exp->concurrency_range.start); @@ -156,6 +186,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 +340,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 +658,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 +693,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 +2025,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 From 9a72d8f427b364cc90e5233e529aeaf581fed019 Mon Sep 17 00:00:00 2001 From: ksmaze Date: Tue, 28 Oct 2025 17:55:14 -0700 Subject: [PATCH 2/3] refactor: remove SSL target name override test cases --- src/test_command_line_parser.cc | 29 ----------------------------- 1 file changed, 29 deletions(-) diff --git a/src/test_command_line_parser.cc b/src/test_command_line_parser.cc index 6d2a20cca..47e3c36ac 100644 --- a/src/test_command_line_parser.cc +++ b/src/test_command_line_parser.cc @@ -88,35 +88,6 @@ CHECK_PARAMS(PAParamsPtr act, PAParamsPtr exp) "Unexpected shape value for: ", act_shape.first, "[", i, "]"); } } - - 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; - } - } CHECK(act->measurement_window_ms == exp->measurement_window_ms); CHECK(act->inference_load_mode == exp->inference_load_mode); CHECK(act->concurrency_range.start == exp->concurrency_range.start); From ed3ade95c9b65098dbdb76115c3a4e6324ab89c5 Mon Sep 17 00:00:00 2001 From: ksmaze Date: Tue, 28 Oct 2025 18:22:24 -0700 Subject: [PATCH 3/3] fix: correct test class scope and brace placement --- src/client_backend/triton/test_triton_client_backend.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/client_backend/triton/test_triton_client_backend.cc b/src/client_backend/triton/test_triton_client_backend.cc index 5f8c6c887..30d509679 100644 --- a/src/client_backend/triton/test_triton_client_backend.cc +++ b/src/client_backend/triton/test_triton_client_backend.cc @@ -45,6 +45,7 @@ class TestTritonClientBackend : public TritonClientBackend { TritonClientBackend::ParseAndStoreMetric( metrics_endpoint_text, metric_id, metric_per_gpu); } +}; TEST_CASE("TritonClientBackend::Create with GRPC and SSL target name override") { @@ -88,7 +89,6 @@ TEST_CASE("TritonClientBackend::Create with GRPC and SSL target name override") CHECK(backend != nullptr); } } -}; TEST_CASE("testing the ParseAndStoreMetric function") {