383 lines
12 KiB
Rust
383 lines
12 KiB
Rust
use std::collections::BTreeMap;
|
|
|
|
use axum::{
|
|
Json, Router,
|
|
extract::{Path, Query},
|
|
http::{HeaderMap, StatusCode},
|
|
response::Redirect,
|
|
routing::{get, post},
|
|
};
|
|
use crank_adapter_rest::{OutboundHttpPolicy, RestAdapter, RestAdapterError, RestRequest};
|
|
use crank_core::{
|
|
HttpMethod, PreparedRequest, ProtocolAdapter, RestTarget, RuntimeRequestContext, Target,
|
|
};
|
|
use opentelemetry::{
|
|
global,
|
|
trace::{TraceContextExt, TracerProvider as _},
|
|
};
|
|
use opentelemetry_sdk::{propagation::TraceContextPropagator, trace::SdkTracerProvider};
|
|
use serde_json::{Value, json};
|
|
use tokio::net::TcpListener;
|
|
use tracing::Instrument;
|
|
use tracing_opentelemetry::OpenTelemetrySpanExt;
|
|
use tracing_subscriber::layer::SubscriberExt;
|
|
|
|
#[tokio::test]
|
|
async fn executes_rest_request_and_normalizes_json_response() {
|
|
let base_url = spawn_test_server().await;
|
|
let adapter = test_adapter();
|
|
let target = RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Post,
|
|
path_template: "/users/{user_id}".to_owned(),
|
|
static_headers: BTreeMap::from([("x-static".to_owned(), "static".to_owned())]),
|
|
};
|
|
let request = RestRequest {
|
|
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
|
query_params: BTreeMap::from([("expand".to_owned(), "true".to_owned())]),
|
|
headers: BTreeMap::from([("x-trace-id".to_owned(), "trace-123".to_owned())]),
|
|
body: Some(json!({ "name": "Ada" })),
|
|
timeout_ms: 1_000,
|
|
};
|
|
|
|
let response = adapter.execute(&target, &request).await.unwrap();
|
|
|
|
assert_eq!(response.status_code, 200);
|
|
assert_eq!(
|
|
response.body,
|
|
json!({
|
|
"id": "42",
|
|
"query": "true",
|
|
"trace": "",
|
|
"static": "static",
|
|
"payload": { "name": "Ada" }
|
|
})
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn protocol_context_overrides_mapped_correlation_headers() {
|
|
let base_url = spawn_test_server().await;
|
|
let adapter = test_adapter();
|
|
let target = Target::Rest(RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Post,
|
|
path_template: "/users/{user_id}".to_owned(),
|
|
static_headers: BTreeMap::from([
|
|
("x-request-id".to_owned(), "static-request".to_owned()),
|
|
(
|
|
"x-correlation-id".to_owned(),
|
|
"static-correlation".to_owned(),
|
|
),
|
|
("x-trace-id".to_owned(), "static-trace".to_owned()),
|
|
]),
|
|
});
|
|
let prepared = PreparedRequest {
|
|
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
|
headers: BTreeMap::from([
|
|
("x-request-id".to_owned(), "mapped-request".to_owned()),
|
|
(
|
|
"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(
|
|
crank_core::RequestId::resolve(Some("req-runtime")),
|
|
crank_core::TraceContext::parse("00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01")
|
|
.unwrap(),
|
|
);
|
|
|
|
let response = adapter
|
|
.invoke_unary(&target, &prepared, &context)
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(response.body["request_id"], "req-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")]
|
|
async fn current_trace_context_overrides_mapped_traceparent() {
|
|
global::set_text_map_propagator(TraceContextPropagator::new());
|
|
let provider = SdkTracerProvider::builder().build();
|
|
let tracer = provider.tracer("rest-propagation-test");
|
|
let subscriber =
|
|
tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer));
|
|
let dispatch = tracing::Dispatch::new(subscriber);
|
|
let _dispatch_guard = tracing::dispatcher::set_default(&dispatch);
|
|
let span = tracing::info_span!("runtime.execute");
|
|
let context = span.context();
|
|
let expected_trace_id = context.span().span_context().trace_id().to_string();
|
|
let base_url = spawn_test_server().await;
|
|
let target = RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Post,
|
|
path_template: "/users/{user_id}".to_owned(),
|
|
static_headers: BTreeMap::new(),
|
|
};
|
|
let request = RestRequest {
|
|
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
|
query_params: BTreeMap::new(),
|
|
headers: BTreeMap::from([(
|
|
"traceparent".to_owned(),
|
|
"00-11111111111111111111111111111111-2222222222222222-01".to_owned(),
|
|
)]),
|
|
body: Some(json!({ "name": "Ada" })),
|
|
timeout_ms: 1_000,
|
|
};
|
|
|
|
let response = test_adapter()
|
|
.execute(&target, &request)
|
|
.instrument(span)
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(
|
|
&response.body["traceparent"].as_str().unwrap()[3..35],
|
|
expected_trace_id
|
|
);
|
|
provider.shutdown().unwrap();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn user_configured_propagation_headers_are_removed_without_trusted_context() {
|
|
let base_url = spawn_test_server().await;
|
|
let target = RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Post,
|
|
path_template: "/users/{user_id}".to_owned(),
|
|
static_headers: BTreeMap::from([
|
|
(
|
|
"traceparent".to_owned(),
|
|
"untrusted\ninvalid-value".to_owned(),
|
|
),
|
|
("tracestate".to_owned(), "vendor=value".to_owned()),
|
|
("baggage".to_owned(), "secret=must-not-leave".to_owned()),
|
|
]),
|
|
};
|
|
let request = RestRequest {
|
|
path_params: BTreeMap::from([("user_id".to_owned(), "42".to_owned())]),
|
|
query_params: BTreeMap::new(),
|
|
headers: BTreeMap::new(),
|
|
body: Some(json!({ "name": "Ada" })),
|
|
timeout_ms: 1_000,
|
|
};
|
|
|
|
let response = test_adapter().execute(&target, &request).await.unwrap();
|
|
|
|
assert!(response.body.get("traceparent").is_none());
|
|
assert!(response.body.get("tracestate").is_none());
|
|
assert!(response.body.get("baggage").is_none());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn returns_unexpected_status_with_normalized_body() {
|
|
let base_url = spawn_test_server().await;
|
|
let adapter = test_adapter();
|
|
let target = RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Get,
|
|
path_template: "/fail".to_owned(),
|
|
static_headers: BTreeMap::new(),
|
|
};
|
|
let request = RestRequest {
|
|
path_params: BTreeMap::new(),
|
|
query_params: BTreeMap::new(),
|
|
headers: BTreeMap::new(),
|
|
body: None,
|
|
timeout_ms: 1_000,
|
|
};
|
|
|
|
let error = adapter.execute(&target, &request).await.unwrap_err();
|
|
|
|
assert!(matches!(
|
|
error,
|
|
RestAdapterError::UnexpectedStatus {
|
|
status: 502,
|
|
body: Value::Object(_)
|
|
}
|
|
));
|
|
}
|
|
|
|
#[test]
|
|
fn rejects_private_targets_by_default() {
|
|
let policy = OutboundHttpPolicy::default();
|
|
|
|
let error = policy
|
|
.validate_base_url("http://127.0.0.1:8080")
|
|
.unwrap_err();
|
|
|
|
assert!(matches!(error, RestAdapterError::TargetNotAllowed { .. }));
|
|
}
|
|
|
|
#[test]
|
|
fn accepts_explicit_private_ipv4_and_ipv6_targets() {
|
|
let policy = OutboundHttpPolicy::allowing_hosts(["192.168.1.10", "::1"]);
|
|
|
|
assert!(policy.validate_base_url("http://192.168.1.10:8080").is_ok());
|
|
assert!(policy.validate_base_url("http://[::1]:8080").is_ok());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn does_not_follow_redirects() {
|
|
let base_url = spawn_test_server().await;
|
|
let adapter = test_adapter();
|
|
let target = RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Get,
|
|
path_template: "/redirect".to_owned(),
|
|
static_headers: BTreeMap::new(),
|
|
};
|
|
|
|
let error = adapter
|
|
.execute(&target, &empty_request())
|
|
.await
|
|
.unwrap_err();
|
|
|
|
assert!(matches!(
|
|
error,
|
|
RestAdapterError::UnexpectedStatus { status: 303, .. }
|
|
));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn rejects_responses_over_the_configured_limit() {
|
|
let base_url = spawn_test_server().await;
|
|
let adapter = RestAdapter::with_policy(
|
|
OutboundHttpPolicy::allowing_hosts(["127.0.0.1"]).with_max_response_bytes(8),
|
|
);
|
|
let target = RestTarget {
|
|
base_url,
|
|
method: HttpMethod::Get,
|
|
path_template: "/large".to_owned(),
|
|
static_headers: BTreeMap::new(),
|
|
};
|
|
|
|
let error = adapter
|
|
.execute(&target, &empty_request())
|
|
.await
|
|
.unwrap_err();
|
|
|
|
assert!(matches!(
|
|
error,
|
|
RestAdapterError::ResponseTooLarge { limit_bytes: 8 }
|
|
));
|
|
}
|
|
|
|
fn empty_request() -> RestRequest {
|
|
RestRequest {
|
|
path_params: BTreeMap::new(),
|
|
query_params: BTreeMap::new(),
|
|
headers: BTreeMap::new(),
|
|
body: None,
|
|
timeout_ms: 1_000,
|
|
}
|
|
}
|
|
|
|
fn test_adapter() -> RestAdapter {
|
|
RestAdapter::with_policy(OutboundHttpPolicy::allowing_hosts(["127.0.0.1"]))
|
|
}
|
|
|
|
async fn spawn_test_server() -> String {
|
|
let app = Router::new()
|
|
.route("/users/{user_id}", post(create_user))
|
|
.route("/fail", get(fail))
|
|
.route("/redirect", get(|| async { Redirect::to("/large") }))
|
|
.route(
|
|
"/large",
|
|
get(|| async { "response larger than eight bytes" }),
|
|
);
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let address = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async move {
|
|
axum::serve(listener, app).await.unwrap();
|
|
});
|
|
|
|
format!("http://{}", address)
|
|
}
|
|
|
|
async fn create_user(
|
|
Path(user_id): Path<String>,
|
|
Query(query): Query<BTreeMap<String, String>>,
|
|
headers: HeaderMap,
|
|
Json(payload): Json<Value>,
|
|
) -> Json<Value> {
|
|
let trace = headers
|
|
.get("x-trace-id")
|
|
.and_then(|value| value.to_str().ok())
|
|
.unwrap_or_default();
|
|
let static_header = headers
|
|
.get("x-static")
|
|
.and_then(|value| value.to_str().ok())
|
|
.unwrap_or_default();
|
|
let request_id = headers
|
|
.get("x-request-id")
|
|
.and_then(|value| value.to_str().ok());
|
|
let correlation_id = headers
|
|
.get("x-correlation-id")
|
|
.and_then(|value| value.to_str().ok());
|
|
let traceparent = headers
|
|
.get("traceparent")
|
|
.and_then(|value| value.to_str().ok());
|
|
let tracestate = headers
|
|
.get("tracestate")
|
|
.and_then(|value| value.to_str().ok());
|
|
let baggage = headers.get("baggage").and_then(|value| value.to_str().ok());
|
|
|
|
let mut response = json!({
|
|
"id": user_id,
|
|
"query": query.get("expand").cloned().unwrap_or_default(),
|
|
"trace": trace,
|
|
"static": static_header,
|
|
"payload": payload
|
|
});
|
|
let response = response.as_object_mut().unwrap();
|
|
if let Some(request_id) = request_id {
|
|
response.insert(
|
|
"request_id".to_owned(),
|
|
Value::String(request_id.to_owned()),
|
|
);
|
|
}
|
|
if let Some(correlation_id) = correlation_id {
|
|
response.insert(
|
|
"correlation_id".to_owned(),
|
|
Value::String(correlation_id.to_owned()),
|
|
);
|
|
}
|
|
if let Some(traceparent) = traceparent {
|
|
response.insert(
|
|
"traceparent".to_owned(),
|
|
Value::String(traceparent.to_owned()),
|
|
);
|
|
}
|
|
if let Some(tracestate) = tracestate {
|
|
response.insert(
|
|
"tracestate".to_owned(),
|
|
Value::String(tracestate.to_owned()),
|
|
);
|
|
}
|
|
if let Some(baggage) = baggage {
|
|
response.insert("baggage".to_owned(), Value::String(baggage.to_owned()));
|
|
}
|
|
|
|
Json(Value::Object(response.clone()))
|
|
}
|
|
|
|
async fn fail() -> (axum::http::StatusCode, Json<Value>) {
|
|
(
|
|
StatusCode::BAD_GATEWAY,
|
|
Json(json!({ "error": "upstream failed" })),
|
|
)
|
|
}
|