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    /// Input token pricing per 1M tokens.
49    pub input_price_per_1m: f64,
50    /// Output token pricing per 1M tokens.
51    pub output_price_per_1m: f64,
52    /// Flat cost per call.
53    pub per_call_price: f64,
54}
55
56impl ResolvedEndpoint {
57    /// Creates a default LLM endpoint from environment variables.
58    #[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/// Registry of all configured endpoints with routing logic.
77#[derive(Debug, Clone)]
78pub struct EndpointRegistry {
79    endpoints: HashMap<String, ResolvedEndpoint>,
80    default_for: HashMap<String, String>,
81}
82
83impl EndpointRegistry {
84    /// Builds a registry from the endpoint definitions in scenario config.
85    ///
86    /// Falls back to a default LLM endpoint derived from `fallback_llm`
87    /// (the runner's effective config — CLI arguments merged over env vars
88    /// and scenario fields) when no `[config.endpoints]` are defined, so
89    /// `--llm-url` / `--llm-model` / `--llm-api-key` are honored even for
90    /// scenarios without an explicit endpoint table. Without a fallback,
91    /// environment variables are used.
92    #[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    /// Resolves an endpoint by explicit name.
166    ///
167    /// Returns `None` if no endpoint with the given name exists.
168    #[must_use]
169    pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
170        self.endpoints.get(name)
171    }
172
173    /// Resolves the best endpoint for a given task type.
174    ///
175    /// Checks for a `default_for` mapping first, then falls back to any LLM
176    /// endpoint, then panics (config error).
177    #[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        // Fallback: first LLM endpoint
186        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    /// Resolves an endpoint: explicit name takes priority, then task-type
193    /// routing, then first LLM endpoint.
194    #[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    /// Returns the number of configured endpoints.
205    #[must_use]
206    pub fn len(&self) -> usize {
207        self.endpoints.len()
208    }
209
210    /// Returns true if no endpoints are configured.
211    #[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}