From 9cfcbc0e0d7ef55e74d0889089e90b4e71119f3b Mon Sep 17 00:00:00 2001 From: Anthony Ronning <101225832+AnthonyRonning@users.noreply.github.com> Date: Fri, 28 Aug 2026 12:55:24 +0000 Subject: [PATCH 1/2] feat: bound Maple inference capacity replay --- frontend/src-tauri/src/agent.rs | 36 +- .../src-tauri/src/agent/developer_tools.rs | 8 +- frontend/src-tauri/src/agent/provider.rs | 1414 +++++++++++++---- .../src-tauri/src/agent/shell_permission.rs | 12 +- .../src-tauri/src/agent/web_permission.rs | 12 +- frontend/src-tauri/src/maple_api.rs | 12 +- frontend/src/components/UnifiedChat.tsx | 94 +- .../src/services/agentThoughtLabels.test.ts | 4 +- frontend/src/services/agentThoughtLabels.ts | 2 - .../services/inferenceCapacityRetry.test.ts | 142 ++ .../src/services/inferenceCapacityRetry.ts | 41 + sdk/opensecret-integration-revision | 2 +- sdk/rust/Cargo.toml | 2 +- sdk/rust/src/client.rs | 203 ++- sdk/rust/src/lib.rs | 5 +- sdk/src/lib/ai.ts | 119 +- sdk/src/lib/api.ts | 2 +- sdk/src/lib/index.ts | 8 +- sdk/src/lib/main.tsx | 4 +- sdk/src/lib/test/customFetch.test.ts | 299 +++- 20 files changed, 2005 insertions(+), 416 deletions(-) create mode 100644 frontend/src/services/inferenceCapacityRetry.test.ts create mode 100644 frontend/src/services/inferenceCapacityRetry.ts diff --git a/frontend/src-tauri/src/agent.rs b/frontend/src-tauri/src/agent.rs index d08b984cb..5851ee6c9 100644 --- a/frontend/src-tauri/src/agent.rs +++ b/frontend/src-tauri/src/agent.rs @@ -112,7 +112,6 @@ const SESSION_TITLE_GENERATION_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(1_500); const SESSION_TITLE_MODEL: &str = "llama3-3-70b"; const SESSION_TITLE_TEMPERATURE: f32 = 0.7; -const SESSION_TITLE_MAX_TOKENS: i32 = 15; const SESSION_TITLE_MAX_INPUT_CHARS: usize = 500; const SESSION_TITLE_SYSTEM_PROMPT: &str = "You are a helpful assistant that generates concise, meaningful titles (3-5 words) for chat conversations based on the user's first message. Return only the title without quotes or explanations."; const DEFAULT_AGENT_SESSION_TITLE: &str = "New task"; @@ -1785,7 +1784,7 @@ async fn generate_agent_session_title( model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(SESSION_TITLE_TEMPERATURE)) - .with_max_tokens(Some(SESSION_TITLE_MAX_TOKENS)); + .with_max_tokens(None); let bounded_prompt = first_prompt .chars() .take(SESSION_TITLE_MAX_INPUT_CHARS) @@ -6763,6 +6762,7 @@ fn maple_model_config( // Maple's authoritative catalog value is per session. Explicitly clear any // process-global Goose context override when metadata is unavailable. model_config.context_limit = context_limit.filter(|limit| *limit > 0); + provider::clear_output_token_limits(&mut model_config); Ok(model_config) } @@ -6796,11 +6796,12 @@ async fn install_maple_provider_config( agent: &Arc, transport: &Arc, session_id: &str, - model_config: goose_providers::model::ModelConfig, + mut model_config: goose_providers::model::ModelConfig, ) -> Result<(), String> where T: provider::MapleInferenceTransport + 'static, { + provider::clear_output_token_limits(&mut model_config); let provider = Arc::new(MapleProvider::new(Arc::clone(transport))); agent .update_provider(provider, model_config, session_id) @@ -10187,6 +10188,7 @@ mod tests { async fn send_inference_request( self: Arc, _request: opensecret::InferenceRequest, + _send_budget: opensecret::InferenceSendBudget, _cancel_token: CancellationToken, ) -> opensecret::Result { Err(opensecret::Error::Other( @@ -10208,6 +10210,7 @@ mod tests { async fn send_inference_request( self: Arc, _request: opensecret::InferenceRequest, + _send_budget: opensecret::InferenceSendBudget, _cancel_token: CancellationToken, ) -> opensecret::Result { match self.0 { @@ -11879,7 +11882,12 @@ mod tests { .unwrap(); let persisted_model_config = goose_providers::model::ModelConfig::new("gemma-3-27b") .with_context_limit(Some(64_321)) - .with_temperature(Some(0.42)); + .with_temperature(Some(0.42)) + .with_max_tokens(Some(4_096)) + .with_merged_request_params(HashMap::from([ + ("max_output_tokens".to_string(), json!(4_096)), + ("include_reasoning".to_string(), json!(false)), + ])); session_manager .update(&session.id) .provider_name(MAPLE_PROVIDER_NAME) @@ -11931,6 +11939,11 @@ mod tests { .and_then(|model| model.temperature), Some(0.42) ); + let restored_config = persisted.model_config.as_ref().unwrap(); + assert_eq!(restored_config.max_tokens, None); + let restored_params = restored_config.request_params.as_ref().unwrap(); + assert!(!restored_params.contains_key("max_output_tokens")); + assert_eq!(restored_params["include_reasoning"], false); drop(manager_result); drop(agent_manager); @@ -13204,6 +13217,19 @@ mod tests { ); } + #[test] + fn maple_model_config_omits_goose_canonical_output_limits() { + let canonical = + ModelConfig::new("deepseek-v4-flash").with_canonical_limits(MAPLE_PROVIDER_NAME); + assert!(canonical.max_tokens.is_some()); + + let config = + maple_model_config("deepseek-v4-flash", Some(1_048_576)).expect("Maple model config"); + assert_eq!(config.model_name, "deepseek-v4-flash"); + assert_eq!(config.context_limit, Some(1_048_576)); + assert_eq!(config.max_tokens, None); + } + #[test] fn agent_session_model_locks_after_first_message() { assert!(validate_session_model_lock(0, Some("glm-5-2"), "gemma4-31b").is_ok()); @@ -17554,7 +17580,7 @@ mod tests { .expect("the provider should capture one title request"); assert_eq!(capture.model_name, SESSION_TITLE_MODEL); assert_eq!(capture.temperature, Some(SESSION_TITLE_TEMPERATURE)); - assert_eq!(capture.max_tokens, Some(SESSION_TITLE_MAX_TOKENS)); + assert_eq!(capture.max_tokens, None); assert_eq!(capture.reasoning, Some(false)); assert!(!capture.request_params_present); assert_eq!(capture.system, SESSION_TITLE_SYSTEM_PROMPT); diff --git a/frontend/src-tauri/src/agent/developer_tools.rs b/frontend/src-tauri/src/agent/developer_tools.rs index 56bc1f513..4cd8412f0 100644 --- a/frontend/src-tauri/src/agent/developer_tools.rs +++ b/frontend/src-tauri/src/agent/developer_tools.rs @@ -61,7 +61,6 @@ const IMAGE_DOWNLOAD_TIMEOUT: Duration = Duration::from_secs(30); const IMAGE_DESCRIPTION_TIMEOUT: Duration = Duration::from_secs(60); const IMAGE_DESCRIPTION_MODEL: &str = "gemma4-31b"; const IMAGE_DESCRIPTION_TEMPERATURE: f32 = 0.0; -const IMAGE_DESCRIPTION_MAX_TOKENS: i32 = 2_048; const IMAGE_DESCRIPTION_CONTEXT_MAX_CHARS: usize = 12_000; pub(super) const EXTERNAL_MCP_TOOL_NAME: &str = "external_mcp"; const IMAGE_DESCRIPTION_SYSTEM_PROMPT: &str = r#"You are the visual perception helper for a coding agent that cannot inspect images directly. @@ -746,7 +745,7 @@ async fn describe_image_for_text_model( model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(IMAGE_DESCRIPTION_TEMPERATURE)) - .with_max_tokens(Some(IMAGE_DESCRIPTION_MAX_TOKENS)); + .with_max_tokens(None); let prompt = contextual_image_prompt(source, image_context); let messages = [Message::user() @@ -3152,7 +3151,7 @@ mod tests { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(IMAGE_DESCRIPTION_TEMPERATURE)) - .with_max_tokens(Some(IMAGE_DESCRIPTION_MAX_TOKENS)); + .with_max_tokens(None); let messages = [Message::user() .with_text(contextual_image_prompt( "icon.png", @@ -3174,7 +3173,8 @@ mod tests { let payload = captured.lock().unwrap().take().unwrap(); assert_eq!(payload["model"], IMAGE_DESCRIPTION_MODEL); assert_eq!(payload["temperature"], IMAGE_DESCRIPTION_TEMPERATURE); - assert_eq!(payload["max_tokens"], IMAGE_DESCRIPTION_MAX_TOKENS); + assert!(payload.get("max_tokens").is_none()); + assert!(payload.get("max_completion_tokens").is_none()); assert_eq!(payload["stream"], true); assert_eq!(payload["include_reasoning"], false); assert_eq!(payload["chat_template_kwargs"]["enable_thinking"], false); diff --git a/frontend/src-tauri/src/agent/provider.rs b/frontend/src-tauri/src/agent/provider.rs index 2824f2e89..9450bc774 100644 --- a/frontend/src-tauri/src/agent/provider.rs +++ b/frontend/src-tauri/src/agent/provider.rs @@ -10,22 +10,28 @@ use goose_providers::formats::openai::{ use goose_providers::images::ImageFormat; use goose_providers::model::ModelConfig; use goose_providers::request_log::{start_log, LoggerHandleExt}; -use goose_providers::retry::{ - should_retry, RetryConfig, DEFAULT_BACKOFF_MULTIPLIER, DEFAULT_INITIAL_RETRY_INTERVAL_MS, - DEFAULT_MAX_RETRY_INTERVAL_MS, +use goose_providers::retry::RetryConfig; +use opensecret::{ + InferenceRequest, InferenceResponse, InferenceSendBudget, OpenSecretClient, + OpenSecretResponseBody, }; -use opensecret::{InferenceRequest, InferenceResponse, OpenSecretClient, OpenSecretResponseBody}; use rmcp::model::Tool; use serde_json::{json, Value}; use std::cell::Cell; -use std::future::{ready, Future}; -use std::sync::Arc; -use std::time::{Duration, SystemTime}; +use std::collections::VecDeque; +use std::future::Future; +use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, +}; +use std::time::Duration; use tokio_util::codec::{FramedRead, LinesCodec, LinesCodecError}; use tokio_util::io::StreamReader; use tokio_util::sync::CancellationToken; const CHAT_COMPLETIONS_PATH: &str = "/v1/chat/completions"; +const OUTPUT_TOKEN_LIMIT_FIELDS: [&str; 3] = + ["max_tokens", "max_completion_tokens", "max_output_tokens"]; pub(super) const MAPLE_PROVIDER_NAME: &str = "maple"; const AUTHENTICATION_ERROR_MESSAGE: &str = "Maple authentication failed"; pub(super) const ATTESTATION_VERIFICATION_ERROR_MESSAGE: &str = @@ -34,12 +40,14 @@ pub(super) const SECURE_CONNECTION_ERROR_MESSAGE: &str = "Maple's encrypted connection could not be recovered"; const ERROR_CONTRACT_HEADER: &str = "x-opensecret-error-contract"; const ERROR_CODE_HEADER: &str = "x-opensecret-error-code"; +const CLIENT_REPLAY_HEADER: &str = "x-opensecret-client-replay"; const ERROR_CONTRACT_VERSION: &[u8] = b"1"; const SESSION_NOT_FOUND_ERROR_CODE: &[u8] = b"session_not_found"; - -/// Extra transient inference attempts after the first failure. Goose's 1s × 2^n -/// backoff, capped at 30s, then covers roughly three minutes of provider blips. -const TRANSIENT_MAX_RETRIES: usize = 10; +const INFERENCE_CAPACITY_ERROR_CODE: &[u8] = b"inference_capacity"; +const CLIENT_REPLAY_SAFE: &[u8] = b"safe"; +const INFERENCE_CAPACITY_ERROR_MESSAGE: &str = "Inference capacity is temporarily unavailable"; +const STREAM_ENDED_BEFORE_COMPLETION_MESSAGE: &str = + "Maple's response stream ended before completion"; const KIMI_K3_MODEL_ID: &str = "kimi-k3"; // Agent Mode forwards the selected catalog ID unchanged for direct model // selections. Keep Gemma's provider-specific opt-in scoped to that explicit @@ -47,7 +55,9 @@ const KIMI_K3_MODEL_ID: &str = "kimi-k3"; const GEMMA4_AGENT_MODEL_ID: &str = "gemma4-31b"; const MAX_ERROR_BODY_BYTES: usize = 16 * 1024; const MAX_STREAM_LINE_BYTES: usize = 16 * 1024 * 1024; -const MAX_RETRY_AFTER_SECS: f64 = 3_600.0; +const DEFAULT_CAPACITY_RETRY_DELAY: Duration = Duration::from_secs(1); +const MAX_CAPACITY_RETRY_DELAY_SECS: u64 = 60; +const MAX_LOGICAL_INFERENCE_SENDS: usize = 2; #[cfg(not(test))] const RESPONSE_START_TIMEOUT: Duration = Duration::from_secs(300); #[cfg(test)] @@ -104,6 +114,11 @@ fn remember_terminal_run_error(error: &ProviderError) { ProviderError::ExecutionError(message) if message == SECURE_CONNECTION_ERROR_MESSAGE => { SECURE_CONNECTION_ERROR_MESSAGE } + ProviderError::ExecutionError(message) + if message == STREAM_ENDED_BEFORE_COMPLETION_MESSAGE => + { + STREAM_ENDED_BEFORE_COMPLETION_MESSAGE + } _ => return, }; @@ -133,6 +148,7 @@ pub(crate) trait MapleInferenceTransport: Send + Sync { async fn send_inference_request( self: Arc, request: InferenceRequest, + send_budget: InferenceSendBudget, cancel_token: CancellationToken, ) -> opensecret::Result; } @@ -145,6 +161,7 @@ impl MapleInferenceTransport for OpenSecretClient { async fn send_inference_request( self: Arc, request: InferenceRequest, + send_budget: InferenceSendBudget, cancel_token: CancellationToken, ) -> opensecret::Result { tokio::select! { @@ -152,15 +169,28 @@ impl MapleInferenceTransport for OpenSecretClient { _ = cancel_token.cancelled() => { Err(opensecret::Error::Other("Inference request was cancelled".to_string())) } - response = OpenSecretClient::send_inference_request(&self, request) => response, + response = OpenSecretClient::send_inference_request_with_budget( + &self, + request, + send_budget, + ) => response, } } } pub(crate) struct MapleProvider { transport: Arc, - #[cfg(test)] - test_retry_config: Option, +} + +pub(super) fn clear_output_token_limits(model_config: &mut ModelConfig) { + // Goose can restore or materialize its own output defaults for Maple's + // sessions and auxiliary calls. Maple uses the model's context window. + model_config.max_tokens = None; + if let Some(params) = model_config.request_params.as_mut() { + for field in OUTPUT_TOKEN_LIMIT_FIELDS { + params.remove(field); + } + } } impl MapleProvider { @@ -168,17 +198,7 @@ impl MapleProvider { where T: MapleInferenceTransport + 'static, { - Self { - transport, - #[cfg(test)] - test_retry_config: None, - } - } - - #[cfg(test)] - fn with_test_retry_config(mut self, retry_config: RetryConfig) -> Self { - self.test_retry_config = Some(retry_config); - self + Self { transport } } fn build_request( @@ -188,7 +208,7 @@ impl MapleProvider { messages: &[Message], tools: &[Tool], ) -> Result { - create_request_with_options( + let mut request = create_request_with_options( model_config, system, messages, @@ -202,7 +222,16 @@ impl MapleProvider { ) .map_err(|error| { ProviderError::RequestFailed(format!("Failed to create Maple request: {error}")) - }) + })?; + // Keep this policy at Maple's final Agent request boundary too: resumed + // sessions and Goose fast-model calls can bypass our config constructors. + // Generic SDK/API caller requests do not pass through this provider. + if let Some(fields) = request.as_object_mut() { + for field in OUTPUT_TOKEN_LIMIT_FIELDS { + fields.remove(field); + } + } + Ok(request) } fn gemma_agent_model_config( @@ -281,6 +310,7 @@ impl MapleProvider { async fn send_attempt( &self, request: InferenceRequest, + send_budget: InferenceSendBudget, cancellation: &CancellationToken, ) -> Result { // The transport owns authentication reconciliation and must get a chance @@ -288,8 +318,11 @@ impl MapleProvider { // take too long. Cancelling and then awaiting the transport future keeps // a rotated SDK JWT from being stranded only in native memory. let transport_cancellation = cancellation.child_token(); - let response = Arc::clone(&self.transport) - .send_inference_request(request, transport_cancellation.clone()); + let response = Arc::clone(&self.transport).send_inference_request( + request, + send_budget, + transport_cancellation.clone(), + ); tokio::pin!(response); let response_start_timeout = tokio::time::sleep(RESPONSE_START_TIMEOUT); tokio::pin!(response_start_timeout); @@ -353,30 +386,135 @@ impl MapleProvider { LinesCodec::new_with_max_length(MAX_STREAM_LINE_BYTES), ) .map_err(anyhow::Error::from); - let parsed = response_to_streaming_message(lines); - - Box::pin(parsed.map(move |result| { - result.map_err(|error| { - if parser_cancellation.is_cancelled() { - cancellation_error() - } else if let Some(error) = secure_connection_stream_error(&error) { - remember_terminal_run_error(&error); - error - } else { - invalid_stream_error() + let saw_done = Arc::new(AtomicBool::new(false)); + let guarded_lines = futures_util::stream::unfold((lines, false, false), { + let saw_done = Arc::clone(&saw_done); + move |(mut lines, stream_saw_done, finished)| { + let saw_done = Arc::clone(&saw_done); + async move { + if finished { + return None; + } + + match lines.next().await { + Some(Ok(line)) => { + let line_is_done = is_done_sse_line(&line); + if line_is_done { + saw_done.store(true, Ordering::Release); + } + Some((Ok(line), (lines, stream_saw_done || line_is_done, false))) + } + Some(Err(error)) => Some((Err(error), (lines, stream_saw_done, true))), + None if stream_saw_done => None, + None => Some(( + Err(anyhow::Error::new(MissingDoneMarker)), + (lines, stream_saw_done, true), + )), + } } - }) - })) + } + }); + let parsed: MessageStream = Box::pin( + response_to_streaming_message(Box::pin(guarded_lines)).map(move |result| { + result.map_err(|error| { + if parser_cancellation.is_cancelled() { + cancellation_error() + } else if let Some(error) = secure_connection_stream_error(&error) { + remember_terminal_run_error(&error); + error + } else if missing_done_marker(&error) { + stream_ended_before_completion_error() + } else { + invalid_stream_error() + } + }) + }), + ); + + Self::enforce_stream_terminal_contract(parsed, saw_done) + } + + fn enforce_stream_terminal_contract( + parsed: MessageStream, + saw_done: Arc, + ) -> MessageStream { + // Goose independently retries empty model turns even when the provider + // retry budget is zero. It also executes a parsed tool request before + // polling the following stream item. Buffer tool requests until DONE, + // and turn any otherwise-empty or unterminated stream into a terminal + // error so neither path can replay or dispatch accepted work. + Box::pin(futures_util::stream::unfold( + (parsed, VecDeque::new(), false, false), + move |(mut parsed, mut buffered, produced_output, finished)| { + let saw_done = Arc::clone(&saw_done); + async move { + if let Some(item) = buffered.pop_front() { + return Some((item, (parsed, buffered, produced_output, finished))); + } + if finished { + return None; + } + + let mut produced_output = produced_output; + loop { + match parsed.next().await { + Some(Ok(item)) => { + produced_output |= stream_item_has_agent_turn_output(&item); + if !buffered.is_empty() + || (stream_item_contains_tool_request(&item) + && !saw_done.load(Ordering::Acquire)) + { + buffered.push_back(Ok(item)); + continue; + } + return Some(( + Ok(item), + (parsed, buffered, produced_output, false), + )); + } + Some(Err(error)) => { + buffered.clear(); + return Some(( + Err(error), + (parsed, buffered, produced_output, true), + )); + } + None if !buffered.is_empty() && saw_done.load(Ordering::Acquire) => { + let item = buffered + .pop_front() + .expect("buffer was checked as non-empty"); + return Some((item, (parsed, buffered, produced_output, true))); + } + None if produced_output && saw_done.load(Ordering::Acquire) => { + return None; + } + None => { + buffered.clear(); + return Some(( + Err(stream_ended_before_completion_error()), + (parsed, buffered, produced_output, true), + )); + } + } + } + } + }, + )) } async fn stream_attempt( &self, payload_bytes: &[u8], + send_budget: InferenceSendBudget, cancellation: &CancellationToken, - ) -> Result { - let request = Self::inference_request(payload_bytes.to_vec())?; - let response = self.send_attempt(request, cancellation).await?; - let response = ensure_success(response).await?; + ) -> Result { + let request = Self::inference_request(payload_bytes.to_vec()) + .map_err(StreamAttemptFailure::without_replay)?; + let response = self + .send_attempt(request, send_budget, cancellation) + .await + .map_err(StreamAttemptFailure::without_replay)?; + let response = ensure_success_for_attempt(response).await?; Ok(Self::message_stream_from_response( response, cancellation.clone(), @@ -388,59 +526,37 @@ impl MapleProvider { payload_bytes: &[u8], cancellation: &CancellationToken, ) -> Result { - let config = Provider::retry_config(self); - let mut attempts = 0; - - loop { - let error = match self.stream_attempt(payload_bytes, cancellation).await { - Ok(mut stream) => { - // TODO(upstream): Remove this Maple-specific bridge from Agent Mode once - // Maple's pinned Goose revision provides equivalent first-item handling: - // https://github.com/aaif-goose/goose/issues/10887 - // If auxiliary complete() calls still need this protection, scope it to - // that path instead. Recovery after any successful item remains out of - // scope here: - // https://github.com/aaif-goose/goose/issues/10897 - let first = tokio::select! { - biased; - _ = cancellation.cancelled() => return Err(cancellation_error()), - first = stream.next() => first, - }; - match first { - Some(Ok(first)) => { - return Ok(Box::pin( - futures_util::stream::once(ready(Ok(first))).chain(stream), - )); - } - Some(Err(error)) => error, - None => return Ok(stream), - } - } - Err(error) => error, - }; + let send_budget = InferenceSendBudget::new(MAX_LOGICAL_INFERENCE_SENDS) + .expect("the fixed inference send budget must be non-zero"); + let first_failure = match self + .stream_attempt(payload_bytes, send_budget.clone(), cancellation) + .await + { + Ok(stream) => return Ok(stream), + Err(failure) => failure, + }; + let Some(delay) = first_failure.replay_delay else { + remember_terminal_run_error(&first_failure.error); + return Err(first_failure.error); + }; + if send_budget.remaining() == 0 { + return Err(first_failure.error); + } - if !should_retry(&error, &config) || attempts >= config.max_retries() { - remember_terminal_run_error(&error); - return Err(error); - } - attempts += 1; - let delay = match &error { - ProviderError::RateLimitExceeded { - retry_delay: Some(provider_delay), - .. - } => *provider_delay, - _ => config.delay_for_attempt(attempts), - }; - let skip_backoff = std::env::var("GOOSE_PROVIDER_SKIP_BACKOFF") - .unwrap_or_default() - .parse::() - .unwrap_or(false); - if !skip_backoff { - tokio::select! { - biased; - _ = cancellation.cancelled() => return Err(cancellation_error()), - _ = tokio::time::sleep(delay) => {} - } + tokio::select! { + biased; + _ = cancellation.cancelled() => return Err(cancellation_error()), + _ = tokio::time::sleep(delay) => {} + } + + match self + .stream_attempt(payload_bytes, send_budget, cancellation) + .await + { + Ok(stream) => Ok(stream), + Err(failure) => { + remember_terminal_run_error(&failure.error); + Err(failure.error) } } } @@ -453,8 +569,9 @@ impl MapleProvider { tools: &[Tool], enable_primary_agent_thinking: bool, ) -> Result { - let effective_model_config = + let mut effective_model_config = Self::gemma_agent_model_config(model_config, enable_primary_agent_thinking); + clear_output_token_limits(&mut effective_model_config); let payload = self.build_request(&effective_model_config, system, messages, tools)?; let payload_bytes = serde_json::to_vec(&payload).map_err(|error| { ProviderError::RequestFailed(format!("Failed to serialize Maple request: {error}")) @@ -496,21 +613,10 @@ impl Provider for MapleProvider { } fn retry_config(&self) -> RetryConfig { - #[cfg(test)] - if let Some(config) = &self.test_retry_config { - return config.clone(); - } - - // Retrying deterministic client failures can repeat side effects and - // causes the SDK to repeat its own stale-session recovery for a 400. One - // shared transient budget covers both setup and pre-first-item failures. - RetryConfig::new( - TRANSIENT_MAX_RETRIES, - DEFAULT_INITIAL_RETRY_INTERVAL_MS, - DEFAULT_BACKOFF_MULTIPLIER, - DEFAULT_MAX_RETRY_INTERVAL_MS, - ) - .transient_only() + // Maple owns the sole replay, and only for OpenSecret's explicit + // pre-acceptance capacity contract. Goose must never add another replay + // for network, HTTP, or pre-first-item stream errors. + RetryConfig::new(0, 0, 1.0, 0).transient_only() } async fn stream( @@ -549,35 +655,108 @@ fn invalid_stream_error() -> ProviderError { ProviderError::NetworkError("Maple's response stream was invalid".to_string()) } -async fn ensure_success(response: InferenceResponse) -> Result { +fn stream_ended_before_completion_error() -> ProviderError { + let error = ProviderError::ExecutionError(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string()); + remember_terminal_run_error(&error); + error +} + +fn stream_item_contains_tool_request(item: &(Option, Option)) -> bool { + item.0.as_ref().is_some_and(Message::is_tool_call) +} + +fn stream_item_has_agent_turn_output(item: &(Option, Option)) -> bool { + item.0.as_ref().is_some_and(|message| { + message.metadata.output_token_limit_reached + || message.content.iter().any(|content| match content { + MessageContent::Text(text) => !text.text.is_empty(), + MessageContent::Image(image) => !image.data.is_empty(), + MessageContent::Thinking(thinking) => { + !thinking.thinking.is_empty() || !thinking.signature.is_empty() + } + MessageContent::RedactedThinking(thinking) => !thinking.data.is_empty(), + MessageContent::SystemNotification(notification) => !notification.msg.is_empty(), + _ => true, + }) + }) +} + +struct StreamAttemptFailure { + error: ProviderError, + replay_delay: Option, +} + +impl StreamAttemptFailure { + fn without_replay(error: ProviderError) -> Self { + Self { + error, + replay_delay: None, + } + } +} + +#[derive(Debug)] +struct MissingDoneMarker; + +impl std::fmt::Display for MissingDoneMarker { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE) + } +} + +impl std::error::Error for MissingDoneMarker {} + +fn missing_done_marker(error: &anyhow::Error) -> bool { + error + .chain() + .any(|cause| cause.downcast_ref::().is_some()) +} + +fn is_done_sse_line(line: &str) -> bool { + line.strip_prefix("data: ") + .or_else(|| line.strip_prefix("data:")) + .is_some_and(|payload| payload.trim() == "[DONE]") +} + +async fn ensure_success_for_attempt( + response: InferenceResponse, +) -> Result { if response.status().is_success() { return Ok(response); } let status = response.status(); + if has_exact_inference_capacity_contract(status, response.headers()) { + let replay_delay = capacity_retry_delay(response.headers()); + return Err(StreamAttemptFailure { + error: ProviderError::RateLimitExceeded { + details: INFERENCE_CAPACITY_ERROR_MESSAGE.to_string(), + retry_delay: replay_delay, + }, + replay_delay, + }); + } + let terminal_session_failure = has_exact_session_not_found_contract(status, response.headers()); - let retry_after_header = response - .headers() - .get("retry-after") - .and_then(|value| value.to_str().ok()) - .map(str::to_owned); let (_parts, body) = response.into_parts(); - let (body, truncated) = collect_bounded_body(body).await?; + let (body, truncated) = collect_bounded_body(body) + .await + .map_err(StreamAttemptFailure::without_replay)?; let payload = error_payload(&body, truncated); - let retry_delay = retry_after_delay(payload.as_ref(), retry_after_header.as_deref()); let error = if terminal_session_failure { ProviderError::ExecutionError(SECURE_CONNECTION_ERROR_MESSAGE.to_string()) } else { map_http_error(status, payload.as_ref()) }; - match error { - ProviderError::RateLimitExceeded { details, .. } => Err(ProviderError::RateLimitExceeded { - details, - retry_delay, - }), - error => Err(error), - } + Err(StreamAttemptFailure::without_replay(error)) +} + +#[cfg(test)] +async fn ensure_success(response: InferenceResponse) -> Result { + ensure_success_for_attempt(response) + .await + .map_err(|failure| failure.error) } fn has_exact_session_not_found_contract( @@ -603,6 +782,51 @@ fn has_exact_session_not_found_contract( code_values.next().is_none() && code.as_bytes() == SESSION_NOT_FOUND_ERROR_CODE } +fn has_exact_inference_capacity_contract( + status: tauri::http::StatusCode, + headers: &tauri::http::HeaderMap, +) -> bool { + matches!( + status, + tauri::http::StatusCode::TOO_MANY_REQUESTS | tauri::http::StatusCode::SERVICE_UNAVAILABLE + ) && has_exact_header(headers, ERROR_CONTRACT_HEADER, ERROR_CONTRACT_VERSION) + && has_exact_header(headers, ERROR_CODE_HEADER, INFERENCE_CAPACITY_ERROR_CODE) + && has_exact_header(headers, CLIENT_REPLAY_HEADER, CLIENT_REPLAY_SAFE) +} + +fn has_exact_header(headers: &tauri::http::HeaderMap, name: &str, expected: &[u8]) -> bool { + let mut values = headers.get_all(name).iter(); + let Some(value) = values.next() else { + return false; + }; + values.next().is_none() && value.as_bytes() == expected +} + +fn capacity_retry_delay(headers: &tauri::http::HeaderMap) -> Option { + let mut values = headers.get_all(tauri::http::header::RETRY_AFTER).iter(); + let Some(value) = values.next() else { + return Some(DEFAULT_CAPACITY_RETRY_DELAY); + }; + if values.next().is_some() { + return Some(DEFAULT_CAPACITY_RETRY_DELAY); + } + let bytes = value.as_bytes(); + let canonical = bytes == b"0" + || (bytes + .first() + .is_some_and(|byte| matches!(byte, b'1'..=b'9')) + && bytes.iter().all(u8::is_ascii_digit)); + if !canonical { + return Some(DEFAULT_CAPACITY_RETRY_DELAY); + } + + let seconds = value + .to_str() + .ok() + .and_then(|value| value.parse::().ok())?; + (seconds <= MAX_CAPACITY_RETRY_DELAY_SECS).then(|| Duration::from_secs(seconds)) +} + async fn collect_bounded_body( mut body: OpenSecretResponseBody, ) -> Result<(Vec, bool), ProviderError> { @@ -663,38 +887,6 @@ fn error_payload(body: &[u8], truncated: bool) -> Option { Some(json!({ "message": message })) } -fn retry_after_delay(payload: Option<&Value>, header: Option<&str>) -> Option { - let body_seconds = payload - .and_then(|payload| payload.get("error")) - .and_then(|error| error.get("metadata")) - .and_then(|metadata| metadata.get("retry_after_seconds")) - .and_then(Value::as_f64); - body_seconds - .and_then(retry_duration_from_seconds) - .or_else(|| header.and_then(parse_retry_after_header)) -} - -fn retry_duration_from_seconds(seconds: f64) -> Option { - if !seconds.is_finite() || seconds < 0.0 { - return None; - } - - Some(Duration::from_secs_f64(seconds.min(MAX_RETRY_AFTER_SECS))) -} - -fn parse_retry_after_header(value: &str) -> Option { - let value = value.trim(); - if let Ok(seconds) = value.parse::() { - return retry_duration_from_seconds(seconds as f64); - } - - let retry_at = httpdate::parse_http_date(value).ok()?; - let delay = retry_at - .duration_since(SystemTime::now()) - .unwrap_or(Duration::ZERO); - retry_duration_from_seconds(delay.as_secs_f64()) -} - fn map_http_error(status: tauri::http::StatusCode, payload: Option<&Value>) -> ProviderError { log::warn!( "Maple inference request failed (http_status_{})", @@ -992,6 +1184,7 @@ mod tests { use goose_providers::retry::should_retry; use rmcp::object; use std::collections::{HashMap, VecDeque}; + use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Mutex; use tokio::sync::Notify; @@ -1006,17 +1199,23 @@ mod tests { struct FakeTransport { requests: Mutex>, + send_limits: Mutex>, responses: Mutex>>, request_notify: Notify, } struct PendingTransport; + struct DoubleSendCapacityTransport { + calls: AtomicUsize, + } + #[async_trait] impl MapleInferenceTransport for PendingTransport { async fn send_inference_request( self: Arc, _request: InferenceRequest, + _send_budget: InferenceSendBudget, cancel_token: CancellationToken, ) -> opensecret::Result { cancel_token.cancelled().await; @@ -1042,6 +1241,7 @@ mod tests { fn with_results(responses: Vec>) -> Self { Self { requests: Mutex::new(Vec::new()), + send_limits: Mutex::new(Vec::new()), responses: Mutex::new(responses.into()), request_notify: Notify::new(), } @@ -1055,6 +1255,10 @@ mod tests { self.responses.lock().expect("response lock").len() } + fn send_limits(&self) -> Vec { + self.send_limits.lock().expect("send limit lock").clone() + } + async fn wait_for_request_count(&self, expected: usize) { loop { let notified = self.request_notify.notified(); @@ -1071,8 +1275,18 @@ mod tests { async fn send_inference_request( self: Arc, request: InferenceRequest, + send_budget: InferenceSendBudget, _cancel_token: CancellationToken, ) -> opensecret::Result { + self.send_limits + .lock() + .expect("send limit lock") + .push(send_budget.remaining()); + if !send_budget.try_reserve_send() { + return Err(opensecret::Error::Other( + "Inference request send budget exhausted".to_string(), + )); + } let (parts, body) = request.into_parts(); let captured = CapturedRequest { method: parts.method.to_string(), @@ -1095,6 +1309,21 @@ mod tests { } } + #[async_trait] + impl MapleInferenceTransport for DoubleSendCapacityTransport { + async fn send_inference_request( + self: Arc, + _request: InferenceRequest, + send_budget: InferenceSendBudget, + _cancel_token: CancellationToken, + ) -> opensecret::Result { + self.calls.fetch_add(1, Ordering::SeqCst); + assert!(send_budget.try_reserve_send()); + assert!(send_budget.try_reserve_send()); + Ok(capacity_response(503, Some("0"))) + } + } + fn response_with_items( status: u16, items: Vec>>, @@ -1119,6 +1348,53 @@ mod tests { response_with_items(status, chunks.into_iter().map(Ok).collect(), retry_after) } + fn add_capacity_contract(response: &mut InferenceResponse) { + response.headers_mut().insert( + ERROR_CONTRACT_HEADER, + tauri::http::HeaderValue::from_static("1"), + ); + response.headers_mut().insert( + ERROR_CODE_HEADER, + tauri::http::HeaderValue::from_static("inference_capacity"), + ); + response.headers_mut().insert( + CLIENT_REPLAY_HEADER, + tauri::http::HeaderValue::from_static("safe"), + ); + } + + fn capacity_response(status: u16, retry_after: Option<&str>) -> InferenceResponse { + let mut response = response( + status, + vec![br#"{"error":{"message":"private upstream capacity detail"}}"#.to_vec()], + retry_after, + ); + add_capacity_contract(&mut response); + response + } + + fn unpolled_capacity_response( + status: u16, + retry_after: Option<&str>, + body_polls: Arc, + ) -> InferenceResponse { + let body: OpenSecretResponseBody = Box::pin(futures_util::stream::once(async move { + body_polls.fetch_add(1, Ordering::SeqCst); + Ok::<_, opensecret::Error>(b"private upstream capacity detail".to_vec().into()) + })); + let mut response = InferenceResponse::new(body); + *response.status_mut() = + tauri::http::StatusCode::from_u16(status).expect("valid fake status"); + if let Some(retry_after) = retry_after { + response.headers_mut().insert( + tauri::http::header::RETRY_AFTER, + tauri::http::HeaderValue::from_str(retry_after).expect("valid retry header"), + ); + } + add_capacity_contract(&mut response); + response + } + fn error_contract_response(contract: Option<&str>, code: Option<&str>) -> InferenceResponse { let mut response = response( 400, @@ -1181,10 +1457,6 @@ mod tests { event } - fn fast_retry_config(max_retries: usize) -> RetryConfig { - RetryConfig::new(max_retries, 0, 1.0, 0).transient_only() - } - fn fragmented_success_response() -> InferenceResponse { let response_bytes = concat!( "data: {\"id\":\"chunk-1\",\"object\":\"chat.completion.chunk\",", @@ -1213,6 +1485,43 @@ mod tests { ) } + fn single_text_response(done_line: Option<&str>) -> InferenceResponse { + let event = json!({ + "id": "single-text", + "object": "chat.completion.chunk", + "created": 1, + "model": "test-model", + "choices": [{ + "index": 0, + "delta": { "role": "assistant", "content": "Hello" }, + "finish_reason": "stop" + }], + "usage": { "prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2 } + }); + let mut sse = format!("data: {event}\n\n"); + if let Some(done_line) = done_line { + sse.push_str(done_line); + sse.push_str("\n\n"); + } + response(200, vec![sse.into_bytes()], None) + } + + fn usage_only_response() -> InferenceResponse { + let event = json!({ + "id": "usage-only", + "object": "chat.completion.chunk", + "created": 1, + "model": "test-model", + "choices": [], + "usage": { "prompt_tokens": 1, "completion_tokens": 0, "total_tokens": 1 } + }); + response( + 200, + vec![format!("data: {event}\n\ndata: [DONE]\n\n").into_bytes()], + None, + ) + } + fn tool_call_response(completion_id: &str, tool_id: &str) -> InferenceResponse { let tool_chunk = json!({ "id": completion_id, @@ -1287,16 +1596,6 @@ mod tests { response } - fn notifying_malformed_response(error_read: Arc) -> InferenceResponse { - let body: OpenSecretResponseBody = Box::pin(futures_util::stream::once(async move { - error_read.notify_one(); - Ok(b"data: transient-invalid-stream\n\n".to_vec().into()) - })); - let mut response = InferenceResponse::new(body); - *response.status_mut() = tauri::http::StatusCode::OK; - response - } - #[tokio::test] async fn formats_openai_request_and_preserves_images_and_thinking() { let transport = Arc::new(FakeTransport::new(fragmented_success_response())); @@ -1354,6 +1653,72 @@ mod tests { ); } + #[tokio::test] + async fn agent_and_auxiliary_requests_omit_inherited_output_token_limits() { + for auxiliary in [false, true] { + let transport = Arc::new(FakeTransport::new(fragmented_success_response())); + let provider = MapleProvider::new(transport.clone()); + let model_config = ModelConfig::new("deepseek-v4-flash") + .with_canonical_limits(MAPLE_PROVIDER_NAME) + .with_context_limit(Some(1_048_576)) + .with_temperature(Some(0.25)) + .with_merged_request_params(HashMap::from([ + ("max_tokens".to_string(), json!(64)), + ("max_completion_tokens".to_string(), json!(64)), + ("max_output_tokens".to_string(), json!(64)), + ("include_reasoning".to_string(), json!(false)), + ])); + assert!(model_config.max_tokens.is_some()); + let messages = [Message::user().with_text("Keep generating within model context")]; + + // This also exercises the final formatter policy independently of + // stream_request's config cleanup, as for a restored config. + let formatted = provider + .build_request(&model_config, "system", &messages, &[]) + .expect("format Maple request"); + for field in OUTPUT_TOKEN_LIMIT_FIELDS { + assert!(formatted.get(field).is_none()); + } + + if auxiliary { + provider + .complete(&model_config, "system", &messages, &[]) + .await + .expect("auxiliary completion"); + } else { + let stream = provider + .stream(&model_config, "system", &messages, &[]) + .await + .expect("Agent stream"); + collect_stream(stream) + .await + .expect("completed Agent stream"); + } + + let requests = transport.requests.lock().expect("request lock"); + assert_eq!(requests.len(), 1); + let request = &requests[0]; + for field in OUTPUT_TOKEN_LIMIT_FIELDS { + assert!(request.body.get(field).is_none()); + } + assert_eq!(request.body["model"], "deepseek-v4-flash"); + assert_eq!(request.body["temperature"], 0.25); + assert_eq!(request.body["include_reasoning"], false); + assert_eq!( + request.body["messages"][1]["content"], + messages[0].as_concat_text() + ); + assert_eq!(model_config.context_limit, Some(1_048_576)); + assert_eq!(model_config.temperature, Some(0.25)); + assert!(model_config.max_tokens.is_some()); + assert!(model_config + .request_params + .as_ref() + .unwrap() + .contains_key("max_output_tokens")); + } + } + #[tokio::test] async fn primary_agent_stream_enables_thinking_only_for_direct_gemma_selection() { let gemma_transport = Arc::new(FakeTransport::new(fragmented_success_response())); @@ -1668,13 +2033,12 @@ mod tests { } #[tokio::test] - async fn retries_invalid_stream_before_first_item_with_the_same_request() { + async fn malformed_stream_before_first_item_is_terminal_without_replay() { let transport = Arc::new(FakeTransport::queued(vec![ malformed_response("transient-invalid-stream"), fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let stream = provider .stream( @@ -1684,25 +2048,17 @@ mod tests { &[], ) .await - .expect("replacement stream should start"); - let (message, usage) = collect_stream(stream) + .expect("successful response headers should start the stream"); + let error = collect_stream(stream) .await - .expect("replacement stream should parse"); - let text = message - .content - .iter() - .filter_map(|content| match content { - MessageContent::Text(text) => Some(text.text.as_str()), - _ => None, - }) - .collect::(); + .expect_err("malformed stream should fail without replay"); - assert_eq!(text, "Hello world"); - assert_eq!(usage.usage.total_tokens, Some(5)); - let requests = transport.requests.lock().expect("request lock"); - assert_eq!(requests.len(), 2); - assert_eq!(requests[0].raw_body, requests[1].raw_body); - assert_eq!(requests[0].body, requests[1].body); + assert_eq!( + error, + ProviderError::NetworkError("Maple's response stream was invalid".to_string()) + ); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); } #[tokio::test] @@ -1729,8 +2085,7 @@ mod tests { interrupted, fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let mut stream = provider .stream( @@ -1789,8 +2144,7 @@ mod tests { interrupted, fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let mut stream = provider .stream( @@ -1818,22 +2172,74 @@ mod tests { } #[tokio::test] - async fn shares_one_retry_budget_across_status_and_first_item_failures() { - let mut responses = vec![response( - 503, - vec![br#"{"error":{"message":"temporarily unavailable"}}"#.to_vec()], - None, - )]; - responses.extend( - (0..TRANSIENT_MAX_RETRIES) - .map(|index| malformed_response(&format!("invalid-stream-{index}"))), - ); - responses.push(fragmented_success_response()); - let transport = Arc::new(FakeTransport::queued(responses)); + async fn exact_capacity_contract_replays_once_with_the_same_unpolled_request() { + let body_polls = Arc::new(AtomicUsize::new(0)); + let transport = Arc::new(FakeTransport::queued(vec![ + unpolled_capacity_response(503, Some("0"), Arc::clone(&body_polls)), + fragmented_success_response(), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let stream = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("hello")], + &[], + ) + .await + .expect("the replay response should start"); + let (_, usage) = collect_stream(stream) + .await + .expect("the replay response should parse"); + + assert_eq!(usage.usage.total_tokens, Some(5)); + assert_eq!(body_polls.load(Ordering::SeqCst), 0); + let requests = transport.requests.lock().expect("request lock"); + assert_eq!(requests.len(), 2); + assert_eq!(requests[0].raw_body, requests[1].raw_body); + assert_eq!(requests[0].body, requests[1].body); + drop(requests); + assert_eq!(transport.send_limits(), vec![2, 1]); + assert_eq!(transport.remaining_response_count(), 1); + assert_eq!(Provider::retry_config(&provider).max_retries(), 0); + } + + #[tokio::test] + async fn internal_repair_exhausting_the_budget_prevents_an_outer_capacity_replay() { + let transport = Arc::new(DoubleSendCapacityTransport { + calls: AtomicUsize::new(0), + }); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let result = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("hello")], + &[], + ) + .await; + + assert!(matches!( + result, + Err(ProviderError::RateLimitExceeded { + retry_delay: Some(Duration::ZERO), + .. + }) + )); + assert_eq!(transport.calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn a_second_capacity_response_stops_after_two_total_sends() { + let transport = Arc::new(FakeTransport::queued(vec![ + capacity_response(429, Some("0")), + capacity_response(503, Some("0")), + fragmented_success_response(), + ])); let provider = MapleProvider::new(Arc::clone(&transport)); - let default_max_retries = Provider::retry_config(&provider).max_retries(); - assert_eq!(default_max_retries, TRANSIENT_MAX_RETRIES); - let provider = provider.with_test_retry_config(fast_retry_config(default_max_retries)); let result = provider .stream( @@ -1844,20 +2250,51 @@ mod tests { ) .await; let error = match result { - Ok(_) => panic!("the shared retry budget should be exhausted"), + Ok(_) => panic!("the second capacity response should stop the send"), Err(error) => error, }; assert_eq!( error, - ProviderError::NetworkError("Maple's response stream was invalid".to_string()) + ProviderError::RateLimitExceeded { + details: INFERENCE_CAPACITY_ERROR_MESSAGE.to_string(), + retry_delay: Some(Duration::ZERO), + } ); - assert_eq!(transport.request_count(), TRANSIENT_MAX_RETRIES + 1); + assert_eq!(transport.request_count(), 2); assert_eq!(transport.remaining_response_count(), 1); } #[tokio::test] - async fn retries_an_incomplete_tool_call_before_it_is_yielded() { + async fn over_budget_capacity_delay_does_not_replay() { + let transport = Arc::new(FakeTransport::queued(vec![ + capacity_response(503, Some("61")), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let result = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("hello")], + &[], + ) + .await; + + assert!(matches!( + result, + Err(ProviderError::RateLimitExceeded { + retry_delay: None, + .. + }) + )); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } + + #[tokio::test] + async fn incomplete_tool_call_before_output_is_terminal_without_replay() { let interrupted = response_with_items( 200, vec![ @@ -1868,8 +2305,7 @@ mod tests { ); let replacement = response(200, vec![complete_tool_call_sse()], None); let transport = Arc::new(FakeTransport::queued(vec![interrupted, replacement])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let stream = provider .stream( @@ -1879,42 +2315,157 @@ mod tests { &[], ) .await - .expect("replacement tool stream should start"); - let (message, usage) = collect_stream(stream) + .expect("successful response headers should start the stream"); + let error = collect_stream(stream) .await - .expect("replacement tool stream should parse"); - let calls = message - .content - .iter() - .filter_map(|content| match content { - MessageContent::ToolRequest(request) => request.tool_call.as_ref().ok(), - _ => None, - }) - .collect::>(); + .expect_err("the malformed tool stream should fail"); + assert_eq!( + error, + ProviderError::NetworkError("Maple's response stream was invalid".to_string()) + ); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } + + #[tokio::test] + async fn valid_completion_without_done_is_terminal_and_not_replayed() { + let transport = Arc::new(FakeTransport::queued(vec![ + single_text_response(None), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + let model_config = ModelConfig::new("test-model"); + let messages = [Message::user().with_text("hello")]; + + let (result, terminal_error) = with_run_cancellation(CancellationToken::new(), async { + let result = match provider + .stream(&model_config, "system", &messages, &[]) + .await + { + Ok(stream) => collect_stream(stream).await.map(|_| ()), + Err(error) => Err(error), + }; + (result, take_terminal_run_error()) + }) + .await; - assert_eq!(calls.len(), 1); - assert_eq!(calls[0].name, "web_search"); assert_eq!( - calls[0] - .arguments - .as_ref() - .and_then(|arguments| arguments.get("query")), - Some(&json!("maple")) + result, + Err(ProviderError::ExecutionError( + STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string() + )) ); - assert_eq!(usage.usage.total_tokens, Some(5)); - assert_eq!(transport.request_count(), 2); + assert_eq!( + terminal_error.as_deref(), + Some(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE) + ); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } + + #[tokio::test] + async fn incomplete_tool_call_without_done_is_terminal_and_not_replayed() { + let transport = Arc::new(FakeTransport::queued(vec![ + response(200, vec![incomplete_tool_call_sse()], None), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let stream = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("search")], + &[], + ) + .await + .expect("successful response headers should start the stream"); + let error = collect_stream(stream) + .await + .expect_err("missing DONE should fail the incomplete tool stream"); + + assert_eq!( + error, + ProviderError::ExecutionError(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string()) + ); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); } #[tokio::test] - async fn leaves_an_incomplete_tool_call_ending_in_done_as_an_empty_stream() { + async fn compact_done_marker_is_accepted() { + let transport = Arc::new(FakeTransport::new(single_text_response(Some( + "data:[DONE]", + )))); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let stream = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("hello")], + &[], + ) + .await + .expect("stream should start"); + let (message, usage) = collect_stream(stream) + .await + .expect("compact DONE marker should complete the stream"); + + assert!(message + .content + .iter() + .any(|content| matches!(content, MessageContent::Text(text) if text.text == "Hello"))); + assert_eq!(usage.usage.total_tokens, Some(2)); + assert_eq!(transport.request_count(), 1); + } + + #[tokio::test] + async fn incomplete_tool_call_ending_in_done_is_terminal_and_not_replayed() { let mut incomplete_then_done = incomplete_tool_call_sse(); incomplete_then_done.extend_from_slice(b"data: [DONE]\n\n"); let transport = Arc::new(FakeTransport::queued(vec![ response(200, vec![incomplete_then_done], None), fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let model_config = ModelConfig::new("test-model"); + let messages = [Message::user().with_text("search")]; + + let (result, terminal_error) = with_run_cancellation(CancellationToken::new(), async { + let result = match provider + .stream(&model_config, "system", &messages, &[]) + .await + { + Ok(stream) => collect_stream(stream).await.map(|_| ()), + Err(error) => Err(error), + }; + (result, take_terminal_run_error()) + }) + .await; + + assert_eq!( + result, + Err(ProviderError::ExecutionError( + STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string() + )) + ); + assert_eq!( + terminal_error.as_deref(), + Some(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE) + ); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } + + #[tokio::test] + async fn complete_tool_call_without_done_is_terminal_before_tool_is_yielded() { + let transport = Arc::new(FakeTransport::queued(vec![ + response(200, vec![complete_tool_call_event()], None), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); let mut stream = provider .stream( @@ -1924,15 +2475,85 @@ mod tests { &[], ) .await - .expect("DONE should preserve Goose's empty-stream recovery path"); + .expect("successful response headers should start the stream"); + let error = stream + .next() + .await + .expect("missing DONE should surface a terminal item") + .expect_err("the tool request must remain withheld"); + assert_eq!( + error, + ProviderError::ExecutionError(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string()) + ); assert!(stream.next().await.is_none()); assert_eq!(transport.request_count(), 1); assert_eq!(transport.remaining_response_count(), 1); } #[tokio::test] - async fn does_not_retry_after_a_complete_tool_call_is_yielded() { + async fn leading_whitespace_pseudo_done_does_not_release_a_tool_call() { + let mut tool_then_pseudo_done = complete_tool_call_event(); + tool_then_pseudo_done.extend_from_slice(b" data: [DONE]\n\n"); + let transport = Arc::new(FakeTransport::queued(vec![ + response(200, vec![tool_then_pseudo_done], None), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let mut stream = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("search")], + &[], + ) + .await + .expect("successful response headers should start the stream"); + let error = stream + .next() + .await + .expect("invalid terminal marker should surface a terminal item") + .expect_err("the tool request must remain withheld"); + + assert_eq!( + error, + ProviderError::ExecutionError(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string()) + ); + assert!(stream.next().await.is_none()); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } + + #[tokio::test] + async fn complete_tool_call_with_done_is_yielded_once() { + let transport = Arc::new(FakeTransport::new(response( + 200, + vec![complete_tool_call_sse()], + None, + ))); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let (message, _) = collect_stream( + provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("search")], + &[], + ) + .await + .expect("successful response headers should start the stream"), + ) + .await + .expect("DONE should release the buffered tool request"); + + assert!(message.is_tool_call()); + assert_eq!(transport.request_count(), 1); + } + + #[tokio::test] + async fn complete_tool_call_is_withheld_before_terminal_stream_error() { let interrupted = response_with_items( 200, vec![ @@ -1947,8 +2568,7 @@ mod tests { interrupted, fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let mut stream = provider .stream( @@ -1959,32 +2579,72 @@ mod tests { ) .await .expect("complete tool item should start the stream"); - let (message, _) = stream - .next() - .await - .expect("tool item") - .expect("tool item should parse"); - let message = message.expect("tool item should contain a message"); - assert!(message - .content - .iter() - .any(|content| matches!(content, MessageContent::ToolRequest(_)))); - let error = stream .next() .await .expect("the interruption should be surfaced") - .expect_err("the interruption should remain an error"); + .expect_err("the tool request must remain withheld"); assert_eq!( error, ProviderError::ExecutionError(SECURE_CONNECTION_ERROR_MESSAGE.to_string()) ); + assert!(stream.next().await.is_none()); assert_eq!(transport.request_count(), 1); assert_eq!(transport.remaining_response_count(), 1); } #[tokio::test] - async fn maps_rate_limit_without_exposing_body_and_preserves_retry_hint() { + async fn usage_only_done_stream_is_terminal_and_not_replayed() { + let transport = Arc::new(FakeTransport::queued(vec![ + usage_only_response(), + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + + let error = collect_stream( + provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("hello")], + &[], + ) + .await + .expect("successful response headers should start the stream"), + ) + .await + .expect_err("usage without assistant output must be terminal"); + + assert_eq!( + error, + ProviderError::ExecutionError(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string()) + ); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } + + #[tokio::test] + async fn empty_message_done_stream_is_terminal() { + let parsed: MessageStream = Box::pin(futures_util::stream::iter([Ok(( + Some(Message::assistant()), + None, + ))])); + let saw_done = Arc::new(AtomicBool::new(true)); + + let error = collect_stream(MapleProvider::enforce_stream_terminal_contract( + parsed, saw_done, + )) + .await + .expect_err("an empty assistant message must not reach Goose as an empty turn"); + + assert_eq!( + error, + ProviderError::ExecutionError(STREAM_ENDED_BEFORE_COMPLETION_MESSAGE.to_string()) + ); + } + + #[tokio::test] + async fn generic_rate_limit_does_not_gain_replay_authority_from_retry_after() { let response = response( 429, vec![br#"{"error":{"message":"private upstream detail"}}"#.to_vec()], @@ -1999,46 +2659,149 @@ mod tests { error, ProviderError::RateLimitExceeded { details: "Maple rate limit exceeded".to_string(), - retry_delay: Some(Duration::from_secs(7)), + retry_delay: None, } ); } - #[test] - fn invalid_body_retry_hint_falls_back_to_retry_after_header() { - let valid = json!({ - "error": { "metadata": { "retry_after_seconds": 2.5 } } - }); - assert_eq!( - retry_after_delay(Some(&valid), Some("7")), - Some(Duration::from_secs_f64(2.5)) + #[tokio::test] + async fn exact_capacity_error_is_fixed_and_does_not_read_its_body() { + for status in [429, 503] { + let body_polls = Arc::new(AtomicUsize::new(0)); + let error = match ensure_success(unpolled_capacity_response( + status, + Some("7"), + Arc::clone(&body_polls), + )) + .await + { + Ok(_) => panic!("capacity response should fail"), + Err(error) => error, + }; + assert_eq!( + error, + ProviderError::RateLimitExceeded { + details: INFERENCE_CAPACITY_ERROR_MESSAGE.to_string(), + retry_delay: Some(Duration::from_secs(7)), + } + ); + assert_eq!(body_polls.load(Ordering::SeqCst), 0); + } + } + + #[tokio::test] + async fn capacity_contract_is_exact_and_fail_closed_without_replay() { + let mut missing_contract = capacity_response(503, Some("0")); + missing_contract.headers_mut().remove(ERROR_CONTRACT_HEADER); + let mut future_contract = capacity_response(503, Some("0")); + future_contract.headers_mut().insert( + ERROR_CONTRACT_HEADER, + tauri::http::HeaderValue::from_static("2"), + ); + let mut duplicate_contract = capacity_response(503, Some("0")); + duplicate_contract.headers_mut().append( + ERROR_CONTRACT_HEADER, + tauri::http::HeaderValue::from_static("1"), ); - let invalid = json!({ - "error": { "metadata": { "retry_after_seconds": "not-a-number" } } - }); - assert_eq!( - retry_after_delay(Some(&invalid), Some("7")), - Some(Duration::from_secs(7)) + let mut missing_code = capacity_response(429, Some("0")); + missing_code.headers_mut().remove(ERROR_CODE_HEADER); + let mut wrong_code = capacity_response(429, Some("0")); + wrong_code.headers_mut().insert( + ERROR_CODE_HEADER, + tauri::http::HeaderValue::from_static("Inference_Capacity"), + ); + let mut duplicate_code = capacity_response(429, Some("0")); + duplicate_code.headers_mut().append( + ERROR_CODE_HEADER, + tauri::http::HeaderValue::from_static("inference_capacity"), ); - let negative = json!({ - "error": { "metadata": { "retry_after_seconds": -1 } } - }); - assert_eq!( - retry_after_delay(Some(&negative), Some("9")), - Some(Duration::from_secs(9)) + let mut missing_replay = capacity_response(503, Some("0")); + missing_replay.headers_mut().remove(CLIENT_REPLAY_HEADER); + let mut wrong_replay = capacity_response(503, Some("0")); + wrong_replay.headers_mut().insert( + CLIENT_REPLAY_HEADER, + tauri::http::HeaderValue::from_static("SAFE"), ); + let mut duplicate_replay = capacity_response(503, Some("0")); + duplicate_replay.headers_mut().append( + CLIENT_REPLAY_HEADER, + tauri::http::HeaderValue::from_static("safe"), + ); + + let invalid = vec![ + missing_contract, + future_contract, + duplicate_contract, + missing_code, + wrong_code, + duplicate_code, + missing_replay, + wrong_replay, + duplicate_replay, + capacity_response(500, Some("0")), + capacity_response(529, Some("0")), + ]; + + for response in invalid { + assert!(!has_exact_inference_capacity_contract( + response.status(), + response.headers() + )); + let transport = Arc::new(FakeTransport::queued(vec![ + response, + fragmented_success_response(), + ])); + let provider = MapleProvider::new(Arc::clone(&transport)); + let result = provider + .stream( + &ModelConfig::new("test-model"), + "system", + &[Message::user().with_text("hello")], + &[], + ) + .await; + + assert!(result.is_err()); + assert_eq!(transport.request_count(), 1); + assert_eq!(transport.remaining_response_count(), 1); + } } #[test] - fn parses_http_date_retry_after_header() { - let retry_at = SystemTime::now() + Duration::from_secs(120); - let header = httpdate::fmt_http_date(retry_at); - let delay = retry_after_delay(None, Some(&header)).expect("HTTP date should parse"); + fn capacity_retry_after_is_canonical_and_bounded() { + let cases = [ + (None, Some(Duration::from_secs(1))), + (Some("0"), Some(Duration::ZERO)), + (Some("7"), Some(Duration::from_secs(7))), + (Some("60"), Some(Duration::from_secs(60))), + (Some("61"), None), + (Some("01"), Some(Duration::from_secs(1))), + (Some("-1"), Some(Duration::from_secs(1))), + (Some("1.5"), Some(Duration::from_secs(1))), + (Some("1e2"), Some(Duration::from_secs(1))), + ( + Some("Wed, 21 Oct 2015 07:28:00 GMT"), + Some(Duration::from_secs(1)), + ), + (Some("999999999999999999999999999999999999"), None), + ]; + + for (value, expected) in cases { + let response = capacity_response(503, value); + assert_eq!(capacity_retry_delay(response.headers()), expected); + } - assert!(delay >= Duration::from_secs(118)); - assert!(delay <= Duration::from_secs(120)); + let mut duplicated = capacity_response(503, Some("7")); + duplicated.headers_mut().append( + tauri::http::header::RETRY_AFTER, + tauri::http::HeaderValue::from_static("9"), + ); + assert_eq!( + capacity_retry_delay(duplicated.headers()), + Some(Duration::from_secs(1)) + ); } #[tokio::test] @@ -2088,7 +2851,7 @@ mod tests { #[test] fn forbidden_errors_are_request_failures_not_authentication() { - let retry_config = fast_retry_config(3); + let retry_config = RetryConfig::new(3, 0, 1.0, 0).transient_only(); let errors = [ map_http_error(tauri::http::StatusCode::FORBIDDEN, None), map_opensecret_error(opensecret::Error::Api { @@ -2158,7 +2921,7 @@ mod tests { #[test] fn secure_connection_sdk_errors_are_fixed_and_non_transient() { - let retry_config = fast_retry_config(3); + let retry_config = RetryConfig::new(3, 0, 1.0, 0).transient_only(); let errors = vec![ opensecret::Error::Session("private session detail".to_string()), opensecret::Error::KeyExchange("private key detail".to_string()), @@ -2190,8 +2953,7 @@ mod tests { response(401, vec![br#"{"message":"Invalid JWT"}"#.to_vec()], None), fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; @@ -2218,8 +2980,7 @@ mod tests { error_contract_response(Some("1"), Some("session_not_found")), fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; @@ -2255,8 +3016,7 @@ mod tests { )), Ok(fragmented_success_response()), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; @@ -2292,8 +3052,7 @@ mod tests { )), Ok(fragmented_success_response()), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; @@ -2334,15 +3093,18 @@ mod tests { failed, fragmented_success_response(), ])); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(fast_retry_config(3)); + let provider = MapleProvider::new(Arc::clone(&transport)); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; let (result, terminal_error) = with_run_cancellation(CancellationToken::new(), async { - let result = provider + let result = match provider .stream(&model_config, "system", &messages, &[]) - .await; + .await + { + Ok(stream) => collect_stream(stream).await.map(|_| ()), + Err(error) => Err(error), + }; (result, take_terminal_run_error()) }) .await; @@ -2424,20 +3186,19 @@ mod tests { const PRIVATE_MALFORMED_LINE: &str = "private-decrypted-malformed-completion"; let provider = MapleProvider::new(Arc::new(FakeTransport::new(malformed_response( PRIVATE_MALFORMED_LINE, - )))) - .with_test_retry_config(fast_retry_config(0)); - let result = provider + )))); + let stream = provider .stream( &ModelConfig::new("test-model"), "system", &[Message::user().with_text("hello")], &[], ) - .await; - let error = match result { - Ok(_) => panic!("malformed completion data should fail"), - Err(error) => error, - }; + .await + .expect("successful response headers should start the stream"); + let error = collect_stream(stream) + .await + .expect_err("malformed completion data should fail"); assert_eq!( error, ProviderError::NetworkError("Maple's response stream was invalid".to_string()) @@ -2467,6 +3228,7 @@ mod tests { assert!(matches!(result, Err(ProviderError::RequestFailed(_)))); assert_eq!(transport.requests.lock().expect("request lock").len(), 1); let retry_config = Provider::retry_config(&provider); + assert_eq!(retry_config.max_retries(), 0); assert!(!should_retry( &ProviderError::RequestFailed("invalid".to_string()), &retry_config @@ -2513,10 +3275,12 @@ mod tests { let cancellation = CancellationToken::new(); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; - let stream = with_run_cancellation( - cancellation.clone(), - provider.stream(&model_config, "system", &messages, &[]), - ); + let stream = with_run_cancellation(cancellation.clone(), async { + let stream = provider + .stream(&model_config, "system", &messages, &[]) + .await?; + collect_stream(stream).await.map(|_| ()) + }); tokio::pin!(stream); tokio::select! { @@ -2533,15 +3297,12 @@ mod tests { } #[tokio::test] - async fn cancellation_interrupts_first_item_retry_backoff() { - let error_read = Arc::new(Notify::new()); + async fn cancellation_interrupts_capacity_contract_delay() { let transport = Arc::new(FakeTransport::queued(vec![ - notifying_malformed_response(Arc::clone(&error_read)), + capacity_response(503, Some("60")), fragmented_success_response(), ])); - let retry_config = RetryConfig::new(3, 60_000, 1.0, 60_000).transient_only(); - let provider = - MapleProvider::new(Arc::clone(&transport)).with_test_retry_config(retry_config); + let provider = MapleProvider::new(Arc::clone(&transport)); let cancellation = CancellationToken::new(); let model_config = ModelConfig::new("test-model"); let messages = [Message::user().with_text("hello")]; @@ -2552,14 +3313,14 @@ mod tests { tokio::pin!(stream); tokio::select! { - _ = error_read.notified() => {} + _ = transport.wait_for_request_count(1) => {} result = &mut stream => panic!("stream unexpectedly finished before cancellation: {}", result.is_ok()), } cancellation.cancel(); let result = tokio::time::timeout(Duration::from_secs(1), stream) .await - .expect("cancellation should interrupt backoff"); + .expect("cancellation should interrupt the capacity delay"); assert!( matches!(result, Err(ProviderError::ExecutionError(message)) if message.contains("cancelled")) ); @@ -2569,16 +3330,17 @@ mod tests { #[tokio::test] async fn stalled_response_stream_has_a_bounded_idle_timeout() { - let provider = MapleProvider::new(Arc::new(FakeTransport::new(pending_success_response()))) - .with_test_retry_config(fast_retry_config(0)); - let result = provider + let provider = MapleProvider::new(Arc::new(FakeTransport::new(pending_success_response()))); + let stream = provider .stream( &ModelConfig::new("test-model"), "system", &[Message::user().with_text("hello")], &[], ) - .await; + .await + .expect("successful response headers should start the stream"); + let result = collect_stream(stream).await; assert!(matches!(result, Err(ProviderError::NetworkError(_)))); } diff --git a/frontend/src-tauri/src/agent/shell_permission.rs b/frontend/src-tauri/src/agent/shell_permission.rs index 971c6bdd6..1785dce40 100644 --- a/frontend/src-tauri/src/agent/shell_permission.rs +++ b/frontend/src-tauri/src/agent/shell_permission.rs @@ -11,7 +11,6 @@ use tokio_util::sync::CancellationToken; const READ_ONLY_MODE: &str = "smart_approve"; const CLASSIFIER_MODEL: &str = "llama3-3-70b"; const CLASSIFIER_TEMPERATURE: f32 = 0.0; -const CLASSIFIER_MAX_TOKENS: i32 = 256; const CLASSIFIER_TOOL_NAME: &str = "maple__classify_shell_permission"; const CLASSIFIER_TIMEOUT: Duration = Duration::from_secs(10); const MAX_COMMAND_CHARS: usize = 32_000; @@ -236,7 +235,7 @@ impl ShellPermissionClassifier { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(CLASSIFIER_TEMPERATURE)) - .with_max_tokens(Some(CLASSIFIER_MAX_TOKENS)); + .with_max_tokens(None); let input = match serde_json::to_string(request) { Ok(input) => input, Err(error) => { @@ -551,10 +550,10 @@ mod tests { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(CLASSIFIER_TEMPERATURE)) - .with_max_tokens(Some(CLASSIFIER_MAX_TOKENS)); + .with_max_tokens(None); assert_eq!(model_config.model_name, CLASSIFIER_MODEL); assert_eq!(model_config.temperature, Some(CLASSIFIER_TEMPERATURE)); - assert_eq!(model_config.max_tokens, Some(CLASSIFIER_MAX_TOKENS)); + assert_eq!(model_config.max_tokens, None); assert_eq!(model_config.reasoning, Some(false)); assert!(model_config.request_params.is_none()); } @@ -607,7 +606,7 @@ mod tests { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(CLASSIFIER_TEMPERATURE)) - .with_max_tokens(Some(CLASSIFIER_MAX_TOKENS)); + .with_max_tokens(None); provider .complete( @@ -623,7 +622,8 @@ mod tests { let payload = captured.lock().unwrap().take().unwrap(); assert_eq!(payload["model"], CLASSIFIER_MODEL); assert_eq!(payload["temperature"], CLASSIFIER_TEMPERATURE); - assert_eq!(payload["max_tokens"], CLASSIFIER_MAX_TOKENS); + assert!(payload.get("max_tokens").is_none()); + assert!(payload.get("max_completion_tokens").is_none()); assert_eq!(payload["stream"], true); assert_eq!(payload["stream_options"]["include_usage"], true); assert!(payload.get("include_reasoning").is_none()); diff --git a/frontend/src-tauri/src/agent/web_permission.rs b/frontend/src-tauri/src/agent/web_permission.rs index 5023885e8..a162d3cc6 100644 --- a/frontend/src-tauri/src/agent/web_permission.rs +++ b/frontend/src-tauri/src/agent/web_permission.rs @@ -12,7 +12,6 @@ use tokio_util::sync::CancellationToken; const READ_ONLY_MODE: &str = "smart_approve"; const CLASSIFIER_MODEL: &str = "llama3-3-70b"; const CLASSIFIER_TEMPERATURE: f32 = 0.0; -const CLASSIFIER_MAX_TOKENS: i32 = 256; const CLASSIFIER_TOOL_NAME: &str = "maple__classify_web_permission"; const CLASSIFIER_TIMEOUT: Duration = Duration::from_secs(10); const MAX_CURRENT_PROMPT_CHARS: usize = 4_096; @@ -184,7 +183,7 @@ impl WebPermissionClassifier { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(CLASSIFIER_TEMPERATURE)) - .with_max_tokens(Some(CLASSIFIER_MAX_TOKENS)); + .with_max_tokens(None); let input = match serde_json::to_string(request) { Ok(input) => input, Err(error) => { @@ -484,10 +483,10 @@ mod tests { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(CLASSIFIER_TEMPERATURE)) - .with_max_tokens(Some(CLASSIFIER_MAX_TOKENS)); + .with_max_tokens(None); assert_eq!(model_config.model_name, CLASSIFIER_MODEL); assert_eq!(model_config.temperature, Some(CLASSIFIER_TEMPERATURE)); - assert_eq!(model_config.max_tokens, Some(CLASSIFIER_MAX_TOKENS)); + assert_eq!(model_config.max_tokens, None); assert_eq!(model_config.reasoning, Some(false)); assert!(model_config.request_params.is_none()); } @@ -540,7 +539,7 @@ mod tests { model_config.reasoning = Some(false); let model_config = model_config .with_temperature(Some(CLASSIFIER_TEMPERATURE)) - .with_max_tokens(Some(CLASSIFIER_MAX_TOKENS)); + .with_max_tokens(None); provider .complete( @@ -556,7 +555,8 @@ mod tests { let payload = captured.lock().unwrap().take().unwrap(); assert_eq!(payload["model"], CLASSIFIER_MODEL); assert_eq!(payload["temperature"], CLASSIFIER_TEMPERATURE); - assert_eq!(payload["max_tokens"], CLASSIFIER_MAX_TOKENS); + assert!(payload.get("max_tokens").is_none()); + assert!(payload.get("max_completion_tokens").is_none()); assert_eq!(payload["stream"], true); assert_eq!(payload["stream_options"]["include_usage"], true); assert!(payload.get("include_reasoning").is_none()); diff --git a/frontend/src-tauri/src/maple_api.rs b/frontend/src-tauri/src/maple_api.rs index 35a2429b8..5802435ad 100644 --- a/frontend/src-tauri/src/maple_api.rs +++ b/frontend/src-tauri/src/maple_api.rs @@ -1,7 +1,7 @@ use crate::open_secret_config::configured_pcr0_environment; use opensecret::{ - InferenceRequest, InferenceResponse, OpenSecretClient, WebExtractRequest, WebExtractResponse, - WebSearchRequest, WebSearchResponse, + InferenceRequest, InferenceResponse, InferenceSendBudget, OpenSecretClient, WebExtractRequest, + WebExtractResponse, WebSearchRequest, WebSearchResponse, }; use rand::RngCore; use serde::{Deserialize, Serialize}; @@ -258,6 +258,7 @@ impl MapleApiSession { pub(crate) async fn send_inference_request( self: Arc, request: InferenceRequest, + send_budget: InferenceSendBudget, cancel_token: CancellationToken, ) -> Result { let snapshot = self @@ -273,7 +274,9 @@ impl MapleApiSession { _ = operation_cancel.cancelled() => { Err(opensecret::Error::Other("Inference request was cancelled".to_string())) } - response = snapshot.client.send_inference_request(request) => response, + response = snapshot + .client + .send_inference_request_with_budget(request, send_budget) => response, }; if let Err(error) = session.record_refresh(&snapshot).await { log::warn!("Failed to reconcile refreshed Maple API credentials: {error}"); @@ -410,9 +413,10 @@ impl crate::agent::provider::MapleInferenceTransport for MapleApiSession { async fn send_inference_request( self: Arc, request: InferenceRequest, + send_budget: InferenceSendBudget, cancel_token: CancellationToken, ) -> opensecret::Result { - MapleApiSession::send_inference_request(self, request, cancel_token).await + MapleApiSession::send_inference_request(self, request, send_budget, cancel_token).await } } diff --git a/frontend/src/components/UnifiedChat.tsx b/frontend/src/components/UnifiedChat.tsx index fb763fde7..0d189108e 100644 --- a/frontend/src/components/UnifiedChat.tsx +++ b/frontend/src/components/UnifiedChat.tsx @@ -69,7 +69,11 @@ import { import { ModelSelector } from "@/components/ModelSelector"; import { useBillingState, useModelState, useSelectedProjectState } from "@/state/useLocalState"; import { isKnownFreePlan } from "@/billing/billingAccess"; -import { useOpenSecret } from "@opensecret/react"; +import { + findOpenSecretInferenceCapacityError, + OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER, + useOpenSecret +} from "@opensecret/react"; import { UpgradePromptDialog } from "@/components/UpgradePromptDialog"; import { DocumentPlatformDialog } from "@/components/DocumentPlatformDialog"; import { ContextLimitDialog } from "@/components/ContextLimitDialog"; @@ -117,6 +121,7 @@ import type { ResponseFunctionWebSearch, ResponseFunctionToolCall, ResponseFunctionToolCallOutputItem, + ResponseCreateParamsStreaming, ResponseOutputItemAddedEvent, ResponseOutputItemDoneEvent, ResponseReasoningItem, @@ -192,6 +197,7 @@ import { unregisterChatOptimisticMessage } from "@/services/chatOptimisticMessageOwnership"; import { toolKindFromName } from "@/services/toolPresentation"; +import { withInferenceCapacityRetry } from "@/services/inferenceCapacityRetry"; const CHAT_ALERT_CLASS = "absolute top-16 left-1/2 z-50 w-full max-w-2xl -translate-x-1/2 px-4"; const STREAM_EVENT_DEBUG_STORAGE_KEY = "maple:sse-debug"; @@ -4440,21 +4446,25 @@ export function UnifiedChat({ isVisible = true }: { isVisible?: boolean }) { return true; }; - const createResponseStream = async ( - targetConversationId: string, - discardOwnedItemsOnError: boolean - ) => { - const stream = await openai.responses.create( - { - conversation: targetConversationId, - model: requestModel, - input: [{ role: "user", content: messageContent }], - metadata: { internal_message_id: localMessageId }, - stream: true, - store: true, - ...(requestWebSearchEnabled && { tools: [{ type: "web_search" }] }) - }, - { signal: run.signal } + const createResponseStream = async (targetConversationId: string) => { + const responseParams: ResponseCreateParamsStreaming = { + conversation: targetConversationId, + model: requestModel, + input: [{ role: "user", content: messageContent }], + metadata: { internal_message_id: localMessageId }, + stream: true, + store: true, + ...(requestWebSearchEnabled && { tools: [{ type: "web_search" }] }) + }; + const stream = await withInferenceCapacityRetry( + (maxInferenceSends) => + openai.responses.create(responseParams, { + signal: run.signal, + headers: { + [OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER]: String(maxInferenceSends) + } + }), + run.signal ); if (!runtimeStore.setAssistantStreaming(runtimeKey, run.token, true)) return null; @@ -4464,7 +4474,7 @@ export function UnifiedChat({ isVisible = true }: { isVisible?: boolean }) { runtimeKey, run.token, localMessageId, - discardOwnedItemsOnError + false ); } finally { runtimeStore.setAssistantStreaming(runtimeKey, run.token, false); @@ -4634,7 +4644,7 @@ export function UnifiedChat({ isVisible = true }: { isVisible?: boolean }) { window.dispatchEvent(new Event("conversationcreated")); } - const terminalState = await createResponseStream(conversationId, isFollowUpConversation); + const terminalState = await createResponseStream(conversationId); completedSuccessfully = terminalState === "completed"; scheduleBillingRefresh(); } catch (error) { @@ -4709,42 +4719,30 @@ export function UnifiedChat({ isVisible = true }: { isVisible?: boolean }) { status403Error.message || "Access denied. Please check your subscription."; } restoreOriginComposer(displayError); + } else if (findOpenSecretInferenceCapacityError(error)) { + restoreOriginComposer( + "Inference capacity is temporarily unavailable. Your message was restored. Please try again shortly." + ); } else if (error instanceof Error && error.name !== "AbortError") { if (isFollowUpConversation && conversationId) { try { - console.log("Waiting 1s before retry..."); - await new Promise((resolve) => setTimeout(resolve, 1000)); - if (!runtimeStore.isRunCurrent(runtimeKey, run.token)) return; - - console.log("Retrying request once..."); - const terminalState = await createResponseStream(conversationId, false); - completedSuccessfully = terminalState === "completed"; - scheduleBillingRefresh(); - console.log("Retry completed successfully"); - return; - } catch (retryError) { - console.error("Retry failed:", retryError); - if (!runtimeStore.isRunCurrent(runtimeKey, run.token)) return; - - try { - const finalCheckResponse = await openai.conversations.items.list(conversationId, { - limit: 5, - order: "desc" - }); - const foundMessage = finalCheckResponse.data.find( - (item) => item.id === localMessageId - ); + const finalCheckResponse = await openai.conversations.items.list(conversationId, { + limit: 5, + order: "desc" + }); + const foundMessage = finalCheckResponse.data.find( + (item) => item.id === localMessageId + ); - if (!foundMessage) { - console.log("Message not found after retry - restoring input"); - restoreOriginComposer("Failed to send message. Please try again."); - } else { - console.log("Message found after retry failure - it actually went through"); - } - } catch (finalCheckError) { - console.error("Final check failed:", finalCheckError); + if (!foundMessage) { + console.log("Message not found after send failure - restoring input"); restoreOriginComposer("Failed to send message. Please try again."); + } else { + console.log("Message found after send failure - it actually went through"); } + } catch (finalCheckError) { + console.error("Final check failed:", finalCheckError); + restoreOriginComposer("Failed to send message. Please try again."); } } else { const optimisticMessageId = getRegisteredChatOptimisticMessage(runtimeStore, run.token); diff --git a/frontend/src/services/agentThoughtLabels.test.ts b/frontend/src/services/agentThoughtLabels.test.ts index 5e80999a3..c037918fb 100644 --- a/frontend/src/services/agentThoughtLabels.test.ts +++ b/frontend/src/services/agentThoughtLabels.test.ts @@ -966,9 +966,11 @@ describe("requestAgentThoughtLabel", () => { expect(requestedBody).toMatchObject({ model: "llama3-3-70b", temperature: 0, - max_tokens: 64, stream: false }); + expect(requestedBody).not.toHaveProperty("max_tokens"); + expect(requestedBody).not.toHaveProperty("max_completion_tokens"); + expect(requestedBody).not.toHaveProperty("max_output_tokens"); expect(requestedBody).not.toHaveProperty("include_reasoning"); expect(requestedBody).not.toHaveProperty("chat_template_kwargs"); expect(requestedBody).not.toHaveProperty("conversation"); diff --git a/frontend/src/services/agentThoughtLabels.ts b/frontend/src/services/agentThoughtLabels.ts index d91afd2de..036d5e031 100644 --- a/frontend/src/services/agentThoughtLabels.ts +++ b/frontend/src/services/agentThoughtLabels.ts @@ -18,7 +18,6 @@ export const AGENT_THOUGHT_LABEL_PROVISIONAL_DEADLINE_MS = 3_000; export const AGENT_THOUGHT_LABEL_MAX_CONCURRENT_PROVISIONAL_REQUESTS = 2; const AGENT_THOUGHT_LABEL_MODEL = "llama3-3-70b"; -const AGENT_THOUGHT_LABEL_MAX_TOKENS = 64; const AGENT_THOUGHT_LABEL_TEMPERATURE = 0; const AGENT_THOUGHT_LABEL_STREAMING_PREMATURE_VERBS = [ "Answering", @@ -436,7 +435,6 @@ export async function requestAgentThoughtLabel( { role: "user", content: input } ], temperature: AGENT_THOUGHT_LABEL_TEMPERATURE, - max_tokens: AGENT_THOUGHT_LABEL_MAX_TOKENS, stream: false }; const response = await client.chat.completions.create(requestBody, { signal }); diff --git a/frontend/src/services/inferenceCapacityRetry.test.ts b/frontend/src/services/inferenceCapacityRetry.test.ts new file mode 100644 index 000000000..f5168f3c8 --- /dev/null +++ b/frontend/src/services/inferenceCapacityRetry.test.ts @@ -0,0 +1,142 @@ +import { describe, expect, test } from "bun:test"; +import { + findOpenSecretInferenceCapacityError, + OpenSecretInferenceCapacityError +} from "@opensecret/react"; +import OpenAI from "openai"; +import { withInferenceCapacityRetry } from "./inferenceCapacityRetry"; + +describe("withInferenceCapacityRetry", () => { + test("performs one initial send and one replay with the same request object", async () => { + let sends = 0; + const request = { + model: "kimi-k3", + metadata: { internal_message_id: "same" } + }; + const sentRequests: (typeof request)[] = []; + const sendLimits: number[] = []; + const wrapped = new Error("OpenAI wrapper") as Error & { cause?: unknown }; + wrapped.cause = new OpenSecretInferenceCapacityError(503, 0); + + const result = await withInferenceCapacityRetry(async (sendLimit) => { + sends += 1; + sendLimits.push(sendLimit); + sentRequests.push(request); + if (sends === 1) throw wrapped; + return "completed"; + }, new AbortController().signal); + + expect(result).toBe("completed"); + expect(sends).toBe(2); + expect(sentRequests).toEqual([request, request]); + expect(sentRequests[0]).toBe(sentRequests[1]); + expect(sendLimits).toEqual([2, 1]); + }); + + test("does not replay generic, structurally spoofed, or over-budget failures", async () => { + const failures: unknown[] = [ + new Error("generic 503"), + { + name: "OpenSecretInferenceCapacityError", + status: 503, + retryDelayMs: 0 + }, + new OpenSecretInferenceCapacityError(429, null) + ]; + + for (const failure of failures) { + let sends = 0; + await expect( + withInferenceCapacityRetry(async () => { + sends += 1; + throw failure; + }, new AbortController().signal) + ).rejects.toBe(failure); + expect(sends).toBe(1); + } + }); + + test("aborting during the delay prevents replay", async () => { + const controller = new AbortController(); + let sends = 0; + const replay = withInferenceCapacityRetry(async () => { + sends += 1; + throw new OpenSecretInferenceCapacityError(503, 60_000); + }, controller.signal); + + controller.abort(); + await expect(replay).rejects.toMatchObject({ name: "AbortError" }); + expect(sends).toBe(1); + }); + + test("propagates the sole replay failure after exactly two total sends", async () => { + let sends = 0; + const retryFailure = new Error("retry failed"); + await expect( + withInferenceCapacityRetry(async () => { + sends += 1; + if (sends === 1) throw new OpenSecretInferenceCapacityError(429, 0); + throw retryFailure; + }, new AbortController().signal) + ).rejects.toBe(retryFailure); + expect(sends).toBe(2); + }); + + test("does not replay capacity after SDK repair consumed both send permits", async () => { + let calls = 0; + const capacity = new OpenSecretInferenceCapacityError(503, 0, 2); + + await expect( + withInferenceCapacityRetry(async () => { + calls += 1; + throw capacity; + }, new AbortController().signal) + ).rejects.toBe(capacity); + + expect(calls).toBe(1); + }); + + test("bounds the real OpenAI wrapper to two transport sends", async () => { + let sends = 0; + const openai = new OpenAI({ + apiKey: "not-a-real-api-key", + baseURL: "https://example.test/v1/", + dangerouslyAllowBrowser: true, + fetch: async () => { + sends += 1; + throw new OpenSecretInferenceCapacityError(503, 0); + }, + maxRetries: 0 + }); + const request = { model: "kimi-k3", input: "hello" }; + + let error: unknown; + try { + await withInferenceCapacityRetry( + () => openai.responses.create(request), + new AbortController().signal + ); + } catch (caught) { + error = caught; + } + + expect(sends).toBe(2); + expect(findOpenSecretInferenceCapacityError(error)).toBeInstanceOf( + OpenSecretInferenceCapacityError + ); + }); + + test("a pre-aborted signal performs no send", async () => { + const controller = new AbortController(); + controller.abort(); + let sends = 0; + + await expect( + withInferenceCapacityRetry(async () => { + sends += 1; + return "unexpected"; + }, controller.signal) + ).rejects.toMatchObject({ name: "AbortError" }); + expect(sends).toBe(0); + }); +}); diff --git a/frontend/src/services/inferenceCapacityRetry.ts b/frontend/src/services/inferenceCapacityRetry.ts new file mode 100644 index 000000000..1958f4172 --- /dev/null +++ b/frontend/src/services/inferenceCapacityRetry.ts @@ -0,0 +1,41 @@ +import { findOpenSecretInferenceCapacityError } from "@opensecret/react"; + +async function waitForRetry(delayMs: number, signal: AbortSignal): Promise { + signal.throwIfAborted(); + if (delayMs === 0) return; + + await new Promise((resolve, reject) => { + const onAbort = () => { + clearTimeout(timeout); + reject(signal.reason ?? new DOMException("The operation was aborted.", "AbortError")); + }; + const timeout = setTimeout(() => { + signal.removeEventListener("abort", onAbort); + resolve(); + }, delayMs); + signal.addEventListener("abort", onAbort, { once: true }); + }); +} + +type InferenceSendLimit = 1 | 2; + +/** Executes at most two inference sends across SDK repair and capacity replay. */ +export async function withInferenceCapacityRetry( + send: (maxInferenceSends: InferenceSendLimit) => Promise, + signal: AbortSignal +): Promise { + signal.throwIfAborted(); + + try { + return await send(2); + } catch (error) { + const capacity = findOpenSecretInferenceCapacityError(error); + if (!capacity || capacity.retryDelayMs === null || capacity.inferenceSendCount !== 1) { + throw error; + } + + await waitForRetry(capacity.retryDelayMs, signal); + signal.throwIfAborted(); + return send(1); + } +} diff --git a/sdk/opensecret-integration-revision b/sdk/opensecret-integration-revision index e7f0a0687..b5e1c2af8 100644 --- a/sdk/opensecret-integration-revision +++ b/sdk/opensecret-integration-revision @@ -1 +1 @@ -3f3c9aff9d4dbcfdcaf945c06ecf0e7ed6dae605 +91ed2e573aa5f358cfd0d59defd04d414be762fe diff --git a/sdk/rust/Cargo.toml b/sdk/rust/Cargo.toml index a1ed097af..c77a1350f 100644 --- a/sdk/rust/Cargo.toml +++ b/sdk/rust/Cargo.toml @@ -13,7 +13,7 @@ categories = ["cryptography", "api-bindings", "web-programming"] [dependencies] # HTTP and async runtime -reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "http2", "charset", "system-proxy"] } +reqwest = { version = "0.12.23", default-features = false, features = ["json", "stream", "rustls-tls", "http2", "charset", "system-proxy"] } http = "1.1" tokio = { version = "1.41", features = ["full"] } async-trait = "0.1" diff --git a/sdk/rust/src/client.rs b/sdk/rust/src/client.rs index 7901dd8ca..ca89c9184 100644 --- a/sdk/rust/src/client.rs +++ b/sdk/rust/src/client.rs @@ -17,7 +17,14 @@ use reqwest::{ Client, }; use serde::{de::DeserializeOwned, Deserialize, Serialize}; -use std::{net::IpAddr, pin::Pin}; +use std::{ + net::IpAddr, + pin::Pin, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, +}; use tokio::sync::Mutex; use uuid::Uuid; @@ -34,6 +41,44 @@ pub type InferenceRequest = HttpRequest; /// A decrypted HTTP response from an OpenSecret inference endpoint. pub type InferenceResponse = HttpResponse; +/// A cloneable, request-local ceiling on inference HTTP sends. +/// +/// Attestation, token refresh, and other control-plane requests do not consume +/// this budget. A permit is reserved only immediately before an inference HTTP +/// request is sent. +#[derive(Clone, Debug)] +pub struct InferenceSendBudget { + remaining: Arc, +} + +impl InferenceSendBudget { + /// Creates a budget that permits at most `max_sends` inference HTTP sends. + pub fn new(max_sends: usize) -> Result { + if max_sends == 0 { + return Err(Error::Configuration( + "Inference send budget must allow at least one send".to_string(), + )); + } + Ok(Self { + remaining: Arc::new(AtomicUsize::new(max_sends)), + }) + } + + /// Returns the number of inference HTTP sends still available. + pub fn remaining(&self) -> usize { + self.remaining.load(Ordering::Acquire) + } + + /// Reserves one inference HTTP send, returning false when exhausted. + pub fn try_reserve_send(&self) -> bool { + self.remaining + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |remaining| { + remaining.checked_sub(1) + }) + .is_ok() + } +} + #[derive(Deserialize)] #[serde(deny_unknown_fields)] struct EncryptedBody { @@ -44,6 +89,7 @@ const MAX_INFERENCE_SSE_LINE_BYTES: usize = 16 * 1024 * 1024; pub struct OpenSecretClient { client: Client, + inference_client: Client, base_url: String, session_manager: SessionManager, refresh_lock: Mutex<()>, @@ -493,6 +539,10 @@ impl OpenSecretClient { Ok(Self { client: Client::new(), + inference_client: Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .retry(reqwest::retry::never()) + .build()?, base_url: base_url.trim_end_matches('/').to_string(), session_manager: SessionManager::new(), refresh_lock: Mutex::new(()), @@ -530,6 +580,10 @@ impl OpenSecretClient { Ok(Self { client: Client::new(), + inference_client: Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .retry(reqwest::retry::never()) + .build()?, base_url: base_url.trim_end_matches('/').to_string(), session_manager: SessionManager::new_with_api_key(api_key), refresh_lock: Mutex::new(()), @@ -939,6 +993,21 @@ impl OpenSecretClient { pub async fn send_inference_request( &self, request: InferenceRequest, + ) -> Result { + self.send_inference_request_with_budget(request, InferenceSendBudget::new(2)?) + .await + } + + /// Sends an inference request while sharing a caller-owned send budget. + /// + /// This is equivalent to [`Self::send_inference_request`], except the SDK's + /// safe authentication or session repair replays consume the same budget as + /// any outer caller retry. The budget is request-local and may be cloned + /// across those nested layers. + pub async fn send_inference_request_with_budget( + &self, + request: InferenceRequest, + send_budget: InferenceSendBudget, ) -> Result { let (parts, body) = request.into_parts(); if parts.uri.scheme().is_some() || parts.uri.authority().is_some() { @@ -973,6 +1042,7 @@ impl OpenSecretClient { &headers, body.clone(), &auth, + &send_budget, ) .await; @@ -990,6 +1060,9 @@ impl OpenSecretClient { match recovery { Some(RecoveryAction::Reattest) => { self.perform_attestation_handshake().await?; + if send_budget.remaining() == 0 { + return self.finish_inference_response(response, session_key).await; + } replayed = true; } Some(RecoveryAction::RefreshAccessToken) => { @@ -1001,6 +1074,11 @@ impl OpenSecretClient { .await, Ok(true) ) { + if send_budget.remaining() == 0 { + return self + .finish_inference_response(response, session_key) + .await; + } replayed = true; } else { return self.finish_inference_response(response, session_key).await; @@ -1025,6 +1103,7 @@ impl OpenSecretClient { caller_headers: &HttpHeaderMap, body: Bytes, auth: &ResolvedAuth, + send_budget: &InferenceSendBudget, ) -> Result<(reqwest::Response, [u8; 32])> { let session = self.session_manager.get_session()?.ok_or_else(|| { Error::Session( @@ -1058,7 +1137,15 @@ impl OpenSecretClient { encrypted: BASE64.encode(encrypted), }) }; - let request = self.client.request(method.clone(), url).headers(headers); + let request = self + .inference_client + .request(method.clone(), url) + .headers(headers); + if !send_budget.try_reserve_send() { + return Err(Error::Other( + "Inference request send budget exhausted".to_string(), + )); + } let response = match encrypted_body { Some(encrypted_body) => request.json(&encrypted_body).send().await?, None => request.send().await?, @@ -4451,7 +4538,11 @@ mod tests { .uri("/v1/chat/completions") .body(request_body) .unwrap(); - let response = client.send_inference_request(request).await.unwrap(); + let send_budget = InferenceSendBudget::new(2).unwrap(); + let response = client + .send_inference_request_with_budget(request, send_budget.clone()) + .await + .unwrap(); assert_eq!(response.status(), http::StatusCode::TOO_MANY_REQUESTS); assert_eq!( @@ -4462,6 +4553,112 @@ mod tests { collect_response_body(response.into_body()).await.unwrap(), error_body ); + assert_eq!(send_budget.remaining(), 0); + } + + #[tokio::test] + async fn exhausted_inference_budget_refreshes_auth_without_a_second_send() { + let mock_server = MockServer::start().await; + let client = OpenSecretClient::new(mock_server.uri()).unwrap(); + let session_id = Uuid::new_v4(); + let session_key = crypto::generate_random_bytes::<32>(); + let request_body = Bytes::from_static(br#"{"model":"test","messages":[]}"#); + client + .session_manager + .set_session(session_id, session_key) + .unwrap(); + client + .session_manager + .set_tokens( + "expired_access".to_string(), + Some("refresh_token".to_string()), + ) + .unwrap(); + + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .and(header("authorization", "Bearer expired_access")) + .respond_with(ResponseTemplate::new(401).set_body_string("jwt expired")) + .expect(1) + .mount(&mock_server) + .await; + Mock::given(method("POST")) + .and(path("/refresh")) + .respond_with(ResponseTemplate::new(200).set_body_json(encrypted_response( + &session_key, + &json!({ + "access_token": "fresh_access", + "refresh_token": "fresh_refresh" + }), + ))) + .expect(1) + .mount(&mock_server) + .await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .and(header("authorization", "Bearer fresh_access")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&mock_server) + .await; + + let request = HttpRequest::builder() + .method(http::Method::POST) + .uri("/v1/chat/completions") + .body(request_body) + .unwrap(); + let send_budget = InferenceSendBudget::new(1).unwrap(); + let response = client + .send_inference_request_with_budget(request, send_budget.clone()) + .await + .unwrap(); + + assert_eq!(response.status(), http::StatusCode::UNAUTHORIZED); + assert_eq!(send_budget.remaining(), 0); + assert_eq!( + client.get_access_token().unwrap().as_deref(), + Some("fresh_access") + ); + mock_server.verify().await; + } + + #[tokio::test] + async fn inference_transport_does_not_follow_redirects_outside_the_send_budget() { + let mock_server = MockServer::start().await; + let client = + OpenSecretClient::new_with_api_key(mock_server.uri(), "api_key".to_string()).unwrap(); + client + .session_manager + .set_session(Uuid::new_v4(), [49u8; 32]) + .unwrap(); + + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(307).insert_header("location", "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/v1/embeddings")) + .expect(1) + .mount(&mock_server) + .await; + Mock::given(method("POST")) + .and(path("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/v1/embeddings")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&mock_server) + .await; + + let request = HttpRequest::builder() + .method(http::Method::POST) + .uri("/v1/chat/completions") + .body(Bytes::from_static(br#"{"model":"test","messages":[]}"#)) + .unwrap(); + let send_budget = InferenceSendBudget::new(2).unwrap(); + let response = client + .send_inference_request_with_budget(request, send_budget.clone()) + .await + .unwrap(); + + assert_eq!(response.status(), http::StatusCode::TEMPORARY_REDIRECT); + assert_eq!(send_budget.remaining(), 1); + mock_server.verify().await; } #[tokio::test] diff --git a/sdk/rust/src/lib.rs b/sdk/rust/src/lib.rs index 4dd2ee9ed..300495cee 100644 --- a/sdk/rust/src/lib.rs +++ b/sdk/rust/src/lib.rs @@ -8,7 +8,10 @@ pub mod push; pub mod session; pub mod types; -pub use client::{InferenceRequest, InferenceResponse, OpenSecretClient, OpenSecretResponseBody}; +pub use client::{ + InferenceRequest, InferenceResponse, InferenceSendBudget, OpenSecretClient, + OpenSecretResponseBody, +}; pub use error::{Error, Result}; pub use pcr::{Pcr0Environment, Pcr0TrustPolicy}; pub use push::*; diff --git a/sdk/src/lib/ai.ts b/sdk/src/lib/ai.ts index f6399b79b..9d75f13b8 100644 --- a/sdk/src/lib/ai.ts +++ b/sdk/src/lib/ai.ts @@ -13,6 +13,94 @@ export interface CustomFetchOptions { pcrConfig?: PcrConfig; } +const INFERENCE_CAPACITY_ERROR_MESSAGE = "Inference capacity is temporarily unavailable."; +const INFERENCE_CAPACITY_CONTRACT_HEADER = "x-opensecret-error-contract"; +const INFERENCE_CAPACITY_CODE_HEADER = "x-opensecret-error-code"; +const INFERENCE_CAPACITY_REPLAY_HEADER = "x-opensecret-client-replay"; +const INFERENCE_CAPACITY_CONTRACT_VERSION = "1"; +const INFERENCE_CAPACITY_ERROR_CODE = "inference_capacity"; +const INFERENCE_CAPACITY_REPLAY_SAFE = "safe"; +const DEFAULT_INFERENCE_CAPACITY_RETRY_DELAY_MS = 1_000; +const MAX_INFERENCE_CAPACITY_RETRY_DELAY_SECS = 60n; +/** Client-only header consumed by createCustomFetch and never forwarded upstream. */ +export const OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER = "x-opensecret-client-inference-send-limit"; + +export class OpenSecretInferenceCapacityError extends Error { + readonly status: 429 | 503; + /** Null means the server delay exceeds Maple's bounded automatic-replay window. */ + readonly retryDelayMs: number | null; + /** Number of inference HTTP sends consumed before this terminal response. */ + readonly inferenceSendCount: number; + + constructor(status: 429 | 503, retryDelayMs: number | null, inferenceSendCount = 1) { + super(INFERENCE_CAPACITY_ERROR_MESSAGE); + this.name = "OpenSecretInferenceCapacityError"; + this.status = status; + this.retryDelayMs = retryDelayMs; + this.inferenceSendCount = inferenceSendCount; + } +} + +/** Finds the SDK-owned capacity error through wrappers such as OpenAI APIConnectionError. */ +export function findOpenSecretInferenceCapacityError( + error: unknown +): OpenSecretInferenceCapacityError | null { + const seen = new Set(); + let current = error; + + for (let depth = 0; depth < 8 && current !== null && current !== undefined; depth += 1) { + if (current instanceof OpenSecretInferenceCapacityError) return current; + if (typeof current !== "object" || seen.has(current)) return null; + seen.add(current); + current = (current as { cause?: unknown }).cause; + } + + return null; +} + +function retryDelayFromCapacityHeaders(headers: Headers): number | null { + const retryAfter = headers.get("retry-after"); + if (retryAfter === null || !/^(0|[1-9]\d*)$/.test(retryAfter)) { + return DEFAULT_INFERENCE_CAPACITY_RETRY_DELAY_MS; + } + + const seconds = BigInt(retryAfter); + if (seconds > MAX_INFERENCE_CAPACITY_RETRY_DELAY_SECS) return null; + return Number(seconds) * 1_000; +} + +function inferenceCapacityError( + response: Response, + inferenceSendCount: number +): OpenSecretInferenceCapacityError | null { + if (response.status !== 429 && response.status !== 503) return null; + if ( + response.headers.get(INFERENCE_CAPACITY_CONTRACT_HEADER) !== + INFERENCE_CAPACITY_CONTRACT_VERSION || + response.headers.get(INFERENCE_CAPACITY_CODE_HEADER) !== INFERENCE_CAPACITY_ERROR_CODE || + response.headers.get(INFERENCE_CAPACITY_REPLAY_HEADER) !== INFERENCE_CAPACITY_REPLAY_SAFE + ) { + return null; + } + + return new OpenSecretInferenceCapacityError( + response.status, + retryDelayFromCapacityHeaders(response.headers), + inferenceSendCount + ); +} + +function takeInferenceSendLimit(headers: Headers): { + maxSends: number; + explicitlyBounded: boolean; +} { + const rawLimit = headers.get(OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER); + headers.delete(OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER); + if (rawLimit === "1") return { maxSends: 1, explicitlyBounded: true }; + if (rawLimit === "2") return { maxSends: 2, explicitlyBounded: true }; + return { maxSends: 2, explicitlyBounded: false }; +} + interface ActiveAttestation { sessionKey: Uint8Array; sessionId: string; @@ -220,8 +308,14 @@ export function createCustomFetchWithDependencies( // already-prepared plaintext request under a different token. let authHeader = getAuthHeader(); const request = await snapshotRequest(requestUrl, init); + const { maxSends: maxInferenceSends, explicitlyBounded } = takeInferenceSendLimit( + request.headers + ); + if (explicitlyBounded) request.options.redirect = "manual"; throwIfAborted(request.signal); + let inferenceSendCount = 0; + const makeRequest = async (attestation: ActiveAttestation) => { const headers = new Headers(request.headers); headers.set("Authorization", authHeader); @@ -243,6 +337,7 @@ export function createCustomFetchWithDependencies( headers.set("Content-Type", "application/json"); } + inferenceSendCount += 1; return { attestation, response: await dependencies.fetch(request.url, requestOptions) @@ -265,8 +360,8 @@ export function createCustomFetchWithDependencies( const recovery = classifyRecovery(attempt.response.status, attempt.response.headers); if (recovery === "refresh_access_token" && !usesApiKey && !replayed) { - replayed = true; - await discardResponse(attempt.response); + const canReplay = inferenceSendCount < maxInferenceSends; + if (canReplay) await discardResponse(attempt.response); throwIfAborted(request.signal); console.warn("Unauthorized, refreshing access token"); await dependencies.refreshToken(); @@ -282,16 +377,26 @@ export function createCustomFetchWithDependencies( attestationIdentity.pcrConfig ) ); + if (!canReplay) { + finalAttempt = attempt; + break; + } + replayed = true; continue; } if (recovery === "renew_session" && !replayed) { - replayed = true; - await discardResponse(attempt.response); + const canReplay = inferenceSendCount < maxInferenceSends; + if (canReplay) await discardResponse(attempt.response); throwIfAborted(request.signal); console.warn("Bad Request, renewing attestation and retrying once"); attestation = await renewAttestation(attempt.attestation.sessionId, attestationIdentity); throwIfAborted(request.signal); + if (!canReplay) { + finalAttempt = attempt; + break; + } + replayed = true; continue; } @@ -302,6 +407,12 @@ export function createCustomFetchWithDependencies( const { response } = finalAttempt; const { sessionKey } = finalAttempt.attestation; + const capacityError = inferenceCapacityError(response, inferenceSendCount); + if (capacityError) { + await discardResponse(response); + throw capacityError; + } + if (!response.ok) { const errorText = await response.text(); console.error( diff --git a/sdk/src/lib/api.ts b/sdk/src/lib/api.ts index cb41a9f2a..e12b1eeec 100644 --- a/sdk/src/lib/api.ts +++ b/sdk/src/lib/api.ts @@ -2045,7 +2045,7 @@ export type ResponsesCreateRequest = { * * NOTE: Prefer using the OpenAI client directly for conversation operations: * ```typescript - * const openai = new OpenAI({ fetch: customFetch }); + * const openai = new OpenAI({ fetch: customFetch, maxRetries: 0 }); * const conversation = await openai.conversations.create({ * metadata: { title: "Product Support", category: "technical" } * }); diff --git a/sdk/src/lib/index.ts b/sdk/src/lib/index.ts index bd18e8879..2a0902b80 100644 --- a/sdk/src/lib/index.ts +++ b/sdk/src/lib/index.ts @@ -114,7 +114,13 @@ export { } from "./api"; // Export AI customization options -export { createCustomFetch, type CustomFetchOptions } from "./ai"; +export { + createCustomFetch, + findOpenSecretInferenceCapacityError, + OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER, + OpenSecretInferenceCapacityError, + type CustomFetchOptions +} from "./ai"; // Re-export Model type from OpenAI for convenience export type { Model } from "openai/resources/models.js"; diff --git a/sdk/src/lib/main.tsx b/sdk/src/lib/main.tsx index e5a23ca88..6d81d7699 100644 --- a/sdk/src/lib/main.tsx +++ b/sdk/src/lib/main.tsx @@ -339,7 +339,9 @@ export type OpenSecretContextType = { * defaultHeaders: { * "Accept-Encoding": "identity" * }, - * fetch: os.aiCustomFetch + * fetch: os.aiCustomFetch, + * // OpenSecret exposes an explicit replay contract; keep transport retries disabled. + * maxRetries: 0 * }); * ``` */ diff --git a/sdk/src/lib/test/customFetch.test.ts b/sdk/src/lib/test/customFetch.test.ts index 5109364b6..c18211a95 100644 --- a/sdk/src/lib/test/customFetch.test.ts +++ b/sdk/src/lib/test/customFetch.test.ts @@ -1,5 +1,12 @@ import { beforeEach, describe, expect, test } from "bun:test"; -import { createCustomFetchWithDependencies, type CustomFetchDependencies } from "../ai"; +import OpenAI from "openai"; +import { + createCustomFetchWithDependencies, + findOpenSecretInferenceCapacityError, + OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER, + OpenSecretInferenceCapacityError, + type CustomFetchDependencies +} from "../ai"; import { getApiPcrConfig, getApiUrl, setApiUrl } from "../api"; import type { Attestation } from "../getAttestation"; import type { PcrConfig } from "../pcr"; @@ -45,6 +52,31 @@ function contractError(status: number, body: string, code?: string): Response { return new Response(body, { status, headers }); } +function capacityContractError( + status: number, + options?: { + contract?: string | null; + code?: string | null; + replay?: string | null; + retryAfter?: string | null; + } +): Response { + const headers = new Headers(); + if (options?.contract !== null) { + headers.set("x-opensecret-error-contract", options?.contract ?? "1"); + } + if (options?.code !== null) { + headers.set("x-opensecret-error-code", options?.code ?? "inference_capacity"); + } + if (options?.replay !== null) { + headers.set("x-opensecret-client-replay", options?.replay ?? "safe"); + } + if (options?.retryAfter !== undefined && options.retryAfter !== null) { + headers.set("retry-after", options.retryAfter); + } + return new Response("private upstream capacity detail", { status, headers }); +} + function dependencies(overrides: Partial): CustomFetchDependencies { return { decryptMessage: decryptForTest, @@ -96,6 +128,271 @@ async function withRequestBodyUnavailable(callback: () => Promise): Promis } } +describe("createCustomFetch inference-capacity contract", () => { + beforeEach(() => { + window.localStorage.clear(); + window.sessionStorage.clear(); + }); + + for (const status of [429, 503] as const) { + test(`classifies exact ${status} without consuming or exposing its body`, async () => { + let bodyRead = false; + let bodyCancelled = false; + const capacityResponse = capacityContractError(status, { retryAfter: "7" }); + const responseBody = capacityResponse.body; + if (!responseBody) throw new Error("capacity test response must have a body"); + const cancelBody = responseBody.cancel.bind(responseBody); + responseBody.cancel = async (reason?: unknown) => { + bodyCancelled = true; + return cancelBody(reason); + }; + capacityResponse.text = async () => { + bodyRead = true; + return "private upstream capacity detail"; + }; + const customFetch = createCustomFetchWithDependencies( + { apiKey: "test-api-key" }, + dependencies({ + fetch: async () => capacityResponse + }) + ); + + let error: unknown; + try { + await customFetch("https://example.test/v1/responses", { + method: "POST", + body: '{"prompt":"hello"}' + }); + } catch (caught) { + error = caught; + } + + expect(error).toBeInstanceOf(OpenSecretInferenceCapacityError); + expect(error).toMatchObject({ + name: "OpenSecretInferenceCapacityError", + message: "Inference capacity is temporarily unavailable.", + status, + retryDelayMs: 7_000, + inferenceSendCount: 1 + }); + expect(String(error)).not.toContain("private upstream"); + expect(bodyRead).toBe(false); + expect(bodyCancelled).toBe(true); + }); + } + + test("uses strict bounded delta-seconds retry hints", async () => { + const cases: Array<[string | undefined, number | null]> = [ + [undefined, 1_000], + ["0", 0], + ["7", 7_000], + ["60", 60_000], + ["61", null], + ["01", 1_000], + ["-1", 1_000], + ["1.5", 1_000], + ["1e2", 1_000], + ["Wed, 21 Oct 2015 07:28:00 GMT", 1_000], + ["7, 9", 1_000], + ["999999999999999999999999999999999999999999", null] + ]; + + for (const [retryAfter, expectedDelay] of cases) { + const customFetch = createCustomFetchWithDependencies( + { apiKey: "test-api-key" }, + dependencies({ + fetch: async () => capacityContractError(503, { retryAfter }) + }) + ); + + let error: unknown; + try { + await customFetch("https://example.test/v1/responses"); + } catch (caught) { + error = caught; + } + expect(error).toBeInstanceOf(OpenSecretInferenceCapacityError); + expect((error as OpenSecretInferenceCapacityError).retryDelayMs).toBe(expectedDelay); + } + }); + + test("rejects missing, future, duplicated, and status-mismatched required headers", async () => { + const invalid = [ + capacityContractError(429, { contract: null }), + capacityContractError(429, { contract: "2" }), + capacityContractError(429, { contract: "1, 1" }), + capacityContractError(429, { code: null }), + capacityContractError(429, { code: "inference_capacity_v2" }), + capacityContractError(429, { code: "inference_capacity, inference_capacity" }), + capacityContractError(429, { replay: null }), + capacityContractError(429, { replay: "true" }), + capacityContractError(429, { replay: "safe, safe" }), + capacityContractError(529), + capacityContractError(500) + ]; + + const duplicateContract = capacityContractError(503); + duplicateContract.headers.append("x-opensecret-error-contract", "1"); + invalid.push(duplicateContract); + const duplicateCode = capacityContractError(503); + duplicateCode.headers.append("x-opensecret-error-code", "inference_capacity"); + invalid.push(duplicateCode); + const duplicateReplay = capacityContractError(503); + duplicateReplay.headers.append("x-opensecret-client-replay", "safe"); + invalid.push(duplicateReplay); + + for (const response of invalid) { + const customFetch = createCustomFetchWithDependencies( + { apiKey: "test-api-key" }, + dependencies({ fetch: async () => response.clone() }) + ); + + let error: unknown; + try { + await customFetch("https://example.test/v1/responses"); + } catch (caught) { + error = caught; + } + expect(findOpenSecretInferenceCapacityError(error)).toBeNull(); + expect(String(error)).toContain(`Request failed with status ${response.status}`); + } + }); + + test("finds only the SDK-owned typed error through a bounded cause chain", () => { + const capacity = new OpenSecretInferenceCapacityError(503, 1_000); + const wrapped = new Error("outer", { cause: new Error("middle", { cause: capacity }) }); + expect(findOpenSecretInferenceCapacityError(wrapped)).toBe(capacity); + expect( + findOpenSecretInferenceCapacityError({ + name: "OpenSecretInferenceCapacityError", + status: 503, + retryDelayMs: 1_000 + }) + ).toBeNull(); + + const cycle: { cause?: unknown } = {}; + cycle.cause = cycle; + expect(findOpenSecretInferenceCapacityError(cycle)).toBeNull(); + }); + + test("survives the real OpenAI wrapper with its transport retries disabled", async () => { + let sends = 0; + const customFetch = createCustomFetchWithDependencies( + { apiKey: "test-api-key" }, + dependencies({ + fetch: async () => { + sends += 1; + return capacityContractError(503, { retryAfter: "0" }); + } + }) + ); + const openai = new OpenAI({ + apiKey: "not-a-real-api-key", + baseURL: "https://example.test/v1/", + dangerouslyAllowBrowser: true, + fetch: customFetch, + maxRetries: 0 + }); + + let error: unknown; + try { + await openai.responses.create({ model: "kimi-k3", input: "hello" }); + } catch (caught) { + error = caught; + } + + expect(sends).toBe(1); + expect(error).toMatchObject({ name: "Error" }); + const capacity = findOpenSecretInferenceCapacityError(error); + expect(capacity).toBeInstanceOf(OpenSecretInferenceCapacityError); + expect(capacity).toMatchObject({ status: 503, retryDelayMs: 0 }); + expect((error as { cause?: unknown }).cause).toBe(capacity); + }); + + test("shares a two-send ceiling across stale-session repair and capacity", async () => { + let currentAttestation = staleAttestation; + let forcedAttestations = 0; + let sends = 0; + const forwardedLimits: Array = []; + const redirects: Array = []; + const customFetch = createCustomFetchWithDependencies( + { apiKey: "test-api-key" }, + dependencies({ + getAttestation: async (forceRefresh) => { + if (forceRefresh) { + forcedAttestations += 1; + currentAttestation = freshAttestation; + } + return currentAttestation; + }, + fetch: async (_input, init) => { + sends += 1; + forwardedLimits.push( + new Headers(init?.headers).get(OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER) + ); + redirects.push(init?.redirect); + if (recordRequest(init).sessionId === staleAttestation.sessionId) { + return contractError(400, "stale session", "session_not_found"); + } + return capacityContractError(503, { retryAfter: "0" }); + } + }) + ); + + let error: unknown; + try { + await customFetch("https://example.test/v1/responses", { + headers: { [OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER]: "2" } + }); + } catch (caught) { + error = caught; + } + + expect(findOpenSecretInferenceCapacityError(error)).toMatchObject({ + status: 503, + retryDelayMs: 0, + inferenceSendCount: 2 + }); + expect(sends).toBe(2); + expect(forcedAttestations).toBe(1); + expect(forwardedLimits).toEqual([null, null]); + expect(redirects).toEqual(["manual", "manual"]); + }); + + test("repairs a stale session but does not exceed a one-send ceiling", async () => { + let currentAttestation = staleAttestation; + let forcedAttestations = 0; + let sends = 0; + const customFetch = createCustomFetchWithDependencies( + { apiKey: "test-api-key" }, + dependencies({ + getAttestation: async (forceRefresh) => { + if (forceRefresh) { + forcedAttestations += 1; + currentAttestation = freshAttestation; + } + return currentAttestation; + }, + fetch: async () => { + sends += 1; + if (sends === 1) return contractError(400, "stale session", "session_not_found"); + return Response.json({ encrypted: '2:{"unexpected":true}' }); + } + }) + ); + + await expect( + customFetch("https://example.test/v1/responses", { + headers: { [OPEN_SECRET_INFERENCE_SEND_LIMIT_HEADER]: "1" } + }) + ).rejects.toMatchObject({ status: 400 }); + + expect(sends).toBe(1); + expect(forcedAttestations).toBe(1); + expect(currentAttestation).toBe(freshAttestation); + }); +}); + describe("createCustomFetch stale-session recovery", () => { beforeEach(() => { window.localStorage.clear(); From 72e8cc3680ae2f050125e606d34cef4450e3ff2a Mon Sep 17 00:00:00 2001 From: Anthony Ronning <101225832+AnthonyRonning@users.noreply.github.com> Date: Sat, 5 Sep 2026 19:12:14 +0000 Subject: [PATCH 2/2] ci(sdk): repin integration backend after transport rebase --- sdk/opensecret-integration-revision | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sdk/opensecret-integration-revision b/sdk/opensecret-integration-revision index b5e1c2af8..85c847cd3 100644 --- a/sdk/opensecret-integration-revision +++ b/sdk/opensecret-integration-revision @@ -1 +1 @@ -91ed2e573aa5f358cfd0d59defd04d414be762fe +d26eb6bd54d50cc8e6b2967f647a94c61da913da