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