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}
57
58impl ResolvedEndpoint {
59 #[must_use]
61 pub fn default_llm() -> Self {
62 Self {
63 name: "default".to_owned(),
64 endpoint_type: EndpointType::Llm,
65 url: crate::llm_base_url(),
66 model: Some(crate::llm_model()),
67 api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
68 headers: crate::parse_headers_env(),
69 command: None,
70 args: Vec::new(),
71 vision: false,
72 input_price_per_1m: 0.0,
73 output_price_per_1m: 0.0,
74 per_call_price: 0.0,
75 }
76 }
77}
78
79#[derive(Debug, Clone)]
81pub struct EndpointRegistry {
82 endpoints: HashMap<String, ResolvedEndpoint>,
83 default_for: HashMap<String, String>,
84}
85
86impl EndpointRegistry {
87 #[must_use]
96 pub fn from_config(
97 endpoints: &HashMap<String, EndpointConfig>,
98 fallback_llm: Option<&crate::LlmConfig>,
99 ) -> Self {
100 if endpoints.is_empty() {
101 let default_llm =
102 fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
103 name: "default".to_owned(),
104 endpoint_type: EndpointType::Llm,
105 url: llm.url.clone(),
106 model: Some(llm.model.clone()),
107 api_key: llm.api_key.clone(),
108 headers: llm.headers.clone(),
109 command: None,
110 args: Vec::new(),
111 vision: false,
112 input_price_per_1m: 0.0,
113 output_price_per_1m: 0.0,
114 per_call_price: 0.0,
115 });
116 let mut map = HashMap::new();
117 let mut default_for = HashMap::new();
118 for tt in &[TaskType::Targeting, TaskType::Assertion] {
119 default_for.insert(tt.as_str().to_owned(), "default".to_owned());
120 }
121 map.insert("default".to_owned(), default_llm);
122 return Self {
123 endpoints: map,
124 default_for,
125 };
126 }
127
128 let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
129 let mut default_for: HashMap<String, String> = HashMap::new();
130
131 for (name, ec) in endpoints {
132 let re = ResolvedEndpoint {
133 name: name.clone(),
134 endpoint_type: ec.endpoint_type.clone(),
135 url: ec
136 .url
137 .clone()
138 .unwrap_or_else(|| match ec.endpoint_type {
139 EndpointType::Llm => crate::llm_base_url(),
140 EndpointType::A2a | EndpointType::Mcp => String::new(),
141 })
142 .trim_end_matches('/')
143 .to_owned(),
144 model: ec.model.clone(),
145 api_key: ec.api_key.clone(),
146 headers: ec.headers.clone(),
147 command: ec.command.clone(),
148 args: ec.args.clone(),
149 vision: ec.vision,
150 input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
151 output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
152 per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
153 };
154
155 for df in &ec.default_for {
156 default_for.insert(df.clone(), name.clone());
157 }
158
159 resolved.insert(name.clone(), re);
160 }
161
162 Self {
163 endpoints: resolved,
164 default_for,
165 }
166 }
167
168 #[must_use]
172 pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
173 self.endpoints.get(name)
174 }
175
176 #[must_use]
181 pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
182 let key = task.as_str();
183 if let Some(name) = self.default_for.get(key) {
184 if let Some(ep) = self.endpoints.get(name) {
185 return ep;
186 }
187 }
188 self.endpoints
190 .values()
191 .find(|ep| ep.endpoint_type == EndpointType::Llm)
192 .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
193 }
194
195 #[must_use]
198 pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
199 if let Some(n) = name {
200 if let Some(ep) = self.endpoints.get(n) {
201 return ep;
202 }
203 }
204 self.resolve_for_task(task)
205 }
206
207 #[must_use]
209 pub fn len(&self) -> usize {
210 self.endpoints.len()
211 }
212
213 #[must_use]
215 pub fn is_empty(&self) -> bool {
216 self.endpoints.is_empty()
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use super::*;
223 use std::collections::HashMap;
224
225 #[test]
226 fn test_task_type_as_str() {
227 assert_eq!(TaskType::Targeting.as_str(), "targeting");
228 assert_eq!(TaskType::Assertion.as_str(), "assertion");
229 }
230
231 #[test]
232 fn test_registry_empty_config() {
233 let endpoints = HashMap::new();
234 let registry = EndpointRegistry::from_config(&endpoints, None);
235 assert_eq!(registry.len(), 1);
236 let ep = registry.get("default").unwrap();
237 assert_eq!(ep.endpoint_type, EndpointType::Llm);
238 }
239
240 #[test]
241 fn test_registry_resolve_by_name() {
242 let mut endpoints = HashMap::new();
243 endpoints.insert(
244 "vision".to_owned(),
245 EndpointConfig {
246 endpoint_type: EndpointType::Llm,
247 url: Some("https://api.openai.com".into()),
248 model: Some("gpt-4o".into()),
249 ..Default::default()
250 },
251 );
252
253 let registry = EndpointRegistry::from_config(&endpoints, None);
254 let ep = registry.get("vision");
255 assert!(ep.is_some());
256 assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
257 }
258
259 #[test]
260 fn test_resolve_for_task_with_default() {
261 let mut endpoints = HashMap::new();
262 let ec = EndpointConfig {
263 endpoint_type: EndpointType::Llm,
264 url: Some("http://localhost:8080".into()),
265 model: Some("deepseek".into()),
266 default_for: vec!["targeting".to_owned()],
267 ..Default::default()
268 };
269 endpoints.insert("main".to_owned(), ec);
270
271 let registry = EndpointRegistry::from_config(&endpoints, None);
272 let ep = registry.resolve_for_task(TaskType::Targeting);
273 assert_eq!(ep.name, "main");
274 }
275
276 #[test]
277 fn test_resolve_explicit_overrides_task() {
278 let mut endpoints = HashMap::new();
279 endpoints.insert(
280 "default".to_owned(),
281 EndpointConfig {
282 endpoint_type: EndpointType::Llm,
283 url: Some("http://default".into()),
284 default_for: vec!["targeting".to_owned()],
285 ..Default::default()
286 },
287 );
288 endpoints.insert(
289 "fast".to_owned(),
290 EndpointConfig {
291 endpoint_type: EndpointType::Llm,
292 url: Some("http://fast".into()),
293 ..Default::default()
294 },
295 );
296
297 let registry = EndpointRegistry::from_config(&endpoints, None);
298 let ep = registry.resolve(Some("fast"), TaskType::Targeting);
299 assert_eq!(ep.name, "fast");
300 }
301
302 #[test]
303 fn test_default_llm_has_env_values() {
304 let ep = ResolvedEndpoint::default_llm();
305 assert_eq!(ep.endpoint_type, EndpointType::Llm);
306 assert!(ep.model.is_some());
307 assert!(!ep.url.is_empty());
308 }
309}