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(|| {
176                        if ec.provider == Provider::Bedrock {
177                            // Bedrock builds its URL from region + model
178                            // (https://bedrock-runtime.<region>.amazonaws.com/
179                            // model/<model>/converse). Substituting the
180                            // OpenAI-compatible base URL here would send the
181                            // SigV4-signed Converse request to the wrong host
182                            // (observed: the exo gateway's FastAPI root answers
183                            // 405 {"detail":"Method Not Allowed"}).
184                            String::new()
185                        } else {
186                            match ec.endpoint_type {
187                                EndpointType::Llm => crate::llm_base_url(),
188                                EndpointType::A2a | EndpointType::Mcp => String::new(),
189                            }
190                        }
191                    })
192                    .trim_end_matches('/')
193                    .to_owned(),
194                model: ec.model.clone(),
195                api_key: ec.api_key.clone(),
196                headers: ec.headers.clone(),
197                command: ec.command.clone(),
198                args: ec.args.clone(),
199                vision: ec.vision,
200                input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
201                output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
202                per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
203                max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
204                fallbacks: ec.fallbacks.clone(),
205                provider: ec.provider,
206                deployment: ec.deployment.clone(),
207                api_version: ec.api_version.clone(),
208                auth: ec.auth.clone(),
209                header_commands: ec.header_commands.clone(),
210                aws: ec.aws.clone(),
211            };
212
213            for df in &ec.default_for {
214                default_for.insert(df.clone(), name.clone());
215            }
216
217            resolved.insert(name.clone(), re);
218        }
219
220        Self {
221            endpoints: resolved,
222            default_for,
223        }
224    }
225
226    /// Resolves an endpoint by explicit name.
227    ///
228    /// Returns `None` if no endpoint with the given name exists.
229    #[must_use]
230    pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
231        self.endpoints.get(name)
232    }
233
234    /// Resolves the best endpoint for a given task type.
235    ///
236    /// Checks for a `default_for` mapping first, then falls back to any LLM
237    /// endpoint, then panics (config error).
238    #[must_use]
239    pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
240        let key = task.as_str();
241        if let Some(name) = self.default_for.get(key) {
242            if let Some(ep) = self.endpoints.get(name) {
243                return ep;
244            }
245        }
246        // Fallback: first LLM endpoint
247        self.endpoints
248            .values()
249            .find(|ep| ep.endpoint_type == EndpointType::Llm)
250            .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
251    }
252
253    /// Resolves an endpoint: explicit name takes priority, then task-type
254    /// routing, then first LLM endpoint.
255    #[must_use]
256    pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
257        if let Some(n) = name {
258            if let Some(ep) = self.endpoints.get(n) {
259                return ep;
260            }
261        }
262        self.resolve_for_task(task)
263    }
264
265    /// Resolves the ordered call chain for a task: the primary endpoint
266    /// followed by its `fallbacks` (LLM endpoints only, deduplicated,
267    /// cycle-guarded, max 8 hops). Every LLM call goes through this chain —
268    /// the primary endpoint gets its own `max_attempts` retry budget, then
269    /// each fallback in turn, until one answers.
270    ///
271    /// Example: a cheap primary (`default`) with a more powerful fallback
272    /// (`pro`) can be declared as
273    /// `fallbacks = ["pro"]` on the `default` endpoint.
274    #[must_use]
275    pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
276        let primary = self.resolve(name, task);
277        let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
278        let mut seen: std::collections::HashSet<&str> =
279            std::collections::HashSet::from([primary.name.as_str()]);
280        let mut cursor = primary;
281        for _ in 0..8 {
282            let next = cursor.fallbacks.iter().find_map(|fb| {
283                let ep = self.endpoints.get(fb)?;
284                (ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
285                    .then_some(ep)
286            });
287            match next {
288                Some(ep) => {
289                    seen.insert(ep.name.as_str());
290                    chain.push(ep);
291                    cursor = ep;
292                }
293                None => break,
294            }
295        }
296        chain
297    }
298
299    /// Returns the number of configured endpoints.
300    #[must_use]
301    pub fn len(&self) -> usize {
302        self.endpoints.len()
303    }
304
305    /// Returns true if no endpoints are configured.
306    #[must_use]
307    pub fn is_empty(&self) -> bool {
308        self.endpoints.is_empty()
309    }
310}
311
312#[cfg(test)]
313mod tests {
314    use super::*;
315    use std::collections::HashMap;
316
317    #[test]
318    fn test_task_type_as_str() {
319        assert_eq!(TaskType::Targeting.as_str(), "targeting");
320        assert_eq!(TaskType::Assertion.as_str(), "assertion");
321    }
322
323    #[test]
324    fn test_registry_empty_config() {
325        let endpoints = HashMap::new();
326        let registry = EndpointRegistry::from_config(&endpoints, None);
327        assert_eq!(registry.len(), 1);
328        let ep = registry.get("default").unwrap();
329        assert_eq!(ep.endpoint_type, EndpointType::Llm);
330    }
331
332    #[test]
333    fn test_registry_resolve_by_name() {
334        let mut endpoints = HashMap::new();
335        endpoints.insert(
336            "vision".to_owned(),
337            EndpointConfig {
338                endpoint_type: EndpointType::Llm,
339                url: Some("https://api.openai.com".into()),
340                model: Some("gpt-4o".into()),
341                ..Default::default()
342            },
343        );
344
345        let registry = EndpointRegistry::from_config(&endpoints, None);
346        let ep = registry.get("vision");
347        assert!(ep.is_some());
348        assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
349    }
350
351    #[test]
352    fn test_resolve_for_task_with_default() {
353        let mut endpoints = HashMap::new();
354        let ec = EndpointConfig {
355            endpoint_type: EndpointType::Llm,
356            url: Some("http://localhost:8080".into()),
357            model: Some("deepseek".into()),
358            default_for: vec!["targeting".to_owned()],
359            ..Default::default()
360        };
361        endpoints.insert("main".to_owned(), ec);
362
363        let registry = EndpointRegistry::from_config(&endpoints, None);
364        let ep = registry.resolve_for_task(TaskType::Targeting);
365        assert_eq!(ep.name, "main");
366    }
367
368    #[test]
369    fn test_resolve_explicit_overrides_task() {
370        let mut endpoints = HashMap::new();
371        endpoints.insert(
372            "default".to_owned(),
373            EndpointConfig {
374                endpoint_type: EndpointType::Llm,
375                url: Some("http://default".into()),
376                default_for: vec!["targeting".to_owned()],
377                ..Default::default()
378            },
379        );
380        endpoints.insert(
381            "fast".to_owned(),
382            EndpointConfig {
383                endpoint_type: EndpointType::Llm,
384                url: Some("http://fast".into()),
385                ..Default::default()
386            },
387        );
388
389        let registry = EndpointRegistry::from_config(&endpoints, None);
390        let ep = registry.resolve(Some("fast"), TaskType::Targeting);
391        assert_eq!(ep.name, "fast");
392    }
393
394    #[test]
395    fn test_resolve_chain_follows_fallbacks() {
396        let mut endpoints = HashMap::new();
397        endpoints.insert(
398            "default".to_owned(),
399            EndpointConfig {
400                endpoint_type: EndpointType::Llm,
401                url: Some("http://default".into()),
402                default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
403                fallbacks: vec!["pro".to_owned()],
404                ..Default::default()
405            },
406        );
407        endpoints.insert(
408            "pro".to_owned(),
409            EndpointConfig {
410                endpoint_type: EndpointType::Llm,
411                url: Some("http://pro".into()),
412                model: Some("gpt-4.1".into()),
413                ..Default::default()
414            },
415        );
416
417        let registry = EndpointRegistry::from_config(&endpoints, None);
418        let chain = registry.resolve_chain(None, TaskType::Assertion);
419        assert_eq!(chain.len(), 2);
420        assert_eq!(chain[0].name, "default");
421        assert_eq!(chain[1].name, "pro");
422    }
423
424    #[test]
425    fn test_resolve_chain_skips_non_llm_and_cycles() {
426        let mut endpoints = HashMap::new();
427        endpoints.insert(
428            "default".to_owned(),
429            EndpointConfig {
430                endpoint_type: EndpointType::Llm,
431                url: Some("http://default".into()),
432                default_for: vec!["assertion".to_owned()],
433                fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
434                ..Default::default()
435            },
436        );
437        // mcp1 is not an LLM endpoint — must be skipped in the chain.
438        endpoints.insert(
439            "mcp1".to_owned(),
440            EndpointConfig {
441                endpoint_type: EndpointType::Mcp,
442                command: Some("npx".into()),
443                ..Default::default()
444            },
445        );
446        // Cycle: pro -> default must terminate.
447        endpoints.insert(
448            "pro".to_owned(),
449            EndpointConfig {
450                endpoint_type: EndpointType::Llm,
451                url: Some("http://pro".into()),
452                fallbacks: vec!["default".to_owned()],
453                ..Default::default()
454            },
455        );
456
457        let registry = EndpointRegistry::from_config(&endpoints, None);
458        let chain = registry.resolve_chain(None, TaskType::Assertion);
459        assert_eq!(chain.len(), 2);
460        assert_eq!(chain[0].name, "default");
461        assert_eq!(chain[1].name, "pro");
462    }
463
464    #[test]
465    fn test_resolve_chain_max_attempts_default() {
466        let mut endpoints = HashMap::new();
467        endpoints.insert(
468            "default".to_owned(),
469            EndpointConfig {
470                endpoint_type: EndpointType::Llm,
471                url: Some("http://default".into()),
472                max_attempts: Some(7),
473                default_for: vec!["assertion".to_owned()],
474                ..Default::default()
475            },
476        );
477        let registry = EndpointRegistry::from_config(&endpoints, None);
478        let ep = registry.resolve_chain(None, TaskType::Assertion);
479        assert_eq!(ep[0].max_attempts, 7);
480    }
481
482    #[test]
483    fn test_default_llm_has_env_values() {
484        let ep = ResolvedEndpoint::default_llm();
485        assert_eq!(ep.endpoint_type, EndpointType::Llm);
486        assert!(ep.model.is_some());
487        assert!(!ep.url.is_empty());
488    }
489}