use std::collections::HashMap;
use serde::{Deserialize, Serialize};
const REDACTED: &str = "[redacted]";
const TOP_P_MIN: f64 = 0.0;
const TOP_P_MAX: f64 = 1.0;
const PENALTY_MIN: f64 = -2.0;
const PENALTY_MAX: f64 = 2.0;
#[derive(Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
pub struct LlmConfig {
pub model: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub api_key: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timeout_secs: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_retries: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub top_p: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub stop: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub seed: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub presence_penalty: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub frequency_penalty: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub reasoning_effort: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub extra_body: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub load_env: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub headers: Option<HashMap<String, String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub providers: Option<Vec<LlmProviderConfig>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub cache: Option<Box<LlmCacheConfig>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub budget: Option<Box<LlmBudgetConfig>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub rate_limit: Option<Box<LlmRateLimitConfig>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub cost_tracking: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub tracing: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub cooldown_secs: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub health_check_secs: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub bedrock: Option<Box<BedrockConfig>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub credential_provider: Option<Box<CredentialProviderConfig>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub max_concurrency: Option<usize>,
}
impl LlmConfig {
pub fn validate(&self) -> crate::Result<()> {
self.validate_sampling_parameters()?;
#[cfg(target_arch = "wasm32")]
validate_wasm_credential_provider(self.credential_provider.as_deref())?;
Ok(())
}
fn validate_sampling_parameters(&self) -> crate::Result<()> {
if let Some(top_p) = self.top_p {
validate_sampling_range("top_p", top_p, TOP_P_MIN, TOP_P_MAX)?;
}
if let Some(presence_penalty) = self.presence_penalty {
validate_sampling_range("presence_penalty", presence_penalty, PENALTY_MIN, PENALTY_MAX)?;
}
if let Some(frequency_penalty) = self.frequency_penalty {
validate_sampling_range("frequency_penalty", frequency_penalty, PENALTY_MIN, PENALTY_MAX)?;
}
Ok(())
}
#[cfg(test)]
pub(crate) fn validate_for_wasm_target(&self) -> crate::Result<()> {
self.validate_sampling_parameters()?;
validate_wasm_credential_provider(self.credential_provider.as_deref())
}
}
#[cfg(any(target_arch = "wasm32", test))]
const WASM_CREDENTIAL_PROVIDER_ERROR: &str =
"credential_provider is not supported on wasm32 targets; use api_key or browser-compatible authentication";
#[cfg(any(target_arch = "wasm32", test))]
fn validate_wasm_credential_provider(provider: Option<&CredentialProviderConfig>) -> crate::Result<()> {
if provider.is_some() {
return Err(crate::XbergError::validation(
WASM_CREDENTIAL_PROVIDER_ERROR.to_string(),
));
}
Ok(())
}
fn validate_sampling_range(field_name: &str, value: f64, min: f64, max: f64) -> crate::Result<()> {
if (min..=max).contains(&value) {
Ok(())
} else {
Err(crate::XbergError::Validation {
message: format!("Invalid LLM {field_name} {value}: expected a value between {min} and {max}"),
source: None,
})
}
}
#[derive(Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub enum CredentialProviderConfig {
AzureAd {
tenant_id: String,
client_id: String,
client_secret: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
scope: Option<String>,
},
VertexOauth2 {
service_account_key_file: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
scope: Option<String>,
},
VertexAdc {
#[serde(default, skip_serializing_if = "Option::is_none")]
scope: Option<String>,
},
BedrockWebIdentity {
role_arn: String,
token_file: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
session_name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
region: Option<String>,
},
}
impl std::fmt::Debug for CredentialProviderConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::AzureAd {
tenant_id,
client_id,
scope,
client_secret: _,
} => f
.debug_struct("AzureAd")
.field("tenant_id", tenant_id)
.field("client_id", client_id)
.field("client_secret", &REDACTED)
.field("scope", scope)
.finish(),
Self::VertexOauth2 {
service_account_key_file,
scope,
} => f
.debug_struct("VertexOauth2")
.field("service_account_key_file", service_account_key_file)
.field("scope", scope)
.finish(),
Self::VertexAdc { scope } => f.debug_struct("VertexAdc").field("scope", scope).finish(),
Self::BedrockWebIdentity {
role_arn,
token_file,
session_name,
region,
} => f
.debug_struct("BedrockWebIdentity")
.field("role_arn", role_arn)
.field("token_file", token_file)
.field("session_name", session_name)
.field("region", region)
.finish(),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub struct LlmProviderConfig {
pub name: String,
pub base_url: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub auth_header: Option<String>,
#[serde(default)]
pub model_prefixes: Vec<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub struct LlmCacheConfig {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_entries: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ttl_seconds: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend_config: Option<HashMap<String, String>>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub struct LlmBudgetConfig {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub global_limit: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub model_limits: Option<HashMap<String, f64>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub enforcement: Option<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub struct LlmRateLimitConfig {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub rpm: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tpm: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub window_seconds: Option<u64>,
}
#[derive(Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[cfg_attr(feature = "alef-meta", alef(since = "1.1.0"))]
pub struct BedrockConfig {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub region: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cross_region_prefix: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub access_key_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub secret_access_key: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub session_token: Option<String>,
}
impl std::fmt::Debug for LlmConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let redacted_headers: Option<Vec<(&str, &str)>> = self
.headers
.as_ref()
.map(|headers| headers.keys().map(|name| (name.as_str(), REDACTED)).collect());
f.debug_struct("LlmConfig")
.field("model", &self.model)
.field("api_key", &self.api_key.as_ref().map(|_| REDACTED))
.field("base_url", &self.base_url)
.field("timeout_secs", &self.timeout_secs)
.field("max_retries", &self.max_retries)
.field("temperature", &self.temperature)
.field("max_tokens", &self.max_tokens)
.field("top_p", &self.top_p)
.field("stop", &self.stop)
.field("seed", &self.seed)
.field("presence_penalty", &self.presence_penalty)
.field("frequency_penalty", &self.frequency_penalty)
.field("reasoning_effort", &self.reasoning_effort)
.field("extra_body", &self.extra_body)
.field("load_env", &self.load_env)
.field("headers", &redacted_headers)
.field("providers", &self.providers)
.field("cache", &self.cache)
.field("budget", &self.budget)
.field("rate_limit", &self.rate_limit)
.field("cost_tracking", &self.cost_tracking)
.field("tracing", &self.tracing)
.field("cooldown_secs", &self.cooldown_secs)
.field("health_check_secs", &self.health_check_secs)
.field("bedrock", &self.bedrock)
.field("credential_provider", &self.credential_provider)
.field("max_concurrency", &self.max_concurrency)
.finish()
}
}
impl std::fmt::Debug for BedrockConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BedrockConfig")
.field("region", &self.region)
.field("cross_region_prefix", &self.cross_region_prefix)
.field("access_key_id", &self.access_key_id.as_ref().map(|_| REDACTED))
.field("secret_access_key", &self.secret_access_key.as_ref().map(|_| REDACTED))
.field("session_token", &self.session_token.as_ref().map(|_| REDACTED))
.finish()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct StructuredExtractionConfig {
pub schema: serde_json::Value,
#[serde(default = "StructuredExtractionConfig::default_schema_name")]
pub schema_name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub schema_description: Option<String>,
#[serde(default)]
pub strict: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt: Option<String>,
pub llm: LlmConfig,
}
impl StructuredExtractionConfig {
pub fn default_schema_name() -> String {
"extraction".to_string()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[serde(rename_all = "snake_case")]
pub enum CallMode {
#[default]
TextOnly,
VisionOnly,
TextPlusVision,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
#[serde(rename_all = "snake_case")]
pub enum MergeMode {
#[default]
ObjectMerge,
ArrayConcat,
ObjectFirst,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_reject_credential_provider_for_wasm_validation() {
let config = LlmConfig {
credential_provider: Some(Box::new(CredentialProviderConfig::VertexAdc { scope: None })),
..Default::default()
};
let error = validate_wasm_credential_provider(config.credential_provider.as_deref())
.expect_err("managed credential providers must not be silently ignored on wasm32");
match error {
crate::XbergError::Validation { message, .. } => {
assert_eq!(message, WASM_CREDENTIAL_PROVIDER_ERROR)
}
other => panic!("expected validation error, got {other:?}"),
}
}
#[test]
fn test_llm_config_default_trait_is_satisfied() {
let cfg = LlmConfig::default();
assert!(cfg.model.is_empty(), "default model should be empty string");
assert!(cfg.api_key.is_none());
assert!(cfg.base_url.is_none());
assert!(cfg.timeout_secs.is_none());
assert!(cfg.max_retries.is_none());
assert!(cfg.max_concurrency.is_none());
assert!(cfg.temperature.is_none());
assert!(cfg.max_tokens.is_none());
assert!(cfg.top_p.is_none());
assert!(cfg.stop.is_none());
assert!(cfg.seed.is_none());
assert!(cfg.presence_penalty.is_none());
assert!(cfg.frequency_penalty.is_none());
assert!(cfg.reasoning_effort.is_none());
assert!(cfg.extra_body.is_none());
assert!(cfg.load_env.is_none());
assert!(cfg.headers.is_none());
assert!(cfg.providers.is_none());
assert!(cfg.cache.is_none());
assert!(cfg.budget.is_none());
assert!(cfg.rate_limit.is_none());
assert!(cfg.cost_tracking.is_none());
assert!(cfg.tracing.is_none());
assert!(cfg.cooldown_secs.is_none());
assert!(cfg.health_check_secs.is_none());
assert!(cfg.bedrock.is_none());
assert!(cfg.credential_provider.is_none());
}
#[test]
fn test_llm_config_struct_update_syntax() {
let cfg = LlmConfig {
model: "openai/gpt-4o-mini".to_string(),
..Default::default()
};
assert_eq!(cfg.model, "openai/gpt-4o-mini");
assert!(cfg.api_key.is_none());
assert!(cfg.base_url.is_none());
assert!(cfg.timeout_secs.is_none());
assert!(cfg.max_retries.is_none());
assert!(cfg.max_concurrency.is_none());
assert!(cfg.temperature.is_none());
assert!(cfg.max_tokens.is_none());
assert!(cfg.top_p.is_none());
assert!(cfg.stop.is_none());
assert!(cfg.seed.is_none());
assert!(cfg.presence_penalty.is_none());
assert!(cfg.frequency_penalty.is_none());
assert!(cfg.reasoning_effort.is_none());
assert!(cfg.extra_body.is_none());
assert!(cfg.load_env.is_none());
assert!(cfg.headers.is_none());
assert!(cfg.providers.is_none());
assert!(cfg.cache.is_none());
assert!(cfg.budget.is_none());
assert!(cfg.rate_limit.is_none());
assert!(cfg.cost_tracking.is_none());
assert!(cfg.tracing.is_none());
assert!(cfg.cooldown_secs.is_none());
assert!(cfg.health_check_secs.is_none());
assert!(cfg.bedrock.is_none());
assert!(cfg.credential_provider.is_none());
}
#[test]
fn test_llm_config_load_env_and_headers_round_trip() {
let toml_src = r#"
model = "openai/gpt-4o"
load_env = true
[headers]
"X-Gateway-Key" = "abc123"
"X-Tenant" = "acme"
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
assert_eq!(cfg.model, "openai/gpt-4o");
assert_eq!(cfg.load_env, Some(true));
let headers = cfg.headers.as_ref().expect("headers present");
assert_eq!(headers.get("X-Gateway-Key").map(String::as_str), Some("abc123"));
assert_eq!(headers.get("X-Tenant").map(String::as_str), Some("acme"));
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_llm_config_max_concurrency_round_trip() {
let cfg: LlmConfig = toml::from_str(
r#"
model = "openai/gpt-4o"
max_concurrency = 3
"#,
)
.expect("deserialize LlmConfig from TOML");
assert_eq!(cfg.max_concurrency, Some(3));
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_llm_config_omits_empty_passthrough_fields() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
..Default::default()
};
let json = serde_json::to_string(&cfg).expect("serialize");
assert!(!json.contains("top_p"), "top_p should be omitted when None: {json}");
assert!(!json.contains("\"stop\""), "stop should be omitted when None: {json}");
assert!(!json.contains("\"seed\""), "seed should be omitted when None: {json}");
assert!(
!json.contains("presence_penalty"),
"presence_penalty should be omitted when None: {json}"
);
assert!(
!json.contains("frequency_penalty"),
"frequency_penalty should be omitted when None: {json}"
);
assert!(
!json.contains("reasoning_effort"),
"reasoning_effort should be omitted when None: {json}"
);
assert!(
!json.contains("extra_body"),
"extra_body should be omitted when None: {json}"
);
assert!(
!json.contains("load_env"),
"load_env should be omitted when None: {json}"
);
assert!(!json.contains("headers"), "headers should be omitted when None: {json}");
assert!(
!json.contains("providers"),
"providers should be omitted when None: {json}"
);
assert!(!json.contains("cache"), "cache should be omitted when None: {json}");
assert!(!json.contains("budget"), "budget should be omitted when None: {json}");
assert!(
!json.contains("rate_limit"),
"rate_limit should be omitted when None: {json}"
);
assert!(
!json.contains("cost_tracking"),
"cost_tracking should be omitted when None: {json}"
);
assert!(!json.contains("tracing"), "tracing should be omitted when None: {json}");
assert!(
!json.contains("cooldown_secs"),
"cooldown_secs should be omitted when None: {json}"
);
assert!(
!json.contains("health_check_secs"),
"health_check_secs should be omitted when None: {json}"
);
assert!(!json.contains("bedrock"), "bedrock should be omitted when None: {json}");
assert!(
!json.contains("credential_provider"),
"credential_provider should be omitted when None: {json}"
);
}
#[test]
fn test_llm_config_full_passthrough_fields_round_trip_through_toml_and_json() {
let toml_src = r#"
model = "openai/gpt-4o"
cost_tracking = true
tracing = false
cooldown_secs = 30
health_check_secs = 60
[[providers]]
name = "my-provider"
base_url = "https://my-llm.example.com/v1"
auth_header = "X-Api-Key"
model_prefixes = ["my-provider/"]
[cache]
max_entries = 512
ttl_seconds = 600
backend = "memory"
[budget]
global_limit = 100.0
enforcement = "hard"
[budget.model_limits]
"openai/gpt-4o" = 25.0
[rate_limit]
rpm = 60
tpm = 100000
window_seconds = 60
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
assert_eq!(cfg.cost_tracking, Some(true));
assert_eq!(cfg.tracing, Some(false));
assert_eq!(cfg.cooldown_secs, Some(30));
assert_eq!(cfg.health_check_secs, Some(60));
let providers = cfg.providers.as_ref().expect("providers present");
assert_eq!(providers.len(), 1);
assert_eq!(providers[0].name, "my-provider");
assert_eq!(providers[0].base_url, "https://my-llm.example.com/v1");
assert_eq!(providers[0].auth_header.as_deref(), Some("X-Api-Key"));
assert_eq!(providers[0].model_prefixes, vec!["my-provider/".to_string()]);
let cache = cfg.cache.as_ref().expect("cache present");
assert_eq!(cache.max_entries, Some(512));
assert_eq!(cache.ttl_seconds, Some(600));
assert_eq!(cache.backend.as_deref(), Some("memory"));
assert_eq!(cache.backend_config, None);
let budget = cfg.budget.as_ref().expect("budget present");
assert_eq!(budget.global_limit, Some(100.0));
assert_eq!(budget.enforcement.as_deref(), Some("hard"));
assert_eq!(
budget.model_limits.as_ref().and_then(|m| m.get("openai/gpt-4o")),
Some(&25.0)
);
let rate_limit = cfg.rate_limit.as_ref().expect("rate_limit present");
assert_eq!(rate_limit.rpm, Some(60));
assert_eq!(rate_limit.tpm, Some(100_000));
assert_eq!(rate_limit.window_seconds, Some(60));
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_llm_config_reasoning_effort_and_extra_body_round_trip_through_toml_and_json() {
let toml_src = r#"
model = "openai/gpt-4o"
reasoning_effort = "high"
[extra_body]
safety_settings = { harassment = "block_none" }
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
assert_eq!(cfg.reasoning_effort.as_deref(), Some("high"));
assert_eq!(
cfg.extra_body,
Some(serde_json::json!({"safety_settings": {"harassment": "block_none"}}))
);
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_llm_config_sampling_fields_round_trip_through_toml_and_json() {
let toml_src = r#"
model = "openai/gpt-4o"
top_p = 0.9
stop = ["\n\n", "[END]"]
seed = 42
presence_penalty = 0.5
frequency_penalty = -0.5
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
assert_eq!(cfg.top_p, Some(0.9));
assert_eq!(cfg.stop, Some(vec!["\n\n".to_string(), "[END]".to_string()]));
assert_eq!(cfg.seed, Some(42));
assert_eq!(cfg.presence_penalty, Some(0.5));
assert_eq!(cfg.frequency_penalty, Some(-0.5));
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_llm_config_omits_empty_sampling_fields() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
..Default::default()
};
let json = serde_json::to_string(&cfg).expect("serialize");
assert_eq!(json, r#"{"model":"openai/gpt-4o"}"#);
}
#[test]
fn test_llm_config_debug_prints_sampling_fields_verbatim() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
top_p: Some(0.9),
stop: Some(vec!["\n\n".to_string()]),
seed: Some(42),
presence_penalty: Some(0.5),
frequency_penalty: Some(-0.5),
..Default::default()
};
let rendered = format!("{cfg:?}");
assert!(rendered.contains("top_p: Some(0.9)"), "{rendered}");
assert!(rendered.contains(r#"stop: Some(["\n\n"])"#), "{rendered}");
assert!(rendered.contains("seed: Some(42)"), "{rendered}");
assert!(rendered.contains("presence_penalty: Some(0.5)"), "{rendered}");
assert!(rendered.contains("frequency_penalty: Some(-0.5)"), "{rendered}");
}
#[test]
fn test_llm_config_validate_accepts_top_p_at_boundaries() {
for value in [0.0, 1.0, 0.5] {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
top_p: Some(value),
..Default::default()
};
assert!(cfg.validate().is_ok(), "top_p {value} should be accepted");
}
}
#[test]
fn test_llm_config_validate_rejects_top_p_out_of_range() {
for value in [-0.1, 1.1] {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
top_p: Some(value),
..Default::default()
};
match cfg.validate() {
Err(crate::XbergError::Validation { message, .. }) => {
assert!(message.contains("top_p"), "{message}");
assert!(message.contains(&value.to_string()), "{message}");
}
other => panic!("expected a Validation error for top_p {value}, got {other:?}"),
}
}
}
#[test]
fn test_llm_config_validate_accepts_penalties_at_boundaries() {
for value in [-2.0, 2.0, 0.0] {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
presence_penalty: Some(value),
frequency_penalty: Some(value),
..Default::default()
};
assert!(cfg.validate().is_ok(), "penalty {value} should be accepted");
}
}
#[test]
fn test_llm_config_validate_rejects_presence_penalty_out_of_range() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
presence_penalty: Some(2.1),
..Default::default()
};
match cfg.validate() {
Err(crate::XbergError::Validation { message, .. }) => {
assert!(message.contains("presence_penalty"), "{message}");
assert!(message.contains("2.1"), "{message}");
}
other => panic!("expected a Validation error, got {other:?}"),
}
}
#[test]
fn test_llm_config_validate_rejects_frequency_penalty_out_of_range() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
frequency_penalty: Some(-2.1),
..Default::default()
};
match cfg.validate() {
Err(crate::XbergError::Validation { message, .. }) => {
assert!(message.contains("frequency_penalty"), "{message}");
assert!(message.contains("-2.1"), "{message}");
}
other => panic!("expected a Validation error, got {other:?}"),
}
}
#[test]
fn test_llm_config_validate_accepts_all_unset_sampling_fields() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
..Default::default()
};
assert!(cfg.validate().is_ok());
}
#[test]
fn test_llm_config_validate_accepts_any_seed_value() {
for value in [i64::MIN, -1, 0, 1, i64::MAX] {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
seed: Some(value),
..Default::default()
};
assert!(cfg.validate().is_ok(), "seed {value} should be accepted");
}
}
#[test]
fn test_llm_config_debug_prints_reasoning_effort_and_extra_body_verbatim() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
reasoning_effort: Some("high".to_string()),
extra_body: Some(serde_json::json!({"foo": "bar"})),
..Default::default()
};
let rendered = format!("{cfg:?}");
assert!(rendered.contains(r#"reasoning_effort: Some("high")"#), "{rendered}");
assert!(rendered.contains(r#"extra_body: Some(Object"#), "{rendered}");
assert!(rendered.contains("bar"), "{rendered}");
}
#[test]
fn test_llm_config_debug_prints_new_passthrough_fields_verbatim() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
cost_tracking: Some(true),
tracing: Some(true),
cooldown_secs: Some(15),
health_check_secs: Some(45),
rate_limit: Some(Box::new(LlmRateLimitConfig {
rpm: Some(30),
tpm: None,
window_seconds: None,
})),
..Default::default()
};
let rendered = format!("{cfg:?}");
assert!(rendered.contains("cost_tracking: Some(true)"), "{rendered}");
assert!(rendered.contains("tracing: Some(true)"), "{rendered}");
assert!(rendered.contains("cooldown_secs: Some(15)"), "{rendered}");
assert!(rendered.contains("health_check_secs: Some(45)"), "{rendered}");
assert!(rendered.contains("rpm: Some(30)"), "{rendered}");
}
#[test]
fn test_llm_config_bedrock_round_trips_through_toml_and_json() {
let toml_src = r#"
model = "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
[bedrock]
region = "eu-central-1"
cross_region_prefix = "eu"
access_key_id = "AKIAEXAMPLE"
secret_access_key = "example-secret"
session_token = "example-token"
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
assert_eq!(cfg.model, "bedrock/anthropic.claude-3-sonnet-20240229-v1:0");
let bedrock = cfg.bedrock.as_ref().expect("bedrock present");
assert_eq!(bedrock.region.as_deref(), Some("eu-central-1"));
assert_eq!(bedrock.cross_region_prefix.as_deref(), Some("eu"));
assert_eq!(bedrock.access_key_id.as_deref(), Some("AKIAEXAMPLE"));
assert_eq!(bedrock.secret_access_key.as_deref(), Some("example-secret"));
assert_eq!(bedrock.session_token.as_deref(), Some("example-token"));
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_llm_config_bedrock_region_only_leaves_credentials_unset() {
let cfg: LlmConfig = toml::from_str(
r#"
model = "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
[bedrock]
region = "us-east-1"
"#,
)
.expect("deserialize LlmConfig from TOML");
let bedrock = cfg.bedrock.as_ref().expect("bedrock present");
assert_eq!(bedrock.region.as_deref(), Some("us-east-1"));
assert_eq!(bedrock.cross_region_prefix, None);
assert_eq!(bedrock.access_key_id, None);
assert_eq!(bedrock.secret_access_key, None);
assert_eq!(bedrock.session_token, None);
}
#[test]
fn test_llm_config_debug_redacts_every_credential() {
let mut headers = HashMap::new();
headers.insert("X-Gateway-Key".to_string(), "header-secret".to_string());
let cfg = LlmConfig {
model: "bedrock/anthropic.claude-3-sonnet-20240229-v1:0".to_string(),
api_key: Some("sk-super-secret".to_string()),
headers: Some(headers),
bedrock: Some(Box::new(BedrockConfig {
region: Some("eu-central-1".to_string()),
cross_region_prefix: Some("eu".to_string()),
access_key_id: Some("AKIAEXAMPLE".to_string()),
secret_access_key: Some("aws-secret".to_string()),
session_token: Some("aws-token".to_string()),
})),
..Default::default()
};
let rendered = format!("{cfg:?}");
for secret in [
"sk-super-secret",
"header-secret",
"AKIAEXAMPLE",
"aws-secret",
"aws-token",
] {
assert!(!rendered.contains(secret), "Debug leaked {secret}: {rendered}");
}
assert!(
rendered.contains("eu-central-1"),
"region should be printed: {rendered}"
);
assert!(rendered.contains("\"eu\""), "prefix should be printed: {rendered}");
assert!(
rendered.contains("X-Gateway-Key"),
"header name should be printed: {rendered}"
);
assert_eq!(
rendered.matches(REDACTED).count(),
5,
"expected 5 redactions: {rendered}"
);
}
#[test]
fn test_bedrock_config_debug_distinguishes_unset_from_redacted() {
let bedrock = BedrockConfig {
region: Some("us-east-1".to_string()),
..Default::default()
};
let rendered = format!("{bedrock:?}");
assert_eq!(rendered.matches(REDACTED).count(), 0, "nothing to redact: {rendered}");
assert_eq!(rendered.matches("None").count(), 4, "4 unset fields: {rendered}");
}
#[test]
fn test_call_mode_serde_round_trip() {
for (mode, wire) in [
(CallMode::TextOnly, "\"text_only\""),
(CallMode::VisionOnly, "\"vision_only\""),
(CallMode::TextPlusVision, "\"text_plus_vision\""),
] {
let json = serde_json::to_string(&mode).expect("serialize");
assert_eq!(json, wire);
let decoded: CallMode = serde_json::from_str(&json).expect("deserialize");
assert_eq!(decoded, mode);
}
assert_eq!(CallMode::default(), CallMode::TextOnly);
}
#[test]
fn test_merge_mode_serde_round_trip() {
for (mode, wire) in [
(MergeMode::ObjectMerge, "\"object_merge\""),
(MergeMode::ArrayConcat, "\"array_concat\""),
(MergeMode::ObjectFirst, "\"object_first\""),
] {
let json = serde_json::to_string(&mode).expect("serialize");
assert_eq!(json, wire);
let decoded: MergeMode = serde_json::from_str(&json).expect("deserialize");
assert_eq!(decoded, mode);
}
assert_eq!(MergeMode::default(), MergeMode::ObjectMerge);
}
#[test]
fn test_credential_provider_azure_ad_round_trips_through_toml_and_json() {
let toml_src = r#"
model = "azure/gpt-4o"
[credential_provider]
type = "azure_ad"
tenant_id = "11111111-1111-1111-1111-111111111111"
client_id = "22222222-2222-2222-2222-222222222222"
client_secret = "example-client-secret"
scope = "https://cognitiveservices.azure.com/.default"
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
match cfg.credential_provider.as_deref() {
Some(CredentialProviderConfig::AzureAd {
tenant_id,
client_id,
client_secret,
scope,
}) => {
assert_eq!(tenant_id, "11111111-1111-1111-1111-111111111111");
assert_eq!(client_id, "22222222-2222-2222-2222-222222222222");
assert_eq!(client_secret, "example-client-secret");
assert_eq!(scope.as_deref(), Some("https://cognitiveservices.azure.com/.default"));
}
other => panic!("expected AzureAd variant, got {other:?}"),
}
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn credential_provider_rejects_unknown_fields() {
let json = r#"{
"type":"azure_ad",
"tenant_id":"tenant",
"client_id":"client",
"client_secret":"secret",
"unexpected_secret":"secret"
}"#;
assert!(serde_json::from_str::<CredentialProviderConfig>(json).is_err());
}
#[test]
fn test_credential_provider_vertex_oauth2_round_trips_through_toml_and_json() {
let toml_src = r#"
model = "vertex_ai/gemini-1.5-pro"
[credential_provider]
type = "vertex_oauth2"
service_account_key_file = "/etc/xberg/vertex-service-account.json"
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
match cfg.credential_provider.as_deref() {
Some(CredentialProviderConfig::VertexOauth2 {
service_account_key_file,
scope,
}) => {
assert_eq!(service_account_key_file, "/etc/xberg/vertex-service-account.json");
assert!(scope.is_none());
}
other => panic!("expected VertexOauth2 variant, got {other:?}"),
}
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_credential_provider_vertex_adc_round_trips_through_toml_and_json() {
let toml_src = r#"
model = "vertex_ai/gemini-1.5-pro"
[credential_provider]
type = "vertex_adc"
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
assert!(matches!(
cfg.credential_provider.as_deref(),
Some(CredentialProviderConfig::VertexAdc { scope: None })
));
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_credential_provider_bedrock_web_identity_round_trips_through_toml_and_json() {
let toml_src = r#"
model = "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
[credential_provider]
type = "bedrock_web_identity"
role_arn = "arn:aws:iam::123456789012:role/xberg-bedrock"
token_file = "/var/run/secrets/eks.amazonaws.com/serviceaccount/token"
session_name = "xberg-session"
region = "eu-central-1"
"#;
let cfg: LlmConfig = toml::from_str(toml_src).expect("deserialize LlmConfig from TOML");
match cfg.credential_provider.as_deref() {
Some(CredentialProviderConfig::BedrockWebIdentity {
role_arn,
token_file,
session_name,
region,
}) => {
assert_eq!(role_arn, "arn:aws:iam::123456789012:role/xberg-bedrock");
assert_eq!(token_file, "/var/run/secrets/eks.amazonaws.com/serviceaccount/token");
assert_eq!(session_name.as_deref(), Some("xberg-session"));
assert_eq!(region.as_deref(), Some("eu-central-1"));
}
other => panic!("expected BedrockWebIdentity variant, got {other:?}"),
}
let round_tripped: LlmConfig =
serde_json::from_str(&serde_json::to_string(&cfg).expect("serialize")).expect("deserialize");
assert_eq!(round_tripped, cfg);
}
#[test]
fn test_credential_provider_debug_redacts_azure_client_secret() {
let provider = CredentialProviderConfig::AzureAd {
tenant_id: "tenant-123".to_string(),
client_id: "client-456".to_string(),
client_secret: "super-secret-value".to_string(),
scope: Some("https://cognitiveservices.azure.com/.default".to_string()),
};
let rendered = format!("{provider:?}");
assert!(!rendered.contains("super-secret-value"), "leaked secret: {rendered}");
assert!(rendered.contains("tenant-123"), "{rendered}");
assert!(rendered.contains("client-456"), "{rendered}");
assert!(rendered.contains(REDACTED), "{rendered}");
}
#[test]
fn test_credential_provider_debug_prints_non_secret_variants_verbatim() {
let vertex_oauth2 = CredentialProviderConfig::VertexOauth2 {
service_account_key_file: "/etc/xberg/vertex.json".to_string(),
scope: None,
};
assert_eq!(
format!("{vertex_oauth2:?}"),
r#"VertexOauth2 { service_account_key_file: "/etc/xberg/vertex.json", scope: None }"#
);
let vertex_adc = CredentialProviderConfig::VertexAdc { scope: None };
assert_eq!(format!("{vertex_adc:?}"), "VertexAdc { scope: None }");
let bedrock = CredentialProviderConfig::BedrockWebIdentity {
role_arn: "arn:aws:iam::123456789012:role/xberg-bedrock".to_string(),
token_file: "/var/run/token".to_string(),
session_name: None,
region: None,
};
let rendered = format!("{bedrock:?}");
assert!(rendered.starts_with("BedrockWebIdentity {"), "{rendered}");
assert!(
rendered.contains(r#"role_arn: "arn:aws:iam::123456789012:role/xberg-bedrock""#),
"{rendered}"
);
assert!(rendered.contains(r#"token_file: "/var/run/token""#), "{rendered}");
assert!(rendered.contains("session_name: None"), "{rendered}");
assert!(rendered.contains("region: None"), "{rendered}");
}
#[test]
fn test_credential_provider_omitted_from_json_when_none() {
let cfg = LlmConfig {
model: "openai/gpt-4o".to_string(),
..LlmConfig::default()
};
let json = serde_json::to_string(&cfg).expect("serialize");
assert_eq!(json, r#"{"model":"openai/gpt-4o"}"#);
}
}