use super::auth::HttpAuthProvider;
use super::{join_url, HttpConnector, HttpConnectorError, Operation, Parameter, ParameterLocation};
use async_trait::async_trait;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct HttpConfig {
#[serde(default = "default_timeout")]
pub timeout_seconds: u64,
#[serde(default = "default_retries")]
pub retries: u32,
#[serde(default = "default_retry_backoff")]
pub retry_backoff_ms: u64,
#[serde(default = "default_user_agent")]
pub user_agent: String,
#[serde(default)]
pub default_headers: HashMap<String, String>,
}
fn default_timeout() -> u64 {
30
}
fn default_retries() -> u32 {
3
}
fn default_retry_backoff() -> u64 {
1000
}
fn default_user_agent() -> String {
format!("pmcp-server-toolkit/{}", env!("CARGO_PKG_VERSION"))
}
impl Default for HttpConfig {
fn default() -> Self {
Self {
timeout_seconds: default_timeout(),
retries: default_retries(),
retry_backoff_ms: default_retry_backoff(),
user_agent: default_user_agent(),
default_headers: HashMap::new(),
}
}
}
pub struct HttpClient {
client: reqwest::Client,
base_url: url::Url,
auth: Arc<dyn HttpAuthProvider>,
http_config: HttpConfig,
policy: Option<Arc<dyn crate::policy::RequestPolicy>>,
}
impl HttpClient {
pub fn new(
client: reqwest::Client,
base_url: String,
auth: Arc<dyn HttpAuthProvider>,
) -> Result<Self, HttpConnectorError> {
Self::with_config(client, base_url, auth, HttpConfig::default())
}
pub fn with_config(
client: reqwest::Client,
base_url: String,
auth: Arc<dyn HttpAuthProvider>,
http_config: HttpConfig,
) -> Result<Self, HttpConnectorError> {
let base_url = url::Url::parse(&base_url)
.map_err(|_| HttpConnectorError::Backend("invalid base URL".to_string()))?;
Ok(Self {
client,
base_url,
auth,
http_config,
policy: None,
})
}
#[must_use]
pub fn with_request_policy(mut self, policy: Arc<dyn crate::policy::RequestPolicy>) -> Self {
self.policy = Some(policy);
self
}
pub fn from_config(
base_url: String,
auth: Arc<dyn HttpAuthProvider>,
http_config: HttpConfig,
) -> Result<Self, HttpConnectorError> {
let mut headers = HeaderMap::new();
if let Ok(ua) = HeaderValue::from_str(&http_config.user_agent) {
headers.insert(reqwest::header::USER_AGENT, ua);
}
for (key, value) in &http_config.default_headers {
if let (Ok(name), Ok(val)) = (
HeaderName::try_from(key.as_str()),
HeaderValue::try_from(value.as_str()),
) {
headers.insert(name, val);
}
}
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(http_config.timeout_seconds))
.redirect(reqwest::redirect::Policy::none())
.default_headers(headers)
.build()
.map_err(|_| HttpConnectorError::Backend("failed to build HTTP client".to_string()))?;
Self::with_config(client, base_url, auth, http_config)
}
fn substitute_path(
operation: &Operation,
args: &serde_json::Map<String, serde_json::Value>,
) -> Result<String, HttpConnectorError> {
let mut rendered: Vec<(String, String)> = Vec::new();
for param in operation.path_parameters() {
let Some(value) = args.get(¶m.name) else {
refuse_missing_path_argument(¶m.name)?;
continue;
};
let value_str = render_scalar(¶m.name, value)?;
check_placeholder_value(param, &value_str)?;
rendered.push((format!("{{{}}}", param.name), value_str));
}
let mut path = operation.path.clone();
for (placeholder, value_str) in &rendered {
path = path.replace(placeholder, value_str);
}
check_composed_path(&path)?;
Ok(path)
}
fn render_query_value(
param_name: &str,
value: &serde_json::Value,
) -> Result<String, HttpConnectorError> {
if let serde_json::Value::Array(arr) = value {
let mut csv = String::new();
for (i, member) in arr.iter().enumerate() {
if i > 0 {
csv.push(',');
}
csv.push_str(&render_scalar(param_name, member)?);
}
Ok(csv)
} else {
render_scalar(param_name, value)
}
}
fn build_query(
operation: &Operation,
args: &serde_json::Map<String, serde_json::Value>,
) -> Result<HashMap<String, String>, HttpConnectorError> {
let mut query = HashMap::new();
for param in operation.query_parameters() {
if let Some(value) = args.get(¶m.name) {
query.insert(
param.name.clone(),
Self::render_query_value(¶m.name, value)?,
);
}
}
Ok(query)
}
fn build_headers(
operation: &Operation,
args: &serde_json::Map<String, serde_json::Value>,
) -> Result<HeaderMap, HttpConnectorError> {
let mut headers = HeaderMap::new();
for param in operation.header_parameters() {
if let Some(value) = args.get(¶m.name) {
let name = HeaderName::try_from(param.name.as_str()).map_err(|_| {
HttpConnectorError::InvalidHeader("invalid header name".to_string())
})?;
let rendered = render_scalar(¶m.name, value)?;
let val = HeaderValue::try_from(rendered).map_err(|_| {
HttpConnectorError::InvalidHeader("invalid header value".to_string())
})?;
headers.insert(name, val);
}
}
Ok(headers)
}
fn build_body(
operation: &Operation,
args: &serde_json::Map<String, serde_json::Value>,
) -> Option<serde_json::Value> {
if !operation.has_request_body {
return None;
}
if let Some(body) = args.get("body") {
return Some(body.clone());
}
let routed_elsewhere: std::collections::HashSet<&str> = operation
.parameters
.iter()
.filter(|p| p.location != ParameterLocation::Body)
.map(|p| p.name.as_str())
.collect();
let body: serde_json::Map<String, serde_json::Value> = args
.iter()
.filter(|(k, _)| !routed_elsewhere.contains(k.as_str()))
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
if body.is_empty() {
None
} else {
Some(serde_json::Value::Object(body))
}
}
fn convert_method(method: &str) -> Result<reqwest::Method, HttpConnectorError> {
match method.to_uppercase().as_str() {
"GET" => Ok(reqwest::Method::GET),
"POST" => Ok(reqwest::Method::POST),
"PUT" => Ok(reqwest::Method::PUT),
"PATCH" => Ok(reqwest::Method::PATCH),
"DELETE" => Ok(reqwest::Method::DELETE),
"HEAD" => Ok(reqwest::Method::HEAD),
"OPTIONS" => Ok(reqwest::Method::OPTIONS),
_ => Err(HttpConnectorError::Backend(
"unknown HTTP method".to_string(),
)),
}
}
async fn send_with_retries(
&self,
request: reqwest::RequestBuilder,
) -> Result<reqwest::Response, HttpConnectorError> {
let max_retries = self.http_config.retries;
let mut last_status: Option<u16> = None;
for attempt in 0..=max_retries {
if attempt > 0 {
let delay = self.http_config.retry_backoff_ms * (1u64 << (attempt - 1));
tokio::time::sleep(Duration::from_millis(delay)).await;
}
let Some(attempt_request) = request.try_clone() else {
return Err(HttpConnectorError::Request(
"request body is not retryable".to_string(),
));
};
match attempt_request.send().await {
Ok(response) => {
let status = response.status();
if status.is_server_error() && attempt < max_retries {
last_status = Some(status.as_u16());
continue;
}
return Ok(response);
},
Err(e) => {
let retryable = e.is_connect() || e.is_timeout();
if retryable && attempt < max_retries {
continue;
}
return Err(HttpConnectorError::Request(
"transport error contacting backend".to_string(),
));
},
}
}
Err(HttpConnectorError::Status {
status: last_status.unwrap_or(0),
})
}
}
fn render_scalar(
param_name: &str,
value: &serde_json::Value,
) -> Result<String, HttpConnectorError> {
match value {
serde_json::Value::String(s) => Ok(s.clone()),
serde_json::Value::Number(n) => Ok(n.to_string()),
serde_json::Value::Bool(b) => Ok(b.to_string()),
serde_json::Value::Null => Ok("null".to_string()),
serde_json::Value::Object(_) | serde_json::Value::Array(_) => {
Err(HttpConnectorError::Backend(format!(
"param '{param_name}' must be a scalar (non-scalar values are \
not supported in path/query/header position)"
)))
},
}
}
#[cfg(feature = "input-validation")]
fn refusal_to_backend_error(
refusal: &pmcp::server::schema_validation::PlaceholderRefusal,
) -> HttpConnectorError {
HttpConnectorError::Backend(format!("{refusal}"))
}
#[cfg(feature = "input-validation")]
fn check_placeholder_value(param: &Parameter, value_str: &str) -> Result<(), HttpConnectorError> {
pmcp::server::schema_validation::validate_path_placeholder(
¶m.name,
value_str,
¶m.placeholder_rules(),
)
.map_err(|refusal| refusal_to_backend_error(&refusal))
}
#[cfg(not(feature = "input-validation"))]
fn check_placeholder_value(_param: &Parameter, _value_str: &str) -> Result<(), HttpConnectorError> {
Ok(())
}
#[cfg(feature = "input-validation")]
fn check_composed_path(path: &str) -> Result<(), HttpConnectorError> {
pmcp::server::schema_validation::validate_resolved_target(path)
.map_err(|refusal| refusal_to_backend_error(&refusal))
}
#[cfg(not(feature = "input-validation"))]
fn check_composed_path(_path: &str) -> Result<(), HttpConnectorError> {
Ok(())
}
#[cfg(feature = "input-validation")]
fn refuse_missing_path_argument(param_name: &str) -> Result<(), HttpConnectorError> {
Err(HttpConnectorError::Backend(format!(
"param '{param_name}' is a declared path parameter and must be supplied"
)))
}
#[cfg(not(feature = "input-validation"))]
fn refuse_missing_path_argument(_param_name: &str) -> Result<(), HttpConnectorError> {
Ok(())
}
impl HttpClient {
async fn run_request_policy(
&self,
tool: &str,
method: &str,
path: &str,
query: &std::collections::HashMap<String, String>,
body: Option<&serde_json::Value>,
) -> Result<(), HttpConnectorError> {
let Some(policy) = self.policy.as_ref() else {
return Ok(());
};
let mut sorted: Vec<(String, String)> =
query.iter().map(|(k, v)| (k.clone(), v.clone())).collect();
sorted.sort();
let method = method.to_uppercase();
let call_id = crate::policy::next_call_id();
let req = crate::policy::OutboundRequest::new(tool, &method, path, &sorted, body)
.with_call_id(&call_id);
policy
.check(&req)
.await
.map_err(|refusal| HttpConnectorError::PolicyRefused(refusal.message().to_string()))
}
async fn execute_inner(
&self,
tool: &str,
operation: &Operation,
args: &serde_json::Value,
) -> Result<serde_json::Value, HttpConnectorError> {
let empty = serde_json::Map::new();
let args_map = args.as_object().unwrap_or(&empty);
let substituted = Self::substitute_path(operation, args_map)?;
let joined = join_url(self.base_url.as_str(), &substituted);
let mut url = url::Url::parse(&joined)
.map_err(|_| HttpConnectorError::Backend("constructed URL is invalid".to_string()))?;
let mut query = Self::build_query(operation, args_map)?;
let mut headers = Self::build_headers(operation, args_map)?;
let request_body = Self::build_body(operation, args_map);
self.run_request_policy(
tool,
&operation.method,
&joined,
&query,
request_body.as_ref(),
)
.await?;
self.auth.apply(&mut headers, &mut query, None).await?;
if !query.is_empty() {
let mut pairs = url.query_pairs_mut();
for (key, value) in &query {
pairs.append_pair(key, value);
}
drop(pairs);
}
let method = Self::convert_method(&operation.method)?;
let mut request = self.client.request(method, url);
request = request.headers(headers);
if let Some(body) = request_body {
request = request.json(&body);
}
let response = self.send_with_retries(request).await?;
let status = response.status();
if !status.is_success() {
return Err(HttpConnectorError::Status {
status: status.as_u16(),
});
}
let body = response
.text()
.await
.map_err(|_| HttpConnectorError::Request("failed to read response body".to_string()))?;
if body.is_empty() {
return Ok(serde_json::Value::Null);
}
serde_json::from_str(&body).map_err(|_| {
HttpConnectorError::Backend("response body was not valid JSON".to_string())
})
}
fn cloned_with_policy(&self, policy: Arc<dyn crate::policy::RequestPolicy>) -> Self {
Self {
client: self.client.clone(),
base_url: self.base_url.clone(),
auth: Arc::clone(&self.auth),
http_config: self.http_config.clone(),
policy: Some(policy),
}
}
}
#[async_trait]
impl HttpConnector for HttpClient {
async fn execute(
&self,
operation: &Operation,
args: &serde_json::Value,
) -> Result<serde_json::Value, HttpConnectorError> {
self.execute_inner("", operation, args).await
}
async fn execute_for_tool(
&self,
tool: &str,
operation: &Operation,
args: &serde_json::Value,
) -> Result<serde_json::Value, HttpConnectorError> {
self.execute_inner(tool, operation, args).await
}
fn has_request_policy(&self) -> bool {
self.policy.is_some()
}
fn governed(
&self,
policy: Arc<dyn crate::policy::RequestPolicy>,
) -> Option<Arc<dyn HttpConnector>> {
Some(Arc::new(self.cloned_with_policy(policy)))
}
fn base_url(&self) -> &str {
self.base_url.as_str()
}
}
#[cfg(all(test, feature = "input-validation"))]
mod d4_support {
use super::{HttpClient, HttpConnectorError, Operation};
use crate::http::{Parameter, ParameterLocation};
pub fn op(path: &str, parameters: Vec<Parameter>) -> Operation {
Operation {
method: "GET".to_string(),
path: path.to_string(),
parameters,
has_request_body: false,
base_url: None,
}
}
pub fn path_param(name: &str) -> Parameter {
Parameter::new(name, ParameterLocation::Path, true)
}
pub fn substitute(
path: &str,
pairs: &[(&str, serde_json::Value)],
) -> Result<String, HttpConnectorError> {
let parameters = pairs.iter().map(|(k, _)| path_param(k)).collect();
let mut args = serde_json::Map::new();
for (k, v) in pairs {
args.insert((*k).to_string(), v.clone());
}
HttpClient::substitute_path(&op(path, parameters), &args)
}
pub fn substitute_one(
path: &str,
name: &str,
value: &str,
) -> Result<String, HttpConnectorError> {
substitute(
path,
&[(name, serde_json::Value::String(value.to_string()))],
)
}
}
#[cfg(all(test, feature = "input-validation"))]
mod placeholder_floor {
use super::d4_support::{op, path_param, substitute, substitute_one};
use super::{HttpClient, HttpConnectorError};
use crate::http::{Parameter, ParameterLocation};
use pmcp::server::schema_validation::PLACEHOLDER_MAX_LENGTH;
fn assert_value_free(err: &HttpConnectorError, param: &str, value: &str, path_fragment: &str) {
assert!(matches!(err, HttpConnectorError::Backend(_)), "{err}");
let rendered = err.to_string();
assert!(
rendered.contains(param),
"the refusal must name the declared parameter: {rendered}"
);
assert!(
!rendered.contains(value),
"the refusal must carry no byte of the value: {rendered}"
);
assert!(
!rendered.contains(path_fragment),
"the refusal must never contain the resolved path: {rendered}"
);
}
#[test]
fn placeholder_floor_refuses_a_query_separator_in_a_value() {
let value = "current?string=x";
let err = substitute_one("/content/{version}/CUI", "version", value).unwrap_err();
assert_value_free(&err, "version", value, "/content/");
}
#[test]
fn placeholder_floor_refuses_traversal_in_a_value() {
let value = "current/../../search/current";
let err = substitute_one("/content/{version}/CUI", "version", value).unwrap_err();
assert_value_free(&err, "version", value, "/content/");
}
#[test]
fn placeholder_floor_refuses_upper_case_encoded_traversal() {
let err = substitute_one("/content/{version}/CUI", "version", "a%2E%2Eb").unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)), "{err}");
}
#[test]
fn placeholder_floor_refuses_a_value_that_is_exactly_a_denied_character() {
let err = substitute_one("/content/{version}/CUI", "version", "?").unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)), "{err}");
}
#[test]
fn placeholder_floor_refuses_a_nul_byte_in_both_forms() {
assert!(substitute_one("/x/{v}", "v", "a\u{0}b").is_err());
assert!(substitute_one("/x/{v}", "v", "a%00b").is_err());
}
#[test]
fn placeholder_floor_refuses_an_empty_value() {
assert!(substitute_one("/x/{v}", "v", "").is_err());
}
#[test]
fn placeholder_floor_accepts_the_cap_and_refuses_one_more() {
let at_cap = "a".repeat(PLACEHOLDER_MAX_LENGTH);
assert_eq!(
substitute_one("/x/{v}", "v", &at_cap).expect("at the cap"),
format!("/x/{at_cap}")
);
let over_cap = "a".repeat(PLACEHOLDER_MAX_LENGTH + 1);
assert!(substitute_one("/x/{v}", "v", &over_cap).is_err());
}
#[test]
fn placeholder_floor_accepts_a_value_matching_its_declared_pattern() {
let parameters = vec![
Parameter::new("cui", ParameterLocation::Path, true).with_rules(
Some("^C[0-9]+$".to_string()),
Some(32),
false,
),
];
let mut args = serde_json::Map::new();
args.insert("cui".to_string(), serde_json::json!("C0018787"));
let resolved = HttpClient::substitute_path(&op("/CUI/{cui}/content", parameters), &args)
.expect("a conforming value must be accepted");
assert_eq!(resolved, "/CUI/C0018787/content");
}
#[test]
fn placeholder_floor_refuses_a_value_failing_its_declared_pattern() {
let parameters = vec![
Parameter::new("cui", ParameterLocation::Path, true).with_rules(
Some("^C[0-9]+$".to_string()),
None,
false,
),
];
let mut args = serde_json::Map::new();
args.insert("cui".to_string(), serde_json::json!("notacui"));
let err = HttpClient::substitute_path(&op("/CUI/{cui}", parameters), &args).unwrap_err();
assert!(err.to_string().contains("cui"), "{err}");
assert!(!err.to_string().contains("notacui"), "{err}");
}
#[test]
fn placeholder_floor_leaves_a_placeholder_free_template_untouched() {
let resolved = HttpClient::substitute_path(
&op("/Line/Mode/tube/Status", vec![]),
&serde_json::Map::new(),
)
.expect("a placeholder-free template must be unaffected");
assert_eq!(resolved, "/Line/Mode/tube/Status");
}
#[test]
fn placeholder_floor_accepts_the_root_path_and_still_refuses_a_trailing_slash() {
let resolved = HttpClient::substitute_path(&op("/", vec![]), &serde_json::Map::new())
.expect("a `GET /` operation must be callable — the root is the shortest legal path");
assert_eq!(resolved, "/");
let err = substitute_one("/search/{v}", "v", "")
.expect_err("an empty tail placeholder must stay refused");
assert!(
matches!(err, HttpConnectorError::Backend(_)),
"the refusal is a Backend error naming the position: {err}"
);
assert!(
HttpClient::substitute_path(&op("/search/", vec![]), &serde_json::Map::new()).is_err(),
"a literal trailing slash in the template stays refused by decision"
);
}
#[test]
fn placeholder_floor_refuses_the_second_of_two_placeholders_without_substituting() {
let err = substitute(
"/a/{first}/b/{second}",
&[
("first", serde_json::json!("ok")),
("second", serde_json::json!("../escape")),
],
)
.unwrap_err();
let rendered = err.to_string();
assert!(rendered.contains("second"), "{rendered}");
assert!(
!rendered.contains("/a/ok/b/"),
"no partially-substituted path may appear anywhere: {rendered}"
);
}
#[test]
fn placeholder_floor_refuses_an_absent_path_argument() {
let err = HttpClient::substitute_path(
&op("/users/{id}/profile", vec![path_param("id")]),
&serde_json::Map::new(),
)
.unwrap_err();
let rendered = err.to_string();
assert!(rendered.contains("id"), "{rendered}");
assert!(
!rendered.contains('{') && !rendered.contains('}'),
"the refusal must not echo the template: {rendered}"
);
assert!(
!rendered.contains("/users/"),
"the refusal must not echo the path: {rendered}"
);
}
#[test]
fn placeholder_floor_refuses_a_composed_segment_over_the_cap() {
let prefix = "p".repeat(100);
let value = "v".repeat(200);
let err = substitute_one(&format!("/x/{prefix}{{id}}"), "id", &value).unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)), "{err}");
}
#[test]
fn placeholder_floor_refuses_a_residual_brace_from_an_unrecognized_template() {
let err = HttpClient::substitute_path(&op("/x/{a}/y/{b}", vec![path_param("a")]), &{
let mut args = serde_json::Map::new();
args.insert("a".to_string(), serde_json::json!("ok"));
args
})
.unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)), "{err}");
}
#[test]
fn placeholder_floor_refuses_traversal_written_into_the_template_literal() {
let err = HttpClient::substitute_path(&op("/a/../b", vec![]), &serde_json::Map::new())
.unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)), "{err}");
}
}
#[cfg(all(test, feature = "input-validation"))]
mod query_separator {
use super::d4_support::substitute_one;
use super::{HttpClient, Operation};
use pmcp::server::schema_validation::PLACEHOLDER_MAX_LENGTH;
fn literal(path: &str) -> Result<String, super::HttpConnectorError> {
HttpClient::substitute_path(
&Operation {
method: "GET".to_string(),
path: path.to_string(),
parameters: vec![],
has_request_body: false,
base_url: None,
},
&serde_json::Map::new(),
)
}
#[test]
fn query_separator_accepts_an_author_written_query_string() {
assert_eq!(
literal("/Line/Mode/tube/Status?detail=true").expect("author query accepted"),
"/Line/Mode/tube/Status?detail=true"
);
}
#[test]
fn query_separator_accepts_a_literal_query_alongside_a_floored_placeholder() {
assert_eq!(
substitute_one("/content/{version}/CUI?string=x", "version", "current")
.expect("author query plus conforming placeholder accepted"),
"/content/current/CUI?string=x"
);
}
#[test]
fn query_separator_accepts_a_graph_style_dollar_projection() {
let resolved = literal(
"/drives/D/items/I/workbook/worksheets/C/range(address='A2:D7')?$select=values",
)
.expect("a Graph $select projection must be accepted");
assert!(resolved.ends_with("?$select=values"), "{resolved}");
}
#[test]
fn query_separator_still_refuses_traversal_in_the_path_portion() {
assert!(
literal("/a/../b?x=1").is_err(),
"appending a query must not launder a traversal"
);
}
#[test]
fn query_separator_still_refuses_traversal_in_the_query_portion() {
let err = literal("/search?next=../../etc/passwd").unwrap_err();
assert!(!err.to_string().contains("passwd"), "{err}");
}
#[test]
fn query_separator_still_refuses_a_control_byte_in_the_query_portion() {
assert!(literal("/search?x=a%00b").is_err());
}
#[test]
fn query_separator_still_refuses_an_over_cap_query_portion() {
let long = "z".repeat(PLACEHOLDER_MAX_LENGTH + 1);
assert!(literal(&format!("/search?q={long}")).is_err());
}
#[test]
fn query_separator_still_refuses_a_second_question_mark() {
assert!(
literal("/search?a=1?b=2").is_err(),
"only the FIRST `?` is split off; one exemption, not a licence"
);
}
#[test]
fn query_separator_still_refuses_an_empty_query_portion() {
assert!(
literal("/search?").is_err(),
"a dangling `?` is the same class as a trailing `/`"
);
}
#[test]
fn query_separator_still_refuses_a_fragment_marker() {
assert!(literal("/search#frag").is_err());
}
#[test]
fn query_separator_still_refuses_an_injected_separator_from_a_value() {
let payload = "2026AA?string=x";
let err = substitute_one("/search/{v}?detail=true", "v", payload).unwrap_err();
let rendered = err.to_string();
assert!(rendered.contains('v'), "{rendered}");
assert!(
!rendered.contains("2026AA") && !rendered.contains('?'),
"the refusal must carry no byte of the value: {rendered}"
);
}
#[test]
fn query_separator_still_refuses_an_injected_traversal_from_a_value() {
assert!(substitute_one("/search/{v}?detail=true", "v", "../../etc/passwd").is_err());
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::http::auth::NoAuth;
use crate::http::{Parameter, ParameterLocation};
fn get_user_op() -> Operation {
Operation {
method: "GET".to_string(),
path: "/users/{id}".to_string(),
parameters: vec![
Parameter::new("id", ParameterLocation::Path, true),
Parameter::new("verbose", ParameterLocation::Query, false),
],
has_request_body: false,
base_url: None,
}
}
#[test]
fn test_build_url_with_path_prefix() {
let client = HttpClient::new(
reqwest::Client::new(),
"https://xxx.execute-api.eu-west-1.amazonaws.com/v1/".to_string(),
Arc::new(NoAuth),
)
.unwrap();
let op = get_user_op();
let mut args = serde_json::Map::new();
args.insert("id".to_string(), serde_json::json!("42"));
let substituted = HttpClient::substitute_path(&op, &args).unwrap();
let joined = join_url(client.base_url(), &substituted);
assert_eq!(
joined,
"https://xxx.execute-api.eu-west-1.amazonaws.com/v1/users/42"
);
}
#[test]
fn test_substitute_path_replaces_placeholder() {
let op = get_user_op();
let mut args = serde_json::Map::new();
args.insert("id".to_string(), serde_json::json!(7));
assert_eq!(HttpClient::substitute_path(&op, &args).unwrap(), "/users/7");
}
#[test]
fn test_build_query_skips_path_params() {
let op = get_user_op();
let mut args = serde_json::Map::new();
args.insert("id".to_string(), serde_json::json!("42"));
args.insert("verbose".to_string(), serde_json::json!(true));
let query = HttpClient::build_query(&op, &args).unwrap();
assert_eq!(query.get("verbose"), Some(&"true".to_string()));
assert!(!query.contains_key("id"));
}
#[test]
fn render_query_value_comma_joins_scalar_array() {
let rendered =
HttpClient::render_query_value("tags", &serde_json::json!(["a", 2, true])).unwrap();
assert_eq!(rendered, "a,2,true");
}
#[test]
fn render_query_value_scalar_passthrough() {
assert_eq!(
HttpClient::render_query_value("q", &serde_json::json!("hi")).unwrap(),
"hi"
);
assert_eq!(
HttpClient::render_query_value("n", &serde_json::json!(7)).unwrap(),
"7"
);
}
#[test]
fn render_scalar_null_is_bare_null() {
assert_eq!(
render_scalar("x", &serde_json::Value::Null).unwrap(),
"null"
);
}
#[test]
fn substitute_path_rejects_object_param() {
let op = get_user_op();
let mut args = serde_json::Map::new();
args.insert("id".to_string(), serde_json::json!({"nested": "x"}));
let err = HttpClient::substitute_path(&op, &args).unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)));
let rendered = err.to_string();
assert!(
rendered.contains("id"),
"error must name the param: {rendered}"
);
for forbidden in ['{', '[', '"'] {
assert!(
!rendered.contains(forbidden),
"must not echo JSON: {rendered}"
);
}
assert!(
!rendered.contains("nested"),
"must not echo the value: {rendered}"
);
}
#[test]
fn build_query_rejects_object_param() {
let op = get_user_op();
let mut args = serde_json::Map::new();
args.insert("verbose".to_string(), serde_json::json!({"k": "v"}));
let err = HttpClient::build_query(&op, &args).unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)));
assert!(err.to_string().contains("verbose"));
}
#[test]
fn render_query_value_rejects_array_with_object_member() {
let err = HttpClient::render_query_value("tags", &serde_json::json!(["ok", {"bad": 1}]))
.unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)));
assert!(err.to_string().contains("tags"));
}
#[test]
fn build_headers_rejects_non_scalar_param() {
let op = Operation {
method: "GET".to_string(),
path: "/x".to_string(),
parameters: vec![Parameter::new("x-trace", ParameterLocation::Header, false)],
has_request_body: false,
base_url: None,
};
let mut args = serde_json::Map::new();
args.insert("x-trace".to_string(), serde_json::json!(["a", "b"]));
let err = HttpClient::build_headers(&op, &args).unwrap_err();
assert!(matches!(err, HttpConnectorError::Backend(_)));
assert!(err.to_string().contains("x-trace"));
let mut args2 = serde_json::Map::new();
args2.insert("x-trace".to_string(), serde_json::json!({"k": "v"}));
let err2 = HttpClient::build_headers(&op, &args2).unwrap_err();
assert!(matches!(err2, HttpConnectorError::Backend(_)));
assert!(err2.to_string().contains("x-trace"));
let mut args3 = serde_json::Map::new();
args3.insert("x-trace".to_string(), serde_json::json!("abc"));
let headers = HttpClient::build_headers(&op, &args3).unwrap();
assert_eq!(headers.get("x-trace").unwrap(), "abc");
}
#[test]
fn test_new_is_lazy_and_rejects_bad_url() {
let err = HttpClient::new(
reqwest::Client::new(),
"not a url".to_string(),
Arc::new(NoAuth),
)
.err()
.expect("bad URL should error");
assert!(matches!(err, HttpConnectorError::Backend(_)));
let rendered = err.to_string();
assert!(!rendered.contains("not a url"), "must not echo the bad URL");
}
#[tokio::test]
async fn http_connector_get_returns_json() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/users/42"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(serde_json::json!({"id": 42, "name": "Ada"})),
)
.mount(&server)
.await;
let client =
HttpClient::new(reqwest::Client::new(), server.uri(), Arc::new(NoAuth)).unwrap();
let op = get_user_op();
let args = serde_json::json!({"id": "42"});
let result = client.execute(&op, &args).await.unwrap();
assert_eq!(result["id"], 42);
assert_eq!(result["name"], "Ada");
}
#[tokio::test]
async fn http_connector_post_sends_body_and_auth() {
use wiremock::matchers::{body_json, header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/items"))
.and(header("authorization", "Bearer tok"))
.and(body_json(serde_json::json!({"name": "widget"})))
.respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({"ok": true})))
.mount(&server)
.await;
let auth = crate::http::auth::create_auth_provider(&crate::http::AuthConfig::Bearer {
token: "tok".to_string(),
required: true,
})
.unwrap();
let client = HttpClient::new(reqwest::Client::new(), server.uri(), auth).unwrap();
let op = Operation {
method: "POST".to_string(),
path: "/items".to_string(),
parameters: vec![],
has_request_body: true,
base_url: None,
};
let args = serde_json::json!({"name": "widget"});
let result = client.execute(&op, &args).await.unwrap();
assert_eq!(result["ok"], true);
}
#[tokio::test]
async fn http_connector_post_sends_declared_body_parameters_as_the_payload() {
use wiremock::matchers::{body_json, method, path, query_param_is_missing};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/items"))
.and(body_json(
serde_json::json!({"name": "widget", "note": "free text"}),
))
.and(query_param_is_missing("name"))
.and(query_param_is_missing("note"))
.respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({"ok": true})))
.mount(&server)
.await;
let client =
HttpClient::new(reqwest::Client::new(), server.uri(), Arc::new(NoAuth)).unwrap();
let op = Operation {
method: "POST".to_string(),
path: "/items".to_string(),
parameters: vec![
Parameter::new("name", ParameterLocation::Body, true),
Parameter::new("note", ParameterLocation::Body, false),
],
has_request_body: true,
base_url: None,
};
let args = serde_json::json!({"name": "widget", "note": "free text"});
let result = client.execute(&op, &args).await.unwrap();
assert_eq!(result["ok"], true);
}
#[test]
fn build_body_withholds_a_query_located_parameter_on_a_post() {
let op = Operation {
method: "POST".to_string(),
path: "/items".to_string(),
parameters: vec![
Parameter::new("dry_run", ParameterLocation::Query, false),
Parameter::new("name", ParameterLocation::Body, true),
],
has_request_body: true,
base_url: None,
};
let args = serde_json::json!({"dry_run": "true", "name": "widget"})
.as_object()
.expect("object")
.clone();
let body = HttpClient::build_body(&op, &args).expect("a body is built");
assert_eq!(body, serde_json::json!({"name": "widget"}));
let query = HttpClient::build_query(&op, &args).expect("a query is built");
assert_eq!(query.get("dry_run").map(String::as_str), Some("true"));
assert!(
!query.contains_key("name"),
"a Body-located parameter must not reach the query string: {query:?}"
);
}
#[tokio::test]
async fn http_connector_maps_non_2xx_to_status_without_url() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/users/42"))
.respond_with(ResponseTemplate::new(404))
.mount(&server)
.await;
let client =
HttpClient::new(reqwest::Client::new(), server.uri(), Arc::new(NoAuth)).unwrap();
let op = get_user_op();
let args = serde_json::json!({"id": "42"});
let err = client.execute(&op, &args).await.unwrap_err();
assert!(matches!(err, HttpConnectorError::Status { status: 404 }));
let rendered = err.to_string();
assert!(rendered.contains("404"));
assert!(
!rendered.contains("http://"),
"status error must not echo the URL"
);
}
}
#[cfg(test)]
mod request_policy_seam {
use super::{HttpClient, HttpConfig, HttpConnectorError};
use crate::http::auth::HttpAuthProvider;
use crate::http::{HttpConnector, Operation, Parameter, ParameterLocation};
use crate::policy::{OutboundRequest, PolicyRefusal, RequestPolicy};
use async_trait::async_trait;
use reqwest::header::{HeaderMap, HeaderValue};
use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
struct RecordingAuth {
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl HttpAuthProvider for RecordingAuth {
async fn apply(
&self,
headers: &mut HeaderMap,
_query: &mut HashMap<String, String>,
_inbound_token: Option<&str>,
) -> Result<(), HttpConnectorError> {
self.calls.fetch_add(1, Ordering::SeqCst);
headers.insert("authorization", HeaderValue::from_static("Bearer tok"));
Ok(())
}
}
type Seen = Arc<Mutex<Vec<(String, String, Vec<(String, String)>, Option<String>)>>>;
struct Recorder {
seen: Seen,
refuse: Option<&'static str>,
}
#[async_trait]
impl RequestPolicy for Recorder {
async fn check(&self, req: &OutboundRequest<'_>) -> Result<(), PolicyRefusal> {
self.seen.lock().expect("lock").push((
req.tool.to_string(),
req.path.to_string(),
req.query.to_vec(),
req.body.map(ToString::to_string),
));
match self.refuse {
Some(msg) => Err(PolicyRefusal::new(msg)),
None => Ok(()),
}
}
}
fn op() -> Operation {
Operation {
method: "GET".to_string(),
path: "/users/{id}".to_string(),
parameters: vec![
Parameter::new("id", ParameterLocation::Path, true),
Parameter::new("q", ParameterLocation::Query, false),
],
has_request_body: false,
base_url: None,
}
}
fn client(
base_url: String,
policy: Option<Arc<dyn RequestPolicy>>,
) -> (HttpClient, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
let auth = Arc::new(RecordingAuth {
calls: Arc::clone(&calls),
});
let cfg = HttpConfig {
retries: 0,
..HttpConfig::default()
};
let c = HttpClient::with_config(reqwest::Client::new(), base_url, auth, cfg)
.expect("client builds");
let c = match policy {
Some(p) => c.with_request_policy(p),
None => c,
};
(c, calls)
}
fn recorder(refuse: Option<&'static str>) -> (Arc<Recorder>, Seen) {
let seen: Seen = Arc::new(Mutex::new(Vec::new()));
(
Arc::new(Recorder {
seen: Arc::clone(&seen),
refuse,
}),
seen,
)
}
#[tokio::test]
async fn a_refusing_policy_stops_the_request_before_auth_and_before_the_send() {
use wiremock::MockServer;
let server = MockServer::start().await;
let (policy, _seen) = recorder(Some("refused by test policy"));
let (client, auth_calls) = client(server.uri(), Some(policy));
let err = client
.execute(&op(), &serde_json::json!({ "id": "42" }))
.await
.expect_err("the policy refuses");
assert_eq!(
err.to_string(),
"outbound request refused by policy: refused by test policy",
"the refusal must carry the policy's own message"
);
assert_eq!(
auth_calls.load(Ordering::SeqCst),
0,
"the auth provider must NOT have been invoked — the hook is before auth"
);
let requests = server
.received_requests()
.await
.expect("wiremock records requests");
assert!(requests.is_empty(), "a refusal must send nothing");
}
#[tokio::test]
async fn the_same_refused_call_twice_is_identical_and_sends_nothing() {
use wiremock::MockServer;
let server = MockServer::start().await;
let (policy, _seen) = recorder(Some("refused by test policy"));
let (client, _auth) = client(server.uri(), Some(policy));
let first = client
.execute(&op(), &serde_json::json!({ "id": "42" }))
.await
.expect_err("refuses");
let second = client
.execute(&op(), &serde_json::json!({ "id": "42" }))
.await
.expect_err("refuses again");
assert_eq!(first.to_string(), second.to_string());
assert!(server
.received_requests()
.await
.expect("recorded")
.is_empty());
}
#[tokio::test]
async fn an_allowing_policy_lets_the_request_through_and_auth_is_applied() {
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/users/42"))
.and(header("authorization", "Bearer tok"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"ok": true})))
.mount(&server)
.await;
let (policy, _seen) = recorder(None);
let (client, auth_calls) = client(server.uri(), Some(policy));
let out = client
.execute(&op(), &serde_json::json!({ "id": "42" }))
.await
.expect("allowed");
assert_eq!(out["ok"], true);
assert_eq!(auth_calls.load(Ordering::SeqCst), 1);
assert_eq!(server.received_requests().await.expect("recorded").len(), 1);
}
#[tokio::test]
async fn the_policy_sees_the_resolved_path_and_the_query_pairs() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/users/42"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({})))
.mount(&server)
.await;
let (policy, seen) = recorder(None);
let (client, _auth) = client(server.uri(), Some(policy));
client
.execute(&op(), &serde_json::json!({ "id": "42", "q": "hay" }))
.await
.expect("allowed");
let seen = seen.lock().expect("lock");
assert_eq!(seen.len(), 1, "exactly one invocation per logical request");
let (_tool, observed_path, query, body) = &seen[0];
assert!(
observed_path.ends_with("/users/42"),
"the policy must see the SUBSTITUTED path, got {observed_path}"
);
assert!(
!observed_path.contains('{'),
"the policy must never see the template"
);
assert_eq!(query.as_slice(), &[("q".to_string(), "hay".to_string())]);
assert!(body.is_none(), "a GET carries no body");
}
#[tokio::test]
async fn no_policy_behaves_exactly_as_before() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/users/42"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"ok": true})))
.mount(&server)
.await;
let (client, auth_calls) = client(server.uri(), None);
assert!(!client.has_request_policy());
let out = client
.execute(&op(), &serde_json::json!({ "id": "42" }))
.await
.expect("succeeds");
assert_eq!(out["ok"], true);
assert_eq!(auth_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn the_policy_is_told_which_tool_the_call_came_from() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/users/42"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({})))
.mount(&server)
.await;
let (policy, seen) = recorder(None);
let (client, _auth) = client(server.uri(), Some(policy));
client
.execute_for_tool("get_user", &op(), &serde_json::json!({ "id": "42" }))
.await
.expect("allowed");
let seen = seen.lock().expect("lock");
assert_eq!(seen[0].0, "get_user");
}
#[test]
fn a_governed_connector_reports_its_policy_through_the_dyn_trait() {
let (policy, _seen) = recorder(None);
let (client, _auth) = client("https://example.test".to_string(), None);
let bare: Arc<dyn HttpConnector> = Arc::new(client);
assert!(
!bare.has_request_policy(),
"a bare connector carries no policy"
);
let governed = bare
.governed(policy)
.expect("HttpClient supports policy attachment");
assert!(
governed.has_request_policy(),
"a registered policy must be observable on the dyn connector, or it could \
look registered while never running"
);
}
}