use std::{collections::BTreeMap, time::Duration};
use reqwest::header::{AUTHORIZATION, CONTENT_TYPE};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::{HttpMethod, HttpRequest, MaxTokensParameter, RetryPolicy};
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum AuthConfig {
Bearer {
token: String,
},
Header {
name: String,
value: String,
},
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct HttpModelConfig {
pub base_url: String,
pub endpoint_path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub auth: Option<AuthConfig>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub headers: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub extra_body: Map<String, Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timeout_ms: Option<u64>,
#[serde(default)]
pub retry_policy: RetryPolicy,
#[serde(default)]
pub max_tokens_parameter: MaxTokensParameter,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub metadata: Map<String, Value>,
}
impl HttpModelConfig {
#[must_use]
pub fn new(base_url: impl Into<String>, endpoint_path: impl Into<String>) -> Self {
Self {
base_url: base_url.into(),
endpoint_path: endpoint_path.into(),
auth: None,
headers: BTreeMap::new(),
extra_body: Map::new(),
timeout_ms: None,
retry_policy: RetryPolicy::default(),
max_tokens_parameter: MaxTokensParameter::Default,
metadata: Map::new(),
}
}
#[must_use]
pub fn endpoint_url(&self) -> String {
let base = self.base_url.trim_end_matches('/');
let path = self.endpoint_path.trim_start_matches('/');
format!("{base}/{path}")
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
pub struct HttpRequestOptions {
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub headers: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub extra_body: Map<String, Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub endpoint_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timeout_ms: Option<u64>,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub metadata: Map<String, Value>,
}
#[must_use]
pub fn merge_extra_body(mut body: Value, extra: &Map<String, Value>) -> Value {
if let Value::Object(object) = &mut body {
for (key, value) in extra {
object.insert(key.clone(), value.clone());
}
}
body
}
pub fn extend_headers_case_insensitive(
headers: &mut BTreeMap<String, String>,
overlay: impl IntoIterator<Item = (String, String)>,
) {
for (key, value) in overlay {
headers.retain(|existing, _| !existing.eq_ignore_ascii_case(&key));
headers.insert(key, value);
}
}
fn merge_metadata(config: &HttpModelConfig, options: &HttpRequestOptions) -> Map<String, Value> {
let mut metadata = config.metadata.clone();
metadata.extend(options.metadata.clone());
metadata
}
#[must_use]
pub fn build_http_request(
config: &HttpModelConfig,
options: &HttpRequestOptions,
body: Value,
) -> HttpRequest {
let mut headers = BTreeMap::from([(
CONTENT_TYPE.as_str().to_string(),
"application/json".to_string(),
)]);
match &config.auth {
Some(AuthConfig::Bearer { token }) => {
extend_headers_case_insensitive(
&mut headers,
[(
AUTHORIZATION.as_str().to_string(),
format!("Bearer {token}"),
)],
);
}
Some(AuthConfig::Header { name, value }) => {
extend_headers_case_insensitive(&mut headers, [(name.clone(), value.clone())]);
}
None => {}
}
extend_headers_case_insensitive(&mut headers, config.headers.clone());
extend_headers_case_insensitive(&mut headers, options.headers.clone());
let body = merge_extra_body(
merge_extra_body(body, &config.extra_body),
&options.extra_body,
);
let timeout_ms = options.timeout_ms.or(config.timeout_ms);
HttpRequest {
method: HttpMethod::Post,
url: options
.endpoint_url
.clone()
.unwrap_or_else(|| config.endpoint_url()),
headers,
body,
timeout: timeout_ms.map(Duration::from_millis),
metadata: merge_metadata(config, options),
cancellation_token: starweaver_core::CancellationToken::default(),
}
}