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::PricingConfig;
10use crate::scenario::Provider;
11
12/// Resolves the cache price per 1M tokens: an explicit price when set,
13/// otherwise the input price scaled by the read/write multiplier (defaults
14/// 0.1× read and 1.25× write).
15fn cache_price(pricing: Option<&PricingConfig>, input_price: f64, read: bool) -> f64 {
16    let Some(p) = pricing else {
17        return 0.0;
18    };
19    let explicit = if read {
20        p.cached_input_per_1m_tokens
21    } else {
22        p.cache_write_per_1m_tokens
23    };
24    explicit.unwrap_or_else(|| {
25        let multiplier = if read {
26            p.cache_read_multiplier.unwrap_or(0.1)
27        } else {
28            p.cache_write_multiplier.unwrap_or(1.25)
29        };
30        input_price * multiplier
31    })
32}
33
34/// Classification of a task for endpoint routing.
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
36pub enum TaskType {
37    /// LLM-based element targeting (resolving CSS selectors from natural
38    /// language).
39    Targeting,
40    /// LLM-based assertion evaluation.
41    Assertion,
42}
43
44impl TaskType {
45    /// Returns the routing key string for this task type.
46    #[must_use]
47    pub const fn as_str(self) -> &'static str {
48        match self {
49            Self::Targeting => "targeting",
50            Self::Assertion => "assertion",
51        }
52    }
53}
54
55/// Resolved endpoint ready for use in calls.
56#[derive(Debug, Clone)]
57pub struct ResolvedEndpoint {
58    /// Endpoint name.
59    pub name: String,
60    /// Endpoint type.
61    pub endpoint_type: EndpointType,
62    /// Base URL for HTTP-based endpoints.
63    pub url: String,
64    /// Model name (LLM endpoints only).
65    pub model: Option<String>,
66    /// API key / bearer token.
67    pub api_key: Option<String>,
68    /// Custom HTTP headers.
69    pub headers: HashMap<String, String>,
70    /// Command for MCP subprocess endpoints.
71    pub command: Option<String>,
72    /// Arguments for MCP subprocess endpoints.
73    pub args: Vec<String>,
74    /// Whether this endpoint accepts image parts (vision).
75    pub vision: bool,
76    /// Input token pricing per 1M tokens.
77    pub input_price_per_1m: f64,
78    /// Output token pricing per 1M tokens.
79    pub output_price_per_1m: f64,
80    /// Cached-input (prompt cache read) pricing per 1M tokens.
81    pub cached_input_price_per_1m: f64,
82    /// Cache-write (cache creation) input pricing per 1M tokens.
83    pub cache_write_price_per_1m: f64,
84    /// Bill cache reads/writes at their cache rates (default `true`).
85    pub cache_pricing: bool,
86    /// Send provider-side prompt-cache markers (default `true`).
87    pub cache_markers: bool,
88    /// Flat cost per call.
89    pub per_call_price: f64,
90    /// Retry budget for a single chat completion (default 3; global env
91    /// override `HARNESS_LLM_CALL_ATTEMPTS`).
92    pub max_attempts: u32,
93    /// Ordered names of fallback endpoints tried when this endpoint
94    /// exhausts its attempts (LLM endpoints only).
95    pub fallbacks: Vec<String>,
96    /// LLM provider protocol.
97    pub provider: Provider,
98    /// `Azure` `OpenAI` deployment name (provider `azure`).
99    pub deployment: Option<String>,
100    /// `Azure` `OpenAI` API version (provider `azure`).
101    pub api_version: Option<String>,
102    /// Authentication configuration.
103    pub auth: AuthConfig,
104    /// Headers produced by running a command per call.
105    pub header_commands: HashMap<String, String>,
106    /// AWS credential settings (provider `bedrock`).
107    pub aws: AwsConfig,
108}
109
110impl ResolvedEndpoint {
111    /// Creates a default LLM endpoint from environment variables.
112    #[must_use]
113    pub fn default_llm() -> Self {
114        Self {
115            name: "default".to_owned(),
116            endpoint_type: EndpointType::Llm,
117            url: crate::llm_base_url(),
118            model: Some(crate::llm_model()),
119            api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
120            headers: crate::parse_headers_env(),
121            command: None,
122            args: Vec::new(),
123            vision: false,
124            input_price_per_1m: 0.0,
125            output_price_per_1m: 0.0,
126            cached_input_price_per_1m: 0.0,
127            cache_write_price_per_1m: 0.0,
128            cache_pricing: true,
129            cache_markers: true,
130            per_call_price: 0.0,
131            max_attempts: crate::default_llm_attempts(),
132            fallbacks: Vec::new(),
133            provider: Provider::Openai,
134            deployment: None,
135            api_version: None,
136            auth: AuthConfig::default(),
137            header_commands: HashMap::new(),
138            aws: AwsConfig::default(),
139        }
140    }
141}
142
143/// Global defaults applied to endpoints that do not override them.
144#[derive(Debug, Clone, Copy, Default)]
145pub struct EndpointDefaults {
146    /// Default for provider-side prompt-cache markers.
147    pub cache: Option<bool>,
148    /// Default for cache-aware cost pricing.
149    pub cache_pricing: Option<bool>,
150}
151
152/// Registry of all configured endpoints with routing logic.
153#[derive(Debug, Clone)]
154pub struct EndpointRegistry {
155    endpoints: HashMap<String, ResolvedEndpoint>,
156    default_for: HashMap<String, String>,
157}
158
159impl EndpointRegistry {
160    /// Builds a registry from the endpoint definitions in scenario config.
161    ///
162    /// Falls back to a default LLM endpoint derived from `fallback_llm`
163    /// (the runner's effective config — CLI arguments merged over env vars
164    /// and scenario fields) when no `[config.endpoints]` are defined, so
165    /// `--llm-url` / `--llm-model` / `--llm-api-key` are honored even for
166    /// scenarios without an explicit endpoint table. Without a fallback,
167    /// environment variables are used.
168    #[must_use]
169    #[allow(clippy::too_many_lines)]
170    pub fn from_config(
171        endpoints: &HashMap<String, EndpointConfig>,
172        fallback_llm: Option<&crate::LlmConfig>,
173        defaults: EndpointDefaults,
174    ) -> Self {
175        if endpoints.is_empty() {
176            let default_llm =
177                fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
178                    name: "default".to_owned(),
179                    endpoint_type: EndpointType::Llm,
180                    url: llm.url.clone(),
181                    model: Some(llm.model.clone()),
182                    api_key: llm.api_key.clone(),
183                    headers: llm.headers.clone(),
184                    command: None,
185                    args: Vec::new(),
186                    vision: false,
187                    input_price_per_1m: 0.0,
188                    output_price_per_1m: 0.0,
189                    cached_input_price_per_1m: 0.0,
190                    cache_write_price_per_1m: 0.0,
191                    cache_pricing: defaults.cache_pricing.unwrap_or(true),
192                    cache_markers: defaults.cache.unwrap_or(true),
193                    per_call_price: 0.0,
194                    max_attempts: llm.max_attempts,
195                    fallbacks: Vec::new(),
196                    provider: llm.provider,
197                    deployment: llm.deployment.clone(),
198                    api_version: llm.api_version.clone(),
199                    auth: llm.auth.clone(),
200                    header_commands: llm.header_commands.clone(),
201                    aws: llm.aws.clone(),
202                });
203            let mut map = HashMap::new();
204            let mut default_for = HashMap::new();
205            for tt in &[TaskType::Targeting, TaskType::Assertion] {
206                default_for.insert(tt.as_str().to_owned(), "default".to_owned());
207            }
208            map.insert("default".to_owned(), default_llm);
209            return Self {
210                endpoints: map,
211                default_for,
212            };
213        }
214
215        let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
216        let mut default_for: HashMap<String, String> = HashMap::new();
217
218        for (name, ec) in endpoints {
219            let re = ResolvedEndpoint {
220                name: name.clone(),
221                endpoint_type: ec.endpoint_type.clone(),
222                url: ec
223                    .url
224                    .clone()
225                    .unwrap_or_else(|| {
226                        if ec.provider == Provider::Bedrock {
227                            // Bedrock builds its URL from region + model
228                            // (https://bedrock-runtime.<region>.amazonaws.com/
229                            // model/<model>/converse). Substituting the
230                            // OpenAI-compatible base URL here would send the
231                            // SigV4-signed Converse request to the wrong host
232                            // (observed: the exo gateway's FastAPI root answers
233                            // 405 {"detail":"Method Not Allowed"}).
234                            String::new()
235                        } else {
236                            match ec.endpoint_type {
237                                EndpointType::Llm => crate::llm_base_url(),
238                                EndpointType::A2a | EndpointType::Mcp => String::new(),
239                            }
240                        }
241                    })
242                    .trim_end_matches('/')
243                    .to_owned(),
244                model: ec.model.clone(),
245                api_key: ec.api_key.clone(),
246                headers: ec.headers.clone(),
247                command: ec.command.clone(),
248                args: ec.args.clone(),
249                vision: ec.vision,
250                input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
251                output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
252                cached_input_price_per_1m: cache_price(
253                    ec.pricing.as_ref(),
254                    ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
255                    true,
256                ),
257                cache_write_price_per_1m: cache_price(
258                    ec.pricing.as_ref(),
259                    ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
260                    false,
261                ),
262                cache_pricing: ec
263                    .pricing
264                    .as_ref()
265                    .and_then(|p| p.cache_pricing)
266                    .or(defaults.cache_pricing)
267                    .unwrap_or(true),
268                cache_markers: ec.cache.or(defaults.cache).unwrap_or(true),
269                per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
270                max_attempts: ec.max_attempts.unwrap_or_else(crate::default_llm_attempts),
271                fallbacks: ec.fallbacks.clone(),
272                provider: ec.provider,
273                deployment: ec.deployment.clone(),
274                api_version: ec.api_version.clone(),
275                auth: ec.auth.clone(),
276                header_commands: ec.header_commands.clone(),
277                aws: ec.aws.clone(),
278            };
279
280            for df in &ec.default_for {
281                default_for.insert(df.clone(), name.clone());
282            }
283
284            resolved.insert(name.clone(), re);
285        }
286
287        Self {
288            endpoints: resolved,
289            default_for,
290        }
291    }
292
293    /// Resolves an endpoint by explicit name.
294    ///
295    /// Returns `None` if no endpoint with the given name exists.
296    #[must_use]
297    pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
298        self.endpoints.get(name)
299    }
300
301    /// Resolves the best endpoint for a given task type.
302    ///
303    /// Checks for a `default_for` mapping first, then falls back to any LLM
304    /// endpoint, then panics (config error).
305    #[must_use]
306    pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
307        let key = task.as_str();
308        if let Some(name) = self.default_for.get(key) {
309            if let Some(ep) = self.endpoints.get(name) {
310                return ep;
311            }
312        }
313        // Fallback: first LLM endpoint
314        self.endpoints
315            .values()
316            .find(|ep| ep.endpoint_type == EndpointType::Llm)
317            .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
318    }
319
320    /// Resolves an endpoint: explicit name takes priority, then task-type
321    /// routing, then first LLM endpoint.
322    #[must_use]
323    pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
324        if let Some(n) = name {
325            if let Some(ep) = self.endpoints.get(n) {
326                return ep;
327            }
328        }
329        self.resolve_for_task(task)
330    }
331
332    /// Resolves the ordered call chain for a task: the primary endpoint
333    /// followed by its `fallbacks` (LLM endpoints only, deduplicated,
334    /// cycle-guarded, max 8 hops). Every LLM call goes through this chain —
335    /// the primary endpoint gets its own `max_attempts` retry budget, then
336    /// each fallback in turn, until one answers.
337    ///
338    /// Example: a cheap primary (`default`) with a more powerful fallback
339    /// (`pro`) can be declared as
340    /// `fallbacks = ["pro"]` on the `default` endpoint.
341    #[must_use]
342    pub fn resolve_chain(&self, name: Option<&str>, task: TaskType) -> Vec<&ResolvedEndpoint> {
343        let primary = self.resolve(name, task);
344        let mut chain: Vec<&ResolvedEndpoint> = vec![primary];
345        let mut seen: std::collections::HashSet<&str> =
346            std::collections::HashSet::from([primary.name.as_str()]);
347        let mut cursor = primary;
348        for _ in 0..8 {
349            let next = cursor.fallbacks.iter().find_map(|fb| {
350                let ep = self.endpoints.get(fb)?;
351                (ep.endpoint_type == EndpointType::Llm && !seen.contains(ep.name.as_str()))
352                    .then_some(ep)
353            });
354            match next {
355                Some(ep) => {
356                    seen.insert(ep.name.as_str());
357                    chain.push(ep);
358                    cursor = ep;
359                }
360                None => break,
361            }
362        }
363        chain
364    }
365
366    /// Returns the number of configured endpoints.
367    #[must_use]
368    pub fn len(&self) -> usize {
369        self.endpoints.len()
370    }
371
372    /// Returns true if no endpoints are configured.
373    #[must_use]
374    pub fn is_empty(&self) -> bool {
375        self.endpoints.is_empty()
376    }
377}
378
379#[cfg(test)]
380mod tests {
381    use super::*;
382    use std::collections::HashMap;
383
384    #[test]
385    fn test_task_type_as_str() {
386        assert_eq!(TaskType::Targeting.as_str(), "targeting");
387        assert_eq!(TaskType::Assertion.as_str(), "assertion");
388    }
389
390    #[test]
391    fn test_registry_empty_config() {
392        let endpoints = HashMap::new();
393        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
394        assert_eq!(registry.len(), 1);
395        let ep = registry.get("default").unwrap();
396        assert_eq!(ep.endpoint_type, EndpointType::Llm);
397    }
398
399    #[test]
400    fn test_registry_resolve_by_name() {
401        let mut endpoints = HashMap::new();
402        endpoints.insert(
403            "vision".to_owned(),
404            EndpointConfig {
405                endpoint_type: EndpointType::Llm,
406                url: Some("https://api.openai.com".into()),
407                model: Some("gpt-4o".into()),
408                ..Default::default()
409            },
410        );
411
412        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
413        let ep = registry.get("vision");
414        assert!(ep.is_some());
415        assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
416    }
417
418    #[test]
419    fn test_resolve_for_task_with_default() {
420        let mut endpoints = HashMap::new();
421        let ec = EndpointConfig {
422            endpoint_type: EndpointType::Llm,
423            url: Some("http://localhost:8080".into()),
424            model: Some("deepseek".into()),
425            default_for: vec!["targeting".to_owned()],
426            ..Default::default()
427        };
428        endpoints.insert("main".to_owned(), ec);
429
430        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
431        let ep = registry.resolve_for_task(TaskType::Targeting);
432        assert_eq!(ep.name, "main");
433    }
434
435    #[test]
436    fn test_resolve_explicit_overrides_task() {
437        let mut endpoints = HashMap::new();
438        endpoints.insert(
439            "default".to_owned(),
440            EndpointConfig {
441                endpoint_type: EndpointType::Llm,
442                url: Some("http://default".into()),
443                default_for: vec!["targeting".to_owned()],
444                ..Default::default()
445            },
446        );
447        endpoints.insert(
448            "fast".to_owned(),
449            EndpointConfig {
450                endpoint_type: EndpointType::Llm,
451                url: Some("http://fast".into()),
452                ..Default::default()
453            },
454        );
455
456        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
457        let ep = registry.resolve(Some("fast"), TaskType::Targeting);
458        assert_eq!(ep.name, "fast");
459    }
460
461    #[test]
462    fn test_resolve_chain_follows_fallbacks() {
463        let mut endpoints = HashMap::new();
464        endpoints.insert(
465            "default".to_owned(),
466            EndpointConfig {
467                endpoint_type: EndpointType::Llm,
468                url: Some("http://default".into()),
469                default_for: vec!["targeting".to_owned(), "assertion".to_owned()],
470                fallbacks: vec!["pro".to_owned()],
471                ..Default::default()
472            },
473        );
474        endpoints.insert(
475            "pro".to_owned(),
476            EndpointConfig {
477                endpoint_type: EndpointType::Llm,
478                url: Some("http://pro".into()),
479                model: Some("gpt-4.1".into()),
480                ..Default::default()
481            },
482        );
483
484        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
485        let chain = registry.resolve_chain(None, TaskType::Assertion);
486        assert_eq!(chain.len(), 2);
487        assert_eq!(chain[0].name, "default");
488        assert_eq!(chain[1].name, "pro");
489    }
490
491    #[test]
492    fn test_resolve_chain_skips_non_llm_and_cycles() {
493        let mut endpoints = HashMap::new();
494        endpoints.insert(
495            "default".to_owned(),
496            EndpointConfig {
497                endpoint_type: EndpointType::Llm,
498                url: Some("http://default".into()),
499                default_for: vec!["assertion".to_owned()],
500                fallbacks: vec!["mcp1".to_owned(), "pro".to_owned()],
501                ..Default::default()
502            },
503        );
504        // mcp1 is not an LLM endpoint — must be skipped in the chain.
505        endpoints.insert(
506            "mcp1".to_owned(),
507            EndpointConfig {
508                endpoint_type: EndpointType::Mcp,
509                command: Some("npx".into()),
510                ..Default::default()
511            },
512        );
513        // Cycle: pro -> default must terminate.
514        endpoints.insert(
515            "pro".to_owned(),
516            EndpointConfig {
517                endpoint_type: EndpointType::Llm,
518                url: Some("http://pro".into()),
519                fallbacks: vec!["default".to_owned()],
520                ..Default::default()
521            },
522        );
523
524        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
525        let chain = registry.resolve_chain(None, TaskType::Assertion);
526        assert_eq!(chain.len(), 2);
527        assert_eq!(chain[0].name, "default");
528        assert_eq!(chain[1].name, "pro");
529    }
530
531    #[test]
532    fn test_resolve_chain_max_attempts_default() {
533        let mut endpoints = HashMap::new();
534        endpoints.insert(
535            "default".to_owned(),
536            EndpointConfig {
537                endpoint_type: EndpointType::Llm,
538                url: Some("http://default".into()),
539                max_attempts: Some(7),
540                default_for: vec!["assertion".to_owned()],
541                ..Default::default()
542            },
543        );
544        let registry = EndpointRegistry::from_config(&endpoints, None, EndpointDefaults::default());
545        let ep = registry.resolve_chain(None, TaskType::Assertion);
546        assert_eq!(ep[0].max_attempts, 7);
547    }
548
549    #[test]
550    fn test_default_llm_has_env_values() {
551        let ep = ResolvedEndpoint::default_llm();
552        assert_eq!(ep.endpoint_type, EndpointType::Llm);
553        assert!(ep.model.is_some());
554        assert_ne!(ep.url, "");
555    }
556}