1use std::collections::HashMap;
4
5use crate::scenario::AuthConfig;
6use crate::scenario::AwsConfig;
7use crate::scenario::EndpointConfig;
8use crate::scenario::EndpointType;
9use crate::scenario::Provider;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub enum TaskType {
14 Targeting,
17 Assertion,
19}
20
21impl TaskType {
22 #[must_use]
24 pub const fn as_str(self) -> &'static str {
25 match self {
26 Self::Targeting => "targeting",
27 Self::Assertion => "assertion",
28 }
29 }
30}
31
32#[derive(Debug, Clone)]
34pub struct ResolvedEndpoint {
35 pub name: String,
37 pub endpoint_type: EndpointType,
39 pub url: String,
41 pub model: Option<String>,
43 pub api_key: Option<String>,
45 pub headers: HashMap<String, String>,
47 pub command: Option<String>,
49 pub args: Vec<String>,
51 pub vision: bool,
53 pub input_price_per_1m: f64,
55 pub output_price_per_1m: f64,
57 pub per_call_price: f64,
59 pub max_attempts: u32,
62 pub fallbacks: Vec<String>,
65 pub provider: Provider,
67 pub deployment: Option<String>,
69 pub api_version: Option<String>,
71 pub auth: AuthConfig,
73 pub header_commands: HashMap<String, String>,
75 pub aws: AwsConfig,
77}
78
79impl ResolvedEndpoint {
80 #[must_use]
82 pub fn default_llm() -> Self {
83 Self {
84 name: "default".to_owned(),
85 endpoint_type: EndpointType::Llm,
86 url: crate::llm_base_url(),
87 model: Some(crate::llm_model()),
88 api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
89 headers: crate::parse_headers_env(),
90 command: None,
91 args: Vec::new(),
92 vision: false,
93 input_price_per_1m: 0.0,
94 output_price_per_1m: 0.0,
95 per_call_price: 0.0,
96 max_attempts: crate::default_llm_attempts(),
97 fallbacks: Vec::new(),
98 provider: Provider::Openai,
99 deployment: None,
100 api_version: None,
101 auth: AuthConfig::default(),
102 header_commands: HashMap::new(),
103 aws: AwsConfig::default(),
104 }
105 }
106}
107
108#[derive(Debug, Clone)]
110pub struct EndpointRegistry {
111 endpoints: HashMap<String, ResolvedEndpoint>,
112 default_for: HashMap<String, String>,
113}
114
115impl EndpointRegistry {
116 #[must_use]
125 pub fn from_config(
126 endpoints: &HashMap<String, EndpointConfig>,
127 fallback_llm: Option<&crate::LlmConfig>,
128 ) -> Self {
129 if endpoints.is_empty() {
130 let default_llm =
131 fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
132 name: "default".to_owned(),
133 endpoint_type: EndpointType::Llm,
134 url: llm.url.clone(),
135 model: Some(llm.model.clone()),
136 api_key: llm.api_key.clone(),
137 headers: llm.headers.clone(),
138 command: None,
139 args: Vec::new(),
140 vision: false,
141 input_price_per_1m: 0.0,
142 output_price_per_1m: 0.0,
143 per_call_price: 0.0,
144 max_attempts: llm.max_attempts,
145 fallbacks: Vec::new(),
146 provider: llm.provider,
147 deployment: llm.deployment.clone(),
148 api_version: llm.api_version.clone(),
149 auth: llm.auth.clone(),
150 header_commands: llm.header_commands.clone(),
151 aws: llm.aws.clone(),
152 });
153 let mut map = HashMap::new();
154 let mut default_for = HashMap::new();
155 for tt in &[TaskType::Targeting, TaskType::Assertion] {
156 default_for.insert(tt.as_str().to_owned(), "default".to_owned());
157 }
158 map.insert("default".to_owned(), default_llm);
159 return Self {
160 endpoints: map,
161 default_for,
162 };
163 }
164
165 let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
166 let mut default_for: HashMap<String, String> = HashMap::new();
167
168 for (name, ec) in endpoints {
169 let re = ResolvedEndpoint {
170 name: name.clone(),
171 endpoint_type: ec.endpoint_type.clone(),
172 url: ec
173 .url
174 .clone()
175 .unwrap_or_else(|| {
176 if ec.provider == Provider::Bedrock {
177 String::new()
185 } else {
186 match ec.endpoint_type {
187 EndpointType::Llm => crate::llm_base_url(),
188 EndpointType::A2a | EndpointType::Mcp => String::new(),
189 }
190 }
191 })
192 .trim_end_matches('/')
193 .to_owned(),
194 model: ec.model.clone(),
195 api_key: ec.api_key.clone(),
196 headers: ec.headers.clone(),
197 command: ec.command.clone(),
198 args: ec.args.clone(),
199 vision: ec.vision,
200 input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
201 output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
202 per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
203 max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
204 fallbacks: ec.fallbacks.clone(),
205 provider: ec.provider,
206 deployment: ec.deployment.clone(),
207 api_version: ec.api_version.clone(),
208 auth: ec.auth.clone(),
209 header_commands: ec.header_commands.clone(),
210 aws: ec.aws.clone(),
211 };
212
213 for df in &ec.default_for {
214 default_for.insert(df.clone(), name.clone());
215 }
216
217 resolved.insert(name.clone(), re);
218 }
219
220 Self {
221 endpoints: resolved,
222 default_for,
223 }
224 }
225
226 #[must_use]
230 pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
231 self.endpoints.get(name)
232 }
233
234 #[must_use]
239 pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
240 let key = task.as_str();
241 if let Some(name) = self.default_for.get(key) {
242 if let Some(ep) = self.endpoints.get(name) {
243 return ep;
244 }
245 }
246 self.endpoints
248 .values()
249 .find(|ep| ep.endpoint_type == EndpointType::Llm)
250 .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
251 }
252
253 #[must_use]
256 pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
257 if let Some(n) = name {
258 if let Some(ep) = self.endpoints.get(n) {
259 return ep;
260 }
261 }
262 self.resolve_for_task(task)
263 }
264
265 #[must_use]
275 pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
276 let primary = self.resolve(name, task);
277 let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
278 let mut seen: std::collections::HashSet<&str> =
279 std::collections::HashSet::from([primary.name.as_str()]);
280 let mut cursor = primary;
281 for _ in 0..8 {
282 let next = cursor.fallbacks.iter().find_map(|fb| {
283 let ep = self.endpoints.get(fb)?;
284 (ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
285 .then_some(ep)
286 });
287 match next {
288 Some(ep) => {
289 seen.insert(ep.name.as_str());
290 chain.push(ep);
291 cursor = ep;
292 }
293 None => break,
294 }
295 }
296 chain
297 }
298
299 #[must_use]
301 pub fn len(&self) -> usize {
302 self.endpoints.len()
303 }
304
305 #[must_use]
307 pub fn is_empty(&self) -> bool {
308 self.endpoints.is_empty()
309 }
310}
311
312#[cfg(test)]
313mod tests {
314 use super::*;
315 use std::collections::HashMap;
316
317 #[test]
318 fn test_task_type_as_str() {
319 assert_eq!(TaskType::Targeting.as_str(), "targeting");
320 assert_eq!(TaskType::Assertion.as_str(), "assertion");
321 }
322
323 #[test]
324 fn test_registry_empty_config() {
325 let endpoints = HashMap::new();
326 let registry = EndpointRegistry::from_config(&endpoints, None);
327 assert_eq!(registry.len(), 1);
328 let ep = registry.get("default").unwrap();
329 assert_eq!(ep.endpoint_type, EndpointType::Llm);
330 }
331
332 #[test]
333 fn test_registry_resolve_by_name() {
334 let mut endpoints = HashMap::new();
335 endpoints.insert(
336 "vision".to_owned(),
337 EndpointConfig {
338 endpoint_type: EndpointType::Llm,
339 url: Some("https://api.openai.com".into()),
340 model: Some("gpt-4o".into()),
341 ..Default::default()
342 },
343 );
344
345 let registry = EndpointRegistry::from_config(&endpoints, None);
346 let ep = registry.get("vision");
347 assert!(ep.is_some());
348 assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
349 }
350
351 #[test]
352 fn test_resolve_for_task_with_default() {
353 let mut endpoints = HashMap::new();
354 let ec = EndpointConfig {
355 endpoint_type: EndpointType::Llm,
356 url: Some("http://localhost:8080".into()),
357 model: Some("deepseek".into()),
358 default_for: vec!["targeting".to_owned()],
359 ..Default::default()
360 };
361 endpoints.insert("main".to_owned(), ec);
362
363 let registry = EndpointRegistry::from_config(&endpoints, None);
364 let ep = registry.resolve_for_task(TaskType::Targeting);
365 assert_eq!(ep.name, "main");
366 }
367
368 #[test]
369 fn test_resolve_explicit_overrides_task() {
370 let mut endpoints = HashMap::new();
371 endpoints.insert(
372 "default".to_owned(),
373 EndpointConfig {
374 endpoint_type: EndpointType::Llm,
375 url: Some("http://default".into()),
376 default_for: vec!["targeting".to_owned()],
377 ..Default::default()
378 },
379 );
380 endpoints.insert(
381 "fast".to_owned(),
382 EndpointConfig {
383 endpoint_type: EndpointType::Llm,
384 url: Some("http://fast".into()),
385 ..Default::default()
386 },
387 );
388
389 let registry = EndpointRegistry::from_config(&endpoints, None);
390 let ep = registry.resolve(Some("fast"), TaskType::Targeting);
391 assert_eq!(ep.name, "fast");
392 }
393
394 #[test]
395 fn test_resolve_chain_follows_fallbacks() {
396 let mut endpoints = HashMap::new();
397 endpoints.insert(
398 "default".to_owned(),
399 EndpointConfig {
400 endpoint_type: EndpointType::Llm,
401 url: Some("http://default".into()),
402 default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
403 fallbacks: vec!["pro".to_owned()],
404 ..Default::default()
405 },
406 );
407 endpoints.insert(
408 "pro".to_owned(),
409 EndpointConfig {
410 endpoint_type: EndpointType::Llm,
411 url: Some("http://pro".into()),
412 model: Some("gpt-4.1".into()),
413 ..Default::default()
414 },
415 );
416
417 let registry = EndpointRegistry::from_config(&endpoints, None);
418 let chain = registry.resolve_chain(None, TaskType::Assertion);
419 assert_eq!(chain.len(), 2);
420 assert_eq!(chain[0].name, "default");
421 assert_eq!(chain[1].name, "pro");
422 }
423
424 #[test]
425 fn test_resolve_chain_skips_non_llm_and_cycles() {
426 let mut endpoints = HashMap::new();
427 endpoints.insert(
428 "default".to_owned(),
429 EndpointConfig {
430 endpoint_type: EndpointType::Llm,
431 url: Some("http://default".into()),
432 default_for: vec!["assertion".to_owned()],
433 fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
434 ..Default::default()
435 },
436 );
437 endpoints.insert(
439 "mcp1".to_owned(),
440 EndpointConfig {
441 endpoint_type: EndpointType::Mcp,
442 command: Some("npx".into()),
443 ..Default::default()
444 },
445 );
446 endpoints.insert(
448 "pro".to_owned(),
449 EndpointConfig {
450 endpoint_type: EndpointType::Llm,
451 url: Some("http://pro".into()),
452 fallbacks: vec!["default".to_owned()],
453 ..Default::default()
454 },
455 );
456
457 let registry = EndpointRegistry::from_config(&endpoints, None);
458 let chain = registry.resolve_chain(None, TaskType::Assertion);
459 assert_eq!(chain.len(), 2);
460 assert_eq!(chain[0].name, "default");
461 assert_eq!(chain[1].name, "pro");
462 }
463
464 #[test]
465 fn test_resolve_chain_max_attempts_default() {
466 let mut endpoints = HashMap::new();
467 endpoints.insert(
468 "default".to_owned(),
469 EndpointConfig {
470 endpoint_type: EndpointType::Llm,
471 url: Some("http://default".into()),
472 max_attempts: Some(7),
473 default_for: vec!["assertion".to_owned()],
474 ..Default::default()
475 },
476 );
477 let registry = EndpointRegistry::from_config(&endpoints, None);
478 let ep = registry.resolve_chain(None, TaskType::Assertion);
479 assert_eq!(ep[0].max_attempts, 7);
480 }
481
482 #[test]
483 fn test_default_llm_has_env_values() {
484 let ep = ResolvedEndpoint::default_llm();
485 assert_eq!(ep.endpoint_type, EndpointType::Llm);
486 assert!(ep.model.is_some());
487 assert!(!ep.url.is_empty());
488 }
489}