Skip to main content

llm_browser_testkit/
endpoints.rs

1//! Endpoint registry — resolves named endpoints and task-type routing.
2
3use 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/// Classification of a task for endpoint routing.
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub enum TaskType {
14    /// LLM-based element targeting (resolving CSS selectors from natural
15    /// language).
16    Targeting,
17    /// LLM-based assertion evaluation.
18    Assertion,
19}
20
21impl TaskType {
22    /// Returns the routing key string for this task type.
23    #[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/// Resolved endpoint ready for use in calls.
33#[derive(Debug, Clone)]
34pub struct ResolvedEndpoint {
35    /// Endpoint name.
36    pub name: String,
37    /// Endpoint type.
38    pub endpoint_type: EndpointType,
39    /// Base URL for HTTP-based endpoints.
40    pub url: String,
41    /// Model name (LLM endpoints only).
42    pub model: Option<String>,
43    /// API key / bearer token.
44    pub api_key: Option<String>,
45    /// Custom HTTP headers.
46    pub headers: HashMap<String, String>,
47    /// Command for MCP subprocess endpoints.
48    pub command: Option<String>,
49    /// Arguments for MCP subprocess endpoints.
50    pub args: Vec<String>,
51    /// Whether this endpoint accepts image parts (vision).
52    pub vision: bool,
53    /// Input token pricing per 1M tokens.
54    pub input_price_per_1m: f64,
55    /// Output token pricing per 1M tokens.
56    pub output_price_per_1m: f64,
57    /// Flat cost per call.
58    pub per_call_price: f64,
59    /// Retry budget for a single chat completion (default 3; global env
60    /// override `HARNESS_LLM_CALL_ATTEMPTS`).
61    pub max_attempts: u32,
62    /// Ordered names of fallback endpoints tried when this endpoint
63    /// exhausts its attempts (LLM endpoints only).
64    pub fallbacks: Vec<String>,
65    /// LLM provider protocol.
66    pub provider: Provider,
67    /// `Azure` `OpenAI` deployment name (provider `azure`).
68    pub deployment: Option<String>,
69    /// `Azure` `OpenAI` API version (provider `azure`).
70    pub api_version: Option<String>,
71    /// Authentication configuration.
72    pub auth: AuthConfig,
73    /// Headers produced by running a command per call.
74    pub header_commands: HashMap<String, String>,
75    /// AWS credential settings (provider `bedrock`).
76    pub aws: AwsConfig,
77}
78
79impl ResolvedEndpoint {
80    /// Creates a default LLM endpoint from environment variables.
81    #[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/// Registry of all configured endpoints with routing logic.
109#[derive(Debug, Clone)]
110pub struct EndpointRegistry {
111    endpoints: HashMap<String, ResolvedEndpoint>,
112    default_for: HashMap<String, String>,
113}
114
115impl EndpointRegistry {
116    /// Builds a registry from the endpoint definitions in scenario config.
117    ///
118    /// Falls back to a default LLM endpoint derived from `fallback_llm`
119    /// (the runner's effective config — CLI arguments merged over env vars
120    /// and scenario fields) when no `[config.endpoints]` are defined, so
121    /// `--llm-url` / `--llm-model` / `--llm-api-key` are honored even for
122    /// scenarios without an explicit endpoint table. Without a fallback,
123    /// environment variables are used.
124    #[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    /// Resolves an endpoint by explicit name.
214    ///
215    /// Returns `None` if no endpoint with the given name exists.
216    #[must_use]
217    pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
218        self.endpoints.get(name)
219    }
220
221    /// Resolves the best endpoint for a given task type.
222    ///
223    /// Checks for a `default_for` mapping first, then falls back to any LLM
224    /// endpoint, then panics (config error).
225    #[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        // Fallback: first LLM endpoint
234        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    /// Resolves an endpoint: explicit name takes priority, then task-type
241    /// routing, then first LLM endpoint.
242    #[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    /// Resolves the ordered call chain for a task: the primary endpoint
253    /// followed by its `fallbacks` (LLM endpoints only, deduplicated,
254    /// cycle-guarded, max 8 hops). Every LLM call goes through this chain —
255    /// the primary endpoint gets its own `max_attempts` retry budget, then
256    /// each fallback in turn, until one answers.
257    ///
258    /// Example: a cheap primary (`default`) with a more powerful fallback
259    /// (`pro`) can be declared as
260    /// `fallbacks = ["pro"]` on the `default` endpoint.
261    #[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    /// Returns the number of configured endpoints.
287    #[must_use]
288    pub fn len(&self) -> usize {
289        self.endpoints.len()
290    }
291
292    /// Returns true if no endpoints are configured.
293    #[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        // mcp1 is not an LLM endpoint — must be skipped in the chain.
425        endpoints.insert(
426            "mcp1".to_owned(),
427            EndpointConfig {
428                endpoint_type: EndpointType::Mcp,
429                command: Some("npx".into()),
430                ..Default::default()
431            },
432        );
433        // Cycle: pro -> default must terminate.
434        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}