diff --git a/apps/admin-api/src/service.rs b/apps/admin-api/src/service.rs index 95d6394..402e7d8 100644 --- a/apps/admin-api/src/service.rs +++ b/apps/admin-api/src/service.rs @@ -33,7 +33,8 @@ use crank_registry::{ UsageSummary, UsageTimelinePoint, WorkspaceMembershipRecord, WorkspaceRecord, }; use crank_runtime::{ - PreparedRequest, ResolvedAuth, RuntimeError, RuntimeExecutor, RuntimeOperation, SecretCrypto, + PreparedRequest, ResolvedAuth, RuntimeError, RuntimeExecutor, RuntimeOperation, + RuntimeRequestContext, SecretCrypto, }; use crank_schema::Schema; use serde::{Deserialize, Serialize}; @@ -2096,6 +2097,7 @@ impl AdminService { payload: TestRunPayload, request_id: &str, ) -> Result { + let runtime_request_context = RuntimeRequestContext::from_request_id(request_id); let record = self .get_operation_version(workspace_id, operation_id, payload.version) .await?; @@ -2149,12 +2151,22 @@ impl AdminService { Ok(resolved_auth) => match mode { ExecutionMode::Unary => self .runtime - .execute_with_auth(&runtime, &payload.input, resolved_auth.as_ref()) + .execute_with_auth_and_context( + &runtime, + &payload.input, + resolved_auth.as_ref(), + Some(&runtime_request_context), + ) .await .map(|output| TestRunOutcome::Unary { output }), ExecutionMode::Window => self .runtime - .execute_window_with_auth(&runtime, &payload.input, resolved_auth.as_ref()) + .execute_window_with_auth_and_context( + &runtime, + &payload.input, + resolved_auth.as_ref(), + Some(&runtime_request_context), + ) .await .map(|output| TestRunOutcome::Window { output }), ExecutionMode::Session => self @@ -2164,11 +2176,18 @@ impl AdminService { &runtime, &payload.input, resolved_auth.as_ref(), + &runtime_request_context, ) .await .map(|output| TestRunOutcome::Session { output }), ExecutionMode::AsyncJob => self - .start_async_job_test(workspace_id, &record.snapshot, &runtime, &payload.input) + .start_async_job_test( + workspace_id, + &record.snapshot, + &runtime, + &payload.input, + &runtime_request_context, + ) .await .map(|output| TestRunOutcome::AsyncJob { output }), }, @@ -2272,6 +2291,7 @@ impl AdminService { runtime: &RuntimeOperation, input: &Value, resolved_auth: Option<&ResolvedAuth>, + request_context: &RuntimeRequestContext, ) -> Result { let streaming = runtime.execution_config.streaming.as_ref().ok_or_else(|| { RuntimeError::MissingStreamingConfig { @@ -2280,7 +2300,12 @@ impl AdminService { })?; let seed = self .runtime - .execute_session_seed_with_auth(runtime, input, resolved_auth) + .execute_session_seed_with_auth_and_context( + runtime, + input, + resolved_auth, + Some(request_context), + ) .await?; let batch_size = streaming.max_items.unwrap_or(10).max(1) as usize; let preview_count = seed.items.len().min(batch_size); @@ -2342,6 +2367,7 @@ impl AdminService { operation: &RegistryOperation, runtime: &RuntimeOperation, input: &Value, + request_context: &RuntimeRequestContext, ) -> Result { let created_at = OffsetDateTime::now_utc(); let streaming = runtime.execution_config.streaming.as_ref().ok_or_else(|| { @@ -2385,6 +2411,7 @@ impl AdminService { let workspace_for_task = workspace_id.clone(); let operation_for_task = runtime.clone(); let input_for_task = input.clone(); + let request_context_for_task = request_context.clone(); let job_id = job.id.clone(); tokio::spawn(async move { let resolved_auth = resolve_runtime_auth_for_task( @@ -2397,10 +2424,11 @@ impl AdminService { let result = match resolved_auth { Ok(resolved_auth) => { runtime_for_task - .execute_with_auth( + .execute_with_auth_and_context( &operation_for_task, &input_for_task, resolved_auth.as_ref(), + Some(&request_context_for_task), ) .await } diff --git a/apps/mcp-server/src/app.rs b/apps/mcp-server/src/app.rs index c985be9..57b072b 100644 --- a/apps/mcp-server/src/app.rs +++ b/apps/mcp-server/src/app.rs @@ -29,7 +29,10 @@ use crank_registry::{ PostgresRegistry, PublishedAgentTool, UpdateAsyncJobStatusRequest, UpdateStreamSessionStateRequest, }; -use crank_runtime::{ResolvedAuth, RuntimeError, RuntimeExecutor, RuntimeOperation, SecretCrypto}; +use crank_runtime::{ + ResolvedAuth, RuntimeError, RuntimeExecutor, RuntimeOperation, RuntimeRequestContext, + SecretCrypto, +}; use futures_util::stream; use serde::{Deserialize, Serialize}; use serde_json::{Value, json}; @@ -628,6 +631,7 @@ async fn handle_base_tool_call( transport_request_id: &str, ) -> Response { let operation = runtime_operation(&tool); + let runtime_request_context = RuntimeRequestContext::from_request_id(transport_request_id); let request_preview = build_request_preview(&state.runtime, &operation, &arguments); let started_at = Instant::now(); let is_window_mode = matches!( @@ -644,7 +648,12 @@ async fn handle_base_tool_call( let result = match resolved_auth { Ok(resolved_auth) if is_window_mode => state .runtime - .execute_window_with_auth(&operation, &arguments, resolved_auth.as_ref()) + .execute_window_with_auth_and_context( + &operation, + &arguments, + resolved_auth.as_ref(), + Some(&runtime_request_context), + ) .await .map(|output| { json!({ @@ -659,7 +668,12 @@ async fn handle_base_tool_call( Ok(resolved_auth) => { state .runtime - .execute_with_auth(&operation, &arguments, resolved_auth.as_ref()) + .execute_with_auth_and_context( + &operation, + &arguments, + resolved_auth.as_ref(), + Some(&runtime_request_context), + ) .await } Err(error) => Err(error), @@ -728,6 +742,7 @@ async fn handle_session_start_call( transport_request_id: &str, ) -> Response { let runtime_operation = runtime_operation(&tool); + let runtime_request_context = RuntimeRequestContext::from_request_id(transport_request_id); let request_preview = build_request_preview(&state.runtime, &runtime_operation, &arguments); let started_at = Instant::now(); let Some(streaming) = runtime_operation.execution_config.streaming.as_ref() else { @@ -751,10 +766,11 @@ async fn handle_session_start_call( Ok(resolved_auth) => { state .runtime - .execute_session_seed_with_auth( + .execute_session_seed_with_auth_and_context( &runtime_operation, &arguments, resolved_auth.as_ref(), + Some(&runtime_request_context), ) .await } @@ -1135,6 +1151,7 @@ async fn handle_async_job_start_call( transport_request_id: &str, ) -> Response { let operation = runtime_operation(&tool); + let runtime_request_context = RuntimeRequestContext::from_request_id(transport_request_id); let request_preview = build_request_preview(&state.runtime, &operation, &arguments); let Some(streaming) = tool.operation.execution_config.streaming.as_ref() else { return tool_error_response( @@ -1178,6 +1195,7 @@ async fn handle_async_job_start_call( let secret_crypto = state.secret_crypto.clone(); let tool_for_task = tool.clone(); let arguments_for_task = arguments.clone(); + let request_context_for_task = runtime_request_context.clone(); let job_id = job.id.clone(); tokio::spawn(async move { let runtime = RuntimeExecutor::new(); @@ -1192,7 +1210,12 @@ async fn handle_async_job_start_call( let result = match resolved_auth { Ok(resolved_auth) => { runtime - .execute_with_auth(&task_operation, &arguments_for_task, resolved_auth.as_ref()) + .execute_with_auth_and_context( + &task_operation, + &arguments_for_task, + resolved_auth.as_ref(), + Some(&request_context_for_task), + ) .await } Err(error) => Err(error), diff --git a/crates/crank-adapter-grpc/src/test_support.rs b/crates/crank-adapter-grpc/src/test_support.rs index c0cae13..db4372a 100644 --- a/crates/crank-adapter-grpc/src/test_support.rs +++ b/crates/crank-adapter-grpc/src/test_support.rs @@ -65,6 +65,62 @@ pub async fn spawn_unary_echo_server() -> String { format!("http://{}", address) } +pub async fn spawn_metadata_echo_server() -> String { + #[derive(Default)] + struct MetadataEchoServiceImpl; + + #[tonic::async_trait] + impl echo::echo_service_server::EchoService for MetadataEchoServiceImpl { + async fn unary_echo( + &self, + request: Request, + ) -> Result, Status> { + let message = request.get_ref().message.clone(); + let request_id = request + .metadata() + .get("x-request-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + let correlation_id = request + .metadata() + .get("x-correlation-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + + Ok(Response::new(echo::EchoResponse { + message: format!("{}|{}|{}", message, request_id, correlation_id), + })) + } + + type ServerEchoStream = std::pin::Pin< + Box> + Send>, + >; + + async fn server_echo( + &self, + _request: Request, + ) -> Result, Status> { + Ok(Response::new(Box::pin(stream::empty()))) + } + } + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let incoming = tonic::transport::server::TcpIncoming::from(listener); + + tokio::spawn(async move { + Server::builder() + .add_service(echo::echo_service_server::EchoServiceServer::new( + MetadataEchoServiceImpl, + )) + .serve_with_incoming(incoming) + .await + .unwrap(); + }); + + format!("http://{}", address) +} + pub fn descriptor_set_b64() -> String { STANDARD.encode(echo::FILE_DESCRIPTOR_SET) } diff --git a/crates/crank-runtime/src/executor.rs b/crates/crank-runtime/src/executor.rs index a373ec3..f1a79eb 100644 --- a/crates/crank-runtime/src/executor.rs +++ b/crates/crank-runtime/src/executor.rs @@ -10,7 +10,7 @@ use serde_json::{Map, Value, json}; use crate::{ AdapterResponse, PreparedRequest, ResolvedAuth, RuntimeError, RuntimeOperation, - WindowExecutionResult, + RuntimeRequestContext, WindowExecutionResult, }; #[derive(Clone, Debug)] @@ -44,7 +44,8 @@ impl RuntimeExecutor { operation: &RuntimeOperation, input: &Value, ) -> Result { - self.execute_with_auth(operation, input, None).await + self.execute_with_auth_and_context(operation, input, None, None) + .await } pub async fn execute_with_auth( @@ -52,10 +53,32 @@ impl RuntimeExecutor { operation: &RuntimeOperation, input: &Value, resolved_auth: Option<&ResolvedAuth>, + ) -> Result { + self.execute_with_auth_and_context(operation, input, resolved_auth, None) + .await + } + + pub async fn execute_with_context( + &self, + operation: &RuntimeOperation, + input: &Value, + request_context: Option<&RuntimeRequestContext>, + ) -> Result { + self.execute_with_auth_and_context(operation, input, None, request_context) + .await + } + + pub async fn execute_with_auth_and_context( + &self, + operation: &RuntimeOperation, + input: &Value, + resolved_auth: Option<&ResolvedAuth>, + request_context: Option<&RuntimeRequestContext>, ) -> Result { let prepared_request = self.prepare_request(operation, input)?; let prepared_request = apply_resolved_auth(prepared_request, resolved_auth); - self.execute_prepared(operation, prepared_request).await + self.execute_prepared(operation, prepared_request, request_context) + .await } pub async fn execute_window( @@ -63,7 +86,8 @@ impl RuntimeExecutor { operation: &RuntimeOperation, input: &Value, ) -> Result { - self.execute_window_with_auth(operation, input, None).await + self.execute_window_with_auth_and_context(operation, input, None, None) + .await } pub async fn execute_window_with_auth( @@ -71,6 +95,27 @@ impl RuntimeExecutor { operation: &RuntimeOperation, input: &Value, resolved_auth: Option<&ResolvedAuth>, + ) -> Result { + self.execute_window_with_auth_and_context(operation, input, resolved_auth, None) + .await + } + + pub async fn execute_window_with_context( + &self, + operation: &RuntimeOperation, + input: &Value, + request_context: Option<&RuntimeRequestContext>, + ) -> Result { + self.execute_window_with_auth_and_context(operation, input, None, request_context) + .await + } + + pub async fn execute_window_with_auth_and_context( + &self, + operation: &RuntimeOperation, + input: &Value, + resolved_auth: Option<&ResolvedAuth>, + request_context: Option<&RuntimeRequestContext>, ) -> Result { let Some(streaming) = operation.execution_config.streaming.as_ref() else { return Err(RuntimeError::MissingStreamingConfig { @@ -94,10 +139,11 @@ impl RuntimeExecutor { operation.target, Target::Rest(_) | Target::Grpc(_) | Target::Websocket(_) ) { - self.execute_window_adapter(operation, prepared_request) + self.execute_window_adapter(operation, prepared_request, request_context) .await? } else { - self.execute_adapter(operation, prepared_request).await? + self.execute_adapter(operation, prepared_request, request_context) + .await? }; crate::aggregation::collect_window_result(&adapter_response.body, streaming) @@ -108,7 +154,7 @@ impl RuntimeExecutor { operation: &RuntimeOperation, input: &Value, ) -> Result { - self.execute_session_seed_with_auth(operation, input, None) + self.execute_session_seed_with_auth_and_context(operation, input, None, None) .await } @@ -117,6 +163,27 @@ impl RuntimeExecutor { operation: &RuntimeOperation, input: &Value, resolved_auth: Option<&ResolvedAuth>, + ) -> Result { + self.execute_session_seed_with_auth_and_context(operation, input, resolved_auth, None) + .await + } + + pub async fn execute_session_seed_with_context( + &self, + operation: &RuntimeOperation, + input: &Value, + request_context: Option<&RuntimeRequestContext>, + ) -> Result { + self.execute_session_seed_with_auth_and_context(operation, input, None, request_context) + .await + } + + pub async fn execute_session_seed_with_auth_and_context( + &self, + operation: &RuntimeOperation, + input: &Value, + resolved_auth: Option<&ResolvedAuth>, + request_context: Option<&RuntimeRequestContext>, ) -> Result { let Some(streaming) = operation.execution_config.streaming.as_ref() else { return Err(RuntimeError::MissingStreamingConfig { @@ -144,8 +211,13 @@ impl RuntimeExecutor { config.max_items = Some(seed_limit); } - self.execute_window_with_auth(&seeded_operation, input, resolved_auth) - .await + self.execute_window_with_auth_and_context( + &seeded_operation, + input, + resolved_auth, + request_context, + ) + .await } pub fn prepare_request( @@ -163,8 +235,11 @@ impl RuntimeExecutor { &self, operation: &RuntimeOperation, prepared_request: PreparedRequest, + request_context: Option<&RuntimeRequestContext>, ) -> Result { - let adapter_response = self.execute_adapter(operation, prepared_request).await?; + let adapter_response = self + .execute_adapter(operation, prepared_request, request_context) + .await?; let finalized_output = finalize_output(operation, &adapter_response)?; operation.output_schema.validate_shape(&finalized_output)?; @@ -176,7 +251,9 @@ impl RuntimeExecutor { &self, operation: &RuntimeOperation, prepared_request: PreparedRequest, + request_context: Option<&RuntimeRequestContext>, ) -> Result { + let context_headers = runtime_context_headers(request_context); match &operation.target { Target::Grpc(target) => { let request = GrpcRequest { @@ -184,6 +261,7 @@ impl RuntimeExecutor { &BTreeMap::new(), &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), body: prepared_request .grpc @@ -206,6 +284,7 @@ impl RuntimeExecutor { &BTreeMap::new(), &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), variables: prepared_request.variables.clone(), timeout_ms: operation.execution_config.timeout_ms, @@ -227,6 +306,7 @@ impl RuntimeExecutor { &target.static_headers, &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), body: prepared_request.body.clone(), timeout_ms: operation.execution_config.timeout_ms, @@ -246,6 +326,7 @@ impl RuntimeExecutor { &BTreeMap::new(), &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), body: prepared_request .body @@ -273,7 +354,9 @@ impl RuntimeExecutor { &self, operation: &RuntimeOperation, prepared_request: PreparedRequest, + request_context: Option<&RuntimeRequestContext>, ) -> Result { + let context_headers = runtime_context_headers(request_context); match &operation.target { Target::Rest(target) => { let Some(streaming) = operation.execution_config.streaming.as_ref() else { @@ -289,6 +372,7 @@ impl RuntimeExecutor { &target.static_headers, &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), body: prepared_request.body.clone(), timeout_ms: streaming @@ -319,6 +403,7 @@ impl RuntimeExecutor { &BTreeMap::new(), &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), body: prepared_request .grpc @@ -356,6 +441,7 @@ impl RuntimeExecutor { &BTreeMap::new(), &operation.execution_config.headers, &prepared_request.headers, + &context_headers, ), window_duration_ms: streaming.window_duration_ms.unwrap_or_default(), max_items: streaming.max_items.map(|value| value.saturating_add(1)), @@ -384,7 +470,10 @@ impl RuntimeExecutor { operation_id: operation.operation_id.as_str().to_owned(), mode: ExecutionMode::Window, }), - _ => self.execute_adapter(operation, prepared_request).await, + _ => { + self.execute_adapter(operation, prepared_request, request_context) + .await + } } } } @@ -444,13 +533,23 @@ fn merge_headers( static_headers: &BTreeMap, execution_headers: &BTreeMap, request_headers: &BTreeMap, + context_headers: &BTreeMap, ) -> BTreeMap { let mut headers = static_headers.clone(); headers.extend(execution_headers.clone()); headers.extend(request_headers.clone()); + headers.extend(context_headers.clone()); headers } +fn runtime_context_headers( + request_context: Option<&RuntimeRequestContext>, +) -> BTreeMap { + request_context + .map(RuntimeRequestContext::outbound_headers) + .unwrap_or_default() +} + fn read_string_map( value: Option<&Value>, field_name: &str, @@ -510,6 +609,7 @@ mod tests { use axum::{ Json, Router, + http::HeaderMap, response::sse::{Event, KeepAlive, Sse}, routing::post, }; @@ -528,9 +628,14 @@ mod tests { use serde_json::{Value, json}; use time::{OffsetDateTime, format_description::well_known::Rfc3339}; use tokio::net::TcpListener; - use tokio_tungstenite::accept_async; + use tokio_tungstenite::{ + accept_async, accept_hdr_async, + tungstenite::handshake::server::{Request, Response}, + }; - use crate::{PreparedRequest, RuntimeError, RuntimeExecutor, RuntimeOperation}; + use crate::{ + PreparedRequest, RuntimeError, RuntimeExecutor, RuntimeOperation, RuntimeRequestContext, + }; fn timestamp(value: &str) -> OffsetDateTime { OffsetDateTime::parse(value, &Rfc3339).unwrap() @@ -550,6 +655,29 @@ mod tests { assert_eq!(output, json!({ "id": "lead_123" })); } + #[tokio::test] + async fn propagates_request_context_to_rest_headers() { + let base_url = spawn_runtime_server().await; + let executor = RuntimeExecutor::new(); + let mut operation = test_rest_operation(&base_url, false, false); + let context = RuntimeRequestContext::new("req_rest_123", "corr_rest_123"); + + if let Target::Rest(target) = &mut operation.target { + target.path_template = "/leads-context".to_owned(); + } + + let output = executor + .execute_with_context( + &operation, + &json!({ "email": "user@example.com" }), + Some(&context), + ) + .await + .unwrap(); + + assert_eq!(output, json!({ "id": "req_rest_123|corr_rest_123" })); + } + #[tokio::test] async fn executes_graphql_operation_end_to_end() { let endpoint = spawn_graphql_server().await; @@ -564,6 +692,25 @@ mod tests { assert_eq!(output, json!({ "id": "lead_123" })); } + #[tokio::test] + async fn propagates_request_context_to_graphql_headers() { + let endpoint = spawn_graphql_context_server().await; + let executor = RuntimeExecutor::new(); + let operation = test_graphql_operation(&endpoint, false); + let context = RuntimeRequestContext::new("req_graphql_123", "corr_graphql_123"); + + let output = executor + .execute_with_context( + &operation, + &json!({ "email": "user@example.com" }), + Some(&context), + ) + .await + .unwrap(); + + assert_eq!(output, json!({ "id": "req_graphql_123|corr_graphql_123" })); + } + #[tokio::test] async fn executes_grpc_operation_end_to_end() { let server_addr = grpc_test_support::spawn_unary_echo_server().await; @@ -578,6 +725,24 @@ mod tests { assert_eq!(output, json!({ "message": "hello" })); } + #[tokio::test] + async fn propagates_request_context_to_grpc_metadata() { + let server_addr = grpc_test_support::spawn_metadata_echo_server().await; + let executor = RuntimeExecutor::new(); + let operation = test_grpc_operation(&server_addr); + let context = RuntimeRequestContext::new("req_grpc_123", "corr_grpc_123"); + + let output = executor + .execute_with_context(&operation, &json!({ "message": "hello" }), Some(&context)) + .await + .unwrap(); + + assert_eq!( + output, + json!({ "message": "hello|req_grpc_123|corr_grpc_123" }) + ); + } + #[tokio::test] async fn executes_soap_operation_end_to_end() { let endpoint = spawn_soap_server().await; @@ -592,6 +757,25 @@ mod tests { assert_eq!(output, json!({ "id": "lead_123" })); } + #[tokio::test] + async fn propagates_request_context_to_soap_headers() { + let endpoint = spawn_soap_context_server().await; + let executor = RuntimeExecutor::new(); + let operation = test_soap_operation(&endpoint); + let context = RuntimeRequestContext::new("req_soap_123", "corr_soap_123"); + + let output = executor + .execute_with_context( + &operation, + &json!({ "email": "user@example.com" }), + Some(&context), + ) + .await + .unwrap(); + + assert_eq!(output, json!({ "id": "req_soap_123|corr_soap_123" })); + } + #[tokio::test] async fn executes_grpc_window_mode_with_server_stream() { let server_addr = grpc_test_support::spawn_unary_echo_server().await; @@ -632,6 +816,23 @@ mod tests { assert!(result.has_more); } + #[tokio::test] + async fn propagates_request_context_to_websocket_headers() { + let target_url = spawn_request_context_websocket_server().await; + let executor = RuntimeExecutor::new(); + let operation = + test_websocket_window_operation(&target_url, AggregationMode::RawItems, Some(2)); + let context = RuntimeRequestContext::new("req_ws_123", "corr_ws_123"); + + let result = executor + .execute_window_with_context(&operation, &json!({}), Some(&context)) + .await + .unwrap(); + + assert_eq!(result.items.len(), 2); + assert_eq!(result.items[0], json!({ "seq": 1, "value": 101 })); + } + #[tokio::test] async fn rejects_invalid_input_shape() { let base_url = spawn_runtime_server().await; @@ -807,6 +1008,7 @@ mod tests { async fn spawn_runtime_server() -> String { let app = Router::new() .route("/leads", post(create_lead)) + .route("/leads-context", post(create_lead_with_context)) .route("/events", post(events_window)) .route("/slow-events", post(slow_events_window)) .route("/events-sse", post(events_stream)) @@ -833,6 +1035,18 @@ mod tests { format!("http://{}", address) } + async fn spawn_graphql_context_server() -> String { + let app = Router::new().route("/", post(graphql_context_handler)); + 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 spawn_soap_server() -> String { let app = Router::new().route("/", post(soap_handler)); let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); @@ -845,6 +1059,18 @@ mod tests { format!("http://{}", address) } + async fn spawn_soap_context_server() -> String { + let app = Router::new().route("/", post(soap_context_handler)); + 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 spawn_websocket_server() -> String { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); @@ -883,6 +1109,61 @@ mod tests { format!("ws://{}", address) } + async fn spawn_request_context_websocket_server() -> String { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + + tokio::spawn(async move { + loop { + let (stream, _) = listener.accept().await.unwrap(); + tokio::spawn(async move { + let websocket = + accept_hdr_async(stream, |request: &Request, response: Response| { + assert_eq!( + request + .headers() + .get("x-request-id") + .and_then(|value| value.to_str().ok()), + Some("req_ws_123") + ); + assert_eq!( + request + .headers() + .get("x-correlation-id") + .and_then(|value| value.to_str().ok()), + Some("corr_ws_123") + ); + Ok(response) + }) + .await + .unwrap(); + let (mut sink, mut source) = websocket.split(); + + if let Some(message) = source.next().await { + let message = message.unwrap(); + if !matches!(message, tokio_tungstenite::tungstenite::Message::Text(_)) { + return; + } + } + + for payload in [ + json!({ "seq": 1, "value": 101 }), + json!({ "seq": 2, "value": 102 }), + ] { + sink.send(tokio_tungstenite::tungstenite::Message::Text( + payload.to_string().into(), + )) + .await + .unwrap(); + } + let _ = sink.close().await; + }); + } + }); + + format!("ws://{}", address) + } + async fn create_lead(Json(payload): Json) -> (axum::http::StatusCode, Json) { let should_fail = payload .get("fail") @@ -902,6 +1183,28 @@ mod tests { ) } + async fn create_lead_with_context( + headers: HeaderMap, + Json(_payload): Json, + ) -> (axum::http::StatusCode, Json) { + ( + axum::http::StatusCode::OK, + Json(json!({ + "id": format!( + "{}|{}", + headers + .get("x-request-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(), + headers + .get("x-correlation-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + ), + })), + ) + } + async fn events_window() -> (axum::http::StatusCode, Json) { ( axum::http::StatusCode::OK, @@ -970,6 +1273,30 @@ mod tests { })) } + async fn graphql_context_handler( + headers: HeaderMap, + Json(_payload): Json, + ) -> Json { + Json(json!({ + "data": { + "createLead": { + "id": format!( + "{}|{}", + headers + .get("x-request-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(), + headers + .get("x-correlation-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + ), + "status": "created" + } + } + })) + } + async fn soap_handler(body: String) -> (axum::http::StatusCode, String) { assert!(body.contains("user@example.com")); @@ -987,6 +1314,35 @@ mod tests { ) } + async fn soap_context_handler( + headers: HeaderMap, + body: String, + ) -> (axum::http::StatusCode, String) { + assert!(body.contains("user@example.com")); + let request_id = headers + .get("x-request-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + let correlation_id = headers + .get("x-correlation-id") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + + ( + axum::http::StatusCode::OK, + format!( + r#" + + + {request_id}|{correlation_id} + created + + + "# + ), + ) + } + fn test_rest_operation( base_url: &str, should_fail: bool, diff --git a/crates/crank-runtime/src/lib.rs b/crates/crank-runtime/src/lib.rs index 46f6e94..9e9c35c 100644 --- a/crates/crank-runtime/src/lib.rs +++ b/crates/crank-runtime/src/lib.rs @@ -4,6 +4,7 @@ mod error; mod executor; mod model; mod redaction; +mod request_context; mod secret_crypto; mod streaming; @@ -11,5 +12,6 @@ pub use auth::ResolvedAuth; pub use error::RuntimeError; pub use executor::RuntimeExecutor; pub use model::{AdapterResponse, PreparedRequest, RuntimeOperation}; +pub use request_context::RuntimeRequestContext; pub use secret_crypto::SecretCrypto; pub use streaming::WindowExecutionResult; diff --git a/crates/crank-runtime/src/request_context.rs b/crates/crank-runtime/src/request_context.rs new file mode 100644 index 0000000..d697b67 --- /dev/null +++ b/crates/crank-runtime/src/request_context.rs @@ -0,0 +1,41 @@ +use std::collections::BTreeMap; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct RuntimeRequestContext { + pub request_id: String, + pub correlation_id: String, +} + +impl RuntimeRequestContext { + pub fn new(request_id: impl Into, correlation_id: impl Into) -> Self { + Self { + request_id: request_id.into(), + correlation_id: correlation_id.into(), + } + } + + pub fn from_request_id(request_id: impl Into) -> Self { + let request_id = request_id.into(); + Self::new(request_id.clone(), request_id) + } + + pub fn outbound_headers(&self) -> BTreeMap { + BTreeMap::from([ + ("x-request-id".to_owned(), self.request_id.clone()), + ("x-correlation-id".to_owned(), self.correlation_id.clone()), + ]) + } +} + +#[cfg(test)] +mod tests { + use super::RuntimeRequestContext; + + #[test] + fn uses_request_id_for_default_correlation_id() { + let context = RuntimeRequestContext::from_request_id("req_123"); + + assert_eq!(context.request_id, "req_123"); + assert_eq!(context.correlation_id, "req_123"); + } +}