Have tracing hooks be owned by GetAuthTokens This ensures BlindSignPerformanceHooks implementations don't need to implement thread safety and complicated bookkeeping if there are multiple concurrent calls to GetAuthTokens. It also eliminates the need for the caller to manage the lifetime of the hooks across callback chains. PiperOrigin-RevId: 800580727
diff --git a/quiche/blind_sign_auth/blind_sign_auth.cc b/quiche/blind_sign_auth/blind_sign_auth.cc index ffb833d..67fe49b 100644 --- a/quiche/blind_sign_auth/blind_sign_auth.cc +++ b/quiche/blind_sign_auth/blind_sign_auth.cc
@@ -33,6 +33,7 @@ #include "quiche/blind_sign_auth/blind_sign_auth_protos.h" #include "quiche/blind_sign_auth/blind_sign_message_interface.h" #include "quiche/blind_sign_auth/blind_sign_message_response.h" +#include "quiche/blind_sign_auth/blind_sign_tracing_hooks.h" #include "quiche/common/platform/api/quiche_logging.h" #include "quiche/common/quiche_random.h" @@ -83,9 +84,10 @@ void BlindSignAuth::GetTokens(std::optional<std::string> oauth_token, int num_tokens, ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, - SignedTokenCallback callback) { - if (hooks_ != nullptr) { - hooks_->OnGetInitialDataStart(); + SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> hooks) { + if (hooks != nullptr) { + hooks->OnGetInitialDataStart(); } // Create GetInitialData RPC. GetInitialDataRequest request; @@ -101,7 +103,7 @@ std::string body_bytes = request.SerializeAsString(); BlindSignMessageCallback initial_data_callback = absl::bind_front( &BlindSignAuth::GetInitialDataCallback, this, oauth_token, num_tokens, - proxy_layer, service_type, std::move(callback)); + proxy_layer, service_type, std::move(callback), std::move(hooks)); fetcher_->DoRequest(BlindSignMessageRequestType::kGetInitialData, oauth_token, body_bytes, std::move(initial_data_callback)); } @@ -109,10 +111,10 @@ void BlindSignAuth::GetInitialDataCallback( std::optional<std::string> oauth_token, int num_tokens, ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, - SignedTokenCallback callback, + SignedTokenCallback callback, std::unique_ptr<BlindSignTracingHooks> hooks, absl::StatusOr<BlindSignMessageResponse> response) { - if (hooks_ != nullptr) { - hooks_->OnGetInitialDataEnd(); + if (hooks != nullptr) { + hooks->OnGetInitialDataEnd(); } absl::StatusOr<GetInitialDataResponse> initial_data_response = ParseGetInitialDataResponseMessage(response); @@ -130,7 +132,7 @@ QUICHE_DVLOG(1) << "Using Privacy Pass client"; GeneratePrivacyPassTokens(*initial_data_response, std::move(oauth_token), num_tokens, proxy_layer, service_type, - std::move(callback)); + std::move(callback), std::move(hooks)); } else { QUICHE_LOG(ERROR) << "Non-Privacy Pass tokens are no longer supported"; std::move(callback)(absl::UnimplementedError( @@ -143,7 +145,8 @@ privacy::ppn::GetInitialDataResponse initial_data_response, std::optional<std::string> oauth_token, int num_tokens, ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, - SignedTokenCallback callback) { + SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> hooks) { absl::StatusOr<PrivacyPassContext> pp_context = CreatePrivacyPassContext(initial_data_response); if (!pp_context.ok()) { @@ -164,15 +167,15 @@ return; } - if (hooks_ != nullptr) { - hooks_->OnGenerateBlindedTokenRequestsStart(); + if (hooks != nullptr) { + hooks->OnGenerateBlindedTokenRequestsStart(); } absl::StatusOr<GeneratedTokenRequests> token_requests_data = GenerateBlindedTokenRequests(num_tokens, *pp_context->rsa_public_key, *token_challenge, pp_context->token_key_id, pp_context->extensions); - if (hooks_ != nullptr) { - hooks_->OnGenerateBlindedTokenRequestsEnd(); + if (hooks != nullptr) { + hooks->OnGenerateBlindedTokenRequestsEnd(); } if (!token_requests_data.ok()) { std::move(callback)(token_requests_data.status()); @@ -193,14 +196,14 @@ sign_request.set_do_not_use_rsa_public_exponent(true); sign_request.set_proxy_layer(QuicheProxyLayerToPpnProxyLayer(proxy_layer)); - if (hooks_ != nullptr) { - hooks_->OnAuthAndSignStart(); + if (hooks != nullptr) { + hooks->OnAuthAndSignStart(); } BlindSignMessageCallback auth_and_sign_callback = absl::bind_front(&BlindSignAuth::PrivacyPassAuthAndSignCallback, this, *std::move(pp_context), std::move(token_requests_data->privacy_pass_clients), - std::move(callback)); + std::move(callback), std::move(hooks)); // TODO(b/304811277): remove other usages of string.data() fetcher_->DoRequest(BlindSignMessageRequestType::kAuthAndSign, oauth_token, sign_request.SerializeAsString(), @@ -212,10 +215,10 @@ std::vector<std::unique_ptr<anonymous_tokens:: PrivacyPassRsaBssaPublicMetadataClient>> privacy_pass_clients, - SignedTokenCallback callback, + SignedTokenCallback callback, std::unique_ptr<BlindSignTracingHooks> hooks, absl::StatusOr<BlindSignMessageResponse> response) { - if (hooks_ != nullptr) { - hooks_->OnAuthAndSignEnd(); + if (hooks != nullptr) { + hooks->OnAuthAndSignEnd(); } // Validate response. if (!response.ok()) { @@ -249,12 +252,12 @@ return; } - if (hooks_ != nullptr) { - hooks_->OnUnblindTokensStart(); + if (hooks != nullptr) { + hooks->OnUnblindTokensStart(); } - absl::Cleanup unblind_tokens_end = [&]() { - if (hooks_ != nullptr) { - hooks_->OnUnblindTokensEnd(); + absl::Cleanup unblind_tokens_end = [hooks = std::move(hooks)]() { + if (hooks != nullptr) { + hooks->OnUnblindTokensEnd(); } };
diff --git a/quiche/blind_sign_auth/blind_sign_auth.h b/quiche/blind_sign_auth/blind_sign_auth.h index 0ae049a..e177128 100644 --- a/quiche/blind_sign_auth/blind_sign_auth.h +++ b/quiche/blind_sign_auth/blind_sign_auth.h
@@ -31,11 +31,8 @@ class QUICHE_EXPORT BlindSignAuth : public BlindSignAuthInterface { public: explicit BlindSignAuth(BlindSignMessageInterface* fetcher, - privacy::ppn::BlindSignAuthOptions auth_options, - BlindSignTracingHooks* hooks = nullptr) - : fetcher_(fetcher), - auth_options_(std::move(auth_options)), - hooks_(hooks) {} + privacy::ppn::BlindSignAuthOptions auth_options) + : fetcher_(fetcher), auth_options_(std::move(auth_options)) {} // Returns signed unblinded tokens, their expiration time, and their geo in a // callback. @@ -45,7 +42,17 @@ // Callers can make multiple concurrent requests to GetTokens. void GetTokens(std::optional<std::string> oauth_token, int num_tokens, ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, - SignedTokenCallback callback) override; + SignedTokenCallback callback) override { + GetTokens(oauth_token, num_tokens, proxy_layer, service_type, + std::move(callback), /*hooks=*/nullptr); + } + + // Same as above, but allows passing tracing hooks which will be used to trace + // this invocation. The hooks will be destroyed after the callback is called. + void GetTokens(std::optional<std::string> oauth_token, int num_tokens, + ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, + SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> hooks) override; // Returns signed unblinded tokens and their expiration time in a // SignedTokenCallback. Errors will be returned in the SignedTokenCallback @@ -87,13 +94,15 @@ std::optional<std::string> oauth_token, int num_tokens, ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> hooks, absl::StatusOr<BlindSignMessageResponse> response); void GeneratePrivacyPassTokens( privacy::ppn::GetInitialDataResponse initial_data_response, std::optional<std::string> oauth_token, int num_tokens, ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, - SignedTokenCallback callback); + SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> hooks); void PrivacyPassAuthAndSignCallback( const PrivacyPassContext& pp_context, @@ -101,6 +110,7 @@ PrivacyPassRsaBssaPublicMetadataClient>> privacy_pass_clients, SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> hooks, absl::StatusOr<BlindSignMessageResponse> response); // Helper functions for GetAttestationTokens flow. @@ -141,7 +151,6 @@ BlindSignMessageInterface* fetcher_ = nullptr; privacy::ppn::BlindSignAuthOptions auth_options_; - BlindSignTracingHooks* hooks_ = nullptr; }; std::string BlindSignAuthServiceTypeToString(
diff --git a/quiche/blind_sign_auth/blind_sign_auth_interface.h b/quiche/blind_sign_auth/blind_sign_auth_interface.h index 9bf1a31..c69ecb2 100644 --- a/quiche/blind_sign_auth/blind_sign_auth_interface.h +++ b/quiche/blind_sign_auth/blind_sign_auth_interface.h
@@ -5,14 +5,17 @@ #ifndef QUICHE_BLIND_SIGN_AUTH_BLIND_SIGN_AUTH_INTERFACE_H_ #define QUICHE_BLIND_SIGN_AUTH_BLIND_SIGN_AUTH_INTERFACE_H_ +#include <memory> #include <optional> #include <string> +#include <utility> #include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "absl/time/time.h" #include "absl/types/span.h" #include "anonymous_tokens/cpp/privacy_pass/token_encodings.h" +#include "quiche/blind_sign_auth/blind_sign_tracing_hooks.h" #include "quiche/common/platform/api/quiche_export.h" #include "quiche/common/quiche_callbacks.h" @@ -76,6 +79,14 @@ ProxyLayer proxy_layer, BlindSignAuthServiceType service_type, SignedTokenCallback callback) = 0; + virtual void GetTokens(std::optional<std::string> oauth_token, int num_tokens, + ProxyLayer proxy_layer, + BlindSignAuthServiceType service_type, + SignedTokenCallback callback, + std::unique_ptr<BlindSignTracingHooks> /*hooks*/) { + GetTokens(oauth_token, num_tokens, proxy_layer, service_type, + std::move(callback)); + } // Returns signed unblinded tokens and their expiration time in a // SignedTokenCallback. Errors will be returned in the SignedTokenCallback
diff --git a/quiche/blind_sign_auth/blind_sign_auth_test.cc b/quiche/blind_sign_auth/blind_sign_auth_test.cc index a910e92..d024cc4 100644 --- a/quiche/blind_sign_auth/blind_sign_auth_test.cc +++ b/quiche/blind_sign_auth/blind_sign_auth_test.cc
@@ -188,8 +188,8 @@ privacy::ppn::BlindSignAuthOptions options; options.set_enable_privacy_pass(true); - blind_sign_auth_ = std::make_unique<BlindSignAuth>( - &mock_message_interface_, options, &mock_tracing_hooks_); + blind_sign_auth_ = + std::make_unique<BlindSignAuth>(&mock_message_interface_, options); } void TearDown() override { blind_sign_auth_.reset(nullptr); } @@ -303,7 +303,6 @@ } MockBlindSignMessageInterface mock_message_interface_; - MockBlindSignTracingHooks mock_tracing_hooks_; std::unique_ptr<BlindSignAuth> blind_sign_auth_; anonymous_tokens::RSABlindSignaturePublicKey public_key_proto_; @@ -433,12 +432,13 @@ } TEST_F(BlindSignAuthTest, TestPrivacyPassGetTokensSucceeds) { + auto mock_tracing_hooks = std::make_unique<MockBlindSignTracingHooks>(); BlindSignMessageResponse fake_public_key_response( absl::StatusCode::kOk, fake_get_initial_data_response_.SerializeAsString()); { InSequence seq; - EXPECT_CALL(mock_tracing_hooks_, OnGetInitialDataStart); + EXPECT_CALL(*mock_tracing_hooks, OnGetInitialDataStart); EXPECT_CALL( mock_message_interface_, DoRequest( @@ -448,10 +448,10 @@ .WillOnce([=](auto&&, auto&&, auto&&, auto get_initial_data_cb) { std::move(get_initial_data_cb)(fake_public_key_response); }); - EXPECT_CALL(mock_tracing_hooks_, OnGetInitialDataEnd); - EXPECT_CALL(mock_tracing_hooks_, OnGenerateBlindedTokenRequestsStart); - EXPECT_CALL(mock_tracing_hooks_, OnGenerateBlindedTokenRequestsEnd); - EXPECT_CALL(mock_tracing_hooks_, OnAuthAndSignStart); + EXPECT_CALL(*mock_tracing_hooks, OnGetInitialDataEnd); + EXPECT_CALL(*mock_tracing_hooks, OnGenerateBlindedTokenRequestsStart); + EXPECT_CALL(*mock_tracing_hooks, OnGenerateBlindedTokenRequestsEnd); + EXPECT_CALL(*mock_tracing_hooks, OnAuthAndSignStart); EXPECT_CALL(mock_message_interface_, DoRequest(Eq(BlindSignMessageRequestType::kAuthAndSign), Eq(oauth_token_), _, _)) @@ -463,9 +463,9 @@ sign_response_.SerializeAsString()); std::move(callback)(response); })); - EXPECT_CALL(mock_tracing_hooks_, OnAuthAndSignEnd); - EXPECT_CALL(mock_tracing_hooks_, OnUnblindTokensStart); - EXPECT_CALL(mock_tracing_hooks_, OnUnblindTokensEnd); + EXPECT_CALL(*mock_tracing_hooks, OnAuthAndSignEnd); + EXPECT_CALL(*mock_tracing_hooks, OnUnblindTokensStart); + EXPECT_CALL(*mock_tracing_hooks, OnUnblindTokensEnd); } int num_tokens = 1; @@ -478,7 +478,8 @@ }; blind_sign_auth_->GetTokens(oauth_token_, num_tokens, ProxyLayer::kProxyA, BlindSignAuthServiceType::kChromeIpBlinding, - std::move(callback)); + std::move(callback), + std::move(mock_tracing_hooks)); done.WaitForNotification(); }