Skip to main content

llm_browser_testkit/
lib.rs

1//! LLM-driven browser test framework.
2//!
3//! Provides reusable building blocks for browser-based test scenarios:
4//! - Browser client via Chrome `DevTools` Protocol (headless)
5//! - LLM client for natural language element targeting and assertions
6//! - A2A agent integration for agent-to-agent communication
7//! - MCP client/server integration for tool-calling
8//! - Cost tracking, token counting, and budget enforcement
9//! - Declarative TOML scenario runner
10//! - `#[browser_test]` macros for `cargo test` integration
11
12#![allow(
13    clippy::expect_used,
14    clippy::unwrap_used,
15    clippy::panic,
16    clippy::missing_panics_doc
17)]
18
19/// A2A agent protocol client.
20pub mod a2a;
21/// Auth token resolution — Entra ID and token commands for LLM endpoints.
22mod auth;
23/// AWS Bedrock provider (`SigV4`-signed Converse calls; feature `aws`).
24#[cfg(feature = "aws")]
25mod bedrock;
26/// Budget tracking and enforcement.
27pub mod budgets;
28/// Cost calculation, usage tracking, and pricing.
29pub mod costs;
30/// Failure diagnostics — page-state capture and artifact writing.
31pub mod diagnostics;
32/// Endpoint registry and routing resolver.
33pub mod endpoints;
34/// Typed run events emitted by the runner.
35pub mod events;
36/// MCP client for connecting to external MCP servers.
37pub mod mcp_client;
38/// MCP server for exposing the framework as an MCP server.
39pub mod mcp_server;
40/// Concurrent execution of multiple scenario files across isolated browsers.
41pub mod parallel;
42pub mod pricing;
43/// Secret redaction for every report sink.
44pub mod redact;
45/// Run reporting: console, NDJSON, JUnit, GitHub and Perfetto sinks.
46pub mod reporting;
47/// Step-by-step scenario executor (navigate, click, type, wait, assert).
48pub mod runner;
49/// Declarative TOML-based test scenario types.
50pub mod scenario;
51/// CSS selector sanitization for LLM-generated selectors.
52pub mod selectors;
53/// Vision support — screenshot capture/downscale/encode for visual asserts.
54pub mod vision;
55
56/// A2A agent server for accepting agent tasks.
57#[cfg(feature = "a2a-server")]
58pub mod a2a_server;
59
60/// `#[browser_test]` macros for `cargo test` integration.
61#[cfg(feature = "macros")]
62pub mod macros;
63
64use std::collections::HashMap;
65use std::time::Duration;
66
67use serde_json::{json, Value};
68
69pub use costs::LlmResponse;
70pub use costs::LlmUsage;
71pub use scenario::AuthConfig;
72pub use scenario::AuthMode;
73pub use scenario::AwsConfig;
74pub use scenario::Provider;
75
76/// Configuration for the LLM client — bundles URL, model, auth, timeouts,
77/// and provider-specific options into a single struct passed everywhere.
78#[derive(Debug, Clone)]
79pub struct LlmConfig {
80    /// OpenAI-compatible API base URL (without trailing `/v1/…`).
81    pub url: String,
82    /// Model name (e.g. `gpt-4o-mini`, `deepseek`).
83    pub model: String,
84    /// API key sent as `Authorization: Bearer <key>`.
85    pub api_key: Option<String>,
86    /// Custom headers appended to every LLM request.
87    pub headers: HashMap<String, String>,
88    /// HTTP timeout.
89    pub timeout: Duration,
90    /// Sampling temperature (0.0–1.0).
91    pub temperature: f64,
92    /// Enable extended thinking / reasoning tokens.
93    /// `None` = don't send any thinking key (provider default).
94    pub thinking: Option<bool>,
95    /// Provider-specific parameters merged into the request body
96    /// (e.g. `effort = "high"` for Anthropic).
97    pub model_params: HashMap<String, Value>,
98    /// Send provider-side prompt-cache markers (default `true`). Applies to
99    /// providers that need explicit markers: AWS Bedrock (`cachePoint`) and
100    /// Anthropic-style OpenAI-compatible models (`cache_control`). Set to
101    /// `false` to disable.
102    pub cache: bool,
103    /// How many times a single call to this endpoint is retried on
104    /// transient failures before giving up (or moving to the next fallback
105    /// endpoint). Default 3; override globally with
106    /// `HARNESS_LLM_CALL_ATTEMPTS`.
107    pub max_attempts: u32,
108    /// LLM provider protocol (defaults to OpenAI-compatible).
109    ///
110    /// `azure` switches to the `Azure` `OpenAI` deployments endpoint, `bedrock`
111    /// to the `SigV4`-signed AWS Bedrock Converse API (feature `aws`).
112    pub provider: Provider,
113    /// `Azure` `OpenAI` deployment name (`Provider::Azure`). Defaults to
114    /// `model` when unset.
115    pub deployment: Option<String>,
116    /// `Azure` `OpenAI` API version (`Provider::Azure`). Defaults to
117    /// `2024-10-21`.
118    pub api_version: Option<String>,
119    /// Authentication configuration (API key, token command, Entra ID).
120    pub auth: AuthConfig,
121    /// Extra HTTP headers produced by running a command per call, keyed by
122    /// header name. Provider-agnostic.
123    pub header_commands: HashMap<String, String>,
124    /// AWS credential settings (`Provider::Bedrock`).
125    pub aws: AwsConfig,
126}
127
128impl Default for LlmConfig {
129    fn default() -> Self {
130        Self {
131            url: String::new(),
132            model: String::new(),
133            api_key: None,
134            headers: HashMap::new(),
135            timeout: Duration::from_secs(60),
136            temperature: 0.0,
137            thinking: None,
138            model_params: HashMap::new(),
139            cache: true,
140            max_attempts: default_llm_attempts(),
141            provider: Provider::Openai,
142            deployment: None,
143            api_version: None,
144            auth: AuthConfig::default(),
145            header_commands: HashMap::new(),
146            aws: AwsConfig::default(),
147        }
148    }
149}
150
151impl LlmConfig {
152    /// Build a config from environment defaults, falling back to safe
153    /// values when no env vars are set.
154    #[must_use]
155    pub fn from_env() -> Self {
156        Self {
157            url: llm_base_url(),
158            model: llm_model(),
159            api_key: std::env::var("HARNESS_LLM_API_KEY").ok(),
160            headers: parse_headers_env(),
161            timeout: Duration::from_secs(60),
162            temperature: 0.0,
163            thinking: None,
164            model_params: HashMap::new(),
165            cache: true,
166            max_attempts: default_llm_attempts(),
167            provider: Provider::Openai,
168            deployment: None,
169            api_version: None,
170            auth: AuthConfig::default(),
171            header_commands: HashMap::new(),
172            aws: AwsConfig::default(),
173        }
174    }
175}
176
177/// Default `Azure` `OpenAI` API version used when an endpoint does not set
178/// `api_version`.
179pub const DEFAULT_AZURE_API_VERSION: &str = "2024-10-21";
180
181/// Builds the `Azure` `OpenAI` chat completions URL: the resource endpoint
182/// (without `/openai`), the deployment name, and the API version.
183#[must_use]
184pub fn build_azure_url(base: &str, deployment: &str, api_version: &str) -> String {
185    let base = base.trim_end_matches('/');
186    let base = base
187        .strip_suffix("/openai")
188        .unwrap_or(base)
189        .trim_end_matches('/');
190    format!("{base}/openai/deployments/{deployment}/chat/completions?api-version={api_version}")
191}
192
193/// Reads `HARNESS_LLM_CALL_ATTEMPTS` (default 3) — how many times a single
194/// chat completion is retried before the endpoint is considered failed.
195#[must_use]
196pub fn default_llm_attempts() -> u32 {
197    std::env::var("HARNESS_LLM_CALL_ATTEMPTS")
198        .ok()
199        .and_then(|v| v.parse().ok())
200        .filter(|n| *n >= 1)
201        .unwrap_or(3)
202}
203
204/// Parses `HARNESS_LLM_HEADERS` env var (JSON object) into a header map.
205///
206/// Exposed for use by `endpoints.rs` and tests.
207#[must_use]
208pub fn parse_headers_env() -> HashMap<String, String> {
209    let Ok(raw) = std::env::var("HARNESS_LLM_HEADERS") else {
210        return HashMap::new();
211    };
212    let Ok(json) = serde_json::from_str::<Value>(&raw) else {
213        return HashMap::new();
214    };
215    let Some(obj) = json.as_object() else {
216        return HashMap::new();
217    };
218    obj.iter()
219        .filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
220        .collect()
221}
222
223/// Returns the target base URL from `HARNESS_BROWSER_BASE_URL` env,
224/// defaulting to `http://localhost:4200`.
225#[must_use]
226pub fn base_url() -> String {
227    std::env::var("HARNESS_BROWSER_BASE_URL").unwrap_or_else(|_| "http://localhost:4200".to_owned())
228}
229
230/// Returns the LLM server base URL from `HARNESS_LLM_TEST_URL` env,
231/// defaulting to `http://localhost:8080`.
232#[must_use]
233pub fn llm_base_url() -> String {
234    std::env::var("HARNESS_LLM_TEST_URL")
235        .unwrap_or_else(|_| "http://localhost:8080".to_owned())
236        .trim_end_matches('/')
237        .to_owned()
238}
239
240/// Returns the LLM model name from `HARNESS_LLM_TEST_MODEL` env,
241/// defaulting to `deepseek`.
242#[must_use]
243pub fn llm_model() -> String {
244    std::env::var("HARNESS_LLM_TEST_MODEL").unwrap_or_else(|_| "deepseek".to_owned())
245}
246
247/// Returns whether to run the browser in headless mode from
248/// `HARNESS_BROWSER_HEADLESS` env, defaulting to `true`.
249#[must_use]
250pub fn browser_headless() -> bool {
251    std::env::var("HARNESS_BROWSER_HEADLESS")
252        .map_or(true, |v| v != "0" && v.to_lowercase() != "false")
253}
254
255/// Builds a `reqwest::Client` with the given timeout.
256#[must_use]
257pub fn http_client(timeout: Duration) -> reqwest::Client {
258    reqwest::Client::builder()
259        .timeout(timeout)
260        .build()
261        .expect("build reqwest client")
262}
263
264/// Sends a chat completion request to the LLM.
265///
266/// Returns `Some(content)` on success, `None` on any error.
267///
268/// Prefer `llm_chat_with_usage` if you need token counting.
269#[must_use]
270pub async fn llm_chat(llm: &LlmConfig, system: &str, user: &str) -> Option<String> {
271    llm_chat_with_usage(llm, system, user)
272        .await
273        .map(|r| r.content)
274        .ok()
275}
276
277/// Sends a chat completion request to the LLM and returns both the content
278/// and token usage from the API response.
279///
280/// Retries transient failures (network errors, HTTP 429/5xx, invalid
281/// responses, and HTTP 200 + empty body — a gateway warm-up signature,
282/// retried with a 3s backoff) up to `llm.max_attempts` times, and returns
283/// the last underlying error instead of collapsing everything into a
284/// generic "server down" message. The error text includes the HTTP status
285/// and a truncated response-body snippet, so a gateway that answers with
286/// an HTML error page is identifiable in CI logs instead of surfacing as a
287/// bare JSON decode error. Deterministic client errors (401/403/404) are
288/// not retried.
289///
290/// # Errors
291///
292/// Returns the last underlying error as a human-readable string when every
293/// attempt fails (transport error, non-success HTTP status, response that is
294/// not valid JSON, an empty HTTP 200 body, or a response missing
295/// `choices[0].message.content`).
296pub async fn llm_chat_with_usage(
297    llm: &LlmConfig,
298    system: &str,
299    user: &str,
300) -> Result<LlmResponse, String> {
301    chat_with_retry(llm, system, user, None).await
302}
303
304/// Sends a vision-enabled chat completion request.
305///
306/// The user message carries the text prompt and one or more screenshots
307/// (JPEG/PNG data URLs, ordered from the top of the page down) as
308/// OpenAI-compatible `image_url` content parts.
309///
310/// Retries and error reporting behave like [`llm_chat_with_usage`].
311///
312/// # Errors
313///
314/// Returns the last underlying error as a human-readable string when every
315/// attempt fails (transport error, non-success HTTP status, response that is
316/// not valid JSON, or a response missing `choices[0].message.content`).
317pub async fn llm_chat_vision_with_usage(
318    llm: &LlmConfig,
319    system: &str,
320    user: &str,
321    image_data_urls: Option<&[String]>,
322) -> Result<LlmResponse, String> {
323    chat_with_retry(llm, system, user, image_data_urls).await
324}
325
326/// Calls a chain of endpoints: the primary [`LlmConfig`] first, then each
327/// fallback in order. Every endpoint gets its own `max_attempts` retry
328/// budget; the first endpoint that answers wins.
329///
330/// Returns the response together with the index of the endpoint that
331/// produced it (0 = primary, 1 = first fallback, …) so the caller can
332/// attribute cost/usage to the right endpoint.
333///
334/// # Errors
335///
336/// Returns an error naming every endpoint that failed.
337pub async fn llm_chat_with_usage_chain(
338    primary: &LlmConfig,
339    fallbacks: &[LlmConfig],
340    system: &str,
341    user: &str,
342) -> Result<(LlmResponse, usize), String> {
343    chat_chain_with_retry(primary, fallbacks, system, user, None).await
344}
345
346/// Vision variant of [`llm_chat_with_usage_chain`].
347///
348/// # Errors
349///
350/// Returns an error naming every endpoint that failed.
351pub async fn llm_chat_vision_with_usage_chain(
352    primary: &LlmConfig,
353    fallbacks: &[LlmConfig],
354    system: &str,
355    user: &str,
356    image_data_urls: Option<&[String]>,
357) -> Result<(LlmResponse, usize), String> {
358    chat_chain_with_retry(primary, fallbacks, system, user, image_data_urls).await
359}
360
361/// Shared chain loop: try each endpoint (primary then fallbacks) with its
362/// own retry budget; first success wins.
363async fn chat_chain_with_retry(
364    primary: &LlmConfig,
365    fallbacks: &[LlmConfig],
366    system: &str,
367    user: &str,
368    image_data_urls: Option<&[String]>,
369) -> Result<(LlmResponse, usize), String> {
370    let mut failures: Vec<String> = Vec::new();
371    for (i, llm) in std::iter::once(primary).chain(fallbacks.iter()).enumerate() {
372        match chat_with_retry(llm, system, user, image_data_urls).await {
373            Ok(resp) => return Ok((resp, i)),
374            Err(e) => failures.push(format!("endpoint '{}' ({:?}): {e}", llm.url, llm.model)),
375        }
376    }
377    let details = failures.iter().fold(String::new(), |mut acc, f| {
378        use std::fmt::Write as _;
379        let _ = writeln!(acc, "  - {f}");
380        acc
381    });
382    Err(format!(
383        "LLM call failed on all {} endpoint(s):\n{details}",
384        failures.len()
385    ))
386}
387
388/// Shared retry loop for text-only and vision chat completions.
389async fn chat_with_retry(
390    llm: &LlmConfig,
391    system: &str,
392    user: &str,
393    image_data_urls: Option<&[String]>,
394) -> Result<LlmResponse, String> {
395    let client = http_client(llm.timeout);
396    let mut last_err = String::from("LLM call failed");
397    let mut attempts: u32 = 0;
398
399    while attempts < llm.max_attempts {
400        attempts += 1;
401        match llm_chat_once(&client, llm, system, user, image_data_urls).await {
402            Ok(resp) => return Ok(resp),
403            Err(err) => {
404                // An HTTP 200 with an empty body is a gateway warm-up
405                // signature: give it a longer window to finish booting
406                // instead of hammering it with short retries.
407                let backoff = match &err {
408                    LlmCallError::EmptyBody { .. } => Duration::from_secs(3),
409                    _ => Duration::from_millis(500 * u64::from(attempts)),
410                };
411                last_err = err.to_string();
412                if attempts >= llm.max_attempts || !err.is_retryable() {
413                    break;
414                }
415                tokio::time::sleep(backoff).await;
416            }
417        }
418    }
419
420    Err(format!(
421        "LLM call failed after {attempts} attempt(s) (endpoint {url}): {last_err}",
422        url = llm.url
423    ))
424}
425
426/// Builds the chat messages array. Text-only messages keep the plain
427/// string `content` shape (maximum provider compatibility); vision calls
428/// use the OpenAI-compatible content-part array with one `image_url` part
429/// per screenshot (ordered from the top of the page down).
430#[must_use]
431fn build_messages(system: &str, user: &str, image_data_urls: Option<&[String]>) -> Value {
432    let user_content = image_data_urls.map_or_else(
433        || Value::String(user.to_owned()),
434        |urls| {
435            let mut parts = vec![json!({"type": "text", "text": user})];
436            for url in urls {
437                parts.push(json!({"type": "image_url", "image_url": {"url": url}}));
438            }
439            Value::Array(parts)
440        },
441    );
442    json!([
443        {"role": "system", "content": system},
444        {"role": "user", "content": user_content}
445    ])
446}
447
448/// Internal error type for a single LLM request attempt, distinguishing
449/// transient failures (worth retrying) from deterministic configuration
450/// errors (fail immediately).
451enum LlmCallError {
452    /// Transport-level failure (connect, timeout, TLS, …).
453    Transport { message: String },
454    /// Non-success HTTP status with a body snippet.
455    Http { status: u16, body: String },
456    /// Success status but the body is not valid JSON.
457    InvalidJson {
458        status: u16,
459        detail: String,
460        body: String,
461    },
462    /// Success status with an EMPTY body — the classic transient gateway
463    /// warm-up signature (HTTP 200, zero bytes). Retried with a longer
464    /// backoff than other errors.
465    EmptyBody { status: u16 },
466    /// Valid JSON but missing `choices[0].message.content`.
467    MissingContent { json: String },
468    /// Authentication failure: token command failed, Entra endpoint
469    /// answered with an error, credentials missing, or an unsupported
470    /// provider/feature combination.
471    Auth { message: String },
472}
473
474impl std::fmt::Display for LlmCallError {
475    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
476        match self {
477            Self::Transport { message } => write!(f, "LLM HTTP request failed: {message}"),
478            Self::Http { status, body } => {
479                write!(
480                    f,
481                    "LLM endpoint returned HTTP {status}: {}",
482                    truncate(body, 300)
483                )
484            }
485            Self::InvalidJson {
486                status,
487                detail,
488                body,
489            } => write!(
490                f,
491                "LLM endpoint returned HTTP {status} with non-JSON body ({detail}): {}",
492                truncate(body, 300)
493            ),
494            Self::EmptyBody { status } => write!(
495                f,
496                "LLM endpoint returned HTTP {status} with an empty response (likely gateway warm-up)"
497            ),
498            Self::MissingContent { json } => write!(
499                f,
500                "LLM response missing choices[0].message.content: {}",
501                truncate(json, 300)
502            ),
503            Self::Auth { message } => write!(f, "LLM authentication failed: {message}"),
504        }
505    }
506}
507
508impl LlmCallError {
509    /// Whether another attempt may succeed. Network errors, rate limits,
510    /// server errors, and 200-with-garbage responses can be transient on
511    /// flaky gateways; auth/not-found errors are deterministic.
512    #[must_use]
513    fn is_retryable(&self) -> bool {
514        match self {
515            Self::Transport { .. }
516            | Self::MissingContent { .. }
517            | Self::EmptyBody { .. }
518            | Self::Auth { .. } => true,
519            Self::Http { status, .. } => {
520                *status == 408 || *status == 429 || (500..600).contains(status)
521            }
522            Self::InvalidJson { status, .. } => {
523                *status == 200 || *status == 408 || *status == 429 || (500..600).contains(status)
524            }
525        }
526    }
527}
528
529/// Single LLM chat request attempt; returns the underlying error as text.
530async fn llm_chat_once(
531    client: &reqwest::Client,
532    llm: &LlmConfig,
533    system: &str,
534    user: &str,
535    image_data_urls: Option<&[String]>,
536) -> Result<LlmResponse, LlmCallError> {
537    match llm.provider {
538        Provider::Openai | Provider::Azure => {
539            chat_openai_compat_once(client, llm, system, user, image_data_urls).await
540        }
541        Provider::Bedrock => {
542            #[cfg(feature = "aws")]
543            let result =
544                crate::bedrock::chat_once(client, llm, system, user, image_data_urls).await;
545            #[cfg(not(feature = "aws"))]
546            let result = Err(LlmCallError::Auth {
547                message:
548                    "provider = \"bedrock\" requires building llm-browser-testkit with the `aws` \
549                     cargo feature"
550                        .to_owned(),
551            });
552            result
553        }
554    }
555}
556
557/// Single attempt against the OpenAI-compatible path, shared by the
558/// `openai` and `azure` providers (`Azure` differs only in the URL and the
559/// API-key header placement).
560async fn chat_openai_compat_once(
561    client: &reqwest::Client,
562    llm: &LlmConfig,
563    system: &str,
564    user: &str,
565    image_data_urls: Option<&[String]>,
566) -> Result<LlmResponse, LlmCallError> {
567    let url = match llm.provider {
568        Provider::Openai => format!("{}/v1/chat/completions", llm.url),
569        Provider::Azure => {
570            let deployment = llm.deployment.clone().unwrap_or_else(|| llm.model.clone());
571            let api_version = llm
572                .api_version
573                .clone()
574                .unwrap_or_else(|| DEFAULT_AZURE_API_VERSION.to_owned());
575            build_azure_url(&llm.url, &deployment, &api_version)
576        }
577        Provider::Bedrock => unreachable!("bedrock is dispatched before this function"),
578    };
579
580    let mut headers: Vec<(String, String)> = Vec::new();
581    if llm.auth.mode == AuthMode::ApiKey {
582        match (&llm.auth.api_key_header, &llm.api_key) {
583            // Explicit header name wins on every provider.
584            (Some(header_name), Some(key)) => {
585                headers.push((header_name.clone(), key.clone()));
586            }
587            (Some(header_name), None) => {
588                return Err(LlmCallError::Auth {
589                    message: format!(
590                        "auth.api_key_header `{header_name}` requires endpoint api_key to be set"
591                    ),
592                });
593            }
594            // Azure classic auth: the key travels in the `api-key` header.
595            (None, Some(key)) if llm.provider == Provider::Azure => {
596                headers.push(("api-key".to_owned(), key.clone()));
597            }
598            // OpenAI-compatible convention.
599            (None, Some(key)) => {
600                headers.push(("Authorization".to_owned(), format!("Bearer {key}")));
601            }
602            (None, None) => {}
603        }
604    } else if let Some(bearer) = auth::resolve_bearer_token(&llm.auth, llm.api_key.as_deref())
605        .await
606        .map_err(|message| LlmCallError::Auth { message })?
607    {
608        headers.push(("Authorization".to_owned(), format!("Bearer {bearer}")));
609    }
610    for (name, value) in &llm.headers {
611        headers.push((name.clone(), value.clone()));
612    }
613    for (name, command) in &llm.header_commands {
614        let value = auth::run_header_command(command)
615            .await
616            .map_err(|e| LlmCallError::Auth {
617                message: format!("header command for `{name}` failed: {e}"),
618            })?;
619        headers.push((name.clone(), value));
620    }
621
622    let payload = build_openai_payload(llm, system, user, image_data_urls);
623
624    let mut req = client.post(&url).header("Content-Type", "application/json");
625
626    for (name, value) in headers {
627        req = req.header(name.as_str(), value.as_str());
628    }
629
630    let resp = req
631        .json(&payload)
632        .send()
633        .await
634        .map_err(|e| LlmCallError::Transport {
635            message: e.to_string(),
636        })?;
637    let status = resp.status();
638    let status_u16 = status.as_u16();
639    let body = resp.text().await.unwrap_or_default();
640    if !status.is_success() {
641        return Err(LlmCallError::Http {
642            status: status_u16,
643            body,
644        });
645    }
646    if body.trim().is_empty() {
647        // HTTP 200 + empty body: a transient gateway hiccup (cold-start /
648        // warm-up), not a client error. Retried with a longer backoff.
649        return Err(LlmCallError::EmptyBody { status: status_u16 });
650    }
651    let json: Value = match serde_json::from_str(&body) {
652        Ok(v) => v,
653        Err(e) => {
654            return Err(LlmCallError::InvalidJson {
655                status: status_u16,
656                detail: e.to_string(),
657                body,
658            });
659        }
660    };
661    let usage = costs::extract_usage(&json);
662    let content = json["choices"][0]["message"]["content"]
663        .as_str()
664        .map(String::from)
665        .ok_or_else(|| LlmCallError::MissingContent {
666            json: json.to_string(),
667        })?;
668
669    Ok(LlmResponse { content, usage })
670}
671
672/// Builds the OpenAI-compatible chat completions request body (shared by the
673/// `openai` and `azure` providers).
674#[must_use]
675fn build_openai_payload(
676    llm: &LlmConfig,
677    system: &str,
678    user: &str,
679    image_data_urls: Option<&[String]>,
680) -> Value {
681    let mut payload = serde_json::json!({
682        "model": llm.model,
683        "messages": build_messages(system, user, image_data_urls),
684        "max_tokens": 4096,
685        "temperature": llm.temperature
686    });
687    if let Some(think) = llm.thinking {
688        if think {
689            payload["thinking"] = serde_json::json!({"type": "enabled"});
690        } else {
691            payload["thinking"] = serde_json::json!({"type": "disabled"});
692        }
693    }
694    // Merge provider-specific parameters into the request body.
695    if !llm.model_params.is_empty() {
696        if let Value::Object(ref mut map) = payload {
697            for (key, val) in &llm.model_params {
698                map.insert(key.clone(), val.clone());
699            }
700        }
701    }
702    // Explicit prompt caching for Anthropic-style models (direct or via a
703    // gateway such as OpenRouter): mark the static system block so the
704    // provider caches everything up to it. OpenAI/Azure/Groq/xAI/DeepSeek
705    // cache automatically and must not receive this marker.
706    if llm.cache && !system.is_empty() && is_anthropic_style_model(&llm.model) {
707        payload["messages"][0]["content"] = serde_json::json!([
708            {"type": "text", "text": system, "cache_control": {"type": "ephemeral"}}
709        ]);
710    }
711    payload
712}
713
714/// Whether a model name suggests an Anthropic-style API that uses explicit
715/// `cache_control` markers (directly or via a gateway such as `OpenRouter`).
716#[must_use]
717fn is_anthropic_style_model(model: &str) -> bool {
718    let model = model.to_ascii_lowercase();
719    model.contains("claude") || model.contains("anthropic")
720}
721
722/// JavaScript to extract interactive elements from the current page.
723/// Returns a JSON array of objects with tag, selector, and label.
724pub const DOM_EXTRACT_JS: &str = r#"
725(() => {
726  const interactive = 'a, button, input, textarea, select, [role="button"], [onclick], [tabindex], [data-testid], [aria-label]';
727  const els = document.querySelectorAll(interactive);
728  const info = [];
729  const seen = new Set();
730  els.forEach((el, i) => {
731    const rect = el.getBoundingClientRect();
732    if (rect.width === 0 || rect.height === 0) return;
733    const tag = el.tagName.toLowerCase();
734    let selector = '';
735    if (el.id) selector = '#' + CSS.escape(el.id);
736    else if (el.getAttribute('data-testid')) selector = '[data-testid="' + el.getAttribute('data-testid') + '"]';
737    else if (el.name) selector = '[name="' + CSS.escape(el.name) + '"]';
738    else if (el.className && typeof el.className === 'string') {
739      const cls = el.className.trim().split(/\\s+/)[0];
740      if (cls) selector = tag + '.' + CSS.escape(cls);
741    }
742    if (!selector) selector = tag;
743    if (seen.has(selector)) return;
744    seen.add(selector);
745
746    let label = '';
747    const aria = el.getAttribute('aria-label');
748    if (aria) {
749      label = aria;
750    } else if (tag === 'input' || tag === 'textarea' || tag === 'select') {
751      label = el.placeholder || el.name || el.getAttribute('aria-label') || '';
752      if (el.type && !label) label = el.type;
753    } else {
754      label = (el.textContent || '').trim().substring(0, 80);
755    }
756
757    info.push(i + ': ' + selector + ' [' + tag + '] "' + label + '"');
758  });
759  return JSON.stringify(info);
760})()
761"#;
762
763/// Truncates a string to the given maximum length, appending a marker with
764/// the number of omitted characters if truncation occurred.
765///
766/// The cut point is always a UTF-8 char boundary, so multi-byte input (umlauts,
767/// emoji, CJK) can never panic the caller.
768#[must_use]
769pub fn truncate(s: &str, max_len: usize) -> String {
770    if s.len() <= max_len {
771        s.to_owned()
772    } else {
773        let cut = floor_char_boundary(s, max_len);
774        let omitted = s[cut..].chars().count();
775        format!("{}...<truncated {omitted} chars>", &s[..cut])
776    }
777}
778
779/// Returns the largest char boundary index in `s` that is `<= index`.
780fn floor_char_boundary(s: &str, index: usize) -> usize {
781    let index = index.min(s.len());
782    let mut i = index;
783    while i > 0 && !s.is_char_boundary(i) {
784        i -= 1;
785    }
786    i
787}
788
789#[cfg(test)]
790mod tests {
791    /// Serializes env-var-mutating tests: they race when the test binary runs
792    /// them in parallel, which intermittently failed CI.
793    static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
794
795    /// Acquires the env lock for one test.
796    fn env_guard() -> std::sync::MutexGuard<'static, ()> {
797        ENV_LOCK
798            .lock()
799            .unwrap_or_else(std::sync::PoisonError::into_inner)
800    }
801
802    use crate::costs::extract_usage;
803    use crate::truncate;
804    use crate::{
805        default_llm_attempts, llm_base_url, llm_model, parse_headers_env, AuthConfig, AwsConfig,
806        LlmConfig, Provider,
807    };
808
809    /// Starts a minimal HTTP server that answers every chat-completions
810    /// request with `status`/`body`. Returns its base URL.
811    fn mock_llm_server(status: u16, body: &'static str) -> String {
812        use std::io::{Read, Write};
813        use std::net::TcpListener;
814        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
815        let addr = listener.local_addr().unwrap();
816        std::thread::spawn(move || {
817            for stream in listener.incoming() {
818                let Ok(mut stream) = stream else { break };
819                let mut buf = [0u8; 4096];
820                let _ = stream.read(&mut buf);
821                let resp = format!(
822                    "HTTP/1.1 {status} {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
823                    if status == 200 { "OK" } else { "ERROR" },
824                    body.len(),
825                    body
826                );
827                let _ = stream.write_all(resp.as_bytes());
828            }
829        });
830        format!("http://{addr}")
831    }
832
833    const PASS_BODY: &str = r#"{"choices":[{"message":{"content":"PASS"}}],"usage":{"prompt_tokens":7,"completion_tokens":2}}"#;
834
835    fn cfg(url: &str, attempts: u32) -> LlmConfig {
836        LlmConfig {
837            url: url.to_owned(),
838            model: "mock".to_owned(),
839            api_key: None,
840            headers: std::collections::HashMap::new(),
841            timeout: std::time::Duration::from_secs(10),
842            temperature: 0.0,
843            thinking: None,
844            model_params: std::collections::HashMap::new(),
845            cache: true,
846            max_attempts: attempts,
847            provider: Provider::Openai,
848            deployment: None,
849            api_version: None,
850            auth: AuthConfig::default(),
851            header_commands: std::collections::HashMap::new(),
852            aws: AwsConfig::default(),
853        }
854    }
855
856    #[tokio::test]
857    async fn test_chain_primary_success_returns_index_zero() {
858        let good = mock_llm_server(200, PASS_BODY);
859        let (resp, idx) = crate::llm_chat_with_usage_chain(&cfg(&good, 2), &[], "s", "u")
860            .await
861            .expect("primary endpoint should answer");
862        assert_eq!(idx, 0);
863        assert_eq!(resp.content, "PASS");
864        assert_eq!(resp.usage.prompt_tokens, 7);
865    }
866
867    #[tokio::test]
868    async fn test_chain_falls_back_on_empty_200() {
869        // Primary always returns HTTP 200 with an EMPTY body (the gateway
870        // warm-up signature) — after `max_attempts` it must hand over to the
871        // fallback, which answers properly.
872        let broken = mock_llm_server(200, "");
873        let good = mock_llm_server(200, PASS_BODY);
874        let (resp, idx) =
875            crate::llm_chat_with_usage_chain(&cfg(&broken, 2), &[cfg(&good, 2)], "s", "u")
876                .await
877                .expect("fallback endpoint should answer");
878        assert_eq!(idx, 1);
879        assert_eq!(resp.content, "PASS");
880    }
881
882    #[tokio::test]
883    async fn test_chain_reports_all_endpoints_on_total_failure() {
884        let broken1 = mock_llm_server(200, "");
885        let broken2 = mock_llm_server(503, "unavailable");
886        let err =
887            crate::llm_chat_with_usage_chain(&cfg(&broken1, 2), &[cfg(&broken2, 2)], "s", "u")
888                .await
889                .expect_err("both endpoints fail");
890        assert!(err.contains("all 2 endpoint(s)"), "got: {err}");
891        assert!(err.contains(&broken1), "primary URL missing: {err}");
892        assert!(err.contains(&broken2), "fallback URL missing: {err}");
893    }
894
895    #[test]
896    fn test_default_llm_attempts_env() {
897        let _env = env_guard();
898        std::env::set_var("HARNESS_LLM_CALL_ATTEMPTS", "7");
899        assert_eq!(default_llm_attempts(), 7);
900        std::env::set_var("HARNESS_LLM_CALL_ATTEMPTS", "0");
901        assert_eq!(default_llm_attempts(), 3, "0 must fall back to default");
902        std::env::set_var("HARNESS_LLM_CALL_ATTEMPTS", "junk");
903        assert_eq!(default_llm_attempts(), 3, "non-numeric must fall back");
904        std::env::remove_var("HARNESS_LLM_CALL_ATTEMPTS");
905        assert_eq!(default_llm_attempts(), 3);
906    }
907
908    #[test]
909    fn test_truncate_short() {
910        assert_eq!(truncate("hello", 10), "hello");
911    }
912
913    #[test]
914    fn test_truncate_long() {
915        let result = truncate("hello world", 5);
916        assert!(result.contains("<truncated 6 chars>"));
917        assert!(result.starts_with("hello"));
918    }
919
920    #[test]
921    fn test_truncate_exact_length() {
922        assert_eq!(truncate("abcde", 5), "abcde");
923    }
924
925    #[test]
926    fn test_truncate_empty() {
927        assert_eq!(truncate("", 5), "");
928    }
929
930    #[test]
931    fn test_parse_headers_env_empty() {
932        let _env = env_guard();
933        std::env::remove_var("HARNESS_LLM_HEADERS");
934        let h = parse_headers_env();
935        assert!(h.is_empty());
936    }
937
938    #[test]
939    fn test_parse_headers_env_valid() {
940        let _env = env_guard();
941        std::env::set_var("HARNESS_LLM_HEADERS", r#"{"X-Org":"acme","X-Version":"1"}"#);
942        let h = parse_headers_env();
943        assert_eq!(h.get("X-Org").map(String::as_str), Some("acme"));
944        assert_eq!(h.get("X-Version").map(String::as_str), Some("1"));
945        std::env::remove_var("HARNESS_LLM_HEADERS");
946    }
947
948    #[test]
949    fn test_parse_headers_env_invalid_json() {
950        let _env = env_guard();
951        std::env::set_var("HARNESS_LLM_HEADERS", "not-json");
952        let h = parse_headers_env();
953        assert!(h.is_empty());
954        std::env::remove_var("HARNESS_LLM_HEADERS");
955    }
956
957    #[test]
958    fn test_llm_config_from_env_defaults() {
959        let _env = env_guard();
960        #[allow(clippy::float_cmp)]
961        {
962            let config = LlmConfig::from_env();
963            assert_eq!(config.temperature, 0.0);
964            assert!(config.thinking.is_none());
965            assert!(config.model_params.is_empty());
966        }
967    }
968
969    #[test]
970    fn test_extract_usage_full() {
971        let json = serde_json::json!({
972            "usage": {
973                "prompt_tokens": 100,
974                "completion_tokens": 200,
975                "total_tokens": 300,
976                "prompt_tokens_details": { "cached_tokens": 40 }
977            }
978        });
979        let usage = extract_usage(&json);
980        assert_eq!(usage.prompt_tokens, 100);
981        assert_eq!(usage.completion_tokens, 200);
982        assert_eq!(usage.total_tokens, 300);
983        assert_eq!(usage.cached_input_tokens, 40);
984        assert_eq!(usage.cache_creation_input_tokens, 0);
985    }
986
987    #[test]
988    fn test_extract_usage_anthropic_cache_fields() {
989        let json = serde_json::json!({
990            "usage": {
991                "prompt_tokens": 100,
992                "completion_tokens": 20,
993                "total_tokens": 120,
994                "cache_read_input_tokens": 30,
995                "cache_creation_input_tokens": 12
996            }
997        });
998        let usage = extract_usage(&json);
999        assert_eq!(usage.cached_input_tokens, 30);
1000        assert_eq!(usage.cache_creation_input_tokens, 12);
1001    }
1002
1003    #[test]
1004    fn test_extract_usage_openai_style_cache_write() {
1005        // OpenRouter / Azure OpenAI (GPT-5.6+) report cache writes under
1006        // prompt_tokens_details.cache_write_tokens alongside cached_tokens.
1007        let json = serde_json::json!({
1008            "usage": {
1009                "prompt_tokens": 1566,
1010                "completion_tokens": 1518,
1011                "total_tokens": 3084,
1012                "prompt_tokens_details": {
1013                    "cached_tokens": 1408,
1014                    "cache_write_tokens": 100
1015                }
1016            }
1017        });
1018        let usage = extract_usage(&json);
1019        assert_eq!(usage.cached_input_tokens, 1408);
1020        assert_eq!(usage.cache_creation_input_tokens, 100);
1021    }
1022
1023    #[test]
1024    fn test_extract_usage_deepseek_cache_hit() {
1025        // DeepSeek reports cache hits under prompt_cache_hit_tokens.
1026        let json = serde_json::json!({
1027            "usage": {
1028                "prompt_tokens": 500,
1029                "completion_tokens": 50,
1030                "total_tokens": 550,
1031                "prompt_cache_hit_tokens": 400,
1032                "prompt_cache_miss_tokens": 100
1033            }
1034        });
1035        let usage = extract_usage(&json);
1036        assert_eq!(usage.cached_input_tokens, 400);
1037        assert_eq!(usage.cache_creation_input_tokens, 0);
1038    }
1039
1040    #[test]
1041    fn test_extract_usage_empty() {
1042        let json = serde_json::json!({});
1043        let usage = extract_usage(&json);
1044        assert_eq!(usage.prompt_tokens, 0);
1045        assert_eq!(usage.completion_tokens, 0);
1046        assert_eq!(usage.total_tokens, 0);
1047    }
1048
1049    #[test]
1050    fn test_truncate_unicode() {
1051        // max_len is a byte budget; the cut lands on a char boundary.
1052        // 'é' is 2 bytes, so index 3 captures "hé" (h=0, é=bytes 1-2)
1053        assert_eq!(truncate("héllo", 3), "hé...<truncated 3 chars>");
1054        // Length 5 captures full string (5 bytes)
1055        assert_eq!(truncate("hello", 5), "hello");
1056    }
1057
1058    #[test]
1059    fn test_truncate_utf8_boundary_mid_char_does_not_panic() {
1060        // Previously: &s[..2] panicked with "not a char boundary" because the
1061        // cut landed inside the 2-byte 'é'. The runner hit this on real pages
1062        // full of umlauts/emoji — inside the diagnostics path that was
1063        // supposed to save the run.
1064        // floor_char_boundary(2) lands BEFORE 'é' (byte 1), keeping "h".
1065        let result = truncate("héllo", 2);
1066        assert_eq!(result, "h...<truncated 4 chars>");
1067
1068        // 4-byte emoji: max_len=4 lands exactly on the first 🎉 boundary? No —
1069        // boundary 0 is the largest <= 4 only if 4 is a boundary; it is, so
1070        // cut=4 keeps "🎉". Check with 5 instead: cut back to 4.
1071        let cut_inside = truncate("🎉🎉🎉 boom", 5);
1072        assert_eq!(cut_inside, "🎉...<truncated 7 chars>");
1073        assert!(is_valid_utf8(&cut_inside), "result must stay valid UTF-8");
1074    }
1075
1076    #[test]
1077    fn test_truncate_utf8_exact_omitted_count() {
1078        // 5 ASCII chars cut at 10 bytes → 5 omitted
1079        assert_eq!(truncate("abcdefghij", 5), "abcde...<truncated 5 chars>");
1080        // 3 multibyte chars cut at exactly their boundary → 0... but the
1081        // guard `len <= max_len` returns the raw string first.
1082        assert_eq!(truncate("ééé", 6), "ééé");
1083        assert_eq!(truncate("ééé", 5), "éé...<truncated 1 chars>");
1084    }
1085
1086    fn is_valid_utf8(s: &str) -> bool {
1087        std::str::from_utf8(s.as_bytes()).is_ok()
1088    }
1089
1090    #[test]
1091    fn test_parse_headers_env_non_object() {
1092        let _env = env_guard();
1093        std::env::set_var("HARNESS_LLM_HEADERS", "[1, 2, 3]");
1094        let h = parse_headers_env();
1095        assert!(h.is_empty());
1096        std::env::remove_var("HARNESS_LLM_HEADERS");
1097    }
1098
1099    #[test]
1100    fn test_parse_headers_env_nested_values_filtered() {
1101        let _env = env_guard();
1102        std::env::set_var(
1103            "HARNESS_LLM_HEADERS",
1104            r#"{"str":"val","num":42,"bool":true}"#,
1105        );
1106        let h = parse_headers_env();
1107        assert_eq!(h.get("str").map(String::as_str), Some("val"));
1108        assert!(!h.contains_key("num"));
1109        assert!(!h.contains_key("bool"));
1110        std::env::remove_var("HARNESS_LLM_HEADERS");
1111    }
1112
1113    #[test]
1114    fn test_llm_config_has_default_model() {
1115        let _env = env_guard();
1116        let config = LlmConfig::from_env();
1117        assert_ne!(config.model, "");
1118    }
1119
1120    #[test]
1121    fn test_llm_base_url_default() {
1122        let _env = env_guard();
1123        std::env::remove_var("HARNESS_LLM_TEST_URL");
1124        let url = llm_base_url();
1125        assert_eq!(url, "http://localhost:8080");
1126    }
1127
1128    #[test]
1129    fn test_llm_base_url_custom() {
1130        let _env = env_guard();
1131        std::env::set_var("HARNESS_LLM_TEST_URL", "https://custom.api.com/v1");
1132        let url = llm_base_url();
1133        assert_eq!(url, "https://custom.api.com/v1");
1134        std::env::remove_var("HARNESS_LLM_TEST_URL");
1135    }
1136
1137    #[test]
1138    fn test_llm_base_url_trailing_slash() {
1139        let _env = env_guard();
1140        std::env::set_var("HARNESS_LLM_TEST_URL", "https://api.com/");
1141        let url = llm_base_url();
1142        assert_eq!(url, "https://api.com");
1143        std::env::remove_var("HARNESS_LLM_TEST_URL");
1144    }
1145
1146    #[test]
1147    fn test_llm_model_default() {
1148        let _env = env_guard();
1149        std::env::remove_var("HARNESS_LLM_TEST_MODEL");
1150        assert_eq!(llm_model(), "deepseek");
1151    }
1152
1153    #[test]
1154    fn test_llm_model_custom() {
1155        let _env = env_guard();
1156        std::env::set_var("HARNESS_LLM_TEST_MODEL", "gpt-4o");
1157        assert_eq!(llm_model(), "gpt-4o");
1158        std::env::remove_var("HARNESS_LLM_TEST_MODEL");
1159    }
1160
1161    #[test]
1162    fn test_extract_usage_partial() {
1163        let _env = env_guard();
1164        let json = serde_json::json!({
1165            "usage": {
1166                "prompt_tokens": 50
1167            }
1168        });
1169        let usage = extract_usage(&json);
1170        assert_eq!(usage.prompt_tokens, 50);
1171        assert_eq!(usage.completion_tokens, 0);
1172        assert_eq!(usage.total_tokens, 0);
1173    }
1174
1175    #[test]
1176    fn test_browser_headless_default() {
1177        let _env = env_guard();
1178        std::env::remove_var("HARNESS_BROWSER_HEADLESS");
1179        assert!(crate::browser_headless());
1180    }
1181}