use std::{collections::BTreeMap, time::Duration}; use crank_core::{HttpMethod, RestTarget}; use reqwest::{ Client, header::{HeaderMap, HeaderName, HeaderValue}, }; use serde_json::Value; use crate::{RestAdapterError, RestRequest, RestResponse}; #[derive(Clone, Debug)] pub struct RestAdapter { client: Client, } impl Default for RestAdapter { fn default() -> Self { Self::new() } } impl RestAdapter { pub fn new() -> Self { Self { client: Client::new(), } } pub async fn execute( &self, target: &RestTarget, request: &RestRequest, ) -> Result { let url = build_url(target, request)?; let headers = build_headers(target, request)?; let mut builder = self .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()); let body = decode_body(response).await?; if !status.is_success() { return Err(RestAdapterError::UnexpectedStatus { status: status.as_u16(), body, }); } Ok(RestResponse { status_code: status.as_u16(), headers, body, }) } } fn build_url(target: &RestTarget, request: &RestRequest) -> Result { let base_url = reqwest::Url::parse(&target.base_url).map_err(|_| RestAdapterError::InvalidBaseUrl { url: target.base_url.clone(), })?; let path = substitute_path_params(&target.path_template, &request.path_params); let mut url = base_url.join(path.trim_start_matches('/')).map_err(|_| { RestAdapterError::InvalidBaseUrl { url: target.base_url.clone(), } })?; { let mut query = url.query_pairs_mut(); for (key, value) in &request.query_params { query.append_pair(key, value); } } Ok(url) } fn substitute_path_params(path_template: &str, path_params: &BTreeMap) -> String { let mut rendered = path_template.to_owned(); for (key, value) in path_params { rendered = rendered.replace(&format!("{{{key}}}"), value); } rendered } fn build_headers( target: &RestTarget, request: &RestRequest, ) -> Result { let mut headers = HeaderMap::new(); for (name, value) in &target.static_headers { insert_header(&mut headers, name, value)?; } for (name, value) in &request.headers { insert_header(&mut headers, name, value)?; } Ok(headers) } fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<(), RestAdapterError> { let header_name = HeaderName::try_from(name).map_err(|_| RestAdapterError::InvalidHeaderName { header: name.to_owned(), })?; let header_value = HeaderValue::try_from(value).map_err(|_| RestAdapterError::InvalidHeaderValue { header: name.to_owned(), })?; headers.insert(header_name, header_value); Ok(()) } async fn decode_body(response: reqwest::Response) -> Result { let bytes = response.bytes().await?; if bytes.is_empty() { return Ok(Value::Null); } match serde_json::from_slice::(&bytes) { Ok(value) => Ok(value), Err(_) => Ok(Value::String( String::from_utf8_lossy(&bytes).trim().to_owned(), )), } } fn normalize_headers(headers: &HeaderMap) -> BTreeMap { headers .iter() .filter_map(|(name, value)| { value .to_str() .ok() .map(|value| (name.as_str().to_owned(), value.to_owned())) }) .collect() } fn to_reqwest_method(method: HttpMethod) -> reqwest::Method { match method { HttpMethod::Get => reqwest::Method::GET, HttpMethod::Post => reqwest::Method::POST, HttpMethod::Put => reqwest::Method::PUT, HttpMethod::Patch => reqwest::Method::PATCH, HttpMethod::Delete => reqwest::Method::DELETE, } } #[cfg(test)] mod tests { use std::collections::BTreeMap; use axum::{ Json, Router, extract::{Path, Query}, http::HeaderMap, routing::{get, post}, }; use crank_core::{HttpMethod, RestTarget}; use serde_json::{Value, json}; use tokio::net::TcpListener; use crate::{RestAdapter, RestAdapterError, RestRequest}; #[tokio::test] async fn executes_rest_request_and_normalizes_json_response() { let base_url = spawn_test_server().await; let adapter = RestAdapter::new(); 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 returns_unexpected_status_with_normalized_body() { let base_url = spawn_test_server().await; let adapter = RestAdapter::new(); 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(_) } )); } async fn spawn_test_server() -> String { let app = Router::new() .route("/users/{user_id}", post(create_user)) .route("/fail", get(fail)); 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(); Json(json!({ "id": user_id, "query": query.get("expand").cloned().unwrap_or_default(), "trace": trace, "static": static_header, "payload": payload })) } async fn fail() -> (axum::http::StatusCode, Json) { ( axum::http::StatusCode::BAD_GATEWAY, Json(json!({ "error": "upstream failed" })), ) } }