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(|| match ec.endpoint_type {
176 EndpointType::Llm => crate::llm_base_url(),
177 EndpointType::A2a | EndpointType::Mcp => String::new(),
178 })
179 .trim_end_matches('/')
180 .to_owned(),
181 model: ec.model.clone(),
182 api_key: ec.api_key.clone(),
183 headers: ec.headers.clone(),
184 command: ec.command.clone(),
185 args: ec.args.clone(),
186 vision: ec.vision,
187 input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
188 output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
189 per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
190 max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
191 fallbacks: ec.fallbacks.clone(),
192 provider: ec.provider,
193 deployment: ec.deployment.clone(),
194 api_version: ec.api_version.clone(),
195 auth: ec.auth.clone(),
196 header_commands: ec.header_commands.clone(),
197 aws: ec.aws.clone(),
198 };
199
200 for df in &ec.default_for {
201 default_for.insert(df.clone(), name.clone());
202 }
203
204 resolved.insert(name.clone(), re);
205 }
206
207 Self {
208 endpoints: resolved,
209 default_for,
210 }
211 }
212
213 #[must_use]
217 pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
218 self.endpoints.get(name)
219 }
220
221 #[must_use]
226 pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
227 let key = task.as_str();
228 if let Some(name) = self.default_for.get(key) {
229 if let Some(ep) = self.endpoints.get(name) {
230 return ep;
231 }
232 }
233 self.endpoints
235 .values()
236 .find(|ep| ep.endpoint_type == EndpointType::Llm)
237 .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
238 }
239
240 #[must_use]
243 pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
244 if let Some(n) = name {
245 if let Some(ep) = self.endpoints.get(n) {
246 return ep;
247 }
248 }
249 self.resolve_for_task(task)
250 }
251
252 #[must_use]
262 pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
263 let primary = self.resolve(name, task);
264 let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
265 let mut seen: std::collections::HashSet<&str> =
266 std::collections::HashSet::from([primary.name.as_str()]);
267 let mut cursor = primary;
268 for _ in 0..8 {
269 let next = cursor.fallbacks.iter().find_map(|fb| {
270 let ep = self.endpoints.get(fb)?;
271 (ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
272 .then_some(ep)
273 });
274 match next {
275 Some(ep) => {
276 seen.insert(ep.name.as_str());
277 chain.push(ep);
278 cursor = ep;
279 }
280 None => break,
281 }
282 }
283 chain
284 }
285
286 #[must_use]
288 pub fn len(&self) -> usize {
289 self.endpoints.len()
290 }
291
292 #[must_use]
294 pub fn is_empty(&self) -> bool {
295 self.endpoints.is_empty()
296 }
297}
298
299#[cfg(test)]
300mod tests {
301 use super::*;
302 use std::collections::HashMap;
303
304 #[test]
305 fn test_task_type_as_str() {
306 assert_eq!(TaskType::Targeting.as_str(), "targeting");
307 assert_eq!(TaskType::Assertion.as_str(), "assertion");
308 }
309
310 #[test]
311 fn test_registry_empty_config() {
312 let endpoints = HashMap::new();
313 let registry = EndpointRegistry::from_config(&endpoints, None);
314 assert_eq!(registry.len(), 1);
315 let ep = registry.get("default").unwrap();
316 assert_eq!(ep.endpoint_type, EndpointType::Llm);
317 }
318
319 #[test]
320 fn test_registry_resolve_by_name() {
321 let mut endpoints = HashMap::new();
322 endpoints.insert(
323 "vision".to_owned(),
324 EndpointConfig {
325 endpoint_type: EndpointType::Llm,
326 url: Some("https://api.openai.com".into()),
327 model: Some("gpt-4o".into()),
328 ..Default::default()
329 },
330 );
331
332 let registry = EndpointRegistry::from_config(&endpoints, None);
333 let ep = registry.get("vision");
334 assert!(ep.is_some());
335 assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
336 }
337
338 #[test]
339 fn test_resolve_for_task_with_default() {
340 let mut endpoints = HashMap::new();
341 let ec = EndpointConfig {
342 endpoint_type: EndpointType::Llm,
343 url: Some("http://localhost:8080".into()),
344 model: Some("deepseek".into()),
345 default_for: vec!["targeting".to_owned()],
346 ..Default::default()
347 };
348 endpoints.insert("main".to_owned(), ec);
349
350 let registry = EndpointRegistry::from_config(&endpoints, None);
351 let ep = registry.resolve_for_task(TaskType::Targeting);
352 assert_eq!(ep.name, "main");
353 }
354
355 #[test]
356 fn test_resolve_explicit_overrides_task() {
357 let mut endpoints = HashMap::new();
358 endpoints.insert(
359 "default".to_owned(),
360 EndpointConfig {
361 endpoint_type: EndpointType::Llm,
362 url: Some("http://default".into()),
363 default_for: vec!["targeting".to_owned()],
364 ..Default::default()
365 },
366 );
367 endpoints.insert(
368 "fast".to_owned(),
369 EndpointConfig {
370 endpoint_type: EndpointType::Llm,
371 url: Some("http://fast".into()),
372 ..Default::default()
373 },
374 );
375
376 let registry = EndpointRegistry::from_config(&endpoints, None);
377 let ep = registry.resolve(Some("fast"), TaskType::Targeting);
378 assert_eq!(ep.name, "fast");
379 }
380
381 #[test]
382 fn test_resolve_chain_follows_fallbacks() {
383 let mut endpoints = HashMap::new();
384 endpoints.insert(
385 "default".to_owned(),
386 EndpointConfig {
387 endpoint_type: EndpointType::Llm,
388 url: Some("http://default".into()),
389 default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
390 fallbacks: vec!["pro".to_owned()],
391 ..Default::default()
392 },
393 );
394 endpoints.insert(
395 "pro".to_owned(),
396 EndpointConfig {
397 endpoint_type: EndpointType::Llm,
398 url: Some("http://pro".into()),
399 model: Some("gpt-4.1".into()),
400 ..Default::default()
401 },
402 );
403
404 let registry = EndpointRegistry::from_config(&endpoints, None);
405 let chain = registry.resolve_chain(None, TaskType::Assertion);
406 assert_eq!(chain.len(), 2);
407 assert_eq!(chain[0].name, "default");
408 assert_eq!(chain[1].name, "pro");
409 }
410
411 #[test]
412 fn test_resolve_chain_skips_non_llm_and_cycles() {
413 let mut endpoints = HashMap::new();
414 endpoints.insert(
415 "default".to_owned(),
416 EndpointConfig {
417 endpoint_type: EndpointType::Llm,
418 url: Some("http://default".into()),
419 default_for: vec!["assertion".to_owned()],
420 fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
421 ..Default::default()
422 },
423 );
424 endpoints.insert(
426 "mcp1".to_owned(),
427 EndpointConfig {
428 endpoint_type: EndpointType::Mcp,
429 command: Some("npx".into()),
430 ..Default::default()
431 },
432 );
433 endpoints.insert(
435 "pro".to_owned(),
436 EndpointConfig {
437 endpoint_type: EndpointType::Llm,
438 url: Some("http://pro".into()),
439 fallbacks: vec!["default".to_owned()],
440 ..Default::default()
441 },
442 );
443
444 let registry = EndpointRegistry::from_config(&endpoints, None);
445 let chain = registry.resolve_chain(None, TaskType::Assertion);
446 assert_eq!(chain.len(), 2);
447 assert_eq!(chain[0].name, "default");
448 assert_eq!(chain[1].name, "pro");
449 }
450
451 #[test]
452 fn test_resolve_chain_max_attempts_default() {
453 let mut endpoints = HashMap::new();
454 endpoints.insert(
455 "default".to_owned(),
456 EndpointConfig {
457 endpoint_type: EndpointType::Llm,
458 url: Some("http://default".into()),
459 max_attempts: Some(7),
460 default_for: vec!["assertion".to_owned()],
461 ..Default::default()
462 },
463 );
464 let registry = EndpointRegistry::from_config(&endpoints, None);
465 let ep = registry.resolve_chain(None, TaskType::Assertion);
466 assert_eq!(ep[0].max_attempts, 7);
467 }
468
469 #[test]
470 fn test_default_llm_has_env_values() {
471 let ep = ResolvedEndpoint::default_llm();
472 assert_eq!(ep.endpoint_type, EndpointType::Llm);
473 assert!(ep.model.is_some());
474 assert!(!ep.url.is_empty());
475 }
476}