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}
57
58impl ResolvedEndpoint {
59    /// Creates a default LLM endpoint from environment variables.
60    #[must_use]
61    pub fn default_llm() -> Self {
62        Self {
63            name: "default".to_owned(),
64            endpoint_type: EndpointType::Llm,
65            url: crate::llm_base_url(),
66            model: Some(crate::llm_model()),
67            api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
68            headers: crate::parse_headers_env(),
69            command: None,
70            args: Vec::new(),
71            vision: false,
72            input_price_per_1m: 0.0,
73            output_price_per_1m: 0.0,
74            per_call_price: 0.0,
75        }
76    }
77}
78
79/// Registry of all configured endpoints with routing logic.
80#[derive(Debug, Clone)]
81pub struct EndpointRegistry {
82    endpoints: HashMap<String, ResolvedEndpoint>,
83    default_for: HashMap<String, String>,
84}
85
86impl EndpointRegistry {
87    /// Builds a registry from the endpoint definitions in scenario config.
88    ///
89    /// Falls back to a default LLM endpoint derived from `fallback_llm`
90    /// (the runner's effective config — CLI arguments merged over env vars
91    /// and scenario fields) when no `[config.endpoints]` are defined, so
92    /// `--llm-url` / `--llm-model` / `--llm-api-key` are honored even for
93    /// scenarios without an explicit endpoint table. Without a fallback,
94    /// environment variables are used.
95    #[must_use]
96    pub fn from_config(
97        endpoints: &HashMap<String, EndpointConfig>,
98        fallback_llm: Option<&crate::LlmConfig>,
99    ) -> Self {
100        if endpoints.is_empty() {
101            let default_llm =
102                fallback_llm.map_or_else(ResolvedEndpoint::default_llm, |llm| ResolvedEndpoint {
103                    name: "default".to_owned(),
104                    endpoint_type: EndpointType::Llm,
105                    url: llm.url.clone(),
106                    model: Some(llm.model.clone()),
107                    api_key: llm.api_key.clone(),
108                    headers: llm.headers.clone(),
109                    command: None,
110                    args: Vec::new(),
111                    vision: false,
112                    input_price_per_1m: 0.0,
113                    output_price_per_1m: 0.0,
114                    per_call_price: 0.0,
115                });
116            let mut map = HashMap::new();
117            let mut default_for = HashMap::new();
118            for tt in &[TaskType::Targeting, TaskType::Assertion] {
119                default_for.insert(tt.as_str().to_owned(), "default".to_owned());
120            }
121            map.insert("default".to_owned(), default_llm);
122            return Self {
123                endpoints: map,
124                default_for,
125            };
126        }
127
128        let mut resolved: HashMap<String, ResolvedEndpoint> = HashMap::new();
129        let mut default_for: HashMap<String, String> = HashMap::new();
130
131        for (name, ec) in endpoints {
132            let re = ResolvedEndpoint {
133                name: name.clone(),
134                endpoint_type: ec.endpoint_type.clone(),
135                url: ec
136                    .url
137                    .clone()
138                    .unwrap_or_else(|| match ec.endpoint_type {
139                        EndpointType::Llm => crate::llm_base_url(),
140                        EndpointType::A2a | EndpointType::Mcp => String::new(),
141                    })
142                    .trim_end_matches('/')
143                    .to_owned(),
144                model: ec.model.clone(),
145                api_key: ec.api_key.clone(),
146                headers: ec.headers.clone(),
147                command: ec.command.clone(),
148                args: ec.args.clone(),
149                vision: ec.vision,
150                input_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.input_per_1m_tokens),
151                output_price_per_1m: ec.pricing.as_ref().map_or(0.0, |p| p.output_per_1m_tokens),
152                per_call_price: ec.pricing.as_ref().map_or(0.0, |p| p.per_call),
153            };
154
155            for df in &ec.default_for {
156                default_for.insert(df.clone(), name.clone());
157            }
158
159            resolved.insert(name.clone(), re);
160        }
161
162        Self {
163            endpoints: resolved,
164            default_for,
165        }
166    }
167
168    /// Resolves an endpoint by explicit name.
169    ///
170    /// Returns `None` if no endpoint with the given name exists.
171    #[must_use]
172    pub fn get(&self, name: &str) -> Option<&ResolvedEndpoint> {
173        self.endpoints.get(name)
174    }
175
176    /// Resolves the best endpoint for a given task type.
177    ///
178    /// Checks for a `default_for` mapping first, then falls back to any LLM
179    /// endpoint, then panics (config error).
180    #[must_use]
181    pub fn resolve_for_task(&self, task: TaskType) -> &ResolvedEndpoint {
182        let key = task.as_str();
183        if let Some(name) = self.default_for.get(key) {
184            if let Some(ep) = self.endpoints.get(name) {
185                return ep;
186            }
187        }
188        // Fallback: first LLM endpoint
189        self.endpoints
190            .values()
191            .find(|ep| ep.endpoint_type == EndpointType::Llm)
192            .unwrap_or_else(|| panic!("no LLM endpoint configured for task {key}"))
193    }
194
195    /// Resolves an endpoint: explicit name takes priority, then task-type
196    /// routing, then first LLM endpoint.
197    #[must_use]
198    pub fn resolve(&self, name: Option<&str>, task: TaskType) -> &ResolvedEndpoint {
199        if let Some(n) = name {
200            if let Some(ep) = self.endpoints.get(n) {
201                return ep;
202            }
203        }
204        self.resolve_for_task(task)
205    }
206
207    /// Returns the number of configured endpoints.
208    #[must_use]
209    pub fn len(&self) -> usize {
210        self.endpoints.len()
211    }
212
213    /// Returns true if no endpoints are configured.
214    #[must_use]
215    pub fn is_empty(&self) -> bool {
216        self.endpoints.is_empty()
217    }
218}
219
220#[cfg(test)]
221mod tests {
222    use super::*;
223    use std::collections::HashMap;
224
225    #[test]
226    fn test_task_type_as_str() {
227        assert_eq!(TaskType::Targeting.as_str(), "targeting");
228        assert_eq!(TaskType::Assertion.as_str(), "assertion");
229    }
230
231    #[test]
232    fn test_registry_empty_config() {
233        let endpoints = HashMap::new();
234        let registry = EndpointRegistry::from_config(&endpoints, None);
235        assert_eq!(registry.len(), 1);
236        let ep = registry.get("default").unwrap();
237        assert_eq!(ep.endpoint_type, EndpointType::Llm);
238    }
239
240    #[test]
241    fn test_registry_resolve_by_name() {
242        let mut endpoints = HashMap::new();
243        endpoints.insert(
244            "vision".to_owned(),
245            EndpointConfig {
246                endpoint_type: EndpointType::Llm,
247                url: Some("https://api.openai.com".into()),
248                model: Some("gpt-4o".into()),
249                ..Default::default()
250            },
251        );
252
253        let registry = EndpointRegistry::from_config(&endpoints, None);
254        let ep = registry.get("vision");
255        assert!(ep.is_some());
256        assert_eq!(ep.unwrap().model.as_deref(), Some("gpt-4o"));
257    }
258
259    #[test]
260    fn test_resolve_for_task_with_default() {
261        let mut endpoints = HashMap::new();
262        let ec = EndpointConfig {
263            endpoint_type: EndpointType::Llm,
264            url: Some("http://localhost:8080".into()),
265            model: Some("deepseek".into()),
266            default_for: vec!["targeting".to_owned()],
267            ..Default::default()
268        };
269        endpoints.insert("main".to_owned(), ec);
270
271        let registry = EndpointRegistry::from_config(&endpoints, None);
272        let ep = registry.resolve_for_task(TaskType::Targeting);
273        assert_eq!(ep.name, "main");
274    }
275
276    #[test]
277    fn test_resolve_explicit_overrides_task() {
278        let mut endpoints = HashMap::new();
279        endpoints.insert(
280            "default".to_owned(),
281            EndpointConfig {
282                endpoint_type: EndpointType::Llm,
283                url: Some("http://default".into()),
284                default_for: vec!["targeting".to_owned()],
285                ..Default::default()
286            },
287        );
288        endpoints.insert(
289            "fast".to_owned(),
290            EndpointConfig {
291                endpoint_type: EndpointType::Llm,
292                url: Some("http://fast".into()),
293                ..Default::default()
294            },
295        );
296
297        let registry = EndpointRegistry::from_config(&endpoints, None);
298        let ep = registry.resolve(Some("fast"), TaskType::Targeting);
299        assert_eq!(ep.name, "fast");
300    }
301
302    #[test]
303    fn test_default_llm_has_env_values() {
304        let ep = ResolvedEndpoint::default_llm();
305        assert_eq!(ep.endpoint_type, EndpointType::Llm);
306        assert!(ep.model.is_some());
307        assert!(!ep.url.is_empty());
308    }
309}