1use std::collections::HashMap;
4
5use crate::scenario::EndpointConfig;
6use crate::scenario::EndpointType;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
10pub enum TaskType {
11 Targeting,
14 Assertion,
16}
17
18impl TaskType {
19 #[must_use]
21 pub const fn as_str(self) -> &'static str {
22 match self {
23 Self::Targeting => "targeting",
24 Self::Assertion => "assertion",
25 }
26 }
27}
28
29#[derive(Debug, Clone)]
31pub struct ResolvedEndpoint {
32 pub name: String,
34 pub endpoint_type: EndpointType,
36 pub url: String,
38 pub model: Option<String>,
40 pub api_key: Option<String>,
42 pub headers: HashMap<String, String>,
44 pub command: Option<String>,
46 pub args: Vec<String>,
48 pub vision: bool,
50 pub input_price_per_1m: f64,
52 pub output_price_per_1m: f64,
54 pub per_call_price: f64,
56 pub max_attempts: u32,
59 pub fallbacks: Vec<String>,
62}
63
64impl ResolvedEndpoint {
65 #[must_use]
67 pub fn default_llm() -> Self {
68 Self {
69 name: "default".to_owned(),
70 endpoint_type: EndpointType::Llm,
71 url: crate::llm_base_url(),
72 model: Some(crate::llm_model()),
73 api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
74 headers: crate::parse_headers_env(),
75 command: None,
76 args: Vec::new(),
77 vision: false,
78 input_price_per_1m: 0.0,
79 output_price_per_1m: 0.0,
80 per_call_price: 0.0,
81 max_attempts: crate::default_llm_attempts(),
82 fallbacks: Vec::new(),
83 }
84 }
85}
86
87#[derive(Debug, Clone)]
89pub struct EndpointRegistry {
90 endpoints: HashMap<String, ResolvedEndpoint>,
91 default_for: HashMap<String, String>,
92}
93
94impl EndpointRegistry {
95 #[must_use]
104 pub fn from_config(
105 endpoints: &HashMap<String, EndpointConfig>,
106 fallback_llm: Option<&crate::LlmConfig>,
107 ) -> Self {
108 if endpoints.is_empty() {
109 let default_llm =
110 fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
111 name: "default".to_owned(),
112 endpoint_type: EndpointType::Llm,
113 url: llm.url.clone(),
114 model: Some(llm.model.clone()),
115 api_key: llm.api_key.clone(),
116 headers: llm.headers.clone(),
117 command: None,
118 args: Vec::new(),
119 vision: false,
120 input_price_per_1m: 0.0,
121 output_price_per_1m: 0.0,
122 per_call_price: 0.0,
123 max_attempts: llm.max_attempts,
124 fallbacks: Vec::new(),
125 });
126 let mut map = HashMap::new();
127 let mut default_for = HashMap::new();
128 for tt in &[TaskType::Targeting, TaskType::Assertion] {
129 default_for.insert(tt.as_str().to_owned(), "default".to_owned());
130 }
131 map.insert("default".to_owned(), default_llm);
132 return Self {
133 endpoints: map,
134 default_for,
135 };
136 }
137
138 let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
139 let mut default_for: HashMap<String, String> = HashMap::new();
140
141 for (name, ec) in endpoints {
142 let re = ResolvedEndpoint {
143 name: name.clone(),
144 endpoint_type: ec.endpoint_type.clone(),
145 url: ec
146 .url
147 .clone()
148 .unwrap_or_else(|| match ec.endpoint_type {
149 EndpointType::Llm => crate::llm_base_url(),
150 EndpointType::A2a | EndpointType::Mcp => String::new(),
151 })
152 .trim_end_matches('/')
153 .to_owned(),
154 model: ec.model.clone(),
155 api_key: ec.api_key.clone(),
156 headers: ec.headers.clone(),
157 command: ec.command.clone(),
158 args: ec.args.clone(),
159 vision: ec.vision,
160 input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
161 output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
162 per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
163 max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
164 fallbacks: ec.fallbacks.clone(),
165 };
166
167 for df in &ec.default_for {
168 default_for.insert(df.clone(), name.clone());
169 }
170
171 resolved.insert(name.clone(), re);
172 }
173
174 Self {
175 endpoints: resolved,
176 default_for,
177 }
178 }
179
180 #[must_use]
184 pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
185 self.endpoints.get(name)
186 }
187
188 #[must_use]
193 pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
194 let key = task.as_str();
195 if let Some(name) = self.default_for.get(key) {
196 if let Some(ep) = self.endpoints.get(name) {
197 return ep;
198 }
199 }
200 self.endpoints
202 .values()
203 .find(|ep| ep.endpoint_type == EndpointType::Llm)
204 .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
205 }
206
207 #[must_use]
210 pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
211 if let Some(n) = name {
212 if let Some(ep) = self.endpoints.get(n) {
213 return ep;
214 }
215 }
216 self.resolve_for_task(task)
217 }
218
219 #[must_use]
229 pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
230 let primary = self.resolve(name, task);
231 let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
232 let mut seen: std::collections::HashSet<&str> =
233 std::collections::HashSet::from([primary.name.as_str()]);
234 let mut cursor = primary;
235 for _ in 0..8 {
236 let next = cursor.fallbacks.iter().find_map(|fb| {
237 let ep = self.endpoints.get(fb)?;
238 (ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
239 .then_some(ep)
240 });
241 match next {
242 Some(ep) => {
243 seen.insert(ep.name.as_str());
244 chain.push(ep);
245 cursor = ep;
246 }
247 None => break,
248 }
249 }
250 chain
251 }
252
253 #[must_use]
255 pub fn len(&self) -> usize {
256 self.endpoints.len()
257 }
258
259 #[must_use]
261 pub fn is_empty(&self) -> bool {
262 self.endpoints.is_empty()
263 }
264}
265
266#[cfg(test)]
267mod tests {
268 use super::*;
269 use std::collections::HashMap;
270
271 #[test]
272 fn test_task_type_as_str() {
273 assert_eq!(TaskType::Targeting.as_str(), "targeting");
274 assert_eq!(TaskType::Assertion.as_str(), "assertion");
275 }
276
277 #[test]
278 fn test_registry_empty_config() {
279 let endpoints = HashMap::new();
280 let registry = EndpointRegistry::from_config(&endpoints, None);
281 assert_eq!(registry.len(), 1);
282 let ep = registry.get("default").unwrap();
283 assert_eq!(ep.endpoint_type, EndpointType::Llm);
284 }
285
286 #[test]
287 fn test_registry_resolve_by_name() {
288 let mut endpoints = HashMap::new();
289 endpoints.insert(
290 "vision".to_owned(),
291 EndpointConfig {
292 endpoint_type: EndpointType::Llm,
293 url: Some("https://api.openai.com".into()),
294 model: Some("gpt-4o".into()),
295 ..Default::default()
296 },
297 );
298
299 let registry = EndpointRegistry::from_config(&endpoints, None);
300 let ep = registry.get("vision");
301 assert!(ep.is_some());
302 assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
303 }
304
305 #[test]
306 fn test_resolve_for_task_with_default() {
307 let mut endpoints = HashMap::new();
308 let ec = EndpointConfig {
309 endpoint_type: EndpointType::Llm,
310 url: Some("http://localhost:8080".into()),
311 model: Some("deepseek".into()),
312 default_for: vec!["targeting".to_owned()],
313 ..Default::default()
314 };
315 endpoints.insert("main".to_owned(), ec);
316
317 let registry = EndpointRegistry::from_config(&endpoints, None);
318 let ep = registry.resolve_for_task(TaskType::Targeting);
319 assert_eq!(ep.name, "main");
320 }
321
322 #[test]
323 fn test_resolve_explicit_overrides_task() {
324 let mut endpoints = HashMap::new();
325 endpoints.insert(
326 "default".to_owned(),
327 EndpointConfig {
328 endpoint_type: EndpointType::Llm,
329 url: Some("http://default".into()),
330 default_for: vec!["targeting".to_owned()],
331 ..Default::default()
332 },
333 );
334 endpoints.insert(
335 "fast".to_owned(),
336 EndpointConfig {
337 endpoint_type: EndpointType::Llm,
338 url: Some("http://fast".into()),
339 ..Default::default()
340 },
341 );
342
343 let registry = EndpointRegistry::from_config(&endpoints, None);
344 let ep = registry.resolve(Some("fast"), TaskType::Targeting);
345 assert_eq!(ep.name, "fast");
346 }
347
348 #[test]
349 fn test_resolve_chain_follows_fallbacks() {
350 let mut endpoints = HashMap::new();
351 endpoints.insert(
352 "default".to_owned(),
353 EndpointConfig {
354 endpoint_type: EndpointType::Llm,
355 url: Some("http://default".into()),
356 default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
357 fallbacks: vec!["pro".to_owned()],
358 ..Default::default()
359 },
360 );
361 endpoints.insert(
362 "pro".to_owned(),
363 EndpointConfig {
364 endpoint_type: EndpointType::Llm,
365 url: Some("http://pro".into()),
366 model: Some("gpt-4.1".into()),
367 ..Default::default()
368 },
369 );
370
371 let registry = EndpointRegistry::from_config(&endpoints, None);
372 let chain = registry.resolve_chain(None, TaskType::Assertion);
373 assert_eq!(chain.len(), 2);
374 assert_eq!(chain[0].name, "default");
375 assert_eq!(chain[1].name, "pro");
376 }
377
378 #[test]
379 fn test_resolve_chain_skips_non_llm_and_cycles() {
380 let mut endpoints = HashMap::new();
381 endpoints.insert(
382 "default".to_owned(),
383 EndpointConfig {
384 endpoint_type: EndpointType::Llm,
385 url: Some("http://default".into()),
386 default_for: vec!["assertion".to_owned()],
387 fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
388 ..Default::default()
389 },
390 );
391 endpoints.insert(
393 "mcp1".to_owned(),
394 EndpointConfig {
395 endpoint_type: EndpointType::Mcp,
396 command: Some("npx".into()),
397 ..Default::default()
398 },
399 );
400 endpoints.insert(
402 "pro".to_owned(),
403 EndpointConfig {
404 endpoint_type: EndpointType::Llm,
405 url: Some("http://pro".into()),
406 fallbacks: vec!["default".to_owned()],
407 ..Default::default()
408 },
409 );
410
411 let registry = EndpointRegistry::from_config(&endpoints, None);
412 let chain = registry.resolve_chain(None, TaskType::Assertion);
413 assert_eq!(chain.len(), 2);
414 assert_eq!(chain[0].name, "default");
415 assert_eq!(chain[1].name, "pro");
416 }
417
418 #[test]
419 fn test_resolve_chain_max_attempts_default() {
420 let mut endpoints = HashMap::new();
421 endpoints.insert(
422 "default".to_owned(),
423 EndpointConfig {
424 endpoint_type: EndpointType::Llm,
425 url: Some("http://default".into()),
426 max_attempts: Some(7),
427 default_for: vec!["assertion".to_owned()],
428 ..Default::default()
429 },
430 );
431 let registry = EndpointRegistry::from_config(&endpoints, None);
432 let ep = registry.resolve_chain(None, TaskType::Assertion);
433 assert_eq!(ep[0].max_attempts, 7);
434 }
435
436 #[test]
437 fn test_default_llm_has_env_values() {
438 let ep = ResolvedEndpoint::default_llm();
439 assert_eq!(ep.endpoint_type, EndpointType::Llm);
440 assert!(ep.model.is_some());
441 assert!(!ep.url.is_empty());
442 }
443}