use std::{collections::BTreeSet, fmt, sync::Arc, time::Duration};
use http::header::{CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue, LOCATION, USER_AGENT};
use url::Url;
use crate::{
Result, VERSION,
config::{Authentication, DEFAULT_BASE_URL, DEFAULT_PATH_PREFIX, RedirectPolicy},
endpoints::QueryEncoder,
error::{ConfigurationErrorKind, Error, Redactor, SafeBody, SecretString},
transport::{HttpExecutor, PreparedRequest, ReqwestExecutor, TransportResponse},
};
pub use crate::endpoints::{EndpointSpec, QueryParameters};
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
const MAX_REDIRECTS: usize = 10;
#[derive(Clone)]
struct AuthMaterial {
header: Option<(HeaderName, HeaderValue)>,
query: Option<(String, SecretString)>,
}
impl fmt::Debug for AuthMaterial {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("AuthMaterial")
.field("header_name", &self.header.as_ref().map(|(name, _)| name))
.field("query_name", &self.query.as_ref().map(|(name, _)| name))
.finish()
}
}
#[derive(Clone)]
pub struct ClientBuilder {
base_url: String,
path_prefix: String,
authentication: Authentication,
authentication_configured: bool,
authentication_conflict: bool,
default_headers: Vec<(String, String)>,
user_agent: String,
timeout: Duration,
connect_timeout: Duration,
redirect_policy: RedirectPolicy,
executor: Option<Arc<dyn HttpExecutor>>,
}
impl Default for ClientBuilder {
fn default() -> Self {
Self {
base_url: DEFAULT_BASE_URL.to_owned(),
path_prefix: DEFAULT_PATH_PREFIX.to_owned(),
authentication: Authentication::None,
authentication_configured: false,
authentication_conflict: false,
default_headers: Vec::new(),
user_agent: format!("libfmp/{VERSION}"),
timeout: DEFAULT_TIMEOUT,
connect_timeout: DEFAULT_CONNECT_TIMEOUT,
redirect_policy: RedirectPolicy::SameOrigin,
executor: None,
}
}
}
impl fmt::Debug for ClientBuilder {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let header_names: Vec<_> = self
.default_headers
.iter()
.map(|(name, _)| name.as_str())
.collect();
formatter
.debug_struct("ClientBuilder")
.field("base_url", &"[CONFIGURED URL]")
.field("path_prefix", &"[CONFIGURED PATH]")
.field("authentication", &self.authentication)
.field("authentication_conflict", &self.authentication_conflict)
.field("default_header_names", &header_names)
.field("user_agent", &"[CONFIGURED USER AGENT]")
.field("timeout", &self.timeout)
.field("connect_timeout", &self.connect_timeout)
.field("redirect_policy", &self.redirect_policy)
.field("custom_executor", &self.executor.is_some())
.finish()
}
}
impl ClientBuilder {
pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
self.base_url = base_url.into();
self
}
pub fn path_prefix(mut self, path_prefix: impl Into<String>) -> Self {
self.path_prefix = path_prefix.into();
self
}
pub fn authentication(mut self, authentication: Authentication) -> Self {
if self.authentication_configured {
self.authentication_conflict = true;
}
self.authentication_configured = true;
self.authentication = authentication;
self
}
pub fn default_header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.default_headers.push((name.into(), value.into()));
self
}
pub fn user_agent(mut self, user_agent: impl Into<String>) -> Self {
self.user_agent = user_agent.into();
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
pub fn connect_timeout(mut self, connect_timeout: Duration) -> Self {
self.connect_timeout = connect_timeout;
self
}
pub fn redirect_policy(mut self, redirect_policy: RedirectPolicy) -> Self {
self.redirect_policy = redirect_policy;
self
}
pub fn executor(mut self, executor: Arc<dyn HttpExecutor>) -> Self {
self.executor = Some(executor);
self
}
pub fn build(self) -> Result<Client> {
if self.authentication_conflict {
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::ConflictingAuthentication,
"authentication can be configured only once",
));
}
let base_url = parse_base_url(&self.base_url)?;
validate_relative_path(&self.path_prefix)?;
let auth = build_auth_material(&self.authentication)?;
let mut protected_headers = reserved_header_names();
if let Some((name, _)) = &auth.header {
protected_headers.insert(name.as_str().to_ascii_lowercase());
}
let mut default_headers = parse_headers(&self.default_headers, &protected_headers)?;
let user_agent = parse_header_value(&self.user_agent)?;
default_headers.insert(USER_AGENT, user_agent);
if matches!(self.authentication, Authentication::None) && is_default_fmp_origin(&base_url) {
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::MissingCredential,
"direct FMP access requires explicit authentication",
));
}
let mut redactor = Redactor::new();
register_auth_redaction(&mut redactor, &self.authentication)?;
redactor.add_secret(&SecretString::new(self.base_url.clone()));
redactor.add_secret(&SecretString::new(base_url.as_str().to_owned()));
redactor.add_secret(&SecretString::new(self.path_prefix.clone()));
redactor.add_secret(&SecretString::new(self.user_agent.clone()));
for (_, value) in &self.default_headers {
redactor.add_secret(&SecretString::new(value.clone()));
}
let executor = match self.executor {
Some(executor) => executor,
None => Arc::new(
ReqwestExecutor::new(self.timeout, self.connect_timeout).map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::HttpClient,
"HTTP client could not be constructed",
)
})?,
),
};
Ok(Client {
inner: Arc::new(ClientInner {
base_url,
path_prefix: self.path_prefix,
auth,
default_headers,
protected_headers,
redactor,
executor,
redirect_policy: self.redirect_policy,
}),
})
}
}
struct ClientInner {
base_url: Url,
path_prefix: String,
auth: AuthMaterial,
default_headers: HeaderMap,
protected_headers: BTreeSet<String>,
redactor: Redactor,
executor: Arc<dyn HttpExecutor>,
redirect_policy: RedirectPolicy,
}
#[derive(Clone)]
pub struct Client {
inner: Arc<ClientInner>,
}
impl fmt::Debug for Client {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Client")
.field("base_url", &"[CONFIGURED URL]")
.field("path_prefix", &"[CONFIGURED PATH]")
.field("authentication", &self.inner.auth)
.field("redirect_policy", &self.inner.redirect_policy)
.field(
"default_header_names",
&self
.inner
.default_headers
.keys()
.map(HeaderName::as_str)
.collect::<Vec<_>>(),
)
.finish_non_exhaustive()
}
}
impl Client {
pub fn builder() -> ClientBuilder {
ClientBuilder::default()
}
pub async fn execute<Q, R>(&self, endpoint: &EndpointSpec<Q, R>) -> Result<R>
where
Q: QueryParameters,
{
self.execute_with_headers(endpoint, &[]).await
}
pub(crate) async fn execute_with_headers<Q, R>(
&self,
endpoint: &EndpointSpec<Q, R>,
request_headers: &[(&str, &str)],
) -> Result<R>
where
Q: QueryParameters,
{
validate_relative_path(endpoint.relative_path())?;
let mut url = build_endpoint_url(
&self.inner.base_url,
&self.inner.path_prefix,
endpoint.relative_path(),
)?;
append_endpoint_query(&mut url, endpoint.query(), self.inner.auth.query.as_ref())?;
let mut headers = self.inner.default_headers.clone();
merge_request_headers(&mut headers, request_headers, &self.inner.protected_headers)?;
if let Some((name, value)) = &self.inner.auth.header {
headers.insert(name.clone(), value.clone());
}
let mut redactor = self.inner.redactor.clone();
for (_, value) in request_headers {
redactor.add_secret(&SecretString::new((*value).to_owned()));
}
self.execute_redirects(endpoint, url, headers, &redactor)
.await
}
async fn execute_redirects<Q, R>(
&self,
endpoint: &EndpointSpec<Q, R>,
mut url: Url,
headers: HeaderMap,
redactor: &Redactor,
) -> Result<R> {
for redirect_count in 0..=MAX_REDIRECTS {
apply_query_auth(&mut url, self.inner.auth.query.as_ref());
let response = self
.inner
.executor
.execute(PreparedRequest::new(
endpoint.method(),
url.clone(),
headers.clone(),
))
.await
.map_err(|_| Error::transport(Some(endpoint.id()), "request execution failed"))?;
if is_redirect(response.status()) {
if self.inner.redirect_policy == RedirectPolicy::None {
return status_error(endpoint.id(), response, redactor);
}
if redirect_count == MAX_REDIRECTS {
return Err(Error::transport(
Some(endpoint.id()),
"same-origin redirect limit exceeded",
));
}
let Some(location) = response.headers().get(LOCATION) else {
return status_error(endpoint.id(), response, redactor);
};
let Ok(location) = location.to_str() else {
return status_error(endpoint.id(), response, redactor);
};
let Ok(mut destination) = url.join(location) else {
return status_error(endpoint.id(), response, redactor);
};
destination.set_fragment(None);
if destination.username() != ""
|| destination.password().is_some()
|| !same_origin(&url, &destination)
{
return status_error(endpoint.id(), response, redactor);
}
url = destination;
continue;
}
if !(200..300).contains(&response.status()) {
return status_error(endpoint.id(), response, redactor);
}
let Some(content_type) = response.headers().get(CONTENT_TYPE) else {
return Err(Error::decode(
Some(endpoint.id()),
Some(response.status()),
Some(safe_body(response.body(), redactor)),
"successful response omitted its content type",
));
};
let Ok(content_type) = content_type.to_str() else {
return Err(Error::decode(
Some(endpoint.id()),
Some(response.status()),
Some(safe_body(response.body(), redactor)),
"successful response used an invalid content type",
));
};
if !endpoint
.response()
.expected_content_type()
.matches(content_type)
{
return Err(Error::decode(
Some(endpoint.id()),
Some(response.status()),
Some(safe_body(response.body(), redactor)),
"successful response used an unexpected content type",
));
}
return endpoint.response().decode(response.body()).map_err(|_| {
Error::decode(
Some(endpoint.id()),
Some(response.status()),
Some(safe_body(response.body(), redactor)),
"successful response could not be decoded",
)
});
}
Err(Error::transport(
Some(endpoint.id()),
"redirect processing ended unexpectedly",
))
}
}
fn parse_base_url(value: &str) -> Result<Url> {
let url = Url::parse(value).map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::InvalidBaseUrl,
"base URL must be an absolute HTTP(S) URL",
)
})?;
if !matches!(url.scheme(), "http" | "https") || url.host_str().is_none() {
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::InvalidBaseUrl,
"base URL must be an absolute HTTP(S) URL",
));
}
if url.username() != ""
|| url.password().is_some()
|| url.query().is_some()
|| url.fragment().is_some()
{
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::UnsafeBaseUrl,
"base URL must not contain credentials, a query, or a fragment",
));
}
Ok(url)
}
fn is_default_fmp_origin(url: &Url) -> bool {
url.scheme() == "https"
&& url.host_str() == Some("financialmodelingprep.com")
&& url.port_or_known_default() == Some(443)
}
fn validate_relative_path(value: &str) -> Result<()> {
if value.chars().any(char::is_control)
|| value.contains(['\\', '?', '#'])
|| value
.trim_matches('/')
.split('/')
.any(|segment| matches!(segment, "." | ".."))
{
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::InvalidPath,
"path must contain only safe relative segments",
));
}
Ok(())
}
fn build_endpoint_url(base: &Url, prefix: &str, relative_path: &str) -> Result<Url> {
let mut url = base.clone();
{
let mut segments = url.path_segments_mut().map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::InvalidBaseUrl,
"base URL cannot contain path segments",
)
})?;
segments.pop_if_empty();
for segment in prefix
.trim_matches('/')
.split('/')
.chain(relative_path.trim_matches('/').split('/'))
.filter(|segment| !segment.is_empty())
{
segments.push(segment);
}
}
Ok(url)
}
fn append_endpoint_query<Q: QueryParameters>(
url: &mut Url,
query: &Q,
auth_query: Option<&(String, SecretString)>,
) -> Result<()> {
let mut failure = None;
query.encode(&mut QueryEncoder::new(&mut |name, value| {
if failure.is_some() {
return;
}
if name.is_empty() || name.chars().any(char::is_control) {
failure = Some(Error::configuration_with_kind(
ConfigurationErrorKind::InvalidQueryName,
"query parameter name must not be empty or contain controls",
));
return;
}
if auth_query.is_some_and(|(protected, _)| name.eq_ignore_ascii_case(protected)) {
failure = Some(Error::configuration_with_kind(
ConfigurationErrorKind::ProtectedFieldCollision,
"endpoint query collides with transport authentication",
));
return;
}
url.query_pairs_mut().append_pair(name, value);
}));
failure.map_or(Ok(()), Err)
}
fn apply_query_auth(url: &mut Url, auth_query: Option<&(String, SecretString)>) {
let Some((protected_name, secret)) = auth_query else {
return;
};
let retained: Vec<(String, String)> = url
.query_pairs()
.filter(|(name, _)| !name.eq_ignore_ascii_case(protected_name))
.map(|(name, value)| (name.into_owned(), value.into_owned()))
.collect();
url.set_query(None);
{
let mut pairs = url.query_pairs_mut();
pairs.extend_pairs(retained);
pairs.append_pair(protected_name, secret.expose_secret());
}
}
fn build_auth_material(authentication: &Authentication) -> Result<AuthMaterial> {
match authentication {
Authentication::None => Ok(AuthMaterial {
header: None,
query: None,
}),
Authentication::FmpHeader(secret) => Ok(AuthMaterial {
header: Some(secret_header("apikey", "", secret)?),
query: None,
}),
Authentication::FmpQuery(secret) => {
validate_credential(secret)?;
Ok(AuthMaterial {
header: None,
query: Some(("apikey".to_owned(), secret.clone())),
})
}
Authentication::Bearer(secret) => Ok(AuthMaterial {
header: Some(secret_header("authorization", "Bearer ", secret)?),
query: None,
}),
Authentication::CustomHeader {
name,
prefix,
secret,
} => Ok(AuthMaterial {
header: Some(secret_header(
name,
prefix.as_deref().unwrap_or(""),
secret,
)?),
query: None,
}),
Authentication::CustomQuery { name, secret } => {
validate_query_name(name)?;
validate_credential(secret)?;
Ok(AuthMaterial {
header: None,
query: Some((name.clone(), secret.clone())),
})
}
}
}
fn secret_header(
name: &str,
prefix: &str,
secret: &SecretString,
) -> Result<(HeaderName, HeaderValue)> {
validate_credential(secret)?;
let name = parse_header_name(name)?;
if is_unsafe_transport_header(&name) {
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::ProtectedFieldCollision,
"authentication header is owned by HTTP framing or transport",
));
}
let mut value = parse_header_value(&format!("{prefix}{}", secret.expose_secret()))?;
value.set_sensitive(true);
Ok((name, value))
}
fn validate_credential(secret: &SecretString) -> Result<()> {
if secret.expose_secret().is_empty() {
Err(Error::configuration_with_kind(
ConfigurationErrorKind::EmptyCredential,
"authentication credential must not be empty",
))
} else {
Ok(())
}
}
fn validate_query_name(name: &str) -> Result<()> {
if name.is_empty() || name.chars().any(char::is_control) {
Err(Error::configuration_with_kind(
ConfigurationErrorKind::InvalidQueryName,
"secret query name must not be empty or contain controls",
))
} else {
Ok(())
}
}
fn parse_header_name(name: &str) -> Result<HeaderName> {
HeaderName::from_bytes(name.as_bytes()).map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::InvalidHeaderName,
"header name is not a valid HTTP field name",
)
})
}
fn parse_header_value(value: &str) -> Result<HeaderValue> {
HeaderValue::from_str(value).map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::InvalidHeaderValue,
"header value is not a valid HTTP field value",
)
})
}
fn parse_headers(headers: &[(String, String)], protected: &BTreeSet<String>) -> Result<HeaderMap> {
let mut parsed = HeaderMap::new();
for (name, value) in headers {
let name = parse_header_name(name)?;
if protected.contains(name.as_str()) {
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::ProtectedFieldCollision,
"default header collides with a transport-owned header",
));
}
let mut value = parse_header_value(value)?;
if is_conventionally_secret_header(&name) {
value.set_sensitive(true);
}
parsed.insert(name, value);
}
Ok(parsed)
}
fn merge_request_headers(
headers: &mut HeaderMap,
request_headers: &[(&str, &str)],
protected: &BTreeSet<String>,
) -> Result<()> {
for (name, value) in request_headers {
let name = parse_header_name(name)?;
if protected.contains(name.as_str()) {
return Err(Error::configuration_with_kind(
ConfigurationErrorKind::ProtectedFieldCollision,
"request header collides with a transport-owned header",
));
}
let mut value = parse_header_value(value)?;
if is_conventionally_secret_header(&name) {
value.set_sensitive(true);
}
headers.insert(name, value);
}
Ok(())
}
fn register_auth_redaction(redactor: &mut Redactor, authentication: &Authentication) -> Result<()> {
match authentication {
Authentication::None => {}
Authentication::FmpHeader(secret) | Authentication::Bearer(secret) => {
redactor.add_secret(secret);
}
Authentication::FmpQuery(secret) => {
redactor.add_secret(secret);
redactor
.add_secret_query_name("apikey")
.map_err(|_| Error::configuration("invalid built-in secret query name"))?;
}
Authentication::CustomHeader { name, secret, .. } => {
redactor.add_secret(secret);
redactor.add_secret_header_name(name.clone()).map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::InvalidHeaderName,
"custom secret header name is invalid",
)
})?;
}
Authentication::CustomQuery { name, secret } => {
redactor.add_secret(secret);
redactor.add_secret_query_name(name.clone()).map_err(|_| {
Error::configuration_with_kind(
ConfigurationErrorKind::InvalidQueryName,
"custom secret query name is invalid",
)
})?;
}
}
Ok(())
}
fn is_conventionally_secret_header(name: &HeaderName) -> bool {
matches!(
name.as_str(),
"authorization" | "proxy-authorization" | "apikey" | "x-api-key"
)
}
fn reserved_header_names() -> BTreeSet<String> {
[
"apikey",
"authorization",
"connection",
"content-length",
"cookie",
"host",
"http2-settings",
"keep-alive",
"proxy-authorization",
"proxy-connection",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"user-agent",
"x-api-key",
]
.into_iter()
.map(str::to_owned)
.collect()
}
fn is_unsafe_transport_header(name: &HeaderName) -> bool {
matches!(
name.as_str(),
"connection"
| "content-length"
| "host"
| "http2-settings"
| "keep-alive"
| "proxy-authorization"
| "proxy-connection"
| "te"
| "trailer"
| "transfer-encoding"
| "upgrade"
| "user-agent"
)
}
fn is_redirect(status: u16) -> bool {
matches!(status, 301 | 302 | 303 | 307 | 308)
}
fn same_origin(left: &Url, right: &Url) -> bool {
left.scheme() == right.scheme()
&& left.host() == right.host()
&& left.port_or_known_default() == right.port_or_known_default()
}
fn status_error<R>(
endpoint: &'static str,
response: TransportResponse,
redactor: &Redactor,
) -> Result<R> {
Err(Error::status(
endpoint,
response.status(),
Some(safe_body(response.body(), redactor)),
))
}
fn safe_body(body: &[u8], redactor: &Redactor) -> SafeBody {
SafeBody::new(&String::from_utf8_lossy(body), redactor)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn request_headers_override_defaults_except_reserved_fields() {
let protected = reserved_header_names();
let mut headers =
parse_headers(&[("x-mode".to_owned(), "default".to_owned())], &protected).unwrap();
merge_request_headers(&mut headers, &[("X-Mode", "request")], &protected).unwrap();
assert_eq!(headers["x-mode"], "request");
for name in [
"AUTHORIZATION",
"apikey",
"X-API-KEY",
"Cookie",
"Host",
"Content-Length",
"Connection",
"Transfer-Encoding",
"User-Agent",
] {
let error = merge_request_headers(&mut headers, &[(name, "replacement")], &protected)
.unwrap_err();
assert_eq!(
error.configuration_kind(),
Some(ConfigurationErrorKind::ProtectedFieldCollision)
);
}
}
#[test]
fn origin_matching_normalizes_default_ports_but_not_scheme() {
let https_default = Url::parse("https://example.test/path").unwrap();
let https_explicit = Url::parse("https://example.test:443/other").unwrap();
let different_port = Url::parse("https://example.test:444/path").unwrap();
let different_scheme = Url::parse("http://example.test:443/path").unwrap();
assert!(same_origin(&https_default, &https_explicit));
assert!(!same_origin(&https_default, &different_port));
assert!(!same_origin(&https_default, &different_scheme));
}
#[test]
fn leading_slashes_never_reset_the_base_path() {
let base = Url::parse("https://example.test/gateway").unwrap();
let url = build_endpoint_url(&base, "/stable/", "/quote-short").unwrap();
assert_eq!(
url.as_str(),
"https://example.test/gateway/stable/quote-short"
);
}
}