//! HTTP client for the xAI sampling APIs. //! //! Owns the `reqwest::Client`, default request headers, and per-method //! defaults. Talks to three backend shapes: //! //! * Chat Completions (`/chat/completions`) //! * Responses API (`/responses`) //! * Anthropic Messages API (`/messages`) //! //! All trace-upload and URL-based header injection is intentionally //! *not* here. The session is responsible for putting any per-request //! headers (proxy auth, OTel context, etc.) //! into [`SamplerConfig::extra_headers`] before constructing the client. use eventsource_stream::Eventsource; use futures_util::StreamExt; use futures_util::stream::BoxStream; use reqwest::header::{ ACCEPT, AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue, USER_AGENT, }; use serde::Serialize; use xai_grok_sampling_types::error::{parse_error_bytes, try_parse_stream_error}; use xai_grok_sampling_types::{ ChatCompletionChunk, ChatCompletionRequest, ChatCompletionResponse, ConversationRequest, ConversationResponse, CreateResponseWrapper, DOOM_LOOP_CHECK_HEADER, MessagesRequestWrapper, ResponseModelMetadata, Result, SamplingError, build_messages_request, is_check_event, messages, rs, }; use crate::config::{AuthScheme, OriginClientInfo, SamplerConfig}; // Re-export ApiBackend from the shared types crate for downstream callers. pub use xai_grok_sampling_types::ApiBackend; /// Process-level fallback for the `x-grok-client-identifier` header. const DEFAULT_CLIENT_IDENTIFIER: &str = "grok-shell"; /// Product identifier baked into User-Agent strings. const AGENT_PRODUCT: &str = "grok-shell"; const ANTHROPIC_DEFAULT_MAX_TOKENS: u32 = 128_000; /// Per-request `x-grok-*` headers. Optional fields are skipped when empty/`None`. struct GrokRequestHeaders<'a> { conv_id: &'a str, req_id: &'a str, model_id: &'a str, session_id: &'a str, turn_idx: Option<&'a str>, agent_id: &'a str, deployment_id: Option<&'a str>, user_id: Option<&'a str>, } impl GrokRequestHeaders<'_> { fn apply(&self, builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder { let mut b = builder .header("x-grok-conv-id", self.conv_id) .header("x-grok-req-id", self.req_id) .header("x-grok-model-override", self.model_id) .header("x-grok-session-id", self.session_id) .header("x-grok-agent-id", self.agent_id); if let Some(idx) = self.turn_idx { b = b.header("x-grok-turn-idx", idx); } if let Some(id) = self.deployment_id.filter(|s| !s.is_empty()) { b = b.header("x-grok-deployment-id", id); } if let Some(id) = self.user_id.filter(|s| !s.is_empty()) { b = b.header("x-grok-user-id", id); } b } } /// Parse the `Retry-After` response header as delta-seconds. /// Our inference backends only emit integer seconds (never HTTP-date), /// so we only handle that form. HTTP-dates silently return `None` and /// the caller falls back to exponential backoff. /// Capped at 120s to prevent absurdly long sleeps from a misbehaving upstream. /// Deserialize a Responses API SSE event, with a fallback for xAI-specific /// tool types (e.g., `x_search`) that `async_openai` can't parse. /// /// The API echoes the request's `tools` array in `ResponseCompleted` and /// `ResponseCreated` events. If we sent `{"type": "x_search"}`, the response /// includes it, and `rs::Tool` deserialization fails. On failure, we strip /// unrecognized tools from the raw JSON and retry. /// /// On `response.completed` / `response.incomplete`, this also rewrites /// `response.usage.total_tokens` in place to the live context length /// (`context_details.input_tokens + context_details.output_tokens`) /// when the API emits the xAI-specific `context_details` field. /// Async-openai's typed `ResponseUsage` doesn't model `context_details`, /// so we peek the raw JSON for it. The cumulative `input_tokens` / /// `output_tokens` / `cached_tokens` continue to flow from the typed /// `ResponseUsage` unchanged so billing telemetry stays correct. When /// the API doesn't emit `context_details` (older deployments) `total_tokens` /// passes through unchanged. fn deserialize_response_event(data: &str) -> Result { let mut event = match serde_json::from_str::(data) { Ok(event) => event, Err(first_err) => { // Try sanitizing: parse as Value, strip unknown tools, retry. if let Ok(mut value) = serde_json::from_str::(data) { // Strip tools that async_openai's rs::Tool can't deserialize // (e.g., xAI-specific "x_search"). Instead of maintaining a // hardcoded allowlist, try deserializing each tool entry — // if it fails, drop it. if let Some(tools) = value .pointer_mut("/response/tools") .and_then(|v| v.as_array_mut()) { tools.retain(|t| serde_json::from_value::(t.clone()).is_ok()); } if let Ok(mut event) = serde_json::from_value::(value) { apply_terminal_event_overrides(&mut event, data); return Ok(event); } } tracing::error!( error = %first_err, raw_data = %data, "Failed to deserialize ResponseStreamEvent from stream" ); return Err(SamplingError::Serialization(first_err)); } }; apply_terminal_event_overrides(&mut event, data); Ok(event) } /// On terminal Responses API events (`response.completed` / /// `response.incomplete`), rewrite `response.usage.total_tokens` to the /// live context length when the wire includes /// `response.usage.context_details.{input_tokens, output_tokens}`. /// /// `total_tokens` drives the CLI's `/context` bar, the auto-compact /// threshold, and `meta.totalTokens` on persisted sessions. Under /// server-side multi-turn loops (e.g. `web_search`, `x_search`) the /// wire's cumulative total inflates as the loop runs; `context_details` /// reports the final turn's prompt + output tokens — the real live /// context the model is sitting in. Billing fields /// (`input_tokens`, `output_tokens`, `input_tokens_details.cached_tokens`, /// `output_tokens_details.reasoning_tokens`) stay on the cumulative /// wire values so telemetry is unaffected. /// /// No-op when: /// - the event is not terminal, /// - `response.usage` is `None`, /// - `context_details` is absent (older backends / non-loop responses), /// - or either of `context_details.{input_tokens, output_tokens}` is /// missing — we don't guess the missing half. fn apply_terminal_event_overrides(event: &mut rs::ResponseStreamEvent, data: &str) { let response = match event { rs::ResponseStreamEvent::ResponseCompleted(e) => &mut e.response, rs::ResponseStreamEvent::ResponseIncomplete(e) => &mut e.response, _ => return, }; // Re-parse for fields async_openai's types omit (context total, cost ticks). let Ok(value) = serde_json::from_str::(data) else { return; }; // Stash cost ticks in metadata for stream_responses. if let Some(ticks) = xai_grok_sampling_types::reported_cost_ticks( value .pointer("/response/usage/cost_in_usd_ticks") .and_then(|v| v.as_i64()), ) { response .metadata .get_or_insert_with(Default::default) .insert(COST_USD_TICKS_METADATA_KEY.to_owned(), ticks.to_string()); } let Some(usage) = response.usage.as_mut() else { return; }; let Some(total) = extract_context_total(&value) else { return; }; usage.total_tokens = total; } /// Metadata key for cost ticks past typed Response events. pub(crate) const COST_USD_TICKS_METADATA_KEY: &str = "xai.cost_usd_ticks"; /// Read `response.usage.context_details.{input_tokens, output_tokens}` /// from the parsed terminal-event JSON and return their sum. Returns `None` /// if either field is missing or out of `u32` range. fn extract_context_total(value: &serde_json::Value) -> Option { let cd = value.pointer("/response/usage/context_details")?; let i = u32::try_from(cd.get("input_tokens")?.as_u64()?).ok()?; let o = u32::try_from(cd.get("output_tokens")?.as_u64()?).ok()?; Some(i.saturating_add(o)) } /// Record `success=false` + `error` on the active inference span when a stream /// request fails before any response (transport/connect/TLS errors). Without /// this the `#[instrument]` span closes with both fields Empty, so an outage /// shows zero `success=false` and error-rate alerts never fire. fn record_stream_request_failure(err: &reqwest::Error) { let span = tracing::Span::current(); span.record("success", false); span.record("error", err.to_string().as_str()); } fn extract_retry_after(headers: &reqwest::header::HeaderMap) -> Option { headers .get(reqwest::header::RETRY_AFTER) .and_then(|v| v.to_str().ok()) .and_then(|s| s.parse::().ok()) .map(|s| s.min(120)) } fn extract_should_retry(headers: &reqwest::header::HeaderMap) -> Option { headers .get("x-should-retry") .and_then(|v| v.to_str().ok()) .and_then(|s| { if s.eq_ignore_ascii_case("true") { Some(true) } else if s.eq_ignore_ascii_case("false") { Some(false) } else { None // unknown value — treat as absent } }) } fn extract_model_metadata(headers: &reqwest::header::HeaderMap) -> Option { let context_window = headers .get("x-grok-context-window") .and_then(|v| v.to_str().ok()) .and_then(|s| s.parse::().ok()); let max_completion_tokens = headers .get("x-grok-max-completion-tokens") .and_then(|v| v.to_str().ok()) .and_then(|s| s.parse::().ok()); let models_etag = headers .get("x-models-etag") .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()); if context_window.is_some() || max_completion_tokens.is_some() || models_etag.is_some() { Some(ResponseModelMetadata { context_window, max_completion_tokens, models_etag, }) } else { None } } /// Wrapper for streaming chat completion requests that adds `stream` and /// `stream_options` fields without modifying the original `ChatCompletionRequest`. /// /// Uses `#[serde(flatten)]` to inline all fields from the inner request, /// allowing single-pass serialization instead of the previous two-pass /// approach (serialize to `Value`, mutate, serialize to bytes). #[derive(Serialize)] struct StreamingChatRequest<'a> { #[serde(flatten)] inner: &'a ChatCompletionRequest, stream: bool, stream_options: StreamOptions, } #[derive(Serialize)] struct StreamOptions { include_usage: bool, } /// HTTP client for sampling. Cheap to clone; carries an `Arc`-backed /// `reqwest::Client` and the default headers/request-defaults computed /// from a [`SamplerConfig`] at construction time. #[derive(Clone)] pub struct SamplingClient { http: reqwest::Client, default_headers: HeaderMap, base_url: String, defaults: ClientDefaults, /// Optional 401-attribution hook. The shell wires this to emit a /// structured event at every UNAUTHORIZED arm so 401s can be /// bucketed by stale-snapshot vs. live-token-rejected. `None` for /// sampler-only callers and tests. attribution_callback: Option, /// Per-request bearer override. See `SamplerConfig::bearer_resolver`. bearer_resolver: Option, /// Per-request header injection (OTel traceparent). header_injector: Option, } impl std::fmt::Debug for SamplingClient { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("SamplingClient") .field("base_url", &self.base_url) .field("defaults", &self.defaults) .field( "has_attribution_callback", &self.attribution_callback.is_some(), ) .field("has_bearer_resolver", &self.bearer_resolver.is_some()) .finish() } } #[derive(Clone, Debug, Default)] struct ClientDefaults { model: String, max_completion_tokens: Option, temperature: Option, top_p: Option, api_backend: ApiBackend, auth_scheme: AuthScheme, stream_tool_calls: bool, doom_loop_recovery: Option, } // ============================================================================= // User-Agent helpers // ============================================================================= #[derive(Clone, Debug, Eq, PartialEq)] struct PlatformInfo { os: String, arch: String, } impl PlatformInfo { fn current() -> Self { let os = match std::env::consts::OS { "macos" => "macos", "windows" => "windows", other => other, } .to_string(); let arch = match std::env::consts::ARCH { "arm64" => "aarch64", "x86_64" => "x86_64", other => other, } .to_string(); Self { os, arch } } } fn agent_version() -> String { xai_grok_version::VERSION.to_string() } /// Render a User-Agent string for the given origin client. /// /// Mirrors the shell's `user_agent_string_for` but uses sampler-local /// constants. The session typically owns the canonical User-Agent /// rendering for process-wide HTTP clients; this helper is for /// per-session sampling clients that want to override it. pub fn user_agent_string_for(origin: &OriginClientInfo) -> String { let agent_version = agent_version(); let platform = PlatformInfo::current(); if origin.product == AGENT_PRODUCT && origin.version.as_deref() == Some(agent_version.as_str()) { return format!( "{}/{} ({}; {})", AGENT_PRODUCT, agent_version, platform.os, platform.arch ); } match origin.version.as_deref() { Some(origin_version) => format!( "{}/{} {}/{} ({}; {})", origin.product, origin_version, AGENT_PRODUCT, agent_version, platform.os, platform.arch ), None => format!( "{} {}/{} ({}; {})", origin.product, AGENT_PRODUCT, agent_version, platform.os, platform.arch ), } } // ============================================================================= // SamplingClient // ============================================================================= impl SamplingClient { /// Construct a sampling client from a [`SamplerConfig`]. /// /// Grabs the process-wide shared `reqwest::Client` (HTTP/2 by /// default, HTTP/1.1 when `config.force_http1` is set) and /// pre-computes the default request headers. This does not perform /// any network I/O. pub fn new(config: SamplerConfig) -> Result { let mut headers = HeaderMap::new(); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); if let Some(ref api_key) = config.api_key { match config.auth_scheme { AuthScheme::XApiKey => { let header_value = HeaderValue::from_str(api_key).map_err(|_| { tracing::debug!( api_key = %api_key, "Invalid api_key: cannot be converted to a valid HTTP header" ); SamplingError::Auth( "Invalid api_key: cannot be converted to a valid HTTP header" .to_string(), ) })?; headers.insert(HeaderName::from_static("x-api-key"), header_value); } AuthScheme::Bearer => { let bearer = format!("Bearer {}", api_key); let header_value = HeaderValue::from_str(&bearer).map_err(|_| { tracing::debug!( api_key = %api_key, "Invalid api_key: cannot be converted to a valid HTTP Authorization header" ); SamplingError::Auth( "Invalid api_key: cannot be converted to a valid HTTP Authorization header" .to_string(), ) })?; headers.insert(AUTHORIZATION, header_value); } } } // Apply all extra headers verbatim. This is the single // injection point for proxy-auth headers and any other URL- or // environment-specific headers the session decides to set. for (key, value) in &config.extra_headers { let header_name = HeaderName::try_from(key.as_str()) .map_err(|_| SamplingError::InvalidConfiguration("Invalid extra header name"))?; let header_value = HeaderValue::from_str(value) .map_err(|_| SamplingError::InvalidConfiguration("Invalid extra header value"))?; headers.insert(header_name, header_value); } // Add x-grok-client-version header for version gating at the proxy. if let Some(client_version) = config.client_version.as_ref() && let Ok(header_value) = HeaderValue::from_str(client_version) { headers.insert( HeaderName::from_static("x-grok-client-version"), header_value, ); } if let Some(deployment_id) = config.deployment_id.as_ref() && let Ok(header_value) = HeaderValue::from_str(deployment_id) { headers.insert( HeaderName::from_static("x-grok-deployment-id"), header_value, ); } if let Some(user_id) = config.user_id.as_ref() && let Ok(header_value) = HeaderValue::from_str(user_id) { headers.insert(HeaderName::from_static("x-grok-user-id"), header_value); } { let client_id = config .client_identifier .clone() .unwrap_or_else(|| DEFAULT_CLIENT_IDENTIFIER.to_string()); if let Ok(header_value) = HeaderValue::from_str(&client_id) { headers.insert( HeaderName::from_static("x-grok-client-identifier"), header_value, ); } } // Always set User-Agent: per-session origin if available, else fallback. { let ua_string = match config.origin_client.as_ref() { Some(origin) => user_agent_string_for(origin), None => user_agent_string_for(&OriginClientInfo { product: AGENT_PRODUCT.to_string(), version: Some(agent_version()), }), }; if let Ok(v) = HeaderValue::from_str(&ua_string) { headers.insert(USER_AGENT, v); } } let http = if config.force_http1 { tracing::info!("Using HTTP/1.1 for sampling client (force_http1=true)"); crate::shared_http::client_http1().map_err(SamplingError::Http)? } else { crate::shared_http::client().map_err(SamplingError::Http)? }; tracing::info!( target: crate::sampling_log::TARGET, event = "client_new", base_url = %config.base_url, model = %config.model, api_backend = ?config.api_backend, auth_scheme = ?config.auth_scheme, // "unset" (not "none"): `ReasoningEffort::None` is a real wire value; // logging the absent Option as "none" looked like we were sending it. reasoning_effort = config.reasoning_effort.map_or("unset", |e| e.as_str()), has_api_key = config.api_key.is_some(), has_bearer_resolver = config.bearer_resolver.is_some(), has_authorization_header = headers.get(AUTHORIZATION).is_some(), has_x_api_key_header = headers.get(HeaderName::from_static("x-api-key")).is_some(), ); let defaults = ClientDefaults { model: config.model, max_completion_tokens: config.max_completion_tokens, temperature: config.temperature, top_p: config.top_p, api_backend: config.api_backend, auth_scheme: config.auth_scheme, stream_tool_calls: config.stream_tool_calls, doom_loop_recovery: config.doom_loop_recovery, }; Ok(Self { http, default_headers: headers, base_url: config.base_url, defaults, attribution_callback: config.attribution_callback, bearer_resolver: config.bearer_resolver, header_injector: config.header_injector, }) } /// The configured API backend for this client. pub fn api_backend(&self) -> ApiBackend { self.defaults.api_backend.clone() } /// POST with default headers. Overrides auth from resolver if wired. fn post(&self, url: impl reqwest::IntoUrl) -> reqwest::RequestBuilder { let mut headers = self.default_headers.clone(); if let Some(resolver) = &self.bearer_resolver && let Some(fresh) = resolver.current_bearer() { match self.defaults.auth_scheme { AuthScheme::XApiKey => { headers.remove(AUTHORIZATION); if let Ok(v) = HeaderValue::from_str(&fresh) { headers.insert(HeaderName::from_static("x-api-key"), v); } } AuthScheme::Bearer => { headers.remove(HeaderName::from_static("x-api-key")); if let Ok(v) = HeaderValue::from_str(&format!("Bearer {fresh}")) { headers.insert(AUTHORIZATION, v); } } } } { let auth_prefix = headers .get(AUTHORIZATION) .and_then(|v| v.to_str().ok()) .map(|s| s.chars().take(20).collect::()); let x_api_key_prefix = headers .get(HeaderName::from_static("x-api-key")) .and_then(|v| v.to_str().ok()) .map(|s| s.chars().take(12).collect::()); tracing::info!( target: crate::sampling_log::TARGET, event = "client_post", base_url = %self.base_url, model = %self.defaults.model, api_backend = ?self.defaults.api_backend, auth_scheme = ?self.defaults.auth_scheme, has_bearer_resolver = self.bearer_resolver.is_some(), has_authorization_header = headers.get(AUTHORIZATION).is_some(), has_x_api_key_header = headers.get(HeaderName::from_static("x-api-key")).is_some(), auth_header_prefix = auth_prefix.as_deref().unwrap_or("none"), x_api_key_prefix = x_api_key_prefix.as_deref().unwrap_or("none"), ); } if let Some(injector) = &self.header_injector { injector.inject(&mut headers); } self.http.post(url).headers(headers) } /// Bearer prefix for 401 attribution. Prefers live resolver, falls back to default_headers. fn current_sent_bearer_prefix(&self) -> Option { self.bearer_resolver .as_ref() .and_then(|r| r.current_bearer()) .or_else(|| self.extract_sent_bearer()) .map(|mut s| { s.truncate(crate::attribution::SENT_BEARER_PREFIX_LEN.min(s.len())); s }) } /// Extract the bearer from `default_headers`, truncated to prefix length. /// Reads `x-api-key` (Anthropic Messages API) or `Authorization` (OpenAI-completions). fn extract_sent_bearer(&self) -> Option { let raw = match self.defaults.auth_scheme { AuthScheme::XApiKey => self .default_headers .get(HeaderName::from_static("x-api-key")) .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()), AuthScheme::Bearer => self .default_headers .get(AUTHORIZATION) .and_then(|v| v.to_str().ok()) .and_then(|s| s.strip_prefix("Bearer ")) .map(|s| s.to_string()), }; raw.map(|mut s| { // Truncate in-place so we never materialize a heap-resident // copy of the full bearer outside the local stack of this // function. `String::truncate` operates on byte indices and // panics on a non-char-boundary cut; bearer tokens are // ASCII (per the `Authorization` and `x-api-key` header // grammars) so the byte index is always safe. s.truncate(crate::attribution::SENT_BEARER_PREFIX_LEN.min(s.len())); s }) } /// Invoke the optional 401 attribution callback for one logical /// 401 response. Each of the six UNAUTHORIZED arms in this file /// calls this helper immediately before returning /// `SamplingError::Auth(...)`. Emit happens at the lowest layer /// that saw the status, so higher layers that react to a 401 must /// not emit a duplicate event. /// /// The bearer passed to the callback is already truncated to /// [`crate::attribution::SENT_BEARER_PREFIX_LEN`] characters by /// [`Self::extract_sent_bearer`]; the trait contract guarantees /// that callers downstream of this crate never see the full /// bearer. fn record_401_attribution(&self, consumer: crate::attribution::SamplingConsumer) { if let Some(cb) = self.attribution_callback.as_ref() { let sent_prefix = self.current_sent_bearer_prefix(); cb.record_401(consumer, sent_prefix.as_deref()); } } pub fn auth_info(&self) -> crate::sampling_log::AuthInfo { let auth_prefix = self.current_sent_bearer_prefix(); let auth_type = match (&self.defaults.auth_scheme, &auth_prefix) { (AuthScheme::XApiKey, Some(_)) => "x-api-key", (AuthScheme::Bearer, Some(_)) => "bearer", (_, None) => "none", }; crate::sampling_log::AuthInfo { auth_type, auth_prefix, } } /// Check if a header name contains sensitive information that should be redacted. fn is_sensitive_header(name: &str) -> bool { let lower = name.to_lowercase(); lower.contains("authorization") || lower.contains("api-key") || lower.contains("apikey") || lower.contains("token") || lower.contains("secret") } /// Format a single header for error messages, redacting sensitive values. fn format_header(name: &str, value: &str) -> String { let display_value = if Self::is_sensitive_header(name) { "[REDACTED]" } else { value }; format!(" {}: {}", name, display_value) } /// Build request headers string for error messages (redacting sensitive values). fn format_request_headers( &self, x_grok_conv_id: &str, x_grok_req_id: &str, model_id: &str, include_accept: bool, ) -> Vec { let mut req_headers: Vec = self .default_headers .iter() .map(|(name, value)| { Self::format_header(name.as_str(), value.to_str().unwrap_or("[non-utf8]")) }) .collect(); req_headers.push(Self::format_header("x-grok-conv-id", x_grok_conv_id)); req_headers.push(Self::format_header("x-grok-req-id", x_grok_req_id)); req_headers.push(Self::format_header("x-grok-model-override", model_id)); if include_accept { req_headers.push(Self::format_header("accept", "text/event-stream")); } req_headers } /// Build response headers string for error messages. fn format_response_headers(response: &reqwest::Response) -> Vec { response .headers() .iter() .map(|(name, value)| Self::format_header(name.as_str(), &format!("{:?}", value))) .collect() } /// Log all headers from a request at debug level (redacting sensitive values). fn log_request_headers(request: &reqwest::Request, endpoint_name: &str) { for (name, value) in request.headers().iter() { let value_str = if Self::is_sensitive_header(name.as_str()) { "[REDACTED]" } else { value.to_str().unwrap_or("[non-utf8]") }; tracing::debug!( header_name = %name, header_value = %value_str, "Request header ({})", endpoint_name ); } } /// Build error context message based on error type and status code. /// Includes relevant request/response details depending on what the error is about. fn build_api_error_message( &self, status: reqwest::StatusCode, server_message: &str, endpoint: &str, req_headers: &[String], resp_headers: Option<&[String]>, ) -> String { let server_message_lower = server_message.to_lowercase(); let mut context_parts = vec![server_message.to_string()]; context_parts.push(format!("\nRequest URL: {}", endpoint)); // Show headers if error mentions headers if server_message_lower.contains("header") { context_parts.push(format!("Request headers:\n{}", req_headers.join("\n"))); } // Always show response headers for server errors if status.is_server_error() && let Some(resp_hdrs) = resp_headers { context_parts.push(format!("Response headers:\n{}", resp_hdrs.join("\n"))); } context_parts.join("\n") } fn endpoint(&self, path: &str) -> String { let base = self.base_url.trim_end_matches('/'); let path = path.trim_start_matches('/'); format!("{base}/{path}") } fn apply_defaults(&self, mut request: ChatCompletionRequest) -> Result { if request.model.is_none() { request.model = Some(self.defaults.model.clone()); } if request.max_tokens.is_none() { request.max_tokens = self.defaults.max_completion_tokens; } if request.temperature.is_none() { request.temperature = self.defaults.temperature; } if request.top_p.is_none() { request.top_p = self.defaults.top_p; } Ok(request) } async fn handle_response(&self, response: reqwest::Response) -> Result { let status = response.status(); let model_metadata = extract_model_metadata(response.headers()); let retry_after_secs = extract_retry_after(response.headers()); let should_retry = extract_should_retry(response.headers()); let bytes = response.bytes().await?; if !status.is_success() { if status == reqwest::StatusCode::UNAUTHORIZED { self.record_401_attribution(crate::attribution::SamplingConsumer::ChatCompletions); let server_message = parse_error_bytes(bytes.as_ref()); return Err(SamplingError::Auth(format!( "Unauthorized (401): {server_message}" ))); } let message = parse_error_bytes(bytes.as_ref()); return Err(SamplingError::Api { status, message, model_metadata, retry_after_secs, should_retry, }); } let completion = serde_json::from_slice::(&bytes).map_err(|e| { let raw_body = String::from_utf8_lossy(&bytes); tracing::error!( error = %e, raw_body = %raw_body, "Failed to deserialize ChatCompletionResponse" ); SamplingError::Serialization(e) })?; Ok(completion) } // ========================================================================= // Chat Completions API // ========================================================================= pub async fn chat_completion( &self, request: ChatCompletionRequest, ) -> Result { let payload = self.apply_defaults(request)?; let x_grok_conv_id = &payload.x_grok_conv_id.clone().unwrap_or_default(); let x_grok_req_id = &payload.x_grok_req_id.clone().unwrap_or_default(); let model_id = payload.model.clone().unwrap_or_default(); tracing::debug!( base_url = %self.base_url, model_id = %model_id, "Sending chat completion request" ); let grok_headers = GrokRequestHeaders { conv_id: x_grok_conv_id, req_id: x_grok_req_id, model_id: &model_id, session_id: payload.x_grok_session_id.as_deref().unwrap_or_default(), turn_idx: payload.x_grok_turn_idx.as_deref(), agent_id: payload.x_grok_agent_id.as_deref().unwrap_or_default(), deployment_id: payload.x_grok_deployment_id.as_deref(), user_id: payload.x_grok_user_id.as_deref(), }; let http_request = grok_headers .apply(self.post(self.endpoint("chat/completions"))) .json(&payload); let response = http_request.send().await.map_err(|e| { // Log at debug level; errors are surfaced to the caller. tracing::debug!("HTTP request failed: {}", e); e })?; self.handle_response(response).await } /// Start a streaming chat completion request. Returns a stream of typed chunks. #[tracing::instrument( name = "http.chat_completion_stream", skip_all, fields( endpoint = %self.endpoint("chat/completions"), model_id = request.model.as_deref().unwrap_or(""), status_code = tracing::field::Empty, success = tracing::field::Empty, error = tracing::field::Empty, ) )] pub async fn chat_completion_stream( &self, request: ChatCompletionRequest, ) -> Result<( BoxStream<'static, Result>, Option, )> { let payload = self.apply_defaults(request)?; let x_grok_conv_id = &payload.x_grok_conv_id.clone().unwrap_or_default(); let x_grok_req_id = &payload.x_grok_req_id.clone().unwrap_or_default(); let model_id = payload.model.clone().unwrap_or_default(); // Wrap the request with streaming fields and serialize once. // Previously this path serialized twice: first to serde_json::Value // (to inject `stream` and `stream_options`), then to HTTP body bytes. let streaming_request = StreamingChatRequest { inner: &payload, stream: true, stream_options: StreamOptions { include_usage: true, }, }; let grok_headers = GrokRequestHeaders { conv_id: x_grok_conv_id, req_id: x_grok_req_id, model_id: &model_id, session_id: payload.x_grok_session_id.as_deref().unwrap_or_default(), turn_idx: payload.x_grok_turn_idx.as_deref(), agent_id: payload.x_grok_agent_id.as_deref().unwrap_or_default(), deployment_id: payload.x_grok_deployment_id.as_deref(), user_id: payload.x_grok_user_id.as_deref(), }; let http_request = grok_headers .apply(self.post(self.endpoint("chat/completions"))) .header(ACCEPT, HeaderValue::from_static("text/event-stream")) .json(&streaming_request); let built_request = http_request.build().map_err(|e| { tracing::error!("Failed to build HTTP request: {}", e); SamplingError::Http(e) })?; tracing::debug!( url = %built_request.url(), method = %built_request.method(), "Sending chat/completions request" ); Self::log_request_headers(&built_request, "chat/completions"); let response = self.http.execute(built_request).await.map_err(|e| { tracing::debug!("HTTP request failed: {}", e); record_stream_request_failure(&e); e })?; let status = response.status(); let span = tracing::Span::current(); span.record("status_code", status.as_u16() as i64); span.record("success", status.is_success()); let model_metadata = extract_model_metadata(response.headers()); let retry_after_secs = extract_retry_after(response.headers()); let should_retry = extract_should_retry(response.headers()); if !status.is_success() { if status == reqwest::StatusCode::UNAUTHORIZED { span.record("error", "unauthorized (401)"); self.record_401_attribution( crate::attribution::SamplingConsumer::ChatCompletionsStream, ); let endpoint = self.endpoint("chat/completions"); let server_message = response.text().await.unwrap_or_default(); return Err(SamplingError::Auth(format!( "Unauthorized (401) from {endpoint}: {server_message}" ))); } let req_headers = self.format_request_headers(x_grok_conv_id, x_grok_req_id, &model_id, true); let resp_headers = Self::format_response_headers(&response); let bytes = response.bytes().await?; let server_message = parse_error_bytes(bytes.as_ref()); let message = self.build_api_error_message( status, &server_message, &self.endpoint("chat/completions"), &req_headers, Some(&resp_headers), ); span.record("error", message.as_str()); tracing::error!( status = %status, error_message = %message, model_id = %model_id, "chat/completions API error" ); return Err(SamplingError::Api { status, message, model_metadata, retry_after_secs, should_retry, }); } // Strip UTF-8 BOM if present: eventsource-stream 0.2.3 incorrectly slices BOM at byte 1 instead of 3. const UTF8_BOM: &[u8] = &[0xEF, 0xBB, 0xBF]; let mut is_first = true; let byte_stream = response.bytes_stream().map(move |result| { result.map(|bytes| { if is_first { is_first = false; if bytes.starts_with(UTF8_BOM) { return bytes.slice(UTF8_BOM.len()..); } } bytes }) }); // Turn raw bytes into SSE events let event_stream = byte_stream.eventsource(); // Map SSE events into ChatCompletionChunk. // Uses `scan` so that `[DONE]` and transport errors both terminate the // stream (`None`). The first transport error is emitted to the consumer, // then subsequent polls return `None` -- preventing an infinite busy-loop // when the HTTP/2 connection drops and h2 keeps producing errors. let chunks = event_stream .scan(false, |had_transport_error, event_res| { if *had_transport_error { return std::future::ready(None); } let item = match event_res { Ok(event) => { let data = &event.data; if data == "[DONE]" { return std::future::ready(None); } tracing::info!( target: crate::sampling_log::TARGET, event = "sse_chunk", backend = "chat_completions", data = %data, ); if let Some(stream_error) = try_parse_stream_error(data) { Some(Err(stream_error)) } else { Some( serde_json::from_str::(data).map_err(|e| { tracing::error!( error = %e, raw_data = %data, "Failed to deserialize ChatCompletionChunk from stream" ); SamplingError::Serialization(e) }), ) } } Err(e) => { *had_transport_error = true; Some(Err(SamplingError::EventStreamError(e.to_string()))) } }; std::future::ready(item) }) .boxed(); Ok((chunks, model_metadata)) } // ========================================================================= // Responses API // ========================================================================= /// Apply default configuration to a Responses API request. fn apply_response_defaults(&self, request: &mut CreateResponseWrapper) -> Result<()> { // Apply model default if not specified if request.inner.model.is_none() { request.inner.model = Some(self.defaults.model.clone()); } // Apply temperature default if not specified if request.inner.temperature.is_none() { request.inner.temperature = self.defaults.temperature; } // Apply top_p default if not specified if request.inner.top_p.is_none() { request.inner.top_p = self.defaults.top_p; } // Apply max_output_tokens default if not specified if request.inner.max_output_tokens.is_none() { request.inner.max_output_tokens = self.defaults.max_completion_tokens; } // Set store to false if not specified (default is true, but that breaks ZDR compliance) if request.inner.store.is_none() { request.inner.store = Some(false); } // Include encrypted reasoning content if not specified let includes = request.inner.include.get_or_insert_with(Vec::new); if !includes.contains(&rs::IncludeEnum::ReasoningEncryptedContent) { includes.push(rs::IncludeEnum::ReasoningEncryptedContent); } Ok(()) } /// Create a response using the Responses API (non-streaming). /// /// This uses the Responses API format which provides a simpler interface /// for multi-turn conversations and tool calling. pub async fn create_response( &self, mut request: CreateResponseWrapper, ) -> Result { self.apply_response_defaults(&mut request)?; let x_grok_conv_id = request.x_grok_conv_id.as_deref().unwrap_or_default(); let x_grok_req_id = request.x_grok_req_id.as_deref().unwrap_or_default(); let model_id = request.inner.model.clone().unwrap_or_default(); // The trace field is process-local: it is consumed by upstream // session code (which may upload a payload artifact) and is not // forwarded by the sampler. Drop it before we send. request.trace.take(); tracing::debug!("create_response: {:?}", &request); tracing::debug!("endpoint: {:?}", self.endpoint("responses")); let grok_headers = GrokRequestHeaders { conv_id: x_grok_conv_id, req_id: x_grok_req_id, model_id: &model_id, session_id: request.x_grok_session_id.as_deref().unwrap_or_default(), turn_idx: request.x_grok_turn_idx.as_deref(), agent_id: request.x_grok_agent_id.as_deref().unwrap_or_default(), deployment_id: request.x_grok_deployment_id.as_deref(), user_id: request.x_grok_user_id.as_deref(), }; let mut request_body = serde_json::to_value(&request.inner).map_err(|e| { tracing::error!("Failed to serialize responses request: {}", e); SamplingError::Serialization(e) })?; // async-openai's ReasoningTextContent struct omits the `type` // discriminator that the Responses API requires on input. Patch // it in post-serialize. This is the last surviving piece of the // old raw_output machinery. xai_grok_sampling_types::patch_reasoning_text_types(&mut request_body); let http_request = grok_headers .apply(self.post(self.endpoint("responses"))) .json(&request_body); let response = http_request.send().await.map_err(|e| { tracing::debug!("HTTP request failed: {}", e); e })?; let status = response.status(); let model_metadata = extract_model_metadata(response.headers()); let retry_after_secs = extract_retry_after(response.headers()); let should_retry = extract_should_retry(response.headers()); let bytes = response.bytes().await?; if !status.is_success() { if status == reqwest::StatusCode::UNAUTHORIZED { self.record_401_attribution(crate::attribution::SamplingConsumer::Responses); let endpoint = self.endpoint("responses"); let server_message = parse_error_bytes(bytes.as_ref()); return Err(SamplingError::Auth(format!( "Unauthorized (401) from {endpoint}: {server_message}" ))); } let req_headers = self.format_request_headers(x_grok_conv_id, x_grok_req_id, &model_id, false); let server_message = parse_error_bytes(bytes.as_ref()); let message = self.build_api_error_message( status, &server_message, &self.endpoint("responses"), &req_headers, None, ); tracing::warn!( status = %status, error_message = %message, model_id = %model_id, "responses API error" ); return Err(SamplingError::Api { status, message, model_metadata, retry_after_secs, should_retry, }); } let response_obj = serde_json::from_slice::(&bytes).map_err(|e| { let raw_body = String::from_utf8_lossy(&bytes); tracing::error!( error = %e, raw_body = %raw_body, "Failed to deserialize rs::Response" ); SamplingError::Serialization(e) })?; Ok(response_obj) } /// Create a streaming response using the Responses API. /// /// Returns a stream of `rs::ResponseStreamEvent` which includes events like: /// - `response.created` - Initial response object /// - `response.output_text.delta` - Text content deltas /// - `response.function_call_arguments.delta` - Function call argument deltas /// - `response.completed` - Final response with all output /// /// The third tuple element is a per-request doom-loop signal collector, /// `Some` only when `SamplerConfig::doom_loop_recovery` is set — the same /// gate that adds the opt-in `x-grok-doom-loop-check` request header, so /// header and parse protection cannot drift apart. It is filled by the /// SSE decoder as the server reports triggers and is meant to be handed /// to `stream_responses` so the signals land on the final /// `ConversationResponse`. #[tracing::instrument( name = "http.create_response_stream", skip_all, fields( endpoint = %self.endpoint("responses"), model_id = request.inner.model.as_deref().unwrap_or(""), status_code = tracing::field::Empty, success = tracing::field::Empty, error = tracing::field::Empty, ) )] #[allow(clippy::type_complexity)] pub async fn create_response_stream( &self, mut request: CreateResponseWrapper, ) -> Result<( BoxStream<'static, Result>, Option, Option, )> { self.apply_response_defaults(&mut request)?; // Enable streaming request.inner.stream = Some(true); let x_grok_conv_id = request.x_grok_conv_id.as_deref().unwrap_or_default(); let x_grok_req_id = request.x_grok_req_id.as_deref().unwrap_or_default(); let model_id = request.inner.model.clone().unwrap_or_default(); // Drop process-local trace data (see note in `create_response`). request.trace.take(); tracing::debug!( base_url = %self.base_url, model_id = model_id.as_str(), "Sending responses API stream request" ); let grok_headers = GrokRequestHeaders { conv_id: x_grok_conv_id, req_id: x_grok_req_id, model_id: &model_id, session_id: request.x_grok_session_id.as_deref().unwrap_or_default(), turn_idx: request.x_grok_turn_idx.as_deref(), agent_id: request.x_grok_agent_id.as_deref().unwrap_or_default(), deployment_id: request.x_grok_deployment_id.as_deref(), user_id: request.x_grok_user_id.as_deref(), }; let extra_raw_tools = std::mem::take(&mut request.extra_raw_tools); let mut request_body = serde_json::to_value(&request.inner).map_err(|e| { tracing::error!("Failed to serialize responses request: {}", e); SamplingError::Serialization(e) })?; // Inject xAI-specific fields not in async-openai's CreateResponse type. if self.defaults.stream_tool_calls { request_body["stream_tool_calls"] = serde_json::json!(true); } // Inject xAI-specific tools (e.g., x_search) that can't be expressed // via async_openai's rs::Tool enum. if !extra_raw_tools.is_empty() { if let Some(tools) = request_body.get_mut("tools").and_then(|v| v.as_array_mut()) { tools.extend(extra_raw_tools); } else { request_body["tools"] = serde_json::Value::Array(extra_raw_tools); } } xai_grok_sampling_types::patch_reasoning_text_types(&mut request_body); // Fresh per attempt so signals never leak across retries; `None` // (check disabled) sends no header and does no peek work per event. let doom_loop = self .defaults .doom_loop_recovery .map(crate::doom_loop::DoomLoopSignalCollector::new); let mut http_request = grok_headers .apply(self.post(self.endpoint("responses"))) .header(ACCEPT, HeaderValue::from_static("text/event-stream")); if doom_loop.is_some() { // Presence opts in; the server ignores the value. http_request = http_request.header(DOOM_LOOP_CHECK_HEADER, "true"); } let http_request = http_request.json(&request_body); let built_request = http_request.build().map_err(|e| { tracing::error!("Failed to build HTTP request: {}", e); SamplingError::Http(e) })?; tracing::debug!( url = %built_request.url(), method = %built_request.method(), "Sending responses API stream request" ); Self::log_request_headers(&built_request, "responses"); let response = self.http.execute(built_request).await.map_err(|e| { tracing::debug!("HTTP request failed: {}", e); record_stream_request_failure(&e); e })?; let status = response.status(); let span = tracing::Span::current(); span.record("status_code", status.as_u16() as i64); span.record("success", status.is_success()); if !status.is_success() { if status == reqwest::StatusCode::UNAUTHORIZED { span.record("error", "unauthorized (401)"); self.record_401_attribution(crate::attribution::SamplingConsumer::ResponsesStream); let endpoint = self.endpoint("responses"); let server_message = response.text().await.unwrap_or_default(); return Err(SamplingError::Auth(format!( "Unauthorized (401) from {endpoint}: {server_message}" ))); } let model_metadata = extract_model_metadata(response.headers()); let retry_after_secs = extract_retry_after(response.headers()); let should_retry = extract_should_retry(response.headers()); let req_headers = self.format_request_headers(x_grok_conv_id, x_grok_req_id, &model_id, true); let resp_headers = Self::format_response_headers(&response); let bytes = response.bytes().await?; let server_message = parse_error_bytes(bytes.as_ref()); let message = self.build_api_error_message( status, &server_message, &self.endpoint("responses"), &req_headers, Some(&resp_headers), ); span.record("error", message.as_str()); tracing::error!( status = %status, error_message = %message, model_id = %model_id, "responses API error" ); return Err(SamplingError::Api { status, message, model_metadata, retry_after_secs, should_retry, }); } let model_metadata = extract_model_metadata(response.headers()); // Strip UTF-8 BOM if present const UTF8_BOM: &[u8] = &[0xEF, 0xBB, 0xBF]; let mut is_first = true; let byte_stream = response.bytes_stream().map(move |result| { result.map(|bytes| { if is_first { is_first = false; if bytes.starts_with(UTF8_BOM) { return bytes.slice(UTF8_BOM.len()..); } } bytes }) }); // Turn raw bytes into SSE events let event_stream = byte_stream.eventsource(); let doom_loop_for_stream = doom_loop.clone(); // The scan item is an `Option`: `Some(None)` skips an absorbed // doom-loop event without terminating the stream (`filter_map` // below), while an outer `None` still ends it. let events = event_stream .scan(false, move |had_transport_error, event_res| { if *had_transport_error { return std::future::ready(None); } let item = match event_res { Ok(event) => { let data = &event.data; if data == "[DONE]" { return std::future::ready(None); } tracing::info!( target: crate::sampling_log::TARGET, event = "sse_chunk", backend = "responses", data = %data, ); // Intercept the non-standard doom-loop event before // typed deserialization; async-openai's event enum // does not know it and would fail to parse it. With // the check disabled, the shared name-or-payload-type // predicate guards against a server emitting it // despite no opt-in (rollout skew), named or not. let swallow = match &doom_loop_for_stream { Some(collector) => collector.absorb(&event.event, data), None => is_check_event(&event.event, data), }; if swallow { Some(None) } else if let Some(stream_error) = try_parse_stream_error(data) { Some(Some(Err(stream_error))) } else { Some(Some(deserialize_response_event(data))) } } Err(e) => { *had_transport_error = true; Some(Some(Err(SamplingError::EventStreamError(e.to_string())))) } }; std::future::ready(item) }) .filter_map(std::future::ready) .boxed(); Ok((events, model_metadata, doom_loop)) } // ========================================================================= // Anthropic Messages API // ========================================================================= /// Apply default configuration to a Messages API request. fn apply_message_defaults(&self, request: &mut MessagesRequestWrapper) -> Result<()> { // Apply model default if not specified if request.inner.model.is_empty() { request.inner.model = self.defaults.model.clone(); } if request.inner.max_tokens == 0 { request.inner.max_tokens = self .defaults .max_completion_tokens .unwrap_or(ANTHROPIC_DEFAULT_MAX_TOKENS); } // Apply temperature default if not specified if request.inner.temperature.is_none() { request.inner.temperature = self.defaults.temperature; } // Apply top_p default if not specified if request.inner.top_p.is_none() { request.inner.top_p = self.defaults.top_p; } Ok(()) } /// Create a message using the Anthropic Messages API (non-streaming). pub async fn create_message( &self, mut request: MessagesRequestWrapper, ) -> Result { self.apply_message_defaults(&mut request)?; let x_grok_conv_id = request.x_grok_conv_id.as_deref().unwrap_or_default(); let x_grok_req_id = request.x_grok_req_id.as_deref().unwrap_or_default(); let model_id = request.inner.model.clone(); // Drop process-local trace data. request.trace.take(); tracing::debug!("create_message: {:?}", &request.inner); tracing::debug!("endpoint: {:?}", self.endpoint("messages")); let grok_headers = GrokRequestHeaders { conv_id: x_grok_conv_id, req_id: x_grok_req_id, model_id: &model_id, session_id: request.x_grok_session_id.as_deref().unwrap_or_default(), turn_idx: request.x_grok_turn_idx.as_deref(), agent_id: request.x_grok_agent_id.as_deref().unwrap_or_default(), deployment_id: request.x_grok_deployment_id.as_deref(), user_id: request.x_grok_user_id.as_deref(), }; let http_request = grok_headers .apply(self.post(self.endpoint("messages"))) .json(&request.inner); let response = http_request.send().await.map_err(|e| { tracing::debug!("HTTP request failed: {}", e); e })?; let status = response.status(); let model_metadata = extract_model_metadata(response.headers()); let retry_after_secs = extract_retry_after(response.headers()); let should_retry = extract_should_retry(response.headers()); let bytes = response.bytes().await?; if !status.is_success() { if status == reqwest::StatusCode::UNAUTHORIZED { self.record_401_attribution(crate::attribution::SamplingConsumer::Messages); let endpoint = self.endpoint("messages"); let server_message = parse_error_bytes(bytes.as_ref()); return Err(SamplingError::Auth(format!( "Unauthorized (401) from {endpoint}: {server_message}" ))); } let req_headers = self.format_request_headers(x_grok_conv_id, x_grok_req_id, &model_id, false); let server_message = parse_error_bytes(bytes.as_ref()); let message = self.build_api_error_message( status, &server_message, &self.endpoint("messages"), &req_headers, None, ); tracing::warn!( status = %status, error_message = %message, model_id = %model_id, "messages API error" ); return Err(SamplingError::Api { status, message, model_metadata, retry_after_secs, should_retry, }); } let response_obj = serde_json::from_slice::(&bytes).map_err(|e| { let raw_body = String::from_utf8_lossy(&bytes); tracing::error!( error = %e, raw_body = %raw_body, "Failed to deserialize MessagesResponse" ); SamplingError::Serialization(e) })?; Ok(response_obj) } /// Create a streaming message using the Anthropic Messages API. /// /// Returns a stream of `MessageStreamEvent` which includes events like: /// - `message_start` - Initial message object /// - `content_block_start` / `content_block_delta` / `content_block_stop` - Content blocks /// - `message_delta` / `message_stop` - Final message with stop reason #[tracing::instrument( name = "http.create_message_stream", skip_all, fields( endpoint = %self.endpoint("messages"), model_id = request.inner.model.as_str(), status_code = tracing::field::Empty, success = tracing::field::Empty, error = tracing::field::Empty, ) )] pub async fn create_message_stream( &self, mut request: MessagesRequestWrapper, ) -> Result<( BoxStream<'static, Result>, Option, )> { self.apply_message_defaults(&mut request)?; // Enable streaming request.inner.stream = Some(true); let x_grok_conv_id = request.x_grok_conv_id.as_deref().unwrap_or_default(); let x_grok_req_id = request.x_grok_req_id.as_deref().unwrap_or_default(); let model_id = request.inner.model.clone(); // Drop process-local trace data. request.trace.take(); tracing::debug!( base_url = %self.base_url, model_id = model_id.as_str(), "Sending Messages API stream request" ); let grok_headers = GrokRequestHeaders { conv_id: x_grok_conv_id, req_id: x_grok_req_id, model_id: &model_id, session_id: request.x_grok_session_id.as_deref().unwrap_or_default(), turn_idx: request.x_grok_turn_idx.as_deref(), agent_id: request.x_grok_agent_id.as_deref().unwrap_or_default(), deployment_id: request.x_grok_deployment_id.as_deref(), user_id: request.x_grok_user_id.as_deref(), }; let http_request = grok_headers .apply(self.post(self.endpoint("messages"))) .header(ACCEPT, HeaderValue::from_static("text/event-stream")) .json(&request.inner); let built_request = http_request.build().map_err(|e| { tracing::error!("Failed to build HTTP request: {}", e); SamplingError::Http(e) })?; tracing::debug!( url = %built_request.url(), method = %built_request.method(), "Sending messages API stream request" ); Self::log_request_headers(&built_request, "messages"); let response = self.http.execute(built_request).await.map_err(|e| { tracing::debug!("HTTP request failed: {}", e); record_stream_request_failure(&e); e })?; let status = response.status(); let span = tracing::Span::current(); span.record("status_code", status.as_u16() as i64); span.record("success", status.is_success()); if !status.is_success() { if status == reqwest::StatusCode::UNAUTHORIZED { span.record("error", "unauthorized (401)"); self.record_401_attribution(crate::attribution::SamplingConsumer::MessagesStream); let endpoint = self.endpoint("messages"); let server_message = response.text().await.unwrap_or_default(); return Err(SamplingError::Auth(format!( "Unauthorized (401) from {endpoint}: {server_message}" ))); } let model_metadata = extract_model_metadata(response.headers()); let retry_after_secs = extract_retry_after(response.headers()); let should_retry = extract_should_retry(response.headers()); let req_headers = self.format_request_headers(x_grok_conv_id, x_grok_req_id, &model_id, true); let resp_headers = Self::format_response_headers(&response); let bytes = response.bytes().await?; let server_message = parse_error_bytes(bytes.as_ref()); let message = self.build_api_error_message( status, &server_message, &self.endpoint("messages"), &req_headers, Some(&resp_headers), ); span.record("error", message.as_str()); tracing::error!( status = %status, error_message = %message, model_id = %model_id, "messages API error" ); return Err(SamplingError::Api { status, message, model_metadata, retry_after_secs, should_retry, }); } let model_metadata = extract_model_metadata(response.headers()); // Strip UTF-8 BOM if present const UTF8_BOM: &[u8] = &[0xEF, 0xBB, 0xBF]; let mut is_first = true; let byte_stream = response.bytes_stream().map(move |result| { result.map(|bytes| { if is_first { is_first = false; if bytes.starts_with(UTF8_BOM) { return bytes.slice(UTF8_BOM.len()..); } } bytes }) }); // Turn raw bytes into SSE events let event_stream = byte_stream.eventsource(); // Map SSE events into MessageStreamEvent. // Uses `scan` so transport errors terminate the stream after the first // error (same pattern as `chat_completion_stream`). let events = event_stream .scan(false, |had_transport_error, event_res| { if *had_transport_error { return std::future::ready(None); } let item = match event_res { Ok(event) => { let data = &event.data; if data == "[DONE]" { return std::future::ready(None); } tracing::info!( target: crate::sampling_log::TARGET, event = "sse_chunk", backend = "messages", data = %data, ); if let Some(stream_error) = try_parse_stream_error(data) { Some(Err(stream_error)) } else { Some( serde_json::from_str::(data).map_err( |e| { tracing::error!( error = %e, raw_data = %data, "Failed to deserialize MessageStreamEvent from stream" ); SamplingError::Serialization(e) }, ), ) } } Err(e) => { *had_transport_error = true; Some(Err(SamplingError::EventStreamError(e.to_string()))) } }; std::future::ready(item) }) .boxed(); Ok((events, model_metadata)) } // ========================================================================= // Unified Conversation API // ========================================================================= /// Apply default configuration to a ConversationRequest. fn apply_conversation_defaults(&self, request: &mut ConversationRequest) -> Result<()> { if request.model.is_none() { request.model = Some(self.defaults.model.clone()); } if request.temperature.is_none() { request.temperature = self.defaults.temperature; } if request.top_p.is_none() { request.top_p = self.defaults.top_p; } if request.max_output_tokens.is_none() { request.max_output_tokens = self.defaults.max_completion_tokens; } Ok(()) } /// Send a conversation request using the Chat Completions API (streaming). /// /// Converts the `ConversationRequest` to `ChatCompletionRequest` internally. /// Returns the stream and any model metadata extracted from response headers. pub async fn conversation_stream( &self, mut request: ConversationRequest, ) -> Result<( BoxStream<'static, Result>, Option, )> { self.apply_conversation_defaults(&mut request)?; let trace = request.trace.take(); let mut chat_request: ChatCompletionRequest = request.into(); if let Some(trace) = trace { chat_request.trace = Some(trace); } self.chat_completion_stream(chat_request).await } /// Send a conversation request using the Chat Completions API (non-streaming). /// /// Converts the `ConversationRequest` to `ChatCompletionRequest` internally. pub async fn conversation( &self, mut request: ConversationRequest, ) -> Result { self.apply_conversation_defaults(&mut request)?; let trace = request.trace.take(); let mut chat_request: ChatCompletionRequest = request.into(); if let Some(trace) = trace { chat_request.trace = Some(trace); } self.chat_completion(chat_request).await } /// Send a conversation request using the Responses API (streaming). /// /// Converts the `ConversationRequest` to Responses API format internally. /// The third tuple element is the per-request doom-loop signal collector /// (see [`Self::create_response_stream`]); callers that don't consume the /// signals can ignore it. #[allow(clippy::type_complexity)] pub async fn conversation_stream_responses( &self, mut request: ConversationRequest, ) -> Result<( BoxStream<'static, Result>, Option, Option, )> { self.apply_conversation_defaults(&mut request)?; let trace = request.trace.take(); let x_grok_conv_id = request.x_grok_conv_id.clone(); let x_grok_req_id = request.x_grok_req_id.clone(); let x_grok_session_id = request.x_grok_session_id.clone(); let x_grok_turn_idx = request.x_grok_turn_idx.clone(); let x_grok_agent_id = request.x_grok_agent_id.clone(); // Collect xAI-specific tools that can't be expressed via rs::Tool // (e.g., x_search). These are injected as raw JSON after serialization. let extra_tools = xai_grok_sampling_types::extra_raw_tools(&request.hosted_tools); let responses_request: rs::CreateResponse = (&request).into(); let mut wrapper = CreateResponseWrapper::new(responses_request); wrapper.x_grok_conv_id = x_grok_conv_id; wrapper.x_grok_req_id = x_grok_req_id; wrapper.x_grok_session_id = x_grok_session_id; wrapper.x_grok_turn_idx = x_grok_turn_idx; wrapper.x_grok_agent_id = x_grok_agent_id; wrapper.extra_raw_tools = extra_tools; if let Some(trace) = trace { wrapper.trace = Some(trace); } self.create_response_stream(wrapper).await } /// Send a conversation request using the Responses API (non-streaming). /// /// Converts the `ConversationRequest` to Responses API format internally. pub async fn conversation_responses( &self, mut request: ConversationRequest, ) -> Result { self.apply_conversation_defaults(&mut request)?; let trace = request.trace.take(); let x_grok_conv_id = request.x_grok_conv_id.clone(); let x_grok_req_id = request.x_grok_req_id.clone(); let x_grok_session_id = request.x_grok_session_id.clone(); let x_grok_turn_idx = request.x_grok_turn_idx.clone(); let x_grok_agent_id = request.x_grok_agent_id.clone(); let responses_request: rs::CreateResponse = (&request).into(); let mut wrapper = CreateResponseWrapper::new(responses_request); wrapper.x_grok_conv_id = x_grok_conv_id; wrapper.x_grok_req_id = x_grok_req_id; wrapper.x_grok_session_id = x_grok_session_id; wrapper.x_grok_turn_idx = x_grok_turn_idx; wrapper.x_grok_agent_id = x_grok_agent_id; if let Some(trace) = trace { wrapper.trace = Some(trace); } self.create_response(wrapper).await } /// Send a conversation request using the Anthropic Messages API (streaming). /// /// Converts the `ConversationRequest` to Messages API format internally. pub async fn conversation_stream_messages( &self, mut request: ConversationRequest, ) -> Result<( BoxStream<'static, Result>, Option, )> { self.apply_conversation_defaults(&mut request)?; let trace = request.trace.take(); let x_grok_conv_id = request.x_grok_conv_id.clone(); let x_grok_req_id = request.x_grok_req_id.clone(); let x_grok_session_id = request.x_grok_session_id.clone(); let x_grok_turn_idx = request.x_grok_turn_idx.clone(); let x_grok_agent_id = request.x_grok_agent_id.clone(); let messages_request = build_messages_request(&request); let mut wrapper = MessagesRequestWrapper::new(messages_request); wrapper.x_grok_conv_id = x_grok_conv_id; wrapper.x_grok_req_id = x_grok_req_id; wrapper.x_grok_session_id = x_grok_session_id; wrapper.x_grok_turn_idx = x_grok_turn_idx; wrapper.x_grok_agent_id = x_grok_agent_id; if let Some(trace) = trace { wrapper.trace = Some(trace); } self.create_message_stream(wrapper).await } /// Send a conversation request using the Anthropic Messages API (non-streaming). /// /// Converts the `ConversationRequest` to Messages API format internally. pub async fn conversation_messages( &self, mut request: ConversationRequest, ) -> Result { self.apply_conversation_defaults(&mut request)?; let trace = request.trace.take(); let x_grok_conv_id = request.x_grok_conv_id.clone(); let x_grok_req_id = request.x_grok_req_id.clone(); let x_grok_session_id = request.x_grok_session_id.clone(); let x_grok_turn_idx = request.x_grok_turn_idx.clone(); let x_grok_agent_id = request.x_grok_agent_id.clone(); let messages_request = build_messages_request(&request); let mut wrapper = MessagesRequestWrapper::new(messages_request); wrapper.x_grok_conv_id = x_grok_conv_id; wrapper.x_grok_req_id = x_grok_req_id; wrapper.x_grok_session_id = x_grok_session_id; wrapper.x_grok_turn_idx = x_grok_turn_idx; wrapper.x_grok_agent_id = x_grok_agent_id; if let Some(trace) = trace { wrapper.trace = Some(trace); } self.create_message(wrapper).await } /// Backend-aware streaming call that collects the full response. pub async fn conversation_collect( &self, request: ConversationRequest, ) -> Result { let request_id = crate::types::RequestId::random(); let idle_timeout = std::time::Duration::from_secs(300); let result = match self.api_backend() { ApiBackend::ChatCompletions => { let (raw, meta) = self.conversation_stream(request).await?; let events = crate::stream::stream_chat_completions(raw, meta, request_id, idle_timeout); crate::stream::collect_response(events).await } ApiBackend::Responses => { let (raw, meta, doom_loop) = self.conversation_stream_responses(request).await?; let events = crate::stream::stream_responses(raw, meta, request_id, idle_timeout, doom_loop); crate::stream::collect_response(events).await } ApiBackend::Messages => { let (raw, meta) = self.conversation_stream_messages(request).await?; let events = crate::stream::stream_messages(raw, meta, request_id, idle_timeout); crate::stream::collect_response(events).await } }; result .map(|(response, _metrics)| response) .map_err(|info| SamplingError::Api { status: info .status_code .and_then(|c| reqwest::StatusCode::from_u16(c).ok()) .unwrap_or(reqwest::StatusCode::INTERNAL_SERVER_ERROR), message: info.message, model_metadata: info.model_metadata, retry_after_secs: info.retry_after_secs, should_retry: None, }) } } #[cfg(test)] mod tests { use super::*; use indexmap::IndexMap; use xai_grok_sampling_types::types::ChatRequestMessage; fn minimal_config() -> SamplerConfig { SamplerConfig { api_key: Some("test-key".to_string()), base_url: "https://example.test".to_string(), model: "test-model".to_string(), max_completion_tokens: None, temperature: None, top_p: None, api_backend: ApiBackend::ChatCompletions, auth_scheme: AuthScheme::Bearer, extra_headers: IndexMap::new(), context_window: 8192, force_http1: false, max_retries: None, stream_tool_calls: false, idle_timeout_secs: None, reasoning_effort: None, origin_client: None, client_identifier: None, deployment_id: None, user_id: None, client_version: None, attribution_callback: None, bearer_resolver: None, supports_backend_search: false, compactions_remaining: None, compaction_at_tokens: None, doom_loop_recovery: None, header_injector: None, } } /// Verify the serialized shape of StreamingChatRequest matches the /// expected wire format: all ChatCompletionRequest fields flattened at /// top level, plus `stream: true` and `stream_options.include_usage: true`. #[test] fn streaming_chat_request_serializes_correctly() { let request = ChatCompletionRequest { model: Some("test-model".into()), messages: vec![ChatRequestMessage::user("hello")], temperature: Some(0.7), max_tokens: None, top_p: None, frequency_penalty: None, presence_penalty: None, user: None, tools: None, tool_choice: None, search_parameters: None, response_format: None, reasoning_effort: None, x_grok_conv_id: None, x_grok_req_id: None, x_grok_session_id: None, x_grok_turn_idx: None, x_grok_agent_id: None, x_grok_deployment_id: None, x_grok_user_id: None, trace: None, }; let wrapper = StreamingChatRequest { inner: &request, stream: true, stream_options: StreamOptions { include_usage: true, }, }; let json: serde_json::Value = serde_json::to_value(&wrapper).unwrap(); let obj = json.as_object().unwrap(); assert_eq!(obj.get("stream").and_then(|v| v.as_bool()), Some(true)); assert_eq!( obj.get("stream_options") .and_then(|v| v.get("include_usage")) .and_then(|v| v.as_bool()), Some(true) ); assert!( obj.get("inner").is_none(), "inner field should be flattened" ); assert_eq!( obj.get("model").and_then(|v| v.as_str()), Some("test-model") ); assert!(obj.get("messages").is_some()); let temp = obj.get("temperature").and_then(|v| v.as_f64()).unwrap(); assert!((temp - 0.7).abs() < 0.001, "temperature should be ~0.7"); assert!(obj.get("max_tokens").is_none()); assert!(obj.get("tools").is_none()); } #[test] fn extract_retry_after_parses_seconds() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert(reqwest::header::RETRY_AFTER, "30".parse().unwrap()); assert_eq!(extract_retry_after(&headers), Some(30)); } #[test] fn extract_retry_after_caps_at_120() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert(reqwest::header::RETRY_AFTER, "3600".parse().unwrap()); assert_eq!(extract_retry_after(&headers), Some(120)); } #[test] fn extract_retry_after_zero_is_valid() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert(reqwest::header::RETRY_AFTER, "0".parse().unwrap()); assert_eq!(extract_retry_after(&headers), Some(0)); } #[test] fn extract_retry_after_ignores_http_date() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert( reqwest::header::RETRY_AFTER, "Fri, 31 Dec 2025 23:59:59 GMT".parse().unwrap(), ); assert_eq!(extract_retry_after(&headers), None); } #[test] fn extract_retry_after_none_when_missing() { let headers = reqwest::header::HeaderMap::new(); assert_eq!(extract_retry_after(&headers), None); } #[test] fn extract_should_retry_true() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert("x-should-retry", "true".parse().unwrap()); assert_eq!(extract_should_retry(&headers), Some(true)); } #[test] fn extract_should_retry_true_case_insensitive() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert("x-should-retry", "TRUE".parse().unwrap()); assert_eq!(extract_should_retry(&headers), Some(true)); } #[test] fn extract_should_retry_false() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert("x-should-retry", "false".parse().unwrap()); assert_eq!(extract_should_retry(&headers), Some(false)); } #[test] fn extract_should_retry_unknown_value_is_none() { let mut headers = reqwest::header::HeaderMap::new(); headers.insert("x-should-retry", "banana".parse().unwrap()); assert_eq!(extract_should_retry(&headers), None); } #[test] fn extract_should_retry_absent_is_none() { let headers = reqwest::header::HeaderMap::new(); assert_eq!(extract_should_retry(&headers), None); } #[test] fn new_with_minimal_config_succeeds() { let client = SamplingClient::new(minimal_config()).expect("client should construct"); assert_eq!(client.api_backend(), ApiBackend::ChatCompletions); } #[test] fn new_applies_extra_headers() { let mut cfg = minimal_config(); cfg.extra_headers .insert("x-test-header".to_string(), "test-value".to_string()); cfg.extra_headers .insert("x-XAI-token-auth".to_string(), "xai-grok-cli".to_string()); let _client = SamplingClient::new(cfg).expect("client with extra headers should construct"); } #[test] fn messages_plus_anthropic_api_key_uses_x_api_key_and_not_authorization() { let cfg = SamplerConfig { api_key: Some("anthropic-key-abc123".to_string()), api_backend: ApiBackend::Messages, auth_scheme: AuthScheme::XApiKey, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); assert!( client .default_headers .get(HeaderName::from_static("x-api-key")) .is_some() ); assert!(client.default_headers.get(AUTHORIZATION).is_none()); } #[test] fn messages_plus_bearer_uses_authorization_and_not_x_api_key() { let cfg = SamplerConfig { api_key: Some("bearer-key-abc123".to_string()), api_backend: ApiBackend::Messages, auth_scheme: AuthScheme::Bearer, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); assert!(client.default_headers.get(AUTHORIZATION).is_some()); assert!( client .default_headers .get(HeaderName::from_static("x-api-key")) .is_none() ); } // Regression: a past change dropped User-Agent from sampling requests. #[test] fn sampling_client_always_has_user_agent() { let client = SamplingClient::new(minimal_config()).expect("build"); assert!(client.default_headers.contains_key(USER_AGENT)); } // Regression: a past change dropped HeaderInjector (traceparent) from sampling requests. #[test] fn header_injector_is_called_in_post() { #[derive(Debug)] struct TestInjector; impl crate::config::HeaderInjector for TestInjector { fn inject(&self, headers: &mut HeaderMap) { headers.insert( HeaderName::from_static("traceparent"), HeaderValue::from_static("00-test-trace-id-00"), ); } } let mut config = minimal_config(); config.header_injector = Some(std::sync::Arc::new(TestInjector)); let client = SamplingClient::new(config).expect("build"); let req = client .post("http://localhost/test") .build() .expect("build request"); assert!( req.headers().contains_key("traceparent"), "HeaderInjector should inject traceparent into post() requests" ); } #[test] fn user_agent_includes_origin_and_agent_product() { let origin = OriginClientInfo { product: "my-client".to_string(), version: Some("1.2.3".to_string()), }; let ua = user_agent_string_for(&origin); assert!(ua.contains("my-client/1.2.3")); assert!(ua.contains(AGENT_PRODUCT)); } #[test] fn user_agent_omits_origin_version_when_absent() { let origin = OriginClientInfo { product: "my-client".to_string(), version: None, }; let ua = user_agent_string_for(&origin); // No slash between product and the grok-shell agent product. assert!(ua.starts_with("my-client grok-shell/")); } #[test] fn user_agent_collapses_when_origin_matches_agent() { let agent_version = xai_grok_version::VERSION.to_string(); let origin = OriginClientInfo { product: AGENT_PRODUCT.to_string(), version: Some(agent_version.clone()), }; let ua = user_agent_string_for(&origin); // Single product/version slot when the origin and agent match. assert!(ua.starts_with(&format!("{}/{}", AGENT_PRODUCT, agent_version))); } /// Counts callbacks for assertions in the tests below. #[derive(Default, Debug)] struct CountingCallback { invocations: std::sync::Mutex)>>, } #[derive(Debug)] struct StaticBearerResolver(&'static str); impl crate::config::BearerResolver for StaticBearerResolver { fn current_bearer(&self) -> Option { Some(self.0.to_string()) } } impl crate::attribution::Auth401AttributionCallback for CountingCallback { fn record_401( &self, consumer: crate::attribution::SamplingConsumer, sent_bearer: Option<&str>, ) { self.invocations .lock() .unwrap() .push((consumer, sent_bearer.map(|s| s.to_string()))); } } /// `extract_sent_bearer` strips the `"Bearer "` prefix off /// `Authorization` for OpenAI-completions backends and truncates the /// remaining bearer to the cross-crate prefix length. #[test] fn extract_sent_bearer_strips_bearer_prefix_for_openai_compat() { let cfg = SamplerConfig { api_key: Some("test-bearer-1234567890".to_string()), api_backend: ApiBackend::ChatCompletions, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); let bearer = client.extract_sent_bearer(); // Bearer is truncated at the crate boundary -- callers // downstream of this method only ever see the prefix. assert_eq!(bearer.as_deref(), Some("test-bearer-")); assert_eq!( bearer.as_deref().map(str::len), Some(crate::attribution::SENT_BEARER_PREFIX_LEN), ); } /// `extract_sent_bearer` reads `x-api-key` for Anthropic Messages API /// and truncates the value to the cross-crate prefix length. #[test] fn extract_sent_bearer_reads_x_api_key_for_messages() { let cfg = SamplerConfig { api_key: Some("anthropic-key-abc123".to_string()), api_backend: ApiBackend::Messages, auth_scheme: AuthScheme::XApiKey, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); let bearer = client.extract_sent_bearer(); assert_eq!(bearer.as_deref(), Some("anthropic-ke")); assert_eq!( bearer.as_deref().map(str::len), Some(crate::attribution::SENT_BEARER_PREFIX_LEN), ); } /// `extract_sent_bearer` returns `None` when no auth header is set. #[test] fn extract_sent_bearer_returns_none_when_no_header() { let cfg = SamplerConfig { api_key: None, api_backend: ApiBackend::ChatCompletions, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); assert!(client.extract_sent_bearer().is_none()); } #[test] fn live_bearer_resolver_uses_authorization_for_messages_plus_bearer() { let cfg = SamplerConfig { api_key: Some("stale-bearer".to_string()), api_backend: ApiBackend::Messages, auth_scheme: AuthScheme::Bearer, bearer_resolver: Some(std::sync::Arc::new(StaticBearerResolver("fresh-bearer"))), ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); let request = client .post("https://example.test/v1/messages") .build() .expect("request should build"); let auth = request .headers() .get(AUTHORIZATION) .and_then(|v| v.to_str().ok()); assert_eq!(auth, Some("Bearer fresh-bearer")); assert!(request.headers().get("x-api-key").is_none()); } /// Regression: when `api_key` (which seeds `default_headers` with an /// `Authorization: Bearer ...`) AND a `bearer_resolver` are both set, /// `post()` must produce **exactly one** `Authorization` header on the /// wire. The pre-fix code used `RequestBuilder::header(AUTHORIZATION, ...)` /// which appends rather than replaces, causing two identical /// `Authorization` headers and a 400 from cli-chat-proxy. #[test] fn post_emits_single_authorization_with_api_key_and_bearer_resolver() { let cfg = SamplerConfig { api_key: Some("stale-bearer".to_string()), api_backend: ApiBackend::Responses, auth_scheme: AuthScheme::Bearer, bearer_resolver: Some(std::sync::Arc::new(StaticBearerResolver("fresh-bearer"))), ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); let request = client .post("https://example.test/v1/responses") .build() .expect("request should build"); let auth_count = request.headers().get_all(AUTHORIZATION).iter().count(); assert_eq!( auth_count, 1, "expected exactly one Authorization header, got {auth_count}" ); assert_eq!( request .headers() .get(AUTHORIZATION) .and_then(|v| v.to_str().ok()), Some("Bearer fresh-bearer"), ); } #[test] fn live_bearer_resolver_uses_x_api_key_for_messages_plus_anthropic_api_key() { let cfg = SamplerConfig { api_key: Some("stale-anthropic".to_string()), api_backend: ApiBackend::Messages, auth_scheme: AuthScheme::XApiKey, bearer_resolver: Some(std::sync::Arc::new(StaticBearerResolver("fresh-anthropic"))), ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); let request = client .post("https://example.test/v1/messages") .build() .expect("request should build"); let api_key = request .headers() .get("x-api-key") .and_then(|v| v.to_str().ok()); assert_eq!(api_key, Some("fresh-anthropic")); assert!(request.headers().get(AUTHORIZATION).is_none()); } /// Bearers shorter than the prefix length pass through unchanged. /// Defensive against the truncation logic inadvertently widening /// short bearers (no panics, no zero-padding). #[test] fn extract_sent_bearer_short_bearer_passes_through_unchanged() { let cfg = SamplerConfig { api_key: Some("abc".to_string()), api_backend: ApiBackend::ChatCompletions, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); assert_eq!(client.extract_sent_bearer().as_deref(), Some("abc")); } /// `record_401_attribution` invokes the wired callback with the /// expected `consumer` and the truncated bearer prefix that the /// wire would carry. The key assertion is that the callback /// receives the prefix only -- the full bearer never crosses the /// crate boundary. #[test] fn record_401_attribution_invokes_callback_with_extracted_bearer() { let cb = std::sync::Arc::new(CountingCallback::default()); let cb_dyn: crate::attribution::SharedAttributionCallback = cb.clone(); let cfg = SamplerConfig { api_key: Some("the-bearer-1234567890-extra-tail".to_string()), api_backend: ApiBackend::ChatCompletions, attribution_callback: Some(cb_dyn), bearer_resolver: None, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); client.record_401_attribution(crate::attribution::SamplingConsumer::ChatCompletionsStream); let calls = cb.invocations.lock().unwrap(); assert_eq!(calls.len(), 1); assert_eq!( calls[0].0, crate::attribution::SamplingConsumer::ChatCompletionsStream ); // Prefix-only -- the `extra-tail` portion of the bearer is // dropped by `extract_sent_bearer` before the callback fires. assert_eq!(calls[0].1.as_deref(), Some("the-bearer-1")); assert_eq!( calls[0].1.as_deref().map(str::len), Some(crate::attribution::SENT_BEARER_PREFIX_LEN), ); } /// Regression test: when a bearer_resolver is wired, `post()` must /// *replace* the Authorization header from `default_headers`, not /// append a second one. Duplicate Authorization headers cause /// Cloudflare to return 400 Bad Request. #[test] fn bearer_resolver_replaces_authorization_header() { #[derive(Debug)] struct StaticResolver(String); impl crate::config::BearerResolver for StaticResolver { fn current_bearer(&self) -> Option { Some(self.0.clone()) } } let resolver: crate::config::SharedBearerResolver = std::sync::Arc::new(StaticResolver("fresh-token".to_string())); let cfg = SamplerConfig { api_key: Some("stale-token".to_string()), api_backend: ApiBackend::Responses, bearer_resolver: Some(resolver), ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); // Build a request to inspect the final headers. let builder = client.post("https://example.test/v1/responses"); let request = builder.body("").build().expect("request should build"); let auth_values: Vec<_> = request.headers().get_all(AUTHORIZATION).iter().collect(); assert_eq!( auth_values.len(), 1, "expected exactly one Authorization header, got {}: {:?}", auth_values.len(), auth_values ); assert_eq!( auth_values[0].to_str().unwrap(), "Bearer fresh-token", "Authorization header should contain the resolver's fresh token" ); } /// `record_401_attribution` is a no-op when `attribution_callback` /// is `None` (the BYOK / sampler-only path). The previous tests /// in this module construct clients without a callback and rely /// on this property holding. #[test] fn record_401_attribution_is_noop_without_callback() { let cfg = SamplerConfig { api_key: Some("bearer".to_string()), api_backend: ApiBackend::ChatCompletions, attribution_callback: None, bearer_resolver: None, ..minimal_config() }; let client = SamplingClient::new(cfg).expect("client should build"); // Must not panic. client.record_401_attribution(crate::attribution::SamplingConsumer::ChatCompletions); } /// `response.completed` carrying /// `usage.context_details.{input_tokens, output_tokens}` rewrites /// `usage.total_tokens` in place to the live context length /// (`ctx.input + ctx.output`). Billing fields stay on the wire's /// cumulative values. #[test] fn deserialize_response_event_overrides_total_tokens_from_context_details() { let sse = r#"{ "type": "response.completed", "sequence_number": 0, "response": { "id": "resp_1", "object": "response", "created_at": 0, "model": "grok-build", "status": "completed", "output": [], "usage": { "input_tokens": 6003, "input_tokens_details": { "cached_tokens": 1984 }, "output_tokens": 711, "output_tokens_details": { "reasoning_tokens": 388 }, "total_tokens": 6714, "context_details": { "input_tokens": 5022, "output_tokens": 571 } } } }"#; let event = deserialize_response_event(sse).expect("parse"); let rs::ResponseStreamEvent::ResponseCompleted(e) = event else { panic!("expected ResponseCompleted"); }; let usage = e.response.usage.expect("usage present"); // Billing fields stay cumulative — unchanged by context_details. assert_eq!(usage.input_tokens, 6003); assert_eq!(usage.output_tokens, 711); assert_eq!(usage.input_tokens_details.cached_tokens, 1984); assert_eq!(usage.output_tokens_details.reasoning_tokens, 388); // total_tokens rewritten to ctx.input + ctx.output (5022 + 571). // NOT the wire's cumulative total (6714). assert_eq!(usage.total_tokens, 5_593); } #[test] fn deserialize_response_event_stashes_cost_in_metadata() { let make = |ticks: i64| { format!( r#"{{ "type": "response.completed", "sequence_number": 0, "response": {{ "id": "resp_1", "object": "response", "created_at": 0, "model": "grok-build", "status": "completed", "output": [], "usage": {{ "input_tokens": 10, "input_tokens_details": {{ "cached_tokens": 0 }}, "output_tokens": 5, "output_tokens_details": {{ "reasoning_tokens": 0 }}, "total_tokens": 15, "cost_in_usd_ticks": {ticks} }} }} }}"# ) }; let event = deserialize_response_event(&make(78)).expect("parse"); let rs::ResponseStreamEvent::ResponseCompleted(e) = event else { panic!("expected ResponseCompleted"); }; assert_eq!( e.response .metadata .as_ref() .and_then(|m| m.get(COST_USD_TICKS_METADATA_KEY)) .map(String::as_str), Some("78") ); // The REST mapper backfills 0 for unbilled requests: no stash. let event = deserialize_response_event(&make(0)).expect("parse"); let rs::ResponseStreamEvent::ResponseCompleted(e) = event else { panic!("expected ResponseCompleted"); }; assert!(e.response.metadata.is_none()); } #[test] fn deserialize_response_event_total_tokens_unchanged_when_context_details_absent() { // Older / non-Responses backends omit `context_details`. // `total_tokens` passes through from the wire unchanged. let sse = r#"{ "type": "response.completed", "sequence_number": 0, "response": { "id": "resp_1", "object": "response", "created_at": 0, "model": "grok-build", "status": "completed", "output": [], "usage": { "input_tokens": 10000, "input_tokens_details": { "cached_tokens": 0 }, "output_tokens": 100, "output_tokens_details": { "reasoning_tokens": 0 }, "total_tokens": 10100 } } }"#; let event = deserialize_response_event(sse).expect("parse"); let rs::ResponseStreamEvent::ResponseCompleted(e) = event else { panic!("expected ResponseCompleted"); }; let usage = e.response.usage.expect("usage present"); assert_eq!(usage.total_tokens, 10_100); } #[test] fn deserialize_response_event_total_tokens_unchanged_when_context_details_partial() { // Defensive: if the backend ever ships only one of the two // context_details fields, we don't have a complete picture of // the live context size, so leave `total_tokens` on the wire's // cumulative value instead of guessing (treating the missing // half as 0 would silently under-report). let sse = r#"{ "type": "response.completed", "sequence_number": 0, "response": { "id": "resp_1", "object": "response", "created_at": 0, "model": "grok-build", "status": "completed", "output": [], "usage": { "input_tokens": 6003, "input_tokens_details": { "cached_tokens": 1984 }, "output_tokens": 711, "output_tokens_details": { "reasoning_tokens": 388 }, "total_tokens": 6714, "context_details": { "input_tokens": 5022 } } } }"#; let event = deserialize_response_event(sse).expect("parse"); let rs::ResponseStreamEvent::ResponseCompleted(e) = event else { panic!("expected ResponseCompleted"); }; let usage = e.response.usage.expect("usage present"); assert_eq!(usage.total_tokens, 6_714); } #[test] fn deserialize_response_event_ignores_context_details_on_non_terminal_events() { // Non-terminal events don't carry final usage; even if the backend ever // echoed `context_details` on one, we don't touch it. let sse = r#"{ "type": "response.output_text.delta", "sequence_number": 0, "item_id": "item-1", "output_index": 0, "content_index": 0, "delta": "hello", "logprobs": [] }"#; let event = deserialize_response_event(sse).expect("non-terminal event parses"); assert!(matches!( event, rs::ResponseStreamEvent::ResponseOutputTextDelta(_) )); } }