From 19ebee129b274c088d95f3844c02224bcafb61fe Mon Sep 17 00:00:00 2001 From: Chili Palmer Date: Sat, 22 Aug 2026 17:59:46 +0200 Subject: [PATCH] Fix Mentra nested Responses endpoints --- Cargo.lock | 2 - Cargo.toml | 6 +- .../tests/live_llm_vision.rs | 17 +- vendor/mentra-provider/Cargo.toml | 99 + vendor/mentra-provider/LICENSE | 21 + vendor/mentra-provider/README.md | 11 + vendor/mentra-provider/src/anthropic.rs | 271 +++ vendor/mentra-provider/src/anthropic/model.rs | 1397 +++++++++++ vendor/mentra-provider/src/anthropic/sse.rs | 460 ++++ .../src/anthropic/stream_model.rs | 232 ++ vendor/mentra-provider/src/auth.rs | 72 + vendor/mentra-provider/src/definition.rs | 483 ++++ vendor/mentra-provider/src/embedding.rs | 175 ++ vendor/mentra-provider/src/error.rs | 271 +++ vendor/mentra-provider/src/gemini.rs | 266 +++ vendor/mentra-provider/src/gemini/model.rs | 945 ++++++++ vendor/mentra-provider/src/gemini/sse.rs | 714 ++++++ vendor/mentra-provider/src/lib.rs | 154 ++ vendor/mentra-provider/src/model.rs | 509 ++++ vendor/mentra-provider/src/registry.rs | 215 ++ vendor/mentra-provider/src/request.rs | 541 +++++ vendor/mentra-provider/src/response.rs | 913 ++++++++ vendor/mentra-provider/src/responses.rs | 354 +++ vendor/mentra-provider/src/responses/model.rs | 1516 ++++++++++++ .../mentra-provider/src/responses/session.rs | 2041 +++++++++++++++++ vendor/mentra-provider/src/responses/sse.rs | 1542 +++++++++++++ .../src/responses/websocket.rs | 705 ++++++ vendor/mentra-provider/src/stream.rs | 101 + vendor/mentra-provider/src/tool.rs | 151 ++ 29 files changed, 14177 insertions(+), 7 deletions(-) create mode 100644 vendor/mentra-provider/Cargo.toml create mode 100644 vendor/mentra-provider/LICENSE create mode 100644 vendor/mentra-provider/README.md create mode 100644 vendor/mentra-provider/src/anthropic.rs create mode 100644 vendor/mentra-provider/src/anthropic/model.rs create mode 100644 vendor/mentra-provider/src/anthropic/sse.rs create mode 100644 vendor/mentra-provider/src/anthropic/stream_model.rs create mode 100644 vendor/mentra-provider/src/auth.rs create mode 100644 vendor/mentra-provider/src/definition.rs create mode 100644 vendor/mentra-provider/src/embedding.rs create mode 100644 vendor/mentra-provider/src/error.rs create mode 100644 vendor/mentra-provider/src/gemini.rs create mode 100644 vendor/mentra-provider/src/gemini/model.rs create mode 100644 vendor/mentra-provider/src/gemini/sse.rs create mode 100644 vendor/mentra-provider/src/lib.rs create mode 100644 vendor/mentra-provider/src/model.rs create mode 100644 vendor/mentra-provider/src/registry.rs create mode 100644 vendor/mentra-provider/src/request.rs create mode 100644 vendor/mentra-provider/src/response.rs create mode 100644 vendor/mentra-provider/src/responses.rs create mode 100644 vendor/mentra-provider/src/responses/model.rs create mode 100644 vendor/mentra-provider/src/responses/session.rs create mode 100644 vendor/mentra-provider/src/responses/sse.rs create mode 100644 vendor/mentra-provider/src/responses/websocket.rs create mode 100644 vendor/mentra-provider/src/stream.rs create mode 100644 vendor/mentra-provider/src/tool.rs diff --git a/Cargo.lock b/Cargo.lock index 98d3b56..9cbf2a0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4556,8 +4556,6 @@ dependencies = [ [[package]] name = "mentra-provider" version = "0.5.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e9f644037cf45558b393becc1ac080f0d55edced3093a155a7cd3431c1180b7f" dependencies = [ "async-trait", "base64", diff --git a/Cargo.toml b/Cargo.toml index d3b9092..bcb9f28 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -25,7 +25,7 @@ members = [ "tools/concurrency-audit", "tools/performance", ] -exclude = ["crates/libremetaverse-openjpeg"] +exclude = ["crates/libremetaverse-openjpeg", "vendor/mentra-provider"] default-members = [ "crates/metacrate", "crates/libremetaverse-types", @@ -78,3 +78,7 @@ inherits = "dev" opt-level = 1 debug = 0 incremental = false + +[patch.crates-io] +# Mentra 0.5.1 duplicates `v1` for nested OpenAI-compatible base URLs. +mentra-provider = { path = "vendor/mentra-provider" } diff --git a/crates/metacrate-grid-agent/tests/live_llm_vision.rs b/crates/metacrate-grid-agent/tests/live_llm_vision.rs index a857468..96fa18d 100644 --- a/crates/metacrate-grid-agent/tests/live_llm_vision.rs +++ b/crates/metacrate-grid-agent/tests/live_llm_vision.rs @@ -50,10 +50,19 @@ async fn luna_describes_live_renderer_evidence() -> Result<(), Box> { ContentBlock::image_url(image), ]) .await?; - for block in response.content { - if let ContentBlock::Text { text } = block { - println!("LIVE_LLM_VISION={text}"); - } + let descriptions = response + .content + .into_iter() + .filter_map(|block| match block { + ContentBlock::Text { text } if !text.trim().is_empty() => Some(text), + _ => None, + }) + .collect::>(); + if descriptions.is_empty() { + return Err("Luna returned no snapshot description".into()); + } + for description in descriptions { + println!("LIVE_LLM_VISION={description}"); } Ok(()) } diff --git a/vendor/mentra-provider/Cargo.toml b/vendor/mentra-provider/Cargo.toml new file mode 100644 index 0000000..06bbe57 --- /dev/null +++ b/vendor/mentra-provider/Cargo.toml @@ -0,0 +1,99 @@ +# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO +# +# When uploading crates to the registry Cargo will automatically +# "normalize" Cargo.toml files for maximal compatibility +# with all versions of Cargo and also rewrite `path` dependencies +# to registry (e.g., crates.io) dependencies. +# +# If you are reading this file be aware that the original Cargo.toml +# will likely look very different (and much more reasonable). +# See Cargo.toml.orig for the original contents. + +[package] +edition = "2024" +rust-version = "1.88" +name = "mentra-provider" +version = "0.5.1" +build = false +autolib = false +autobins = false +autoexamples = false +autotests = false +autobenches = false +description = "Shared provider core for Mentra" +homepage = "https://github.com/oops-rs/mentra" +documentation = "https://docs.rs/mentra-provider" +readme = "README.md" +license = "MIT" +repository = "https://github.com/oops-rs/mentra" + +[features] +default = ["responses-websocket"] +responses-websocket = [ + "dep:tokio-tungstenite", + "futures-util/sink", +] + +[lib] +name = "mentra_provider" +path = "src/lib.rs" + +[dependencies.async-trait] +version = "0.1.89" + +[dependencies.base64] +version = "0.22.1" + +[dependencies.futures-util] +version = "0.3.31" + +[dependencies.http] +version = "1.3.1" + +[dependencies.reqwest] +version = "0.12" +features = [ + "json", + "rustls-tls", + "stream", +] +default-features = false + +[dependencies.serde] +version = "1.0.228" +features = ["derive"] + +[dependencies.serde_json] +version = "1.0.149" + +[dependencies.strum] +version = "0.27" +features = ["derive"] + +[dependencies.thiserror] +version = "2.0.18" + +[dependencies.time] +version = "0.3" +features = [ + "formatting", + "parsing", + "serde", +] + +[dependencies.tokio] +version = "1.50.0" +features = [ + "macros", + "sync", +] + +[dependencies.tokio-tungstenite] +version = "0.28" +optional = true + +[dependencies.url] +version = "2.5" + +[dependencies.zstd] +version = "0.13" diff --git a/vendor/mentra-provider/LICENSE b/vendor/mentra-provider/LICENSE new file mode 100644 index 0000000..62fd3c7 --- /dev/null +++ b/vendor/mentra-provider/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Wendell Wang + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/vendor/mentra-provider/README.md b/vendor/mentra-provider/README.md new file mode 100644 index 0000000..14033aa --- /dev/null +++ b/vendor/mentra-provider/README.md @@ -0,0 +1,11 @@ +# mentra-provider + +`mentra-provider` is Mentra's publishable provider-core crate. + +It contains provider-neutral request, response, model, streaming, and tool +schema types that can be reused without depending on the full Mentra runtime. + +For most application code, depend on +[`mentra`](https://crates.io/crates/mentra) instead. The `mentra` crate +re-exports these provider-core types and adds the runtime, tooling, +persistence, and collaboration layers. diff --git a/vendor/mentra-provider/src/anthropic.rs b/vendor/mentra-provider/src/anthropic.rs new file mode 100644 index 0000000..5e1ce96 --- /dev/null +++ b/vendor/mentra-provider/src/anthropic.rs @@ -0,0 +1,271 @@ +use async_trait::async_trait; +use serde_json::Value; +use std::collections::HashMap; +use std::sync::Arc; + +pub(crate) mod model; +pub(crate) mod sse; +pub(crate) mod stream_model; + +use crate::AuthScheme; +use crate::BuiltinProvider; +use crate::CompactionRequest; +use crate::CompactionResponse; +use crate::CredentialSource; +use crate::ModelCatalog; +use crate::ModelInfo; +use crate::ProviderCapabilities; +use crate::ProviderDefinition; +use crate::ProviderError; +use crate::ProviderEventStream; +use crate::ProviderSession; +use crate::ProviderSessionFactory; +use crate::RegisteredProvider; +use crate::Request; +use crate::StaticCredentialSource; +use crate::WireApi; + +const DEFAULT_BASE_URL: &str = "https://api.anthropic.com"; +const ANTHROPIC_VERSION: &str = "2023-06-01"; + +/// Returns the default Anthropic-compatible provider definition. +pub fn definition() -> ProviderDefinition { + let mut definition = ProviderDefinition::new(BuiltinProvider::Anthropic); + definition.descriptor.display_name = Some("Anthropic".to_string()); + definition.descriptor.description = Some("Anthropic Messages API provider".to_string()); + definition.wire_api = WireApi::AnthropicMessages; + definition.auth_scheme = AuthScheme::Header { + name: "x-api-key".to_string(), + }; + definition.capabilities = ProviderCapabilities { + supports_model_listing: true, + supports_streaming: true, + supports_websockets: false, + supports_tool_calls: true, + supports_images: true, + supports_history_compaction: true, + supports_memory_summarization: true, + supports_deferred_tools: true, + supports_hosted_tool_search: true, + supports_hosted_web_search: false, + supports_image_generation: false, + supports_reasoning_effort: true, + reports_reasoning_tokens: false, + reports_thoughts_tokens: false, + supports_structured_tool_results: false, + supports_embeddings: false, + }; + definition.base_url = Some(DEFAULT_BASE_URL.to_string()); + definition.headers = Some(HashMap::from([( + "anthropic-version".to_string(), + ANTHROPIC_VERSION.to_string(), + )])); + definition +} + +pub struct AnthropicProvider { + client: reqwest::Client, + credential_source: Arc, + definition: ProviderDefinition, +} + +impl Clone for AnthropicProvider { + fn clone(&self) -> Self { + Self { + client: self.client.clone(), + credential_source: Arc::clone(&self.credential_source), + definition: self.definition.clone(), + } + } +} + +impl AnthropicProvider { + pub fn new(api_key: impl Into) -> Self { + Self::with_credential_source(StaticCredentialSource::new(api_key)) + } +} + +impl AnthropicProvider +where + C: CredentialSource + 'static, +{ + pub fn with_credential_source(credential_source: C) -> Self { + Self::with_shared_credential_source(Arc::new(credential_source)) + } + + pub fn with_shared_credential_source(credential_source: Arc) -> Self { + Self::with_definition_and_shared_credential_source(definition(), credential_source) + } + + pub fn with_definition_and_credential_source( + definition: ProviderDefinition, + credential_source: C, + ) -> Self { + Self::with_definition_and_shared_credential_source(definition, Arc::new(credential_source)) + } + + pub fn with_definition_and_shared_credential_source( + definition: ProviderDefinition, + credential_source: Arc, + ) -> Self { + let client = reqwest::Client::builder() + .build() + .expect("Failed to build client"); + + Self { + client, + credential_source, + definition, + } + } +} + +#[async_trait] +impl ModelCatalog for AnthropicProvider +where + C: CredentialSource + 'static, +{ + async fn list_models(&self) -> Result, ProviderError> { + let mut models = Vec::new(); + let mut after_id = None; + + loop { + let credentials = self.credential_source.credentials().await?; + let request = self + .client + .get( + self.definition + .request_url_with_auth_for_path("v1/models", &credentials)?, + ) + .headers(self.definition.build_headers(&credentials)?) + .query(&[ + ("limit", "1000"), + ("after_id", after_id.as_deref().unwrap_or("")), + ]); + + let response = request.send().await.map_err(ProviderError::Transport)?; + + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + let page = response + .json::() + .await + .map_err(ProviderError::Decode)?; + + after_id = page.last_id.clone(); + models.extend(page.data.into_iter().map(|model| model.into())); + + if !page.has_more { + break; + } + } + + Ok(models) + } +} + +#[async_trait] +impl ProviderSessionFactory for AnthropicProvider +where + C: CredentialSource + 'static, +{ + async fn create_session(&self) -> Result, ProviderError> { + Ok(Box::new((*self).clone())) + } +} + +#[async_trait] +impl ProviderSession for AnthropicProvider +where + C: CredentialSource + 'static, +{ + async fn stream(&self, request: Request<'_>) -> Result { + let requested_model = request.model.to_string(); + let provider = self.definition.provider_id().clone(); + let response = self.send_message(request, true).await?; + Ok(sse::spawn_event_stream(response, provider, requested_model)) + } + + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + let request = request.into_model_request()?; + let response = ProviderSession::send(self, request).await?; + Ok(response.into_compaction_response()) + } + + async fn summarize_memories( + &self, + request: crate::MemorySummarizeRequest<'_>, + ) -> Result { + let request = request.into_model_request()?; + let response = ProviderSession::send(self, request).await?; + response.into_memory_summarize_response() + } +} + +#[async_trait] +impl RegisteredProvider for AnthropicProvider +where + C: CredentialSource + 'static, +{ + fn definition(&self) -> ProviderDefinition { + self.definition.clone() + } +} + +impl AnthropicProvider +where + C: CredentialSource + 'static, +{ + async fn send_message( + &self, + request: Request<'_>, + stream: bool, + ) -> Result { + let session = request.provider_request_options.session.clone(); + let request = model::AnthropicRequest::try_from_with_provider( + request, + self.definition.provider_id(), + )?; + let mut body = serde_json::to_value(request).map_err(ProviderError::Serialize)?; + if stream { + body["stream"] = Value::Bool(true); + } + let credentials = self.credential_source.credentials().await?; + let response = self + .client + .post( + self.definition + .request_url_with_auth_for_path("v1/messages", &credentials)?, + ) + .headers(self.definition.build_headers_for_session( + &credentials, + Some(&session), + None, + )?) + .json(&body) + .send() + .await + .map_err(ProviderError::Transport)?; + + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + Ok(response) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn definition_advertises_history_compaction_support() { + assert!(definition().capabilities.supports_history_compaction); + } +} diff --git a/vendor/mentra-provider/src/anthropic/model.rs b/vendor/mentra-provider/src/anthropic/model.rs new file mode 100644 index 0000000..094d7f7 --- /dev/null +++ b/vendor/mentra-provider/src/anthropic/model.rs @@ -0,0 +1,1397 @@ +use base64::{Engine as _, engine::general_purpose::STANDARD}; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use time::{OffsetDateTime, format_description::well_known::Rfc3339}; + +use crate::{ + BuiltinProvider, ContentBlock, ImageSource, Message, ModelInfo, ProviderError, ProviderId, + ProviderToolKind, ReasoningEffort, ReasoningFormat, ReasoningProvenance, Request, Response, + Role, TokenUsage, ToolChoice, ToolLoadingPolicy, ToolResultContent, ToolSearchMode, ToolSpec, +}; + +#[derive(Deserialize)] +pub(crate) struct AnthropicModelsPage { + pub(crate) data: Vec, + pub(crate) has_more: bool, + pub(crate) last_id: Option, +} + +#[derive(Deserialize)] +pub(crate) struct AnthropicModel { + pub(crate) id: String, + #[serde(default)] + pub(crate) display_name: Option, + #[serde(default)] + pub(crate) created_at: Option, +} + +impl From for ModelInfo { + fn from(model: AnthropicModel) -> Self { + ModelInfo { + id: model.id, + provider: BuiltinProvider::Anthropic.into(), + display_name: model.display_name, + description: None, + created_at: model + .created_at + .as_deref() + .and_then(|value| OffsetDateTime::parse(value, &Rfc3339).ok()), + } + } +} + +#[derive(Serialize)] +pub(crate) struct AnthropicRequest { + model: String, + #[serde(skip_serializing_if = "Option::is_none")] + system: Option>, + messages: Vec, + #[serde(skip_serializing_if = "Vec::is_empty")] + tools: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + tool_choice: Option, + #[serde(skip_serializing_if = "Option::is_none")] + temperature: Option, + #[serde(rename = "max_tokens", skip_serializing_if = "Option::is_none")] + max_output_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + disable_parallel_tool_use: Option, + #[serde(skip_serializing_if = "Option::is_none")] + thinking: Option, + #[serde(skip_serializing_if = "Option::is_none")] + output_config: Option, +} + +/// A system prompt block with optional cache control. +#[derive(Serialize)] +pub(crate) struct AnthropicSystemBlock { + #[serde(rename = "type")] + kind: &'static str, + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + cache_control: Option, +} + +/// Cache control marker for Anthropic prompt caching. +#[derive(Serialize)] +pub(crate) struct AnthropicCacheControl { + #[serde(rename = "type")] + kind: &'static str, +} + +impl AnthropicCacheControl { + fn ephemeral() -> Self { + Self { kind: "ephemeral" } + } +} + +/// Build system blocks from a system prompt string, with cache_control on +/// the final block to enable prompt caching. +fn build_system_blocks(system: String) -> Vec { + vec![AnthropicSystemBlock { + kind: "text", + text: system, + cache_control: Some(AnthropicCacheControl::ephemeral()), + }] +} + +#[derive(Deserialize)] +pub(crate) struct AnthropicResponse { + pub(crate) id: String, + pub(crate) model: String, + pub(crate) role: String, + #[serde(default)] + pub(crate) usage: Option, + content: Vec, + stop_reason: Option, +} + +impl TryFrom for Response { + type Error = ProviderError; + + fn try_from(response: AnthropicResponse) -> Result { + let provider = ProviderId::from(BuiltinProvider::Anthropic); + let requested_model = response.model.clone(); + response.try_into_response_with_provider(&provider, &requested_model) + } +} + +impl AnthropicResponse { + fn try_into_response_with_provider( + self, + provider: &ProviderId, + requested_model: &str, + ) -> Result { + let provenance = ReasoningProvenance { + provider: provider.clone(), + model: requested_model.to_string(), + format: ReasoningFormat::AnthropicSigned, + }; + Ok(Response { + id: self.id, + model: self.model, + role: match self.role.as_str() { + "user" => Role::User, + "assistant" => Role::Assistant, + _ => Role::Unknown(self.role), + }, + content: self + .content + .into_iter() + .map(|block| ContentBlock::try_from((block, provenance.clone()))) + .collect::, _>>()?, + stop_reason: self.stop_reason, + usage: self.usage.and_then(|usage| usage.into_token_usage()), + }) + } +} + +#[derive(Debug, Clone, Deserialize)] +pub(crate) struct AnthropicUsage { + #[serde(default)] + pub(crate) input_tokens: Option, + #[serde(default)] + pub(crate) output_tokens: Option, + #[serde(default)] + pub(crate) cache_read_input_tokens: Option, + #[serde(default)] + pub(crate) cache_creation_input_tokens: Option, + #[serde(default)] + pub(crate) total_tokens: Option, +} + +impl AnthropicUsage { + pub(crate) fn into_token_usage(self) -> Option { + let usage = TokenUsage { + input_tokens: self.input_tokens, + output_tokens: self.output_tokens, + total_tokens: self.total_tokens, + cache_read_input_tokens: self.cache_read_input_tokens, + cache_creation_input_tokens: self.cache_creation_input_tokens, + reasoning_tokens: None, + thoughts_tokens: None, + tool_input_tokens: None, + }; + + (!usage.is_empty()).then_some(usage) + } +} + +impl<'a> TryFrom> for AnthropicRequest { + type Error = ProviderError; + + fn try_from(value: Request<'a>) -> Result { + let provider = ProviderId::from(BuiltinProvider::Anthropic); + Self::try_from_with_provider(value, &provider) + } +} + +impl AnthropicRequest { + pub(crate) fn try_from_with_provider( + value: Request<'_>, + target_provider: &ProviderId, + ) -> Result { + let reasoning_effort = value + .provider_request_options + .reasoning + .as_ref() + .and_then(|reasoning| reasoning.effort); + + let effort_capabilities = reasoning_effort + .map(|effort| { + let capabilities = + anthropic_effort_capabilities(&value.model).ok_or_else(|| { + ProviderError::InvalidRequest(format!( + "Anthropic reasoning effort is not supported by model '{}'", + value.model + )) + })?; + if matches!(effort, ReasoningEffort::Max) && !capabilities.max { + return Err(ProviderError::InvalidRequest(format!( + "Anthropic max reasoning effort is not supported by model '{}'", + value.model + ))); + } + if matches!(effort, ReasoningEffort::XHigh) && !capabilities.xhigh { + return Err(ProviderError::InvalidRequest(format!( + "Anthropic xhigh reasoning effort is not supported by model '{}'", + value.model + ))); + } + Ok(capabilities) + }) + .transpose()?; + + let target_model = value.model.to_string(); + + Ok(AnthropicRequest { + model: value.model.into_owned(), + system: value.system.map(|s| build_system_blocks(s.into_owned())), + messages: value + .messages + .iter() + .map(|message| { + AnthropicMessage::try_from_with_target(message, target_provider, &target_model) + }) + .collect::, _>>()?, + tools: build_anthropic_tools( + value.tools.as_ref(), + value.tool_choice.as_ref(), + value.provider_request_options.tool_search_mode, + )?, + tool_choice: value.tool_choice.map(AnthropicToolChoice::from), + temperature: value.temperature, + max_output_tokens: value.max_output_tokens, + disable_parallel_tool_use: value + .provider_request_options + .anthropic + .disable_parallel_tool_use, + thinking: effort_capabilities + .filter(|capabilities| capabilities.adaptive_thinking) + .map(|_| AnthropicThinkingConfig::adaptive()), + output_config: reasoning_effort.map(AnthropicOutputConfig::new), + }) + } +} + +#[derive(Serialize)] +struct AnthropicThinkingConfig { + #[serde(rename = "type")] + kind: &'static str, +} + +impl AnthropicThinkingConfig { + fn adaptive() -> Self { + Self { kind: "adaptive" } + } +} + +#[derive(Serialize)] +struct AnthropicOutputConfig { + effort: AnthropicReasoningEffort, +} + +impl AnthropicOutputConfig { + fn new(effort: ReasoningEffort) -> Self { + Self { + effort: effort.into(), + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "snake_case")] +enum AnthropicReasoningEffort { + Low, + Medium, + High, + #[serde(rename = "xhigh")] + XHigh, + Max, +} + +impl From for AnthropicReasoningEffort { + fn from(value: ReasoningEffort) -> Self { + match value { + ReasoningEffort::Low => Self::Low, + ReasoningEffort::Medium => Self::Medium, + ReasoningEffort::High => Self::High, + ReasoningEffort::XHigh => Self::XHigh, + ReasoningEffort::Max => Self::Max, + } + } +} + +#[derive(Clone, Copy)] +struct AnthropicEffortCapabilities { + adaptive_thinking: bool, + max: bool, + xhigh: bool, +} + +impl AnthropicEffortCapabilities { + const BASIC: Self = Self { + adaptive_thinking: false, + max: false, + xhigh: false, + }; + const ADAPTIVE_WITH_MAX: Self = Self { + adaptive_thinking: true, + max: true, + xhigh: false, + }; + const ALL: Self = Self { + adaptive_thinking: true, + max: true, + xhigh: true, + }; +} + +fn anthropic_effort_capabilities(model: &str) -> Option { + let model = model.strip_prefix("models/").unwrap_or(model); + if matches_anthropic_model(model, "claude-opus-4-5") { + Some(AnthropicEffortCapabilities::BASIC) + } else if matches_anthropic_model(model, "claude-mythos-preview") + || matches_anthropic_model(model, "claude-opus-4-6") + || matches_anthropic_model(model, "claude-sonnet-4-6") + { + Some(AnthropicEffortCapabilities::ADAPTIVE_WITH_MAX) + } else if matches_anthropic_model(model, "claude-opus-4-7") + || matches_anthropic_model(model, "claude-opus-4-8") + || matches_anthropic_model(model, "claude-opus-5") + || matches_anthropic_model(model, "claude-sonnet-5") + || matches_anthropic_model(model, "claude-fable-5") + || matches_anthropic_model(model, "claude-mythos-5") + { + Some(AnthropicEffortCapabilities::ALL) + } else { + None + } +} + +fn matches_anthropic_model(model: &str, canonical: &str) -> bool { + model == canonical + || model + .strip_prefix(canonical) + .and_then(|suffix| suffix.strip_prefix('-')) + .is_some_and(|snapshot| { + snapshot.len() == 8 && snapshot.bytes().all(|byte| byte.is_ascii_digit()) + }) +} + +#[derive(Serialize)] +struct AnthropicMessage { + role: String, + content: Vec, +} + +impl TryFrom for AnthropicMessage { + type Error = ProviderError; + + fn try_from(message: Message) -> Result { + AnthropicMessage::try_from(&message) + } +} + +impl TryFrom<&Message> for AnthropicMessage { + type Error = ProviderError; + + fn try_from(message: &Message) -> Result { + if !matches!(message.role, Role::User) && message_has_image(message) { + return Err(ProviderError::InvalidRequest( + "Anthropic image inputs are only supported in user messages".to_string(), + )); + } + + Ok(AnthropicMessage { + role: message.role.to_string(), + content: message + .content + .iter() + .map(AnthropicContentBlock::from_without_replay_target) + .collect(), + }) + } +} + +impl AnthropicMessage { + fn try_from_with_target( + message: &Message, + provider: &ProviderId, + model: &str, + ) -> Result { + if !matches!(message.role, Role::User) && message_has_image(message) { + return Err(ProviderError::InvalidRequest( + "Anthropic image inputs are only supported in user messages".to_string(), + )); + } + + let target = AnthropicReplayTarget { + provider, + model, + role: &message.role, + }; + Ok(Self { + role: message.role.to_string(), + content: message + .content + .iter() + .map(|block| AnthropicContentBlock::from_with_replay_target(block, &target)) + .collect(), + }) + } +} + +struct AnthropicReplayTarget<'a> { + provider: &'a ProviderId, + model: &'a str, + role: &'a Role, +} + +#[derive(Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum AnthropicContentBlock { + Text { + text: String, + }, + Thinking { + thinking: String, + signature: String, + }, + RedactedThinking { + data: String, + }, + Image { + source: AnthropicImageSource, + }, + ToolUse { + id: String, + name: String, + input: Value, + }, + ToolResult { + tool_use_id: String, + content: String, + is_error: bool, + }, +} + +#[derive(Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum AnthropicImageSource { + Base64 { media_type: String, data: String }, + Url { url: String }, +} + +impl From for AnthropicContentBlock { + fn from(block: ContentBlock) -> Self { + AnthropicContentBlock::from(&block) + } +} + +impl From<&ContentBlock> for AnthropicContentBlock { + fn from(block: &ContentBlock) -> Self { + Self::from_without_replay_target(block) + } +} + +impl AnthropicContentBlock { + fn from_without_replay_target(block: &ContentBlock) -> Self { + match block { + ContentBlock::Thinking { .. } => AnthropicContentBlock::Text { + text: block + .thinking_fallback_text() + .expect("thinking block has fallback text"), + }, + _ => Self::from_non_thinking(block), + } + } + + fn from_with_replay_target(block: &ContentBlock, target: &AnthropicReplayTarget<'_>) -> Self { + match block { + ContentBlock::Thinking { + thinking, + signature: Some(signature), + provenance: Some(provenance), + redacted, + .. + } if matches!(target.role, Role::Assistant) + && !signature.is_empty() + && provenance.provider == *target.provider + && provenance.model == target.model + && provenance.format == ReasoningFormat::AnthropicSigned => + { + if *redacted { + Self::RedactedThinking { + data: signature.clone(), + } + } else { + Self::Thinking { + thinking: thinking.clone(), + signature: signature.clone(), + } + } + } + _ => Self::from_without_replay_target(block), + } + } + + fn from_non_thinking(block: &ContentBlock) -> Self { + match block { + ContentBlock::Text { text } => AnthropicContentBlock::Text { text: text.clone() }, + ContentBlock::Thinking { .. } => { + unreachable!("thinking blocks are handled before non-thinking projection") + } + ContentBlock::Image { source } => AnthropicContentBlock::Image { + source: source.into(), + }, + ContentBlock::ToolUse { id, name, input } => AnthropicContentBlock::ToolUse { + id: id.clone(), + name: name.clone(), + input: input.clone(), + }, + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => AnthropicContentBlock::ToolResult { + tool_use_id: tool_use_id.clone(), + content: content.to_display_string(), + is_error: *is_error, + }, + ContentBlock::HostedToolSearch { call } => AnthropicContentBlock::ToolUse { + id: call.id.clone(), + name: "tool_search".to_string(), + input: serde_json::json!({ "query": call.query }), + }, + ContentBlock::HostedWebSearch { call } => AnthropicContentBlock::ToolUse { + id: call.id.clone(), + name: "web_search".to_string(), + input: serde_json::to_value(call.action.clone()).unwrap_or(serde_json::Value::Null), + }, + ContentBlock::ImageGeneration { call } => AnthropicContentBlock::ToolUse { + id: call.id.clone(), + name: "image_generation".to_string(), + input: serde_json::json!({ + "status": call.status, + "revised_prompt": call.revised_prompt, + }), + }, + } + } +} + +impl TryFrom<(AnthropicContentBlock, ReasoningProvenance)> for ContentBlock { + type Error = ProviderError; + + fn try_from( + (block, provenance): (AnthropicContentBlock, ReasoningProvenance), + ) -> Result { + Ok(match block { + AnthropicContentBlock::Text { text } => ContentBlock::Text { text }, + AnthropicContentBlock::Thinking { + thinking, + signature, + } => ContentBlock::Thinking { + thinking, + signature: Some(signature), + encrypted_content: None, + id: None, + provenance: Some(provenance), + redacted: false, + }, + AnthropicContentBlock::RedactedThinking { data } => ContentBlock::Thinking { + thinking: String::new(), + signature: Some(data), + encrypted_content: None, + id: None, + provenance: Some(provenance), + redacted: true, + }, + AnthropicContentBlock::Image { source } => ContentBlock::Image { + source: source.try_into()?, + }, + AnthropicContentBlock::ToolUse { id, name, input } => { + ContentBlock::ToolUse { id, name, input } + } + AnthropicContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => ContentBlock::ToolResult { + tool_use_id, + content: ToolResultContent::Text(content), + is_error, + }, + }) + } +} + +impl From<&ImageSource> for AnthropicImageSource { + fn from(value: &ImageSource) -> Self { + match value { + ImageSource::Bytes { media_type, data } => AnthropicImageSource::Base64 { + media_type: media_type.clone(), + data: STANDARD.encode(data), + }, + ImageSource::Url { url } => AnthropicImageSource::Url { url: url.clone() }, + } + } +} + +impl From for AnthropicImageSource { + fn from(value: ImageSource) -> Self { + AnthropicImageSource::from(&value) + } +} + +impl TryFrom for ImageSource { + type Error = ProviderError; + + fn try_from(value: AnthropicImageSource) -> Result { + match value { + AnthropicImageSource::Base64 { media_type, data } => { + let data = STANDARD.decode(data).map_err(|error| { + ProviderError::InvalidResponse(format!( + "invalid Anthropic image payload for media type {media_type}: {error}" + )) + })?; + Ok(ImageSource::Bytes { media_type, data }) + } + AnthropicImageSource::Url { url } => Ok(ImageSource::Url { url }), + } + } +} + +#[derive(Serialize)] +#[serde(untagged)] +enum AnthropicTool { + Custom(AnthropicCustomTool), + HostedSearch(AnthropicHostedSearchTool), +} + +#[derive(Serialize)] +struct AnthropicCustomTool { + name: String, + #[serde(skip_serializing_if = "Option::is_none")] + description: Option, + input_schema: Value, + #[serde(skip_serializing_if = "std::ops::Not::not")] + defer_loading: bool, + #[serde(skip_serializing_if = "Option::is_none")] + cache_control: Option, +} + +#[derive(Serialize)] +struct AnthropicHostedSearchTool { + #[serde(rename = "type")] + kind: &'static str, + name: &'static str, +} + +impl AnthropicTool { + fn custom(tool: &ToolSpec, force_immediate: bool, is_last: bool) -> Self { + Self::Custom(AnthropicCustomTool { + name: tool.name.clone(), + description: tool.description.clone(), + input_schema: tool.input_schema.clone(), + defer_loading: tool.loading_policy == ToolLoadingPolicy::Deferred && !force_immediate, + cache_control: if is_last { + Some(AnthropicCacheControl::ephemeral()) + } else { + None + }, + }) + } + + fn hosted_search() -> Self { + Self::HostedSearch(AnthropicHostedSearchTool { + kind: "tool_search_tool_bm25_20251119", + name: "tool_search_tool_bm25", + }) + } +} + +fn build_anthropic_tools( + tools: &[ToolSpec], + tool_choice: Option<&ToolChoice>, + tool_search_mode: ToolSearchMode, +) -> Result, ProviderError> { + if let Some(tool) = tools + .iter() + .find(|tool| tool.kind != ProviderToolKind::Function) + { + return Err(ProviderError::InvalidRequest(format!( + "Anthropic does not support provider tool kind {:?} for '{}'", + tool.kind, tool.name + ))); + } + + let forced_tool_name = match tool_choice { + Some(ToolChoice::Tool { name }) => Some(name.as_str()), + _ => None, + }; + + let has_deferred_tools = tools.iter().any(|tool| { + tool.loading_policy == ToolLoadingPolicy::Deferred + && forced_tool_name != Some(tool.name.as_str()) + }); + + if has_deferred_tools && tool_search_mode != ToolSearchMode::Hosted { + return Err(ProviderError::InvalidRequest( + "Anthropic deferred tools require hosted tool search".to_string(), + )); + } + + let tool_count = tools.len(); + let mut provider_tools = tools + .iter() + .enumerate() + .map(|(i, tool)| { + let is_last = i == tool_count - 1 && !has_deferred_tools; + AnthropicTool::custom(tool, forced_tool_name == Some(tool.name.as_str()), is_last) + }) + .collect::>(); + + if has_deferred_tools { + provider_tools.push(AnthropicTool::hosted_search()); + } + + Ok(provider_tools) +} + +#[derive(Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum AnthropicToolChoice { + Auto, + Any, + Tool { name: String }, +} + +impl From for AnthropicToolChoice { + fn from(choice: ToolChoice) -> Self { + match choice { + ToolChoice::Auto => AnthropicToolChoice::Auto, + ToolChoice::Any => AnthropicToolChoice::Any, + ToolChoice::Tool { name } => AnthropicToolChoice::Tool { name }, + } + } +} + +fn message_has_image(message: &Message) -> bool { + message + .content + .iter() + .any(|block| matches!(block, ContentBlock::Image { .. })) +} + +#[cfg(test)] +mod tests { + use std::{borrow::Cow, collections::BTreeMap}; + + use time::{OffsetDateTime, format_description::well_known::Rfc3339}; + + use crate::{ + AnthropicRequestOptions, ContentBlock, ImageSource, Message, ModelInfo, ProviderError, + ProviderId, ProviderRequestOptions, ReasoningEffort, ReasoningFormat, ReasoningOptions, + ReasoningProvenance, Request, Role, ToolChoice, ToolLoadingPolicy, ToolResultContent, + ToolSearchMode, ToolSpec, + }; + + use super::{ + AnthropicContentBlock, AnthropicImageSource, AnthropicModel, AnthropicRequest, + AnthropicResponse, + }; + + fn request_with_message(model: &str, message: Message) -> Request<'static> { + Request { + model: Cow::Owned(model.to_string()), + system: None, + messages: Cow::Owned(vec![message]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: Some(512), + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + } + } + + fn request_with_effort(model: &str, effort: ReasoningEffort) -> Request<'static> { + Request { + model: Cow::Owned(model.to_string()), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: Some(512), + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(ReasoningOptions { + effort: Some(effort), + summary: None, + }), + ..Default::default() + }, + } + } + + fn anthropic_thinking( + thinking: &str, + signature: Option<&str>, + provider: &str, + model: &str, + redacted: bool, + ) -> ContentBlock { + ContentBlock::Thinking { + thinking: thinking.to_string(), + signature: signature.map(str::to_string), + encrypted_content: None, + id: None, + provenance: Some(ReasoningProvenance { + provider: ProviderId::new(provider), + model: model.to_string(), + format: ReasoningFormat::AnthropicSigned, + }), + redacted, + } + } + + #[test] + fn converts_rfc3339_timestamp_to_offset_datetime() { + let raw = "2025-03-04T12:34:56Z"; + let model = AnthropicModel { + id: "claude-test".to_string(), + display_name: None, + created_at: Some(raw.to_string()), + }; + + let info = ModelInfo::from(model); + + assert_eq!( + info.created_at, + Some(OffsetDateTime::parse(raw, &Rfc3339).expect("valid rfc3339")) + ); + } + + #[test] + fn serializes_inline_images_into_anthropic_content_blocks() { + let request = Request { + model: Cow::Borrowed("claude-sonnet"), + system: None, + messages: Cow::Owned(vec![Message { + role: Role::User, + content: vec![ + ContentBlock::text("Describe this"), + ContentBlock::image_bytes("image/png", [1_u8, 2, 3]), + ContentBlock::ToolResult { + tool_use_id: "call_1".to_string(), + content: ToolResultContent::text("ok"), + is_error: false, + }, + ], + }]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: Some(0.1), + max_output_tokens: Some(512), + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["messages"][0]["role"], "user"); + assert_eq!(payload["messages"][0]["content"][0]["type"], "text"); + assert_eq!( + payload["messages"][0]["content"][0]["text"], + "Describe this" + ); + assert_eq!(payload["messages"][0]["content"][1]["type"], "image"); + assert_eq!( + payload["messages"][0]["content"][1]["source"]["type"], + "base64" + ); + assert_eq!( + payload["messages"][0]["content"][1]["source"]["media_type"], + "image/png" + ); + assert_eq!( + payload["messages"][0]["content"][1]["source"]["data"], + "AQID" + ); + assert_eq!(payload["messages"][0]["content"][2]["type"], "tool_result"); + assert_eq!(payload["max_tokens"], 512); + let temperature = payload["temperature"] + .as_f64() + .expect("temperature should be numeric"); + assert!((temperature - 0.1).abs() < 1e-6); + } + + #[test] + fn rejects_invalid_base64_image_payloads() { + let error = ImageSource::try_from(AnthropicImageSource::Base64 { + media_type: "image/png".to_string(), + data: "!not-base64!".to_string(), + }) + .expect_err("invalid base64 should fail"); + + match error { + ProviderError::InvalidResponse(message) => { + assert!(message.contains("invalid Anthropic image payload")); + assert!(message.contains("image/png")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn replays_signed_and_redacted_thinking_only_to_exact_provider_and_model() { + let provider = ProviderId::new("anthropic-edge"); + let request = request_with_message( + "claude-requested", + Message { + role: Role::Assistant, + content: vec![ + anthropic_thinking( + "private chain", + Some("opaque-signature"), + "anthropic-edge", + "claude-requested", + false, + ), + anthropic_thinking( + "", + Some("opaque-redacted-data"), + "anthropic-edge", + "claude-requested", + true, + ), + ], + }, + ); + + let payload = serde_json::to_value( + AnthropicRequest::try_from_with_provider(request, &provider).unwrap(), + ) + .unwrap(); + + assert_eq!(payload["messages"][0]["content"][0]["type"], "thinking"); + assert_eq!( + payload["messages"][0]["content"][0]["thinking"], + "private chain" + ); + assert_eq!( + payload["messages"][0]["content"][0]["signature"], + "opaque-signature" + ); + assert_eq!( + payload["messages"][0]["content"][1]["type"], + "redacted_thinking" + ); + assert_eq!( + payload["messages"][0]["content"][1]["data"], + "opaque-redacted-data" + ); + } + + #[test] + fn downgrades_unreplayable_thinking_to_nonempty_text() { + let provider = ProviderId::new("anthropic-edge"); + let cases = [ + anthropic_thinking( + "wrong provider", + Some("signature"), + "anthropic-other", + "claude-requested", + false, + ), + anthropic_thinking( + "wrong model", + Some("signature"), + "anthropic-edge", + "claude-other", + false, + ), + anthropic_thinking( + "missing signature", + None, + "anthropic-edge", + "claude-requested", + false, + ), + anthropic_thinking( + "empty signature", + Some(""), + "anthropic-edge", + "claude-requested", + false, + ), + anthropic_thinking("", None, "anthropic-edge", "claude-requested", true), + ]; + let request = request_with_message( + "claude-requested", + Message { + role: Role::Assistant, + content: cases.to_vec(), + }, + ); + + let payload = serde_json::to_value( + AnthropicRequest::try_from_with_provider(request, &provider).unwrap(), + ) + .unwrap(); + let content = payload["messages"][0]["content"].as_array().unwrap(); + + assert_eq!(content.len(), cases.len()); + assert!(content.iter().all(|block| block["type"] == "text")); + assert!( + content + .iter() + .all(|block| { block["text"].as_str().is_some_and(|text| !text.is_empty()) }) + ); + assert_eq!(content[4]["text"], "[redacted reasoning]"); + } + + #[test] + fn downgrades_user_role_thinking_even_with_matching_provenance() { + let provider = ProviderId::new("anthropic-edge"); + let request = request_with_message( + "claude-requested", + Message::user(anthropic_thinking( + "private chain", + Some("opaque-signature"), + "anthropic-edge", + "claude-requested", + false, + )), + ); + + let payload = serde_json::to_value( + AnthropicRequest::try_from_with_provider(request, &provider).unwrap(), + ) + .unwrap(); + + assert_eq!(payload["messages"][0]["content"][0]["type"], "text"); + assert_eq!( + payload["messages"][0]["content"][0]["text"], + "private chain" + ); + } + + #[test] + fn non_stream_response_captures_thinking_and_redacted_data() { + let response = AnthropicResponse { + id: "msg-1".to_string(), + model: "claude-resolved".to_string(), + role: "assistant".to_string(), + usage: None, + content: vec![ + AnthropicContentBlock::Thinking { + thinking: "private chain".to_string(), + signature: "opaque-signature".to_string(), + }, + AnthropicContentBlock::RedactedThinking { + data: "opaque-redacted-data".to_string(), + }, + ], + stop_reason: Some("end_turn".to_string()), + }; + + let converted = response + .try_into_response_with_provider(&ProviderId::new("anthropic-edge"), "claude-requested") + .unwrap(); + + assert_eq!( + converted.content, + vec![ + anthropic_thinking( + "private chain", + Some("opaque-signature"), + "anthropic-edge", + "claude-requested", + false, + ), + anthropic_thinking( + "", + Some("opaque-redacted-data"), + "anthropic-edge", + "claude-requested", + true, + ), + ] + ); + } + + #[test] + fn serializes_disable_parallel_tool_use_option() { + let request = Request { + model: Cow::Borrowed("claude-sonnet"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Disabled, + reasoning: None, + responses: Default::default(), + anthropic: AnthropicRequestOptions { + disable_parallel_tool_use: Some(true), + }, + gemini: Default::default(), + session: Default::default(), + }, + }; + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["disable_parallel_tool_use"], true); + } + + #[test] + fn nests_reasoning_effort_under_output_config_with_adaptive_thinking() { + let request = request_with_effort("claude-sonnet-4-6", ReasoningEffort::Medium); + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["thinking"]["type"], "adaptive"); + assert_eq!(payload["output_config"]["effort"], "medium"); + assert!(payload.get("effort").is_none()); + } + + #[test] + fn opus_4_5_serializes_effort_without_adaptive_thinking() { + let request = request_with_effort("claude-opus-4-5-20251101", ReasoningEffort::Medium); + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["output_config"]["effort"], "medium"); + assert!(payload.get("thinking").is_none()); + assert!(payload.get("effort").is_none()); + } + + #[test] + fn mythos_preview_supports_max_with_adaptive_thinking() { + let request = request_with_effort("claude-mythos-preview", ReasoningEffort::Max); + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["output_config"]["effort"], "max"); + assert_eq!(payload["thinking"]["type"], "adaptive"); + assert!(payload.get("effort").is_none()); + } + + #[test] + fn omits_anthropic_reasoning_fields_without_an_effort() { + let request = request_with_message( + "claude-sonnet-4-6", + Message::user(ContentBlock::text("hello")), + ); + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert!(payload.get("thinking").is_none()); + assert!(payload.get("output_config").is_none()); + assert!(payload.get("effort").is_none()); + } + + #[test] + fn serializes_all_anthropic_effort_tiers_exactly() { + let cases = [ + (ReasoningEffort::Low, "low", "claude-opus-4-5"), + (ReasoningEffort::Medium, "medium", "claude-opus-4-5"), + (ReasoningEffort::High, "high", "claude-opus-4-5"), + (ReasoningEffort::XHigh, "xhigh", "claude-opus-5"), + (ReasoningEffort::Max, "max", "claude-sonnet-4-6"), + ]; + + for (effort, expected, model) in cases { + let request = request_with_effort(model, effort); + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["output_config"]["effort"], expected); + assert!(payload.get("effort").is_none()); + } + } + + #[test] + fn supports_xhigh_on_documented_anthropic_models() { + for model in [ + "claude-opus-4-7", + "claude-opus-4-8", + "claude-opus-5", + "claude-sonnet-5", + "claude-fable-5", + "claude-mythos-5", + ] { + let request = request_with_effort(model, ReasoningEffort::XHigh); + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["output_config"]["effort"], "xhigh"); + } + } + + #[test] + fn rejects_xhigh_for_claude_4_6() { + let request = request_with_effort("claude-opus-4-6", ReasoningEffort::XHigh); + + let error = AnthropicRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("xhigh")); + assert!(message.contains("claude-opus-4-6")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn rejects_max_and_xhigh_for_opus_4_5() { + for effort in [ReasoningEffort::Max, ReasoningEffort::XHigh] { + let request = request_with_effort("claude-opus-4-5", effort); + + let error = AnthropicRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("not supported")); + assert!(message.contains("claude-opus-4-5")); + } + other => panic!("unexpected error: {other:?}"), + } + } + } + + #[test] + fn rejects_reasoning_effort_for_unsupported_anthropic_models() { + let request = request_with_effort("claude-sonnet-4-5", ReasoningEffort::Low); + + let error = AnthropicRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("not supported")); + assert!(message.contains("claude-sonnet-4-5")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn rejects_reasoning_effort_for_unknown_anthropic_models() { + let request = request_with_effort("claude-opus-6", ReasoningEffort::Low); + + let error = AnthropicRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("not supported")); + assert!(message.contains("claude-opus-6")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn hosted_tool_search_adds_search_tool_for_deferred_tools() { + let request = Request { + model: Cow::Borrowed("claude-sonnet"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hello"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "lookup_order".to_string(), + description: Some("Look up an order".to_string()), + input_schema: serde_json::json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + ..Default::default() + }, + }; + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["tools"][0]["name"], "lookup_order"); + assert_eq!(payload["tools"][0]["defer_loading"], true); + assert_eq!( + payload["tools"][1]["type"], + "tool_search_tool_bm25_20251119" + ); + assert_eq!(payload["tools"][1]["name"], "tool_search_tool_bm25"); + } + + #[test] + fn rejects_deferred_tools_without_hosted_tool_search() { + let request = Request { + model: Cow::Borrowed("claude-sonnet"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![ToolSpec { + name: "lookup_order".to_string(), + description: None, + input_schema: serde_json::json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let error = AnthropicRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("deferred tools require hosted tool search")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn forced_deferred_tool_serializes_as_immediate() { + let request = Request { + model: Cow::Borrowed("claude-sonnet"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![ToolSpec { + name: "lookup_order".to_string(), + description: Some("Look up an order".to_string()), + input_schema: serde_json::json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Tool { + name: "lookup_order".to_string(), + }), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(AnthropicRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["tools"][0]["name"], "lookup_order"); + assert!(payload["tools"][0].get("defer_loading").is_none()); + assert!(payload["tools"].get(1).is_none()); + assert_eq!(payload["tool_choice"]["name"], "lookup_order"); + } +} diff --git a/vendor/mentra-provider/src/anthropic/sse.rs b/vendor/mentra-provider/src/anthropic/sse.rs new file mode 100644 index 0000000..54518d5 --- /dev/null +++ b/vendor/mentra-provider/src/anthropic/sse.rs @@ -0,0 +1,460 @@ +use std::collections::HashMap; +use std::collections::HashSet; + +use futures_util::StreamExt; +use tokio::sync::mpsc; + +use crate::{ + ProviderError, ProviderEvent, ProviderEventStream, ProviderId, ReasoningFormat, + ReasoningProvenance, TokenUsage, +}; + +use super::stream_model::AnthropicContentBlockDelta; +use super::stream_model::AnthropicStreamContentBlock; +use super::stream_model::AnthropicStreamEvent; + +pub(crate) fn spawn_event_stream( + response: reqwest::Response, + provider: ProviderId, + requested_model: String, +) -> ProviderEventStream { + let (tx, rx) = mpsc::unbounded_channel(); + + tokio::spawn(async move { + if let Err(error) = forward_events(response, tx.clone(), provider, requested_model).await { + let _ = tx.send(Err(error)); + } + }); + + rx +} + +async fn forward_events( + response: reqwest::Response, + tx: mpsc::UnboundedSender>, + provider: ProviderId, + requested_model: String, +) -> Result<(), ProviderError> { + let mut bytes_stream = response.bytes_stream(); + let mut buffer = Vec::new(); + let mut state = StreamState::new(provider, requested_model); + + while let Some(chunk) = bytes_stream.next().await { + let chunk = chunk.map_err(ProviderError::Transport)?; + buffer.extend_from_slice(&chunk); + + while let Some((frame_end, delimiter_len)) = find_frame_boundary(&buffer) { + let frame = buffer.drain(..frame_end).collect::>(); + buffer.drain(..delimiter_len); + + for event in parse_frame(&frame, &mut state)? { + if tx.send(Ok(event)).is_err() { + return Ok(()); + } + } + } + } + + if !buffer.is_empty() { + for event in parse_frame(&buffer, &mut state)? { + let _ = tx.send(Ok(event)); + } + } + + Ok(()) +} + +struct StreamState { + ignored_blocks: HashSet, + latest_usage: Option, + block_kinds: HashMap, + reasoning_provenance: ReasoningProvenance, +} + +impl StreamState { + fn new(provider: ProviderId, requested_model: String) -> Self { + Self { + ignored_blocks: HashSet::new(), + latest_usage: None, + block_kinds: HashMap::new(), + reasoning_provenance: ReasoningProvenance { + provider, + model: requested_model, + format: ReasoningFormat::AnthropicSigned, + }, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum StreamingBlockKind { + ToolUse, + HostedToolSearch, +} + +fn parse_frame(frame: &[u8], state: &mut StreamState) -> Result, ProviderError> { + let frame = std::str::from_utf8(frame) + .map_err(|error| ProviderError::MalformedStream(error.to_string()))?; + let mut data_lines = Vec::new(); + + for raw_line in frame.lines() { + let line = raw_line.strip_suffix('\r').unwrap_or(raw_line); + if line.is_empty() || line.starts_with(':') { + continue; + } + + if let Some(rest) = line.strip_prefix("data:") { + data_lines.push(rest.trim_start().to_string()); + } + } + + if data_lines.is_empty() { + return Ok(Vec::new()); + } + + let data = data_lines.join("\n"); + let event: AnthropicStreamEvent = + serde_json::from_str(&data).map_err(ProviderError::Deserialize)?; + + if let AnthropicStreamEvent::ContentBlockStart { + index, + content_block, + } = &event + { + match content_block { + AnthropicStreamContentBlock::ToolUse { .. } => { + state + .block_kinds + .insert(*index, StreamingBlockKind::ToolUse); + } + AnthropicStreamContentBlock::ServerToolUse { name, .. } + if name.starts_with("tool_search") => + { + state + .block_kinds + .insert(*index, StreamingBlockKind::HostedToolSearch); + } + _ => {} + } + } + + match &event { + AnthropicStreamEvent::ContentBlockStart { + index, + content_block, + } if !content_block.is_supported() => { + state.ignored_blocks.insert(*index); + return Ok(Vec::new()); + } + AnthropicStreamEvent::ContentBlockDelta { index, .. } + | AnthropicStreamEvent::ContentBlockStop { index } + if state.ignored_blocks.contains(index) => + { + if matches!(event, AnthropicStreamEvent::ContentBlockStop { .. }) { + state.ignored_blocks.remove(index); + } + return Ok(Vec::new()); + } + AnthropicStreamEvent::ContentBlockDelta { + index, + delta: AnthropicContentBlockDelta::InputJsonDelta { partial_json }, + } => { + if matches!( + state.block_kinds.get(index), + Some(StreamingBlockKind::HostedToolSearch) + ) { + return Ok(vec![ProviderEvent::ContentBlockDelta { + index: *index, + delta: ProviderEventDeltaExt::hosted_tool_search_delta(partial_json), + }]); + } + } + _ => {} + } + + let events = event + .into_provider_events(&state.reasoning_provenance) + .map_err(|error| { + ProviderError::MalformedStream(format!( + "anthropic stream error ({}): {}", + error.kind, error.message + )) + })?; + + Ok(events + .into_iter() + .map(|event| match event { + ProviderEvent::ContentBlockStopped { index } => { + state.block_kinds.remove(&index); + ProviderEvent::ContentBlockStopped { index } + } + ProviderEvent::MessageDelta { stop_reason, usage } => { + let usage = merge_usage(state.latest_usage.clone(), usage); + state.latest_usage = usage.clone(); + ProviderEvent::MessageDelta { stop_reason, usage } + } + other => other, + }) + .collect()) +} + +struct ProviderEventDeltaExt; + +impl ProviderEventDeltaExt { + fn hosted_tool_search_delta(partial_json: &str) -> crate::ContentBlockDelta { + crate::ContentBlockDelta::HostedToolSearchQuery( + extract_tool_search_query(partial_json).unwrap_or_else(|| partial_json.to_string()), + ) + } +} + +fn extract_tool_search_query(partial_json: &str) -> Option { + serde_json::from_str::(partial_json) + .ok() + .and_then(|value| { + value + .get("query") + .and_then(serde_json::Value::as_str) + .map(str::to_string) + }) +} + +fn merge_usage(base: Option, update: Option) -> Option { + match (base, update) { + (Some(base), Some(update)) => { + let merged = TokenUsage { + input_tokens: update.input_tokens.or(base.input_tokens), + output_tokens: update.output_tokens.or(base.output_tokens), + total_tokens: update.total_tokens.or(base.total_tokens), + cache_read_input_tokens: update + .cache_read_input_tokens + .or(base.cache_read_input_tokens), + cache_creation_input_tokens: update + .cache_creation_input_tokens + .or(base.cache_creation_input_tokens), + reasoning_tokens: update.reasoning_tokens.or(base.reasoning_tokens), + thoughts_tokens: update.thoughts_tokens.or(base.thoughts_tokens), + tool_input_tokens: update.tool_input_tokens.or(base.tool_input_tokens), + }; + Some(merged) + } + (Some(base), None) => Some(base), + (None, Some(update)) => Some(update), + (None, None) => None, + } +} + +fn find_frame_boundary(buffer: &[u8]) -> Option<(usize, usize)> { + for (index, window) in buffer.windows(2).enumerate() { + if window == b"\n\n" { + return Some((index, 2)); + } + } + + for (index, window) in buffer.windows(4).enumerate() { + if window == b"\r\n\r\n" { + return Some((index, 4)); + } + } + + None +} + +#[cfg(test)] +mod tests { + use super::{StreamState, parse_frame}; + use crate::{ProviderEvent, ProviderId, Role, TokenUsage}; + + fn stream_state() -> StreamState { + StreamState::new( + ProviderId::new("anthropic-edge"), + "claude-requested".to_string(), + ) + } + + #[test] + fn merges_anthropic_usage_updates_into_cumulative_totals() { + let mut state = stream_state(); + + let started = parse_frame( + br#"data: {"type":"message_start","message":{"id":"msg_1","model":"claude-sonnet","role":"assistant","content":[],"usage":{"input_tokens":10,"cache_read_input_tokens":2}}}"#, + &mut state, + ) + .expect("message start should parse"); + assert_eq!( + started, + vec![ + ProviderEvent::MessageStarted { + id: "msg_1".to_string(), + model: "claude-sonnet".to_string(), + role: Role::Assistant, + }, + ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(TokenUsage { + input_tokens: Some(10), + output_tokens: None, + total_tokens: None, + cache_read_input_tokens: Some(2), + cache_creation_input_tokens: None, + reasoning_tokens: None, + thoughts_tokens: None, + tool_input_tokens: None, + }), + }, + ] + ); + + let delta = parse_frame( + br#"data: {"type":"message_delta","delta":{"stop_reason":"end_turn","usage":{"output_tokens":3}}}"#, + &mut state, + ) + .expect("message delta should parse"); + assert_eq!( + delta, + vec![ProviderEvent::MessageDelta { + stop_reason: Some("end_turn".to_string()), + usage: Some(TokenUsage { + input_tokens: Some(10), + output_tokens: Some(3), + total_tokens: None, + cache_read_input_tokens: Some(2), + cache_creation_input_tokens: None, + reasoning_tokens: None, + thoughts_tokens: None, + tool_input_tokens: None, + }), + }] + ); + } + + #[test] + fn parses_hosted_tool_search_bookkeeping_blocks() { + let mut state = stream_state(); + + let started = parse_frame( + br#"data: {"type":"content_block_start","index":1,"content_block":{"type":"server_tool_use","id":"srvtoolu_1","name":"tool_search_tool_bm25"}}"#, + &mut state, + ) + .expect("server tool use should parse"); + assert_eq!( + started, + vec![ProviderEvent::ContentBlockStarted { + index: 1, + kind: crate::ContentBlockStart::HostedToolSearch { + call: crate::HostedToolSearchCall { + id: "srvtoolu_1".to_string(), + status: Some("in_progress".to_string()), + query: None, + }, + }, + }] + ); + + let delta = parse_frame( + br#"data: {"type":"content_block_delta","index":1,"delta":{"type":"input_json_delta","partial_json":"{\"query\":\"weather\"}"}}"#, + &mut state, + ) + .expect("hosted search delta should parse"); + assert_eq!( + delta, + vec![ProviderEvent::ContentBlockDelta { + index: 1, + delta: crate::ContentBlockDelta::HostedToolSearchQuery("weather".to_string()), + }] + ); + + let stopped = parse_frame( + br#"data: {"type":"content_block_stop","index":1}"#, + &mut state, + ) + .expect("hosted search stop should parse"); + assert_eq!( + stopped, + vec![ProviderEvent::ContentBlockStopped { index: 1 }] + ); + } + + #[test] + fn parses_signed_and_redacted_thinking_with_requested_provenance() { + let mut state = stream_state(); + + let started = parse_frame( + br#"data: {"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":""}}"#, + &mut state, + ) + .expect("thinking start should parse"); + assert_eq!( + started, + vec![ProviderEvent::ContentBlockStarted { + index: 0, + kind: crate::ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: Some(crate::ReasoningProvenance { + provider: ProviderId::new("anthropic-edge"), + model: "claude-requested".to_string(), + format: crate::ReasoningFormat::AnthropicSigned, + }), + redacted: false, + }, + }] + ); + + let thinking = parse_frame( + br#"data: {"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"private chain"}}"#, + &mut state, + ) + .expect("thinking delta should parse"); + assert_eq!( + thinking, + vec![ProviderEvent::ContentBlockDelta { + index: 0, + delta: crate::ContentBlockDelta::ThinkingText("private chain".to_string()), + }] + ); + + let signature = parse_frame( + br#"data: {"type":"content_block_delta","index":0,"delta":{"type":"signature_delta","signature":"opaque-signature"}}"#, + &mut state, + ) + .expect("signature delta should parse"); + assert_eq!( + signature, + vec![ProviderEvent::ContentBlockDelta { + index: 0, + delta: crate::ContentBlockDelta::ThinkingSignature("opaque-signature".to_string()), + }] + ); + + let redacted = parse_frame( + br#"data: {"type":"content_block_start","index":1,"content_block":{"type":"redacted_thinking","data":"opaque-redacted-data"}}"#, + &mut state, + ) + .expect("redacted thinking should parse"); + assert_eq!( + redacted, + vec![ + ProviderEvent::ContentBlockStarted { + index: 1, + kind: crate::ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: Some(crate::ReasoningProvenance { + provider: ProviderId::new("anthropic-edge"), + model: "claude-requested".to_string(), + format: crate::ReasoningFormat::AnthropicSigned, + }), + redacted: true, + }, + }, + ProviderEvent::ContentBlockDelta { + index: 1, + delta: crate::ContentBlockDelta::ThinkingSignature( + "opaque-redacted-data".to_string() + ), + }, + ] + ); + } +} diff --git a/vendor/mentra-provider/src/anthropic/stream_model.rs b/vendor/mentra-provider/src/anthropic/stream_model.rs new file mode 100644 index 0000000..c2268ba --- /dev/null +++ b/vendor/mentra-provider/src/anthropic/stream_model.rs @@ -0,0 +1,232 @@ +use serde::Deserialize; + +use crate::{ + ContentBlockDelta, ContentBlockStart, HostedToolSearchCall, ProviderEvent, ReasoningProvenance, + Role, +}; + +use super::model::{AnthropicResponse, AnthropicUsage}; + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum AnthropicStreamEvent { + MessageStart { + message: AnthropicResponse, + }, + ContentBlockStart { + index: usize, + content_block: AnthropicStreamContentBlock, + }, + ContentBlockDelta { + index: usize, + delta: AnthropicContentBlockDelta, + }, + ContentBlockStop { + index: usize, + }, + MessageDelta { + delta: AnthropicMessageDelta, + }, + MessageStop, + Ping, + Error { + error: AnthropicStreamError, + }, +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum AnthropicStreamContentBlock { + Text {}, + Thinking { + #[serde(default)] + thinking: String, + }, + RedactedThinking { + data: String, + }, + ToolUse { + id: String, + name: String, + }, + ServerToolUse { + id: String, + name: String, + }, + #[serde(other)] + Unsupported, +} + +impl AnthropicStreamContentBlock { + pub(crate) fn into_provider_events( + self, + index: usize, + provenance: &ReasoningProvenance, + ) -> Vec { + match self { + AnthropicStreamContentBlock::Text {} => vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Text, + }], + AnthropicStreamContentBlock::Thinking { thinking } => { + let mut events = vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: Some(provenance.clone()), + redacted: false, + }, + }]; + if !thinking.is_empty() { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ThinkingText(thinking), + }); + } + events + } + AnthropicStreamContentBlock::RedactedThinking { data } => vec![ + ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: None, + provenance: Some(provenance.clone()), + redacted: true, + }, + }, + ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ThinkingSignature(data), + }, + ], + AnthropicStreamContentBlock::ToolUse { id, name } => { + vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolUse { id, name }, + }] + } + AnthropicStreamContentBlock::ServerToolUse { id, name } => name + .starts_with("tool_search") + .then(|| ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::HostedToolSearch { + call: HostedToolSearchCall { + id, + status: Some("in_progress".to_string()), + query: None, + }, + }, + }) + .into_iter() + .collect(), + AnthropicStreamContentBlock::Unsupported => Vec::new(), + } + } + + pub(crate) fn is_supported(&self) -> bool { + !matches!(self, AnthropicStreamContentBlock::Unsupported) + } +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum AnthropicContentBlockDelta { + TextDelta { + text: String, + }, + InputJsonDelta { + partial_json: String, + }, + ThinkingDelta { + thinking: String, + }, + SignatureDelta { + signature: String, + }, + #[serde(other)] + Unsupported, +} + +impl AnthropicContentBlockDelta { + pub(crate) fn into_provider_delta(self) -> Option { + match self { + AnthropicContentBlockDelta::TextDelta { text } => Some(ContentBlockDelta::Text(text)), + AnthropicContentBlockDelta::InputJsonDelta { partial_json } => { + Some(ContentBlockDelta::ToolUseInputJson(partial_json)) + } + AnthropicContentBlockDelta::ThinkingDelta { thinking } => { + Some(ContentBlockDelta::ThinkingText(thinking)) + } + AnthropicContentBlockDelta::SignatureDelta { signature } => { + Some(ContentBlockDelta::ThinkingSignature(signature)) + } + AnthropicContentBlockDelta::Unsupported => None, + } + } +} + +#[derive(Deserialize)] +pub(crate) struct AnthropicMessageDelta { + pub(crate) stop_reason: Option, + #[serde(default)] + pub(crate) usage: Option, +} + +#[derive(Deserialize)] +pub(crate) struct AnthropicStreamError { + #[serde(rename = "type")] + pub(crate) kind: String, + pub(crate) message: String, +} + +impl AnthropicStreamEvent { + pub(crate) fn into_provider_events( + self, + reasoning_provenance: &ReasoningProvenance, + ) -> Result, AnthropicStreamError> { + match self { + AnthropicStreamEvent::MessageStart { message } => { + let usage = message + .usage + .clone() + .and_then(AnthropicUsage::into_token_usage); + let mut events = vec![ProviderEvent::MessageStarted { + id: message.id, + model: message.model, + role: match message.role.as_str() { + "user" => Role::User, + "assistant" => Role::Assistant, + _ => Role::Unknown(message.role), + }, + }]; + if let Some(usage) = usage { + events.push(ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(usage), + }); + } + Ok(events) + } + AnthropicStreamEvent::ContentBlockStart { + index, + content_block, + } => Ok(content_block.into_provider_events(index, reasoning_provenance)), + AnthropicStreamEvent::ContentBlockDelta { index, delta } => Ok(delta + .into_provider_delta() + .map(|delta| vec![ProviderEvent::ContentBlockDelta { index, delta }]) + .unwrap_or_default()), + AnthropicStreamEvent::ContentBlockStop { index } => { + Ok(vec![ProviderEvent::ContentBlockStopped { index }]) + } + AnthropicStreamEvent::MessageDelta { delta } => Ok(vec![ProviderEvent::MessageDelta { + stop_reason: delta.stop_reason, + usage: delta.usage.and_then(AnthropicUsage::into_token_usage), + }]), + AnthropicStreamEvent::MessageStop => Ok(vec![ProviderEvent::MessageStopped]), + AnthropicStreamEvent::Ping => Ok(Vec::new()), + AnthropicStreamEvent::Error { error } => Err(error), + } + } +} diff --git a/vendor/mentra-provider/src/auth.rs b/vendor/mentra-provider/src/auth.rs new file mode 100644 index 0000000..afd105f --- /dev/null +++ b/vendor/mentra-provider/src/auth.rs @@ -0,0 +1,72 @@ +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::sync::Arc; + +use crate::error::ProviderError; + +/// What kind of auth material a provider expects. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum AuthScheme { + #[default] + None, + BearerToken, + Header { + name: String, + }, + QueryParam { + name: String, + }, +} + +/// Provider credentials resolved at runtime. +/// +/// `bearer_token` is the provider's primary auth secret and is applied according +/// to the provider definition's [`AuthScheme`]. For example, Responses-family +/// providers send it as `Authorization: Bearer ...`, while header/query auth +/// providers can reuse the same resolved secret in a different location. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct ProviderCredentials { + pub bearer_token: Option, + pub account_id: Option, + #[serde(default)] + pub headers: HashMap, +} + +/// Supplies credentials on demand for a provider instance. +#[async_trait] +pub trait CredentialSource: Send + Sync { + async fn credentials(&self) -> Result; + + async fn bearer_token(&self) -> Result { + self.credentials() + .await? + .bearer_token + .ok_or_else(|| ProviderError::InvalidRequest("missing bearer token".to_string())) + } +} + +/// Supplies a fixed auth secret to the provider. +#[derive(Clone)] +pub struct StaticCredentialSource { + secret: Arc, +} + +impl StaticCredentialSource { + pub fn new(secret: impl Into) -> Self { + Self { + secret: Arc::from(secret.into()), + } + } +} + +#[async_trait] +impl CredentialSource for StaticCredentialSource { + async fn credentials(&self) -> Result { + Ok(ProviderCredentials { + bearer_token: Some(self.secret.to_string()), + account_id: None, + headers: HashMap::new(), + }) + } +} diff --git a/vendor/mentra-provider/src/definition.rs b/vendor/mentra-provider/src/definition.rs new file mode 100644 index 0000000..6e7a7c4 --- /dev/null +++ b/vendor/mentra-provider/src/definition.rs @@ -0,0 +1,483 @@ +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use http::header; +use serde::Deserialize; +use serde::Serialize; +use std::borrow::Cow; +use std::collections::HashMap; +use std::fmt::Display; +use std::time::Duration; +use strum::Display as StrumDisplay; +use strum::IntoStaticStr; +use url::Url; + +use crate::request::SessionRequestOptions; + +/// Builtin provider families Mentra can construct from presets. +#[derive(Debug, Clone, Copy, PartialEq, Eq, StrumDisplay, IntoStaticStr)] +#[strum(serialize_all = "lowercase")] +pub enum BuiltinProvider { + Anthropic, + Gemini, + OpenAI, + OpenRouter, + Ollama, + LmStudio, +} + +impl From for ProviderId { + fn from(value: BuiltinProvider) -> Self { + Self(Cow::Borrowed(value.into())) + } +} + +/// Stable identifier for a registered provider implementation. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, PartialOrd, Ord)] +pub struct ProviderId(Cow<'static, str>); + +impl ProviderId { + pub fn new(id: impl Into) -> Self { + Self(Cow::Owned(id.into())) + } + + pub fn as_str(&self) -> &str { + self.0.as_ref() + } +} + +impl Display for ProviderId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(self.as_str()) + } +} + +impl From<&str> for ProviderId { + fn from(value: &str) -> Self { + Self::new(value) + } +} + +impl From for ProviderId { + fn from(value: String) -> Self { + Self(Cow::Owned(value)) + } +} + +impl From<&String> for ProviderId { + fn from(value: &String) -> Self { + Self::new(value.as_str()) + } +} + +/// Human-facing metadata about a provider. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ProviderDescriptor { + pub id: ProviderId, + pub display_name: Option, + pub description: Option, +} + +impl ProviderDescriptor { + pub fn new(id: impl Into) -> Self { + Self { + id: id.into(), + display_name: None, + description: None, + } + } +} + +/// Capabilities advertised by a provider instance. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct ProviderCapabilities { + pub supports_model_listing: bool, + pub supports_streaming: bool, + pub supports_websockets: bool, + pub supports_tool_calls: bool, + pub supports_images: bool, + pub supports_history_compaction: bool, + pub supports_memory_summarization: bool, + pub supports_deferred_tools: bool, + pub supports_hosted_tool_search: bool, + pub supports_hosted_web_search: bool, + pub supports_image_generation: bool, + pub supports_reasoning_effort: bool, + pub reports_reasoning_tokens: bool, + pub reports_thoughts_tokens: bool, + pub supports_structured_tool_results: bool, + pub supports_embeddings: bool, +} + +/// Wire protocol supported by a provider. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum WireApi { + #[default] + Responses, + AnthropicMessages, + GeminiGenerateContent, +} + +impl Display for WireApi { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let value = match self { + Self::Responses => "responses", + Self::AnthropicMessages => "anthropic_messages", + Self::GeminiGenerateContent => "gemini_generate_content", + }; + f.write_str(value) + } +} + +/// Retry configuration for provider transport calls. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct RetryPolicy { + pub max_attempts: u64, + pub base_delay: Duration, + pub retry_429: bool, + pub retry_5xx: bool, + pub retry_transport: bool, +} + +impl Default for RetryPolicy { + fn default() -> Self { + Self { + max_attempts: 5, + base_delay: Duration::from_millis(200), + retry_429: false, + retry_5xx: true, + retry_transport: true, + } + } +} + +/// Serializable provider definition used by runtime and adapter layers. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ProviderDefinition { + pub descriptor: ProviderDescriptor, + #[serde(default)] + pub wire_api: WireApi, + #[serde(default)] + pub auth_scheme: crate::AuthScheme, + #[serde(default)] + pub capabilities: ProviderCapabilities, + pub base_url: Option, + #[serde(default)] + pub query_params: Option>, + #[serde(default)] + pub headers: Option>, + #[serde(default)] + pub retry: RetryPolicy, + #[serde(default = "default_stream_idle_timeout")] + pub stream_idle_timeout: Duration, + #[serde(default = "default_websocket_connect_timeout")] + pub websocket_connect_timeout: Duration, +} + +fn default_stream_idle_timeout() -> Duration { + Duration::from_millis(300_000) +} + +fn default_websocket_connect_timeout() -> Duration { + Duration::from_millis(15_000) +} + +impl ProviderDefinition { + pub fn new(id: impl Into) -> Self { + Self { + descriptor: ProviderDescriptor::new(id), + wire_api: WireApi::default(), + auth_scheme: crate::AuthScheme::default(), + capabilities: ProviderCapabilities { + supports_model_listing: true, + supports_streaming: true, + supports_websockets: false, + supports_tool_calls: true, + supports_images: true, + supports_history_compaction: false, + supports_memory_summarization: false, + supports_deferred_tools: false, + supports_hosted_tool_search: false, + supports_hosted_web_search: false, + supports_image_generation: false, + supports_reasoning_effort: false, + reports_reasoning_tokens: false, + reports_thoughts_tokens: false, + supports_structured_tool_results: false, + supports_embeddings: false, + }, + base_url: None, + query_params: None, + headers: None, + retry: RetryPolicy::default(), + stream_idle_timeout: default_stream_idle_timeout(), + websocket_connect_timeout: default_websocket_connect_timeout(), + } + } + + pub fn descriptor(&self) -> ProviderDescriptor { + self.descriptor.clone() + } + + pub fn provider_id(&self) -> &ProviderId { + &self.descriptor.id + } + + pub fn url_for_path(&self, path: &str) -> String { + let base = self + .base_url + .as_deref() + .unwrap_or_default() + .trim_end_matches('/'); + let path = path.trim_start_matches('/'); + let mut url = if path.is_empty() { + base.to_string() + } else { + format!("{base}/{path}") + }; + + if let Some(params) = self + .query_params + .as_ref() + .filter(|params| !params.is_empty()) + { + let qs = params + .iter() + .map(|(key, value)| format!("{key}={value}")) + .collect::>() + .join("&"); + url.push('?'); + url.push_str(&qs); + } + + url + } + + pub fn build_headers( + &self, + credentials: &crate::ProviderCredentials, + ) -> Result { + let mut headers = HeaderMap::new(); + + if let Some(configured_headers) = &self.headers { + for (name, value) in configured_headers { + insert_header(&mut headers, name, value)?; + } + } + + for (name, value) in &credentials.headers { + insert_header(&mut headers, name, value)?; + } + + match &self.auth_scheme { + crate::AuthScheme::None | crate::AuthScheme::QueryParam { .. } => {} + crate::AuthScheme::BearerToken => { + let token = required_auth_value(credentials)?; + let auth_value = + HeaderValue::from_str(&format!("Bearer {token}")).map_err(|error| { + crate::ProviderError::InvalidRequest(format!( + "invalid bearer token header: {error}" + )) + })?; + headers.insert(header::AUTHORIZATION, auth_value); + } + crate::AuthScheme::Header { name } => { + let token = required_auth_value(credentials)?; + insert_header(&mut headers, name, token)?; + } + } + + Ok(headers) + } + + pub fn build_headers_for_session( + &self, + credentials: &crate::ProviderCredentials, + session: Option<&SessionRequestOptions>, + fallback_turn_state: Option<&str>, + ) -> Result { + let mut headers = self.build_headers(credentials)?; + + if let Some(value) = session + .and_then(|session| session.sticky_turn_state.as_deref()) + .or(fallback_turn_state) + .and_then(|turn_state| HeaderValue::from_str(turn_state).ok()) + { + headers.insert("x-mentra-turn-state", value.clone()); + headers.insert("x-codex-turn-state", value); + } + if let Some(value) = session + .and_then(|session| session.turn_metadata.as_deref()) + .and_then(|value| HeaderValue::from_str(value).ok()) + { + headers.insert("x-mentra-turn-metadata", value.clone()); + headers.insert("x-codex-turn-metadata", value); + } + if let Some(value) = session + .and_then(|session| session.session_affinity.as_deref()) + .and_then(|value| HeaderValue::from_str(value).ok()) + { + headers.insert("x-mentra-session-affinity", value); + } + if let Some(prefer_connection_reuse) = + session.and_then(|session| session.prefer_connection_reuse) + { + headers.insert( + "x-mentra-connection-reuse", + HeaderValue::from_static(if prefer_connection_reuse { + "prefer-reuse" + } else { + "prefer-fresh" + }), + ); + } + if let Some(value) = session + .and_then(|session| session.subagent.as_deref()) + .and_then(|value| HeaderValue::from_str(value).ok()) + { + headers.insert("x-openai-subagent", value); + } + if let Some(extra_headers) = session.map(|session| &session.extra_headers) { + for (name, value) in extra_headers { + if let (Ok(name), Ok(value)) = ( + name.parse::(), + HeaderValue::from_str(value), + ) { + headers.insert(name, value); + } + } + } + + Ok(headers) + } + + pub fn request_url_with_auth_for_path( + &self, + path: &str, + credentials: &crate::ProviderCredentials, + ) -> Result { + let mut url = Url::parse(&self.url_for_path(path)) + .map_err(|error| crate::ProviderError::InvalidRequest(error.to_string()))?; + + if let crate::AuthScheme::QueryParam { name } = &self.auth_scheme { + let token = required_auth_value(credentials)?; + url.query_pairs_mut().append_pair(name, token); + } + + Ok(url) + } + + pub fn websocket_url_for_path(&self, path: &str) -> Result { + let mut url = Url::parse(&self.url_for_path(path))?; + + let scheme = match url.scheme() { + "http" => "ws", + "https" => "wss", + "ws" | "wss" => return Ok(url), + _ => return Ok(url), + }; + let _ = url.set_scheme(scheme); + Ok(url) + } + + pub fn websocket_url_with_auth_for_path( + &self, + path: &str, + credentials: &crate::ProviderCredentials, + ) -> Result { + let mut url = self + .websocket_url_for_path(path) + .map_err(|error| crate::ProviderError::InvalidRequest(error.to_string()))?; + + if let crate::AuthScheme::QueryParam { name } = &self.auth_scheme { + let token = required_auth_value(credentials)?; + url.query_pairs_mut().append_pair(name, token); + } + + Ok(url) + } +} + +fn insert_header( + headers: &mut HeaderMap, + name: &str, + value: &str, +) -> Result<(), crate::ProviderError> { + let header_name = HeaderName::from_bytes(name.as_bytes()).map_err(|error| { + crate::ProviderError::InvalidRequest(format!( + "invalid provider header name {name:?}: {error}" + )) + })?; + let header_value = HeaderValue::from_str(value).map_err(|error| { + crate::ProviderError::InvalidRequest(format!( + "invalid provider header value for {name:?}: {error}" + )) + })?; + headers.insert(header_name, header_value); + Ok(()) +} + +fn required_auth_value( + credentials: &crate::ProviderCredentials, +) -> Result<&str, crate::ProviderError> { + credentials.bearer_token.as_deref().ok_or_else(|| { + crate::ProviderError::InvalidRequest("missing provider auth credential".to_string()) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn build_headers_applies_bearer_auth_and_static_headers() { + let mut definition = ProviderDefinition::new("test"); + definition.auth_scheme = crate::AuthScheme::BearerToken; + definition.headers = Some(HashMap::from([( + "x-provider-header".to_string(), + "static".to_string(), + )])); + + let headers = definition + .build_headers(&crate::ProviderCredentials { + bearer_token: Some("secret".to_string()), + account_id: None, + headers: HashMap::from([("x-runtime-header".to_string(), "dynamic".to_string())]), + }) + .expect("headers should build"); + + assert_eq!(headers["x-provider-header"], "static"); + assert_eq!(headers["x-runtime-header"], "dynamic"); + assert_eq!(headers[header::AUTHORIZATION], "Bearer secret"); + } + + #[test] + fn request_url_with_auth_appends_query_param_auth() { + let mut definition = ProviderDefinition::new("test"); + definition.base_url = Some("https://example.com/v1".to_string()); + definition.query_params = Some(HashMap::from([( + "api-version".to_string(), + "2026".to_string(), + )])); + definition.auth_scheme = crate::AuthScheme::QueryParam { + name: "api-key".to_string(), + }; + + let url = definition + .request_url_with_auth_for_path( + "responses", + &crate::ProviderCredentials { + bearer_token: Some("secret".to_string()), + account_id: None, + headers: HashMap::new(), + }, + ) + .expect("url should build"); + + assert_eq!( + url.as_str(), + "https://example.com/v1/responses?api-version=2026&api-key=secret" + ); + } +} diff --git a/vendor/mentra-provider/src/embedding.rs b/vendor/mentra-provider/src/embedding.rs new file mode 100644 index 0000000..658d384 --- /dev/null +++ b/vendor/mentra-provider/src/embedding.rs @@ -0,0 +1,175 @@ +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; + +use crate::ProviderError; + +/// Static metadata about an embedding model. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct EmbeddingModelInfo { + pub id: String, + pub dimensions: usize, + pub max_tokens: usize, +} + +impl EmbeddingModelInfo { + pub fn new(id: impl Into, dimensions: usize, max_tokens: usize) -> Self { + Self { + id: id.into(), + dimensions, + max_tokens, + } + } +} + +/// Input variants for an embedding request. +/// +/// Only `Serialize` is derived — these types are sent in requests but never +/// deserialized from responses, so `Deserialize` is not needed. +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(untagged)] +pub enum EmbeddingInput<'a> { + Single(&'a str), + Batch(&'a [&'a str]), +} + +/// Request payload sent to the embeddings endpoint. +#[derive(Debug, Clone, Serialize)] +pub struct EmbeddingRequest<'a> { + pub model: &'a str, + pub input: EmbeddingInput<'a>, +} + +impl<'a> EmbeddingRequest<'a> { + pub fn single(model: &'a str, text: &'a str) -> Self { + Self { + model, + input: EmbeddingInput::Single(text), + } + } + + pub fn batch(model: &'a str, texts: &'a [&'a str]) -> Self { + Self { + model, + input: EmbeddingInput::Batch(texts), + } + } +} + +/// A single embedding vector with its position in the batch. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct EmbeddingData { + pub index: usize, + pub embedding: Vec, +} + +/// Token usage reported by the embeddings endpoint. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct EmbeddingUsage { + pub prompt_tokens: u32, + pub total_tokens: u32, +} + +/// Response returned from the embeddings endpoint. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct EmbeddingResponse { + pub data: Vec, + pub model: String, + pub usage: EmbeddingUsage, +} + +/// Trait implemented by providers that support vector embeddings. +/// +/// `EmbeddingProvider` is intentionally separate from `Provider` because not all +/// LLM providers expose an embeddings endpoint (e.g. Anthropic and Gemini do not). +#[async_trait] +pub trait EmbeddingProvider: Send + Sync { + /// Embed a single piece of text. + async fn embed(&self, model: &str, text: &str) -> Result, ProviderError> { + let texts = [text]; + let mut response = self.embed_batch(model, &texts).await?; + response + .data + .pop() + .map(|d| d.embedding) + .ok_or_else(|| ProviderError::InvalidResponse("empty embedding response".to_string())) + } + + /// Embed a batch of texts, returning one vector per input in order. + async fn embed_batch( + &self, + model: &str, + texts: &[&str], + ) -> Result; + + /// Returns metadata about the embedding models available from this provider. + fn embedding_models(&self) -> Vec; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn embedding_request_single_serializes_as_string() { + let req = EmbeddingRequest::single("text-embedding-3-small", "hello world"); + let json = serde_json::to_value(&req).unwrap(); + + assert_eq!(json["model"], "text-embedding-3-small"); + assert_eq!(json["input"], "hello world"); + } + + #[test] + fn embedding_request_batch_serializes_as_array() { + let texts = ["hello", "world"]; + let req = EmbeddingRequest::batch("text-embedding-3-small", &texts); + let json = serde_json::to_value(&req).unwrap(); + + assert_eq!(json["model"], "text-embedding-3-small"); + assert_eq!(json["input"], serde_json::json!(["hello", "world"])); + } + + #[test] + fn embedding_response_deserializes_correctly() { + let raw = serde_json::json!({ + "data": [ + { "index": 0, "embedding": [0.1, 0.2, 0.3] }, + { "index": 1, "embedding": [0.4, 0.5, 0.6] } + ], + "model": "text-embedding-3-small", + "usage": { "prompt_tokens": 5, "total_tokens": 5 } + }); + + let response: EmbeddingResponse = serde_json::from_value(raw).unwrap(); + + assert_eq!(response.model, "text-embedding-3-small"); + assert_eq!(response.data.len(), 2); + assert_eq!(response.data[0].index, 0); + assert_eq!(response.data[0].embedding, vec![0.1f32, 0.2, 0.3]); + assert_eq!(response.data[1].index, 1); + assert_eq!(response.usage.prompt_tokens, 5); + assert_eq!(response.usage.total_tokens, 5); + } + + #[test] + fn embedding_model_info_stores_fields() { + let info = EmbeddingModelInfo::new("text-embedding-3-large", 3072, 8191); + assert_eq!(info.id, "text-embedding-3-large"); + assert_eq!(info.dimensions, 3072); + assert_eq!(info.max_tokens, 8191); + } + + #[test] + fn embedding_input_single_serializes_as_bare_string() { + let input = EmbeddingInput::Single("test text"); + let serialized = serde_json::to_string(&input).unwrap(); + assert_eq!(serialized, r#""test text""#); + } + + #[test] + fn embedding_input_batch_serializes_as_array() { + let texts = ["a", "b"]; + let input = EmbeddingInput::Batch(&texts); + let json = serde_json::to_value(&input).unwrap(); + assert_eq!(json, serde_json::json!(["a", "b"])); + } +} diff --git a/vendor/mentra-provider/src/error.rs b/vendor/mentra-provider/src/error.rs new file mode 100644 index 0000000..e94721e --- /dev/null +++ b/vendor/mentra-provider/src/error.rs @@ -0,0 +1,271 @@ +use std::time::Duration; + +use thiserror::Error; +use time::{OffsetDateTime, PrimitiveDateTime}; + +/// Errors returned by provider implementations and stream adapters. +#[derive(Debug, Error)] +pub enum ProviderError { + #[error("provider transport error: {0}")] + Transport(#[source] reqwest::Error), + #[error("retryable provider error: {message}")] + Retryable { + message: String, + delay: Option, + }, + #[error("provider does not support capability: {0}")] + UnsupportedCapability(String), + #[error("{message}", message = provider_http_error(.status, .body))] + Http { + status: reqwest::StatusCode, + body: String, + /// How long the server asked the caller to wait, from the response's + /// `Retry-After` header, or `None` when it sent none. + /// + /// A rate limit is the one failure whose recovery time the server + /// knows and the client cannot guess: an exponential backoff shaped + /// for a connection blip retries five times inside the window and + /// then gives up while the limit is still in force. Capturing the + /// header here, where the response is turned into an error, is what + /// lets a caller wait the interval the server named instead — + /// nothing further up the stack ever sees the headers. + /// + /// Read it back through [`retry_after`](ProviderError::retry_after), + /// which answers the same question for + /// [`Retryable`](ProviderError::Retryable) too. Build one of these + /// from a live response with + /// [`from_http_response`](ProviderError::from_http_response) rather + /// than filling the field in by hand. + retry_after: Option, + }, + #[error("failed to decode provider response: {0}")] + Decode(#[source] reqwest::Error), + #[error("failed to serialize provider request: {0}")] + Serialize(#[source] serde_json::Error), + #[error("failed to deserialize provider payload: {0}")] + Deserialize(#[source] serde_json::Error), + #[error("invalid provider request: {0}")] + InvalidRequest(String), + #[error("invalid provider response: {0}")] + InvalidResponse(String), + #[error("malformed provider stream: {0}")] + MalformedStream(String), +} + +impl ProviderError { + /// Turns an unsuccessful HTTP response into an [`Http`](ProviderError::Http) + /// error, reading `Retry-After` before the body is consumed. + /// + /// This is the constructor every provider in this crate uses, and the one + /// a custom provider should use: the status and the retry hint both come + /// off the response, so neither can be forgotten at a call site. + pub async fn from_http_response(response: reqwest::Response) -> Self { + let status = response.status(); + let retry_after = retry_after_from_headers(response.headers()); + Self::Http { + status, + body: response.text().await.unwrap_or_default(), + retry_after, + } + } + + /// How long the provider asked the caller to wait before trying again, or + /// `None` when it asked for nothing. + /// + /// This is a request from the server, not a promise about the schedule: a + /// caller decides what to do with it, including refusing an interval long + /// enough to be an outage rather than a rate limit. + pub fn retry_after(&self) -> Option { + match self { + Self::Http { retry_after, .. } => *retry_after, + Self::Retryable { delay, .. } => *delay, + _ => None, + } + } +} + +fn provider_http_error(status: &reqwest::StatusCode, body: &str) -> String { + if body.trim().is_empty() { + format!("provider returned HTTP {status}") + } else { + format!("provider returned HTTP {status}: {body}") + } +} + +/// Reads `Retry-After` off a response's headers, in whichever of its two forms +/// the server chose. +pub(crate) fn retry_after_from_headers(headers: &reqwest::header::HeaderMap) -> Option { + retry_after_from_header_value(headers.get(reqwest::header::RETRY_AFTER)?.to_str().ok()?) +} + +/// Reads one already-extracted `Retry-After` value, for transports that carry +/// their headers as something other than a [`HeaderMap`](reqwest::header::HeaderMap). +pub(crate) fn retry_after_from_header_value(value: &str) -> Option { + parse_retry_after(value, OffsetDateTime::now_utc()) +} + +/// Parses a `Retry-After` value against a known `now`. +/// +/// RFC 9110 allows two spellings — a count of seconds, and an HTTP-date — and +/// providers use both, so a parser that understands only one silently ignores +/// half the rate limits it is meant to honor. `now` is a parameter rather than +/// read from the clock so the date form can be tested without waiting. +/// +/// A date already in the past yields [`Duration::ZERO`] ("retry now") rather +/// than nothing: the server did answer, and the answer was that the wait is +/// over. Only the IMF-fixdate spelling that RFC 9110 requires senders to emit +/// is understood; the two obsolete formats it only requires recipients to +/// tolerate are read as no hint at all, which costs nothing but the schedule's +/// own delay. +fn parse_retry_after(value: &str, now: OffsetDateTime) -> Option { + let value = value.trim(); + if value.is_empty() { + return None; + } + + if let Ok(seconds) = value.parse::() { + return Some(Duration::from_secs(seconds)); + } + + let deadline = parse_http_date(value)?; + Some((deadline - now).try_into().unwrap_or(Duration::ZERO)) +} + +/// Parses an IMF-fixdate, the `Sun, 06 Nov 1994 08:49:37 GMT` spelling. +fn parse_http_date(value: &str) -> Option { + // `parse_borrowed` rather than `parse`: the latter is deprecated from + // time 0.3.55 and a downstream resolving a newer `time` would see the + // warning in this crate. Version 2 of the description syntax, which the + // spelling below is already written in. + let format = time::format_description::parse_borrowed::<2>( + "[weekday repr:short], [day] [month repr:short] [year] [hour]:[minute]:[second] GMT", + ) + .ok()?; + PrimitiveDateTime::parse(value, format.as_slice()) + .ok() + .map(PrimitiveDateTime::assume_utc) +} + +#[cfg(test)] +mod tests { + use super::*; + + /// The instant in RFC 9110's own `Retry-After` example, built without the + /// `time` macros feature so the crate's dependency set stays as it is. + fn now() -> OffsetDateTime { + time::Date::from_calendar_date(1994, time::Month::November, 6) + .expect("a real date") + .with_hms(8, 49, 37) + .expect("a real time") + .assume_utc() + } + + #[test] + fn a_delay_in_seconds_is_read_as_seconds() { + assert_eq!( + parse_retry_after("30", now()), + Some(Duration::from_secs(30)) + ); + assert_eq!( + parse_retry_after(" 30 ", now()), + Some(Duration::from_secs(30)), + "surrounding whitespace is not part of the value" + ); + } + + #[test] + fn an_http_date_is_read_as_the_wait_until_it() { + // The other spelling RFC 9110 allows. A parser that understood only + // seconds would return None here and fall back to its own schedule, + // which is exactly the rate limit it was supposed to wait out. + assert_eq!( + parse_retry_after("Sun, 06 Nov 1994 08:50:37 GMT", now()), + Some(Duration::from_secs(60)) + ); + } + + #[test] + fn an_http_date_that_has_passed_means_retry_now() { + assert_eq!( + parse_retry_after("Sun, 06 Nov 1994 08:49:00 GMT", now()), + Some(Duration::ZERO), + "the server answered; the answer is that the wait is over" + ); + } + + #[test] + fn an_unparseable_value_is_no_hint_at_all() { + assert_eq!(parse_retry_after("", now()), None); + assert_eq!(parse_retry_after("soon", now()), None); + assert_eq!(parse_retry_after("-5", now()), None); + } + + #[test] + fn a_rate_limited_response_reports_the_interval_it_named() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert(reqwest::header::RETRY_AFTER, "42".parse().expect("valid")); + + assert_eq!( + retry_after_from_headers(&headers), + Some(Duration::from_secs(42)) + ); + } + + #[tokio::test] + async fn a_429_response_becomes_an_error_that_still_knows_the_interval() { + // The whole point of reading the header at construction time: nothing + // above this call ever sees the response, so a hint not captured here + // is a hint lost. + let response = http::Response::builder() + .status(429) + .header("retry-after", "60") + .body("rate limit exceeded") + .expect("a response"); + + let error = ProviderError::from_http_response(reqwest::Response::from(response)).await; + + let ProviderError::Http { + status, + body, + retry_after, + } = error + else { + panic!("an unsuccessful response is an Http error"); + }; + assert_eq!(status, reqwest::StatusCode::TOO_MANY_REQUESTS); + assert_eq!(body, "rate limit exceeded"); + assert_eq!(retry_after, Some(Duration::from_secs(60))); + } + + #[tokio::test] + async fn a_response_without_the_header_asks_for_nothing() { + let response = http::Response::builder() + .status(503) + .body("upstream is restarting") + .expect("a response"); + + let error = ProviderError::from_http_response(reqwest::Response::from(response)).await; + + assert_eq!(error.retry_after(), None); + } + + #[test] + fn both_retryable_shapes_answer_the_same_question() { + // A caller shaping a backoff asks one question and must not have to + // know which variant carried the answer. + let http = ProviderError::Http { + status: reqwest::StatusCode::TOO_MANY_REQUESTS, + body: String::new(), + retry_after: Some(Duration::from_secs(20)), + }; + let retryable = ProviderError::Retryable { + message: "connection closed".to_string(), + delay: Some(Duration::from_millis(750)), + }; + let silent = ProviderError::InvalidRequest("bad model".to_string()); + + assert_eq!(http.retry_after(), Some(Duration::from_secs(20))); + assert_eq!(retryable.retry_after(), Some(Duration::from_millis(750))); + assert_eq!(silent.retry_after(), None); + } +} diff --git a/vendor/mentra-provider/src/gemini.rs b/vendor/mentra-provider/src/gemini.rs new file mode 100644 index 0000000..e59e243 --- /dev/null +++ b/vendor/mentra-provider/src/gemini.rs @@ -0,0 +1,266 @@ +use async_trait::async_trait; +use std::sync::Arc; + +pub(crate) mod model; +pub(crate) mod sse; + +use crate::AuthScheme; +use crate::BuiltinProvider; +use crate::CompactionRequest; +use crate::CompactionResponse; +use crate::CredentialSource; +use crate::ModelCatalog; +use crate::ModelInfo; +use crate::ProviderCapabilities; +use crate::ProviderDefinition; +use crate::ProviderError; +use crate::ProviderEventStream; +use crate::ProviderSession; +use crate::ProviderSessionFactory; +use crate::RegisteredProvider; +use crate::Request; +use crate::StaticCredentialSource; +use crate::WireApi; + +const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/"; + +pub struct GeminiProvider { + client: reqwest::Client, + credential_source: Arc, + definition: ProviderDefinition, +} + +impl Clone for GeminiProvider { + fn clone(&self) -> Self { + Self { + client: self.client.clone(), + credential_source: Arc::clone(&self.credential_source), + definition: self.definition.clone(), + } + } +} + +impl GeminiProvider { + pub fn new(api_key: impl Into) -> Self { + Self::with_credential_source(StaticCredentialSource::new(api_key)) + } +} + +impl GeminiProvider +where + C: CredentialSource + 'static, +{ + pub fn with_credential_source(credential_source: C) -> Self { + Self::with_shared_credential_source(Arc::new(credential_source)) + } + + pub fn with_shared_credential_source(credential_source: Arc) -> Self { + Self::with_definition_and_shared_credential_source(Self::definition(), credential_source) + } + + pub fn with_definition_and_credential_source( + definition: ProviderDefinition, + credential_source: C, + ) -> Self { + Self::with_definition_and_shared_credential_source(definition, Arc::new(credential_source)) + } + + pub fn with_definition_and_shared_credential_source( + definition: ProviderDefinition, + credential_source: Arc, + ) -> Self { + let client = reqwest::Client::builder() + .build() + .expect("Failed to build client"); + + Self { + client, + credential_source, + definition, + } + } + + fn definition() -> ProviderDefinition { + let mut definition = ProviderDefinition::new(BuiltinProvider::Gemini); + definition.descriptor.display_name = Some("Gemini".to_string()); + definition.descriptor.description = + Some("Google Gemini Developer API provider".to_string()); + definition.wire_api = WireApi::GeminiGenerateContent; + definition.auth_scheme = AuthScheme::Header { + name: "x-goog-api-key".to_string(), + }; + definition.capabilities = ProviderCapabilities { + supports_model_listing: true, + supports_streaming: true, + supports_websockets: false, + supports_tool_calls: true, + supports_images: true, + supports_history_compaction: true, + supports_memory_summarization: true, + supports_deferred_tools: false, + supports_hosted_tool_search: false, + supports_hosted_web_search: false, + supports_image_generation: false, + supports_reasoning_effort: true, + reports_reasoning_tokens: false, + reports_thoughts_tokens: true, + supports_structured_tool_results: false, + supports_embeddings: false, + }; + definition.base_url = Some(DEFAULT_BASE_URL.to_string()); + definition + } +} + +#[async_trait] +impl ModelCatalog for GeminiProvider +where + C: CredentialSource + 'static, +{ + async fn list_models(&self) -> Result, ProviderError> { + let mut models = Vec::new(); + let mut page_token = None::; + + loop { + let credentials = self.credential_source.credentials().await?; + let mut request = self + .client + .get( + self.definition + .request_url_with_auth_for_path("v1beta/models", &credentials)?, + ) + .headers(self.definition.build_headers(&credentials)?) + .query(&[("pageSize", "1000")]); + + if let Some(token) = page_token.as_deref() { + request = request.query(&[("pageToken", token)]); + } + + let response = request.send().await.map_err(ProviderError::Transport)?; + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + let page = response + .json::() + .await + .map_err(ProviderError::Decode)?; + + models.extend( + page.models + .into_iter() + .filter(|model| model.supports_generate_content()) + .map(ModelInfo::from), + ); + + page_token = page.next_page_token; + if page_token.is_none() { + break; + } + } + + Ok(models) + } +} + +#[async_trait] +impl ProviderSessionFactory for GeminiProvider +where + C: CredentialSource + 'static, +{ + async fn create_session(&self) -> Result, ProviderError> { + Ok(Box::new((*self).clone())) + } +} + +#[async_trait] +impl ProviderSession for GeminiProvider +where + C: CredentialSource + 'static, +{ + async fn stream(&self, request: Request<'_>) -> Result { + let session = request.provider_request_options.session.clone(); + let model_name = request.model.to_string(); + let request = model::GeminiGenerateContentRequest::try_from(request)?; + let credentials = self.credential_source.credentials().await?; + let response = self + .client + .post(self.definition.request_url_with_auth_for_path( + &format!( + "v1beta/{}:streamGenerateContent", + normalize_model_name(&model_name) + ), + &credentials, + )?) + .headers(self.definition.build_headers_for_session( + &credentials, + Some(&session), + None, + )?) + .query(&[("alt", "sse")]) + .json(&request) + .send() + .await + .map_err(ProviderError::Transport)?; + + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + Ok(sse::spawn_event_stream(response, model_name)) + } + + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + let request = request.into_model_request()?; + let response = ProviderSession::send(self, request).await?; + Ok(response.into_compaction_response()) + } + + async fn summarize_memories( + &self, + request: crate::MemorySummarizeRequest<'_>, + ) -> Result { + let request = request.into_model_request()?; + let response = ProviderSession::send(self, request).await?; + response.into_memory_summarize_response() + } +} + +#[async_trait] +impl RegisteredProvider for GeminiProvider +where + C: CredentialSource + 'static, +{ + fn definition(&self) -> ProviderDefinition { + self.definition.clone() + } +} + +fn normalize_model_name(model: &str) -> String { + if model.starts_with("models/") { + model.to_string() + } else { + format!("models/{model}") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::RegisteredProvider; + + #[test] + fn definition_advertises_history_compaction_support() { + let provider = GeminiProvider::new("test-key"); + + assert!( + provider + .definition() + .capabilities + .supports_history_compaction + ); + } +} diff --git a/vendor/mentra-provider/src/gemini/model.rs b/vendor/mentra-provider/src/gemini/model.rs new file mode 100644 index 0000000..5279431 --- /dev/null +++ b/vendor/mentra-provider/src/gemini/model.rs @@ -0,0 +1,945 @@ +use std::collections::BTreeMap; + +use base64::{Engine as _, engine::general_purpose::STANDARD}; +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; + +use crate::{ + BuiltinProvider, ContentBlock, ImageSource, Message, ModelInfo, ProviderError, + ProviderToolKind, ReasoningEffort, Request, Role, ToolChoice, ToolLoadingPolicy, + ToolSearchMode, ToolSpec, +}; + +#[derive(Deserialize)] +pub(crate) struct GeminiModelsPage { + #[serde(default)] + pub(crate) models: Vec, + #[serde(default, rename = "nextPageToken", alias = "next_page_token")] + pub(crate) next_page_token: Option, +} + +#[derive(Deserialize)] +pub(crate) struct GeminiModel { + pub(crate) name: String, + #[serde(default, rename = "baseModelId", alias = "base_model_id")] + pub(crate) base_model_id: Option, + #[serde(default, rename = "displayName", alias = "display_name")] + pub(crate) display_name: Option, + #[serde(default)] + pub(crate) description: Option, + #[serde( + default, + rename = "supportedGenerationMethods", + alias = "supported_generation_methods" + )] + supported_generation_methods: Vec, +} + +impl GeminiModel { + pub(crate) fn supports_generate_content(&self) -> bool { + self.supported_generation_methods + .iter() + .any(|method| matches!(method.as_str(), "generateContent" | "streamGenerateContent")) + } +} + +impl From for ModelInfo { + fn from(model: GeminiModel) -> Self { + let id = model.base_model_id.unwrap_or_else(|| { + model + .name + .strip_prefix("models/") + .unwrap_or(&model.name) + .to_string() + }); + + ModelInfo { + id, + provider: BuiltinProvider::Gemini.into(), + display_name: model.display_name, + description: model.description, + created_at: None, + } + } +} + +#[derive(Serialize)] +pub(crate) struct GeminiGenerateContentRequest { + #[serde(rename = "systemInstruction", skip_serializing_if = "Option::is_none")] + system_instruction: Option, + contents: Vec, + #[serde(skip_serializing_if = "Vec::is_empty")] + tools: Vec, + #[serde(rename = "toolConfig", skip_serializing_if = "Option::is_none")] + tool_config: Option, + #[serde(rename = "generationConfig", skip_serializing_if = "Option::is_none")] + generation_config: Option, +} + +impl<'a> TryFrom> for GeminiGenerateContentRequest { + type Error = ProviderError; + + fn try_from(value: Request<'a>) -> Result { + let generation_config = GeminiGenerationConfig::from_request(&value)?; + let tool_name_by_id = collect_tool_name_by_id(value.messages.as_ref()); + let contents = value + .messages + .iter() + .map(|message| GeminiContent::try_from_message(message, &tool_name_by_id)) + .collect::, _>>()? + .into_iter() + .filter(|content| !content.parts.is_empty()) + .collect::>(); + validate_gemini_tools( + value.tools.as_ref(), + value.tool_choice.as_ref(), + value.provider_request_options.tool_search_mode, + )?; + let tools = if value.tools.is_empty() { + Vec::new() + } else { + vec![GeminiTool { + function_declarations: value + .tools + .iter() + .map(GeminiFunctionDeclaration::from) + .collect(), + }] + }; + + Ok(GeminiGenerateContentRequest { + system_instruction: value.system.map(|system| GeminiInstruction { + parts: vec![GeminiPart::Text { + text: system.into_owned(), + }], + }), + contents, + tool_config: value + .tool_choice + .filter(|_| !tools.is_empty()) + .map(Into::into), + tools, + generation_config, + }) + } +} + +fn validate_gemini_tools( + tools: &[ToolSpec], + tool_choice: Option<&ToolChoice>, + tool_search_mode: ToolSearchMode, +) -> Result<(), ProviderError> { + if let Some(tool) = tools + .iter() + .find(|tool| tool.kind != ProviderToolKind::Function) + { + return Err(ProviderError::InvalidRequest(format!( + "Gemini does not support provider tool kind {:?} for '{}'", + tool.kind, tool.name + ))); + } + + let forced_tool_name = match tool_choice { + Some(ToolChoice::Tool { name }) => Some(name.as_str()), + _ => None, + }; + + let has_deferred_tools = tools.iter().any(|tool| { + tool.loading_policy == ToolLoadingPolicy::Deferred + && forced_tool_name != Some(tool.name.as_str()) + }); + + if !has_deferred_tools { + return Ok(()); + } + + let message = match tool_search_mode { + ToolSearchMode::Hosted => { + "Gemini does not support hosted tool search for deferred custom tools" + } + ToolSearchMode::Disabled => { + "Gemini does not support deferred custom tools without hosted tool search" + } + }; + + Err(ProviderError::InvalidRequest(message.to_string())) +} + +fn collect_tool_name_by_id(messages: &[Message]) -> BTreeMap { + let mut names = BTreeMap::new(); + + for message in messages { + for block in &message.content { + if let ContentBlock::ToolUse { id, name, .. } = block { + names.insert(id.clone(), name.clone()); + } + } + } + + names +} + +#[derive(Serialize)] +struct GeminiInstruction { + parts: Vec, +} + +#[derive(Serialize)] +struct GeminiContent { + role: String, + parts: Vec, +} + +impl GeminiContent { + fn try_from_message( + message: &Message, + tool_name_by_id: &BTreeMap, + ) -> Result { + let role = match &message.role { + Role::User | Role::Assistant => message.role.to_string(), + Role::Unknown(role) => { + return Err(ProviderError::InvalidRequest(format!( + "Gemini message role '{role}' is not supported" + ))); + } + }; + + let mut parts = Vec::with_capacity(message.content.len()); + for block in &message.content { + parts.push(GeminiPart::try_from_block( + block, + &message.role, + tool_name_by_id, + )?); + } + + Ok(GeminiContent { role, parts }) + } +} + +#[derive(Serialize)] +#[serde(untagged)] +enum GeminiPart { + Text { + text: String, + }, + InlineData { + #[serde(rename = "inlineData")] + inline_data: GeminiInlineData, + }, + FunctionCall { + #[serde(rename = "functionCall")] + function_call: GeminiFunctionCall, + }, + FunctionResponse { + #[serde(rename = "functionResponse")] + function_response: GeminiFunctionResponse, + }, +} + +impl GeminiPart { + fn try_from_block( + block: &ContentBlock, + role: &Role, + tool_name_by_id: &BTreeMap, + ) -> Result { + match block { + ContentBlock::Text { text } => Ok(GeminiPart::Text { text: text.clone() }), + ContentBlock::Thinking { .. } => Ok(GeminiPart::Text { + text: block + .thinking_fallback_text() + .expect("thinking block has fallback text"), + }), + ContentBlock::Image { source } => { + if !matches!(role, Role::User) { + return Err(ProviderError::InvalidRequest( + "Gemini image inputs are only supported in user messages".to_string(), + )); + } + + match source { + ImageSource::Bytes { media_type, data } => Ok(GeminiPart::InlineData { + inline_data: GeminiInlineData { + mime_type: media_type.clone(), + data: STANDARD.encode(data), + }, + }), + ImageSource::Url { .. } => Err(ProviderError::InvalidRequest( + "Gemini image URL inputs are not supported without a file upload flow" + .to_string(), + )), + } + } + ContentBlock::ToolUse { name, input, .. } => Ok(GeminiPart::FunctionCall { + function_call: GeminiFunctionCall { + name: name.clone(), + args: input.clone(), + }, + }), + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => { + let name = tool_name_by_id.get(tool_use_id).cloned().ok_or_else(|| { + ProviderError::InvalidRequest(format!( + "Gemini tool result references unknown tool_use_id '{tool_use_id}'" + )) + })?; + + Ok(GeminiPart::FunctionResponse { + function_response: GeminiFunctionResponse { + name, + response: json!({ + "content": content.to_display_string(), + "is_error": is_error, + }), + }, + }) + } + ContentBlock::HostedToolSearch { call } => Ok(GeminiPart::FunctionCall { + function_call: GeminiFunctionCall { + name: "tool_search".to_string(), + args: json!({ "query": call.query }), + }, + }), + ContentBlock::HostedWebSearch { call } => Ok(GeminiPart::FunctionCall { + function_call: GeminiFunctionCall { + name: "web_search".to_string(), + args: serde_json::to_value(call.action.clone()).unwrap_or(Value::Null), + }, + }), + ContentBlock::ImageGeneration { call } => Ok(GeminiPart::FunctionCall { + function_call: GeminiFunctionCall { + name: "image_generation".to_string(), + args: json!({ + "status": call.status, + "revised_prompt": call.revised_prompt, + }), + }, + }), + } + } +} + +#[derive(Serialize)] +struct GeminiInlineData { + #[serde(rename = "mimeType")] + mime_type: String, + data: String, +} + +#[derive(Serialize)] +struct GeminiFunctionCall { + name: String, + args: Value, +} + +#[derive(Serialize)] +struct GeminiFunctionResponse { + name: String, + response: Value, +} + +#[derive(Serialize)] +struct GeminiTool { + #[serde(rename = "functionDeclarations")] + function_declarations: Vec, +} + +#[derive(Serialize)] +struct GeminiFunctionDeclaration { + name: String, + #[serde(skip_serializing_if = "Option::is_none")] + description: Option, + parameters: Value, +} + +impl From<&ToolSpec> for GeminiFunctionDeclaration { + fn from(tool: &ToolSpec) -> Self { + GeminiFunctionDeclaration { + name: tool.name.clone(), + description: tool.description.clone(), + parameters: tool.input_schema.clone(), + } + } +} + +#[derive(Serialize)] +struct GeminiToolConfig { + #[serde(rename = "functionCallingConfig")] + function_calling_config: GeminiFunctionCallingConfig, +} + +impl From for GeminiToolConfig { + fn from(choice: ToolChoice) -> Self { + let function_calling_config = match choice { + ToolChoice::Auto => GeminiFunctionCallingConfig { + mode: GeminiFunctionCallingMode::Auto, + allowed_function_names: Vec::new(), + }, + ToolChoice::Any => GeminiFunctionCallingConfig { + mode: GeminiFunctionCallingMode::Any, + allowed_function_names: Vec::new(), + }, + ToolChoice::Tool { name } => GeminiFunctionCallingConfig { + mode: GeminiFunctionCallingMode::Any, + allowed_function_names: vec![name], + }, + }; + + GeminiToolConfig { + function_calling_config, + } + } +} + +#[derive(Serialize)] +struct GeminiFunctionCallingConfig { + mode: GeminiFunctionCallingMode, + #[serde(rename = "allowedFunctionNames", skip_serializing_if = "Vec::is_empty")] + allowed_function_names: Vec, +} + +#[derive(Serialize)] +enum GeminiFunctionCallingMode { + #[serde(rename = "AUTO")] + Auto, + #[serde(rename = "ANY")] + Any, +} + +#[derive(Serialize)] +struct GeminiGenerationConfig { + #[serde(skip_serializing_if = "Option::is_none")] + temperature: Option, + #[serde(rename = "maxOutputTokens", skip_serializing_if = "Option::is_none")] + max_output_tokens: Option, + #[serde(rename = "thinkingConfig", skip_serializing_if = "Option::is_none")] + thinking_config: Option, +} + +impl GeminiGenerationConfig { + fn from_request(request: &Request<'_>) -> Result, ProviderError> { + let thinking_config = + if let Some(reasoning) = request.provider_request_options.reasoning.as_ref() { + let Some(effort) = reasoning.effort else { + return Ok(None); + }; + if !supports_gemini_thinking_level(&request.model) { + return Err(ProviderError::InvalidRequest(format!( + "Gemini reasoning effort requires a Gemini 3 model, got '{}'", + request.model + ))); + } + + Some(GeminiThinkingConfig { + thinking_level: effort.try_into()?, + }) + } else { + None + }; + + let config = GeminiGenerationConfig { + temperature: request.temperature, + max_output_tokens: request.max_output_tokens, + thinking_config, + }; + + Ok((!config.is_empty()).then_some(config)) + } + + fn is_empty(&self) -> bool { + self.temperature.is_none() + && self.max_output_tokens.is_none() + && self.thinking_config.is_none() + } +} + +#[derive(Serialize)] +struct GeminiThinkingConfig { + #[serde(rename = "thinkingLevel")] + thinking_level: GeminiThinkingLevel, +} + +#[derive(Serialize)] +#[serde(rename_all = "snake_case")] +enum GeminiThinkingLevel { + Low, + Medium, + High, +} + +impl TryFrom for GeminiThinkingLevel { + type Error = ProviderError; + + fn try_from(value: ReasoningEffort) -> Result { + match value { + ReasoningEffort::Low => Ok(Self::Low), + ReasoningEffort::Medium => Ok(Self::Medium), + ReasoningEffort::High => Ok(Self::High), + ReasoningEffort::XHigh => Err(ProviderError::InvalidRequest( + "Gemini does not support reasoning effort 'xhigh'".to_string(), + )), + ReasoningEffort::Max => Err(ProviderError::InvalidRequest( + "Gemini does not support reasoning effort 'max'".to_string(), + )), + } + } +} + +fn supports_gemini_thinking_level(model: &str) -> bool { + let model = model.strip_prefix("models/").unwrap_or(model); + model.starts_with("gemini-3") +} + +#[cfg(test)] +mod tests { + use std::{borrow::Cow, collections::BTreeMap}; + + use serde_json::json; + + use crate::{ + BuiltinProvider, ContentBlock, Message, ModelInfo, ProviderError, ProviderRequestOptions, + ReasoningEffort, ReasoningOptions, Request, Role, ToolChoice, ToolLoadingPolicy, + ToolResultContent, ToolSearchMode, ToolSpec, + }; + + use super::{GeminiGenerateContentRequest, GeminiModel}; + + #[test] + fn converts_model_name_to_base_model_id() { + let model = GeminiModel { + name: "models/gemini-3-flash".to_string(), + base_model_id: Some("gemini-3-flash".to_string()), + display_name: Some("Gemini 3 Flash".to_string()), + description: Some("Test".to_string()), + supported_generation_methods: vec!["generateContent".to_string()], + }; + + let info = ModelInfo::from(model); + + assert_eq!(info.id, "gemini-3-flash"); + assert_eq!(info.provider, BuiltinProvider::Gemini.into()); + assert_eq!(info.display_name.as_deref(), Some("Gemini 3 Flash")); + } + + #[test] + fn converts_request_to_gemini_payload() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: Some(Cow::Borrowed("Be helpful.")), + messages: Cow::Owned(vec![ + Message::user(ContentBlock::text("What files changed?")), + Message::assistant(ContentBlock::ToolUse { + id: "call_1".to_string(), + name: "files".to_string(), + input: json!({ "operations": [{ "op": "read", "path": "README.md" }] }), + }), + Message::user(ContentBlock::ToolResult { + tool_use_id: "call_1".to_string(), + content: ToolResultContent::text("README contents"), + is_error: false, + }), + ]), + tools: Cow::Owned(vec![ToolSpec { + name: "files".to_string(), + description: Some("Read and edit files".to_string()), + input_schema: json!({ + "type": "object", + "properties": { + "operations": { "type": "array" } + } + }), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Immediate, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Tool { + name: "files".to_string(), + }), + temperature: Some(0.2), + max_output_tokens: Some(256), + metadata: Cow::Owned(BTreeMap::from([( + "agent".to_string(), + "mentra".to_string(), + )])), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!( + payload["systemInstruction"]["parts"][0]["text"], + "Be helpful." + ); + assert_eq!(payload["contents"][0]["role"], "user"); + assert_eq!( + payload["contents"][0]["parts"][0]["text"], + "What files changed?" + ); + assert_eq!( + payload["contents"][1]["parts"][0]["functionCall"]["name"], + "files" + ); + assert_eq!( + payload["contents"][2]["parts"][0]["functionResponse"]["name"], + "files" + ); + assert_eq!( + payload["contents"][2]["parts"][0]["functionResponse"]["response"]["content"], + "README contents" + ); + assert_eq!( + payload["tools"][0]["functionDeclarations"][0]["name"], + "files" + ); + assert_eq!( + payload["toolConfig"]["functionCallingConfig"]["mode"], + "ANY" + ); + assert_eq!( + payload["toolConfig"]["functionCallingConfig"]["allowedFunctionNames"][0], + "files" + ); + let temperature = payload["generationConfig"]["temperature"] + .as_f64() + .expect("temperature should be numeric"); + assert!((temperature - 0.2).abs() < 1e-6); + assert_eq!(payload["generationConfig"]["maxOutputTokens"], 256); + assert!(payload.get("metadata").is_none()); + } + + #[test] + fn serializes_inline_images_into_inline_data_parts() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message { + role: Role::User, + content: vec![ + ContentBlock::text("Describe this"), + ContentBlock::image_bytes("image/png", [1_u8, 2, 3]), + ], + }]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["contents"][0]["parts"][0]["text"], "Describe this"); + assert_eq!( + payload["contents"][0]["parts"][1]["inlineData"]["mimeType"], + "image/png" + ); + assert_eq!( + payload["contents"][0]["parts"][1]["inlineData"]["data"], + "AQID" + ); + } + + #[test] + fn rejects_url_images() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::image_url( + "https://example.com/image.png", + ))]), + tools: Cow::Owned(vec![]), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let error = GeminiGenerateContentRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("image URL inputs are not supported")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn serializes_tool_choice_modes() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "echo".to_string(), + description: None, + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Immediate, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Any), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + let any_payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + assert_eq!( + any_payload["toolConfig"]["functionCallingConfig"]["mode"], + "ANY" + ); + + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "echo".to_string(), + description: None, + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Immediate, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + let auto_payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + assert_eq!( + auto_payload["toolConfig"]["functionCallingConfig"]["mode"], + "AUTO" + ); + } + + #[test] + fn omits_tool_config_when_tool_choice_is_unset() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "echo".to_string(), + description: None, + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Immediate, + strict: None, + options: None, + }]), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert!(payload.get("toolConfig").is_none()); + } + + #[test] + fn serializes_shared_reasoning_effort_for_gemini_3_models() { + for (effort, expected) in [ + (ReasoningEffort::Low, "low"), + (ReasoningEffort::Medium, "medium"), + (ReasoningEffort::High, "high"), + ] { + let request = Request { + model: Cow::Borrowed("gemini-3-flash-preview"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(ReasoningOptions { + effort: Some(effort), + summary: None, + }), + ..Default::default() + }, + }; + + let payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!( + payload["generationConfig"]["thinkingConfig"]["thinkingLevel"], + expected + ); + } + } + + #[test] + fn rejects_reasoning_effort_for_gemini_2_5_models() { + let request = Request { + model: Cow::Borrowed("gemini-2.5-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(ReasoningOptions { + effort: Some(ReasoningEffort::Low), + summary: None, + }), + ..Default::default() + }, + }; + + let error = GeminiGenerateContentRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("Gemini 3")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn rejects_extended_reasoning_effort_for_gemini() { + for (effort, expected) in [ + (ReasoningEffort::XHigh, "xhigh"), + (ReasoningEffort::Max, "max"), + ] { + let request = Request { + model: Cow::Borrowed("gemini-3-flash-preview"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(ReasoningOptions { + effort: Some(effort), + summary: None, + }), + ..Default::default() + }, + }; + + let error = GeminiGenerateContentRequest::try_from(request) + .err() + .expect("extended Gemini effort should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains(expected)); + } + other => panic!("unexpected error: {other:?}"), + } + } + } + + #[test] + fn rejects_hosted_tool_search_with_deferred_tools() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "echo".to_string(), + description: None, + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + ..Default::default() + }, + }; + + let error = GeminiGenerateContentRequest::try_from(request) + .err() + .expect("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("does not support hosted tool search")); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn forced_deferred_tool_still_serializes_as_function_declaration() { + let request = Request { + model: Cow::Borrowed("gemini-2.0-flash"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "echo".to_string(), + description: None, + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Tool { + name: "echo".to_string(), + }), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + ..Default::default() + }, + }; + + let payload = + serde_json::to_value(GeminiGenerateContentRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!( + payload["tools"][0]["functionDeclarations"][0]["name"], + "echo" + ); + } +} diff --git a/vendor/mentra-provider/src/gemini/sse.rs b/vendor/mentra-provider/src/gemini/sse.rs new file mode 100644 index 0000000..05cb988 --- /dev/null +++ b/vendor/mentra-provider/src/gemini/sse.rs @@ -0,0 +1,714 @@ +use std::collections::{BTreeSet, HashMap}; + +use futures_util::StreamExt; +use serde::Deserialize; +use serde_json::Value; +use tokio::sync::mpsc; + +use crate::{ + ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent, ProviderEventStream, Role, + TokenUsage, +}; + +pub(crate) fn spawn_event_stream( + response: reqwest::Response, + request_model: String, +) -> ProviderEventStream { + let (tx, rx) = mpsc::unbounded_channel(); + + tokio::spawn(async move { + if let Err(error) = forward_events(response, request_model, tx.clone()).await { + let _ = tx.send(Err(error)); + } + }); + + rx +} + +async fn forward_events( + response: reqwest::Response, + request_model: String, + tx: mpsc::UnboundedSender>, +) -> Result<(), ProviderError> { + let mut bytes_stream = response.bytes_stream(); + let mut buffer = Vec::new(); + let mut state = StreamState::new(request_model); + + while let Some(chunk) = bytes_stream.next().await { + let chunk = chunk.map_err(ProviderError::Transport)?; + buffer.extend_from_slice(&chunk); + + while let Some((frame_end, delimiter_len)) = find_frame_boundary(&buffer) { + let frame = buffer.drain(..frame_end).collect::>(); + buffer.drain(..delimiter_len); + + for event in parse_frame(&frame, &mut state)? { + if tx.send(Ok(event)).is_err() { + return Ok(()); + } + } + } + } + + if !buffer.is_empty() { + for event in parse_frame(&buffer, &mut state)? { + let _ = tx.send(Ok(event)); + } + } + + if state.started && !state.stopped { + return Err(ProviderError::MalformedStream( + "Gemini stream ended before MessageStopped".to_string(), + )); + } + + Ok(()) +} + +struct StreamState { + request_model: String, + response_id: Option, + model_version: Option, + started: bool, + stopped: bool, + latest_usage: Option, + open_blocks: BTreeSet, + text_snapshots: HashMap, + tool_snapshots: HashMap, + tool_call_ids: HashMap, +} + +impl StreamState { + fn new(request_model: String) -> Self { + Self { + request_model, + response_id: None, + model_version: None, + started: false, + stopped: false, + latest_usage: None, + open_blocks: BTreeSet::new(), + text_snapshots: HashMap::new(), + tool_snapshots: HashMap::new(), + tool_call_ids: HashMap::new(), + } + } + + fn ensure_message_started(&mut self, chunk: &GeminiStreamChunk) -> Option { + if self.started { + return None; + } + + self.started = true; + self.response_id = chunk + .response_id + .clone() + .or_else(|| Some(format!("gemini-{}", self.request_model))); + self.model_version = chunk.model_version.clone(); + + Some(ProviderEvent::MessageStarted { + id: self + .response_id + .clone() + .unwrap_or_else(|| "gemini".to_string()), + model: self + .model_version + .clone() + .unwrap_or_else(|| self.request_model.clone()), + role: Role::Assistant, + }) + } + + fn ensure_text_block_started(&mut self, index: usize) -> Option { + if self.open_blocks.insert(index) { + Some(ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Text, + }) + } else { + None + } + } + + fn ensure_tool_block_started( + &mut self, + index: usize, + function_call: &GeminiFunctionCall, + ) -> Option { + if self.open_blocks.insert(index) { + let response_id = self + .response_id + .clone() + .unwrap_or_else(|| format!("gemini-{}", self.request_model)); + let id = format!("{response_id}-{index}-{}", function_call.name); + self.tool_call_ids.insert(index, id.clone()); + Some(ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolUse { + id, + name: function_call.name.clone(), + }, + }) + } else { + None + } + } + + fn close_all_blocks(&mut self) -> Vec { + let indices = self.open_blocks.iter().copied().collect::>(); + self.open_blocks.clear(); + self.text_snapshots.clear(); + self.tool_snapshots.clear(); + self.tool_call_ids.clear(); + + indices + .into_iter() + .map(|index| ProviderEvent::ContentBlockStopped { index }) + .collect() + } + + fn update_usage(&mut self, usage: Option) -> Option { + match usage { + Some(usage) if self.latest_usage.as_ref() != Some(&usage) => { + self.latest_usage = Some(usage.clone()); + Some(usage) + } + Some(usage) => { + self.latest_usage = Some(usage); + None + } + None => None, + } + } +} + +fn parse_frame(frame: &[u8], state: &mut StreamState) -> Result, ProviderError> { + let frame = std::str::from_utf8(frame) + .map_err(|error| ProviderError::MalformedStream(error.to_string()))?; + let mut data_lines = Vec::new(); + + for raw_line in frame.lines() { + let line = raw_line.strip_suffix('\r').unwrap_or(raw_line); + if line.is_empty() || line.starts_with(':') || line.starts_with("event:") { + continue; + } + + if let Some(rest) = line.strip_prefix("data:") { + data_lines.push(rest.trim_start().to_string()); + } + } + + if data_lines.is_empty() { + return Ok(Vec::new()); + } + + let data = data_lines.join("\n"); + let chunk: GeminiStreamChunk = + serde_json::from_str(&data).map_err(ProviderError::Deserialize)?; + + if let Some(error) = chunk.error { + return Err(ProviderError::MalformedStream( + error + .message + .unwrap_or_else(|| "gemini stream error".to_string()), + )); + } + + let mut events = Vec::new(); + let latest_usage = chunk + .usage_metadata + .as_ref() + .and_then(GeminiUsageMetadata::to_token_usage); + let usage_changed = state.update_usage(latest_usage.clone()); + + if let Some(candidate) = chunk.candidates.first() { + if let Some(event) = state.ensure_message_started(&chunk) { + events.push(event); + } + + if let Some(content) = candidate.content.as_ref() { + for (index, part) in content.parts.iter().enumerate() { + if let Some(text) = part.text.as_ref() { + if let Some(event) = state.ensure_text_block_started(index) { + events.push(event); + } + if let Some(delta) = merge_chunk( + state.text_snapshots.entry(index).or_default(), + text.as_str(), + ) { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::Text(delta), + }); + } + } else if let Some(function_call) = part.function_call.as_ref() { + if let Some(event) = state.ensure_tool_block_started(index, function_call) { + events.push(event); + } + if let Some(delta) = merge_chunk( + state.tool_snapshots.entry(index).or_default(), + &serde_json::to_string(&function_call.args) + .expect("function call args should serialize"), + ) { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ToolUseInputJson(delta), + }); + } + } + } + } + + if let Some(stop_reason) = candidate.finish_reason.clone() { + events.extend(state.close_all_blocks()); + events.push(ProviderEvent::MessageDelta { + stop_reason: Some(stop_reason), + usage: latest_usage.or_else(|| state.latest_usage.clone()), + }); + events.push(ProviderEvent::MessageStopped); + state.stopped = true; + } else if let Some(usage) = usage_changed { + events.push(ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(usage), + }); + } + } else if let Some(prompt_feedback) = chunk.prompt_feedback.as_ref() { + if let Some(event) = state.ensure_message_started(&chunk) { + events.push(event); + } + let _ = state.update_usage(latest_usage.clone()); + events.extend(state.close_all_blocks()); + events.push(ProviderEvent::MessageDelta { + stop_reason: Some(prompt_feedback.stop_reason()), + usage: latest_usage.or_else(|| state.latest_usage.clone()), + }); + events.push(ProviderEvent::MessageStopped); + state.stopped = true; + } + + Ok(events) +} + +fn merge_chunk(previous: &mut String, current: &str) -> Option { + if current.is_empty() { + return None; + } + + if previous.is_empty() { + *previous = current.to_string(); + return Some(current.to_string()); + } + + if current == previous { + return None; + } + + if current.starts_with(previous.as_str()) { + let delta = current[previous.len()..].to_string(); + *previous = current.to_string(); + return (!delta.is_empty()).then_some(delta); + } + + previous.push_str(current); + Some(current.to_string()) +} + +fn find_frame_boundary(buffer: &[u8]) -> Option<(usize, usize)> { + for (index, window) in buffer.windows(2).enumerate() { + if window == b"\n\n" { + return Some((index, 2)); + } + } + + for (index, window) in buffer.windows(4).enumerate() { + if window == b"\r\n\r\n" { + return Some((index, 4)); + } + } + + None +} + +#[derive(Deserialize)] +struct GeminiStreamChunk { + #[serde(default)] + candidates: Vec, + #[serde(default, rename = "promptFeedback", alias = "prompt_feedback")] + prompt_feedback: Option, + #[serde(default, rename = "usageMetadata", alias = "usage_metadata")] + usage_metadata: Option, + #[serde(default, rename = "responseId", alias = "response_id")] + response_id: Option, + #[serde(default, rename = "modelVersion", alias = "model_version")] + model_version: Option, + #[serde(default)] + error: Option, +} + +#[derive(Deserialize)] +struct GeminiCandidate { + #[serde(default)] + content: Option, + #[serde(default, rename = "finishReason", alias = "finish_reason")] + finish_reason: Option, +} + +#[derive(Deserialize)] +struct GeminiContent { + #[allow(dead_code)] + #[serde(default)] + role: Option, + #[serde(default)] + parts: Vec, +} + +#[derive(Deserialize)] +struct GeminiPart { + #[serde(default)] + text: Option, + #[serde(default, rename = "functionCall", alias = "function_call")] + function_call: Option, +} + +#[derive(Deserialize)] +struct GeminiFunctionCall { + name: String, + #[serde(default)] + args: Value, +} + +#[derive(Deserialize)] +struct GeminiErrorBody { + #[serde(default)] + message: Option, +} + +#[derive(Deserialize)] +struct GeminiPromptFeedback { + #[serde(default, rename = "blockReason", alias = "block_reason")] + block_reason: Option, +} + +impl GeminiPromptFeedback { + fn stop_reason(&self) -> String { + self.block_reason + .clone() + .unwrap_or_else(|| "BLOCKED".to_string()) + } +} + +#[derive(Deserialize)] +struct GeminiUsageMetadata { + #[serde(default, rename = "promptTokenCount", alias = "prompt_token_count")] + prompt_token_count: Option, + #[serde( + default, + rename = "candidatesTokenCount", + alias = "candidates_token_count" + )] + candidates_token_count: Option, + #[serde(default, rename = "totalTokenCount", alias = "total_token_count")] + total_token_count: Option, + #[serde( + default, + rename = "cachedContentTokenCount", + alias = "cached_content_token_count" + )] + cached_content_token_count: Option, + #[serde(default, rename = "thoughtsTokenCount", alias = "thoughts_token_count")] + thoughts_token_count: Option, + #[serde( + default, + rename = "toolUsePromptTokenCount", + alias = "tool_use_prompt_token_count" + )] + tool_use_prompt_token_count: Option, +} + +impl GeminiUsageMetadata { + fn to_token_usage(&self) -> Option { + let usage = TokenUsage { + input_tokens: self.prompt_token_count, + output_tokens: self.candidates_token_count, + total_tokens: self.total_token_count, + cache_read_input_tokens: self.cached_content_token_count, + cache_creation_input_tokens: None, + reasoning_tokens: None, + thoughts_tokens: self.thoughts_token_count, + tool_input_tokens: self.tool_use_prompt_token_count, + }; + + (!usage.is_empty()).then_some(usage) + } +} + +#[cfg(test)] +mod tests { + use crate::{ContentBlockDelta, ContentBlockStart, ProviderEvent, TokenUsage}; + + use super::{StreamState, parse_frame}; + + #[test] + fn streams_text_and_completion_events() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + + let events = parse_frame( + br#"data: {"responseId":"resp-1","modelVersion":"gemini-2.0-flash-001","candidates":[{"content":{"role":"model","parts":[{"text":"Hel"}]}}]}"#, + &mut state, + ) + .expect("frame should parse"); + + assert_eq!( + events, + vec![ + ProviderEvent::MessageStarted { + id: "resp-1".to_string(), + model: "gemini-2.0-flash-001".to_string(), + role: crate::Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hel".to_string()), + }, + ] + ); + + let events = parse_frame( + br#"data: {"candidates":[{"content":{"parts":[{"text":"lo"}]}}]}"#, + &mut state, + ) + .expect("frame should parse"); + assert_eq!( + events, + vec![ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("lo".to_string()), + }] + ); + + let events = parse_frame( + br#"data: {"candidates":[{"finishReason":"STOP"}]}"#, + &mut state, + ) + .expect("frame should parse"); + assert_eq!( + events, + vec![ + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: Some("STOP".to_string()), + usage: None, + }, + ProviderEvent::MessageStopped, + ] + ); + } + + #[test] + fn streams_function_calls() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + + let events = parse_frame( + br#"data: {"responseId":"resp-1","candidates":[{"content":{"parts":[{"functionCall":{"name":"read_file","args":{"path":"README.md"}}}]}}]}"#, + &mut state, + ) + .expect("frame should parse"); + + assert_eq!( + events, + vec![ + ProviderEvent::MessageStarted { + id: "resp-1".to_string(), + model: "gemini-2.0-flash".to_string(), + role: crate::Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "resp-1-0-read_file".to_string(), + name: "read_file".to_string(), + }, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson( + "{\"path\":\"README.md\"}".to_string() + ), + }, + ] + ); + } + + #[test] + fn ignores_duplicate_full_function_call_payloads() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + parse_frame( + br#"data: {"responseId":"resp-1","candidates":[{"content":{"parts":[{"functionCall":{"name":"read_file","args":{"path":"README.md"}}}]}}]}"#, + &mut state, + ) + .expect("first frame should parse"); + + let events = parse_frame( + br#"data: {"candidates":[{"content":{"parts":[{"functionCall":{"name":"read_file","args":{"path":"README.md"}}}]}}]}"#, + &mut state, + ) + .expect("second frame should parse"); + + assert!(events.is_empty()); + } + + #[test] + fn ignores_unsupported_parts_without_breaking_indexes() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + + let events = parse_frame( + br#"data: {"responseId":"resp-1","candidates":[{"content":{"parts":[{"fileData":{"mimeType":"image/png","fileUri":"files/1"}},{"text":"Done"}]}}]}"#, + &mut state, + ) + .expect("frame should parse"); + + assert_eq!( + events, + vec![ + ProviderEvent::MessageStarted { + id: "resp-1".to_string(), + model: "gemini-2.0-flash".to_string(), + role: crate::Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 1, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 1, + delta: ContentBlockDelta::Text("Done".to_string()), + }, + ] + ); + } + + #[test] + fn surfaces_stream_errors() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + let error = parse_frame(br#"data: {"error":{"message":"boom"}}"#, &mut state) + .expect_err("frame should fail"); + + match error { + crate::ProviderError::MalformedStream(message) => { + assert_eq!(message, "boom"); + } + other => panic!("unexpected error: {other:?}"), + } + } + + #[test] + fn treats_prompt_feedback_only_chunks_as_terminal() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + + let events = parse_frame( + br#"data: {"responseId":"resp-2","promptFeedback":{"blockReason":"SAFETY"},"usageMetadata":{"promptTokenCount":11,"totalTokenCount":11,"cachedContentTokenCount":4}}"#, + &mut state, + ) + .expect("frame should parse"); + + assert_eq!( + events, + vec![ + ProviderEvent::MessageStarted { + id: "resp-2".to_string(), + model: "gemini-2.0-flash".to_string(), + role: crate::Role::Assistant, + }, + ProviderEvent::MessageDelta { + stop_reason: Some("SAFETY".to_string()), + usage: Some(TokenUsage { + input_tokens: Some(11), + output_tokens: None, + total_tokens: Some(11), + cache_read_input_tokens: Some(4), + cache_creation_input_tokens: None, + reasoning_tokens: None, + thoughts_tokens: None, + tool_input_tokens: None, + }), + }, + ProviderEvent::MessageStopped, + ] + ); + } + + #[test] + fn emits_usage_metadata_updates_and_final_usage() { + let mut state = StreamState::new("gemini-2.0-flash".to_string()); + + let events = parse_frame( + br#"data: {"responseId":"resp-3","candidates":[{"content":{"parts":[{"text":"Hi"}]}}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":1,"totalTokenCount":9,"thoughtsTokenCount":2,"toolUsePromptTokenCount":1}}"#, + &mut state, + ) + .expect("frame should parse"); + + assert_eq!( + events, + vec![ + ProviderEvent::MessageStarted { + id: "resp-3".to_string(), + model: "gemini-2.0-flash".to_string(), + role: crate::Role::Assistant, + }, + ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Text, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::Text("Hi".to_string()), + }, + ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(TokenUsage { + input_tokens: Some(8), + output_tokens: Some(1), + total_tokens: Some(9), + cache_read_input_tokens: None, + cache_creation_input_tokens: None, + reasoning_tokens: None, + thoughts_tokens: Some(2), + tool_input_tokens: Some(1), + }), + }, + ] + ); + + let events = parse_frame( + br#"data: {"candidates":[{"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":2,"totalTokenCount":10,"thoughtsTokenCount":2,"toolUsePromptTokenCount":1}}"#, + &mut state, + ) + .expect("frame should parse"); + + assert_eq!( + events, + vec![ + ProviderEvent::ContentBlockStopped { index: 0 }, + ProviderEvent::MessageDelta { + stop_reason: Some("STOP".to_string()), + usage: Some(TokenUsage { + input_tokens: Some(8), + output_tokens: Some(2), + total_tokens: Some(10), + cache_read_input_tokens: None, + cache_creation_input_tokens: None, + reasoning_tokens: None, + thoughts_tokens: Some(2), + tool_input_tokens: Some(1), + }), + }, + ProviderEvent::MessageStopped, + ] + ); + } +} diff --git a/vendor/mentra-provider/src/lib.rs b/vendor/mentra-provider/src/lib.rs new file mode 100644 index 0000000..0e8bd87 --- /dev/null +++ b/vendor/mentra-provider/src/lib.rs @@ -0,0 +1,154 @@ +pub mod anthropic; +mod auth; +mod definition; +pub mod embedding; +mod error; +pub mod gemini; +mod model; +mod registry; +mod request; +mod response; +pub mod responses; +mod stream; +mod tool; + +pub use auth::AuthScheme; +pub use auth::CredentialSource; +pub use auth::ProviderCredentials; +pub use auth::StaticCredentialSource; +pub use definition::BuiltinProvider; +pub use definition::ProviderCapabilities; +pub use definition::ProviderDefinition; +pub use definition::ProviderDescriptor; +pub use definition::ProviderId; +pub use definition::RetryPolicy; +pub use definition::WireApi; +pub use embedding::EmbeddingData; +pub use embedding::EmbeddingModelInfo; +pub use embedding::EmbeddingProvider; +pub use embedding::EmbeddingRequest; +pub use embedding::EmbeddingResponse; +pub use embedding::EmbeddingUsage; +pub use error::ProviderError; +pub use model::ContentBlock; +pub use model::HostedToolSearchCall; +pub use model::HostedWebSearchCall; +pub use model::ImageGenerationCall; +pub use model::ImageGenerationResult; +pub use model::ImageSource; +pub use model::Message; +pub use model::ModelInfo; +pub use model::ModelSelector; +pub use model::ReasoningFormat; +pub use model::ReasoningProvenance; +pub use model::Role; +pub use model::TokenUsage; +pub use model::ToolChoice; +pub use model::ToolResultContent; +pub use model::WebSearchAction; +pub use registry::ModelCatalog; +pub use registry::Provider; +pub use registry::ProviderRegistry; +pub use registry::ProviderSession; +pub use registry::ProviderSessionFactory; +pub use registry::RegisteredProvider; +pub use request::AnthropicRequestOptions; +pub use request::CompactionInputItem; +pub use request::CompactionRequest; +pub use request::GeminiRequestOptions; +pub use request::MemorySummarizeRequest; +pub use request::ProviderRequestOptions; +pub use request::RawMemory; +pub use request::RawMemoryMetadata; +pub use request::ReasoningEffort; +pub use request::ReasoningOptions; +pub use request::ReasoningSummary; +pub use request::Request; +pub use request::ResponsesRequestCompression; +pub use request::ResponsesRequestOptions; +pub use request::ResponsesStateMode; +pub use request::ResponsesTextControls; +pub use request::ResponsesTextFormat; +pub use request::ResponsesTextFormatType; +pub use request::ResponsesTransport; +pub use request::ResponsesVerbosity; +pub use request::SessionRequestOptions; +pub use request::ToolSearchMode; +pub use response::CompactionResponse; +pub use response::MemorySummarizeOutput; +pub use response::MemorySummarizeResponse; +pub use response::Response; +pub use response::collect_response_from_stream; +pub use response::provider_event_stream_from_response; +pub use stream::ContentBlockDelta; +pub use stream::ContentBlockStart; +pub use stream::ProviderEvent; +pub use stream::ProviderEventStream; +pub use stream::ResponseHeaders; +pub use tool::ProviderToolKind; +pub use tool::ToolLoadingPolicy; +pub use tool::ToolSpec; +pub use tool::ToolSpecBuilder; + +pub type OpenAIRequestOptions = ResponsesRequestOptions; + +pub mod provider { + pub use crate::Provider; + + pub mod model { + pub use crate::AnthropicRequestOptions; + pub use crate::ContentBlock; + pub use crate::ContentBlockDelta; + pub use crate::ContentBlockStart; + pub use crate::HostedToolSearchCall; + pub use crate::HostedWebSearchCall; + pub use crate::ImageGenerationCall; + pub use crate::ImageGenerationResult; + pub use crate::ImageSource; + pub use crate::Message; + pub use crate::ModelInfo; + pub use crate::OpenAIRequestOptions; + pub use crate::ProviderError; + pub use crate::ProviderEvent; + pub use crate::ProviderEventStream; + pub use crate::ProviderId; + pub use crate::ProviderRequestOptions; + pub use crate::ReasoningEffort; + pub use crate::ReasoningFormat; + pub use crate::ReasoningOptions; + pub use crate::ReasoningProvenance; + pub use crate::ReasoningSummary; + pub use crate::Request; + pub use crate::Response; + pub use crate::ResponsesTextControls; + pub use crate::ResponsesTextFormat; + pub use crate::ResponsesTextFormatType; + pub use crate::ResponsesVerbosity; + pub use crate::Role; + pub use crate::SessionRequestOptions; + pub use crate::TokenUsage; + pub use crate::ToolChoice; + pub use crate::ToolResultContent; + pub use crate::ToolSearchMode; + pub use crate::WebSearchAction; + pub use crate::collect_response_from_stream; + pub use crate::provider_event_stream_from_response; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn provider_definition_defaults_to_responses_wire_api_and_websockets_disabled() { + let definition = ProviderDefinition::new(BuiltinProvider::OpenAI); + + assert_eq!( + definition.descriptor.id, + ProviderId::from(BuiltinProvider::OpenAI) + ); + assert_eq!(definition.wire_api, WireApi::Responses); + assert!(!definition.capabilities.supports_websockets); + } +} diff --git a/vendor/mentra-provider/src/model.rs b/vendor/mentra-provider/src/model.rs new file mode 100644 index 0000000..bdcdf1e --- /dev/null +++ b/vendor/mentra-provider/src/model.rs @@ -0,0 +1,509 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::fmt::Display; +use time::OffsetDateTime; + +/// Metadata describing a model available from a provider. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ModelInfo { + pub id: String, + pub provider: crate::ProviderId, + pub display_name: Option, + pub description: Option, + pub created_at: Option, +} + +impl ModelInfo { + pub fn new(id: impl Into, provider: impl Into) -> Self { + Self { + id: id.into(), + provider: provider.into(), + display_name: None, + description: None, + created_at: None, + } + } +} + +/// Selection strategy used when resolving a model from a provider. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ModelSelector { + Id(String), + NewestAvailable, +} + +/// Provider-neutral token usage metadata for a completed or in-progress response. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct TokenUsage { + pub input_tokens: Option, + pub output_tokens: Option, + pub total_tokens: Option, + pub cache_read_input_tokens: Option, + pub cache_creation_input_tokens: Option, + pub reasoning_tokens: Option, + pub thoughts_tokens: Option, + pub tool_input_tokens: Option, +} + +impl TokenUsage { + pub fn is_empty(&self) -> bool { + self.input_tokens.is_none() + && self.output_tokens.is_none() + && self.total_tokens.is_none() + && self.cache_read_input_tokens.is_none() + && self.cache_creation_input_tokens.is_none() + && self.reasoning_tokens.is_none() + && self.thoughts_tokens.is_none() + && self.tool_input_tokens.is_none() + } +} + +/// Provider-neutral chat role labels. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum Role { + User, + Assistant, + Unknown(String), +} + +impl Display for Role { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let value = match self { + Self::User => "user", + Self::Assistant => "assistant", + Self::Unknown(role) => role.as_str(), + }; + f.write_str(value) + } +} + +/// Image payload supported by model providers. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum ImageSource { + Bytes { media_type: String, data: Vec }, + Url { url: String }, +} + +impl ImageSource { + pub fn bytes(media_type: impl Into, data: impl Into>) -> Self { + Self::Bytes { + media_type: media_type.into(), + data: data.into(), + } + } + + pub fn url(url: impl Into) -> Self { + Self::Url { url: url.into() } + } +} + +/// Tool result payloads supported by provider streams and history replay. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum ToolResultContent { + Text(String), + Structured(Value), +} + +impl ToolResultContent { + pub fn text(value: impl Into) -> Self { + Self::Text(value.into()) + } + + pub fn len(&self) -> usize { + match self { + Self::Text(text) => text.len(), + Self::Structured(value) => value.to_string().len(), + } + } + + pub fn is_empty(&self) -> bool { + self.len() == 0 + } + + pub fn clear(&mut self) { + *self = Self::Text(String::new()); + } + + pub fn as_str(&self) -> &str { + match self { + Self::Text(text) => text.as_str(), + Self::Structured(_) => panic!("ToolResultContent::as_str requires text content"), + } + } + + pub fn contains(&self, pattern: &str) -> bool { + match self { + Self::Text(text) => text.contains(pattern), + Self::Structured(value) => value.to_string().contains(pattern), + } + } + + pub fn starts_with(&self, pattern: &str) -> bool { + match self { + Self::Text(text) => text.starts_with(pattern), + Self::Structured(value) => value.to_string().starts_with(pattern), + } + } + + pub fn push_str(&mut self, value: &str) { + match self { + Self::Text(text) => text.push_str(value), + Self::Structured(existing) => { + let mut text = existing.to_string(); + text.push_str(value); + *self = Self::Text(text); + } + } + } + + pub fn to_display_string(&self) -> String { + match self { + Self::Text(text) => text.clone(), + Self::Structured(value) => value.to_string(), + } + } +} + +impl Default for ToolResultContent { + fn default() -> Self { + Self::Text(String::new()) + } +} + +impl From for ToolResultContent { + fn from(value: String) -> Self { + Self::Text(value) + } +} + +impl Display for ToolResultContent { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(&self.to_display_string()) + } +} + +impl PartialEq<&str> for ToolResultContent { + fn eq(&self, other: &&str) -> bool { + self.to_display_string() == *other + } +} + +impl PartialEq for ToolResultContent { + fn eq(&self, other: &str) -> bool { + self.to_display_string() == other + } +} + +impl PartialEq for &str { + fn eq(&self, other: &ToolResultContent) -> bool { + *self == other.to_display_string() + } +} + +impl PartialEq for str { + fn eq(&self, other: &ToolResultContent) -> bool { + self == other.to_display_string() + } +} + +impl From<&str> for ToolResultContent { + fn from(value: &str) -> Self { + Self::Text(value.to_string()) + } +} + +/// Provider-neutral hosted tool search action. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct HostedToolSearchCall { + pub id: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub status: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub query: Option, +} + +/// Provider-neutral hosted web search actions. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum WebSearchAction { + Search { + #[serde(default, skip_serializing_if = "Option::is_none")] + query: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + queries: Option>, + }, + OpenPage { + #[serde(default, skip_serializing_if = "Option::is_none")] + url: Option, + }, + FindInPage { + #[serde(default, skip_serializing_if = "Option::is_none")] + url: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pattern: Option, + }, +} + +/// Provider-neutral hosted web search call. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct HostedWebSearchCall { + pub id: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub status: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub action: Option, +} + +/// Provider-neutral image generation result. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ImageGenerationResult { + Image { source: ImageSource }, + ArtifactRef { artifact_id: String }, +} + +/// Provider-neutral image generation call. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ImageGenerationCall { + pub id: String, + pub status: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub revised_prompt: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub result: Option, +} + +/// Provider-specific format carried by a provider-neutral reasoning block. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ReasoningFormat { + AnthropicSigned, + OpenAiEncrypted, + GeminiThought, +} + +/// Origin required to decide whether opaque reasoning metadata is safe to replay. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ReasoningProvenance { + pub provider: crate::ProviderId, + pub model: String, + pub format: ReasoningFormat, +} + +/// A provider-neutral content block exchanged with models. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum ContentBlock { + Text { + text: String, + }, + Thinking { + #[serde(default, skip_serializing_if = "String::is_empty")] + thinking: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + signature: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + encrypted_content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + provenance: Option, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + redacted: bool, + }, + Image { + source: ImageSource, + }, + ToolUse { + id: String, + name: String, + input: Value, + }, + ToolResult { + tool_use_id: String, + content: ToolResultContent, + is_error: bool, + }, + HostedToolSearch { + call: HostedToolSearchCall, + }, + HostedWebSearch { + call: HostedWebSearchCall, + }, + ImageGeneration { + call: ImageGenerationCall, + }, +} + +impl ContentBlock { + pub fn text(text: impl Into) -> Self { + Self::Text { text: text.into() } + } + + pub fn thinking(thinking: impl Into) -> Self { + Self::Thinking { + thinking: thinking.into(), + signature: None, + encrypted_content: None, + id: None, + provenance: None, + redacted: false, + } + } + + pub(crate) fn thinking_fallback_text(&self) -> Option { + let Self::Thinking { + thinking, redacted, .. + } = self + else { + return None; + }; + + if !thinking.is_empty() { + Some(thinking.clone()) + } else if *redacted { + Some("[redacted reasoning]".to_string()) + } else { + Some("[reasoning unavailable]".to_string()) + } + } + + pub fn image_bytes(media_type: impl Into, data: impl Into>) -> Self { + Self::Image { + source: ImageSource::bytes(media_type, data), + } + } + + pub fn image_url(url: impl Into) -> Self { + Self::Image { + source: ImageSource::url(url), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn thinking_serde_is_externally_tagged_and_omits_empty_optional_fields() { + let block = ContentBlock::Thinking { + thinking: "private chain".to_string(), + signature: Some("opaque-signature".to_string()), + encrypted_content: None, + id: None, + provenance: Some(ReasoningProvenance { + provider: crate::ProviderId::new("anthropic-edge"), + model: "claude-test".to_string(), + format: ReasoningFormat::AnthropicSigned, + }), + redacted: false, + }; + + let json = serde_json::to_value(&block).expect("thinking block should serialize"); + assert_eq!(json["Thinking"]["thinking"], "private chain"); + assert_eq!(json["Thinking"]["signature"], "opaque-signature"); + assert_eq!(json["Thinking"]["provenance"]["provider"], "anthropic-edge"); + assert_eq!(json["Thinking"]["provenance"]["format"], "anthropic_signed"); + assert!(json["Thinking"].get("encrypted_content").is_none()); + assert!(json["Thinking"].get("id").is_none()); + assert!(json["Thinking"].get("redacted").is_none()); + assert_eq!( + serde_json::from_value::(json).expect("thinking block should load"), + block + ); + } + + #[test] + fn thinking_serde_defaults_omitted_payload_fields() { + let block: ContentBlock = serde_json::from_value(serde_json::json!({ + "Thinking": { + "provenance": { + "provider": "anthropic", + "model": "claude-test", + "format": "anthropic_signed" + } + } + })) + .expect("omitted thinking payload fields should default"); + + assert_eq!( + block, + ContentBlock::Thinking { + thinking: String::new(), + signature: None, + encrypted_content: None, + id: None, + provenance: Some(ReasoningProvenance { + provider: crate::ProviderId::new("anthropic"), + model: "claude-test".to_string(), + format: ReasoningFormat::AnthropicSigned, + }), + redacted: false, + } + ); + } + + #[test] + fn pre_thinking_content_block_json_still_deserializes_unchanged() { + let json = serde_json::json!({"Text":{"text":"legacy"}}); + + assert_eq!( + serde_json::from_value::(json).expect("legacy block should load"), + ContentBlock::text("legacy") + ); + } +} + +/// Provider-neutral chat message content. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct Message { + pub role: Role, + pub content: Vec, +} + +impl Message { + pub fn user(content: ContentBlock) -> Self { + Self { + role: Role::User, + content: vec![content], + } + } + + pub fn assistant(content: ContentBlock) -> Self { + Self { + role: Role::Assistant, + content: vec![content], + } + } + + pub fn unknown(role: impl Into, content: ContentBlock) -> Self { + Self { + role: Role::Unknown(role.into()), + content: vec![content], + } + } + + pub fn text(&self) -> String { + self.content + .iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text.as_str()), + _ => None, + }) + .collect::>() + .join("") + } +} + +/// Provider-neutral tool choice hint passed to model APIs. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolChoice { + #[default] + Auto, + Any, + Tool { + name: String, + }, +} diff --git a/vendor/mentra-provider/src/registry.rs b/vendor/mentra-provider/src/registry.rs new file mode 100644 index 0000000..f4cf901 --- /dev/null +++ b/vendor/mentra-provider/src/registry.rs @@ -0,0 +1,215 @@ +use async_trait::async_trait; +use std::collections::HashMap; +use std::sync::Arc; + +use crate::definition::ProviderDefinition; +use crate::definition::ProviderDescriptor; +use crate::definition::ProviderId; +use crate::error::ProviderError; +use crate::model::ModelInfo; +use crate::request::CompactionRequest; +use crate::request::MemorySummarizeRequest; +use crate::request::Request; +use crate::response::CompactionResponse; +use crate::response::MemorySummarizeResponse; +use crate::response::Response; +use crate::response::collect_response_from_stream; +use crate::stream::ProviderEventStream; + +/// Lists models available from a provider. +#[async_trait] +pub trait ModelCatalog: Send + Sync { + async fn list_models(&self) -> Result, ProviderError>; +} + +/// Creates a provider session on demand. +#[async_trait] +pub trait ProviderSessionFactory: Send + Sync { + async fn create_session(&self) -> Result, ProviderError>; +} + +/// Transport-neutral session used to stream model responses. +#[async_trait] +pub trait ProviderSession: Send + Sync { + async fn stream(&self, request: Request<'_>) -> Result; + + async fn send(&self, request: Request<'_>) -> Result { + collect_response_from_stream(self.stream(request).await?).await + } + + async fn compact( + &self, + _request: CompactionRequest<'_>, + ) -> Result { + Err(ProviderError::UnsupportedCapability( + "history_compaction".to_string(), + )) + } + + async fn summarize_memories( + &self, + _request: MemorySummarizeRequest<'_>, + ) -> Result { + Err(ProviderError::UnsupportedCapability( + "memory_summarization".to_string(), + )) + } +} + +/// Transport-neutral provider registration interface. +#[async_trait] +pub trait Provider: ModelCatalog + ProviderSessionFactory { + fn definition(&self) -> ProviderDefinition; + + fn descriptor(&self) -> ProviderDescriptor { + self.definition().descriptor + } + + async fn stream(&self, request: Request<'_>) -> Result { + self.create_session().await?.stream(request).await + } + + async fn send(&self, request: Request<'_>) -> Result { + collect_response_from_stream(self.stream(request).await?).await + } + + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + self.create_session().await?.compact(request).await + } + + async fn summarize_memories( + &self, + request: MemorySummarizeRequest<'_>, + ) -> Result { + self.create_session() + .await? + .summarize_memories(request) + .await + } +} + +pub use Provider as RegisteredProvider; + +#[derive(Default)] +pub struct ProviderRegistry { + default_provider: Option, + providers: HashMap>, +} + +impl ProviderRegistry { + pub fn register_provider_instance

(&mut self, provider: P) + where + P: Provider + 'static, + { + let definition = provider.definition(); + let id = definition.descriptor.id.clone(); + + if self.default_provider.is_none() { + self.default_provider = Some(id.clone()); + } + + self.providers.insert(id, Arc::new(provider)); + } + + pub fn get_provider(&self, id: Option<&ProviderId>) -> Option> { + match id { + Some(id) => self.providers.get(id).cloned(), + None => self + .default_provider + .as_ref() + .and_then(|id| self.providers.get(id).cloned()), + } + } + + pub fn definitions(&self) -> Vec { + self.providers + .values() + .map(|provider| provider.definition()) + .collect() + } + + pub fn descriptors(&self) -> Vec { + self.providers + .values() + .map(|provider| provider.descriptor()) + .collect() + } + + pub fn is_empty(&self) -> bool { + self.providers.is_empty() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use async_trait::async_trait; + use tokio::sync::mpsc; + + #[derive(Clone)] + struct TestProvider { + definition: ProviderDefinition, + models: Vec, + } + + struct TestSession; + + #[async_trait] + impl ModelCatalog for TestProvider { + async fn list_models(&self) -> Result, ProviderError> { + Ok(self.models.clone()) + } + } + + #[async_trait] + impl ProviderSessionFactory for TestProvider { + async fn create_session(&self) -> Result, ProviderError> { + Ok(Box::new(TestSession)) + } + } + + #[async_trait] + impl Provider for TestProvider { + fn definition(&self) -> ProviderDefinition { + self.definition.clone() + } + } + + #[async_trait] + impl ProviderSession for TestSession { + async fn stream( + &self, + _request: Request<'_>, + ) -> Result { + let (_tx, rx) = mpsc::unbounded_channel(); + Ok(rx) + } + } + + #[tokio::test] + async fn registry_returns_registered_provider_descriptors() { + let mut registry = ProviderRegistry::default(); + let provider = TestProvider { + definition: ProviderDefinition::new("test-provider"), + models: vec![ModelInfo::new("model-1", "test-provider")], + }; + + registry.register_provider_instance(provider); + + assert_eq!(registry.descriptors().len(), 1); + assert_eq!(registry.definitions().len(), 1); + assert_eq!( + registry + .get_provider(None) + .expect("provider should exist") + .definition() + .descriptor + .id + .as_str(), + "test-provider" + ); + } +} diff --git a/vendor/mentra-provider/src/request.rs b/vendor/mentra-provider/src/request.rs new file mode 100644 index 0000000..f2f8720 --- /dev/null +++ b/vendor/mentra-provider/src/request.rs @@ -0,0 +1,541 @@ +use serde::Deserialize; +use serde::Serialize; +use std::borrow::Cow; +use std::collections::BTreeMap; + +use crate::ContentBlock; +use crate::Message; +use crate::ProviderError; +use crate::model::ToolChoice; +use crate::tool::ToolSpec; + +/// Provider-neutral reasoning controls supported across multiple providers. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ReasoningOptions { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub effort: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub summary: Option, +} + +/// Shared reasoning effort levels supported by Mentra's public API. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +#[non_exhaustive] +pub enum ReasoningEffort { + Low, + Medium, + High, + #[serde(rename = "xhigh")] + XHigh, + Max, +} + +/// Shared reasoning summary levels used by Responses-family providers. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ReasoningSummary { + Auto, + Concise, + Detailed, +} + +/// Provider-neutral tool search behavior requested for a model call. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum ToolSearchMode { + #[default] + Disabled, + Hosted, +} + +/// Responses-compatible verbosity controls for text output. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "lowercase")] +pub enum ResponsesVerbosity { + Low, + #[default] + Medium, + High, +} + +/// Transport-level request compression supported by Responses-family HTTP calls. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "lowercase")] +pub enum ResponsesRequestCompression { + #[default] + None, + Zstd, +} + +/// Provider-side conversation state strategy for Responses-family providers. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum ResponsesStateMode { + /// Send the complete local transcript and do not attach provider-side state. + ReplayOnly, + /// Keep local replay as the source of truth while opportunistically chaining provider state. + /// + /// An unknown HTTP endpoint may receive one capability probe before Hybrid + /// learns that `previous_response_id` is unsupported. Hosts that already + /// know this can disable the probe on [`crate::responses::ResponsesProvider`]. + #[default] + Hybrid, + /// Require provider-side state chaining once a previous response id is available. + Stateful, +} + +impl ResponsesStateMode { + pub fn uses_provider_state(self) -> bool { + matches!(self, Self::Hybrid | Self::Stateful) + } +} + +/// Streaming transport used by Responses-family providers. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum ResponsesTransport { + /// Standard HTTP request with Server-Sent Events response streaming. + #[default] + HttpSse, + /// Long-lived WebSocket connection driven by `response.create` frames. + WebSocket, +} + +/// Responses-compatible format discriminator for structured text output. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum ResponsesTextFormatType { + #[default] + JsonSchema, +} + +/// Structured text output format controls for Responses-family providers. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct ResponsesTextFormat { + #[serde(default)] + pub r#type: ResponsesTextFormatType, + #[serde(default)] + pub strict: bool, + pub schema: serde_json::Value, + pub name: String, +} + +/// Responses-compatible text controls combining verbosity and output schemas. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct ResponsesTextControls { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub verbosity: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub format: Option, +} + +/// Shared Responses-family request options. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ResponsesRequestOptions { + #[serde(default)] + pub parallel_tool_calls: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub previous_response_id: Option, + #[serde(default)] + pub state_mode: ResponsesStateMode, + #[serde(default)] + pub transport: ResponsesTransport, + #[serde(default)] + pub store: Option, + #[serde(default)] + pub stream: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub include: Vec, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub service_tier: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub prompt_cache_key: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub text: Option, + #[serde(default)] + pub compression: ResponsesRequestCompression, +} + +impl Default for ResponsesRequestOptions { + fn default() -> Self { + Self { + parallel_tool_calls: None, + previous_response_id: None, + state_mode: ResponsesStateMode::Hybrid, + transport: ResponsesTransport::HttpSse, + store: None, + stream: Some(true), + include: Vec::new(), + service_tier: None, + prompt_cache_key: None, + text: None, + compression: ResponsesRequestCompression::None, + } + } +} + +/// Anthropic-specific request options. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct AnthropicRequestOptions { + #[serde(default)] + pub disable_parallel_tool_use: Option, +} + +/// Gemini-specific request options. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct GeminiRequestOptions { + #[serde(default)] + pub thoughts: Option, +} + +/// Provider-neutral session metadata and affinity hints. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct SessionRequestOptions { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub sticky_turn_state: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub turn_metadata: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub subagent: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub prefer_connection_reuse: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub session_affinity: Option, + #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] + pub extra_headers: BTreeMap, +} + +/// Provider-specific request options that should be forwarded on the wire. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +pub struct ProviderRequestOptions { + #[serde(default)] + pub tool_search_mode: ToolSearchMode, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning: Option, + #[serde(default)] + pub responses: ResponsesRequestOptions, + #[serde(default)] + pub anthropic: AnthropicRequestOptions, + #[serde(default)] + pub gemini: GeminiRequestOptions, + #[serde(default)] + pub session: SessionRequestOptions, +} + +/// Provider request assembled by the runtime before dispatch. +#[derive(Debug, Clone)] +pub struct Request<'a> { + pub model: Cow<'a, str>, + pub system: Option>, + pub messages: Cow<'a, [Message]>, + pub tools: Cow<'a, [ToolSpec]>, + pub tool_choice: Option, + pub temperature: Option, + pub max_output_tokens: Option, + pub metadata: Cow<'a, BTreeMap>, + pub provider_request_options: ProviderRequestOptions, +} + +impl Request<'_> { + pub fn into_owned(self) -> Request<'static> { + Request { + model: Cow::Owned(self.model.into_owned()), + system: self.system.map(|system| Cow::Owned(system.into_owned())), + messages: Cow::Owned(self.messages.into_owned()), + tools: Cow::Owned(self.tools.into_owned()), + tool_choice: self.tool_choice, + temperature: self.temperature, + max_output_tokens: self.max_output_tokens, + metadata: Cow::Owned(self.metadata.into_owned()), + provider_request_options: self.provider_request_options, + } + } +} + +/// Provider-neutral transcript item used for history compaction. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum CompactionInputItem { + UserTurn { + content: String, + }, + AssistantTurn { + content: String, + }, + ToolExchange { + #[serde(default, skip_serializing_if = "Option::is_none")] + request: Option, + result: String, + is_error: bool, + }, + CanonicalContext { + content: String, + }, + MemoryRecall { + content: String, + }, + DelegationResult { + agent_id: String, + agent_name: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + role: Option, + status: String, + content: String, + }, + CompactionSummary { + content: String, + }, +} + +/// Provider-neutral request assembled for history compaction. +#[derive(Debug, Clone)] +pub struct CompactionRequest<'a> { + pub model: Cow<'a, str>, + pub instructions: Cow<'a, str>, + pub input: Cow<'a, [CompactionInputItem]>, + pub metadata: Cow<'a, BTreeMap>, + pub provider_request_options: ProviderRequestOptions, +} + +impl CompactionRequest<'_> { + /// Converts a compaction request into an ordinary model request. + pub fn into_model_request(self) -> Result, ProviderError> { + let input_json = + serde_json::to_string(self.input.as_ref()).map_err(ProviderError::Serialize)?; + + Ok(Request { + model: Cow::Owned(self.model.into_owned()), + system: Some(Cow::Owned(self.instructions.into_owned())), + messages: Cow::Owned(vec![Message::user(ContentBlock::text(format!( + "Compaction input JSON:\n{input_json}" + )))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(self.metadata.into_owned()), + provider_request_options: self.provider_request_options, + }) + } +} + +/// Canonical raw memory payload used by memory summarization requests. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct RawMemory { + pub id: String, + pub metadata: RawMemoryMetadata, + pub items: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct RawMemoryMetadata { + pub source_path: String, +} + +/// Provider-neutral request assembled for trace memory summarization. +#[derive(Debug, Clone)] +pub struct MemorySummarizeRequest<'a> { + pub model: Cow<'a, str>, + pub raw_memories: Cow<'a, [RawMemory]>, + pub reasoning: Option, + pub metadata: Cow<'a, BTreeMap>, + pub provider_request_options: ProviderRequestOptions, +} + +impl MemorySummarizeRequest<'_> { + /// Converts a memory summarize request into an ordinary model request. + pub fn into_model_request(self) -> Result, ProviderError> { + let raw_memories_json = + serde_json::to_string(self.raw_memories.as_ref()).map_err(ProviderError::Serialize)?; + + Ok(Request { + model: Cow::Owned(self.model.into_owned()), + system: Some(Cow::Borrowed(MEMORY_SUMMARIZE_SYSTEM_PROMPT)), + messages: Cow::Owned(vec![Message::user(ContentBlock::text(format!( + "Memory summarize input JSON:\n{raw_memories_json}" + )))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(self.metadata.into_owned()), + provider_request_options: ProviderRequestOptions { + reasoning: self.reasoning, + ..self.provider_request_options + }, + }) + } +} + +const MEMORY_SUMMARIZE_SYSTEM_PROMPT: &str = concat!( + "You summarize trace memories for Codex.\n", + "Return valid JSON only.\n", + "The output must be a JSON array with one object per input trace, in the same order.\n", + "Each object must have exactly these string fields: `raw_memory` and `memory_summary`.\n", + "`raw_memory` should be a concrete, detailed summary of the trace contents.\n", + "`memory_summary` should be a shorter durable takeaway focused on reusable context.\n", + "Use empty strings when information is unavailable.\n", + "Do not include markdown fences or extra commentary.\n", +); + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::Value; + + #[test] + fn reasoning_effort_uses_the_five_public_spellings() { + for (effort, expected) in [ + (ReasoningEffort::Low, "low"), + (ReasoningEffort::Medium, "medium"), + (ReasoningEffort::High, "high"), + (ReasoningEffort::XHigh, "xhigh"), + (ReasoningEffort::Max, "max"), + ] { + let serialized = serde_json::to_string(&effort).expect("effort should serialize"); + assert_eq!(serialized, format!("\"{expected}\"")); + assert_eq!( + serde_json::from_str::(&serialized) + .expect("effort should deserialize"), + effort + ); + } + } + + #[test] + fn compaction_request_into_model_request_serializes_input_as_prompt_text() { + let request = CompactionRequest { + model: Cow::Borrowed("gpt-5"), + instructions: Cow::Borrowed("Summarize the transcript."), + input: Cow::Owned(vec![ + CompactionInputItem::UserTurn { + content: "hello".to_string(), + }, + CompactionInputItem::AssistantTurn { + content: "world".to_string(), + }, + ]), + metadata: Cow::Owned(BTreeMap::from([("scope".to_string(), "test".to_string())])), + provider_request_options: ProviderRequestOptions { + session: SessionRequestOptions { + sticky_turn_state: Some("sticky".to_string()), + turn_metadata: None, + subagent: Some("compact".to_string()), + prefer_connection_reuse: Some(true), + session_affinity: None, + extra_headers: BTreeMap::new(), + }, + ..ProviderRequestOptions::default() + }, + }; + + let model_request = request + .into_model_request() + .expect("compaction request should convert"); + + assert_eq!(model_request.model.as_ref(), "gpt-5"); + assert_eq!( + model_request.system.as_deref(), + Some("Summarize the transcript.") + ); + assert_eq!(model_request.metadata["scope"], "test"); + assert_eq!( + model_request + .provider_request_options + .session + .sticky_turn_state + .as_deref(), + Some("sticky") + ); + assert_eq!( + model_request + .provider_request_options + .session + .subagent + .as_deref(), + Some("compact") + ); + assert_eq!(model_request.messages.len(), 1); + + let prompt = model_request.messages[0].text(); + assert!(prompt.starts_with("Compaction input JSON:\n")); + let payload = prompt + .strip_prefix("Compaction input JSON:\n") + .expect("prompt should contain the compaction prefix"); + let input: Vec = serde_json::from_str(payload).expect("prompt should be json"); + assert_eq!(input[0]["type"], "user_turn"); + assert_eq!(input[0]["content"], "hello"); + assert_eq!(input[1]["type"], "assistant_turn"); + assert_eq!(input[1]["content"], "world"); + } + + #[test] + fn memory_summarize_request_into_model_request_serializes_input_as_prompt_text() { + let request = MemorySummarizeRequest { + model: Cow::Borrowed("gpt-5"), + raw_memories: Cow::Owned(vec![RawMemory { + id: "memory-1".to_string(), + metadata: RawMemoryMetadata { + source_path: "/tmp/trace.jsonl".to_string(), + }, + items: vec![serde_json::json!({"type":"message","role":"user"})], + }]), + reasoning: Some(ReasoningOptions { + effort: Some(ReasoningEffort::Medium), + summary: None, + }), + metadata: Cow::Owned(BTreeMap::from([("scope".to_string(), "test".to_string())])), + provider_request_options: ProviderRequestOptions { + session: SessionRequestOptions { + sticky_turn_state: None, + turn_metadata: Some("{\"turn_id\":\"t1\"}".to_string()), + subagent: None, + prefer_connection_reuse: Some(true), + session_affinity: Some("thread-1".to_string()), + extra_headers: BTreeMap::new(), + }, + ..ProviderRequestOptions::default() + }, + }; + + let model_request = request + .into_model_request() + .expect("memory summarize request should convert"); + + assert_eq!(model_request.model.as_ref(), "gpt-5"); + assert_eq!( + model_request.system.as_deref(), + Some(MEMORY_SUMMARIZE_SYSTEM_PROMPT) + ); + assert_eq!(model_request.metadata["scope"], "test"); + assert_eq!( + model_request + .provider_request_options + .session + .turn_metadata + .as_deref(), + Some("{\"turn_id\":\"t1\"}") + ); + assert_eq!( + model_request + .provider_request_options + .reasoning + .as_ref() + .expect("reasoning options") + .effort, + Some(ReasoningEffort::Medium) + ); + assert_eq!(model_request.messages.len(), 1); + + let prompt = model_request.messages[0].text(); + assert!(prompt.starts_with("Memory summarize input JSON:\n")); + let payload = prompt + .strip_prefix("Memory summarize input JSON:\n") + .expect("prompt should contain the memory summarize prefix"); + let input: Vec = serde_json::from_str(payload).expect("prompt should be json"); + assert_eq!(input[0].id, "memory-1"); + assert_eq!(input[0].metadata.source_path, "/tmp/trace.jsonl"); + assert_eq!(input[0].items[0]["role"], "user"); + } +} diff --git a/vendor/mentra-provider/src/response.rs b/vendor/mentra-provider/src/response.rs new file mode 100644 index 0000000..86620ef --- /dev/null +++ b/vendor/mentra-provider/src/response.rs @@ -0,0 +1,913 @@ +use serde::Deserialize; +use serde::Serialize; + +use crate::error::ProviderError; +use crate::model::ContentBlock; +use crate::model::HostedToolSearchCall; +use crate::model::HostedWebSearchCall; +use crate::model::ImageGenerationCall; +use crate::model::ReasoningProvenance; +use crate::model::Role; +use crate::model::TokenUsage; +use crate::model::ToolResultContent; +use crate::request::CompactionInputItem; +use crate::stream::ContentBlockDelta; +use crate::stream::ContentBlockStart; +use crate::stream::ProviderEvent; +use crate::stream::ProviderEventStream; + +/// A complete response collected from a provider stream. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct Response { + pub id: String, + pub model: String, + pub role: Role, + pub content: Vec, + pub stop_reason: Option, + pub usage: Option, +} + +/// A complete history-compaction response collected from a provider. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct CompactionResponse { + pub output: Vec, +} + +/// A complete memory summarize response collected from a provider. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemorySummarizeResponse { + pub output: Vec, +} + +/// One summary object for a single raw memory input. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct MemorySummarizeOutput { + #[serde(rename = "trace_summary", alias = "raw_memory")] + pub raw_memory: String, + pub memory_summary: String, +} + +/// Rebuilds a full response from a provider event stream. +pub async fn collect_response_from_stream( + mut stream: ProviderEventStream, +) -> Result { + let mut builder = StreamingResponseBuilder::default(); + + while let Some(event) = stream.recv().await { + builder.apply(event?)?; + } + + builder.build() +} + +/// Converts a response into a provider event stream. +pub fn provider_event_stream_from_response(response: Response) -> ProviderEventStream { + let events = response.into_provider_events(); + let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); + + for event in events { + if tx.send(Ok(event)).is_err() { + break; + } + } + + rx +} + +impl Response { + pub fn into_provider_events(self) -> Vec { + let mut events = vec![ProviderEvent::MessageStarted { + id: self.id, + model: self.model, + role: self.role, + }]; + + for (index, block) in self.content.into_iter().enumerate() { + events.extend(block.into_provider_events(index)); + } + + events.push(ProviderEvent::MessageDelta { + stop_reason: self.stop_reason, + usage: self.usage, + }); + events.push(ProviderEvent::MessageStopped); + events + } + + /// Converts a normal response into a compaction response by collecting its text output. + pub fn into_compaction_response(self) -> CompactionResponse { + let text = self + .content + .into_iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text), + _ => None, + }) + .collect::>() + .join("\n") + .trim() + .to_string(); + + CompactionResponse::from_text(text) + } + + /// Converts a normal response into a memory summarize response by parsing its text output. + pub fn into_memory_summarize_response(self) -> Result { + let text = self + .content + .into_iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text), + _ => None, + }) + .collect::>() + .join("\n") + .trim() + .to_string(); + + MemorySummarizeResponse::from_text(&text) + } +} + +impl CompactionResponse { + /// Wraps a text summary into a single compaction-summary item. + pub fn from_text(text: impl Into) -> Self { + Self { + output: vec![CompactionInputItem::CompactionSummary { + content: text.into(), + }], + } + } +} + +impl MemorySummarizeResponse { + pub fn from_text(text: &str) -> Result { + let text = strip_markdown_code_fence(text); + + if let Ok(output) = serde_json::from_str::>(text) { + return Ok(Self { output }); + } + + if let Ok(response) = serde_json::from_str::(text) { + return Ok(response); + } + + Err(ProviderError::InvalidResponse( + "failed to parse memory summarize response json".to_string(), + )) + } +} + +fn strip_markdown_code_fence(text: &str) -> &str { + let trimmed = text.trim(); + let Some(rest) = trimmed.strip_prefix("```") else { + return trimmed; + }; + let rest = rest + .strip_prefix("json") + .or_else(|| rest.strip_prefix("JSON")) + .unwrap_or(rest); + rest.trim() + .strip_suffix("```") + .map(str::trim) + .unwrap_or(trimmed) +} + +#[derive(Default)] +struct StreamingResponseBuilder { + id: Option, + model: Option, + role: Option, + blocks: std::collections::BTreeMap, + stop_reason: Option, + usage: Option, + stopped: bool, +} + +impl StreamingResponseBuilder { + fn apply(&mut self, event: ProviderEvent) -> Result<(), ProviderError> { + match event { + ProviderEvent::ResponseHeaders(_) | ProviderEvent::ResponseCreated => {} + ProviderEvent::MessageStarted { id, model, role } => { + self.id = Some(id); + self.model = Some(model); + self.role = Some(role); + } + ProviderEvent::ContentBlockStarted { index, kind } => { + self.blocks.insert(index, StreamingContentBlock::from(kind)); + } + ProviderEvent::ContentBlockDelta { index, delta } => { + let block = self.blocks.get_mut(&index).ok_or_else(|| { + ProviderError::MalformedStream(format!( + "content block delta received before start for index {index}" + )) + })?; + block.apply_delta(delta)?; + } + ProviderEvent::ContentBlockStopped { index } => { + let block = self.blocks.get_mut(&index).ok_or_else(|| { + ProviderError::MalformedStream(format!( + "content block stop received before start for index {index}" + )) + })?; + block.mark_complete(); + } + ProviderEvent::MessageDelta { stop_reason, usage } => { + self.stop_reason = stop_reason; + self.usage = usage; + } + ProviderEvent::ReasoningSummaryDelta { .. } + | ProviderEvent::ReasoningContentDelta { .. } + | ProviderEvent::ReasoningSummaryPartAdded { .. } => {} + ProviderEvent::MessageStopped => { + self.stopped = true; + } + } + + Ok(()) + } + + fn build(self) -> Result { + if !self.stopped { + return Err(ProviderError::MalformedStream( + "message stream ended before MessageStopped".to_string(), + )); + } + + let id = self + .id + .ok_or_else(|| ProviderError::MalformedStream("missing message id".to_string()))?; + let model = self + .model + .ok_or_else(|| ProviderError::MalformedStream("missing model id".to_string()))?; + let role = self + .role + .ok_or_else(|| ProviderError::MalformedStream("missing message role".to_string()))?; + let mut content = Vec::with_capacity(self.blocks.len()); + + for (index, block) in self.blocks { + if !block.is_complete() { + return Err(ProviderError::MalformedStream(format!( + "content block {index} did not complete" + ))); + } + content.push(block.try_into_content_block()?); + } + + Ok(Response { + id, + model, + role, + content, + stop_reason: self.stop_reason, + usage: self.usage, + }) + } +} + +enum StreamingContentBlock { + Text { + text: String, + complete: bool, + }, + Thinking { + thinking: String, + signature: Option, + encrypted_content: Option, + id: Option, + provenance: Option, + redacted: bool, + complete: bool, + }, + Image { + source: crate::model::ImageSource, + complete: bool, + }, + ToolUse { + id: String, + name: String, + input_json: String, + complete: bool, + }, + ToolResult { + tool_use_id: String, + content: Option, + is_error: bool, + complete: bool, + }, + HostedToolSearch { + call: HostedToolSearchCall, + complete: bool, + }, + HostedWebSearch { + call: HostedWebSearchCall, + complete: bool, + }, + ImageGeneration { + call: ImageGenerationCall, + complete: bool, + }, +} + +impl StreamingContentBlock { + fn apply_delta(&mut self, delta: ContentBlockDelta) -> Result<(), ProviderError> { + match (self, delta) { + (StreamingContentBlock::Text { text, .. }, ContentBlockDelta::Text(delta)) => { + text.push_str(&delta); + Ok(()) + } + ( + StreamingContentBlock::Thinking { thinking, .. }, + ContentBlockDelta::ThinkingText(delta), + ) => { + thinking.push_str(&delta); + Ok(()) + } + ( + StreamingContentBlock::Thinking { signature, .. }, + ContentBlockDelta::ThinkingSignature(delta), + ) => { + signature.get_or_insert_with(String::new).push_str(&delta); + Ok(()) + } + ( + StreamingContentBlock::Thinking { + encrypted_content, .. + }, + ContentBlockDelta::ThinkingEncryptedContent(value), + ) => { + *encrypted_content = Some(value); + Ok(()) + } + ( + StreamingContentBlock::ToolUse { input_json, .. }, + ContentBlockDelta::ToolUseInputJson(delta), + ) => { + input_json.push_str(&delta); + Ok(()) + } + ( + StreamingContentBlock::ToolResult { content, .. }, + ContentBlockDelta::ToolResultContent(delta), + ) => { + *content = Some(match (content.take(), delta) { + (_, ToolResultContent::Structured(value)) => { + ToolResultContent::Structured(value) + } + ( + Some(ToolResultContent::Structured(value)), + ToolResultContent::Text(delta), + ) => ToolResultContent::Structured(merge_structured_text(value, delta)), + (Some(ToolResultContent::Text(existing)), ToolResultContent::Text(delta)) => { + ToolResultContent::Text(format!("{existing}{delta}")) + } + (None, delta) => delta, + }); + Ok(()) + } + ( + StreamingContentBlock::HostedToolSearch { call, .. }, + ContentBlockDelta::HostedToolSearchQuery(delta), + ) => { + let query = call.query.get_or_insert_with(String::new); + query.push_str(&delta); + Ok(()) + } + ( + StreamingContentBlock::HostedToolSearch { call, .. }, + ContentBlockDelta::HostedToolSearchStatus(status), + ) => { + call.status = Some(status); + Ok(()) + } + ( + StreamingContentBlock::HostedWebSearch { call, .. }, + ContentBlockDelta::HostedWebSearchAction(action), + ) => { + call.action = Some(action); + Ok(()) + } + ( + StreamingContentBlock::HostedWebSearch { call, .. }, + ContentBlockDelta::HostedWebSearchStatus(status), + ) => { + call.status = Some(status); + Ok(()) + } + ( + StreamingContentBlock::ImageGeneration { call, .. }, + ContentBlockDelta::ImageGenerationStatus(status), + ) => { + call.status = status; + Ok(()) + } + ( + StreamingContentBlock::ImageGeneration { call, .. }, + ContentBlockDelta::ImageGenerationRevisedPrompt(delta), + ) => { + let revised_prompt = call.revised_prompt.get_or_insert_with(String::new); + revised_prompt.push_str(&delta); + Ok(()) + } + ( + StreamingContentBlock::ImageGeneration { call, .. }, + ContentBlockDelta::ImageGenerationResult(result), + ) => { + call.result = Some(result); + Ok(()) + } + (block, delta) => Err(ProviderError::MalformedStream(format!( + "delta {delta:?} is not valid for block {}", + block.kind_name() + ))), + } + } + + fn mark_complete(&mut self) { + match self { + StreamingContentBlock::Text { complete, .. } + | StreamingContentBlock::Thinking { complete, .. } + | StreamingContentBlock::Image { complete, .. } + | StreamingContentBlock::ToolUse { complete, .. } + | StreamingContentBlock::ToolResult { complete, .. } + | StreamingContentBlock::HostedToolSearch { complete, .. } + | StreamingContentBlock::HostedWebSearch { complete, .. } + | StreamingContentBlock::ImageGeneration { complete, .. } => *complete = true, + } + } + + fn is_complete(&self) -> bool { + match self { + StreamingContentBlock::Text { complete, .. } + | StreamingContentBlock::Thinking { complete, .. } + | StreamingContentBlock::Image { complete, .. } + | StreamingContentBlock::ToolUse { complete, .. } + | StreamingContentBlock::ToolResult { complete, .. } + | StreamingContentBlock::HostedToolSearch { complete, .. } + | StreamingContentBlock::HostedWebSearch { complete, .. } + | StreamingContentBlock::ImageGeneration { complete, .. } => *complete, + } + } + + fn try_into_content_block(self) -> Result { + match self { + StreamingContentBlock::Text { text, .. } => Ok(ContentBlock::Text { text }), + StreamingContentBlock::Thinking { + thinking, + signature, + encrypted_content, + id, + provenance, + redacted, + .. + } => Ok(ContentBlock::Thinking { + thinking, + signature, + encrypted_content, + id, + provenance, + redacted, + }), + StreamingContentBlock::Image { source, .. } => Ok(ContentBlock::Image { source }), + StreamingContentBlock::ToolUse { + id, + name, + input_json, + .. + } => Ok(ContentBlock::ToolUse { + id, + name, + input: serde_json::from_str(&input_json).map_err(ProviderError::Deserialize)?, + }), + StreamingContentBlock::ToolResult { + tool_use_id, + content, + is_error, + .. + } => Ok(ContentBlock::ToolResult { + tool_use_id, + content: content.unwrap_or_else(|| ToolResultContent::text(String::new())), + is_error, + }), + StreamingContentBlock::HostedToolSearch { call, .. } => { + Ok(ContentBlock::HostedToolSearch { call }) + } + StreamingContentBlock::HostedWebSearch { call, .. } => { + Ok(ContentBlock::HostedWebSearch { call }) + } + StreamingContentBlock::ImageGeneration { call, .. } => { + Ok(ContentBlock::ImageGeneration { call }) + } + } + } + + fn kind_name(&self) -> &'static str { + match self { + StreamingContentBlock::Text { .. } => "text", + StreamingContentBlock::Thinking { .. } => "thinking", + StreamingContentBlock::Image { .. } => "image", + StreamingContentBlock::ToolUse { .. } => "tool_use", + StreamingContentBlock::ToolResult { .. } => "tool_result", + StreamingContentBlock::HostedToolSearch { .. } => "hosted_tool_search", + StreamingContentBlock::HostedWebSearch { .. } => "hosted_web_search", + StreamingContentBlock::ImageGeneration { .. } => "image_generation", + } + } +} + +impl From for StreamingContentBlock { + fn from(value: ContentBlockStart) -> Self { + match value { + ContentBlockStart::Text => StreamingContentBlock::Text { + text: String::new(), + complete: false, + }, + ContentBlockStart::Thinking { + encrypted_content, + id, + provenance, + redacted, + } => StreamingContentBlock::Thinking { + thinking: String::new(), + signature: None, + encrypted_content, + id, + provenance, + redacted, + complete: false, + }, + ContentBlockStart::Image { source } => StreamingContentBlock::Image { + source, + complete: false, + }, + ContentBlockStart::ToolUse { id, name } => StreamingContentBlock::ToolUse { + id, + name, + input_json: String::new(), + complete: false, + }, + ContentBlockStart::ToolResult { + tool_use_id, + is_error, + content, + } => StreamingContentBlock::ToolResult { + tool_use_id, + content, + is_error, + complete: false, + }, + ContentBlockStart::HostedToolSearch { call } => { + StreamingContentBlock::HostedToolSearch { + call, + complete: false, + } + } + ContentBlockStart::HostedWebSearch { call } => StreamingContentBlock::HostedWebSearch { + call, + complete: false, + }, + ContentBlockStart::ImageGeneration { call } => StreamingContentBlock::ImageGeneration { + call, + complete: false, + }, + } + } +} + +impl ContentBlock { + fn into_provider_events(self, index: usize) -> Vec { + match self { + ContentBlock::Text { text } => { + let mut events = vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Text, + }]; + if !text.is_empty() { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::Text(text), + }); + } + events.push(ProviderEvent::ContentBlockStopped { index }); + events + } + ContentBlock::Thinking { + thinking, + signature, + encrypted_content, + id, + provenance, + redacted, + } => { + let mut events = vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Thinking { + encrypted_content, + id, + provenance, + redacted, + }, + }]; + if !thinking.is_empty() { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ThinkingText(thinking), + }); + } + if let Some(signature) = signature { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ThinkingSignature(signature), + }); + } + events.push(ProviderEvent::ContentBlockStopped { index }); + events + } + ContentBlock::Image { source } => vec![ + ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::Image { source }, + }, + ProviderEvent::ContentBlockStopped { index }, + ], + ContentBlock::ToolUse { id, name, input } => { + let mut events = vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolUse { id, name }, + }]; + let input_json = input.to_string(); + if !input_json.is_empty() { + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ToolUseInputJson(input_json), + }); + } + events.push(ProviderEvent::ContentBlockStopped { index }); + events + } + ContentBlock::ToolResult { + tool_use_id, + content, + is_error, + } => { + let mut events = vec![ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ToolResult { + tool_use_id, + is_error, + content: Some(content.clone()), + }, + }]; + events.push(ProviderEvent::ContentBlockStopped { index }); + events + } + ContentBlock::HostedToolSearch { call } => { + vec![ + ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::HostedToolSearch { call }, + }, + ProviderEvent::ContentBlockStopped { index }, + ] + } + ContentBlock::HostedWebSearch { call } => { + vec![ + ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::HostedWebSearch { call }, + }, + ProviderEvent::ContentBlockStopped { index }, + ] + } + ContentBlock::ImageGeneration { call } => { + vec![ + ProviderEvent::ContentBlockStarted { + index, + kind: ContentBlockStart::ImageGeneration { call }, + }, + ProviderEvent::ContentBlockStopped { index }, + ] + } + } + } +} + +fn merge_structured_text(value: serde_json::Value, delta: String) -> serde_json::Value { + match value { + serde_json::Value::String(existing) => { + serde_json::Value::String(format!("{existing}{delta}")) + } + other => other, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn response_round_trip_preserves_usage() { + let response = Response { + id: "resp-1".to_string(), + model: "model".to_string(), + role: Role::Assistant, + content: vec![ContentBlock::text("hello")], + stop_reason: Some("stop".to_string()), + usage: Some(TokenUsage { + input_tokens: Some(10), + output_tokens: Some(3), + total_tokens: Some(13), + cache_read_input_tokens: Some(2), + cache_creation_input_tokens: None, + reasoning_tokens: Some(1), + thoughts_tokens: None, + tool_input_tokens: None, + }), + }; + + let rebuilt = + collect_response_from_stream(provider_event_stream_from_response(response.clone())) + .await + .expect("response should rebuild"); + + assert_eq!(rebuilt, response); + } + + #[tokio::test] + async fn response_round_trip_preserves_signed_and_redacted_thinking() { + let provenance = ReasoningProvenance { + provider: crate::ProviderId::new("anthropic-edge"), + model: "claude-test".to_string(), + format: crate::ReasoningFormat::AnthropicSigned, + }; + let response = Response { + id: "resp-thinking".to_string(), + model: "claude-test".to_string(), + role: Role::Assistant, + content: vec![ + ContentBlock::Thinking { + thinking: "private chain".to_string(), + signature: Some("opaque-signature".to_string()), + encrypted_content: None, + id: None, + provenance: Some(provenance.clone()), + redacted: false, + }, + ContentBlock::Thinking { + thinking: String::new(), + signature: Some("opaque-redacted-data".to_string()), + encrypted_content: None, + id: None, + provenance: Some(provenance), + redacted: true, + }, + ], + stop_reason: Some("end_turn".to_string()), + usage: None, + }; + + let rebuilt = + collect_response_from_stream(provider_event_stream_from_response(response.clone())) + .await + .expect("thinking response should rebuild"); + + assert_eq!(rebuilt, response); + } + + #[tokio::test] + async fn response_round_trip_preserves_hosted_actions_and_structured_tool_results() { + let response = Response { + id: "resp-2".to_string(), + model: "model".to_string(), + role: Role::Assistant, + content: vec![ + ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: ToolResultContent::Structured(serde_json::json!({"ok":true})), + is_error: false, + }, + ContentBlock::HostedToolSearch { + call: HostedToolSearchCall { + id: "search-1".to_string(), + status: Some("completed".to_string()), + query: Some("weather".to_string()), + }, + }, + ContentBlock::HostedWebSearch { + call: HostedWebSearchCall { + id: "web-1".to_string(), + status: Some("completed".to_string()), + action: Some(crate::model::WebSearchAction::Search { + query: Some("weather".to_string()), + queries: None, + }), + }, + }, + ContentBlock::ImageGeneration { + call: ImageGenerationCall { + id: "image-1".to_string(), + status: "completed".to_string(), + revised_prompt: Some("A blue square".to_string()), + result: Some(crate::model::ImageGenerationResult::ArtifactRef { + artifact_id: "artifact-1".to_string(), + }), + }, + }, + ], + stop_reason: Some("stop".to_string()), + usage: None, + }; + + let rebuilt = + collect_response_from_stream(provider_event_stream_from_response(response.clone())) + .await + .expect("response should rebuild"); + + assert_eq!(rebuilt, response); + } + + #[test] + fn response_into_compaction_response_collects_text_blocks() { + let response = Response { + id: "resp-3".to_string(), + model: "model".to_string(), + role: Role::Assistant, + content: vec![ + ContentBlock::text("first"), + ContentBlock::ToolResult { + tool_use_id: "call-1".to_string(), + content: ToolResultContent::text("ignored"), + is_error: false, + }, + ContentBlock::text("second"), + ], + stop_reason: None, + usage: None, + }; + + let compaction = response.into_compaction_response(); + + assert_eq!( + compaction.output, + vec![CompactionInputItem::CompactionSummary { + content: "first\nsecond".to_string(), + }] + ); + } + + #[test] + fn memory_summarize_response_from_text_accepts_json_array() { + let response = MemorySummarizeResponse::from_text( + r#"[{"raw_memory":"Detailed summary","memory_summary":"Short summary"}]"#, + ) + .expect("memory summarize response should parse"); + + assert_eq!( + response, + MemorySummarizeResponse { + output: vec![MemorySummarizeOutput { + raw_memory: "Detailed summary".to_string(), + memory_summary: "Short summary".to_string(), + }], + } + ); + } + + #[test] + fn memory_summarize_response_from_text_accepts_markdown_fence_and_trace_alias() { + let response = MemorySummarizeResponse::from_text( + "```json\n[{\"trace_summary\":\"Detailed summary\",\"memory_summary\":\"Short summary\"}]\n```", + ) + .expect("memory summarize response should parse"); + + assert_eq!( + response.output[0], + MemorySummarizeOutput { + raw_memory: "Detailed summary".to_string(), + memory_summary: "Short summary".to_string(), + } + ); + } + + #[test] + fn response_into_memory_summarize_response_collects_text_content() { + let response = Response { + id: "resp-4".to_string(), + model: "model".to_string(), + role: Role::Assistant, + content: vec![ContentBlock::text( + "[{\"raw_memory\":\"Detailed summary\",\"memory_summary\":\"Short summary\"}]", + )], + stop_reason: None, + usage: None, + }; + + let summarize = response + .into_memory_summarize_response() + .expect("memory summarize response should parse"); + + assert_eq!(summarize.output.len(), 1); + assert_eq!(summarize.output[0].raw_memory, "Detailed summary"); + assert_eq!(summarize.output[0].memory_summary, "Short summary"); + } +} diff --git a/vendor/mentra-provider/src/responses.rs b/vendor/mentra-provider/src/responses.rs new file mode 100644 index 0000000..3af19c7 --- /dev/null +++ b/vendor/mentra-provider/src/responses.rs @@ -0,0 +1,354 @@ +pub mod model; +pub mod session; +pub mod sse; +/// The `response.create` websocket transport. Compiled in with the +/// `responses-websocket` feature; without it, a request that selects +/// [`ResponsesTransport::WebSocket`](crate::ResponsesTransport::WebSocket) +/// fails rather than falling back to HTTP. +#[cfg(feature = "responses-websocket")] +pub mod websocket; + +use std::collections::HashMap; +use std::sync::Arc; + +use async_trait::async_trait; + +use crate::AuthScheme; +use crate::BuiltinProvider; +use crate::CredentialSource; +use crate::ModelCatalog; +use crate::ModelInfo; +use crate::ProviderCapabilities; +use crate::ProviderDefinition; +use crate::ProviderError; +use crate::ProviderSessionFactory; +use crate::RegisteredProvider; +use crate::RetryPolicy; +use crate::StaticCredentialSource; +use crate::WireApi; +use crate::embedding::EmbeddingModelInfo; +use crate::embedding::EmbeddingProvider; +use crate::embedding::EmbeddingRequest; +use crate::embedding::EmbeddingResponse; + +use self::session::ResponsesEndpointCapabilities; +use self::session::ResponsesSession; +use self::session::ResponsesSessionState; + +pub(crate) type SharedTurnState = Arc>>; + +const DEFAULT_OPENAI_BASE_URL: &str = "https://api.openai.com/"; +const DEFAULT_OPENROUTER_BASE_URL: &str = "https://openrouter.ai/api/"; + +pub fn openai(api_key: impl Into) -> ResponsesProvider { + ResponsesProvider::openai(api_key) +} + +pub fn openrouter(api_key: impl Into) -> ResponsesProvider { + ResponsesProvider::openrouter(api_key) +} + +pub fn openai_with_credential_source(credential_source: C) -> ResponsesProvider +where + C: CredentialSource + 'static, +{ + ResponsesProvider::openai_with_credential_source(credential_source) +} + +pub fn openrouter_with_credential_source(credential_source: C) -> ResponsesProvider +where + C: CredentialSource + 'static, +{ + ResponsesProvider::openrouter_with_credential_source(credential_source) +} + +/// Shared Responses-family provider implementation. +/// +/// This type owns the provider definition, credential source, client, and transport state while +/// the request mapping and SSE decoding live in the sibling modules. +#[derive(Clone)] +pub struct ResponsesProvider { + definition: ProviderDefinition, + credential_source: Arc, + client: reqwest::Client, + session_state: Arc, + endpoint_capabilities: Arc, + hybrid_http_previous_response_id: bool, +} + +impl ResponsesProvider +where + C: CredentialSource + 'static, +{ + pub fn new(definition: ProviderDefinition, credential_source: C) -> Self { + Self::with_shared_credential_source(definition, Arc::new(credential_source)) + } + + pub fn with_shared_credential_source( + definition: ProviderDefinition, + credential_source: Arc, + ) -> Self { + let client = reqwest::Client::builder() + .build() + .expect("failed to build reqwest client"); + Self { + definition, + credential_source, + client, + session_state: Arc::new(ResponsesSessionState::default()), + endpoint_capabilities: Arc::new(ResponsesEndpointCapabilities::default()), + hybrid_http_previous_response_id: true, + } + } + + pub fn definition(&self) -> &ProviderDefinition { + &self.definition + } + + /// Disables opportunistic `previous_response_id` chaining for Hybrid HTTP + /// requests made by this provider. + /// + /// Use this when the endpoint is already known not to accept that optional + /// Responses parameter. Hybrid requests retain their complete local replay, + /// so disabling the optimization avoids a known-failing discovery request. + /// Stateful requests, explicit response ids, and WebSocket transport keep + /// their existing behavior. + pub fn without_hybrid_http_previous_response_id(mut self) -> Self { + self.hybrid_http_previous_response_id = false; + self + } + + pub fn session(&self) -> ResponsesSession { + ResponsesSession::new( + self.definition.clone(), + Arc::clone(&self.credential_source), + self.client.clone(), + Arc::clone(&self.session_state), + Arc::clone(&self.endpoint_capabilities), + self.hybrid_http_previous_response_id, + ) + } + + pub fn openai_with_credential_source(credential_source: C) -> Self { + Self::with_shared_credential_source(openai_definition(), Arc::new(credential_source)) + } + + pub fn openrouter_with_credential_source(credential_source: C) -> Self { + Self::with_shared_credential_source(openrouter_definition(), Arc::new(credential_source)) + } +} + +impl ResponsesProvider { + pub fn openai(api_key: impl Into) -> Self { + Self::openai_with_credential_source(StaticCredentialSource::new(api_key)) + } + + pub fn openrouter(api_key: impl Into) -> Self { + Self::openrouter_with_credential_source(StaticCredentialSource::new(api_key)) + } +} + +pub fn openai_definition() -> ProviderDefinition { + build_definition( + BuiltinProvider::OpenAI, + "OpenAI", + "OpenAI Responses API provider", + DEFAULT_OPENAI_BASE_URL, + ) +} + +pub fn openrouter_definition() -> ProviderDefinition { + build_definition( + BuiltinProvider::OpenRouter, + "OpenRouter", + "OpenRouter Responses API provider", + DEFAULT_OPENROUTER_BASE_URL, + ) +} + +fn build_definition( + builtin: BuiltinProvider, + display_name: &str, + description: &str, + base_url: &str, +) -> ProviderDefinition { + let mut definition = ProviderDefinition::new(builtin); + definition.descriptor.display_name = Some(display_name.to_string()); + definition.descriptor.description = Some(description.to_string()); + definition.wire_api = WireApi::Responses; + definition.auth_scheme = AuthScheme::BearerToken; + definition.capabilities = ProviderCapabilities { + supports_model_listing: true, + supports_streaming: true, + supports_websockets: true, + supports_tool_calls: true, + supports_images: true, + supports_history_compaction: true, + supports_memory_summarization: true, + supports_deferred_tools: true, + supports_hosted_tool_search: true, + supports_hosted_web_search: true, + supports_image_generation: true, + supports_reasoning_effort: true, + reports_reasoning_tokens: true, + reports_thoughts_tokens: false, + supports_structured_tool_results: true, + supports_embeddings: true, + }; + definition.base_url = Some(base_url.to_string()); + definition.headers = Some(HashMap::new()); + definition.retry = RetryPolicy::default(); + definition +} + +#[async_trait] +impl ModelCatalog for ResponsesProvider +where + C: CredentialSource + 'static, +{ + async fn list_models(&self) -> Result, ProviderError> { + let credentials = self.credential_source.credentials().await?; + let request = self + .client + .get( + self.definition + .request_url_with_auth_for_path("v1/models", &credentials)?, + ) + .headers(self.definition.build_headers(&credentials)?); + + let response = request.send().await.map_err(ProviderError::Transport)?; + + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + let models = response + .json::() + .await + .map_err(ProviderError::Decode)?; + + Ok(models.into_model_info(self.definition.descriptor.id.clone())) + } +} + +#[async_trait] +impl ProviderSessionFactory for ResponsesProvider +where + C: CredentialSource + 'static, +{ + async fn create_session(&self) -> Result, ProviderError> { + Ok(Box::new(self.session())) + } +} + +#[async_trait] +impl RegisteredProvider for ResponsesProvider +where + C: CredentialSource + 'static, +{ + fn definition(&self) -> ProviderDefinition { + self.definition.clone() + } +} + +#[async_trait] +impl EmbeddingProvider for ResponsesProvider +where + C: CredentialSource + 'static, +{ + async fn embed_batch( + &self, + model: &str, + texts: &[&str], + ) -> Result { + let credentials = self.credential_source.credentials().await?; + let url = self + .definition + .request_url_with_auth_for_path("v1/embeddings", &credentials)?; + let headers = self.definition.build_headers(&credentials)?; + let body = EmbeddingRequest::batch(model, texts); + + let response = self + .client + .post(url) + .headers(headers) + .json(&body) + .send() + .await + .map_err(ProviderError::Transport)?; + + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + response + .json::() + .await + .map_err(ProviderError::Decode) + } + + fn embedding_models(&self) -> Vec { + // Available embedding models depend on the specific provider instance and its + // configuration. Callers should use the /v1/models endpoint for discovery + // rather than relying on a static list that would only be accurate for OpenAI. + vec![] + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ProviderId; + + #[test] + fn openai_preset_uses_responses_wire_api() { + let provider = openai("test-key"); + let definition = provider.definition(); + + assert_eq!( + definition.descriptor.id, + ProviderId::from(BuiltinProvider::OpenAI) + ); + assert_eq!( + definition.descriptor.display_name.as_deref(), + Some("OpenAI") + ); + assert_eq!(definition.wire_api, WireApi::Responses); + assert!(definition.capabilities.supports_websockets); + assert!(definition.capabilities.supports_history_compaction); + assert_eq!( + definition.base_url.as_deref(), + Some(DEFAULT_OPENAI_BASE_URL) + ); + } + + #[test] + fn openrouter_preset_uses_openrouter_base_url() { + let provider = openrouter("test-key"); + let definition = provider.definition(); + + assert_eq!( + definition.descriptor.id, + ProviderId::from(BuiltinProvider::OpenRouter) + ); + assert_eq!( + definition.descriptor.display_name.as_deref(), + Some("OpenRouter") + ); + assert_eq!(definition.wire_api, WireApi::Responses); + assert!(definition.capabilities.supports_history_compaction); + assert_eq!( + definition.base_url.as_deref(), + Some(DEFAULT_OPENROUTER_BASE_URL) + ); + } + + #[test] + fn disabling_hybrid_http_state_on_a_clone_does_not_reconfigure_the_original() { + let original = openai("test-key"); + let disabled = original.clone().without_hybrid_http_previous_response_id(); + + assert!(original.hybrid_http_previous_response_id); + assert!(!disabled.hybrid_http_previous_response_id); + } +} diff --git a/vendor/mentra-provider/src/responses/model.rs b/vendor/mentra-provider/src/responses/model.rs new file mode 100644 index 0000000..7247c40 --- /dev/null +++ b/vendor/mentra-provider/src/responses/model.rs @@ -0,0 +1,1516 @@ +use std::borrow::Cow; +use std::collections::BTreeMap; + +use base64::{Engine as _, engine::general_purpose::STANDARD}; +use serde::{Deserialize, Serialize}; +use time::OffsetDateTime; + +use crate::ProviderId; +use crate::{ + BuiltinProvider, ContentBlock, HostedToolSearchCall, HostedWebSearchCall, ImageGenerationCall, + ImageGenerationResult, ImageSource, Message, ModelInfo, ProviderError, ReasoningEffort, + ReasoningFormat, ReasoningOptions, ReasoningSummary, Request, ResponsesTextControls, Role, + ToolChoice, ToolSearchMode, WebSearchAction, +}; + +use crate::tool::{ProviderToolKind, ToolLoadingPolicy, ToolSpec}; + +pub(crate) fn encode_responses_tool_use_id(call_id: &str, item_id: Option<&str>) -> String { + match item_id.filter(|id| !id.is_empty()) { + Some(item_id) => format!("{call_id}|{item_id}"), + None => call_id.to_string(), + } +} + +pub(crate) fn decode_responses_tool_use_id(id: &str) -> (&str, Option<&str>) { + match id.split_once('|') { + Some((call_id, item_id)) if !item_id.is_empty() => (call_id, Some(item_id)), + _ => (id, None), + } +} + +/// Page returned by the Responses-compatible models endpoint. +#[derive(Debug, Deserialize)] +pub struct ResponsesModelsPage { + data: Vec, +} + +impl ResponsesModelsPage { + pub fn into_model_info(self, provider: ProviderId) -> Vec { + self.data + .into_iter() + .map(|model| model.into_model_info(provider.clone())) + .collect() + } +} + +#[derive(Debug, Deserialize)] +pub struct ResponsesModel { + pub id: String, + #[serde(default)] + pub name: Option, + #[serde(default)] + pub description: Option, + #[serde(default)] + pub owned_by: Option, + #[serde(default)] + pub created: Option, +} + +impl ResponsesModel { + pub fn into_model_info(self, provider: ProviderId) -> ModelInfo { + ModelInfo { + id: self.id, + provider, + display_name: self.name, + description: self + .description + .or_else(|| self.owned_by.map(|owner| format!("Owned by {owner}"))), + created_at: self + .created + .and_then(|timestamp| OffsetDateTime::from_unix_timestamp(timestamp as i64).ok()), + } + } +} + +#[derive(Debug, Serialize)] +pub struct ResponsesRequest { + model: String, + #[serde(skip_serializing_if = "Option::is_none")] + instructions: Option, + input: Vec, + #[serde(skip_serializing_if = "Vec::is_empty")] + tools: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + tool_choice: Option, + #[serde(skip_serializing_if = "Option::is_none")] + temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + max_output_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + previous_response_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + parallel_tool_calls: Option, + #[serde(skip_serializing_if = "Option::is_none")] + store: Option, + #[serde(skip_serializing_if = "Option::is_none")] + stream: Option, + #[serde(skip_serializing_if = "Vec::is_empty")] + include: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + service_tier: Option, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_key: Option, + #[serde(skip_serializing_if = "Option::is_none")] + text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + reasoning: Option, + #[serde(skip_serializing_if = "BTreeMap::is_empty")] + metadata: BTreeMap, +} + +impl<'a> TryFrom> for ResponsesRequest { + type Error = ProviderError; + + fn try_from(value: Request<'a>) -> Result { + Self::try_from_request_with_provider( + value, + "Responses", + &ProviderId::from(BuiltinProvider::OpenAI), + ) + } +} + +impl ResponsesRequest { + pub fn try_from_request<'a>( + value: Request<'a>, + provider_name: &str, + ) -> Result { + Self::try_from_request_with_provider( + value, + provider_name, + &ProviderId::from(BuiltinProvider::OpenAI), + ) + } + + pub(crate) fn try_from_request_with_provider( + value: Request<'_>, + provider_name: &str, + target_provider: &ProviderId, + ) -> Result { + let mut input = Vec::new(); + let target_model = value.model.to_string(); + let reasoning_enabled = value.provider_request_options.reasoning.is_some(); + + for message in value.messages.iter() { + input.extend(ResponsesInputItem::from_message( + message, + target_provider, + &target_model, + )?); + } + + let mut include = value.provider_request_options.responses.include.clone(); + if reasoning_enabled + && !include + .iter() + .any(|item| item == "reasoning.encrypted_content") + { + include.push("reasoning.encrypted_content".to_string()); + } + + Ok(Self { + model: value.model.into_owned(), + instructions: value.system.map(Cow::into_owned), + input, + tools: build_responses_tools( + value.tools.as_ref(), + value.tool_choice.as_ref(), + value.provider_request_options.tool_search_mode, + provider_name, + )?, + tool_choice: value.tool_choice.map(Into::into), + temperature: value.temperature, + max_output_tokens: value.max_output_tokens, + previous_response_id: value + .provider_request_options + .responses + .previous_response_id + .clone(), + parallel_tool_calls: value.provider_request_options.responses.parallel_tool_calls, + store: value.provider_request_options.responses.store, + stream: value.provider_request_options.responses.stream, + include, + service_tier: value + .provider_request_options + .responses + .service_tier + .clone(), + prompt_cache_key: value + .provider_request_options + .responses + .prompt_cache_key + .clone(), + text: value.provider_request_options.responses.text.clone(), + reasoning: value + .provider_request_options + .reasoning + .map(ResponsesReasoning::from), + metadata: value.metadata.into_owned(), + }) + } + + pub(crate) fn previous_response_id(&self) -> Option<&str> { + self.previous_response_id.as_deref() + } + + pub(crate) fn clear_previous_response_id(&mut self) { + self.previous_response_id = None; + } +} + +#[derive(Debug, Serialize)] +struct ResponsesReasoning { + #[serde(skip_serializing_if = "Option::is_none")] + effort: Option, + #[serde(skip_serializing_if = "Option::is_none")] + summary: Option, +} + +impl From for ResponsesReasoning { + fn from(value: ReasoningOptions) -> Self { + Self { + effort: value.effort.map(Into::into), + summary: value.summary.map(Into::into), + } + } +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "snake_case")] +enum ResponsesReasoningEffort { + Low, + Medium, + High, + #[serde(rename = "xhigh")] + XHigh, + Max, +} + +impl From for ResponsesReasoningEffort { + fn from(value: ReasoningEffort) -> Self { + match value { + ReasoningEffort::Low => Self::Low, + ReasoningEffort::Medium => Self::Medium, + ReasoningEffort::High => Self::High, + ReasoningEffort::XHigh => Self::XHigh, + ReasoningEffort::Max => Self::Max, + } + } +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "lowercase")] +enum ResponsesReasoningSummary { + Auto, + Concise, + Detailed, +} + +impl From for ResponsesReasoningSummary { + fn from(value: ReasoningSummary) -> Self { + match value { + ReasoningSummary::Auto => Self::Auto, + ReasoningSummary::Concise => Self::Concise, + ReasoningSummary::Detailed => Self::Detailed, + } + } +} + +#[derive(Debug, Serialize)] +#[serde(untagged)] +pub enum ResponsesInputItem { + Message(ResponsesMessageInput), + Reasoning(ResponsesReasoningInput), + FunctionCall(ResponsesFunctionCallInput), + FunctionCallOutput(ResponsesFunctionCallOutputInput), + ToolSearchCall(ResponsesToolSearchCallInput), + WebSearchCall(ResponsesWebSearchCallInput), + ImageGenerationCall(ResponsesImageGenerationCallInput), +} + +impl ResponsesInputItem { + fn from_message( + message: &Message, + target_provider: &ProviderId, + target_model: &str, + ) -> Result, ProviderError> { + let mut items = Vec::new(); + let mut content = Vec::new(); + let mut text_buffer = String::new(); + let mut replayed_reasoning = false; + + for block in &message.content { + match block { + ContentBlock::Text { text } => text_buffer.push_str(text), + ContentBlock::Thinking { + thinking, + encrypted_content, + id, + provenance, + .. + } => { + let replayable = matches!(message.role, Role::Assistant) + && encrypted_content + .as_deref() + .is_some_and(|value| !value.is_empty()) + && id.as_deref().is_some_and(|value| !value.is_empty()) + && provenance.as_ref().is_some_and(|provenance| { + provenance.provider == *target_provider + && provenance.model == target_model + && provenance.format == ReasoningFormat::OpenAiEncrypted + }); + if replayable { + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + items.push(Self::Reasoning(ResponsesReasoningInput { + kind: "reasoning", + id: id.clone().expect("replayable reasoning has an id"), + encrypted_content: encrypted_content + .clone() + .expect("replayable reasoning has encrypted content"), + summary: (!thinking.is_empty()) + .then(|| ResponsesReasoningSummaryInput { + kind: "summary_text", + text: thinking.clone(), + }) + .into_iter() + .collect(), + })); + replayed_reasoning = true; + } else { + text_buffer.push_str( + &block + .thinking_fallback_text() + .expect("thinking block has fallback text"), + ); + } + } + ContentBlock::Image { source } => { + Self::flush_text(&mut text_buffer, &message.role, &mut content)?; + content.push(ResponsesMessageContentPart::try_from(( + source, + &message.role, + ))?); + } + ContentBlock::ToolUse { id, name, input } => { + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + let (call_id, item_id) = decode_responses_tool_use_id(id); + items.push(Self::FunctionCall(ResponsesFunctionCallInput { + kind: "function_call", + id: if replayed_reasoning { + item_id.map(str::to_string) + } else { + None + }, + call_id: call_id.to_string(), + name: name.clone(), + arguments: input.to_string(), + })); + } + ContentBlock::ToolResult { + tool_use_id, + content: tool_output, + is_error, + } => { + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + let (call_id, _) = decode_responses_tool_use_id(tool_use_id); + items.push(Self::FunctionCallOutput(ResponsesFunctionCallOutputInput { + kind: "function_call_output", + call_id: call_id.to_string(), + output: render_tool_output(&tool_output.to_display_string(), *is_error), + })); + } + ContentBlock::HostedToolSearch { call } => { + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + items.push(Self::ToolSearchCall(ResponsesToolSearchCallInput::from( + call.clone(), + ))); + } + ContentBlock::HostedWebSearch { call } => { + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + items.push(Self::WebSearchCall(ResponsesWebSearchCallInput::from( + call.clone(), + ))); + } + ContentBlock::ImageGeneration { call } => { + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + items.push(Self::ImageGenerationCall( + ResponsesImageGenerationCallInput::from(call.clone()), + )); + } + } + } + + Self::flush_message(message, &mut text_buffer, &mut content, &mut items)?; + Ok(items) + } + + fn flush_text( + text_buffer: &mut String, + role: &Role, + content: &mut Vec, + ) -> Result<(), ProviderError> { + if text_buffer.is_empty() { + return Ok(()); + } + + content.push(ResponsesMessageContentPart::text_for_role( + role, + std::mem::take(text_buffer), + )?); + Ok(()) + } + + fn flush_message( + message: &Message, + text_buffer: &mut String, + content: &mut Vec, + items: &mut Vec, + ) -> Result<(), ProviderError> { + Self::flush_text(text_buffer, &message.role, content)?; + if content.is_empty() { + return Ok(()); + } + + items.push(Self::Message(ResponsesMessageInput { + role: message.role.to_string(), + content: std::mem::take(content), + })); + Ok(()) + } +} + +#[derive(Debug, Serialize)] +pub struct ResponsesReasoningInput { + #[serde(rename = "type")] + kind: &'static str, + id: String, + encrypted_content: String, + summary: Vec, +} + +#[derive(Debug, Serialize)] +pub struct ResponsesReasoningSummaryInput { + #[serde(rename = "type")] + kind: &'static str, + text: String, +} + +fn render_tool_output(content: &str, is_error: bool) -> String { + if is_error { + format!("Tool error: {content}") + } else { + content.to_string() + } +} + +#[derive(Debug, Serialize)] +pub struct ResponsesMessageInput { + role: String, + content: Vec, +} + +#[derive(Debug, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesMessageContentPart { + InputText { text: String }, + OutputText { text: String }, + InputImage { image_url: String }, +} + +impl ResponsesMessageContentPart { + fn text_for_role(role: &Role, text: String) -> Result { + match role { + Role::User => Ok(Self::InputText { text }), + Role::Assistant => Ok(Self::OutputText { text }), + Role::Unknown(role) => Err(ProviderError::InvalidRequest(format!( + "Responses message role '{role}' is not supported for text content" + ))), + } + } +} + +impl TryFrom<(&ImageSource, &Role)> for ResponsesMessageContentPart { + type Error = ProviderError; + + fn try_from(value: (&ImageSource, &Role)) -> Result { + let (source, role) = value; + if !matches!(role, Role::User) { + return Err(ProviderError::InvalidRequest( + "Responses image inputs are only supported in user messages".to_string(), + )); + } + + let image_url = match source { + ImageSource::Bytes { media_type, data } => { + format!("data:{media_type};base64,{}", STANDARD.encode(data)) + } + ImageSource::Url { url } => url.clone(), + }; + + Ok(ResponsesMessageContentPart::InputImage { image_url }) + } +} + +#[derive(Debug, Serialize)] +pub struct ResponsesFunctionCallInput { + #[serde(rename = "type")] + kind: &'static str, + #[serde(skip_serializing_if = "Option::is_none")] + id: Option, + call_id: String, + name: String, + arguments: String, +} + +#[derive(Debug, Serialize)] +pub struct ResponsesFunctionCallOutputInput { + #[serde(rename = "type")] + kind: &'static str, + call_id: String, + output: String, +} + +#[derive(Debug, Serialize)] +pub struct ResponsesToolSearchCallInput { + #[serde(rename = "type")] + kind: &'static str, + call_id: String, + #[serde(skip_serializing_if = "Option::is_none")] + status: Option, + execution: &'static str, + arguments: serde_json::Value, +} + +impl From for ResponsesToolSearchCallInput { + fn from(value: HostedToolSearchCall) -> Self { + Self { + kind: "tool_search_call", + call_id: value.id, + status: value.status, + execution: "client", + arguments: serde_json::json!({ + "query": value.query, + }), + } + } +} + +#[derive(Debug, Serialize)] +pub struct ResponsesWebSearchCallInput { + #[serde(rename = "type")] + kind: &'static str, + id: String, + #[serde(skip_serializing_if = "Option::is_none")] + status: Option, + #[serde(skip_serializing_if = "Option::is_none")] + action: Option, +} + +impl From for ResponsesWebSearchCallInput { + fn from(value: HostedWebSearchCall) -> Self { + Self { + kind: "web_search_call", + id: value.id, + status: value.status, + action: value.action, + } + } +} + +#[derive(Debug, Serialize)] +pub struct ResponsesImageGenerationCallInput { + #[serde(rename = "type")] + kind: &'static str, + id: String, + status: String, + #[serde(skip_serializing_if = "Option::is_none")] + revised_prompt: Option, + #[serde(skip_serializing_if = "Option::is_none")] + result: Option, +} + +impl From for ResponsesImageGenerationCallInput { + fn from(value: ImageGenerationCall) -> Self { + Self { + kind: "image_generation_call", + id: value.id, + status: value.status, + revised_prompt: value.revised_prompt, + result: value.result.map(render_image_generation_result), + } + } +} + +fn render_image_generation_result(result: ImageGenerationResult) -> String { + match result { + ImageGenerationResult::ArtifactRef { artifact_id } => artifact_id, + ImageGenerationResult::Image { source } => match source { + ImageSource::Bytes { data, .. } => STANDARD.encode(data), + ImageSource::Url { url } => url, + }, + } +} + +#[derive(Debug, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesTool { + Function { + name: String, + #[serde(skip_serializing_if = "Option::is_none")] + description: Option, + parameters: serde_json::Value, + strict: bool, + #[serde(skip_serializing_if = "std::ops::Not::not")] + defer_loading: bool, + }, + ToolSearch {}, + WebSearch { + #[serde(skip_serializing_if = "Option::is_none")] + external_web_access: Option, + #[serde(skip_serializing_if = "Option::is_none")] + filters: Option, + #[serde(skip_serializing_if = "Option::is_none")] + user_location: Option, + #[serde(skip_serializing_if = "Option::is_none")] + search_context_size: Option, + #[serde(skip_serializing_if = "Option::is_none")] + search_content_types: Option>, + }, + ImageGeneration { + output_format: String, + }, +} + +#[derive(Debug, Deserialize)] +struct ResponsesWebSearchToolOptions { + #[serde(default)] + external_web_access: Option, + #[serde(default)] + filters: Option, + #[serde(default)] + user_location: Option, + #[serde(default)] + search_context_size: Option, + #[serde(default)] + search_content_types: Option>, +} + +#[derive(Debug, Deserialize)] +struct ResponsesImageGenerationToolOptions { + output_format: String, +} + +impl ResponsesTool { + fn function(tool: &ToolSpec, force_immediate: bool) -> Self { + Self::Function { + name: tool.name.clone(), + description: tool.description.clone(), + parameters: tool.input_schema.clone(), + strict: tool.strict.unwrap_or(false), + defer_loading: tool.loading_policy == ToolLoadingPolicy::Deferred && !force_immediate, + } + } + + fn tool_search() -> Self { + Self::ToolSearch {} + } + + fn web_search(tool: &ToolSpec) -> Result { + let options = tool.options.clone().ok_or_else(|| { + ProviderError::InvalidRequest("web_search tools require provider options".to_string()) + })?; + let options: ResponsesWebSearchToolOptions = + serde_json::from_value(options).map_err(|error| { + ProviderError::InvalidRequest(format!("invalid web_search tool options: {error}")) + })?; + Ok(Self::WebSearch { + external_web_access: options.external_web_access, + filters: options.filters, + user_location: options.user_location, + search_context_size: options.search_context_size, + search_content_types: options.search_content_types, + }) + } + + fn image_generation(tool: &ToolSpec) -> Result { + let options = tool.options.clone().ok_or_else(|| { + ProviderError::InvalidRequest( + "image_generation tools require provider options".to_string(), + ) + })?; + let options: ResponsesImageGenerationToolOptions = serde_json::from_value(options) + .map_err(|error| { + ProviderError::InvalidRequest(format!( + "invalid image_generation tool options: {error}" + )) + })?; + Ok(Self::ImageGeneration { + output_format: options.output_format, + }) + } +} + +fn build_responses_tools( + tools: &[ToolSpec], + tool_choice: Option<&ToolChoice>, + tool_search_mode: ToolSearchMode, + provider_name: &str, +) -> Result, ProviderError> { + let forced_tool_name = match tool_choice { + Some(ToolChoice::Tool { name }) => Some(name.as_str()), + _ => None, + }; + + let has_deferred_tools = tools.iter().any(|tool| { + tool.kind == ProviderToolKind::Function + && tool.loading_policy == ToolLoadingPolicy::Deferred + && forced_tool_name != Some(tool.name.as_str()) + }); + + if has_deferred_tools && tool_search_mode != ToolSearchMode::Hosted { + return Err(ProviderError::InvalidRequest(format!( + "{provider_name} deferred tools require hosted tool search" + ))); + } + + let mut provider_tools = Vec::with_capacity(tools.len() + usize::from(has_deferred_tools)); + for tool in tools { + let provider_tool = match tool.kind { + ProviderToolKind::Function => { + ResponsesTool::function(tool, forced_tool_name == Some(tool.name.as_str())) + } + ProviderToolKind::HostedWebSearch => ResponsesTool::web_search(tool)?, + ProviderToolKind::ImageGeneration => ResponsesTool::image_generation(tool)?, + }; + provider_tools.push(provider_tool); + } + + if has_deferred_tools { + provider_tools.push(ResponsesTool::tool_search()); + } + + Ok(provider_tools) +} + +#[derive(Debug, Serialize)] +#[serde(untagged)] +pub enum ResponsesToolChoice { + Mode(ResponsesToolChoiceMode), + Function(ResponsesToolChoiceFunction), +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum ResponsesToolChoiceMode { + Auto, + Required, +} + +#[derive(Debug, Serialize)] +pub struct ResponsesToolChoiceFunction { + #[serde(rename = "type")] + kind: &'static str, + name: String, +} + +impl From for ResponsesToolChoice { + fn from(choice: ToolChoice) -> Self { + match choice { + ToolChoice::Auto => ResponsesToolChoice::Mode(ResponsesToolChoiceMode::Auto), + ToolChoice::Any => ResponsesToolChoice::Mode(ResponsesToolChoiceMode::Required), + ToolChoice::Tool { name } => { + ResponsesToolChoice::Function(ResponsesToolChoiceFunction { + kind: "function", + name, + }) + } + } + } +} + +#[cfg(test)] +mod tests { + use std::{borrow::Cow, collections::BTreeMap}; + + use serde_json::json; + use time::OffsetDateTime; + + use crate::{ + ContentBlock, HostedToolSearchCall, HostedWebSearchCall, ImageGenerationCall, + ImageGenerationResult, Message, ProviderError, ProviderId, ProviderRequestOptions, + ReasoningEffort, ReasoningFormat, ReasoningOptions, ReasoningProvenance, ReasoningSummary, + Request, ResponsesRequestCompression, ResponsesRequestOptions, ResponsesStateMode, + ResponsesTextControls, ResponsesTextFormat, ResponsesTransport, ResponsesVerbosity, Role, + ToolChoice, ToolLoadingPolicy, ToolResultContent, ToolSearchMode, ToolSpec, + WebSearchAction, + }; + + use super::{ResponsesModel, ResponsesModelsPage, ResponsesRequest}; + + fn responses_thinking( + thinking: &str, + encrypted_content: Option<&str>, + id: Option<&str>, + provider: &str, + model: &str, + ) -> ContentBlock { + ContentBlock::Thinking { + thinking: thinking.to_string(), + signature: None, + encrypted_content: encrypted_content.map(str::to_string), + id: id.map(str::to_string), + provenance: Some(ReasoningProvenance { + provider: ProviderId::new(provider), + model: model.to_string(), + format: ReasoningFormat::OpenAiEncrypted, + }), + redacted: false, + } + } + + fn replay_request(model: &str, messages: Vec) -> Request<'static> { + Request { + model: Cow::Owned(model.to_string()), + system: None, + messages: Cow::Owned(messages), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + } + } + + #[test] + fn converts_request_to_responses_payload() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: Some(Cow::Borrowed("Be helpful.")), + messages: Cow::Owned(vec![ + Message::user(ContentBlock::text("What files changed?")), + Message::assistant(ContentBlock::text("I'll inspect that.")), + Message::assistant(ContentBlock::ToolUse { + id: "call_123".to_string(), + name: "files".to_string(), + input: json!({ "operations": [{ "op": "read", "path": "README.md" }] }), + }), + Message::assistant(ContentBlock::ToolResult { + tool_use_id: "call_123".to_string(), + content: ToolResultContent::text("README contents"), + is_error: false, + }), + ]), + tools: Cow::Owned(vec![ToolSpec { + name: "files".to_string(), + description: Some("Read and edit files".to_string()), + input_schema: json!({ + "type": "object", + "properties": { + "operations": { "type": "array" } + } + }), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: crate::tool::ToolLoadingPolicy::Immediate, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Tool { + name: "files".to_string(), + }), + temperature: Some(0.2), + max_output_tokens: Some(256), + metadata: Cow::Owned(BTreeMap::from([( + "agent".to_string(), + "mentra".to_string(), + )])), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["model"], "gpt-5"); + assert_eq!(payload["instructions"], "Be helpful."); + assert_eq!(payload["input"][0]["role"], "user"); + assert_eq!(payload["input"][0]["content"][0]["type"], "input_text"); + assert_eq!( + payload["input"][0]["content"][0]["text"], + "What files changed?" + ); + assert_eq!(payload["input"][1]["role"], "assistant"); + assert_eq!(payload["input"][1]["content"][0]["type"], "output_text"); + assert_eq!( + payload["input"][1]["content"][0]["text"], + "I'll inspect that." + ); + assert_eq!(payload["input"][2]["type"], "function_call"); + assert_eq!(payload["input"][2]["call_id"], "call_123"); + assert_eq!(payload["input"][2]["name"], "files"); + assert_eq!(payload["input"][3]["type"], "function_call_output"); + assert_eq!(payload["input"][3]["output"], "README contents"); + assert_eq!(payload["tools"][0]["type"], "function"); + assert_eq!(payload["tools"][0]["strict"], false); + assert!(payload["tools"][0].get("defer_loading").is_none()); + assert_eq!(payload["tool_choice"]["type"], "function"); + assert_eq!(payload["tool_choice"]["name"], "files"); + let temperature = payload["temperature"] + .as_f64() + .expect("temperature should be numeric"); + assert!((temperature - 0.2).abs() < 1e-6); + assert_eq!(payload["max_output_tokens"], 256); + assert_eq!(payload["metadata"]["agent"], "mentra"); + } + + #[test] + fn replays_reasoning_and_paired_function_ids_for_exact_provider_and_model() { + let request = replay_request( + "gpt-requested", + vec![ + Message { + role: Role::Assistant, + content: vec![ + responses_thinking( + "short summary", + Some("encrypted-1"), + Some("rs_1"), + "openai-edge", + "gpt-requested", + ), + ContentBlock::ToolUse { + id: "call_1|fc_1".to_string(), + name: "read_file".to_string(), + input: json!({"path":"README.md"}), + }, + ], + }, + Message::user(ContentBlock::ToolResult { + tool_use_id: "call_1|fc_1".to_string(), + content: ToolResultContent::text("contents"), + is_error: false, + }), + ], + ); + + let payload = serde_json::to_value( + ResponsesRequest::try_from_request_with_provider( + request, + "OpenAI edge", + &ProviderId::new("openai-edge"), + ) + .unwrap(), + ) + .unwrap(); + + assert_eq!(payload["input"][0]["type"], "reasoning"); + assert_eq!(payload["input"][0]["id"], "rs_1"); + assert_eq!(payload["input"][0]["encrypted_content"], "encrypted-1"); + assert_eq!(payload["input"][0]["summary"][0]["text"], "short summary"); + assert_eq!(payload["input"][1]["type"], "function_call"); + assert_eq!(payload["input"][1]["id"], "fc_1"); + assert_eq!(payload["input"][1]["call_id"], "call_1"); + assert_eq!(payload["input"][2]["type"], "function_call_output"); + assert_eq!(payload["input"][2]["call_id"], "call_1"); + } + + #[test] + fn cross_model_reasoning_downgrades_and_omits_function_item_id() { + let request = replay_request( + "gpt-new", + vec![Message { + role: Role::Assistant, + content: vec![ + responses_thinking( + "short summary", + Some("encrypted-1"), + Some("rs_1"), + "openai-edge", + "gpt-old", + ), + ContentBlock::ToolUse { + id: "call_1|fc_1".to_string(), + name: "read_file".to_string(), + input: json!({}), + }, + ], + }], + ); + + let payload = serde_json::to_value( + ResponsesRequest::try_from_request_with_provider( + request, + "OpenAI edge", + &ProviderId::new("openai-edge"), + ) + .unwrap(), + ) + .unwrap(); + + assert_eq!(payload["input"][0]["role"], "assistant"); + assert_eq!(payload["input"][0]["content"][0]["text"], "short summary"); + assert_eq!(payload["input"][1]["type"], "function_call"); + assert!(payload["input"][1].get("id").is_none()); + assert_eq!(payload["input"][1]["call_id"], "call_1"); + } + + #[test] + fn cross_provider_user_and_incomplete_reasoning_downgrade_to_text() { + let blocks = vec![ + responses_thinking( + "cross provider", + Some("encrypted-1"), + Some("rs_1"), + "openai-other", + "gpt-requested", + ), + responses_thinking( + "missing encrypted content", + None, + Some("rs_2"), + "openai-edge", + "gpt-requested", + ), + responses_thinking( + "missing id", + Some("encrypted-3"), + None, + "openai-edge", + "gpt-requested", + ), + ]; + let request = replay_request( + "gpt-requested", + vec![Message { + role: Role::User, + content: blocks, + }], + ); + + let payload = serde_json::to_value( + ResponsesRequest::try_from_request_with_provider( + request, + "OpenAI edge", + &ProviderId::new("openai-edge"), + ) + .unwrap(), + ) + .unwrap(); + + assert_eq!(payload["input"].as_array().unwrap().len(), 1); + assert_eq!(payload["input"][0]["role"], "user"); + assert_eq!( + payload["input"][0]["content"][0]["text"], + "cross providermissing encrypted contentmissing id" + ); + } + + #[test] + fn serializes_explicit_strict_function_tool() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![ + ToolSpec::builder("strict_tool") + .input_schema(json!({ + "type": "object", + "properties": { + "path": { "type": "string" } + }, + "required": ["path"], + "additionalProperties": false + })) + .strict(true) + .build(), + ]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["tools"][0]["type"], "function"); + assert_eq!(payload["tools"][0]["strict"], true); + } + + #[test] + fn prefixes_tool_errors_in_function_call_output() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::ToolResult { + tool_use_id: "call_456".to_string(), + content: ToolResultContent::text("No such file"), + is_error: true, + })]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["input"][0]["output"], "Tool error: No such file"); + assert_eq!(payload["tool_choice"], "auto"); + } + + #[test] + fn preserves_hosted_actions_in_responses_history_replay() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![Message { + role: Role::Assistant, + content: vec![ + ContentBlock::HostedToolSearch { + call: HostedToolSearchCall { + id: "search_1".to_string(), + status: Some("completed".to_string()), + query: Some("weather".to_string()), + }, + }, + ContentBlock::HostedWebSearch { + call: HostedWebSearchCall { + id: "web_1".to_string(), + status: Some("completed".to_string()), + action: Some(WebSearchAction::Search { + query: Some("weather".to_string()), + queries: None, + }), + }, + }, + ContentBlock::ImageGeneration { + call: ImageGenerationCall { + id: "image_1".to_string(), + status: "completed".to_string(), + revised_prompt: Some("A blue square".to_string()), + result: Some(ImageGenerationResult::ArtifactRef { + artifact_id: "artifact_1".to_string(), + }), + }, + }, + ], + }]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["input"][0]["type"], "tool_search_call"); + assert_eq!(payload["input"][0]["call_id"], "search_1"); + assert_eq!(payload["input"][1]["type"], "web_search_call"); + assert_eq!(payload["input"][1]["id"], "web_1"); + assert_eq!(payload["input"][2]["type"], "image_generation_call"); + assert_eq!(payload["input"][2]["id"], "image_1"); + assert_eq!(payload["input"][2]["result"], "artifact_1"); + } + + #[test] + fn serializes_inline_images_as_input_image_parts() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![Message { + role: Role::User, + content: vec![ + ContentBlock::text("What is in this image?"), + ContentBlock::image_bytes("image/png", [1_u8, 2, 3]), + ], + }]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["input"][0]["role"], "user"); + assert_eq!(payload["input"][0]["content"][0]["type"], "input_text"); + assert_eq!( + payload["input"][0]["content"][0]["text"], + "What is in this image?" + ); + assert_eq!(payload["input"][0]["content"][1]["type"], "input_image"); + assert_eq!( + payload["input"][0]["content"][1]["image_url"], + "data:image/png;base64,AQID" + ); + } + + #[test] + fn serializes_assistant_text_as_output_text() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![Message::assistant(ContentBlock::text("Done."))]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["input"][0]["role"], "assistant"); + assert_eq!(payload["input"][0]["content"][0]["type"], "output_text"); + assert_eq!(payload["input"][0]["content"][0]["text"], "Done."); + } + + #[test] + fn converts_unix_timestamp_to_offset_datetime() { + let model = ResponsesModel { + id: "gpt-5".to_string(), + name: None, + description: None, + owned_by: None, + created: Some(1_741_049_700), + }; + + let info = model.into_model_info(crate::ProviderId::new("openai")); + + assert_eq!( + info.created_at, + Some(OffsetDateTime::from_unix_timestamp(1_741_049_700).expect("valid timestamp")) + ); + } + + #[test] + fn converts_model_list_page_into_model_info() { + let page: ResponsesModelsPage = serde_json::from_str( + r#"{ + "data": [ + { + "id": "gpt-5", + "name": "GPT-5", + "description": "General-purpose model", + "created": 1741049700 + }, + { + "id": "gpt-5-mini", + "owned_by": "openai" + } + ] + }"#, + ) + .expect("page should parse"); + + let models = page.into_model_info(crate::ProviderId::new("openai")); + + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "gpt-5"); + assert_eq!(models[0].display_name.as_deref(), Some("GPT-5")); + assert_eq!( + models[0].description.as_deref(), + Some("General-purpose model") + ); + assert_eq!(models[1].description.as_deref(), Some("Owned by openai")); + } + + #[test] + fn serializes_parallel_tool_calls_option() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Disabled, + reasoning: None, + responses: ResponsesRequestOptions { + parallel_tool_calls: Some(true), + ..Default::default() + }, + anthropic: Default::default(), + gemini: Default::default(), + session: Default::default(), + }, + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["parallel_tool_calls"], true); + } + + #[test] + fn serializes_previous_response_id_option() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + responses: ResponsesRequestOptions { + previous_response_id: Some("resp_previous".to_string()), + state_mode: ResponsesStateMode::Stateful, + ..Default::default() + }, + ..Default::default() + }, + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["previous_response_id"], "resp_previous"); + } + + #[test] + fn serializes_all_reasoning_effort_options() { + for (effort, expected) in [ + (ReasoningEffort::Low, "low"), + (ReasoningEffort::Medium, "medium"), + (ReasoningEffort::High, "high"), + (ReasoningEffort::XHigh, "xhigh"), + (ReasoningEffort::Max, "max"), + ] { + let request = Request { + model: Cow::Borrowed("gpt-5.6-luna"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(ReasoningOptions { + effort: Some(effort), + summary: None, + }), + ..Default::default() + }, + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["reasoning"]["effort"], expected); + assert_eq!(payload["include"][0], "reasoning.encrypted_content"); + } + } + + #[test] + fn reasoning_encrypted_content_include_is_added_once() { + let mut request = replay_request("gpt-5", Vec::new()); + request.provider_request_options.reasoning = Some(ReasoningOptions { + effort: None, + summary: Some(ReasoningSummary::Auto), + }); + request.provider_request_options.responses.include = vec![ + "message.output_text.logprobs".to_string(), + "reasoning.encrypted_content".to_string(), + ]; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()).unwrap(); + assert_eq!( + payload["include"], + json!([ + "message.output_text.logprobs", + "reasoning.encrypted_content" + ]) + ); + } + + #[test] + fn serializes_advanced_responses_request_controls() { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: Some(Cow::Borrowed("Be structured.")), + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(ReasoningOptions { + effort: None, + summary: Some(ReasoningSummary::Detailed), + }), + responses: ResponsesRequestOptions { + parallel_tool_calls: Some(true), + previous_response_id: None, + state_mode: ResponsesStateMode::Hybrid, + transport: ResponsesTransport::HttpSse, + store: Some(false), + stream: Some(true), + include: vec!["reasoning.encrypted_content".to_string()], + service_tier: Some("priority".to_string()), + prompt_cache_key: Some("thread-123".to_string()), + text: Some(ResponsesTextControls { + verbosity: Some(ResponsesVerbosity::High), + format: Some(ResponsesTextFormat { + r#type: crate::ResponsesTextFormatType::JsonSchema, + strict: true, + schema: json!({"type":"object"}), + name: "codex_output_schema".to_string(), + }), + }), + compression: ResponsesRequestCompression::None, + }, + ..Default::default() + }, + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["store"], false); + assert_eq!(payload["stream"], true); + assert_eq!(payload["include"][0], "reasoning.encrypted_content"); + assert_eq!(payload["service_tier"], "priority"); + assert_eq!(payload["prompt_cache_key"], "thread-123"); + assert_eq!(payload["reasoning"]["summary"], "detailed"); + assert_eq!(payload["text"]["verbosity"], "high"); + assert_eq!(payload["text"]["format"]["name"], "codex_output_schema"); + } + + #[test] + fn hosted_tool_search_adds_search_tool_for_deferred_tools() { + let request = Request { + model: Cow::Borrowed("gpt-5.4"), + system: None, + messages: Cow::Owned(vec![Message::user(ContentBlock::text("hello"))]), + tools: Cow::Owned(vec![ToolSpec { + name: "lookup_order".to_string(), + description: Some("Look up an order".to_string()), + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + tool_search_mode: ToolSearchMode::Hosted, + ..Default::default() + }, + }; + + let payload = serde_json::to_value(ResponsesRequest::try_from(request).unwrap()) + .expect("request should serialize"); + + assert_eq!(payload["tools"][0]["type"], "function"); + assert_eq!(payload["tools"][0]["name"], "lookup_order"); + assert_eq!(payload["tools"][0]["defer_loading"], true); + assert_eq!(payload["tools"][1]["type"], "tool_search"); + } + + #[test] + fn rejects_deferred_tools_without_hosted_tool_search() { + let request = Request { + model: Cow::Borrowed("gpt-5.4"), + system: None, + messages: Cow::Owned(vec![]), + tools: Cow::Owned(vec![ToolSpec { + name: "lookup_order".to_string(), + description: None, + input_schema: json!({"type":"object"}), + output_schema: None, + kind: crate::ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Deferred, + strict: None, + options: None, + }]), + tool_choice: Some(ToolChoice::Auto), + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let error = ResponsesRequest::try_from(request).expect_err("request should fail"); + match error { + ProviderError::InvalidRequest(message) => { + assert!(message.contains("deferred tools require hosted tool search")); + } + other => panic!("unexpected error: {other:?}"), + } + } +} diff --git a/vendor/mentra-provider/src/responses/session.rs b/vendor/mentra-provider/src/responses/session.rs new file mode 100644 index 0000000..c4f165e --- /dev/null +++ b/vendor/mentra-provider/src/responses/session.rs @@ -0,0 +1,2041 @@ +use std::collections::HashSet; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use http::HeaderMap; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::oneshot::error::TryRecvError; +use url::Url; + +use crate::CompactionRequest; +use crate::CompactionResponse; +use crate::CredentialSource; +use crate::MemorySummarizeRequest; +use crate::MemorySummarizeResponse; +use crate::ModelInfo; +use crate::ProviderCredentials; +use crate::ProviderDefinition; +use crate::ProviderError; +use crate::ProviderEvent; +use crate::ProviderEventStream; +use crate::ProviderSession; +use crate::ReasoningFormat; +use crate::ReasoningProvenance; +use crate::Request; +use crate::Response; +use crate::ResponsesTransport; +use crate::SessionRequestOptions; +use crate::request::ResponsesRequestCompression; + +use super::SharedTurnState; +use super::model::ResponsesModelsPage; +use super::model::ResponsesRequest; +use super::sse::spawn_event_stream_with_provenance; +#[cfg(feature = "responses-websocket")] +use super::websocket::ResponsesWebsocketConnection; +#[cfg(feature = "responses-websocket")] +use super::websocket::ResponsesWebsocketTelemetry; +#[cfg(feature = "responses-websocket")] +use super::websocket::merge_request_headers; +#[cfg(feature = "responses-websocket")] +use super::websocket::response_create_frame; + +/// Session-scoped Responses transport state. +/// +/// This is intentionally lightweight and keeps the pieces needed for websocket prewarm and +/// HTTP fallback without binding the provider to any higher-level runtime. +pub struct ResponsesSession { + definition: ProviderDefinition, + credential_source: Arc, + client: reqwest::Client, + state: Arc, + endpoint_capabilities: Arc, + hybrid_http_previous_response_id: bool, +} + +#[derive(Default)] +struct WebsocketSession { + connection_reused: StdMutex, + #[cfg(feature = "responses-websocket")] + connection: Option, + _last_request: Option, + last_response_rx: Option>, +} + +impl WebsocketSession { + fn set_connection_reused(&self, connection_reused: bool) { + *self + .connection_reused + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = connection_reused; + } + + fn connection_reused(&self) -> bool { + *self + .connection_reused + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + #[cfg(feature = "responses-websocket")] + fn connection(&self) -> Option { + self.connection.clone() + } + + #[cfg(feature = "responses-websocket")] + fn store_connection(&mut self, connection: ResponsesWebsocketConnection) { + self.connection = Some(connection); + } + + #[cfg(feature = "responses-websocket")] + fn clear_connection(&mut self) { + self.connection = None; + } +} + +pub(crate) struct ResponsesSessionState { + disable_websockets: AtomicBool, + websocket_session: StdMutex, + turn_state: SharedTurnState, + latest_response_id: StdMutex>, +} + +impl Default for ResponsesSessionState { + fn default() -> Self { + Self { + disable_websockets: AtomicBool::new(false), + websocket_session: StdMutex::new(WebsocketSession::default()), + turn_state: Arc::new(StdMutex::new(None)), + latest_response_id: StdMutex::new(None), + } + } +} + +impl ResponsesSessionState { + fn turn_state(&self) -> Option { + self.turn_state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + + fn set_turn_state(&self, turn_state: impl Into) { + *self + .turn_state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(turn_state.into()); + } + + fn latest_response_id(&self) -> Option { + self.latest_response_id + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + + fn set_latest_response_id(&self, response_id: impl Into) { + *self + .latest_response_id + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(response_id.into()); + } + + fn clear_latest_response_id(&self) { + *self + .latest_response_id + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = None; + } +} + +#[derive(Default)] +pub(crate) struct ResponsesEndpointCapabilities { + http_previous_response_id_unsupported_models: StdMutex>, +} + +impl ResponsesEndpointCapabilities { + fn http_previous_response_id_is_unsupported(&self, model: &str) -> bool { + self.http_previous_response_id_unsupported_models + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .contains(model) + } + + fn mark_http_previous_response_id_unsupported(&self, model: impl Into) { + self.http_previous_response_id_unsupported_models + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .insert(model.into()); + } +} + +impl ResponsesSession +where + C: CredentialSource + 'static, +{ + pub(crate) fn new( + definition: ProviderDefinition, + credential_source: Arc, + client: reqwest::Client, + state: Arc, + endpoint_capabilities: Arc, + hybrid_http_previous_response_id: bool, + ) -> Self { + Self { + definition, + credential_source, + client, + state, + endpoint_capabilities, + hybrid_http_previous_response_id, + } + } + + pub fn websocket_connect_timeout(&self) -> Duration { + self.definition.websocket_connect_timeout + } + + pub fn stream_idle_timeout(&self) -> Duration { + self.definition.stream_idle_timeout + } + + pub fn websocket_url_for_path(&self, path: &str) -> Result { + self.definition + .websocket_url_for_path(path) + .map_err(|error| ProviderError::InvalidRequest(error.to_string())) + } + + pub fn request_url_for_path(&self, path: &str) -> Result { + Url::parse(&self.definition.url_for_path(path)) + .map_err(|error| ProviderError::InvalidRequest(error.to_string())) + } + + fn responses_endpoint_path(&self) -> &'static str { + responses_endpoint_path_for_base(self.definition.base_url.as_deref()) + } + + pub fn disable_websockets(&self) { + self.state.disable_websockets.store(true, Ordering::Relaxed); + } + + pub fn websockets_enabled(&self) -> bool { + self.definition.capabilities.supports_websockets + && !self.state.disable_websockets.load(Ordering::Relaxed) + } + + pub fn set_connection_reused(&self, connection_reused: bool) { + self.state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .set_connection_reused(connection_reused); + } + + pub fn connection_reused(&self) -> bool { + self.state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .connection_reused() + } + + pub fn build_websocket_headers( + &self, + credentials: &ProviderCredentials, + turn_metadata_header: Option<&str>, + ) -> Result { + let session = SessionRequestOptions { + sticky_turn_state: self.state.turn_state(), + turn_metadata: turn_metadata_header.map(str::to_string), + subagent: None, + prefer_connection_reuse: Some(self.connection_reused()), + session_affinity: None, + extra_headers: std::collections::BTreeMap::new(), + }; + self.build_websocket_headers_for_session(credentials, Some(&session)) + } + + pub fn build_websocket_headers_for_session( + &self, + credentials: &ProviderCredentials, + session: Option<&SessionRequestOptions>, + ) -> Result { + let fallback_turn_state = self.state.turn_state(); + self.definition.build_headers_for_session( + credentials, + session, + fallback_turn_state.as_deref(), + ) + } + + pub fn set_turn_state(&self, turn_state: impl Into) -> bool { + self.state.set_turn_state(turn_state); + true + } + + pub fn turn_state(&self) -> Option { + self.state.turn_state() + } + + pub fn latest_response_id(&self) -> Option { + self.state.latest_response_id() + } + + pub fn clear_latest_response_id(&self) { + self.state.clear_latest_response_id(); + } + + /// Whether the cached websocket connection is gone or was never opened. + /// + /// This and the four methods below are the session's half of the + /// `responses-websocket` transport; without the feature there is no + /// connection to cache, so they are not compiled. + #[cfg(feature = "responses-websocket")] + pub async fn websocket_connection_is_closed(&self) -> bool { + let connection = self + .state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .connection(); + match connection { + Some(connection) => connection.is_closed().await, + None => true, + } + } + + #[cfg(feature = "responses-websocket")] + pub async fn connect_websocket( + &self, + extra_headers: HeaderMap, + default_headers: HeaderMap, + turn_state: Option, + telemetry: Option>, + ) -> Result<(), ProviderError> { + let credentials = self.credential_source.credentials().await?; + let provider_headers = self.definition.build_headers(&credentials)?; + let headers = merge_request_headers(&provider_headers, extra_headers, default_headers); + let url = self + .definition + .websocket_url_with_auth_for_path(self.responses_endpoint_path(), &credentials)?; + let connection = ResponsesWebsocketConnection::connect( + url, + headers, + turn_state.or_else(|| Some(Arc::clone(&self.state.turn_state))), + self.stream_idle_timeout(), + telemetry, + ) + .await?; + self.state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .store_connection(connection); + Ok(()) + } + + #[cfg(feature = "responses-websocket")] + pub async fn stream_websocket_request( + &self, + request_body: serde_json::Value, + ) -> Result { + let (connection, connection_reused) = { + let websocket_session = self + .state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + ( + websocket_session.connection(), + websocket_session.connection_reused(), + ) + }; + let Some(connection) = connection else { + return Err(ProviderError::MalformedStream( + "websocket connection is unavailable".to_string(), + )); + }; + connection + .stream_request(request_body, connection_reused) + .await + } + + #[cfg(feature = "responses-websocket")] + async fn stream_websocket_request_with_provenance( + &self, + request_body: serde_json::Value, + provenance: ReasoningProvenance, + ) -> Result { + let (connection, connection_reused) = { + let websocket_session = self + .state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + ( + websocket_session.connection(), + websocket_session.connection_reused(), + ) + }; + let Some(connection) = connection else { + return Err(ProviderError::MalformedStream( + "websocket connection is unavailable".to_string(), + )); + }; + connection + .stream_request_with_provenance(request_body, connection_reused, provenance) + .await + } + + #[cfg(feature = "responses-websocket")] + pub fn clear_websocket_connection(&self) { + self.state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clear_connection(); + } + + pub async fn list_models(&self) -> Result, ProviderError> { + let credentials = self.credential_source.credentials().await?; + let response = self + .client + .get( + self.definition + .request_url_with_auth_for_path("v1/models", &credentials)?, + ) + .headers(self.definition.build_headers(&credentials)?) + .send() + .await + .map_err(ProviderError::Transport)?; + + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + + let models = response + .json::() + .await + .map_err(ProviderError::Decode)?; + + Ok(models.into_model_info(self.definition.descriptor.id.clone())) + } + + pub async fn stream_response<'a>( + &self, + mut request: Request<'a>, + ) -> Result { + let provider_name = self + .definition + .descriptor + .display_name + .as_deref() + .unwrap_or(self.definition.descriptor.id.as_str()); + let target_provider = self.definition.descriptor.id.clone(); + let requested_model = request.model.to_string(); + let reasoning_provenance = ReasoningProvenance { + provider: target_provider.clone(), + model: requested_model.clone(), + format: ReasoningFormat::OpenAiEncrypted, + }; + let session = request.provider_request_options.session.clone(); + let state_mode = request.provider_request_options.responses.state_mode; + let transport = request.provider_request_options.responses.transport; + if request + .provider_request_options + .responses + .previous_response_id + .is_none() + && state_mode.uses_provider_state() + && !(state_mode == crate::ResponsesStateMode::Hybrid + && transport == ResponsesTransport::HttpSse + && (!self.hybrid_http_previous_response_id + || self + .endpoint_capabilities + .http_previous_response_id_is_unsupported(&requested_model))) + { + request + .provider_request_options + .responses + .previous_response_id = self.state.latest_response_id(); + } + let compression = request.provider_request_options.responses.compression; + let request = ResponsesRequest::try_from_request_with_provider( + request, + provider_name, + &target_provider, + )?; + let credentials = self.credential_source.credentials().await?; + + match transport { + ResponsesTransport::HttpSse => { + self.stream_http_response( + request, + compression, + &credentials, + &session, + state_mode, + reasoning_provenance, + ) + .await + } + ResponsesTransport::WebSocket => { + self.stream_websocket_response( + request, + &credentials, + &session, + reasoning_provenance, + ) + .await + } + } + } + + async fn stream_http_response( + &self, + mut request: ResponsesRequest, + compression: ResponsesRequestCompression, + credentials: &ProviderCredentials, + session: &SessionRequestOptions, + state_mode: crate::ResponsesStateMode, + reasoning_provenance: ReasoningProvenance, + ) -> Result { + let response = self + .send_http_responses_request(&request, compression, credentials, session) + .await?; + + if !response.status().is_success() { + let error = ProviderError::from_http_response(response).await; + if state_mode == crate::ResponsesStateMode::Hybrid + && request.previous_response_id().is_some() + && let Some(rejection) = previous_response_state_rejection(&error) + { + if rejection == PreviousResponseStateRejection::ParameterUnsupported { + self.endpoint_capabilities + .mark_http_previous_response_id_unsupported( + reasoning_provenance.model.clone(), + ); + } + self.state.clear_latest_response_id(); + request.clear_previous_response_id(); + let response = self + .send_http_responses_request(&request, compression, credentials, session) + .await?; + if !response.status().is_success() { + return Err(ProviderError::from_http_response(response).await); + } + return Ok( + self.track_response_state(spawn_event_stream_with_provenance( + response, + reasoning_provenance.provider, + reasoning_provenance.model, + )), + ); + } + return Err(error); + } + + if let Some(turn_state) = response + .headers() + .get("x-codex-turn-state") + .and_then(|value| value.to_str().ok()) + { + self.state.set_turn_state(turn_state); + } + + Ok( + self.track_response_state(spawn_event_stream_with_provenance( + response, + reasoning_provenance.provider, + reasoning_provenance.model, + )), + ) + } + + #[cfg(feature = "responses-websocket")] + async fn stream_websocket_response( + &self, + request: ResponsesRequest, + credentials: &ProviderCredentials, + session: &SessionRequestOptions, + reasoning_provenance: ReasoningProvenance, + ) -> Result { + if !self.websockets_enabled() { + return Err(ProviderError::UnsupportedCapability( + "responses_websocket".to_string(), + )); + } + + if self.websocket_connection_is_closed().await { + let headers = self.build_websocket_headers_for_session(credentials, Some(session))?; + let url = self + .definition + .websocket_url_with_auth_for_path(self.responses_endpoint_path(), credentials)?; + let connection = ResponsesWebsocketConnection::connect( + url, + headers, + Some(Arc::clone(&self.state.turn_state)), + self.stream_idle_timeout(), + None, + ) + .await?; + self.state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .store_connection(connection); + } + + let response = serde_json::to_value(request).map_err(ProviderError::Serialize)?; + let stream = self + .stream_websocket_request_with_provenance( + response_create_frame(response), + reasoning_provenance, + ) + .await?; + Ok(self.track_response_state(stream)) + } + + /// The same entry point when the transport was not compiled in. + /// + /// Selecting [`ResponsesTransport::WebSocket`] is an explicit choice by the + /// caller, so a build without the transport reports that it cannot honor it + /// and names the feature that would. Quietly answering over HTTP+SSE instead + /// would return a stream the caller never asked for, and hide a + /// misconfigured build behind a working one. + #[cfg(not(feature = "responses-websocket"))] + async fn stream_websocket_response( + &self, + _request: ResponsesRequest, + _credentials: &ProviderCredentials, + _session: &SessionRequestOptions, + _reasoning_provenance: ReasoningProvenance, + ) -> Result { + Err(ProviderError::UnsupportedCapability( + "responses_websocket (not compiled in: rebuild mentra-provider with the \ + `responses-websocket` feature)" + .to_string(), + )) + } + + pub async fn send_response<'a>(&self, request: Request<'a>) -> Result { + crate::collect_response_from_stream(self.stream_response(request).await?).await + } + + async fn send_http_responses_request( + &self, + request: &ResponsesRequest, + compression: ResponsesRequestCompression, + credentials: &ProviderCredentials, + session: &SessionRequestOptions, + ) -> Result { + let request_builder = self + .client + .post( + self.definition + .request_url_with_auth_for_path(self.responses_endpoint_path(), credentials)?, + ) + .headers(self.build_http_headers_for_session(credentials, Some(session))?) + .header(reqwest::header::ACCEPT, "text/event-stream"); + + match compression { + ResponsesRequestCompression::None => request_builder + .json(request) + .send() + .await + .map_err(ProviderError::Transport), + ResponsesRequestCompression::Zstd => { + let body = serde_json::to_vec(request).map_err(ProviderError::Serialize)?; + let compressed = + zstd::stream::encode_all(std::io::Cursor::new(body), 3).map_err(|error| { + ProviderError::InvalidRequest(format!( + "failed to compress responses request: {error}" + )) + })?; + request_builder + .header(reqwest::header::CONTENT_ENCODING, "zstd") + .header(reqwest::header::CONTENT_TYPE, "application/json") + .body(compressed) + .send() + .await + .map_err(ProviderError::Transport) + } + } + } + + fn track_response_state(&self, mut stream: ProviderEventStream) -> ProviderEventStream { + let (tx, rx) = mpsc::unbounded_channel(); + let state = Arc::clone(&self.state); + + tokio::spawn(async move { + while let Some(event) = stream.recv().await { + if let Ok(ProviderEvent::MessageStarted { id, .. }) = &event { + state.set_latest_response_id(id.clone()); + } + + if tx.send(event).is_err() { + break; + } + } + }); + + rx + } + + pub fn take_turn_state(&self) -> SharedTurnState { + Arc::clone(&self.state.turn_state) + } + + pub fn last_response_rx_ready(&self) -> bool { + let mut session = self + .state + .websocket_session + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + session + .last_response_rx + .as_mut() + .is_some_and(|rx| matches!(rx.try_recv(), Ok(_) | Err(TryRecvError::Closed))) + } + + fn build_http_headers_for_session( + &self, + credentials: &ProviderCredentials, + session: Option<&SessionRequestOptions>, + ) -> Result { + let mut headers = self.build_websocket_headers_for_session(credentials, session)?; + if let Some(value) = session + .and_then(|session| session.session_affinity.as_deref()) + .and_then(|session_id| http::HeaderValue::from_str(session_id).ok()) + { + headers.insert("x-client-request-id", value.clone()); + headers.insert("session_id", value); + } + if let Some(value) = session + .and_then(|session| session.subagent.as_deref()) + .and_then(|subagent| http::HeaderValue::from_str(subagent).ok()) + { + headers.insert("x-openai-subagent", value); + } + Ok(headers) + } +} + +fn responses_endpoint_path_for_base(base_url: Option<&str>) -> &'static str { + let Some(base_url) = base_url else { + return "v1/responses"; + }; + let Ok(url) = Url::parse(base_url) else { + return "v1/responses"; + }; + + let path = url.path().trim_end_matches('/'); + if path.ends_with("/v1") || path == "/backend-api/codex" { + "responses" + } else { + "v1/responses" + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum PreviousResponseStateRejection { + ReferenceUnavailable, + ParameterUnsupported, +} + +fn previous_response_parameter_is_unsupported(body: &str) -> bool { + let normalized = body + .chars() + .map(|character| { + if character.is_ascii_alphanumeric() || character == '_' { + character.to_ascii_lowercase() + } else { + ' ' + } + }) + .collect::(); + let words = normalized.split_whitespace().collect::>(); + + words.windows(3).any(|window| { + matches!( + window, + [ + "unsupported" | "unknown" | "unrecognized", + "parameter", + "previous_response_id" + ] + ) + }) || words.windows(5).any(|window| { + matches!( + window, + [ + "parameter", + "previous_response_id", + "is", + "not", + "supported" + ] | [ + "previous_response_id", + "parameter", + "is", + "not", + "supported" + ] + ) + }) || words + .windows(4) + .any(|window| matches!(window, ["previous_response_id", "is", "not", "supported"])) +} + +fn previous_response_state_rejection( + error: &ProviderError, +) -> Option { + let ProviderError::Http { status, body, .. } = error else { + return None; + }; + if !(*status == reqwest::StatusCode::BAD_REQUEST || *status == reqwest::StatusCode::NOT_FOUND) { + return None; + } + + let body = body.to_ascii_lowercase(); + if previous_response_parameter_is_unsupported(&body) { + return Some(PreviousResponseStateRejection::ParameterUnsupported); + } + + (body.contains("previous_response_id") + || body.contains("previous response") + || (body.contains("response") && body.contains("not found")) + || (body.contains("response") && body.contains("expired"))) + .then_some(PreviousResponseStateRejection::ReferenceUnavailable) +} + +#[async_trait::async_trait] +impl ProviderSession for ResponsesSession +where + C: CredentialSource + 'static, +{ + async fn stream(&self, request: Request<'_>) -> Result { + self.stream_response(request).await + } + + async fn compact( + &self, + request: CompactionRequest<'_>, + ) -> Result { + let request = request.into_model_request()?; + let response = self.send_response(request).await?; + Ok(response.into_compaction_response()) + } + + async fn summarize_memories( + &self, + request: MemorySummarizeRequest<'_>, + ) -> Result { + let request = request.into_model_request()?; + let response = self.send_response(request).await?; + response.into_memory_summarize_response() + } +} + +#[cfg(test)] +mod tests { + use std::borrow::Cow; + use std::collections::BTreeMap; + use std::io::Read; + use std::io::Write; + use std::net::TcpListener; + use std::thread; + + use super::*; + use crate::ProviderId; + use crate::ProviderRequestOptions; + use crate::StaticCredentialSource; + use crate::responses::ResponsesProvider; + + fn spawn_single_response_server( + response_body: &'static str, + ) -> (String, thread::JoinHandle) { + spawn_single_response_server_with_headers(response_body, "") + } + + fn spawn_single_response_server_with_headers( + response_body: &'static str, + extra_headers: &'static str, + ) -> (String, thread::JoinHandle) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server"); + let addr = listener.local_addr().expect("read listener addr"); + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("accept request"); + let mut buffer = Vec::new(); + let mut temp = [0_u8; 1024]; + let mut header_end = None; + let mut content_length = 0_usize; + + loop { + let read = stream.read(&mut temp).expect("read request"); + if read == 0 { + break; + } + buffer.extend_from_slice(&temp[..read]); + let discovered_header = if header_end.is_none() { + buffer.windows(4).position(|window| window == b"\r\n\r\n") + } else { + None + }; + if let Some(index) = discovered_header { + let end = index + 4; + header_end = Some(end); + let headers = String::from_utf8_lossy(&buffer[..end]); + content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length").then(|| { + value.trim().parse::().expect("parse content-length") + }) + }) + .unwrap_or_default(); + } + if header_end.is_some_and(|end| buffer.len() >= end + content_length) { + break; + } + } + + let response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "content-type: text/event-stream\r\n", + "{}", + "content-length: {}\r\n\r\n", + "{}" + ), + extra_headers, + response_body.len(), + response_body, + ); + stream + .write_all(response.as_bytes()) + .expect("write response"); + String::from_utf8(buffer).expect("request should be utf8") + }); + + (format!("http://{addr}/"), handle) + } + + fn read_http_request(stream: &mut std::net::TcpStream) -> String { + let mut buffer = Vec::new(); + let mut temp = [0_u8; 1024]; + let mut header_end = None; + let mut content_length = 0_usize; + + loop { + let read = stream.read(&mut temp).expect("read request"); + if read == 0 { + break; + } + buffer.extend_from_slice(&temp[..read]); + let discovered_header = if header_end.is_none() { + buffer.windows(4).position(|window| window == b"\r\n\r\n") + } else { + None + }; + if let Some(index) = discovered_header { + let end = index + 4; + header_end = Some(end); + let headers = String::from_utf8_lossy(&buffer[..end]); + content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().expect("parse content-length")) + }) + .unwrap_or_default(); + } + if header_end.is_some_and(|end| buffer.len() >= end + content_length) { + break; + } + } + + String::from_utf8(buffer).expect("request should be utf8") + } + + fn request_body(captured: &str) -> &str { + captured.split("\r\n\r\n").nth(1).unwrap_or_default() + } + + fn spawn_hybrid_fallback_server( + rejection_body: &'static str, + successful_requests: usize, + ) -> (String, thread::JoinHandle>) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server"); + let addr = listener.local_addr().expect("read listener addr"); + let handle = thread::spawn(move || { + let (mut first_stream, _) = listener.accept().expect("accept first request"); + let first = read_http_request(&mut first_stream); + let first_response = format!( + concat!( + "HTTP/1.1 400 Bad Request\r\n", + "connection: close\r\n", + "content-type: application/json\r\n", + "content-length: {}\r\n\r\n", + "{}" + ), + rejection_body.len(), + rejection_body + ); + first_stream + .write_all(first_response.as_bytes()) + .expect("write first response"); + drop(first_stream); + + let mut captured = vec![first]; + for response_index in 1..=successful_requests { + let (mut stream, _) = listener.accept().expect("accept successful request"); + captured.push(read_http_request(&mut stream)); + let response_id = format!("resp_fresh_{response_index}"); + let response_body = format!( + concat!( + "data: {{\"type\":\"response.created\",\"response\":{{\"id\":\"{}\",\"model\":\"gpt-5\",\"status\":\"in_progress\"}}}}\n\n", + "data: {{\"type\":\"response.completed\",\"response\":{{\"id\":\"{}\",\"model\":\"gpt-5\",\"status\":\"completed\"}}}}\n\n" + ), + response_id, response_id + ); + let response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "connection: close\r\n", + "content-type: text/event-stream\r\n", + "content-length: {}\r\n\r\n", + "{}" + ), + response_body.len(), + response_body + ); + stream + .write_all(response.as_bytes()) + .expect("write successful response"); + } + + captured + }); + + (format!("http://{addr}/"), handle) + } + + async fn consume_stream(mut stream: ProviderEventStream) { + while let Some(event) = stream.recv().await { + event.expect("stream event should decode"); + } + } + + fn spawn_two_turn_state_server() -> (String, thread::JoinHandle<(String, String)>) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server"); + let addr = listener.local_addr().expect("read listener addr"); + let handle = thread::spawn(move || { + let (mut first_stream, _) = listener.accept().expect("accept first request"); + let first = read_http_request(&mut first_stream); + let first_body = concat!( + "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"in_progress\"}}\n\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n" + ); + let first_response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "connection: close\r\n", + "content-type: text/event-stream\r\n", + "x-codex-turn-state: state-1\r\n", + "content-length: {}\r\n\r\n", + "{}" + ), + first_body.len(), + first_body + ); + first_stream + .write_all(first_response.as_bytes()) + .expect("write first response"); + drop(first_stream); + + let (mut second_stream, _) = listener.accept().expect("accept second request"); + let second = read_http_request(&mut second_stream); + let second_body = concat!( + "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_2\",\"model\":\"gpt-5\",\"status\":\"in_progress\"}}\n\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_2\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n" + ); + let second_response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "connection: close\r\n", + "content-type: text/event-stream\r\n", + "x-codex-turn-state: state-2\r\n", + "content-length: {}\r\n\r\n", + "{}" + ), + second_body.len(), + second_body + ); + second_stream + .write_all(second_response.as_bytes()) + .expect("write second response"); + + (first, second) + }); + + (format!("http://{addr}/"), handle) + } + + fn spawn_compaction_response_server( + response_body: &'static str, + ) -> (String, thread::JoinHandle<(String, String)>) { + let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server"); + let addr = listener.local_addr().expect("read listener addr"); + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("accept request"); + let mut buffer = Vec::new(); + let mut temp = [0_u8; 1024]; + let mut header_end = None; + let mut content_length = 0_usize; + + loop { + let read = stream.read(&mut temp).expect("read request"); + if read == 0 { + break; + } + buffer.extend_from_slice(&temp[..read]); + let discovered_header = if header_end.is_none() { + buffer.windows(4).position(|window| window == b"\r\n\r\n") + } else { + None + }; + if let Some(index) = discovered_header { + let end = index + 4; + header_end = Some(end); + let headers = String::from_utf8_lossy(&buffer[..end]); + content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length").then(|| { + value.trim().parse::().expect("parse content-length") + }) + }) + .unwrap_or_default(); + } + if header_end.is_some_and(|end| buffer.len() >= end + content_length) { + break; + } + } + + let response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "content-type: text/event-stream\r\n", + "content-length: {}\r\n\r\n", + "{}" + ), + response_body.len(), + response_body + ); + stream + .write_all(response.as_bytes()) + .expect("write response"); + let captured = String::from_utf8(buffer).expect("request should be utf8"); + let body = captured + .split("\r\n\r\n") + .nth(1) + .unwrap_or_default() + .to_string(); + (captured, body) + }); + + (format!("http://{addr}/"), handle) + } + + #[tokio::test] + async fn stream_response_honors_session_request_options_on_http_path() { + let sse_body = "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n"; + let (base_url, handle) = spawn_single_response_server(sse_body); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + session.set_turn_state("sticky-turn-state"); + + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: Some(Cow::Borrowed("system")), + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + session: SessionRequestOptions { + sticky_turn_state: None, + turn_metadata: Some("{\"turn_id\":\"turn-123\"}".to_string()), + subagent: Some("memory_consolidation".to_string()), + prefer_connection_reuse: Some(true), + session_affinity: Some("session-affinity-123".to_string()), + extra_headers: BTreeMap::new(), + }, + ..ProviderRequestOptions::default() + }, + }; + + let _stream = session + .stream_response(request) + .await + .expect("stream response should succeed"); + + let captured = handle.join().expect("server should capture request"); + assert!(captured.contains("x-codex-turn-state: sticky-turn-state\r\n")); + assert!(captured.contains("x-codex-turn-metadata: {\"turn_id\":\"turn-123\"}\r\n")); + assert!(captured.contains("x-mentra-turn-metadata: {\"turn_id\":\"turn-123\"}\r\n")); + assert!(captured.contains("x-mentra-session-affinity: session-affinity-123\r\n")); + assert!(captured.contains("x-mentra-connection-reuse: prefer-reuse\r\n")); + assert!(captured.contains("x-client-request-id: session-affinity-123\r\n")); + assert!(captured.contains("session_id: session-affinity-123\r\n")); + assert!(captured.contains("x-openai-subagent: memory_consolidation\r\n")); + } + + #[tokio::test] + async fn live_http_path_threads_custom_provider_and_requested_model_provenance() { + let sse_body = concat!( + "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-resolved\",\"status\":\"in_progress\"}}\n\n", + "data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"reasoning\",\"id\":\"rs_new\",\"summary\":[]}}\n\n", + "data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"reasoning\",\"id\":\"rs_new\",\"encrypted_content\":\"encrypted-new\",\"summary\":[{\"type\":\"summary_text\",\"text\":\"new summary\"}]}}\n\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-resolved\",\"status\":\"completed\"}}\n\n" + ); + let (base_url, handle) = spawn_single_response_server(sse_body); + let provider_id = ProviderId::new("openai-edge"); + let requested_model = "gpt-requested"; + let mut definition = super::super::openai_definition(); + definition.descriptor.id = provider_id.clone(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + let request = Request { + model: Cow::Borrowed(requested_model), + system: None, + messages: Cow::Owned(vec![crate::Message::assistant( + crate::ContentBlock::Thinking { + thinking: "old summary".to_string(), + signature: None, + encrypted_content: Some("encrypted-old".to_string()), + id: Some("rs_old".to_string()), + provenance: Some(crate::ReasoningProvenance { + provider: provider_id.clone(), + model: requested_model.to_string(), + format: crate::ReasoningFormat::OpenAiEncrypted, + }), + redacted: false, + }, + )]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + reasoning: Some(crate::ReasoningOptions { + effort: Some(crate::ReasoningEffort::Medium), + summary: Some(crate::ReasoningSummary::Auto), + }), + ..ProviderRequestOptions::default() + }, + }; + + let response = session + .send_response(request) + .await + .expect("custom Responses request should complete"); + let captured = handle.join().expect("server should capture request"); + let body: serde_json::Value = + serde_json::from_str(request_body(&captured)).expect("request body should be JSON"); + + assert_eq!(body["input"][0]["type"], "reasoning"); + assert_eq!(body["input"][0]["id"], "rs_old"); + assert_eq!(body["include"][0], "reasoning.encrypted_content"); + assert_eq!( + response.content, + vec![crate::ContentBlock::Thinking { + thinking: "new summary".to_string(), + signature: None, + encrypted_content: Some("encrypted-new".to_string()), + id: Some("rs_new".to_string()), + provenance: Some(crate::ReasoningProvenance { + provider: provider_id, + model: requested_model.to_string(), + format: crate::ReasoningFormat::OpenAiEncrypted, + }), + redacted: false, + }] + ); + } + + #[tokio::test] + async fn stream_response_captures_turn_state_from_http_response_headers() { + let sse_body = "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n"; + let (base_url, _handle) = spawn_single_response_server_with_headers( + sse_body, + "x-codex-turn-state: next-turn-state\r\n", + ); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: Some(Cow::Borrowed("system")), + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let _stream = session + .stream_response(request) + .await + .expect("stream response should succeed"); + + assert_eq!(session.turn_state().as_deref(), Some("next-turn-state")); + } + + #[tokio::test] + async fn http_turn_state_updates_across_multiple_turns() { + let (base_url, handle) = spawn_two_turn_state_server(); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + for message in ["first", "second"] { + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + message, + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let mut stream = session + .stream_response(request) + .await + .expect("stream response should succeed"); + while let Some(event) = stream.recv().await { + event.expect("stream event should decode"); + } + } + + let (first, second) = handle.join().expect("server should capture requests"); + assert!(!first.contains("x-codex-turn-state:")); + assert!(second.contains("x-codex-turn-state: state-1\r\n")); + assert_eq!(session.turn_state().as_deref(), Some("state-2")); + assert_eq!(session.latest_response_id().as_deref(), Some("resp_2")); + } + + #[tokio::test] + async fn predeclared_unsupported_endpoint_skips_the_first_hybrid_http_probe() { + let (base_url, handle) = spawn_two_turn_state_server(); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let provider = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .without_hybrid_http_previous_response_id(); + + for message in ["first", "second"] { + let session = provider.session(); + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + message, + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + consume_stream( + session + .stream_response(request) + .await + .expect("hybrid replay should not probe unsupported provider state"), + ) + .await; + } + + let (first, second) = handle.join().expect("server should capture requests"); + let first_payload: serde_json::Value = + serde_json::from_str(request_body(&first)).expect("first body should be json"); + let second_payload: serde_json::Value = + serde_json::from_str(request_body(&second)).expect("second body should be json"); + assert!(first_payload.get("previous_response_id").is_none()); + assert!(second_payload.get("previous_response_id").is_none()); + assert_eq!( + provider.session().latest_response_id().as_deref(), + Some("resp_2") + ); + } + + #[tokio::test] + async fn stream_response_tracks_latest_response_id_from_http_events() { + let sse_body = "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"in_progress\"}}\n\n\ +data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n"; + let (base_url, _handle) = spawn_single_response_server(sse_body); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let mut stream = session + .stream_response(request) + .await + .expect("stream response should succeed"); + while let Some(event) = stream.recv().await { + event.expect("stream event should decode"); + } + + assert_eq!(session.latest_response_id().as_deref(), Some("resp_1")); + } + + #[test] + fn responses_endpoint_path_supports_openai_and_xipe_base_urls() { + assert_eq!( + responses_endpoint_path_for_base(Some("https://api.openai.com/")), + "v1/responses" + ); + assert_eq!( + responses_endpoint_path_for_base(Some("https://api.openai.com/v1")), + "responses" + ); + assert_eq!( + responses_endpoint_path_for_base(Some("https://opencode.ai/zen/go/v1")), + "responses" + ); + assert_eq!( + responses_endpoint_path_for_base(Some("https://chatgpt.com/backend-api/codex")), + "responses" + ); + } + + #[cfg(feature = "responses-websocket")] + #[tokio::test] + async fn websocket_transport_sends_response_create_frame() { + use futures_util::{SinkExt, StreamExt}; + use http::HeaderValue; + use tokio_tungstenite::accept_hdr_async; + use tokio_tungstenite::tungstenite::Message as WsMessage; + use tokio_tungstenite::tungstenite::handshake::server::{ + ErrorResponse, Request as WsHandshakeRequest, Response as WsHandshakeResponse, + }; + + #[expect( + clippy::result_large_err, + reason = "tokio-tungstenite fixes the handshake callback error type" + )] + fn add_turn_state_header( + _request: &WsHandshakeRequest, + mut response: WsHandshakeResponse, + ) -> Result { + response + .headers_mut() + .insert("x-codex-turn-state", HeaderValue::from_static("ws-state")); + Ok(response) + } + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind websocket test server"); + let addr = listener.local_addr().expect("read websocket server addr"); + let (tx_frame, rx_frame) = tokio::sync::oneshot::channel::(); + + tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept websocket"); + let mut ws = accept_hdr_async(stream, add_turn_state_header) + .await + .expect("upgrade websocket"); + let frame = ws + .next() + .await + .expect("client should send a frame") + .expect("websocket frame should be valid") + .into_text() + .expect("request frame should be text") + .to_string(); + tx_frame.send(frame).expect("send captured frame"); + + ws.send(WsMessage::Text( + serde_json::json!({ + "type": "response.created", + "response": { + "id": "resp_ws", + "model": "gpt-resolved", + "status": "in_progress" + } + }) + .to_string() + .into(), + )) + .await + .expect("send response.created"); + ws.send(WsMessage::Text( + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 0, + "item": { + "type": "reasoning", + "id": "rs_ws", + "summary": [] + } + }) + .to_string() + .into(), + )) + .await + .expect("send reasoning start"); + ws.send(WsMessage::Text( + serde_json::json!({ + "type": "response.output_item.done", + "output_index": 0, + "item": { + "type": "reasoning", + "id": "rs_ws", + "encrypted_content": "encrypted-ws", + "summary": [{"type": "summary_text", "text": "ws summary"}] + } + }) + .to_string() + .into(), + )) + .await + .expect("send reasoning completion"); + ws.send(WsMessage::Text( + serde_json::json!({ + "type": "response.completed", + "response": { + "id": "resp_ws", + "model": "gpt-resolved", + "status": "completed" + } + }) + .to_string() + .into(), + )) + .await + .expect("send response.completed"); + }); + + let mut definition = super::super::openai_definition(); + definition.descriptor.id = ProviderId::new("openai-ws-edge"); + definition.base_url = Some(format!("http://{addr}/v1")); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + let request = Request { + model: Cow::Borrowed("gpt-requested"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + responses: crate::ResponsesRequestOptions { + transport: crate::ResponsesTransport::WebSocket, + ..Default::default() + }, + ..ProviderRequestOptions::default() + }, + }; + + let mut stream = session + .stream_response(request) + .await + .expect("websocket transport should stream"); + let mut events = Vec::new(); + while let Some(event) = stream.recv().await { + events.push(event.expect("websocket event should parse")); + } + + let frame: serde_json::Value = + serde_json::from_str(&rx_frame.await.expect("server should capture frame")) + .expect("frame should be json"); + assert_eq!(frame["type"], "response.create"); + assert_eq!(frame["model"], "gpt-requested"); + assert_eq!(frame["input"][0]["role"], "user"); + assert_eq!(frame["instructions"], ""); + assert!(frame.get("response").is_none()); + assert_eq!(session.turn_state().as_deref(), Some("ws-state")); + assert_eq!(session.latest_response_id().as_deref(), Some("resp_ws")); + assert!(events.iter().any(|event| matches!( + event, + ProviderEvent::ContentBlockStarted { + kind: crate::ContentBlockStart::Thinking { + id: Some(id), + provenance: Some(crate::ReasoningProvenance { + provider, + model, + format: crate::ReasoningFormat::OpenAiEncrypted, + }), + .. + }, + .. + } if id == "rs_ws" + && provider == &ProviderId::new("openai-ws-edge") + && model == "gpt-requested" + ))); + } + + #[tokio::test] + async fn hybrid_state_falls_back_to_replay_when_previous_response_id_is_rejected() { + let (base_url, handle) = spawn_hybrid_fallback_server( + r#"{"error":{"message":"previous_response_id not found"}}"#, + 2, + ); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + session.state.set_latest_response_id("resp_stale"); + + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let stream = session + .stream_response(request.clone()) + .await + .expect("hybrid fallback should retry without provider state"); + consume_stream(stream).await; + + let stream = session + .stream_response(request) + .await + .expect("fresh provider state should remain usable"); + consume_stream(stream).await; + + let captured = handle.join().expect("server should capture requests"); + let first_payload: serde_json::Value = + serde_json::from_str(request_body(&captured[0])).expect("first body should be json"); + let second_payload: serde_json::Value = + serde_json::from_str(request_body(&captured[1])).expect("second body should be json"); + let third_payload: serde_json::Value = + serde_json::from_str(request_body(&captured[2])).expect("third body should be json"); + assert_eq!(first_payload["previous_response_id"], "resp_stale"); + assert!(second_payload.get("previous_response_id").is_none()); + assert_eq!(third_payload["previous_response_id"], "resp_fresh_1"); + assert_eq!( + session.latest_response_id().as_deref(), + Some("resp_fresh_2") + ); + } + + #[test] + fn classifies_previous_response_state_rejections_conservatively() { + let unsupported = ProviderError::Http { + status: reqwest::StatusCode::BAD_REQUEST, + body: r#"{"detail":"Unsupported parameter: previous_response_id"}"#.to_string(), + retry_after: None, + }; + assert_eq!( + previous_response_state_rejection(&unsupported), + Some(PreviousResponseStateRejection::ParameterUnsupported) + ); + + let stale = ProviderError::Http { + status: reqwest::StatusCode::NOT_FOUND, + body: r#"{"error":{"message":"previous_response_id expired"}}"#.to_string(), + retry_after: None, + }; + assert_eq!( + previous_response_state_rejection(&stale), + Some(PreviousResponseStateRejection::ReferenceUnavailable) + ); + + let unrelated = ProviderError::Http { + status: reqwest::StatusCode::BAD_REQUEST, + body: r#"{"detail":"Unsupported parameter: temperature"}"#.to_string(), + retry_after: None, + }; + assert_eq!(previous_response_state_rejection(&unrelated), None); + + let echoed_previous_response_id = ProviderError::Http { + status: reqwest::StatusCode::BAD_REQUEST, + body: r#"{"detail":"Unsupported parameter: temperature","request":{ + "previous_response_id":"resp_1"}}"# + .to_string(), + retry_after: None, + }; + assert_eq!( + previous_response_state_rejection(&echoed_previous_response_id), + Some(PreviousResponseStateRejection::ReferenceUnavailable) + ); + } + + #[tokio::test] + async fn hybrid_state_remembers_when_previous_response_id_is_unsupported() { + let (base_url, handle) = spawn_hybrid_fallback_server( + r#"{"detail":"Unsupported parameter: previous_response_id"}"#, + 4, + ); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let provider = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ); + let session = provider.session(); + session.state.set_latest_response_id("resp_stale"); + + let request = Request { + model: Cow::Borrowed("gpt-5"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions::default(), + }; + + let stream = session + .stream_response(request.clone()) + .await + .expect("hybrid fallback should retry without unsupported provider state"); + consume_stream(stream).await; + + let next_session = provider.session(); + let stream = next_session + .stream_response(request.clone()) + .await + .expect("same-model replay should skip unsupported provider state"); + consume_stream(stream).await; + + let mut other_model_request = request.clone(); + other_model_request.model = Cow::Borrowed("gpt-5-mini"); + let stream = next_session + .stream_response(other_model_request) + .await + .expect("another model should keep opportunistic provider state"); + consume_stream(stream).await; + + let mut stateful_request = request; + stateful_request + .provider_request_options + .responses + .state_mode = crate::ResponsesStateMode::Stateful; + let stream = next_session + .stream_response(stateful_request) + .await + .expect("stateful mode should ignore the hybrid capability cache"); + consume_stream(stream).await; + + let captured = handle.join().expect("server should capture requests"); + let payloads = captured + .iter() + .map(|request| { + serde_json::from_str::(request_body(request)) + .expect("request body should be json") + }) + .collect::>(); + assert_eq!(payloads[0]["previous_response_id"], "resp_stale"); + assert!(payloads[1].get("previous_response_id").is_none()); + assert!(payloads[2].get("previous_response_id").is_none()); + assert_eq!(payloads[3]["previous_response_id"], "resp_fresh_2"); + assert_eq!(payloads[4]["previous_response_id"], "resp_fresh_3"); + assert_eq!( + next_session.latest_response_id().as_deref(), + Some("resp_fresh_4") + ); + } + + #[tokio::test] + async fn compact_sends_normal_model_request_and_wraps_summary_text() { + let sse_body = concat!( + "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"in_progress\"}}\n\n", + "data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"message\",\"content\":[]}}\n\n", + "data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"{\\\"goal\\\":\\\"keep going\\\"}\"}]}}\n\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n" + ); + let (base_url, handle) = spawn_compaction_response_server(sse_body); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + let request = crate::CompactionRequest { + model: Cow::Borrowed("gpt-5"), + instructions: Cow::Borrowed("Summarize the transcript."), + input: Cow::Owned(vec![crate::CompactionInputItem::UserTurn { + content: "hello".to_string(), + }]), + metadata: Cow::Owned(BTreeMap::from([("scope".to_string(), "test".to_string())])), + provider_request_options: crate::ProviderRequestOptions::default(), + }; + + let response = session.compact(request).await.expect("compaction succeeds"); + let captured = handle.join().expect("server should capture request"); + + let payload: serde_json::Value = + serde_json::from_str(&captured.1).expect("request body should be json"); + assert_eq!(payload["model"], "gpt-5"); + assert_eq!(payload["instructions"], "Summarize the transcript."); + assert_eq!(payload["metadata"]["scope"], "test"); + assert_eq!(payload["input"][0]["content"][0]["type"], "input_text"); + assert!( + payload["input"][0]["content"][0]["text"] + .as_str() + .expect("prompt text should be a string") + .starts_with("Compaction input JSON:\n") + ); + + assert_eq!(response.output.len(), 1); + assert_eq!( + response.output[0], + crate::CompactionInputItem::CompactionSummary { + content: "{\"goal\":\"keep going\"}".to_string() + } + ); + } + + #[tokio::test] + async fn summarize_memories_sends_normal_model_request_and_parses_json_output() { + let sse_body = concat!( + "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_2\",\"model\":\"gpt-5\",\"status\":\"in_progress\"}}\n\n", + "data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"message\",\"content\":[]}}\n\n", + "data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"[{\\\"raw_memory\\\":\\\"Detailed summary\\\",\\\"memory_summary\\\":\\\"Short summary\\\"}]\"}]}}\n\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_2\",\"model\":\"gpt-5\",\"status\":\"completed\"}}\n\n" + ); + let (base_url, handle) = spawn_compaction_response_server(sse_body); + + let mut definition = super::super::openai_definition(); + definition.base_url = Some(base_url); + let session = ResponsesProvider::with_shared_credential_source( + definition, + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + let request = crate::MemorySummarizeRequest { + model: Cow::Borrowed("gpt-5"), + raw_memories: Cow::Owned(vec![crate::RawMemory { + id: "memory-1".to_string(), + metadata: crate::RawMemoryMetadata { + source_path: "/tmp/trace.jsonl".to_string(), + }, + items: vec![serde_json::json!({"type":"message","role":"user"})], + }]), + reasoning: Some(crate::ReasoningOptions { + effort: Some(crate::ReasoningEffort::Medium), + summary: None, + }), + metadata: Cow::Owned(BTreeMap::from([("scope".to_string(), "test".to_string())])), + provider_request_options: crate::ProviderRequestOptions { + session: crate::SessionRequestOptions { + sticky_turn_state: None, + turn_metadata: Some("{\"turn_id\":\"turn-321\"}".to_string()), + subagent: Some("compact".to_string()), + prefer_connection_reuse: Some(true), + session_affinity: Some("session-affinity-321".to_string()), + extra_headers: BTreeMap::new(), + }, + ..crate::ProviderRequestOptions::default() + }, + }; + + let response = session + .summarize_memories(request) + .await + .expect("memory summarization succeeds"); + let captured = handle.join().expect("server should capture request"); + + let payload: serde_json::Value = + serde_json::from_str(&captured.1).expect("request body should be json"); + assert_eq!(payload["model"], "gpt-5"); + assert_eq!(payload["reasoning"]["effort"], "medium"); + assert_eq!(payload["metadata"]["scope"], "test"); + assert_eq!(payload["input"][0]["content"][0]["type"], "input_text"); + assert!( + payload["input"][0]["content"][0]["text"] + .as_str() + .expect("prompt text should be a string") + .starts_with("Memory summarize input JSON:\n") + ); + + assert_eq!(response.output.len(), 1); + assert_eq!(response.output[0].raw_memory, "Detailed summary"); + assert_eq!(response.output[0].memory_summary, "Short summary"); + assert!( + captured + .0 + .contains("x-codex-turn-metadata: {\"turn_id\":\"turn-321\"}\r\n") + ); + assert!( + captured + .0 + .contains("x-mentra-session-affinity: session-affinity-321\r\n") + ); + assert!( + captured + .0 + .contains("x-client-request-id: session-affinity-321\r\n") + ); + assert!(captured.0.contains("session_id: session-affinity-321\r\n")); + assert!(captured.0.contains("x-openai-subagent: compact\r\n")); + } + + #[cfg(feature = "responses-websocket")] + #[tokio::test] + async fn websocket_connection_is_closed_without_cached_connection() { + let session = ResponsesProvider::with_shared_credential_source( + super::super::openai_definition(), + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + assert!(session.websocket_connection_is_closed().await); + session.clear_websocket_connection(); + assert!(session.websocket_connection_is_closed().await); + } + + /// A build without the transport must say so, and name the way back. + /// + /// The failure has to be visible at the call that asked for it: no panic, + /// and no silent demotion to HTTP+SSE, either of which would let a build + /// that cannot do what was configured pass for one that can. `openai` is + /// the preset whose capabilities advertise websockets, so this reaches the + /// missing transport rather than the runtime `websockets_enabled` check. + #[cfg(not(feature = "responses-websocket"))] + #[tokio::test] + async fn websocket_transport_without_the_feature_is_an_honest_error() { + let session = ResponsesProvider::with_shared_credential_source( + super::super::openai_definition(), + Arc::new(StaticCredentialSource::new("test-key")), + ) + .session(); + + let request = Request { + model: Cow::Borrowed("gpt-requested"), + system: None, + messages: Cow::Owned(vec![crate::Message::user(crate::ContentBlock::text( + "hello", + ))]), + tools: Cow::Owned(Vec::new()), + tool_choice: None, + temperature: None, + max_output_tokens: None, + metadata: Cow::Owned(BTreeMap::new()), + provider_request_options: ProviderRequestOptions { + responses: crate::ResponsesRequestOptions { + transport: crate::ResponsesTransport::WebSocket, + ..Default::default() + }, + ..ProviderRequestOptions::default() + }, + }; + + let error = session + .stream_response(request) + .await + .err() + .expect("a transport that was not compiled in cannot stream"); + + assert!( + matches!(error, ProviderError::UnsupportedCapability(_)), + "expected an unsupported-capability error, got {error:?}" + ); + assert_eq!( + error.to_string(), + "provider does not support capability: responses_websocket (not compiled in: \ + rebuild mentra-provider with the `responses-websocket` feature)", + "the message is the only thing that tells an operator how to fix the build" + ); + } +} diff --git a/vendor/mentra-provider/src/responses/sse.rs b/vendor/mentra-provider/src/responses/sse.rs new file mode 100644 index 0000000..b7726cf --- /dev/null +++ b/vendor/mentra-provider/src/responses/sse.rs @@ -0,0 +1,1542 @@ +use std::collections::{HashMap, HashSet}; + +use futures_util::StreamExt; +use serde::Deserialize; +use tokio::sync::mpsc; + +use crate::{ + ContentBlockDelta, ContentBlockStart, HostedToolSearchCall, HostedWebSearchCall, + ImageGenerationCall, ImageGenerationResult, ProviderError, ProviderEvent, ProviderEventStream, + ProviderId, ReasoningFormat, ReasoningProvenance, ResponseHeaders, Role, TokenUsage, + WebSearchAction, +}; + +use super::model::encode_responses_tool_use_id; + +/// Spawns an event stream that decodes Responses SSE frames. +pub fn spawn_event_stream(response: reqwest::Response) -> ProviderEventStream { + spawn_event_stream_with_provenance(response, ProviderId::new("openai"), String::new()) +} + +pub(crate) fn spawn_event_stream_with_provenance( + response: reqwest::Response, + provider: ProviderId, + requested_model: String, +) -> ProviderEventStream { + let (tx, rx) = mpsc::unbounded_channel(); + + tokio::spawn(async move { + if let Err(error) = forward_events(response, tx.clone(), provider, requested_model).await { + let _ = tx.send(Err(error)); + } + }); + + rx +} + +async fn forward_events( + response: reqwest::Response, + tx: mpsc::UnboundedSender>, + provider: ProviderId, + requested_model: String, +) -> Result<(), ProviderError> { + let receiver_closed = response_headers_event(response.headers()) + .map(|headers| tx.send(Ok(headers)).is_err()) + .unwrap_or(false); + if receiver_closed { + return Ok(()); + } + + let mut bytes_stream = response.bytes_stream(); + let mut buffer = Vec::new(); + let mut state = StreamState::new(provider, requested_model); + + while let Some(chunk) = bytes_stream.next().await { + let chunk = chunk.map_err(ProviderError::Transport)?; + buffer.extend_from_slice(&chunk); + + while let Some((frame_end, delimiter_len)) = find_frame_boundary(&buffer) { + let frame = buffer.drain(..frame_end).collect::>(); + buffer.drain(..delimiter_len); + + for event in parse_frame(&frame, &mut state)? { + if tx.send(Ok(event)).is_err() { + return Ok(()); + } + } + } + } + + if !buffer.is_empty() { + for event in parse_frame(&buffer, &mut state)? { + let _ = tx.send(Ok(event)); + } + } + + Ok(()) +} + +pub(crate) fn response_headers_event(headers: &http::HeaderMap) -> Option { + let values = headers + .iter() + .filter_map(|(name, value)| { + value + .to_str() + .ok() + .map(|value| (name.as_str().to_string(), value.to_string())) + }) + .collect::>(); + + (!values.is_empty()).then_some(ProviderEvent::ResponseHeaders(ResponseHeaders { values })) +} + +pub(crate) struct StreamState { + ignored_output_indices: HashSet, + text_delta_seen: HashSet, + function_delta_seen: HashSet, + tool_search_delta_seen: HashSet, + reasoning_delta_seen: HashSet, + reasoning_encrypted_content_seen: HashSet, + reasoning_output_indices: HashMap, + reasoning_provenance: ReasoningProvenance, +} + +impl Default for StreamState { + fn default() -> Self { + Self::new(ProviderId::new("openai"), String::new()) + } +} + +impl StreamState { + pub(crate) fn new(provider: ProviderId, requested_model: String) -> Self { + Self { + ignored_output_indices: HashSet::new(), + text_delta_seen: HashSet::new(), + function_delta_seen: HashSet::new(), + tool_search_delta_seen: HashSet::new(), + reasoning_delta_seen: HashSet::new(), + reasoning_encrypted_content_seen: HashSet::new(), + reasoning_output_indices: HashMap::new(), + reasoning_provenance: ReasoningProvenance { + provider, + model: requested_model, + format: ReasoningFormat::OpenAiEncrypted, + }, + } + } + + fn backfill_reasoning_encrypted_content( + &mut self, + output: &[ResponsesOutputItem], + ) -> Vec { + let mut events = Vec::new(); + for item in output { + let Some((id, encrypted_content)) = item.reasoning_encrypted_content() else { + continue; + }; + let Some(index) = self.reasoning_output_indices.get(id).copied() else { + continue; + }; + if !self.reasoning_encrypted_content_seen.insert(index) { + continue; + } + events.push(ProviderEvent::ContentBlockDelta { + index, + delta: ContentBlockDelta::ThinkingEncryptedContent(encrypted_content.to_string()), + }); + } + events + } +} + +fn parse_frame(frame: &[u8], state: &mut StreamState) -> Result, ProviderError> { + let frame = std::str::from_utf8(frame) + .map_err(|error| ProviderError::MalformedStream(error.to_string()))?; + let mut data_lines = Vec::new(); + + for raw_line in frame.lines() { + let line = raw_line.strip_suffix('\r').unwrap_or(raw_line); + if line.is_empty() || line.starts_with(':') || line.starts_with("event:") { + continue; + } + + if let Some(rest) = line.strip_prefix("data:") { + data_lines.push(rest.trim_start().to_string()); + } + } + + if data_lines.is_empty() { + return Ok(Vec::new()); + } + + let data = data_lines.join("\n"); + parse_json_event(&data, state) +} + +pub(crate) fn parse_json_event( + data: &str, + state: &mut StreamState, +) -> Result, ProviderError> { + if data == "[DONE]" { + return Ok(Vec::new()); + } + + let event: ResponsesStreamEvent = + serde_json::from_str(data).map_err(ProviderError::Deserialize)?; + + match event { + ResponsesStreamEvent::ResponseCreated { response } => { + if state.reasoning_provenance.model.is_empty() { + state.reasoning_provenance.model = response.model.clone(); + } + let model = if response.model.is_empty() { + state.reasoning_provenance.model.clone() + } else { + response.model + }; + Ok(vec![ + ProviderEvent::ResponseCreated, + ProviderEvent::MessageStarted { + id: response.id, + model, + role: Role::Assistant, + }, + ]) + } + ResponsesStreamEvent::ResponseOutputItemAdded { output_index, item } => { + if let Some(id) = item.reasoning_id() { + state + .reasoning_output_indices + .insert(id.to_string(), output_index); + } + if item.reasoning_encrypted_content().is_some() { + state.reasoning_encrypted_content_seen.insert(output_index); + } + let pair_function_item_ids = !state.reasoning_output_indices.is_empty(); + if let Some(kind) = + item.into_provider_start(&state.reasoning_provenance, pair_function_item_ids) + { + Ok(vec![ProviderEvent::ContentBlockStarted { + index: output_index, + kind, + }]) + } else { + state.ignored_output_indices.insert(output_index); + Ok(Vec::new()) + } + } + ResponsesStreamEvent::ResponseOutputTextDelta { + output_index, + delta, + .. + } => { + if state.ignored_output_indices.contains(&output_index) { + return Ok(Vec::new()); + } + + state.text_delta_seen.insert(output_index); + Ok(vec![ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::Text(delta), + }]) + } + ResponsesStreamEvent::ResponseReasoningSummaryTextDelta { + output_index, + delta, + summary_index, + } => { + let mut events = vec![ProviderEvent::ReasoningSummaryDelta { + delta: delta.clone(), + summary_index, + }]; + let output_index = output_index + .filter(|output_index| !state.ignored_output_indices.contains(output_index)); + if let Some(output_index) = output_index { + state.reasoning_delta_seen.insert(output_index); + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ThinkingText(delta), + }); + } + Ok(events) + } + ResponsesStreamEvent::ResponseReasoningTextDelta { + output_index, + delta, + content_index, + } => { + let mut events = vec![ProviderEvent::ReasoningContentDelta { + delta: delta.clone(), + content_index, + }]; + let output_index = output_index + .filter(|output_index| !state.ignored_output_indices.contains(output_index)); + if let Some(output_index) = output_index { + state.reasoning_delta_seen.insert(output_index); + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ThinkingText(delta), + }); + } + Ok(events) + } + ResponsesStreamEvent::ResponseReasoningSummaryPartAdded { summary_index, .. } => { + Ok(vec![ProviderEvent::ReasoningSummaryPartAdded { + summary_index, + }]) + } + ResponsesStreamEvent::ResponseReasoningSummaryPartDone { output_index } => { + let output_index = output_index + .filter(|output_index| !state.ignored_output_indices.contains(output_index)); + if let Some(output_index) = output_index { + state.reasoning_delta_seen.insert(output_index); + return Ok(vec![ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ThinkingText("\n\n".to_string()), + }]); + } + Ok(Vec::new()) + } + ResponsesStreamEvent::ResponseFunctionCallArgumentsDelta { + output_index, + delta, + .. + } => { + if state.ignored_output_indices.contains(&output_index) { + return Ok(Vec::new()); + } + + state.function_delta_seen.insert(output_index); + Ok(vec![ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ToolUseInputJson(delta), + }]) + } + ResponsesStreamEvent::ResponseToolSearchCallDelta { + output_index, + delta, + } => { + if state.ignored_output_indices.contains(&output_index) { + return Ok(Vec::new()); + } + + state.tool_search_delta_seen.insert(output_index); + Ok(vec![ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::HostedToolSearchQuery(delta), + }]) + } + ResponsesStreamEvent::ResponseOutputItemDone { output_index, item } => { + if state.ignored_output_indices.remove(&output_index) { + return Ok(Vec::new()); + } + + let mut events = Vec::new(); + + if let Some(id) = item.reasoning_id() { + state + .reasoning_output_indices + .insert(id.to_string(), output_index); + } + + let completed_text = if state.text_delta_seen.remove(&output_index) { + None + } else { + item.completed_text().filter(|text| !text.is_empty()) + }; + if let Some(text) = completed_text { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::Text(text), + }); + } + + let completed_arguments = if state.function_delta_seen.remove(&output_index) { + None + } else { + item.completed_arguments() + .filter(|arguments| !arguments.is_empty()) + }; + if let Some(arguments) = completed_arguments { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ToolUseInputJson(arguments), + }); + } + + let completed_query = if state.tool_search_delta_seen.remove(&output_index) { + None + } else { + item.completed_tool_search_query() + .filter(|query| !query.is_empty()) + }; + if let Some(query) = completed_query { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::HostedToolSearchQuery(query), + }); + } + + let reasoning = if state.reasoning_delta_seen.remove(&output_index) { + None + } else { + item.completed_reasoning_text() + .filter(|reasoning| !reasoning.is_empty()) + }; + if let Some(reasoning) = reasoning { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ThinkingText(reasoning), + }); + } + + if item.reasoning_encrypted_content().is_some() { + state.reasoning_encrypted_content_seen.insert(output_index); + } + + events.extend(item.completion_deltas(output_index)); + + if item.is_supported() { + events.push(ProviderEvent::ContentBlockStopped { + index: output_index, + }); + } + + Ok(events) + } + ResponsesStreamEvent::ResponseCompleted { response } + | ResponsesStreamEvent::ResponseIncomplete { response } => { + let mut events = state.backfill_reasoning_encrypted_content(&response.output); + events.push(ProviderEvent::MessageDelta { + stop_reason: response.stop_reason(), + usage: response.usage(), + }); + events.push(ProviderEvent::MessageStopped); + Ok(events) + } + ResponsesStreamEvent::ResponseFailed { response } => { + Err(ProviderError::MalformedStream(format!( + "responses response failed{}", + response + .error_message() + .map(|message| format!(": {message}")) + .unwrap_or_default() + ))) + } + ResponsesStreamEvent::Error { message, error } => Err(ProviderError::MalformedStream( + error + .and_then(|error| error.message) + .or(message) + .unwrap_or_else(|| "responses stream error".to_string()), + )), + ResponsesStreamEvent::Unknown => Ok(Vec::new()), + } +} + +fn find_frame_boundary(buffer: &[u8]) -> Option<(usize, usize)> { + for (index, window) in buffer.windows(2).enumerate() { + if window == b"\n\n" { + return Some((index, 2)); + } + } + + for (index, window) in buffer.windows(4).enumerate() { + if window == b"\r\n\r\n" { + return Some((index, 4)); + } + } + + None +} + +#[derive(Deserialize)] +#[serde(tag = "type")] +enum ResponsesStreamEvent { + #[serde(rename = "response.created")] + ResponseCreated { response: ResponsesResponseEnvelope }, + #[serde(rename = "response.output_item.added")] + ResponseOutputItemAdded { + output_index: usize, + item: ResponsesOutputItem, + }, + #[serde(rename = "response.output_text.delta")] + ResponseOutputTextDelta { + output_index: usize, + delta: String, + #[allow(dead_code)] + content_index: Option, + }, + #[serde(rename = "response.reasoning_summary_text.delta")] + ResponseReasoningSummaryTextDelta { + #[serde(default)] + output_index: Option, + delta: String, + summary_index: i64, + }, + #[serde(rename = "response.reasoning_text.delta")] + ResponseReasoningTextDelta { + #[serde(default)] + output_index: Option, + delta: String, + content_index: i64, + }, + #[serde(rename = "response.reasoning_summary_part.added")] + ResponseReasoningSummaryPartAdded { + #[serde(default, rename = "output_index")] + _output_index: Option, + summary_index: i64, + }, + #[serde(rename = "response.reasoning_summary_part.done")] + ResponseReasoningSummaryPartDone { + #[serde(default)] + output_index: Option, + }, + #[serde(rename = "response.function_call_arguments.delta")] + ResponseFunctionCallArgumentsDelta { + output_index: usize, + delta: String, + #[allow(dead_code)] + item_id: Option, + }, + #[serde(rename = "response.tool_search_call.delta")] + ResponseToolSearchCallDelta { output_index: usize, delta: String }, + #[serde(rename = "response.output_item.done")] + ResponseOutputItemDone { + output_index: usize, + item: ResponsesOutputItem, + }, + #[serde(rename = "response.completed")] + ResponseCompleted { response: ResponsesResponseEnvelope }, + #[serde(rename = "response.incomplete")] + ResponseIncomplete { response: ResponsesResponseEnvelope }, + #[serde(rename = "response.failed")] + ResponseFailed { response: ResponsesResponseEnvelope }, + Error { + #[serde(default)] + message: Option, + #[serde(default)] + error: Option, + }, + #[serde(other)] + Unknown, +} + +#[derive(Deserialize)] +struct ResponsesResponseEnvelope { + #[serde(default)] + id: String, + #[serde(default)] + model: String, + #[serde(default)] + status: Option, + #[serde(default)] + usage: Option, + #[serde(default)] + incomplete_details: Option, + #[serde(default)] + error: Option, + #[serde(default)] + output: Vec, +} + +impl ResponsesResponseEnvelope { + fn stop_reason(&self) -> Option { + if let Some(reason) = self + .incomplete_details + .as_ref() + .and_then(|details| details.reason.as_ref()) + { + return Some(reason.clone()); + } + + match self.status.as_deref() { + Some("completed") | Some("in_progress") => None, + Some(status) => Some(status.to_string()), + None => None, + } + } + + fn error_message(&self) -> Option { + self.error.as_ref().and_then(|error| error.message.clone()) + } + + fn usage(&self) -> Option { + self.usage.as_ref().and_then(ResponsesUsage::to_token_usage) + } +} + +#[derive(Deserialize)] +struct ResponsesIncompleteDetails { + #[serde(default)] + reason: Option, +} + +#[derive(Deserialize)] +struct ResponsesErrorBody { + #[serde(default)] + message: Option, +} + +#[derive(Deserialize)] +struct ResponsesUsage { + #[serde(default)] + input_tokens: Option, + #[serde(default)] + output_tokens: Option, + #[serde(default)] + total_tokens: Option, + #[serde(default)] + input_tokens_details: Option, + #[serde(default)] + output_tokens_details: Option, +} + +impl ResponsesUsage { + fn to_token_usage(&self) -> Option { + let usage = TokenUsage { + input_tokens: self.input_tokens, + output_tokens: self.output_tokens, + total_tokens: self.total_tokens, + cache_read_input_tokens: self + .input_tokens_details + .as_ref() + .and_then(|details| details.cached_tokens), + cache_creation_input_tokens: None, + reasoning_tokens: self + .output_tokens_details + .as_ref() + .and_then(|details| details.reasoning_tokens), + thoughts_tokens: None, + tool_input_tokens: None, + }; + + (!usage.is_empty()).then_some(usage) + } +} + +#[derive(Deserialize)] +struct ResponsesInputTokenDetails { + #[serde(default)] + cached_tokens: Option, +} + +#[derive(Deserialize)] +struct ResponsesOutputTokenDetails { + #[serde(default)] + reasoning_tokens: Option, +} + +#[derive(Deserialize)] +#[serde(tag = "type")] +enum ResponsesOutputItem { + #[serde(rename = "message")] + Message { + #[serde(default)] + content: Vec, + }, + #[serde(rename = "reasoning")] + Reasoning { + id: String, + #[serde(default)] + encrypted_content: Option, + #[serde(default)] + summary: Vec, + #[serde(default)] + content: Vec, + }, + #[serde(rename = "function_call")] + FunctionCall { + #[serde(default)] + id: Option, + call_id: String, + name: String, + #[serde(default)] + arguments: String, + }, + #[serde(rename = "tool_search_call")] + ToolSearchCall { + #[serde(default)] + id: Option, + #[serde(default)] + call_id: Option, + #[serde(default)] + status: Option, + #[serde(default)] + _execution: Option, + #[serde(default)] + arguments: Option, + }, + #[serde(rename = "web_search_call")] + WebSearchCall { + #[serde(default)] + id: Option, + #[serde(default)] + status: Option, + #[serde(default)] + action: Option, + }, + #[serde(rename = "image_generation_call")] + ImageGenerationCall { + id: String, + status: String, + #[serde(default)] + revised_prompt: Option, + #[serde(default)] + result: Option, + }, + #[serde(other)] + Unsupported, +} + +impl ResponsesOutputItem { + fn into_provider_start( + self, + reasoning_provenance: &ReasoningProvenance, + pair_function_item_ids: bool, + ) -> Option { + match self { + ResponsesOutputItem::Message { .. } => Some(ContentBlockStart::Text), + ResponsesOutputItem::Reasoning { + id, + encrypted_content, + .. + } => Some(ContentBlockStart::Thinking { + encrypted_content, + id: Some(id), + provenance: Some(reasoning_provenance.clone()), + redacted: false, + }), + ResponsesOutputItem::FunctionCall { + id, call_id, name, .. + } => { + let item_id = if pair_function_item_ids { + id.as_deref() + } else { + None + }; + let id = encode_responses_tool_use_id(&call_id, item_id); + Some(ContentBlockStart::ToolUse { id, name }) + } + ResponsesOutputItem::ToolSearchCall { + id, + call_id, + status, + arguments, + .. + } => Some(ContentBlockStart::HostedToolSearch { + call: HostedToolSearchCall { + id: call_id + .or(id) + .unwrap_or_else(|| "tool_search_call".to_string()), + status, + query: arguments + .as_ref() + .and_then(extract_tool_search_query_from_value), + }, + }), + ResponsesOutputItem::WebSearchCall { id, status, action } => { + Some(ContentBlockStart::HostedWebSearch { + call: HostedWebSearchCall { + id: id.unwrap_or_else(|| "web_search_call".to_string()), + status, + action, + }, + }) + } + ResponsesOutputItem::ImageGenerationCall { + id, + status, + revised_prompt, + result, + } => Some(ContentBlockStart::ImageGeneration { + call: ImageGenerationCall { + id, + status, + revised_prompt, + result: result.map(|result| ImageGenerationResult::ArtifactRef { + artifact_id: result, + }), + }, + }), + ResponsesOutputItem::Unsupported => None, + } + } + + fn completed_text(&self) -> Option { + match self { + ResponsesOutputItem::Message { content } => { + let text = content + .iter() + .filter_map(ResponsesMessageContent::text) + .collect::>() + .join(""); + Some(text) + } + _ => None, + } + } + + fn completed_reasoning_text(&self) -> Option { + match self { + ResponsesOutputItem::Reasoning { + summary, content, .. + } => { + let summary = summary + .iter() + .map(|part| part.text.as_str()) + .collect::>() + .join("\n\n"); + if !summary.is_empty() { + return Some(summary); + } + Some( + content + .iter() + .map(|part| part.text.as_str()) + .collect::>() + .join("\n\n"), + ) + } + _ => None, + } + } + + fn reasoning_id(&self) -> Option<&str> { + match self { + ResponsesOutputItem::Reasoning { id, .. } => Some(id), + _ => None, + } + } + + fn reasoning_encrypted_content(&self) -> Option<(&str, &str)> { + match self { + ResponsesOutputItem::Reasoning { + id, + encrypted_content: Some(encrypted_content), + .. + } if !encrypted_content.is_empty() => Some((id, encrypted_content)), + _ => None, + } + } + + fn completed_arguments(&self) -> Option { + match self { + ResponsesOutputItem::FunctionCall { arguments, .. } => Some(arguments.clone()), + _ => None, + } + } + + fn completed_tool_search_query(&self) -> Option { + match self { + ResponsesOutputItem::ToolSearchCall { arguments, .. } => arguments + .as_ref() + .and_then(extract_tool_search_query_from_value), + _ => None, + } + } + + fn completion_deltas(&self, output_index: usize) -> Vec { + match self { + ResponsesOutputItem::ToolSearchCall { status, .. } => status + .clone() + .filter(|status| !status.is_empty()) + .map(|status| ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::HostedToolSearchStatus(status), + }) + .into_iter() + .collect(), + ResponsesOutputItem::Reasoning { + encrypted_content: Some(encrypted_content), + .. + } if !encrypted_content.is_empty() => { + vec![ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ThinkingEncryptedContent(encrypted_content.clone()), + }] + } + ResponsesOutputItem::WebSearchCall { status, action, .. } => { + let mut events = Vec::new(); + if let Some(action) = action.clone() { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::HostedWebSearchAction(action), + }); + } + if let Some(status) = status.clone().filter(|status| !status.is_empty()) { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::HostedWebSearchStatus(status), + }); + } + events + } + ResponsesOutputItem::ImageGenerationCall { + status, + revised_prompt, + result, + .. + } => { + let mut events = Vec::new(); + if let Some(revised_prompt) = revised_prompt + .clone() + .filter(|revised_prompt| !revised_prompt.is_empty()) + { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ImageGenerationRevisedPrompt(revised_prompt), + }); + } + if let Some(result) = result.clone() { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ImageGenerationResult( + ImageGenerationResult::ArtifactRef { + artifact_id: result, + }, + ), + }); + } + if !status.is_empty() { + events.push(ProviderEvent::ContentBlockDelta { + index: output_index, + delta: ContentBlockDelta::ImageGenerationStatus(status.clone()), + }); + } + events + } + _ => Vec::new(), + } + } + + fn is_supported(&self) -> bool { + !matches!(self, ResponsesOutputItem::Unsupported) + } +} + +#[derive(Deserialize)] +struct ResponsesReasoningText { + #[serde(default)] + text: String, +} + +fn extract_tool_search_query_from_value(value: &serde_json::Value) -> Option { + value + .get("query") + .and_then(serde_json::Value::as_str) + .map(str::to_string) +} + +#[derive(Deserialize)] +#[serde(tag = "type")] +enum ResponsesMessageContent { + #[serde(rename = "output_text")] + OutputText { text: String }, + #[serde(rename = "input_text")] + InputText { text: String }, + #[serde(other)] + Unsupported, +} + +impl ResponsesMessageContent { + fn text(&self) -> Option { + match self { + ResponsesMessageContent::OutputText { text } + | ResponsesMessageContent::InputText { text } => Some(text.clone()), + ResponsesMessageContent::Unsupported => None, + } + } +} + +#[cfg(test)] +mod tests { + use http::HeaderMap; + use http::HeaderValue; + + use crate::{ + ContentBlock, ContentBlockDelta, ContentBlockStart, ProviderError, ProviderEvent, + ProviderId, ReasoningFormat, ReasoningProvenance, Role, TokenUsage, + }; + + use super::{StreamState, parse_frame, response_headers_event}; + + #[test] + fn emits_response_headers_event_for_metadata_headers() { + let mut headers = HeaderMap::new(); + headers.insert("openai-model", HeaderValue::from_static("gpt-5")); + headers.insert("x-models-etag", HeaderValue::from_static("etag-123")); + headers.insert( + "x-ratelimit-limit-requests", + HeaderValue::from_static("1000"), + ); + + let event = response_headers_event(&headers).expect("headers event should be emitted"); + assert_eq!( + event, + ProviderEvent::ResponseHeaders(crate::ResponseHeaders { + values: vec![ + ("openai-model".to_string(), "gpt-5".to_string()), + ("x-models-etag".to_string(), "etag-123".to_string()), + ("x-ratelimit-limit-requests".to_string(), "1000".to_string()), + ], + }) + ); + } + + #[test] + fn parses_reasoning_delta_stream_events() { + let mut state = StreamState::default(); + + let summary_delta = parse_frame( + br#"data: {"type":"response.reasoning_summary_text.delta","summary_index":2,"delta":"short summary"}"#, + &mut state, + ) + .expect("reasoning summary delta should parse"); + assert_eq!( + summary_delta, + vec![ProviderEvent::ReasoningSummaryDelta { + delta: "short summary".to_string(), + summary_index: 2, + }] + ); + + let part_added = parse_frame( + br#"data: {"type":"response.reasoning_summary_part.added","summary_index":2}"#, + &mut state, + ) + .expect("reasoning summary part should parse"); + assert_eq!( + part_added, + vec![ProviderEvent::ReasoningSummaryPartAdded { summary_index: 2 }] + ); + + let reasoning_delta = parse_frame( + br#"data: {"type":"response.reasoning_text.delta","content_index":7,"delta":"internal chain"}"#, + &mut state, + ) + .expect("reasoning content delta should parse"); + assert_eq!( + reasoning_delta, + vec![ProviderEvent::ReasoningContentDelta { + delta: "internal chain".to_string(), + content_index: 7, + }] + ); + } + + #[test] + fn captures_reasoning_items_and_pairs_function_item_ids() { + let mut state = + StreamState::new(ProviderId::new("openai-edge"), "gpt-requested".to_string()); + + let added = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":0,"item":{"type":"reasoning","id":"rs_1","encrypted_content":null,"summary":[]}}"#, + &mut state, + ) + .expect("reasoning item start should parse"); + assert_eq!( + added, + vec![ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::Thinking { + encrypted_content: None, + id: Some("rs_1".to_string()), + provenance: Some(ReasoningProvenance { + provider: ProviderId::new("openai-edge"), + model: "gpt-requested".to_string(), + format: ReasoningFormat::OpenAiEncrypted, + }), + redacted: false, + }, + }] + ); + + let delta = parse_frame( + br#"data: {"type":"response.reasoning_summary_text.delta","output_index":0,"summary_index":0,"delta":"short summary"}"#, + &mut state, + ) + .expect("reasoning summary delta should parse"); + assert_eq!( + delta, + vec![ + ProviderEvent::ReasoningSummaryDelta { + delta: "short summary".to_string(), + summary_index: 0, + }, + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingText("short summary".to_string()), + }, + ] + ); + + let done = parse_frame( + br#"data: {"type":"response.output_item.done","output_index":0,"item":{"type":"reasoning","id":"rs_1","encrypted_content":"encrypted-1","summary":[{"type":"summary_text","text":"short summary"}]}}"#, + &mut state, + ) + .expect("reasoning item completion should parse"); + assert_eq!( + done, + vec![ + ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ThinkingEncryptedContent("encrypted-1".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 0 }, + ] + ); + + let tool = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":1,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"read_file","arguments":""}}"#, + &mut state, + ) + .expect("function call should parse"); + assert_eq!( + tool, + vec![ProviderEvent::ContentBlockStarted { + index: 1, + kind: ContentBlockStart::ToolUse { + id: "call_1|fc_1".to_string(), + name: "read_file".to_string(), + }, + }] + ); + } + + #[test] + fn function_item_id_stays_out_of_non_reasoning_tool_use_ids() { + let mut state = StreamState::default(); + + let tool = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"read_file","arguments":""}}"#, + &mut state, + ) + .expect("function call should parse"); + + assert_eq!( + tool, + vec![ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "call_1".to_string(), + name: "read_file".to_string(), + }, + }] + ); + } + + #[tokio::test] + async fn completed_response_backfills_late_azure_encrypted_content() { + let mut state = + StreamState::new(ProviderId::new("azure-openai"), "gpt-requested".to_string()); + let frames: &[&[u8]] = &[ + br#"data: {"type":"response.created","response":{"id":"resp_1","model":"gpt-resolved","status":"in_progress"}}"#, + br#"data: {"type":"response.output_item.added","output_index":0,"item":{"type":"reasoning","id":"rs_1","summary":[]}}"#, + br#"data: {"type":"response.output_item.done","output_index":0,"item":{"type":"reasoning","id":"rs_1","summary":[{"type":"summary_text","text":"summary"}]}}"#, + br#"data: {"type":"response.completed","response":{"id":"resp_1","model":"gpt-resolved","status":"completed","output":[{"type":"reasoning","id":"rs_1","encrypted_content":"from-completed","summary":[{"type":"summary_text","text":"summary"}]}]}}"#, + ]; + let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); + for frame in frames { + for event in parse_frame(frame, &mut state).expect("frame should parse") { + tx.send(Ok(event)).unwrap(); + } + } + drop(tx); + + let response = crate::collect_response_from_stream(rx) + .await + .expect("response should rebuild after metadata backfill"); + + assert_eq!( + response.content, + vec![ContentBlock::Thinking { + thinking: "summary".to_string(), + signature: None, + encrypted_content: Some("from-completed".to_string()), + id: Some("rs_1".to_string()), + provenance: Some(ReasoningProvenance { + provider: ProviderId::new("azure-openai"), + model: "gpt-requested".to_string(), + format: ReasoningFormat::OpenAiEncrypted, + }), + redacted: false, + }] + ); + } + + #[tokio::test] + async fn completed_response_does_not_replace_output_item_encrypted_content() { + let mut state = + StreamState::new(ProviderId::new("azure-openai"), "gpt-requested".to_string()); + let frames: &[&[u8]] = &[ + br#"data: {"type":"response.created","response":{"id":"resp_1","model":"gpt-resolved","status":"in_progress"}}"#, + br#"data: {"type":"response.output_item.added","output_index":0,"item":{"type":"reasoning","id":"rs_1","summary":[]}}"#, + br#"data: {"type":"response.output_item.done","output_index":0,"item":{"type":"reasoning","id":"rs_1","encrypted_content":"from-done","summary":[]}}"#, + br#"data: {"type":"response.completed","response":{"id":"resp_1","model":"gpt-resolved","status":"completed","output":[{"type":"reasoning","id":"rs_1","encrypted_content":"from-completed","summary":[]}]}}"#, + ]; + let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); + for frame in frames { + for event in parse_frame(frame, &mut state).expect("frame should parse") { + tx.send(Ok(event)).unwrap(); + } + } + drop(tx); + + let response = crate::collect_response_from_stream(rx) + .await + .expect("response should rebuild"); + let ContentBlock::Thinking { + encrypted_content, .. + } = &response.content[0] + else { + panic!("expected thinking block"); + }; + assert_eq!(encrypted_content.as_deref(), Some("from-done")); + } + + #[test] + fn parses_tool_call_stream_events() { + let mut state = StreamState::default(); + + let created = parse_frame( + br#"data: {"type":"response.created","response":{"id":"resp_1","model":"gpt-5","status":"in_progress"}}"#, + &mut state, + ) + .expect("created event should parse"); + assert_eq!( + created, + vec![ + ProviderEvent::ResponseCreated, + ProviderEvent::MessageStarted { + id: "resp_1".to_string(), + model: "gpt-5".to_string(), + role: Role::Assistant, + }, + ] + ); + + let added = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"call_1","name":"read_file","arguments":""}}"#, + &mut state, + ) + .expect("tool call start should parse"); + assert_eq!( + added, + vec![ProviderEvent::ContentBlockStarted { + index: 0, + kind: ContentBlockStart::ToolUse { + id: "call_1".to_string(), + name: "read_file".to_string(), + }, + }] + ); + + let delta = parse_frame( + br#"data: {"type":"response.function_call_arguments.delta","output_index":0,"delta":"{\"path\":\"README.md\"}"}"#, + &mut state, + ) + .expect("tool arguments delta should parse"); + assert_eq!( + delta, + vec![ProviderEvent::ContentBlockDelta { + index: 0, + delta: ContentBlockDelta::ToolUseInputJson("{\"path\":\"README.md\"}".to_string()), + }] + ); + + let done = parse_frame( + br#"data: {"type":"response.output_item.done","output_index":0,"item":{"type":"function_call","call_id":"call_1","name":"read_file","arguments":"{\"path\":\"README.md\"}"}}"#, + &mut state, + ) + .expect("tool call completion should parse"); + assert_eq!(done, vec![ProviderEvent::ContentBlockStopped { index: 0 }]); + } + + #[test] + fn falls_back_to_completed_message_text_when_no_text_delta_arrives() { + let mut state = StreamState::default(); + + let _ = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":1,"item":{"type":"message","content":[]}}"#, + &mut state, + ) + .expect("message start should parse"); + + let done = parse_frame( + br#"data: {"type":"response.output_item.done","output_index":1,"item":{"type":"message","content":[{"type":"output_text","text":"Hello"}]}}"#, + &mut state, + ) + .expect("message completion should parse"); + assert_eq!( + done, + vec![ + ProviderEvent::ContentBlockDelta { + index: 1, + delta: ContentBlockDelta::Text("Hello".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 1 }, + ] + ); + + let completed = parse_frame( + br#"data: {"type":"response.completed","response":{"id":"resp_1","model":"gpt-5","status":"completed"}}"#, + &mut state, + ) + .expect("completion should parse"); + assert_eq!( + completed, + vec![ + ProviderEvent::MessageDelta { + stop_reason: None, + usage: None, + }, + ProviderEvent::MessageStopped, + ] + ); + } + + #[test] + fn parses_final_usage_from_completed_response() { + let mut state = StreamState::default(); + + let completed = parse_frame( + br#"data: {"type":"response.completed","response":{"id":"resp_1","model":"gpt-5","status":"completed","usage":{"input_tokens":328,"input_tokens_details":{"cached_tokens":12},"output_tokens":52,"output_tokens_details":{"reasoning_tokens":7},"total_tokens":380}}}"#, + &mut state, + ) + .expect("completion should parse"); + + assert_eq!( + completed, + vec![ + ProviderEvent::MessageDelta { + stop_reason: None, + usage: Some(TokenUsage { + input_tokens: Some(328), + output_tokens: Some(52), + total_tokens: Some(380), + cache_read_input_tokens: Some(12), + cache_creation_input_tokens: None, + reasoning_tokens: Some(7), + thoughts_tokens: None, + tool_input_tokens: None, + }), + }, + ProviderEvent::MessageStopped, + ] + ); + } + + #[test] + fn surfaces_the_provider_error_from_a_failed_response_without_a_model() { + let mut state = StreamState::default(); + + let error = parse_frame( + br#"data: {"type":"response.failed","response":{"error":{"message":"upstream exploded"}}}"#, + &mut state, + ) + .expect_err("failed response should surface the provider error"); + + assert!( + matches!(&error, ProviderError::MalformedStream(message) if message.contains("upstream exploded")), + "unexpected error: {error}" + ); + } + + #[test] + fn terminates_the_message_for_an_incomplete_response_without_a_model() { + let mut state = StreamState::default(); + + let incomplete = parse_frame( + br#"data: {"type":"response.incomplete","response":{"status":"max_output_tokens"}}"#, + &mut state, + ) + .expect("incomplete response should parse"); + + assert_eq!( + incomplete, + vec![ + ProviderEvent::MessageDelta { + stop_reason: Some("max_output_tokens".to_string()), + usage: None, + }, + ProviderEvent::MessageStopped, + ] + ); + } + + #[test] + fn falls_back_to_the_requested_model_when_a_created_response_omits_it() { + let mut state = StreamState::new(ProviderId::new("openai"), "gpt-requested".to_string()); + + let created = parse_frame( + br#"data: {"type":"response.created","response":{"id":"resp_1"}}"#, + &mut state, + ) + .expect("created event should parse"); + + assert_eq!( + created, + vec![ + ProviderEvent::ResponseCreated, + ProviderEvent::MessageStarted { + id: "resp_1".to_string(), + model: "gpt-requested".to_string(), + role: Role::Assistant, + }, + ] + ); + assert_eq!(state.reasoning_provenance.model, "gpt-requested"); + } + + #[test] + fn parses_hosted_tool_search_output_items() { + let mut state = StreamState::default(); + + let added = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":3,"item":{"type":"tool_search_call","id":"search_1","status":"in_progress"}}"#, + &mut state, + ) + .expect("hosted search start should parse"); + assert_eq!( + added, + vec![ProviderEvent::ContentBlockStarted { + index: 3, + kind: ContentBlockStart::HostedToolSearch { + call: crate::HostedToolSearchCall { + id: "search_1".to_string(), + status: Some("in_progress".to_string()), + query: None, + }, + }, + }] + ); + + let delta = parse_frame( + br#"data: {"type":"response.tool_search_call.delta","output_index":3,"delta":"weather"}"#, + &mut state, + ) + .expect("hosted search delta should parse"); + assert_eq!( + delta, + vec![ProviderEvent::ContentBlockDelta { + index: 3, + delta: ContentBlockDelta::HostedToolSearchQuery("weather".to_string()), + }] + ); + + let done = parse_frame( + br#"data: {"type":"response.output_item.done","output_index":3,"item":{"type":"tool_search_call","id":"search_1","status":"completed","arguments":{"query":"weather"}}}"#, + &mut state, + ) + .expect("hosted search completion should parse"); + assert_eq!( + done, + vec![ + ProviderEvent::ContentBlockDelta { + index: 3, + delta: ContentBlockDelta::HostedToolSearchStatus("completed".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 3 }, + ] + ); + } + + #[test] + fn parses_web_search_and_image_generation_output_items() { + let mut state = StreamState::default(); + + let added = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":4,"item":{"type":"web_search_call","id":"ws_1","status":"in_progress"}}"#, + &mut state, + ) + .expect("web search start should parse"); + assert_eq!( + added, + vec![ProviderEvent::ContentBlockStarted { + index: 4, + kind: ContentBlockStart::HostedWebSearch { + call: crate::HostedWebSearchCall { + id: "ws_1".to_string(), + status: Some("in_progress".to_string()), + action: None, + }, + }, + }] + ); + + let done = parse_frame( + br#"data: {"type":"response.output_item.done","output_index":4,"item":{"type":"web_search_call","id":"ws_1","status":"completed","action":{"type":"search","query":"weather seattle"}}}"#, + &mut state, + ) + .expect("web search done should parse"); + assert_eq!( + done, + vec![ + ProviderEvent::ContentBlockDelta { + index: 4, + delta: ContentBlockDelta::HostedWebSearchAction( + crate::WebSearchAction::Search { + query: Some("weather seattle".to_string()), + queries: None, + } + ), + }, + ProviderEvent::ContentBlockDelta { + index: 4, + delta: ContentBlockDelta::HostedWebSearchStatus("completed".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 4 }, + ] + ); + + let image_added = parse_frame( + br#"data: {"type":"response.output_item.added","output_index":5,"item":{"type":"image_generation_call","id":"ig_1","status":"in_progress"}}"#, + &mut state, + ) + .expect("image generation start should parse"); + assert_eq!( + image_added, + vec![ProviderEvent::ContentBlockStarted { + index: 5, + kind: ContentBlockStart::ImageGeneration { + call: crate::ImageGenerationCall { + id: "ig_1".to_string(), + status: "in_progress".to_string(), + revised_prompt: None, + result: None, + }, + }, + }] + ); + + let image_done = parse_frame( + br#"data: {"type":"response.output_item.done","output_index":5,"item":{"type":"image_generation_call","id":"ig_1","status":"completed","revised_prompt":"A blue square","result":"artifact_1"}}"#, + &mut state, + ) + .expect("image generation done should parse"); + assert_eq!( + image_done, + vec![ + ProviderEvent::ContentBlockDelta { + index: 5, + delta: ContentBlockDelta::ImageGenerationRevisedPrompt( + "A blue square".to_string() + ), + }, + ProviderEvent::ContentBlockDelta { + index: 5, + delta: ContentBlockDelta::ImageGenerationResult( + crate::ImageGenerationResult::ArtifactRef { + artifact_id: "artifact_1".to_string(), + } + ), + }, + ProviderEvent::ContentBlockDelta { + index: 5, + delta: ContentBlockDelta::ImageGenerationStatus("completed".to_string()), + }, + ProviderEvent::ContentBlockStopped { index: 5 }, + ] + ); + } +} diff --git a/vendor/mentra-provider/src/responses/websocket.rs b/vendor/mentra-provider/src/responses/websocket.rs new file mode 100644 index 0000000..107d002 --- /dev/null +++ b/vendor/mentra-provider/src/responses/websocket.rs @@ -0,0 +1,705 @@ +use std::sync::Arc; +use std::time::Duration; + +use futures_util::SinkExt; +use futures_util::StreamExt; +use http::HeaderMap; +use http::HeaderValue; +use serde::Deserialize; +use serde_json::Value; +use serde_json::json; +use serde_json::map::Map as JsonMap; +use tokio::net::TcpStream; +use tokio::sync::Mutex; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::Instant; +use tokio_tungstenite::MaybeTlsStream; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite::Error as WsError; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use url::Url; + +use crate::ProviderError; +use crate::ProviderEvent; +use crate::ProviderEventStream; +use crate::ProviderId; +use crate::ReasoningFormat; +use crate::ReasoningProvenance; +use crate::ResponseHeaders; +use crate::error::retry_after_from_header_value; +use crate::error::retry_after_from_headers; + +use super::SharedTurnState; +use super::sse::StreamState; +use super::sse::parse_json_event; + +const X_CODEX_TURN_STATE_HEADER: &str = "x-codex-turn-state"; +const WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE: &str = "websocket_connection_limit_reached"; +const WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE: &str = "Responses websocket connection limit reached (60 minutes). Create a new websocket connection to continue."; + +pub trait ResponsesWebsocketTelemetry: Send + Sync { + fn on_ws_request( + &self, + duration: Duration, + error: Option<&ProviderError>, + connection_reused: bool, + ); + + fn on_ws_event( + &self, + result: &Result>, ProviderError>, + duration: Duration, + ); +} + +struct WsStream { + tx_command: mpsc::Sender, + rx_message: mpsc::UnboundedReceiver>, + pump_task: tokio::task::JoinHandle<()>, +} + +enum WsCommand { + Send { + message: Message, + tx_result: oneshot::Sender>, + }, +} + +impl WsStream { + fn new(inner: WebSocketStream>) -> Self { + let (tx_command, mut rx_command) = mpsc::channel::(32); + let (tx_message, rx_message) = mpsc::unbounded_channel::>(); + + let pump_task = tokio::spawn(async move { + let mut inner = inner; + loop { + tokio::select! { + command = rx_command.recv() => { + let Some(command) = command else { + break; + }; + match command { + WsCommand::Send { message, tx_result } => { + let result = inner.send(message).await; + let should_break = result.is_err(); + let _ = tx_result.send(result); + if should_break { + break; + } + } + } + } + message = inner.next() => { + let Some(message) = message else { + break; + }; + match message { + Ok(Message::Ping(payload)) => { + if let Err(err) = inner.send(Message::Pong(payload)).await { + let _ = tx_message.send(Err(err)); + break; + } + } + Ok(Message::Pong(_)) => {} + Ok(message @ (Message::Text(_) + | Message::Binary(_) + | Message::Close(_) + | Message::Frame(_))) => { + let is_close = matches!(message, Message::Close(_)); + if tx_message.send(Ok(message)).is_err() { + break; + } + if is_close { + break; + } + } + Err(err) => { + let _ = tx_message.send(Err(err)); + break; + } + } + } + } + } + }); + + Self { + tx_command, + rx_message, + pump_task, + } + } + + async fn send(&self, message: Message) -> Result<(), WsError> { + let (tx_result, rx_result) = oneshot::channel(); + if self + .tx_command + .send(WsCommand::Send { message, tx_result }) + .await + .is_err() + { + return Err(WsError::ConnectionClosed); + } + rx_result.await.unwrap_or(Err(WsError::ConnectionClosed)) + } + + async fn next(&mut self) -> Option> { + self.rx_message.recv().await + } +} + +impl Drop for WsStream { + fn drop(&mut self) { + self.pump_task.abort(); + } +} + +#[derive(Clone)] +pub struct ResponsesWebsocketConnection { + stream: Arc>>, + idle_timeout: Duration, + response_headers: ResponseHeaders, + telemetry: Option>, +} + +impl std::fmt::Debug for ResponsesWebsocketConnection { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ResponsesWebsocketConnection") + .field("stream", &"") + .field("idle_timeout", &self.idle_timeout) + .field("response_headers", &self.response_headers) + .field("telemetry", &self.telemetry.as_ref().map(|_| "")) + .finish() + } +} + +impl ResponsesWebsocketConnection { + pub async fn connect( + url: Url, + headers: HeaderMap, + turn_state: Option, + idle_timeout: Duration, + telemetry: Option>, + ) -> Result { + let mut request = url + .as_str() + .into_client_request() + .map_err(|error| ProviderError::InvalidRequest(error.to_string()))?; + request.headers_mut().extend(headers); + + let (stream, response) = connect_async(request) + .await + .map_err(|error| map_ws_error(error, &url))?; + + let header_value = response + .headers() + .get(X_CODEX_TURN_STATE_HEADER) + .and_then(|value| value.to_str().ok()); + if let (Some(turn_state), Some(header_value)) = (turn_state, header_value) { + *turn_state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + Some(header_value.to_string()); + } + + let response_headers = ResponseHeaders { + values: response + .headers() + .iter() + .filter_map(|(name, value)| { + value + .to_str() + .ok() + .map(|value| (name.as_str().to_string(), value.to_string())) + }) + .collect(), + }; + + Ok(Self { + stream: Arc::new(Mutex::new(Some(WsStream::new(stream)))), + idle_timeout, + response_headers, + telemetry, + }) + } + + pub async fn is_closed(&self) -> bool { + self.stream.lock().await.is_none() + } + + pub async fn stream_request( + &self, + request_body: Value, + connection_reused: bool, + ) -> Result { + let requested_model = request_body + .get("model") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + self.stream_request_with_provenance( + request_body, + connection_reused, + ReasoningProvenance { + provider: ProviderId::new("openai"), + model: requested_model, + format: ReasoningFormat::OpenAiEncrypted, + }, + ) + .await + } + + pub(crate) async fn stream_request_with_provenance( + &self, + request_body: Value, + connection_reused: bool, + provenance: ReasoningProvenance, + ) -> Result { + let (tx_event, rx_event) = + mpsc::unbounded_channel::>(); + let stream = Arc::clone(&self.stream); + let idle_timeout = self.idle_timeout; + let response_headers = self.response_headers.clone(); + let telemetry = self.telemetry.clone(); + + let request_text = + serde_json::to_string(&request_body).map_err(ProviderError::Serialize)?; + + tokio::spawn(async move { + if tx_event + .send(Ok(ProviderEvent::ResponseHeaders(response_headers))) + .is_err() + { + return; + } + + let mut guard = stream.lock().await; + let result = { + let Some(ws_stream) = guard.as_mut() else { + let _ = tx_event.send(Err(ProviderError::MalformedStream( + "websocket connection is closed".to_string(), + ))); + return; + }; + run_websocket_response_stream( + ws_stream, + tx_event.clone(), + request_text, + idle_timeout, + telemetry, + connection_reused, + StreamState::new(provenance.provider, provenance.model), + ) + .await + }; + + if let Err(err) = result { + let failed_stream = guard.take(); + drop(guard); + drop(failed_stream); + let _ = tx_event.send(Err(err)); + } + }); + + Ok(rx_event) + } +} + +pub fn merge_request_headers( + provider_headers: &HeaderMap, + extra_headers: HeaderMap, + default_headers: HeaderMap, +) -> HeaderMap { + let mut headers = provider_headers.clone(); + headers.extend(extra_headers); + for (name, value) in &default_headers { + if let http::header::Entry::Vacant(entry) = headers.entry(name) { + entry.insert(value.clone()); + } + } + headers +} + +pub(crate) fn response_create_frame(mut response: Value) -> Value { + let Value::Object(response_object) = &mut response else { + return json!({ + "type": "response.create", + "response": response, + }); + }; + + response_object.insert( + "type".to_string(), + Value::String("response.create".to_string()), + ); + response_object + .entry("instructions") + .or_insert_with(|| Value::String(String::new())); + response +} + +async fn run_websocket_response_stream( + ws_stream: &mut WsStream, + tx_event: mpsc::UnboundedSender>, + request_text: String, + idle_timeout: Duration, + telemetry: Option>, + connection_reused: bool, + mut state: StreamState, +) -> Result<(), ProviderError> { + let request_start = Instant::now(); + let send_result = ws_stream.send(Message::Text(request_text.into())).await; + let send_error = send_result + .as_ref() + .err() + .map(|error| ProviderError::MalformedStream(error.to_string())); + if let Some(t) = telemetry.as_ref() { + t.on_ws_request( + request_start.elapsed(), + send_error.as_ref(), + connection_reused, + ); + } + send_result.map_err(|error| ProviderError::MalformedStream(error.to_string()))?; + + loop { + let poll_start = Instant::now(); + let message_result = tokio::time::timeout(idle_timeout, ws_stream.next()) + .await + .map_err(|_| { + ProviderError::MalformedStream("idle timeout waiting for websocket".into()) + }); + if let Some(t) = telemetry.as_ref() { + t.on_ws_event(&message_result, poll_start.elapsed()); + } + let message = match message_result { + Ok(Some(Ok(message))) => message, + Ok(Some(Err(error))) => { + return Err(ProviderError::MalformedStream(error.to_string())); + } + Ok(None) => { + return Err(ProviderError::MalformedStream( + "stream closed before response.completed".to_string(), + )); + } + Err(error) => return Err(error), + }; + + match message { + Message::Text(text) => { + if let Some(mapped) = parse_wrapped_websocket_error_event(&text) { + let receiver_closed = mapped + .headers + .map(|headers| { + tx_event + .send(Ok(ProviderEvent::ResponseHeaders(headers))) + .is_err() + }) + .unwrap_or(false); + if receiver_closed { + return Ok(()); + } + return Err(mapped.error); + } + + let mut saw_message_stopped = false; + for event in parse_json_event(&text, &mut state)? { + if matches!(event, ProviderEvent::MessageStopped) { + saw_message_stopped = true; + } + if tx_event.send(Ok(event)).is_err() { + return Ok(()); + } + } + if saw_message_stopped { + break; + } + } + Message::Binary(_) => { + return Err(ProviderError::MalformedStream( + "unexpected binary websocket event".to_string(), + )); + } + Message::Close(_) => { + return Err(ProviderError::MalformedStream( + "websocket closed by server before response.completed".to_string(), + )); + } + Message::Frame(_) => {} + Message::Ping(_) | Message::Pong(_) => {} + } + } + + Ok(()) +} + +/// Delay hint for connection-level websocket failures. Short enough that a +/// tunnel blip or service restart is retried promptly, long enough to avoid +/// hammering a target that is still coming back up. +const WS_CONNECTION_RETRY_DELAY: Duration = Duration::from_millis(750); + +fn map_ws_error(error: WsError, url: &Url) -> ProviderError { + match error { + WsError::Http(response) => { + let status = response.status(); + let retry_after = retry_after_from_headers(response.headers()); + let body = response + .body() + .as_ref() + .and_then(|bytes| String::from_utf8(bytes.clone()).ok()) + .unwrap_or_else(|| format!("websocket connection failed for {url}")); + ProviderError::Http { + status, + body, + retry_after, + } + } + // Connection lost mid-handshake: the peer went away before the upgrade + // completed. That is a transient liveness failure (restart, tunnel + // teardown), not a request the caller must fix — surface it as retryable. + WsError::ConnectionClosed | WsError::AlreadyClosed => ProviderError::Retryable { + message: format!("websocket connection closed while connecting to {url}"), + delay: Some(WS_CONNECTION_RETRY_DELAY), + }, + // Transport-level failure establishing the connection (e.g. connection + // refused when an SSH tunnel is down, DNS blips). These clear on their + // own once the path is back, so the whole turn is worth retrying rather + // than falling straight through to a terminal failure. + WsError::Io(error) => ProviderError::Retryable { + message: format!("websocket connection failed for {url}: {error}"), + delay: Some(WS_CONNECTION_RETRY_DELAY), + }, + other => ProviderError::InvalidResponse(other.to_string()), + } +} + +#[derive(Debug)] +struct MappedWebsocketError { + error: ProviderError, + headers: Option, +} + +#[derive(Debug, Deserialize)] +struct WrappedWebsocketError { + code: Option, + message: Option, +} + +#[derive(Debug, Deserialize)] +struct WrappedWebsocketErrorEvent { + #[serde(rename = "type")] + kind: String, + #[serde(alias = "status_code")] + status: Option, + #[serde(default)] + error: Option, + #[serde(default)] + headers: Option>, +} + +fn parse_wrapped_websocket_error_event(payload: &str) -> Option { + let event: WrappedWebsocketErrorEvent = serde_json::from_str(payload).ok()?; + if event.kind != "error" { + return None; + } + + if let Some(error) = event + .error + .as_ref() + .filter(|error| error.code.as_deref() == Some(WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE)) + { + return Some(MappedWebsocketError { + error: ProviderError::Retryable { + message: error + .message + .clone() + .unwrap_or_else(|| WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE.to_string()), + delay: None, + }, + headers: event.headers.map(response_headers_from_json), + }); + } + + let status = reqwest::StatusCode::from_u16(event.status?).ok()?; + let body = payload.to_string(); + let headers = event.headers.map(response_headers_from_json); + // The upgrade's own headers are long gone by the time a rate limit arrives + // as a wrapped frame, so the only `Retry-After` on this path is the one the + // frame echoed. Read it here or it is lost. + let retry_after = headers.as_ref().and_then(retry_after_from_response_headers); + Some(MappedWebsocketError { + error: ProviderError::Http { + status, + body, + retry_after, + }, + headers, + }) +} + +/// Reads `Retry-After` out of headers a frame carried as JSON. +fn retry_after_from_response_headers(headers: &ResponseHeaders) -> Option { + headers + .values + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("retry-after")) + .and_then(|(_, value)| retry_after_from_header_value(value)) +} + +fn response_headers_from_json(headers: JsonMap) -> ResponseHeaders { + ResponseHeaders { + values: headers + .into_iter() + .filter_map(|(name, value)| { + json_header_value(value) + .and_then(|value| value.to_str().ok().map(|value| (name, value.to_string()))) + }) + .collect(), + } +} + +fn json_header_value(value: Value) -> Option { + let value = match value { + Value::String(value) => value, + Value::Number(value) => value.to_string(), + Value::Bool(value) => value.to_string(), + _ => return None, + }; + HeaderValue::from_str(&value).ok() +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn merge_request_headers_matches_http_precedence() { + let mut provider_headers = HeaderMap::new(); + provider_headers.insert( + "originator", + HeaderValue::from_static("provider-originator"), + ); + provider_headers.insert("x-priority", HeaderValue::from_static("provider")); + + let mut extra_headers = HeaderMap::new(); + extra_headers.insert("x-priority", HeaderValue::from_static("extra")); + + let mut default_headers = HeaderMap::new(); + default_headers.insert("originator", HeaderValue::from_static("default-originator")); + default_headers.insert("x-priority", HeaderValue::from_static("default")); + default_headers.insert("x-default-only", HeaderValue::from_static("default-only")); + + let merged = merge_request_headers(&provider_headers, extra_headers, default_headers); + + assert_eq!( + merged.get("originator"), + Some(&HeaderValue::from_static("provider-originator")) + ); + assert_eq!( + merged.get("x-priority"), + Some(&HeaderValue::from_static("extra")) + ); + assert_eq!( + merged.get("x-default-only"), + Some(&HeaderValue::from_static("default-only")) + ); + } + + #[test] + fn response_create_frame_adds_type_to_responses_request_payload() { + let frame = response_create_frame(json!({ + "model": "gpt-5", + "input": [] + })); + + assert_eq!(frame["type"], "response.create"); + assert_eq!(frame["model"], "gpt-5"); + assert_eq!(frame["input"], json!([])); + assert_eq!(frame["instructions"], ""); + assert!(frame.get("response").is_none()); + } + + #[test] + fn wrapped_websocket_error_preserves_rate_limit_headers() { + let payload = json!({ + "type": "error", + "status": 429, + "error": { + "type": "usage_limit_reached", + "message": "The usage limit has been reached" + }, + "headers": { + "x-codex-primary-used-percent": "100.0" + } + }) + .to_string(); + + let mapped = parse_wrapped_websocket_error_event(&payload) + .expect("error payload should map to a provider error"); + let ProviderError::Http { status, .. } = mapped.error else { + panic!("expected ProviderError::Http"); + }; + assert_eq!(status, reqwest::StatusCode::TOO_MANY_REQUESTS); + assert_eq!( + mapped + .headers + .expect("rate limit headers should be present") + .values, + vec![( + "x-codex-primary-used-percent".to_string(), + "100.0".to_string() + )] + ); + } + + #[test] + fn ws_io_error_maps_to_retryable_with_delay() { + let url = Url::parse("wss://example.test/responses").expect("url parses"); + let io = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "Connection refused"); + let mapped = map_ws_error(WsError::Io(io), &url); + let ProviderError::Retryable { message, delay } = mapped else { + panic!("expected ProviderError::Retryable for a connection-level io error"); + }; + assert_eq!(delay, Some(WS_CONNECTION_RETRY_DELAY)); + // The io error string is preserved so operators can see the real cause. + assert!(message.contains("Connection refused"), "message: {message}"); + } + + #[test] + fn ws_connection_closed_maps_to_retryable() { + let url = Url::parse("wss://example.test/responses").expect("url parses"); + for error in [WsError::ConnectionClosed, WsError::AlreadyClosed] { + let mapped = map_ws_error(error, &url); + assert!( + matches!(mapped, ProviderError::Retryable { delay: Some(_), .. }), + "connection-closed handshake failure should be retryable", + ); + } + } + + #[test] + fn websocket_connection_limit_maps_to_retryable() { + let payload = json!({ + "type": "error", + "status": 400, + "error": { + "type": "invalid_request_error", + "code": "websocket_connection_limit_reached", + "message": WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE + } + }) + .to_string(); + + let mapped = parse_wrapped_websocket_error_event(&payload) + .expect("connection limit payload should map"); + let ProviderError::Retryable { message, delay } = mapped.error else { + panic!("expected ProviderError::Retryable"); + }; + assert_eq!(message, WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE); + assert_eq!(delay, None); + } +} diff --git a/vendor/mentra-provider/src/stream.rs b/vendor/mentra-provider/src/stream.rs new file mode 100644 index 0000000..b3f0d9e --- /dev/null +++ b/vendor/mentra-provider/src/stream.rs @@ -0,0 +1,101 @@ +use tokio::sync::mpsc; + +use crate::{ + ReasoningProvenance, model::HostedToolSearchCall, model::HostedWebSearchCall, + model::ImageGenerationCall, model::ImageGenerationResult, model::ImageSource, model::Role, + model::TokenUsage, model::ToolResultContent, model::WebSearchAction, +}; + +pub type ProviderEventStream = mpsc::UnboundedReceiver>; + +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub struct ResponseHeaders { + pub values: Vec<(String, String)>, +} + +#[derive(Debug, Clone, PartialEq)] +pub enum ProviderEvent { + ResponseHeaders(ResponseHeaders), + ResponseCreated, + MessageStarted { + id: String, + model: String, + role: Role, + }, + ContentBlockStarted { + index: usize, + kind: ContentBlockStart, + }, + ContentBlockDelta { + index: usize, + delta: ContentBlockDelta, + }, + ContentBlockStopped { + index: usize, + }, + MessageDelta { + stop_reason: Option, + usage: Option, + }, + ReasoningSummaryDelta { + delta: String, + summary_index: i64, + }, + ReasoningContentDelta { + delta: String, + content_index: i64, + }, + ReasoningSummaryPartAdded { + summary_index: i64, + }, + MessageStopped, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ContentBlockStart { + Text, + Thinking { + encrypted_content: Option, + id: Option, + provenance: Option, + redacted: bool, + }, + Image { + source: ImageSource, + }, + ToolUse { + id: String, + name: String, + }, + ToolResult { + tool_use_id: String, + is_error: bool, + content: Option, + }, + HostedToolSearch { + call: HostedToolSearchCall, + }, + HostedWebSearch { + call: HostedWebSearchCall, + }, + ImageGeneration { + call: ImageGenerationCall, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ContentBlockDelta { + Text(String), + ThinkingText(String), + ThinkingSignature(String), + ThinkingEncryptedContent(String), + ToolUseInputJson(String), + ToolResultContent(ToolResultContent), + HostedToolSearchQuery(String), + HostedToolSearchStatus(String), + HostedWebSearchAction(WebSearchAction), + HostedWebSearchStatus(String), + ImageGenerationStatus(String), + ImageGenerationRevisedPrompt(String), + ImageGenerationResult(ImageGenerationResult), +} diff --git a/vendor/mentra-provider/src/tool.rs b/vendor/mentra-provider/src/tool.rs new file mode 100644 index 0000000..f231ceb --- /dev/null +++ b/vendor/mentra-provider/src/tool.rs @@ -0,0 +1,151 @@ +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; + +/// Declares whether a tool is loaded eagerly or deferred for provider-native tool search. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +pub enum ToolLoadingPolicy { + #[default] + Immediate, + Deferred, +} + +/// Provider-visible tool kinds supported by Mentra-backed providers. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum ProviderToolKind { + #[default] + Function, + HostedWebSearch, + ImageGeneration, +} + +/// Provider-facing description of a tool and its schemas. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ToolSpec { + pub name: String, + pub description: Option, + pub input_schema: Value, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub output_schema: Option, + #[serde(default)] + pub kind: ProviderToolKind, + #[serde(default)] + pub loading_policy: ToolLoadingPolicy, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub strict: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub options: Option, +} + +impl ToolSpec { + pub fn builder(name: impl Into) -> ToolSpecBuilder { + ToolSpecBuilder { + name: name.into(), + description: None, + input_schema: json!({ + "type": "object", + "properties": {} + }), + output_schema: None, + kind: ProviderToolKind::Function, + loading_policy: ToolLoadingPolicy::Immediate, + strict: None, + options: None, + } + } +} + +#[derive(Debug, Clone)] +pub struct ToolSpecBuilder { + name: String, + description: Option, + input_schema: Value, + output_schema: Option, + kind: ProviderToolKind, + loading_policy: ToolLoadingPolicy, + strict: Option, + options: Option, +} + +impl ToolSpecBuilder { + pub fn description(mut self, description: impl Into) -> Self { + self.description = Some(description.into()); + self + } + + pub fn input_schema(mut self, input_schema: Value) -> Self { + self.input_schema = input_schema; + self + } + + pub fn output_schema(mut self, output_schema: Value) -> Self { + self.output_schema = Some(output_schema); + self + } + + pub fn kind(mut self, kind: ProviderToolKind) -> Self { + self.kind = kind; + self + } + + pub fn options(mut self, options: Value) -> Self { + self.options = Some(options); + self + } + + pub fn loading_policy(mut self, loading_policy: ToolLoadingPolicy) -> Self { + self.loading_policy = loading_policy; + self + } + + pub fn strict(mut self, strict: bool) -> Self { + self.strict = Some(strict); + self + } + + pub fn non_strict(self) -> Self { + self.strict(false) + } + + pub fn defer_loading(self, defer_loading: bool) -> Self { + self.loading_policy(if defer_loading { + ToolLoadingPolicy::Deferred + } else { + ToolLoadingPolicy::Immediate + }) + } + + pub fn build(self) -> ToolSpec { + ToolSpec { + name: self.name, + description: self.description, + input_schema: self.input_schema, + output_schema: self.output_schema, + kind: self.kind, + loading_policy: self.loading_policy, + strict: self.strict, + options: self.options, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn builder_leaves_function_tool_strictness_unset_by_default() { + let spec = ToolSpec::builder("echo").build(); + + assert_eq!(spec.strict, None); + } + + #[test] + fn builder_sets_explicit_strictness() { + let strict = ToolSpec::builder("strict").strict(true).build(); + let non_strict = ToolSpec::builder("loose").non_strict().build(); + + assert_eq!(strict.strict, Some(true)); + assert_eq!(non_strict.strict, Some(false)); + } +}