use std::sync::Arc;
use reqwest::header::HeaderValue;
use reqwest::{RequestBuilder, Url};
use serde::Serialize;
use serde::de::DeserializeOwned;
use url::{Host, form_urlencoded};
use crate::error::{Result, TreetopError};
use crate::token::UploadToken;
use crate::types::{
AuthorizeBriefResponse, AuthorizeDetailedResponse, AuthorizeRequest, BatchResult,
DecisionBrief, PoliciesDownload, PoliciesMetadata, Request, SchemaDownload, StatusResponse,
UserPolicies, VersionInfo,
};
use super::builder::ClientBuilder;
const CORRELATION_HEADER: &str = "x-correlation-id";
const ERROR_BODY_LIMIT: usize = 64 * 1024;
#[derive(Debug, Clone)]
pub(super) struct BaseUrl(Url);
impl BaseUrl {
pub(super) fn parse(value: &str) -> Result<Self> {
let mut url = Url::parse(value.trim())?;
if !matches!(url.scheme(), "http" | "https") {
return Err(TreetopError::Configuration(
"base URL scheme must be http or https".to_string(),
));
}
if !url.has_host() || url.cannot_be_a_base() {
return Err(TreetopError::Configuration(
"base URL must include a host".to_string(),
));
}
if !url.username().is_empty() || url.password().is_some() {
return Err(TreetopError::Configuration(
"base URL must not contain credentials".to_string(),
));
}
if url.query().is_some() || url.fragment().is_some() {
return Err(TreetopError::Configuration(
"base URL must not contain a query string or fragment".to_string(),
));
}
let normalized_path = format!("{}/", url.path().trim_end_matches('/'));
url.set_path(&normalized_path);
Ok(Self(url))
}
pub(super) fn validate_upload_transport(&self, allow_insecure: bool) -> Result<()> {
if self.0.scheme() == "https" || allow_insecure || self.is_loopback() {
Ok(())
} else {
Err(TreetopError::Configuration(
"refusing to send an upload token over plaintext HTTP; use HTTPS or explicitly enable danger_allow_insecure_uploads".to_string(),
))
}
}
fn is_loopback(&self) -> bool {
match self.0.host() {
Some(Host::Domain(domain)) => domain.eq_ignore_ascii_case("localhost"),
Some(Host::Ipv4(address)) => address.is_loopback(),
Some(Host::Ipv6(address)) => address.is_loopback(),
None => false,
}
}
fn api_base(&self) -> Url {
self.0
.join("api/v1/")
.expect("a validated HTTP base URL must support relative joins")
}
fn root_endpoint(&self, path: &str) -> Url {
self.0
.join(path)
.expect("static root endpoint paths must be valid relative URLs")
}
fn as_str(&self) -> &str {
self.0.as_str().trim_end_matches('/')
}
}
#[derive(Debug, Clone)]
pub(super) struct CorrelationId(HeaderValue);
impl CorrelationId {
pub(super) fn parse(value: String) -> Result<Self> {
if value.is_empty() {
return Err(TreetopError::Configuration(
"correlation ID must not be empty".to_string(),
));
}
let value = HeaderValue::try_from(value).map_err(|_| {
TreetopError::Configuration(
"correlation ID contains invalid HTTP header characters".to_string(),
)
})?;
Ok(Self(value))
}
fn as_header(&self) -> &HeaderValue {
&self.0
}
}
#[derive(Debug, Clone, Copy)]
pub(super) struct ResponseSizeLimit(usize);
impl ResponseSizeLimit {
pub(super) fn new(value: usize) -> Result<Self> {
if value == 0 {
Err(TreetopError::Configuration(
"maximum response size must be greater than zero".to_string(),
))
} else {
Ok(Self(value))
}
}
fn get(self) -> usize {
self.0
}
}
struct ClientState {
http: reqwest::Client,
base_url: BaseUrl,
api_base: Url,
upload_token: Option<UploadToken>,
max_response_bytes: ResponseSizeLimit,
}
impl ClientState {
fn endpoint(&self, path: &str) -> Url {
self.api_base
.join(path.trim_start_matches('/'))
.expect("static API endpoint paths must be valid relative URLs")
}
}
#[derive(Clone)]
pub struct Client {
state: Arc<ClientState>,
correlation_id: Option<CorrelationId>,
}
impl Client {
pub(super) fn new(
http: reqwest::Client,
base_url: BaseUrl,
upload_token: Option<UploadToken>,
correlation_id: Option<CorrelationId>,
max_response_bytes: ResponseSizeLimit,
) -> Self {
let api_base = base_url.api_base();
Self {
state: Arc::new(ClientState {
http,
base_url,
api_base,
upload_token,
max_response_bytes,
}),
correlation_id,
}
}
pub fn builder(base_url: impl Into<String>) -> ClientBuilder {
ClientBuilder::new(base_url)
}
pub fn with_correlation_id(&self, id: impl Into<String>) -> Result<Client> {
Ok(Client {
state: Arc::clone(&self.state),
correlation_id: Some(CorrelationId::parse(id.into())?),
})
}
pub fn without_correlation_id(&self) -> Client {
Client {
state: Arc::clone(&self.state),
correlation_id: None,
}
}
fn apply_headers(&self, builder: RequestBuilder) -> RequestBuilder {
if let Some(cid) = &self.correlation_id {
builder.header(CORRELATION_HEADER, cid.as_header())
} else {
builder
}
}
async fn read_body(&self, mut response: reqwest::Response, limit: usize) -> Result<Vec<u8>> {
if response
.content_length()
.is_some_and(|content_length| content_length > limit as u64)
{
return Err(TreetopError::ResponseTooLarge { limit });
}
let capacity = response
.content_length()
.and_then(|length| usize::try_from(length).ok())
.unwrap_or_default()
.min(limit);
let mut body = Vec::with_capacity(capacity);
while let Some(chunk) = response.chunk().await.map_err(TreetopError::Transport)? {
if chunk.len() > limit.saturating_sub(body.len()) {
return Err(TreetopError::ResponseTooLarge { limit });
}
body.extend_from_slice(&chunk);
}
Ok(body)
}
async fn api_error(&self, response: reqwest::Response) -> TreetopError {
let status = response.status();
let body = match self.read_body(response, ERROR_BODY_LIMIT).await {
Ok(body) => body,
Err(TreetopError::ResponseTooLarge { .. }) => {
return TreetopError::Api {
status,
message: format!("response error body exceeded {ERROR_BODY_LIMIT} bytes"),
};
}
Err(error) => return error,
};
#[derive(serde::Deserialize)]
struct ErrorEnvelope {
error: String,
}
let message = serde_json::from_slice::<ErrorEnvelope>(&body)
.map(|envelope| envelope.error)
.unwrap_or_else(|_| String::from_utf8_lossy(&body).into_owned());
TreetopError::Api { status, message }
}
async fn handle_response<T: DeserializeOwned>(&self, resp: reqwest::Response) -> Result<T> {
let status = resp.status();
if status.is_success() {
let body = self
.read_body(resp, self.state.max_response_bytes.get())
.await?;
serde_json::from_slice(&body).map_err(TreetopError::Deserialization)
} else {
Err(self.api_error(resp).await)
}
}
async fn handle_text_response(&self, resp: reqwest::Response) -> Result<String> {
let status = resp.status();
if status.is_success() {
let body = self
.read_body(resp, self.state.max_response_bytes.get())
.await?;
Ok(String::from_utf8_lossy(&body).into_owned())
} else {
Err(self.api_error(resp).await)
}
}
async fn get<T: DeserializeOwned>(&self, path: &str) -> Result<T> {
let resp = self
.apply_headers(self.state.http.get(self.state.endpoint(path)))
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
async fn get_text(&self, url: Url) -> Result<String> {
let resp = self
.apply_headers(self.state.http.get(url))
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_text_response(resp).await
}
async fn post_json<T: DeserializeOwned, B: Serialize>(
&self,
path: &str,
body: &B,
) -> Result<T> {
let resp = self
.apply_headers(self.state.http.post(self.state.endpoint(path)).json(body))
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
pub async fn health(&self) -> Result<()> {
let resp = self
.apply_headers(self.state.http.get(self.state.endpoint("health")))
.send()
.await
.map_err(TreetopError::Transport)?;
let status = resp.status();
if status.is_success() {
Ok(())
} else {
Err(self.api_error(resp).await)
}
}
pub async fn version(&self) -> Result<VersionInfo> {
self.get("/version").await
}
pub async fn status(&self) -> Result<StatusResponse> {
self.get("/status").await
}
pub async fn authorize(&self, request: &AuthorizeRequest) -> Result<AuthorizeBriefResponse> {
request.validate()?;
let response: AuthorizeBriefResponse =
self.post_json("authorize?detail=brief", request).await?;
response.validate(request.len())?;
Ok(response)
}
pub async fn authorize_detailed(
&self,
request: &AuthorizeRequest,
) -> Result<AuthorizeDetailedResponse> {
request.validate()?;
let response: AuthorizeDetailedResponse =
self.post_json("authorize?detail=full", request).await?;
response.validate(request.len())?;
Ok(response)
}
pub async fn is_allowed(&self, request: Request) -> Result<bool> {
let batch = AuthorizeRequest::single(request);
let resp = self.authorize(&batch).await?;
let result = resp.results().first().ok_or_else(|| {
TreetopError::InvalidResponse("empty response from authorize endpoint".to_string())
})?;
match &result.result {
BatchResult::Success { data } => Ok(matches!(data.decision, DecisionBrief::Allow)),
BatchResult::Failed { message } => Err(TreetopError::Evaluation(message.clone())),
}
}
pub async fn get_policies(&self) -> Result<PoliciesDownload> {
self.get("/policies").await
}
pub async fn get_policies_raw(&self) -> Result<String> {
self.get_text(self.state.endpoint("policies?format=raw"))
.await
}
pub async fn get_schema(&self) -> Result<SchemaDownload> {
self.get("/schema").await
}
pub async fn get_schema_raw(&self) -> Result<String> {
self.get_text(self.state.endpoint("schema?format=raw"))
.await
}
pub async fn upload_policies_raw(&self, content: &str) -> Result<PoliciesMetadata> {
let token =
self.state.upload_token.as_ref().ok_or_else(|| {
TreetopError::Configuration("no upload token configured".to_string())
})?;
let resp = self
.apply_headers(
self.state
.http
.post(self.state.endpoint("policies"))
.header("Content-Type", "text/plain")
.header("X-Upload-Token", token.header_value())
.body(content.to_string()),
)
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
pub async fn upload_policies_json(&self, content: &str) -> Result<PoliciesMetadata> {
let token =
self.state.upload_token.as_ref().ok_or_else(|| {
TreetopError::Configuration("no upload token configured".to_string())
})?;
#[derive(Serialize)]
struct Upload<'a> {
policies: &'a str,
}
let resp = self
.apply_headers(
self.state
.http
.post(self.state.endpoint("policies"))
.header("X-Upload-Token", token.header_value())
.json(&Upload { policies: content }),
)
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
pub async fn upload_schema_raw(&self, content: &str) -> Result<PoliciesMetadata> {
let token =
self.state.upload_token.as_ref().ok_or_else(|| {
TreetopError::Configuration("no upload token configured".to_string())
})?;
let resp = self
.apply_headers(
self.state
.http
.post(self.state.endpoint("schema"))
.header("Content-Type", "text/plain")
.header("X-Upload-Token", token.header_value())
.body(content.to_string()),
)
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
pub async fn upload_schema_json(&self, content: &str) -> Result<PoliciesMetadata> {
let token =
self.state.upload_token.as_ref().ok_or_else(|| {
TreetopError::Configuration("no upload token configured".to_string())
})?;
#[derive(Serialize)]
struct Upload<'a> {
schema: &'a str,
}
let resp = self
.apply_headers(
self.state
.http
.post(self.state.endpoint("schema"))
.header("X-Upload-Token", token.header_value())
.json(&Upload { schema: content }),
)
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
pub async fn get_user_policies(
&self,
user: &str,
groups: &[String],
namespaces: &[String],
) -> Result<UserPolicies> {
let url = self.build_user_policies_url(user, groups, namespaces, false);
let resp = self
.apply_headers(self.state.http.get(url))
.send()
.await
.map_err(TreetopError::Transport)?;
self.handle_response(resp).await
}
pub async fn get_user_policies_raw(
&self,
user: &str,
groups: &[String],
namespaces: &[String],
) -> Result<String> {
let url = self.build_user_policies_url(user, groups, namespaces, true);
self.get_text(url).await
}
pub async fn metrics(&self) -> Result<String> {
self.get_text(self.state.base_url.root_endpoint("metrics"))
.await
}
fn build_user_policies_url(
&self,
user: &str,
groups: &[String],
namespaces: &[String],
raw: bool,
) -> Url {
let encoded_user: String = form_urlencoded::byte_serialize(user.as_bytes()).collect();
let mut url = self
.state
.endpoint("policies/")
.join(&encoded_user)
.expect("a form-encoded path segment must be a valid relative URL");
let mut query = url.query_pairs_mut();
for ns in namespaces {
query.append_pair("namespaces[]", ns);
}
for group in groups {
query.append_pair("groups[]", group);
}
if raw {
query.append_pair("format", "raw");
}
drop(query);
url
}
}
impl std::fmt::Debug for Client {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Client")
.field("base_url", &self.state.base_url.as_str())
.field(
"upload_token",
&self.state.upload_token.as_ref().map(|_| "[SET]"),
)
.field("correlation_id", &self.correlation_id)
.finish()
}
}