1use std::collections::HashMap;
4
5use crate::scenario::AuthConfig;
6use crate::scenario::AwsConfig;
7use crate::scenario::EndpointConfig;
8use crate::scenario::EndpointType;
9use crate::scenario::PricingConfig;
10use crate::scenario::Provider;
11
12fn cache_price(pricing: Option<&PricingConfig>, input_price: f64, read: bool) -> f64 {
16 let Some(p) = pricing else {
17 return 0.0;
18 };
19 let explicit = if read {
20 p.cached_input_per_1m_tokens
21 } else {
22 p.cache_write_per_1m_tokens
23 };
24 explicit.unwrap_or_else(|| {
25 let multiplier = if read {
26 p.cache_read_multiplier.unwrap_or(0.1)
27 } else {
28 p.cache_write_multiplier.unwrap_or(1.25)
29 };
30 input_price * multiplier
31 })
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
36pub enum TaskType {
37 Targeting,
40 Assertion,
42}
43
44impl TaskType {
45 #[must_use]
47 pub const fn as_str(self) -> &'static str {
48 match self {
49 Self::Targeting => "targeting",
50 Self::Assertion => "assertion",
51 }
52 }
53}
54
55#[derive(Debug, Clone)]
57pub struct ResolvedEndpoint {
58 pub name: String,
60 pub endpoint_type: EndpointType,
62 pub url: String,
64 pub model: Option<String>,
66 pub api_key: Option<String>,
68 pub headers: HashMap<String, String>,
70 pub command: Option<String>,
72 pub args: Vec<String>,
74 pub vision: bool,
76 pub input_price_per_1m: f64,
78 pub output_price_per_1m: f64,
80 pub cached_input_price_per_1m: f64,
82 pub cache_write_price_per_1m: f64,
84 pub cache_pricing: bool,
86 pub cache_markers: bool,
88 pub per_call_price: f64,
90 pub max_attempts: u32,
93 pub fallbacks: Vec<String>,
96 pub provider: Provider,
98 pub deployment: Option<String>,
100 pub api_version: Option<String>,
102 pub auth: AuthConfig,
104 pub header_commands: HashMap<String, String>,
106 pub aws: AwsConfig,
108}
109
110impl ResolvedEndpoint {
111 #[must_use]
113 pub fn default_llm() -> Self {
114 Self {
115 name: "default".to_owned(),
116 endpoint_type: EndpointType::Llm,
117 url: crate::llm_base_url(),
118 model: Some(crate::llm_model()),
119 api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
120 headers: crate::parse_headers_env(),
121 command: None,
122 args: Vec::new(),
123 vision: false,
124 input_price_per_1m: 0.0,
125 output_price_per_1m: 0.0,
126 cached_input_price_per_1m: 0.0,
127 cache_write_price_per_1m: 0.0,
128 cache_pricing: true,
129 cache_markers: true,
130 per_call_price: 0.0,
131 max_attempts: crate::default_llm_attempts(),
132 fallbacks: Vec::new(),
133 provider: Provider::Openai,
134 deployment: None,
135 api_version: None,
136 auth: AuthConfig::default(),
137 header_commands: HashMap::new(),
138 aws: AwsConfig::default(),
139 }
140 }
141}
142
143#[derive(Debug, Clone, Copy, Default)]
145pub struct EndpointDefaults {
146 pub cache: Option<bool>,
148 pub cache_pricing: Option<bool>,
150}
151
152#[derive(Debug, Clone)]
154pub struct EndpointRegistry {
155 endpoints: HashMap<String, ResolvedEndpoint>,
156 default_for: HashMap<String, String>,
157}
158
159impl EndpointRegistry {
160 #[must_use]
169 #[allow(clippy::too_many_lines)]
170 pub fn from_config(
171 endpoints: &HashMap<String, EndpointConfig>,
172 fallback_llm: Option<&crate::LlmConfig>,
173 defaults: EndpointDefaults,
174 ) -> Self {
175 if endpoints.is_empty() {
176 let default_llm =
177 fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
178 name: "default".to_owned(),
179 endpoint_type: EndpointType::Llm,
180 url: llm.url.clone(),
181 model: Some(llm.model.clone()),
182 api_key: llm.api_key.clone(),
183 headers: llm.headers.clone(),
184 command: None,
185 args: Vec::new(),
186 vision: false,
187 input_price_per_1m: 0.0,
188 output_price_per_1m: 0.0,
189 cached_input_price_per_1m: 0.0,
190 cache_write_price_per_1m: 0.0,
191 cache_pricing: defaults.cache_pricing.unwrap_or(true),
192 cache_markers: defaults.cache.unwrap_or(true),
193 per_call_price: 0.0,
194 max_attempts: llm.max_attempts,
195 fallbacks: Vec::new(),
196 provider: llm.provider,
197 deployment: llm.deployment.clone(),
198 api_version: llm.api_version.clone(),
199 auth: llm.auth.clone(),
200 header_commands: llm.header_commands.clone(),
201 aws: llm.aws.clone(),
202 });
203 let mut map = HashMap::new();
204 let mut default_for = HashMap::new();
205 for tt in &[TaskType::Targeting, TaskType::Assertion] {
206 default_for.insert(tt.as_str().to_owned(), "default".to_owned());
207 }
208 map.insert("default".to_owned(), default_llm);
209 return Self {
210 endpoints: map,
211 default_for,
212 };
213 }
214
215 let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
216 let mut default_for: HashMap<String, String> = HashMap::new();
217
218 for (name, ec) in endpoints {
219 let re = ResolvedEndpoint {
220 name: name.clone(),
221 endpoint_type: ec.endpoint_type.clone(),
222 url: ec
223 .url
224 .clone()
225 .unwrap_or_else(|| {
226 if ec.provider == Provider::Bedrock {
227 String::new()
235 } else {
236 match ec.endpoint_type {
237 EndpointType::Llm => crate::llm_base_url(),
238 EndpointType::A2a | EndpointType::Mcp => String::new(),
239 }
240 }
241 })
242 .trim_end_matches('/')
243 .to_owned(),
244 model: ec.model.clone(),
245 api_key: ec.api_key.clone(),
246 headers: ec.headers.clone(),
247 command: ec.command.clone(),
248 args: ec.args.clone(),
249 vision: ec.vision,
250 input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
251 output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
252 cached_input_price_per_1m: cache_price(
253 ec.pricing.as_ref(),
254 ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
255 true,
256 ),
257 cache_write_price_per_1m: cache_price(
258 ec.pricing.as_ref(),
259 ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
260 false,
261 ),
262 cache_pricing: ec
263 .pricing
264 .as_ref()
265 .and_then(|p| p.cache_pricing)
266 .or(defaults.cache_pricing)
267 .unwrap_or(true),
268 cache_markers: ec.cache.or(defaults.cache).unwrap_or(true),
269 per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
270 max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
271 fallbacks: ec.fallbacks.clone(),
272 provider: ec.provider,
273 deployment: ec.deployment.clone(),
274 api_version: ec.api_version.clone(),
275 auth: ec.auth.clone(),
276 header_commands: ec.header_commands.clone(),
277 aws: ec.aws.clone(),
278 };
279
280 for df in &ec.default_for {
281 default_for.insert(df.clone(), name.clone());
282 }
283
284 resolved.insert(name.clone(), re);
285 }
286
287 Self {
288 endpoints: resolved,
289 default_for,
290 }
291 }
292
293 #[must_use]
297 pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
298 self.endpoints.get(name)
299 }
300
301 #[must_use]
306 pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
307 let key = task.as_str();
308 if let Some(name) = self.default_for.get(key) {
309 if let Some(ep) = self.endpoints.get(name) {
310 return ep;
311 }
312 }
313 self.endpoints
315 .values()
316 .find(|ep| ep.endpoint_type == EndpointType::Llm)
317 .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
318 }
319
320 #[must_use]
323 pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
324 if let Some(n) = name {
325 if let Some(ep) = self.endpoints.get(n) {
326 return ep;
327 }
328 }
329 self.resolve_for_task(task)
330 }
331
332 #[must_use]
342 pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
343 let primary = self.resolve(name, task);
344 let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
345 let mut seen: std::collections::HashSet<&str> =
346 std::collections::HashSet::from([primary.name.as_str()]);
347 let mut cursor = primary;
348 for _ in 0..8 {
349 let next = cursor.fallbacks.iter().find_map(|fb| {
350 let ep = self.endpoints.get(fb)?;
351 (ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
352 .then_some(ep)
353 });
354 match next {
355 Some(ep) => {
356 seen.insert(ep.name.as_str());
357 chain.push(ep);
358 cursor = ep;
359 }
360 None => break,
361 }
362 }
363 chain
364 }
365
366 #[must_use]
368 pub fn len(&self) -> usize {
369 self.endpoints.len()
370 }
371
372 #[must_use]
374 pub fn is_empty(&self) -> bool {
375 self.endpoints.is_empty()
376 }
377}
378
379#[cfg(test)]
380mod tests {
381 use super::*;
382 use std::collections::HashMap;
383
384 #[test]
385 fn test_task_type_as_str() {
386 assert_eq!(TaskType::Targeting.as_str(), "targeting");
387 assert_eq!(TaskType::Assertion.as_str(), "assertion");
388 }
389
390 #[test]
391 fn test_registry_empty_config() {
392 let endpoints = HashMap::new();
393 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
394 assert_eq!(registry.len(), 1);
395 let ep = registry.get("default").unwrap();
396 assert_eq!(ep.endpoint_type, EndpointType::Llm);
397 }
398
399 #[test]
400 fn test_registry_resolve_by_name() {
401 let mut endpoints = HashMap::new();
402 endpoints.insert(
403 "vision".to_owned(),
404 EndpointConfig {
405 endpoint_type: EndpointType::Llm,
406 url: Some("https://api.openai.com".into()),
407 model: Some("gpt-4o".into()),
408 ..Default::default()
409 },
410 );
411
412 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
413 let ep = registry.get("vision");
414 assert!(ep.is_some());
415 assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
416 }
417
418 #[test]
419 fn test_resolve_for_task_with_default() {
420 let mut endpoints = HashMap::new();
421 let ec = EndpointConfig {
422 endpoint_type: EndpointType::Llm,
423 url: Some("http://localhost:8080".into()),
424 model: Some("deepseek".into()),
425 default_for: vec!["targeting".to_owned()],
426 ..Default::default()
427 };
428 endpoints.insert("main".to_owned(), ec);
429
430 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
431 let ep = registry.resolve_for_task(TaskType::Targeting);
432 assert_eq!(ep.name, "main");
433 }
434
435 #[test]
436 fn test_resolve_explicit_overrides_task() {
437 let mut endpoints = HashMap::new();
438 endpoints.insert(
439 "default".to_owned(),
440 EndpointConfig {
441 endpoint_type: EndpointType::Llm,
442 url: Some("http://default".into()),
443 default_for: vec!["targeting".to_owned()],
444 ..Default::default()
445 },
446 );
447 endpoints.insert(
448 "fast".to_owned(),
449 EndpointConfig {
450 endpoint_type: EndpointType::Llm,
451 url: Some("http://fast".into()),
452 ..Default::default()
453 },
454 );
455
456 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
457 let ep = registry.resolve(Some("fast"), TaskType::Targeting);
458 assert_eq!(ep.name, "fast");
459 }
460
461 #[test]
462 fn test_resolve_chain_follows_fallbacks() {
463 let mut endpoints = HashMap::new();
464 endpoints.insert(
465 "default".to_owned(),
466 EndpointConfig {
467 endpoint_type: EndpointType::Llm,
468 url: Some("http://default".into()),
469 default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
470 fallbacks: vec!["pro".to_owned()],
471 ..Default::default()
472 },
473 );
474 endpoints.insert(
475 "pro".to_owned(),
476 EndpointConfig {
477 endpoint_type: EndpointType::Llm,
478 url: Some("http://pro".into()),
479 model: Some("gpt-4.1".into()),
480 ..Default::default()
481 },
482 );
483
484 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
485 let chain = registry.resolve_chain(None, TaskType::Assertion);
486 assert_eq!(chain.len(), 2);
487 assert_eq!(chain[0].name, "default");
488 assert_eq!(chain[1].name, "pro");
489 }
490
491 #[test]
492 fn test_resolve_chain_skips_non_llm_and_cycles() {
493 let mut endpoints = HashMap::new();
494 endpoints.insert(
495 "default".to_owned(),
496 EndpointConfig {
497 endpoint_type: EndpointType::Llm,
498 url: Some("http://default".into()),
499 default_for: vec!["assertion".to_owned()],
500 fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
501 ..Default::default()
502 },
503 );
504 endpoints.insert(
506 "mcp1".to_owned(),
507 EndpointConfig {
508 endpoint_type: EndpointType::Mcp,
509 command: Some("npx".into()),
510 ..Default::default()
511 },
512 );
513 endpoints.insert(
515 "pro".to_owned(),
516 EndpointConfig {
517 endpoint_type: EndpointType::Llm,
518 url: Some("http://pro".into()),
519 fallbacks: vec!["default".to_owned()],
520 ..Default::default()
521 },
522 );
523
524 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
525 let chain = registry.resolve_chain(None, TaskType::Assertion);
526 assert_eq!(chain.len(), 2);
527 assert_eq!(chain[0].name, "default");
528 assert_eq!(chain[1].name, "pro");
529 }
530
531 #[test]
532 fn test_resolve_chain_max_attempts_default() {
533 let mut endpoints = HashMap::new();
534 endpoints.insert(
535 "default".to_owned(),
536 EndpointConfig {
537 endpoint_type: EndpointType::Llm,
538 url: Some("http://default".into()),
539 max_attempts: Some(7),
540 default_for: vec!["assertion".to_owned()],
541 ..Default::default()
542 },
543 );
544 let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
545 let ep = registry.resolve_chain(None, TaskType::Assertion);
546 assert_eq!(ep[0].max_attempts, 7);
547 }
548
549 #[test]
550 fn test_default_llm_has_env_values() {
551 let ep = ResolvedEndpoint::default_llm();
552 assert_eq!(ep.endpoint_type, EndpointType::Llm);
553 assert!(ep.model.is_some());
554 assert_ne!(ep.url, "");
555 }
556}