use std::time::Duration;
use dataflow_rs::engine::error::DataflowError;
use serde_json::Value;
use crate::connector::{AuthConfig, HttpConnectorConfig};
pub fn build_url(base: &str, path: Option<&str>) -> String {
match path {
Some(p) if !p.is_empty() => {
let base = base.trim_end_matches('/');
let path = p.trim_start_matches('/');
format!("{base}/{path}")
}
_ => base.to_string(),
}
}
pub fn apply_auth(req: reqwest::RequestBuilder, auth: &AuthConfig) -> reqwest::RequestBuilder {
match auth {
AuthConfig::Bearer { token } => req.header("authorization", format!("Bearer {token}")),
AuthConfig::Basic { username, password } => req.basic_auth(username, Some(password)),
AuthConfig::ApiKey { header, key } => req.header(header, key),
}
}
const MAX_REDIRECTS: usize = 5;
#[tracing::instrument(skip(client, task_headers, http_config, body))]
pub async fn execute_request(
client: &reqwest::Client,
method: &reqwest::Method,
url: &str,
task_headers: Option<&std::collections::HashMap<String, String>>,
http_config: &HttpConnectorConfig,
body: Option<&Value>,
timeout: Duration,
) -> dataflow_rs::Result<Value> {
let original = url::Url::parse(url)
.map_err(|e| DataflowError::Validation(format!("Invalid URL '{url}': {e}")))?;
let mut current = original.clone();
let mut method = method.clone();
let mut body = body;
for _ in 0..=MAX_REDIRECTS {
let own_endpoint = same_endpoint(¤t, &original);
if !(http_config.allow_private_urls && own_endpoint)
&& let Err(msg) = crate::validation::validate_url_not_private(current.as_str()).await
{
return Err(DataflowError::function_execution(
format!("SSRF protection: {msg}"),
None,
));
}
let mut req = client
.request(method.clone(), current.clone())
.timeout(timeout);
{
let mut trace_headers = std::collections::HashMap::new();
crate::server::trace_context::inject_trace_context(&mut trace_headers);
for (k, v) in &trace_headers {
req = req.header(k, v);
}
}
if own_endpoint {
for (k, v) in &http_config.headers {
req = req.header(k, v);
}
if let Some(ref auth) = http_config.auth {
req = apply_auth(req, auth);
}
}
if let Some(b) = body {
req = req.header("content-type", "application/json").json(b);
}
if own_endpoint && let Some(headers) = task_headers {
for (k, v) in headers {
req = req.header(k, v);
}
}
let response = req.send().await.map_err(|e| {
if e.is_timeout() {
DataflowError::Timeout(format!("HTTP request to {current} timed out"))
} else {
DataflowError::Io(format!("HTTP request to {current} failed: {e}"))
}
})?;
if let Some(next) = redirect_target(&response, ¤t)? {
if matches!(response.status().as_u16(), 301..=303)
&& method != reqwest::Method::GET
&& method != reqwest::Method::HEAD
{
method = reqwest::Method::GET;
body = None;
}
current = next;
continue;
}
return read_json_response(response, ¤t, http_config.max_response_size).await;
}
Err(DataflowError::function_execution(
format!("Stopped after {MAX_REDIRECTS} redirects requesting {url}"),
None,
))
}
fn same_endpoint(a: &url::Url, b: &url::Url) -> bool {
a.host_str().is_some()
&& a.host_str() == b.host_str()
&& a.port_or_known_default() == b.port_or_known_default()
}
fn redirect_target(
response: &reqwest::Response,
current: &url::Url,
) -> dataflow_rs::Result<Option<url::Url>> {
if !matches!(response.status().as_u16(), 301 | 302 | 303 | 307 | 308) {
return Ok(None);
}
let Some(location) = response.headers().get(reqwest::header::LOCATION) else {
return Ok(None);
};
let location = location.to_str().map_err(|_| {
DataflowError::function_execution(
format!("Redirect from {current} has a non-ASCII Location header"),
None,
)
})?;
let next = current.join(location).map_err(|e| {
DataflowError::function_execution(
format!("Redirect from {current} has invalid Location '{location}': {e}"),
None,
)
})?;
if !matches!(next.scheme(), "http" | "https") {
return Err(DataflowError::function_execution(
format!(
"Redirect from {current} targets unsupported scheme '{}'",
next.scheme()
),
None,
));
}
Ok(Some(next))
}
async fn read_json_response(
mut response: reqwest::Response,
url: &url::Url,
max_size: usize,
) -> dataflow_rs::Result<Value> {
let status = response.status();
if let Some(content_length) = response.content_length()
&& content_length as usize > max_size
{
return Err(DataflowError::function_execution(
format!(
"Response from {url} declared Content-Length {content_length} exceeds limit of {max_size} bytes"
),
None,
));
}
if !status.is_success() {
let mut body_bytes = Vec::new();
while let Some(chunk) = response.chunk().await.ok().flatten() {
let room = max_size.saturating_sub(body_bytes.len());
let take = chunk.len().min(room);
body_bytes.extend_from_slice(&chunk[..take]);
if take < chunk.len() {
break;
}
}
let body_text = String::from_utf8_lossy(&body_bytes);
return Err(DataflowError::http(
status.as_u16(),
format!("HTTP {status} from {url}: {body_text}"),
));
}
let mut body_bytes = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|e| {
DataflowError::function_execution(
format!("Failed to read response body from {url}: {e}"),
None,
)
})? {
if body_bytes.len() + chunk.len() > max_size {
return Err(DataflowError::function_execution(
format!("Response body from {url} exceeds limit of {max_size} bytes"),
None,
));
}
body_bytes.extend_from_slice(&chunk);
}
let response_body: Value = serde_json::from_slice(&body_bytes).map_err(|e| {
DataflowError::function_execution(
format!("Failed to parse response from {url} as JSON: {e}"),
None,
)
})?;
Ok(response_body)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_url() {
assert_eq!(
build_url("https://api.example.com", Some("/users")),
"https://api.example.com/users"
);
assert_eq!(
build_url("https://api.example.com/", Some("/users")),
"https://api.example.com/users"
);
assert_eq!(
build_url("https://api.example.com", None),
"https://api.example.com"
);
}
#[test]
fn test_build_url_no_path() {
assert_eq!(
build_url("https://api.example.com", None),
"https://api.example.com"
);
}
#[test]
fn test_build_url_empty_path() {
assert_eq!(
build_url("https://api.example.com", Some("")),
"https://api.example.com"
);
}
#[test]
fn test_build_url_trims_slashes() {
assert_eq!(
build_url("https://api.example.com///", Some("///path")),
"https://api.example.com/path"
);
}
#[test]
fn test_apply_auth_bearer() {
let client = reqwest::Client::new();
let auth = AuthConfig::Bearer {
token: "tok123".to_string(),
};
let req = apply_auth(client.get("http://localhost"), &auth);
let built = req.build().expect("test");
assert_eq!(
built
.headers()
.get("authorization")
.expect("test")
.to_str()
.expect("test"),
"Bearer tok123"
);
}
#[test]
fn test_apply_auth_api_key() {
let client = reqwest::Client::new();
let auth = AuthConfig::ApiKey {
header: "x-api-key".to_string(),
key: "secret123".to_string(),
};
let req = apply_auth(client.get("http://localhost"), &auth);
let built = req.build().expect("test");
assert_eq!(
built
.headers()
.get("x-api-key")
.expect("test")
.to_str()
.expect("test"),
"secret123"
);
}
#[tokio::test]
async fn test_execute_request_success() {
let mock_app = axum::Router::new().route(
"/test",
axum::routing::get(|| async { axum::Json(serde_json::json!({"result": "success"})) }),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, mock_app).await.expect("test");
});
let client = reqwest::Client::new();
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
};
let result = execute_request(
&client,
&reqwest::Method::GET,
&format!("http://{}/test", addr),
None,
&http_config,
None,
std::time::Duration::from_secs(5),
)
.await;
assert!(result.is_ok());
let val = result.expect("test");
assert_eq!(val["result"], "success");
}
#[tokio::test]
async fn test_execute_request_with_headers_auth_and_body() {
let mock_app = axum::Router::new().route(
"/post-test",
axum::routing::post(|| async { axum::Json(serde_json::json!({"received": true})) }),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, mock_app).await.expect("test");
});
let client = reqwest::Client::new();
let mut headers = std::collections::HashMap::new();
headers.insert("x-custom".to_string(), "custom-value".to_string());
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::from([(
"x-connector-header".to_string(),
"conn-val".to_string(),
)]),
auth: Some(AuthConfig::Bearer {
token: "test-token".to_string(),
}),
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
};
let body = serde_json::json!({"data": "payload"});
let result = execute_request(
&client,
&reqwest::Method::POST,
&format!("http://{}/post-test", addr),
Some(&headers),
&http_config,
Some(&body),
std::time::Duration::from_secs(5),
)
.await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_execute_request_non_success_status() {
let mock_app = axum::Router::new().route(
"/error",
axum::routing::get(|| async { (axum::http::StatusCode::BAD_REQUEST, "Bad Request") }),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, mock_app).await.expect("test");
});
let client = reqwest::Client::new();
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
};
let result = execute_request(
&client,
&reqwest::Method::GET,
&format!("http://{}/error", addr),
None,
&http_config,
None,
std::time::Duration::from_secs(5),
)
.await;
assert!(result.is_err());
let err = result.expect_err("test");
assert!(err.to_string().contains("400"));
}
#[tokio::test]
async fn test_execute_request_non_json_response() {
let mock_app = axum::Router::new().route(
"/text",
axum::routing::get(|| async { "plain text response" }),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, mock_app).await.expect("test");
});
let client = reqwest::Client::new();
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
};
let result = execute_request(
&client,
&reqwest::Method::GET,
&format!("http://{}/text", addr),
None,
&http_config,
None,
std::time::Duration::from_secs(5),
)
.await;
assert!(result.is_err());
assert!(result.expect_err("test").to_string().contains("parse"));
}
#[tokio::test]
async fn test_execute_request_response_too_large() {
let mock_app = axum::Router::new().route(
"/large",
axum::routing::get(|| async {
axum::Json(serde_json::json!({"data": "x".repeat(200)}))
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, mock_app).await.expect("test");
});
let client = reqwest::Client::new();
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10, allow_private_urls: true, operations: Default::default(),
};
let result = execute_request(
&client,
&reqwest::Method::GET,
&format!("http://{}/large", addr),
None,
&http_config,
None,
std::time::Duration::from_secs(5),
)
.await;
assert!(result.is_err());
assert!(result.expect_err("test").to_string().contains("exceed"));
}
#[tokio::test]
async fn test_execute_request_timeout() {
let mock_app = axum::Router::new().route(
"/slow",
axum::routing::get(|| async {
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
axum::Json(serde_json::json!({"slow": true}))
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, mock_app).await.expect("test");
});
let client = reqwest::Client::new();
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
};
let result = execute_request(
&client,
&reqwest::Method::GET,
&format!("http://{}/slow", addr),
None,
&http_config,
None,
std::time::Duration::from_millis(100), )
.await;
assert!(result.is_err());
assert!(result.expect_err("test").to_string().contains("timed out"));
}
#[tokio::test]
async fn test_execute_request_connection_refused() {
let client = reqwest::Client::new();
let http_config = HttpConnectorConfig {
retry_non_idempotent: false,
url: "http://127.0.0.1:1".to_string(),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
};
let result = execute_request(
&client,
&reqwest::Method::GET,
"http://127.0.0.1:1/test",
None,
&http_config,
None,
std::time::Duration::from_secs(1),
)
.await;
assert!(result.is_err());
assert!(result.expect_err("test").to_string().contains("failed"));
}
fn redirectless_client() -> reqwest::Client {
reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test")
}
fn localhost_config(addr: std::net::SocketAddr) -> HttpConnectorConfig {
HttpConnectorConfig {
retry_non_idempotent: false,
url: format!("http://{}", addr),
method: String::new(),
headers: std::collections::HashMap::new(),
auth: None,
retry: crate::connector::RetryConfig::default(),
max_response_size: 10 * 1024 * 1024,
allow_private_urls: true, operations: Default::default(),
}
}
async fn spawn_mock(app: axum::Router) -> std::net::SocketAddr {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("test");
let addr = listener.local_addr().expect("test");
tokio::spawn(async move {
axum::serve(listener, app).await.expect("test");
});
addr
}
#[tokio::test]
async fn test_redirect_to_private_target_refused() {
let mock_app = axum::Router::new().route(
"/redirect",
axum::routing::get(|| async {
(
axum::http::StatusCode::FOUND,
[(
axum::http::header::LOCATION,
"http://169.254.169.254/latest/meta-data",
)],
)
}),
);
let addr = spawn_mock(mock_app).await;
let result = execute_request(
&redirectless_client(),
&reqwest::Method::GET,
&format!("http://{}/redirect", addr),
None,
&localhost_config(addr),
None,
std::time::Duration::from_secs(5),
)
.await;
let err = result.expect_err("test").to_string();
assert!(err.contains("SSRF protection"), "unexpected error: {err}");
}
#[tokio::test]
async fn test_redirect_followed_within_own_endpoint() {
let mock_app = axum::Router::new()
.route(
"/a",
axum::routing::get(|| async {
(
axum::http::StatusCode::FOUND,
[(axum::http::header::LOCATION, "/b")],
)
}),
)
.route(
"/b",
axum::routing::get(|| async { axum::Json(serde_json::json!({"hop": "b"})) }),
);
let addr = spawn_mock(mock_app).await;
let result = execute_request(
&redirectless_client(),
&reqwest::Method::GET,
&format!("http://{}/a", addr),
None,
&localhost_config(addr),
None,
std::time::Duration::from_secs(5),
)
.await;
assert_eq!(result.expect("test")["hop"], "b");
}
#[tokio::test]
async fn test_redirect_loop_is_capped() {
let mock_app = axum::Router::new().route(
"/loop",
axum::routing::get(|| async {
(
axum::http::StatusCode::FOUND,
[(axum::http::header::LOCATION, "/loop")],
)
}),
);
let addr = spawn_mock(mock_app).await;
let result = execute_request(
&redirectless_client(),
&reqwest::Method::GET,
&format!("http://{}/loop", addr),
None,
&localhost_config(addr),
None,
std::time::Duration::from_secs(5),
)
.await;
let err = result.expect_err("test").to_string();
assert!(err.contains("redirects"), "unexpected error: {err}");
}
#[tokio::test]
async fn test_redirect_303_downgrades_post_to_get() {
let mock_app = axum::Router::new()
.route(
"/submit",
axum::routing::post(|| async {
(
axum::http::StatusCode::SEE_OTHER,
[(axum::http::header::LOCATION, "/done")],
)
}),
)
.route(
"/done",
axum::routing::get(|| async { axum::Json(serde_json::json!({"done": true})) }),
);
let addr = spawn_mock(mock_app).await;
let body = serde_json::json!({"data": "payload"});
let result = execute_request(
&redirectless_client(),
&reqwest::Method::POST,
&format!("http://{}/submit", addr),
None,
&localhost_config(addr),
Some(&body),
std::time::Duration::from_secs(5),
)
.await;
assert_eq!(result.expect("test")["done"], true);
}
#[test]
fn test_apply_auth_basic() {
let client = reqwest::Client::new();
let auth = AuthConfig::Basic {
username: "user".to_string(),
password: "pass".to_string(),
};
let req = apply_auth(client.get("http://localhost"), &auth);
let built = req.build().expect("test");
let auth_header = built
.headers()
.get("authorization")
.expect("test")
.to_str()
.expect("test");
assert!(auth_header.starts_with("Basic "));
}
}