feat: harden community production foundation through story 1.5
This commit is contained in:
@@ -1,16 +1,20 @@
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
env, io,
|
||||
io,
|
||||
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
|
||||
sync::Arc,
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use crank_core::{HttpMethod, RestTarget};
|
||||
use crank_core::{HttpMethod, RestTarget, RuntimeRequestContext};
|
||||
use crank_metrics::{UpstreamOperationKind, UpstreamOutcome, UpstreamRequestMetrics};
|
||||
use crank_trace::{ErrorCategory, Stage, StageOutcome};
|
||||
use futures_util::StreamExt;
|
||||
use opentelemetry::{global, propagation::Injector, trace::TraceContextExt};
|
||||
use opentelemetry::{
|
||||
Context, global,
|
||||
propagation::Injector,
|
||||
trace::{SpanContext, SpanId, TraceContextExt, TraceFlags, TraceId, TraceState},
|
||||
};
|
||||
use reqwest::{
|
||||
Client,
|
||||
dns::{Addrs, Name, Resolve, Resolving},
|
||||
@@ -49,10 +53,6 @@ impl RestAdapter {
|
||||
Self::with_policy(OutboundHttpPolicy::default())
|
||||
}
|
||||
|
||||
pub fn from_env() -> Result<Self, RestAdapterError> {
|
||||
Ok(Self::with_policy(OutboundHttpPolicy::from_env()?))
|
||||
}
|
||||
|
||||
pub fn with_policy(policy: OutboundHttpPolicy) -> Self {
|
||||
let resolver = Arc::new(PolicyDnsResolver {
|
||||
policy: policy.clone(),
|
||||
@@ -73,7 +73,23 @@ impl RestAdapter {
|
||||
request: &RestRequest,
|
||||
) -> Result<RestResponse, RestAdapterError> {
|
||||
let request_metrics = UpstreamRequestMetrics::start(UpstreamOperationKind::Rest);
|
||||
let result = self.execute_inner(target, request).await;
|
||||
let result = self.execute_inner(target, request, None).await;
|
||||
let outcome = match &result {
|
||||
Ok(_) => UpstreamOutcome::Success,
|
||||
Err(error) => upstream_outcome(error),
|
||||
};
|
||||
request_metrics.complete(outcome);
|
||||
result
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_with_context(
|
||||
&self,
|
||||
target: &RestTarget,
|
||||
request: &RestRequest,
|
||||
context: &RuntimeRequestContext,
|
||||
) -> Result<RestResponse, RestAdapterError> {
|
||||
let request_metrics = UpstreamRequestMetrics::start(UpstreamOperationKind::Rest);
|
||||
let result = self.execute_inner(target, request, Some(context)).await;
|
||||
let outcome = match &result {
|
||||
Ok(_) => UpstreamOutcome::Success,
|
||||
Err(error) => upstream_outcome(error),
|
||||
@@ -86,28 +102,48 @@ impl RestAdapter {
|
||||
&self,
|
||||
target: &RestTarget,
|
||||
request: &RestRequest,
|
||||
trusted_context: Option<&RuntimeRequestContext>,
|
||||
) -> Result<RestResponse, RestAdapterError> {
|
||||
let url = build_url(target, request)?;
|
||||
self.policy.validate_url(&url)?;
|
||||
let mut headers = build_headers(target, request)?;
|
||||
apply_current_trace_context(&mut headers);
|
||||
let client =
|
||||
self.client
|
||||
.as_ref()
|
||||
.map_err(|details| RestAdapterError::InvalidConfiguration {
|
||||
details: details.to_string(),
|
||||
})?;
|
||||
let mut builder = client
|
||||
.request(to_reqwest_method(target.method), url)
|
||||
.headers(headers)
|
||||
.timeout(Duration::from_millis(request.timeout_ms));
|
||||
|
||||
if let Some(body) = &request.body {
|
||||
builder = builder.json(body);
|
||||
}
|
||||
|
||||
let upstream_span = Stage::UpstreamHttp.span();
|
||||
if let Some(context) = trusted_context {
|
||||
set_span_parent_from_traceparent(&upstream_span, context.trace_context.traceparent());
|
||||
}
|
||||
let result = async {
|
||||
let url = build_url(target, request)?;
|
||||
self.policy.validate_url(&url)?;
|
||||
let mut headers = build_headers(target, request)?;
|
||||
if let Some(context) = trusted_context {
|
||||
for (name, value) in context.outbound_headers() {
|
||||
let (Ok(name), Ok(value)) =
|
||||
(HeaderName::try_from(name), HeaderValue::try_from(value))
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
headers.insert(name, value);
|
||||
}
|
||||
}
|
||||
apply_current_trace_context(&mut headers);
|
||||
if !headers.contains_key("traceparent")
|
||||
&& let Some(context) = trusted_context
|
||||
&& let Ok(value) = HeaderValue::from_str(context.trace_context.traceparent())
|
||||
{
|
||||
headers.insert("traceparent", value);
|
||||
}
|
||||
let client =
|
||||
self.client
|
||||
.as_ref()
|
||||
.map_err(|details| RestAdapterError::InvalidConfiguration {
|
||||
details: details.to_string(),
|
||||
})?;
|
||||
let mut builder = client
|
||||
.request(to_reqwest_method(target.method), url)
|
||||
.headers(headers)
|
||||
.timeout(Duration::from_millis(request.timeout_ms));
|
||||
|
||||
if let Some(body) = &request.body {
|
||||
builder = builder.json(body);
|
||||
}
|
||||
|
||||
let response = builder.send().await?;
|
||||
let status = response.status();
|
||||
let headers = normalize_headers(response.headers());
|
||||
@@ -174,32 +210,19 @@ impl Default for OutboundHttpPolicy {
|
||||
}
|
||||
|
||||
impl OutboundHttpPolicy {
|
||||
pub fn from_env() -> Result<Self, RestAdapterError> {
|
||||
let max_response_bytes = match env::var("CRANK_OUTBOUND_MAX_RESPONSE_BYTES") {
|
||||
Ok(value) => {
|
||||
value
|
||||
.parse::<usize>()
|
||||
.map_err(|_| RestAdapterError::InvalidConfiguration {
|
||||
details: "CRANK_OUTBOUND_MAX_RESPONSE_BYTES must be a positive integer"
|
||||
.to_owned(),
|
||||
})?
|
||||
}
|
||||
Err(env::VarError::NotPresent) => DEFAULT_MAX_RESPONSE_BYTES,
|
||||
Err(error) => {
|
||||
return Err(RestAdapterError::InvalidConfiguration {
|
||||
details: error.to_string(),
|
||||
});
|
||||
}
|
||||
};
|
||||
pub fn try_new(
|
||||
allowed_hosts: Vec<String>,
|
||||
denied_hosts: Vec<String>,
|
||||
max_response_bytes: usize,
|
||||
) -> Result<Self, RestAdapterError> {
|
||||
if max_response_bytes == 0 {
|
||||
return Err(RestAdapterError::InvalidConfiguration {
|
||||
details: "CRANK_OUTBOUND_MAX_RESPONSE_BYTES must be greater than zero".to_owned(),
|
||||
details: "outbound response limit must be greater than zero".to_owned(),
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
allowed_hosts: host_patterns_from_env("CRANK_OUTBOUND_ALLOWED_HOSTS")?,
|
||||
denied_hosts: host_patterns_from_env("CRANK_OUTBOUND_DENIED_HOSTS")?,
|
||||
allowed_hosts: validate_host_patterns(allowed_hosts)?,
|
||||
denied_hosts: validate_host_patterns(denied_hosts)?,
|
||||
max_response_bytes,
|
||||
})
|
||||
}
|
||||
@@ -308,20 +331,11 @@ fn boxed_io_error(message: String) -> Box<dyn std::error::Error + Send + Sync> {
|
||||
Box::new(io::Error::new(io::ErrorKind::PermissionDenied, message))
|
||||
}
|
||||
|
||||
fn host_patterns_from_env(name: &str) -> Result<Vec<String>, RestAdapterError> {
|
||||
let value = match env::var(name) {
|
||||
Ok(value) => value,
|
||||
Err(env::VarError::NotPresent) => return Ok(Vec::new()),
|
||||
Err(error) => {
|
||||
return Err(RestAdapterError::InvalidConfiguration {
|
||||
details: error.to_string(),
|
||||
});
|
||||
}
|
||||
};
|
||||
value
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
fn validate_host_patterns(
|
||||
values: impl IntoIterator<Item = String>,
|
||||
) -> Result<Vec<String>, RestAdapterError> {
|
||||
values
|
||||
.into_iter()
|
||||
.map(|value| {
|
||||
let wildcard = value.starts_with("*.");
|
||||
let normalized = normalize_host(value.trim_start_matches("*."));
|
||||
@@ -332,7 +346,7 @@ fn host_patterns_from_env(name: &str) -> Result<Vec<String>, RestAdapterError> {
|
||||
|| (wildcard && normalized.parse::<IpAddr>().is_ok())
|
||||
{
|
||||
return Err(RestAdapterError::InvalidConfiguration {
|
||||
details: format!("{name} contains an invalid host pattern: {value}"),
|
||||
details: "outbound host pattern is invalid".to_owned(),
|
||||
});
|
||||
}
|
||||
Ok(if wildcard {
|
||||
@@ -465,7 +479,7 @@ fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<(),
|
||||
HeaderName::try_from(name).map_err(|_| RestAdapterError::InvalidHeaderName {
|
||||
header: name.to_owned(),
|
||||
})?;
|
||||
if is_trace_propagation_header(&header_name) {
|
||||
if is_reserved_correlation_header(&header_name) {
|
||||
return Ok(());
|
||||
}
|
||||
let header_value =
|
||||
@@ -477,12 +491,53 @@ fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<(),
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn is_trace_propagation_header(name: &HeaderName) -> bool {
|
||||
matches!(name.as_str(), "traceparent" | "tracestate" | "baggage")
|
||||
fn is_reserved_correlation_header(name: &HeaderName) -> bool {
|
||||
matches!(
|
||||
name.as_str(),
|
||||
"traceparent"
|
||||
| "tracestate"
|
||||
| "baggage"
|
||||
| "x-request-id"
|
||||
| "x-trace-id"
|
||||
| "x-correlation-id"
|
||||
)
|
||||
}
|
||||
|
||||
fn set_span_parent_from_traceparent(span: &Span, traceparent: &str) -> bool {
|
||||
let mut parts = traceparent.split('-');
|
||||
let (Some("00"), Some(trace_id), Some(parent_id), Some(flags), None) = (
|
||||
parts.next(),
|
||||
parts.next(),
|
||||
parts.next(),
|
||||
parts.next(),
|
||||
parts.next(),
|
||||
) else {
|
||||
return false;
|
||||
};
|
||||
let (Ok(trace_id), Ok(parent_id)) = (TraceId::from_hex(trace_id), SpanId::from_hex(parent_id))
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
let trace_flags = if flags == "01" {
|
||||
TraceFlags::SAMPLED
|
||||
} else if flags == "00" {
|
||||
TraceFlags::default()
|
||||
} else {
|
||||
return false;
|
||||
};
|
||||
let parent = SpanContext::new(
|
||||
trace_id,
|
||||
parent_id,
|
||||
trace_flags,
|
||||
true,
|
||||
TraceState::default(),
|
||||
);
|
||||
span.set_parent(Context::new().with_remote_span_context(parent))
|
||||
.is_ok()
|
||||
}
|
||||
|
||||
fn apply_current_trace_context(headers: &mut HeaderMap) {
|
||||
for header in ["traceparent", "tracestate", "baggage"] {
|
||||
for header in ["tracestate", "baggage"] {
|
||||
headers.remove(header);
|
||||
}
|
||||
|
||||
@@ -490,6 +545,7 @@ fn apply_current_trace_context(headers: &mut HeaderMap) {
|
||||
if !context.span().span_context().is_valid() {
|
||||
return;
|
||||
}
|
||||
headers.remove("traceparent");
|
||||
global::get_text_map_propagator(|propagator| {
|
||||
propagator.inject_context(&context, &mut ReqwestHeaderInjector(headers));
|
||||
});
|
||||
|
||||
@@ -29,16 +29,14 @@ impl ProtocolAdapter for RestAdapter {
|
||||
context: &RuntimeRequestContext,
|
||||
) -> Result<AdapterResponse, ProtocolAdapterError> {
|
||||
let target = rest_target(target)?;
|
||||
let mut headers = prepared.headers.clone();
|
||||
headers.extend(context.outbound_headers());
|
||||
let request = RestRequest {
|
||||
path_params: prepared.path_params.clone(),
|
||||
query_params: prepared.query_params.clone(),
|
||||
headers,
|
||||
headers: prepared.headers.clone(),
|
||||
body: prepared.body.clone(),
|
||||
timeout_ms: prepared.timeout_ms,
|
||||
};
|
||||
let response = self.execute(target, &request).await?;
|
||||
let response = self.execute_with_context(target, &request, context).await?;
|
||||
|
||||
Ok(AdapterResponse {
|
||||
status_code: response.status_code,
|
||||
|
||||
@@ -48,7 +48,7 @@ async fn executes_rest_request_and_normalizes_json_response() {
|
||||
json!({
|
||||
"id": "42",
|
||||
"query": "true",
|
||||
"trace": "trace-123",
|
||||
"trace": "",
|
||||
"static": "static",
|
||||
"payload": { "name": "Ada" }
|
||||
})
|
||||
@@ -69,6 +69,7 @@ async fn protocol_context_overrides_mapped_correlation_headers() {
|
||||
"x-correlation-id".to_owned(),
|
||||
"static-correlation".to_owned(),
|
||||
),
|
||||
("x-trace-id".to_owned(), "static-trace".to_owned()),
|
||||
]),
|
||||
});
|
||||
let prepared = PreparedRequest {
|
||||
@@ -79,12 +80,17 @@ async fn protocol_context_overrides_mapped_correlation_headers() {
|
||||
"x-correlation-id".to_owned(),
|
||||
"mapped-correlation".to_owned(),
|
||||
),
|
||||
("x-trace-id".to_owned(), "mapped-trace".to_owned()),
|
||||
]),
|
||||
body: Some(json!({ "name": "Ada" })),
|
||||
timeout_ms: 1_000,
|
||||
..PreparedRequest::default()
|
||||
};
|
||||
let context = RuntimeRequestContext::new("req-runtime", "corr-runtime");
|
||||
let context = RuntimeRequestContext::new(
|
||||
crank_core::RequestId::resolve(Some("req-runtime")),
|
||||
crank_core::TraceContext::parse("00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01")
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let response = adapter
|
||||
.invoke_unary(&target, &prepared, &context)
|
||||
@@ -92,7 +98,12 @@ async fn protocol_context_overrides_mapped_correlation_headers() {
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.body["request_id"], "req-runtime");
|
||||
assert_eq!(response.body["correlation_id"], "corr-runtime");
|
||||
assert_eq!(response.body["correlation_id"], "req-runtime");
|
||||
assert_eq!(response.body["trace"], "0af7651916cd43dd8448eb211c80319c");
|
||||
assert_eq!(
|
||||
&response.body["traceparent"].as_str().unwrap()[3..35],
|
||||
"0af7651916cd43dd8448eb211c80319c"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
|
||||
Reference in New Issue
Block a user