use std::collections::HashMap;
use crate::scenario::AuthConfig;
use crate::scenario::AwsConfig;
use crate::scenario::EndpointConfig;
use crate::scenario::EndpointType;
use crate::scenario::PricingConfig;
use crate::scenario::Provider;
fn cache_price(pricing: Option<&PricingConfig>, input_price: f64, read: bool) -> f64 {
let Some(p) = pricing else {
return 0.0;
};
let explicit = if read {
p.cached_input_per_1m_tokens
} else {
p.cache_write_per_1m_tokens
};
explicit.unwrap_or_else(|| {
let multiplier = if read {
p.cache_read_multiplier.unwrap_or(0.1)
} else {
p.cache_write_multiplier.unwrap_or(1.25)
};
input_price * multiplier
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum TaskType {
Targeting,
Assertion,
}
impl TaskType {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Targeting => "targeting",
Self::Assertion => "assertion",
}
}
}
#[derive(Debug, Clone)]
pub struct ResolvedEndpoint {
pub name: String,
pub endpoint_type: EndpointType,
pub url: String,
pub model: Option<String>,
pub api_key: Option<String>,
pub headers: HashMap<String, String>,
pub command: Option<String>,
pub args: Vec<String>,
pub vision: bool,
pub input_price_per_1m: f64,
pub output_price_per_1m: f64,
pub cached_input_price_per_1m: f64,
pub cache_write_price_per_1m: f64,
pub cache_pricing: bool,
pub cache_markers: bool,
pub per_call_price: f64,
pub max_attempts: u32,
pub fallbacks: Vec<String>,
pub provider: Provider,
pub deployment: Option<String>,
pub api_version: Option<String>,
pub auth: AuthConfig,
pub header_commands: HashMap<String, String>,
pub aws: AwsConfig,
}
impl ResolvedEndpoint {
#[must_use]
pub fn default_llm() -> Self {
Self {
name: "default".to_owned(),
endpoint_type: EndpointType::Llm,
url: crate::llm_base_url(),
model: Some(crate::llm_model()),
api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
headers: crate::parse_headers_env(),
command: None,
args: Vec::new(),
vision: false,
input_price_per_1m: 0.0,
output_price_per_1m: 0.0,
cached_input_price_per_1m: 0.0,
cache_write_price_per_1m: 0.0,
cache_pricing: true,
cache_markers: true,
per_call_price: 0.0,
max_attempts: crate::default_llm_attempts(),
fallbacks: Vec::new(),
provider: Provider::Openai,
deployment: None,
api_version: None,
auth: AuthConfig::default(),
header_commands: HashMap::new(),
aws: AwsConfig::default(),
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct EndpointDefaults {
pub cache: Option<bool>,
pub cache_pricing: Option<bool>,
}
#[derive(Debug, Clone)]
pub struct EndpointRegistry {
endpoints: HashMap<String, ResolvedEndpoint>,
default_for: HashMap<String, String>,
}
impl EndpointRegistry {
#[must_use]
#[allow(clippy::too_many_lines)]
pub fn from_config(
endpoints: &HashMap<String, EndpointConfig>,
fallback_llm: Option<&crate::LlmConfig>,
defaults: EndpointDefaults,
) -> Self {
if endpoints.is_empty() {
let default_llm =
fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
name: "default".to_owned(),
endpoint_type: EndpointType::Llm,
url: llm.url.clone(),
model: Some(llm.model.clone()),
api_key: llm.api_key.clone(),
headers: llm.headers.clone(),
command: None,
args: Vec::new(),
vision: false,
input_price_per_1m: 0.0,
output_price_per_1m: 0.0,
cached_input_price_per_1m: 0.0,
cache_write_price_per_1m: 0.0,
cache_pricing: defaults.cache_pricing.unwrap_or(true),
cache_markers: defaults.cache.unwrap_or(true),
per_call_price: 0.0,
max_attempts: llm.max_attempts,
fallbacks: Vec::new(),
provider: llm.provider,
deployment: llm.deployment.clone(),
api_version: llm.api_version.clone(),
auth: llm.auth.clone(),
header_commands: llm.header_commands.clone(),
aws: llm.aws.clone(),
});
let mut map = HashMap::new();
let mut default_for = HashMap::new();
for tt in &[TaskType::Targeting, TaskType::Assertion] {
default_for.insert(tt.as_str().to_owned(), "default".to_owned());
}
map.insert("default".to_owned(), default_llm);
return Self {
endpoints: map,
default_for,
};
}
let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
let mut default_for: HashMap<String, String> = HashMap::new();
for (name, ec) in endpoints {
let re = ResolvedEndpoint {
name: name.clone(),
endpoint_type: ec.endpoint_type.clone(),
url: ec
.url
.clone()
.unwrap_or_else(|| {
if ec.provider == Provider::Bedrock {
String::new()
} else {
match ec.endpoint_type {
EndpointType::Llm => crate::llm_base_url(),
EndpointType::A2a | EndpointType::Mcp => String::new(),
}
}
})
.trim_end_matches('/')
.to_owned(),
model: ec.model.clone(),
api_key: ec.api_key.clone(),
headers: ec.headers.clone(),
command: ec.command.clone(),
args: ec.args.clone(),
vision: ec.vision,
input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
cached_input_price_per_1m: cache_price(
ec.pricing.as_ref(),
ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
true,
),
cache_write_price_per_1m: cache_price(
ec.pricing.as_ref(),
ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
false,
),
cache_pricing: ec
.pricing
.as_ref()
.and_then(|p| p.cache_pricing)
.or(defaults.cache_pricing)
.unwrap_or(true),
cache_markers: ec.cache.or(defaults.cache).unwrap_or(true),
per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
fallbacks: ec.fallbacks.clone(),
provider: ec.provider,
deployment: ec.deployment.clone(),
api_version: ec.api_version.clone(),
auth: ec.auth.clone(),
header_commands: ec.header_commands.clone(),
aws: ec.aws.clone(),
};
for df in &ec.default_for {
default_for.insert(df.clone(), name.clone());
}
resolved.insert(name.clone(), re);
}
Self {
endpoints: resolved,
default_for,
}
}
#[must_use]
pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
self.endpoints.get(name)
}
#[must_use]
pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
let key = task.as_str();
if let Some(name) = self.default_for.get(key) {
if let Some(ep) = self.endpoints.get(name) {
return ep;
}
}
self.endpoints
.values()
.find(|ep| ep.endpoint_type == EndpointType::Llm)
.unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
}
#[must_use]
pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
if let Some(n) = name {
if let Some(ep) = self.endpoints.get(n) {
return ep;
}
}
self.resolve_for_task(task)
}
#[must_use]
pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
let primary = self.resolve(name, task);
let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
let mut seen: std::collections::HashSet<&str> =
std::collections::HashSet::from([primary.name.as_str()]);
let mut cursor = primary;
for _ in 0..8 {
let next = cursor.fallbacks.iter().find_map(|fb| {
let ep = self.endpoints.get(fb)?;
(ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
.then_some(ep)
});
match next {
Some(ep) => {
seen.insert(ep.name.as_str());
chain.push(ep);
cursor = ep;
}
None => break,
}
}
chain
}
#[must_use]
pub fn len(&self) -> usize {
self.endpoints.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.endpoints.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn test_task_type_as_str() {
assert_eq!(TaskType::Targeting.as_str(), "targeting");
assert_eq!(TaskType::Assertion.as_str(), "assertion");
}
#[test]
fn test_registry_empty_config() {
let endpoints = HashMap::new();
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
assert_eq!(registry.len(), 1);
let ep = registry.get("default").unwrap();
assert_eq!(ep.endpoint_type, EndpointType::Llm);
}
#[test]
fn test_registry_resolve_by_name() {
let mut endpoints = HashMap::new();
endpoints.insert(
"vision".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("https://api.openai.com".into()),
model: Some("gpt-4o".into()),
..Default::default()
},
);
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
let ep = registry.get("vision");
assert!(ep.is_some());
assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
}
#[test]
fn test_resolve_for_task_with_default() {
let mut endpoints = HashMap::new();
let ec = EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://localhost:8080".into()),
model: Some("deepseek".into()),
default_for: vec!["targeting".to_owned()],
..Default::default()
};
endpoints.insert("main".to_owned(), ec);
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
let ep = registry.resolve_for_task(TaskType::Targeting);
assert_eq!(ep.name, "main");
}
#[test]
fn test_resolve_explicit_overrides_task() {
let mut endpoints = HashMap::new();
endpoints.insert(
"default".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://default".into()),
default_for: vec!["targeting".to_owned()],
..Default::default()
},
);
endpoints.insert(
"fast".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://fast".into()),
..Default::default()
},
);
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
let ep = registry.resolve(Some("fast"), TaskType::Targeting);
assert_eq!(ep.name, "fast");
}
#[test]
fn test_resolve_chain_follows_fallbacks() {
let mut endpoints = HashMap::new();
endpoints.insert(
"default".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://default".into()),
default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
fallbacks: vec!["pro".to_owned()],
..Default::default()
},
);
endpoints.insert(
"pro".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://pro".into()),
model: Some("gpt-4.1".into()),
..Default::default()
},
);
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
let chain = registry.resolve_chain(None, TaskType::Assertion);
assert_eq!(chain.len(), 2);
assert_eq!(chain[0].name, "default");
assert_eq!(chain[1].name, "pro");
}
#[test]
fn test_resolve_chain_skips_non_llm_and_cycles() {
let mut endpoints = HashMap::new();
endpoints.insert(
"default".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://default".into()),
default_for: vec!["assertion".to_owned()],
fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
..Default::default()
},
);
endpoints.insert(
"mcp1".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Mcp,
command: Some("npx".into()),
..Default::default()
},
);
endpoints.insert(
"pro".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://pro".into()),
fallbacks: vec!["default".to_owned()],
..Default::default()
},
);
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
let chain = registry.resolve_chain(None, TaskType::Assertion);
assert_eq!(chain.len(), 2);
assert_eq!(chain[0].name, "default");
assert_eq!(chain[1].name, "pro");
}
#[test]
fn test_resolve_chain_max_attempts_default() {
let mut endpoints = HashMap::new();
endpoints.insert(
"default".to_owned(),
EndpointConfig {
endpoint_type: EndpointType::Llm,
url: Some("http://default".into()),
max_attempts: Some(7),
default_for: vec!["assertion".to_owned()],
..Default::default()
},
);
let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
let ep = registry.resolve_chain(None, TaskType::Assertion);
assert_eq!(ep[0].max_attempts, 7);
}
#[test]
fn test_default_llm_has_env_values() {
let ep = ResolvedEndpoint::default_llm();
assert_eq!(ep.endpoint_type, EndpointType::Llm);
assert!(ep.model.is_some());
assert_ne!(ep.url, "");
}
}