Skip to content
47 changes: 47 additions & 0 deletions src/gax-internal/src/api_header.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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<HeaderValue, Error> {
if let Some(h) = extensions.get::<XGoogApiClient>() {
return HeaderValue::from_str(&h.grpc_header_value()).map_err(Error::ser);
}
Ok(HeaderValue::from_static(default_header))
}

#[cfg(test)]
mod tests {
use super::*;
Expand Down Expand Up @@ -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}");
}
}
21 changes: 20 additions & 1 deletion src/gax-internal/src/grpc/grpc_helpers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand Down Expand Up @@ -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;

Expand Down
12 changes: 10 additions & 2 deletions src/gax-internal/src/http.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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,
Expand All @@ -74,6 +73,7 @@ pub struct ReqwestClient {
_tracing_enabled: bool,
universe_domain: String,
transport_metric: Option<crate::observability::TransportMetric>,
extensions: crate::options::Extensions,
}

impl ReqwestClient {
Expand Down Expand Up @@ -138,6 +138,7 @@ impl ReqwestClient {
universe_domain,
attempt_timeout: config.attempt_timeout,
transport_metric: None,
extensions: config.extensions,
})
}

Expand Down Expand Up @@ -452,6 +453,13 @@ impl ReqwestClient {
);
}

if let Some(h) = self.extensions.get::<crate::api_header::XGoogApiClient>() {
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)
Expand Down
25 changes: 25 additions & 0 deletions src/gax-internal/tests/grpc_simple_request.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(())
}
}
89 changes: 89 additions & 0 deletions src/gax-internal/tests/http_custom_header.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<String> {
response
.as_object()
Expand Down
Loading