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": "trace-123", "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(), ), ]), }); 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(), ), ]), body: Some(json!({ "name": "Ada" })), timeout_ms: 1_000, ..PreparedRequest::default() }; let context = RuntimeRequestContext::new("req-runtime", "corr-runtime"); let response = adapter .invoke_unary(&target, &prepared, &context) .await .unwrap(); assert_eq!(response.body["request_id"], "req-runtime"); assert_eq!(response.body["correlation_id"], "corr-runtime"); } #[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, Query(query): Query>, headers: HeaderMap, Json(payload): Json, ) -> Json { 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) { ( StatusCode::BAD_GATEWAY, Json(json!({ "error": "upstream failed" })), ) }