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::EndpointConfig;
6use crate::scenario::EndpointType;
7
8/// Classification of a task for endpoint routing.
9#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
10pub enum TaskType {
11    /// LLM-based element targeting (resolving CSS selectors from natural
12    /// language).
13    Targeting,
14    /// LLM-based assertion evaluation.
15    Assertion,
16}
17
18impl TaskType {
19    /// Returns the routing key string for this task type.
20    #[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/// Resolved endpoint ready for use in calls.
30#[derive(Debug, Clone)]
31pub struct ResolvedEndpoint {
32    /// Endpoint name.
33    pub name: String,
34    /// Endpoint type.
35    pub endpoint_type: EndpointType,
36    /// Base URL for HTTP-based endpoints.
37    pub url: String,
38    /// Model name (LLM endpoints only).
39    pub model: Option<String>,
40    /// API key / bearer token.
41    pub api_key: Option<String>,
42    /// Custom HTTP headers.
43    pub headers: HashMap<String, String>,
44    /// Command for MCP subprocess endpoints.
45    pub command: Option<String>,
46    /// Arguments for MCP subprocess endpoints.
47    pub args: Vec<String>,
48    /// Whether this endpoint accepts image parts (vision).
49    pub vision: bool,
50    /// Input token pricing per 1M tokens.
51    pub input_price_per_1m: f64,
52    /// Output token pricing per 1M tokens.
53    pub output_price_per_1m: f64,
54    /// Flat cost per call.
55    pub per_call_price: f64,
56    /// Retry budget for a single chat completion (default 3; global env
57    /// override `HARNESS_LLM_CALL_ATTEMPTS`).
58    pub max_attempts: u32,
59    /// Ordered names of fallback endpoints tried when this endpoint
60    /// exhausts its attempts (LLM endpoints only).
61    pub fallbacks: Vec<String>,
62}
63
64impl ResolvedEndpoint {
65    /// Creates a default LLM endpoint from environment variables.
66    #[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/// Registry of all configured endpoints with routing logic.
88#[derive(Debug, Clone)]
89pub struct EndpointRegistry {
90    endpoints: HashMap<String, ResolvedEndpoint>,
91    default_for: HashMap<String, String>,
92}
93
94impl EndpointRegistry {
95    /// Builds a registry from the endpoint definitions in scenario config.
96    ///
97    /// Falls back to a default LLM endpoint derived from `fallback_llm`
98    /// (the runner's effective config — CLI arguments merged over env vars
99    /// and scenario fields) when no `[config.endpoints]` are defined, so
100    /// `--llm-url` / `--llm-model` / `--llm-api-key` are honored even for
101    /// scenarios without an explicit endpoint table. Without a fallback,
102    /// environment variables are used.
103    #[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    /// Resolves an endpoint by explicit name.
181    ///
182    /// Returns `None` if no endpoint with the given name exists.
183    #[must_use]
184    pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
185        self.endpoints.get(name)
186    }
187
188    /// Resolves the best endpoint for a given task type.
189    ///
190    /// Checks for a `default_for` mapping first, then falls back to any LLM
191    /// endpoint, then panics (config error).
192    #[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        // Fallback: first LLM endpoint
201        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    /// Resolves an endpoint: explicit name takes priority, then task-type
208    /// routing, then first LLM endpoint.
209    #[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    /// Resolves the ordered call chain for a task: the primary endpoint
220    /// followed by its `fallbacks` (LLM endpoints only, deduplicated,
221    /// cycle-guarded, max 8 hops). Every LLM call goes through this chain —
222    /// the primary endpoint gets its own `max_attempts` retry budget, then
223    /// each fallback in turn, until one answers.
224    ///
225    /// Example: a cheap primary (`default`) with a more powerful fallback
226    /// (`pro`) can be declared as
227    /// `fallbacks = ["pro"]` on the `default` endpoint.
228    #[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    /// Returns the number of configured endpoints.
254    #[must_use]
255    pub fn len(&self) -> usize {
256        self.endpoints.len()
257    }
258
259    /// Returns true if no endpoints are configured.
260    #[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        // mcp1 is not an LLM endpoint — must be skipped in the chain.
392        endpoints.insert(
393            "mcp1".to_owned(),
394            EndpointConfig {
395                endpoint_type: EndpointType::Mcp,
396                command: Some("npx".into()),
397                ..Default::default()
398            },
399        );
400        // Cycle: pro -> default must terminate.
401        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}