diff --git a/src/gax-internal/src/api_header.rs b/src/gax-internal/src/api_header.rs index ddddfbb9c7..31f14727af 100644 --- a/src/gax-internal/src/api_header.rs +++ b/src/gax-internal/src/api_header.rs @@ -14,6 +14,10 @@ //! Telemetry header helpers. +use google_cloud_gax::client_builder::internal::Extensions; +use google_cloud_gax::error::Error; +use http::HeaderValue; + /// Generated libraries create one static instance of this struct and use it /// to lazy initialize (via [std::sync::LazyLock]) the x-goog-api-client header /// value. @@ -72,6 +76,22 @@ impl XGoogApiClient { } } +/// Resolves the effective `x-goog-api-client` header value for a gRPC request. +/// +/// Precedence order: +/// 1. Client-level `XGoogApiClient` +/// 2. Default static header value +#[cfg(any(test, feature = "_internal-grpc-client"))] +pub(crate) fn resolve_grpc_header_value( + extensions: &Extensions, + default_header: &'static str, +) -> Result { + if let Some(h) = extensions.get::() { + return HeaderValue::from_str(&h.grpc_header_value()).map_err(Error::ser); + } + Ok(HeaderValue::from_static(default_header)) +} + #[cfg(test)] mod tests { use super::*; @@ -147,4 +167,31 @@ mod tests { "mismatched rustc version {want} and {got:?}" ); } + + #[test] + fn test_resolve_grpc_default() { + let extensions = Extensions::new(); + let resolved = resolve_grpc_header_value(&extensions, "default-header") + .expect("should resolve default header"); + assert_eq!( + resolved.to_str().expect("valid header string"), + "default-header" + ); + } + + #[test] + fn test_resolve_grpc_client_extension() { + let mut extensions = Extensions::new(); + extensions.insert(XGoogApiClient { + name: "storage", + library_type: GCCL, + version: "0.1.0", + }); + + let resolved = resolve_grpc_header_value(&extensions, "default-header") + .expect("should resolve client extension header"); + let val = resolved.to_str().expect("valid header string"); + assert!(val.contains("gccl/0.1.0"), "{val}"); + assert!(val.contains("grpc/"), "{val}"); + } } diff --git a/src/gax-internal/src/grpc/grpc_helpers.rs b/src/gax-internal/src/grpc/grpc_helpers.rs index e6f2485554..a2c03c62ed 100644 --- a/src/gax-internal/src/grpc/grpc_helpers.rs +++ b/src/gax-internal/src/grpc/grpc_helpers.rs @@ -108,7 +108,7 @@ pub(crate) fn make_headers( headers.insert( X_GOOG_API_CLIENT, - http::header::HeaderValue::from_static(api_client_header), + crate::api_header::resolve_grpc_header_value(extensions, api_client_header)?, ); if !request_params.is_empty() { @@ -619,6 +619,25 @@ mod tests { Ok(()) } + #[test] + fn make_headers_with_client_x_goog_api_client() -> TestResult { + let options = RequestOptions::default(); + let mut extensions = crate::options::Extensions::new(); + extensions.insert(crate::api_header::XGoogApiClient { + name: "spanner", + library_type: crate::api_header::GCCL, + version: "1.2.3", + }); + + let headers = make_headers(&extensions, API_CLIENT_HEADER, "", &options)?; + let val = headers.get(X_GOOG_API_CLIENT).unwrap().to_str().unwrap(); + assert!(val.contains("gccl/1.2.3"), "{val}"); + assert!(val.contains("grpc/"), "{val}"); + assert!(!val.contains("gapic/"), "{val}"); + + Ok(()) + } + use google_cloud_test_utils::test_layer::{AttributeValue, TestLayer}; use std::collections::HashMap; diff --git a/src/gax-internal/src/http.rs b/src/gax-internal/src/http.rs index 342d5480a5..bb23ccc26a 100644 --- a/src/gax-internal/src/http.rs +++ b/src/gax-internal/src/http.rs @@ -26,6 +26,7 @@ pub mod reqwest; use crate::as_inner::as_inner; use crate::attempt_info::AttemptInfo; +use crate::headers::{X_GOOG_API_CLIENT, X_GOOG_USER_PROJECT, sanitize_custom_headers}; use crate::observability::{HttpResultExt, RequestRecorder, create_http_attempt_span}; use crate::universe_domain::DEFAULT_UNIVERSE_DOMAIN; use ::reqwest::Url; @@ -56,8 +57,6 @@ use std::sync::Arc; use std::time::Duration; use tracing::Instrument; -use crate::headers::{X_GOOG_USER_PROJECT, sanitize_custom_headers}; - #[derive(Clone, Debug)] pub struct ReqwestClient { inner: ::reqwest::Client, @@ -74,6 +73,7 @@ pub struct ReqwestClient { _tracing_enabled: bool, universe_domain: String, transport_metric: Option, + extensions: crate::options::Extensions, } impl ReqwestClient { @@ -138,6 +138,7 @@ impl ReqwestClient { universe_domain, attempt_timeout: config.attempt_timeout, transport_metric: None, + extensions: config.extensions, }) } @@ -452,6 +453,13 @@ impl ReqwestClient { ); } + if let Some(h) = self.extensions.get::() { + headers.insert( + X_GOOG_API_CLIENT, + http::header::HeaderValue::from_str(&h.rest_header_value()).map_err(Error::ser)?, + ); + } + builder = builder.headers(headers); builder.build().map_err(map_send_error) diff --git a/src/gax-internal/tests/grpc_simple_request.rs b/src/gax-internal/tests/grpc_simple_request.rs index a35a2d5352..a3c22cee0e 100644 --- a/src/gax-internal/tests/grpc_simple_request.rs +++ b/src/gax-internal/tests/grpc_simple_request.rs @@ -634,4 +634,29 @@ mod tests { assert!(addresses.len() > 1, "{addresses:?}"); Ok(()) } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn api_client_header_client_extension_wire() -> anyhow::Result<()> { + let (endpoint, _server) = start_echo_server().await?; + + let client = builder(endpoint) + .with_credentials(test_credentials()) + .with_extension(google_cloud_gax_internal::api_header::XGoogApiClient { + name: "storage", + library_type: google_cloud_gax_internal::api_header::GCCL, + version: "1.0.0", + }) + .build() + .await?; + + let response = send_request(client, "test message", "").await?; + let header = response + .metadata + .get("x-goog-api-client") + .expect("header present"); + assert!(header.contains("gccl/1.0.0"), "{header}"); + assert!(header.contains("grpc/"), "{header}"); + assert!(!header.contains("gapic/"), "{header}"); + Ok(()) + } } diff --git a/src/gax-internal/tests/http_custom_header.rs b/src/gax-internal/tests/http_custom_header.rs index f8154cc129..e12c91e423 100644 --- a/src/gax-internal/tests/http_custom_header.rs +++ b/src/gax-internal/tests/http_custom_header.rs @@ -227,6 +227,95 @@ mod tests { Ok(()) } + #[tokio::test] + async fn api_client_header_default_builder_header_preserved() -> anyhow::Result<()> { + let (endpoint, _server) = echo_server::start().await?; + + let client = echo_server::builder(endpoint) + .with_credentials(Credentials::from(mock_credentials())) + .build() + .await?; + + let builder = client + .builder(reqwest::Method::GET, "/echo".into()) + .header("x-goog-api-client", "gapic/1.2.3"); + let options = RequestOptions::default(); + + let response: serde_json::Value = client + .execute(builder, Some(json!({})), options) + .await? + .into_body(); + + let val = get_header_value(&response, "x-goog-api-client").expect("header present"); + assert_eq!(val, "gapic/1.2.3"); + + Ok(()) + } + + #[tokio::test] + async fn api_client_header_client_extension() -> anyhow::Result<()> { + let (endpoint, _server) = echo_server::start().await?; + + let client = echo_server::builder(endpoint) + .with_credentials(Credentials::from(mock_credentials())) + .with_extension(google_cloud_gax_internal::api_header::XGoogApiClient { + name: "storage", + library_type: google_cloud_gax_internal::api_header::GCCL, + version: "1.0.0", + }) + .build() + .await?; + + let builder = client + .builder(reqwest::Method::GET, "/echo".into()) + .header("x-goog-api-client", "gapic/0.0.0 rest/1.2.3"); + let options = RequestOptions::default(); + + let response: serde_json::Value = client + .execute(builder, Some(json!({})), options) + .await? + .into_body(); + + let val = get_header_value(&response, "x-goog-api-client").expect("header present"); + assert!(val.contains("gccl/1.0.0"), "{val}"); + assert!(val.contains("rest/"), "{val}"); + assert!(!val.contains("gapic/"), "{val}"); + + Ok(()) + } + + #[tokio::test] + async fn custom_header_does_not_override_api_client() -> anyhow::Result<()> { + let (endpoint, _server) = echo_server::start().await?; + + let client = echo_server::builder(endpoint) + .with_credentials(Credentials::from(mock_credentials())) + .with_extension(google_cloud_gax_internal::api_header::XGoogApiClient { + name: "trusted-veneer", + library_type: google_cloud_gax_internal::api_header::GCCL, + version: "1.0.0", + }) + .build() + .await?; + + let builder = client + .builder(reqwest::Method::GET, "/echo".into()) + .header("x-goog-api-client", "gapic/0.0.0 rest/1.2.3"); + let mut options = RequestOptions::default(); + options = with_custom_header(options, "x-goog-api-client", "malicious-header"); + + let response: serde_json::Value = client + .execute(builder, Some(json!({})), options) + .await? + .into_body(); + + let val = get_header_value(&response, "x-goog-api-client").expect("header present"); + assert!(val.contains("gccl/1.0.0"), "{val}"); + assert!(!val.contains("malicious-header"), "{val}"); + + Ok(()) + } + fn get_header_value(response: &serde_json::Value, name: &str) -> Option { response .as_object()