horizon-sdk 16.0.0

Canonical Rust data access layer for the Horizon platform
//! HTTP client wrapper for the inference service's coverage endpoints.

use std::time::Duration;

use reqwest::Client;
use reqwest::header::{AUTHORIZATION, HeaderMap, HeaderValue};
use serde_json::Value;
use tracing::{debug, error};

use crate::types::config::CoverageConfig;
use crate::types::error::{ConfigError, HttpError, Result};

/// Async HTTP client for the coverage endpoints of the inference service.
///
/// Provides typed wrappers for the three coverage endpoints:
/// course-of-action generation, coverage heatmaps, and predictive heatmaps.
#[derive(Debug, Clone)]
pub struct CoverageClient {
    /// Base URL of the inference service (without trailing slash).
    base_url: String,
    /// HTTP client configured with timeout.
    client: Client,
}

impl CoverageClient {
    /// The base URL this client connects to.
    #[must_use]
    pub fn base_url(&self) -> &str {
        &self.base_url
    }

    /// Generate course of action candidates for the provided objective.
    ///
    /// # Errors
    ///
    /// Returns `HttpError::ServiceUnreachable` / `HttpError::Timeout` /
    /// `HttpError::RequestFailed` on transport failures, or
    /// `HttpError::UnexpectedStatus` / `HttpError::DecodeFailed` on a
    /// successful-but-unusable response.
    pub async fn get_coa(&self, request: &Value) -> Result<Value> {
        self.post_coverage("/coverage/generate_course_of_actions", request)
            .await
    }

    /// Get coverage for the requested area.
    ///
    /// # Errors
    ///
    /// Returns `HttpError::ServiceUnreachable` / `HttpError::Timeout` /
    /// `HttpError::RequestFailed` on transport failures, or
    /// `HttpError::UnexpectedStatus` / `HttpError::DecodeFailed` on a
    /// successful-but-unusable response.
    pub async fn get_coverage_heatmap(&self, request: &Value) -> Result<Value> {
        self.post_coverage("/coverage/get_coverage", request).await
    }

    /// Get a predictive heatmap for the requested area.
    ///
    /// # Errors
    ///
    /// Returns `HttpError::ServiceUnreachable` / `HttpError::Timeout` /
    /// `HttpError::RequestFailed` on transport failures, or
    /// `HttpError::UnexpectedStatus` / `HttpError::DecodeFailed` on a
    /// successful-but-unusable response.
    pub async fn get_predictive_heatmap(&self, request: &Value) -> Result<Value> {
        self.post_coverage("/coverage/get_predictive_heatmap", request)
            .await
    }

    /// Create a new coverage client from SDK configuration.
    ///
    /// # Errors
    ///
    /// Returns `ConfigError::InvalidUrl` if the base URL is malformed, or
    /// `HttpError::ClientBuild` if the HTTP client cannot be created (an API key
    /// that is not a valid header value counts as a build failure).
    pub fn new(config: &CoverageConfig) -> Result<Self> {
        let base_url = config.base_url.trim_end_matches('/').to_owned();
        url::Url::parse(&base_url).map_err(|source| ConfigError::InvalidUrl {
            field: "coverage base",
            url: base_url.clone(),
            source,
        })?;

        let mut default_headers = HeaderMap::new();
        if let Some(api_key) = &config.api_key {
            let mut authorization = HeaderValue::from_str(&format!("Bearer {api_key}"))
                .map_err(|_invalid_header_value| HttpError::InvalidApiKey)?;
            authorization.set_sensitive(true);
            default_headers.insert(AUTHORIZATION, authorization);
        }

        let client = Client::builder()
            .default_headers(default_headers)
            .timeout(Duration::from_secs(config.timeout_secs))
            .build()
            .map_err(|source| HttpError::ClientBuild { source })?;

        debug!(base_url = base_url.as_str(), "initialized CoverageClient");

        Ok(Self { base_url, client })
    }

    /// Send a POST request to a coverage service endpoint; the routes take the model under a `request` key.
    async fn post_coverage(&self, endpoint: &str, payload: &Value) -> Result<Value> {
        let url = format!("{}{endpoint}", self.base_url);

        let response = self
            .client
            .post(&url)
            .json(&serde_json::json!({ "request": payload }))
            .send()
            .await
            .map_err(|source| {
                HttpError::from_send_error("horizon-inference", &self.base_url, endpoint, source)
            })?;

        let status = response.status();
        if !status.is_success() {
            error!(
                status = status.as_u16(),
                endpoint, "coverage service returned error"
            );
            return Err(HttpError::UnexpectedStatus {
                service: "horizon-inference",
                status: status.as_u16(),
                endpoint: endpoint.to_owned(),
            }
            .into());
        }

        response.json().await.map_err(|source| {
            error!(endpoint, error = %source, "failed to decode coverage response");
            HttpError::DecodeFailed {
                endpoint: endpoint.to_owned(),
                source,
            }
            .into()
        })
    }
}

#[cfg(test)]
#[allow(clippy::unwrap_used, reason = "unit tests use unwrap for assertions")]
mod tests {
    use super::*;
    use crate::types::error::HorizonError;

    #[test]
    fn coverage_client_preserves_clean_url() {
        let config = CoverageConfig {
            api_key: None,
            base_url: "https://example.com".to_owned(),
            timeout_secs: 10,
        };
        let client = CoverageClient::new(&config).unwrap();
        assert_eq!(client.base_url(), "https://example.com");
    }

    #[test]
    fn coverage_client_trims_trailing_slash() {
        let config = CoverageConfig {
            api_key: None,
            base_url: "http://localhost:3000/".to_owned(),
            timeout_secs: 30,
        };
        let client = CoverageClient::new(&config).unwrap();
        assert_eq!(client.base_url(), "http://localhost:3000");
    }

    #[test]
    fn debug_impl_shows_base_url() {
        let config = CoverageConfig {
            api_key: Some("secret-token".to_owned()),
            base_url: "http://localhost:3000".to_owned(),
            timeout_secs: 30,
        };
        let client = CoverageClient::new(&config).unwrap();
        let debug = format!("{client:?}");
        assert!(debug.contains("http://localhost:3000"));
        assert!(!debug.contains("secret-token"));
    }

    #[test]
    fn coverage_client_rejects_api_key_with_invalid_header_characters() {
        let config = CoverageConfig {
            api_key: Some("bad\nkey".to_owned()),
            base_url: "http://localhost:3000".to_owned(),
            timeout_secs: 30,
        };
        let error = CoverageClient::new(&config).unwrap_err();
        assert!(matches!(
            error,
            HorizonError::Http(HttpError::InvalidApiKey)
        ));
    }
}