mirror of
https://github.com/block/goose.git
synced 2026-07-17 12:56:20 +02:00
OpenRouter & Xai streaming (#5873)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Generated
+10
-1
@@ -4123,6 +4123,15 @@ dependencies = [
|
||||
"either",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "itertools"
|
||||
version = "0.13.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "413ee7dfc52ee1a4949ceeb7dbc8a33f2d6c088194d9f922fb8318faf1f01186"
|
||||
dependencies = [
|
||||
"either",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "itertools"
|
||||
version = "0.14.0"
|
||||
@@ -5695,7 +5704,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"itertools 0.14.0",
|
||||
"itertools 0.13.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.111",
|
||||
|
||||
@@ -149,7 +149,14 @@ impl Provider for AzureProvider {
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||
let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
|
||||
let payload = create_request(
|
||||
model_config,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
let response = self
|
||||
.with_retry(|| async {
|
||||
let payload_clone = payload.clone();
|
||||
|
||||
@@ -1,13 +1,8 @@
|
||||
use anyhow::Result;
|
||||
use async_stream::try_stream;
|
||||
use async_trait::async_trait;
|
||||
use futures::TryStreamExt;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::io;
|
||||
use std::time::Duration;
|
||||
use tokio::pin;
|
||||
use tokio_util::io::StreamReader;
|
||||
|
||||
use super::api_client::{ApiClient, AuthMethod, AuthProvider};
|
||||
use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
@@ -17,21 +12,19 @@ use super::formats::databricks::{create_request, response_to_message};
|
||||
use super::oauth;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
get_model, handle_response_openai_compat, map_http_error_to_provider_error, ImageFormat,
|
||||
RequestLog,
|
||||
get_model, handle_response_openai_compat, map_http_error_to_provider_error,
|
||||
stream_openai_compat, ImageFormat, RequestLog,
|
||||
};
|
||||
use crate::config::ConfigError;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{get_usage, response_to_streaming_message};
|
||||
use crate::providers::formats::openai::get_usage;
|
||||
use crate::providers::retry::{
|
||||
RetryConfig, DEFAULT_BACKOFF_MULTIPLIER, DEFAULT_INITIAL_RETRY_INTERVAL_MS,
|
||||
DEFAULT_MAX_RETRIES, DEFAULT_MAX_RETRY_INTERVAL_MS,
|
||||
};
|
||||
use rmcp::model::Tool;
|
||||
use serde_json::json;
|
||||
use tokio_stream::StreamExt;
|
||||
use tokio_util::codec::{FramedRead, LinesCodec};
|
||||
|
||||
const DEFAULT_CLIENT_ID: &str = "databricks-cli";
|
||||
const DEFAULT_REDIRECT_URL: &str = "http://localhost";
|
||||
@@ -347,20 +340,7 @@ impl Provider for DatabricksProvider {
|
||||
let _ = log.error(e);
|
||||
})?;
|
||||
|
||||
let stream = response.bytes_stream().map_err(io::Error::other);
|
||||
|
||||
Ok(Box::pin(try_stream! {
|
||||
let stream_reader = StreamReader::new(stream);
|
||||
let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from);
|
||||
|
||||
let message_stream = response_to_streaming_message(framed);
|
||||
pin!(message_stream);
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
stream_openai_compat(response, log)
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
|
||||
@@ -613,6 +613,7 @@ pub fn create_request(
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
image_format: &ImageFormat,
|
||||
for_streaming: bool,
|
||||
) -> anyhow::Result<Value, Error> {
|
||||
if model_config.model_name.starts_with("o1-mini") {
|
||||
return Err(anyhow!(
|
||||
@@ -652,13 +653,8 @@ pub fn create_request(
|
||||
});
|
||||
|
||||
let messages_spec = format_messages(messages, image_format);
|
||||
let mut tools_spec = if !tools.is_empty() {
|
||||
format_tools(tools)?
|
||||
} else {
|
||||
vec![]
|
||||
};
|
||||
let mut tools_spec = format_tools(tools)?;
|
||||
|
||||
// Validate tool schemas
|
||||
validate_tool_schemas(&mut tools_spec);
|
||||
|
||||
let mut messages_array = vec![system_message];
|
||||
@@ -670,25 +666,17 @@ pub fn create_request(
|
||||
});
|
||||
|
||||
if let Some(effort) = reasoning_effort {
|
||||
payload
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("reasoning_effort".to_string(), json!(effort));
|
||||
payload["reasoning_effort"] = json!(effort);
|
||||
}
|
||||
|
||||
if !tools_spec.is_empty() {
|
||||
payload
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("tools".to_string(), json!(tools_spec));
|
||||
payload["tools"] = json!(tools_spec);
|
||||
}
|
||||
|
||||
// o1, o3 models currently don't support temperature
|
||||
if !is_ox_model {
|
||||
if let Some(temp) = model_config.temperature {
|
||||
payload
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("temperature".to_string(), json!(temp));
|
||||
payload["temperature"] = json!(temp);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -704,6 +692,12 @@ pub fn create_request(
|
||||
.unwrap()
|
||||
.insert(key.to_string(), json!(tokens));
|
||||
}
|
||||
|
||||
if for_streaming {
|
||||
payload["stream"] = json!(true);
|
||||
payload["stream_options"] = json!({"include_usage": true});
|
||||
}
|
||||
|
||||
Ok(payload)
|
||||
}
|
||||
|
||||
@@ -1277,7 +1271,14 @@ mod tests {
|
||||
toolshim_model: None,
|
||||
fast_model: None,
|
||||
};
|
||||
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
||||
let request = create_request(
|
||||
&model_config,
|
||||
"system",
|
||||
&[],
|
||||
&[],
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
let obj = request.as_object().unwrap();
|
||||
let expected = json!({
|
||||
"model": "gpt-4o",
|
||||
@@ -1309,7 +1310,14 @@ mod tests {
|
||||
toolshim_model: None,
|
||||
fast_model: None,
|
||||
};
|
||||
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
||||
let request = create_request(
|
||||
&model_config,
|
||||
"system",
|
||||
&[],
|
||||
&[],
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
let obj = request.as_object().unwrap();
|
||||
let expected = json!({
|
||||
"model": "o1",
|
||||
@@ -1342,7 +1350,14 @@ mod tests {
|
||||
toolshim_model: None,
|
||||
fast_model: None,
|
||||
};
|
||||
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
||||
let request = create_request(
|
||||
&model_config,
|
||||
"system",
|
||||
&[],
|
||||
&[],
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
let obj = request.as_object().unwrap();
|
||||
let expected = json!({
|
||||
"model": "o3-mini",
|
||||
|
||||
@@ -460,7 +460,14 @@ impl Provider for GithubCopilotProvider {
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||
let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
|
||||
let payload = create_request(
|
||||
model_config,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
let mut log = RequestLog::start(model_config, &payload)?;
|
||||
|
||||
// Make request with retry
|
||||
|
||||
@@ -179,6 +179,7 @@ impl Provider for LiteLLMProvider {
|
||||
messages,
|
||||
tools,
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
|
||||
if self.supports_cache_control().await {
|
||||
|
||||
@@ -3,7 +3,8 @@ use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, Provider
|
||||
use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
get_model, handle_response_openai_compat, handle_status_openai_compat, RequestLog,
|
||||
get_model, handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat,
|
||||
RequestLog,
|
||||
};
|
||||
use crate::config::declarative_providers::DeclarativeProviderConfig;
|
||||
use crate::config::GooseMode;
|
||||
@@ -11,23 +12,14 @@ use crate::conversation::message::Message;
|
||||
use crate::conversation::Conversation;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{
|
||||
create_request, get_usage, response_to_message, response_to_streaming_message,
|
||||
};
|
||||
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use crate::utils::safe_truncate;
|
||||
use anyhow::Result;
|
||||
use async_stream::try_stream;
|
||||
use async_trait::async_trait;
|
||||
use futures::TryStreamExt;
|
||||
use regex::Regex;
|
||||
use rmcp::model::Tool;
|
||||
use serde_json::{json, Value};
|
||||
use std::io;
|
||||
use serde_json::Value;
|
||||
use std::time::Duration;
|
||||
use tokio::pin;
|
||||
use tokio_stream::StreamExt;
|
||||
use tokio_util::codec::{FramedRead, LinesCodec};
|
||||
use tokio_util::io::StreamReader;
|
||||
use url::Url;
|
||||
|
||||
pub const OLLAMA_HOST: &str = "localhost";
|
||||
@@ -200,6 +192,7 @@ impl Provider for OllamaProvider {
|
||||
messages,
|
||||
filtered_tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
|
||||
let mut log = RequestLog::start(model_config, &payload)?;
|
||||
@@ -262,17 +255,14 @@ impl Provider for OllamaProvider {
|
||||
tools
|
||||
};
|
||||
|
||||
let mut payload = create_request(
|
||||
let payload = create_request(
|
||||
&self.model,
|
||||
system,
|
||||
messages,
|
||||
filtered_tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
true,
|
||||
)?;
|
||||
payload["stream"] = json!(true);
|
||||
payload["stream_options"] = json!({
|
||||
"include_usage": true,
|
||||
});
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
|
||||
let response = self
|
||||
@@ -287,19 +277,7 @@ impl Provider for OllamaProvider {
|
||||
.inspect_err(|e| {
|
||||
let _ = log.error(e);
|
||||
})?;
|
||||
let stream = response.bytes_stream().map_err(io::Error::other);
|
||||
|
||||
Ok(Box::pin(try_stream! {
|
||||
let stream_reader = StreamReader::new(stream);
|
||||
let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from);
|
||||
let message_stream = response_to_streaming_message(framed);
|
||||
pin!(message_stream);
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
|
||||
log.write(&message, usage.as_ref().map(|f| &f.usage))?;
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
stream_openai_compat(response, log)
|
||||
}
|
||||
|
||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
||||
|
||||
@@ -1,33 +1,30 @@
|
||||
use anyhow::Result;
|
||||
use async_stream::try_stream;
|
||||
use async_trait::async_trait;
|
||||
use futures::TryStreamExt;
|
||||
use reqwest::StatusCode;
|
||||
use serde_json::{json, Value};
|
||||
use std::collections::HashMap;
|
||||
use std::io;
|
||||
use tokio::pin;
|
||||
use tokio_stream::StreamExt;
|
||||
use tokio_util::codec::{FramedRead, LinesCodec};
|
||||
use tokio_util::io::StreamReader;
|
||||
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::base::{ConfigKey, ModelInfo, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse};
|
||||
use super::errors::ProviderError;
|
||||
use super::formats::openai::{
|
||||
create_request, get_usage, response_to_message, response_to_streaming_message,
|
||||
};
|
||||
use super::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use super::formats::openai_responses::{
|
||||
create_responses_request, get_responses_usage, responses_api_to_message,
|
||||
responses_api_to_streaming_message, ResponsesApiResponse,
|
||||
};
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
get_model, handle_response_openai_compat, handle_status_openai_compat, ImageFormat,
|
||||
get_model, handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat,
|
||||
ImageFormat,
|
||||
};
|
||||
use crate::config::declarative_providers::DeclarativeProviderConfig;
|
||||
use crate::conversation::message::Message;
|
||||
use anyhow::Result;
|
||||
use async_stream::try_stream;
|
||||
use async_trait::async_trait;
|
||||
use futures::{StreamExt, TryStreamExt};
|
||||
use reqwest::StatusCode;
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
use std::io;
|
||||
use tokio::pin;
|
||||
use tokio_util::codec::{FramedRead, LinesCodec};
|
||||
use tokio_util::io::StreamReader;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::base::MessageStream;
|
||||
@@ -286,8 +283,14 @@ impl Provider for OpenAiProvider {
|
||||
log.write(&json_response, Some(&usage))?;
|
||||
Ok((message, ProviderUsage::new(model, usage)))
|
||||
} else {
|
||||
let payload =
|
||||
create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
|
||||
let payload = create_request(
|
||||
model_config,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
let json_response = self
|
||||
@@ -404,12 +407,14 @@ impl Provider for OpenAiProvider {
|
||||
}
|
||||
}))
|
||||
} else {
|
||||
let mut payload =
|
||||
create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?;
|
||||
payload["stream"] = serde_json::Value::Bool(true);
|
||||
payload["stream_options"] = json!({
|
||||
"include_usage": true,
|
||||
});
|
||||
let payload = create_request(
|
||||
&self.model,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&ImageFormat::OpenAi,
|
||||
true,
|
||||
)?;
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
|
||||
let response = self
|
||||
@@ -425,20 +430,7 @@ impl Provider for OpenAiProvider {
|
||||
let _ = log.error(e);
|
||||
})?;
|
||||
|
||||
let stream = response.bytes_stream().map_err(io::Error::other);
|
||||
|
||||
Ok(Box::pin(try_stream! {
|
||||
let stream_reader = StreamReader::new(stream);
|
||||
let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from);
|
||||
|
||||
let message_stream = response_to_streaming_message(framed);
|
||||
pin!(message_stream);
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
stream_openai_compat(response, log)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,12 +3,12 @@ use async_trait::async_trait;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
get_model, handle_response_google_compat, handle_response_openai_compat, is_google_model,
|
||||
RequestLog,
|
||||
get_model, handle_response_google_compat, handle_response_openai_compat,
|
||||
handle_status_openai_compat, is_google_model, stream_openai_compat, RequestLog,
|
||||
};
|
||||
use crate::conversation::message::Message;
|
||||
|
||||
@@ -40,6 +40,7 @@ pub struct OpenRouterProvider {
|
||||
#[serde(skip)]
|
||||
api_client: ApiClient,
|
||||
model: ModelConfig,
|
||||
supports_streaming: bool,
|
||||
#[serde(skip)]
|
||||
name: String,
|
||||
}
|
||||
@@ -62,6 +63,7 @@ impl OpenRouterProvider {
|
||||
Ok(Self {
|
||||
api_client,
|
||||
model,
|
||||
supports_streaming: true,
|
||||
name: Self::metadata().name,
|
||||
})
|
||||
}
|
||||
@@ -211,13 +213,13 @@ async fn create_request_based_on_model(
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
|
||||
if provider.supports_cache_control().await {
|
||||
payload = update_request_for_anthropic(&payload);
|
||||
}
|
||||
|
||||
// Always add transforms: ["middle-out"] for OpenRouter to handle prompts > context size
|
||||
payload
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
@@ -371,4 +373,50 @@ impl Provider for OpenRouterProvider {
|
||||
.model_name
|
||||
.starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC)
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
self.supports_streaming
|
||||
}
|
||||
|
||||
async fn stream(
|
||||
&self,
|
||||
system: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<MessageStream, ProviderError> {
|
||||
let mut payload = create_request(
|
||||
&self.model,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
true,
|
||||
)?;
|
||||
|
||||
if self.supports_cache_control().await {
|
||||
payload = update_request_for_anthropic(&payload);
|
||||
}
|
||||
|
||||
payload
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("transforms".to_string(), json!(["middle-out"]));
|
||||
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
|
||||
let response = self
|
||||
.with_retry(|| async {
|
||||
let resp = self
|
||||
.api_client
|
||||
.response_post("api/v1/chat/completions", &payload)
|
||||
.await?;
|
||||
handle_status_openai_compat(resp).await
|
||||
})
|
||||
.await
|
||||
.inspect_err(|e| {
|
||||
let _ = log.error(e);
|
||||
})?;
|
||||
|
||||
stream_openai_compat(response, log)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,25 +1,16 @@
|
||||
use anyhow::Result;
|
||||
use async_stream::try_stream;
|
||||
use async_trait::async_trait;
|
||||
use futures::TryStreamExt;
|
||||
use serde_json::{json, Value};
|
||||
use std::io;
|
||||
use tokio::pin;
|
||||
use tokio_stream::StreamExt;
|
||||
use tokio_util::codec::{FramedRead, LinesCodec};
|
||||
use tokio_util::io::StreamReader;
|
||||
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::errors::ProviderError;
|
||||
use super::formats::openai::response_to_streaming_message;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
get_model, handle_response_google_compat, handle_response_openai_compat,
|
||||
handle_status_openai_compat, is_google_model, RequestLog,
|
||||
handle_status_openai_compat, is_google_model, stream_openai_compat, RequestLog,
|
||||
};
|
||||
use crate::config::signup_tetrate::TETRATE_DEFAULT_MODEL;
|
||||
use crate::conversation::message::Message;
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
|
||||
@@ -178,6 +169,7 @@ impl Provider for TetrateProvider {
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
let mut log = RequestLog::start(model_config, &payload)?;
|
||||
|
||||
@@ -206,41 +198,31 @@ impl Provider for TetrateProvider {
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<MessageStream, ProviderError> {
|
||||
let mut payload = create_request(
|
||||
let payload = create_request(
|
||||
&self.model,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
true,
|
||||
)?;
|
||||
|
||||
payload["stream"] = json!(true);
|
||||
payload["stream_options"] = json!({
|
||||
"include_usage": true,
|
||||
});
|
||||
|
||||
let resp = self
|
||||
.api_client
|
||||
.response_post("v1/chat/completions", &payload)
|
||||
.await?;
|
||||
|
||||
let response = handle_status_openai_compat(resp).await?;
|
||||
|
||||
let stream = response.bytes_stream().map_err(io::Error::other);
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
|
||||
Ok(Box::pin(try_stream! {
|
||||
let stream_reader = StreamReader::new(stream);
|
||||
let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from);
|
||||
let response = self
|
||||
.with_retry(|| async {
|
||||
let resp = self
|
||||
.api_client
|
||||
.response_post("v1/chat/completions", &payload)
|
||||
.await?;
|
||||
handle_status_openai_compat(resp).await
|
||||
})
|
||||
.await
|
||||
.inspect_err(|e| {
|
||||
let _ = log.error(e);
|
||||
})?;
|
||||
|
||||
let message_stream = response_to_streaming_message(framed);
|
||||
pin!(message_stream);
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
stream_openai_compat(response, log)
|
||||
}
|
||||
|
||||
/// Fetch supported models from Tetrate Agent Router Service API (only models with tool support)
|
||||
|
||||
@@ -162,6 +162,7 @@ impl OllamaInterpreter {
|
||||
&messages,
|
||||
&[], // No tools
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
|
||||
payload["stream"] = json!(false); // needed for the /api/chat endpoint to work
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
use super::base::Usage;
|
||||
use super::base::{MessageStream, Usage};
|
||||
use super::errors::GoogleErrorCode;
|
||||
use crate::config::paths::Paths;
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::errors::ProviderError;
|
||||
use crate::providers::formats::openai::response_to_streaming_message;
|
||||
use anyhow::{anyhow, Result};
|
||||
use async_stream::try_stream;
|
||||
use base64::Engine;
|
||||
use futures::TryStreamExt;
|
||||
use regex::Regex;
|
||||
use reqwest::{Response, StatusCode};
|
||||
use rmcp::model::{AnnotateAble, ImageContent, RawImageContent};
|
||||
@@ -12,9 +15,14 @@ use serde::{Deserialize, Serialize};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::fmt::Display;
|
||||
use std::fs::File;
|
||||
use std::io;
|
||||
use std::io::{BufWriter, Read, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Duration;
|
||||
use tokio::pin;
|
||||
use tokio_stream::StreamExt;
|
||||
use tokio_util::codec::{FramedRead, LinesCodec};
|
||||
use tokio_util::io::StreamReader;
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(Debug, Copy, Clone, Serialize, Deserialize)]
|
||||
@@ -178,19 +186,36 @@ pub async fn handle_response_openai_compat(response: Response) -> Result<Value,
|
||||
})
|
||||
}
|
||||
|
||||
/// Check if the model is a Google model based on the "model" field in the payload.
|
||||
///
|
||||
/// ### Arguments
|
||||
/// - `payload`: The JSON payload as a `serde_json::Value`.
|
||||
///
|
||||
/// ### Returns
|
||||
/// - `bool`: Returns `true` if the model is a Google model, otherwise `false`.
|
||||
pub fn stream_openai_compat(
|
||||
response: Response,
|
||||
mut log: RequestLog,
|
||||
) -> Result<MessageStream, ProviderError> {
|
||||
let stream = response.bytes_stream().map_err(io::Error::other);
|
||||
|
||||
Ok(Box::pin(try_stream! {
|
||||
let stream_reader = StreamReader::new(stream);
|
||||
let framed = FramedRead::new(stream_reader, LinesCodec::new())
|
||||
.map_err(anyhow::Error::from);
|
||||
|
||||
let message_stream = response_to_streaming_message(framed);
|
||||
pin!(message_stream);
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e|
|
||||
ProviderError::RequestFailed(format!("Stream decode error: {}", e))
|
||||
)?;
|
||||
log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?;
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
pub fn is_google_model(payload: &Value) -> bool {
|
||||
if let Some(model) = payload.get("model").and_then(|m| m.as_str()) {
|
||||
// Check if the model name contains "google"
|
||||
return model.to_lowercase().contains("google");
|
||||
}
|
||||
false
|
||||
payload
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or("")
|
||||
.to_lowercase()
|
||||
.contains("google")
|
||||
}
|
||||
|
||||
/// Extracts `StatusCode` from response status or payload error code.
|
||||
|
||||
@@ -1,17 +1,20 @@
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{get_model, handle_response_openai_compat, RequestLog};
|
||||
use super::utils::{
|
||||
get_model, handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat,
|
||||
RequestLog,
|
||||
};
|
||||
use crate::conversation::message::Message;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use crate::providers::base::{
|
||||
ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage,
|
||||
};
|
||||
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use rmcp::model::Tool;
|
||||
use serde_json::Value;
|
||||
|
||||
pub const XAI_API_HOST: &str = "https://api.x.ai/v1";
|
||||
pub const XAI_DEFAULT_MODEL: &str = "grok-code-fast-1";
|
||||
pub const XAI_KNOWN_MODELS: &[&str] = &[
|
||||
@@ -42,6 +45,7 @@ pub struct XaiProvider {
|
||||
#[serde(skip)]
|
||||
api_client: ApiClient,
|
||||
model: ModelConfig,
|
||||
supports_streaming: bool,
|
||||
#[serde(skip)]
|
||||
name: String,
|
||||
}
|
||||
@@ -60,13 +64,12 @@ impl XaiProvider {
|
||||
Ok(Self {
|
||||
api_client,
|
||||
model,
|
||||
supports_streaming: true,
|
||||
name: Self::metadata().name,
|
||||
})
|
||||
}
|
||||
|
||||
async fn post(&self, payload: Value) -> Result<Value, ProviderError> {
|
||||
tracing::debug!("xAI request model: {:?}", self.model.model_name);
|
||||
|
||||
let response = self
|
||||
.api_client
|
||||
.response_post("chat/completions", &payload)
|
||||
@@ -118,6 +121,7 @@ impl Provider for XaiProvider {
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
false,
|
||||
)?;
|
||||
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
@@ -132,4 +136,40 @@ impl Provider for XaiProvider {
|
||||
log.write(&response, Some(&usage))?;
|
||||
Ok((message, ProviderUsage::new(response_model, usage)))
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
self.supports_streaming
|
||||
}
|
||||
|
||||
async fn stream(
|
||||
&self,
|
||||
system: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<MessageStream, ProviderError> {
|
||||
let payload = create_request(
|
||||
&self.model,
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
&super::utils::ImageFormat::OpenAi,
|
||||
true,
|
||||
)?;
|
||||
let mut log = RequestLog::start(&self.model, &payload)?;
|
||||
|
||||
let response = self
|
||||
.with_retry(|| async {
|
||||
let resp = self
|
||||
.api_client
|
||||
.response_post("chat/completions", &payload)
|
||||
.await?;
|
||||
handle_status_openai_compat(resp).await
|
||||
})
|
||||
.await
|
||||
.inspect_err(|e| {
|
||||
let _ = log.error(e);
|
||||
})?;
|
||||
|
||||
stream_openai_compat(response, log)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -100,81 +100,81 @@ export default function ExternalBackendSection() {
|
||||
<Card className="pb-2">
|
||||
<CardHeader className="pb-0">
|
||||
<CardTitle>Goose Server</CardTitle>
|
||||
<CardDescription>
|
||||
By default goose launches a server for you, use this to connect to an external goose
|
||||
server
|
||||
</CardDescription>
|
||||
</CardHeader>
|
||||
<CardContent className="pt-4 space-y-4 px-4">
|
||||
<div className="flex items-center justify-between">
|
||||
<div>
|
||||
<h3 className="text-text-default text-xs">Use external server</h3>
|
||||
<p className="text-xs text-text-muted max-w-md mt-[2px]">
|
||||
Connect to a goose server running elsewhere (requires app restart)
|
||||
</p>
|
||||
</div>
|
||||
<div className="flex items-center">
|
||||
<Switch
|
||||
checked={config.enabled}
|
||||
onCheckedChange={(checked) => saveConfig(updateField('enabled', checked))}
|
||||
disabled={isSaving}
|
||||
variant="mono"
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{config.enabled && (
|
||||
<>
|
||||
<div className="space-y-2">
|
||||
<label htmlFor="external-url" className="text-text-default text-xs">
|
||||
Server URL
|
||||
</label>
|
||||
<Input
|
||||
id="external-url"
|
||||
type="url"
|
||||
placeholder="http://127.0.0.1:3000"
|
||||
value={config.url}
|
||||
onChange={(e) => handleUrlChange(e.target.value)}
|
||||
onBlur={handleUrlBlur}
|
||||
<CardDescription>
|
||||
By default goose launches a server for you, use this to connect to an external goose
|
||||
server
|
||||
</CardDescription>
|
||||
</CardHeader>
|
||||
<CardContent className="pt-4 space-y-4 px-4">
|
||||
<div className="flex items-center justify-between">
|
||||
<div>
|
||||
<h3 className="text-text-default text-xs">Use external server</h3>
|
||||
<p className="text-xs text-text-muted max-w-md mt-[2px]">
|
||||
Connect to a goose server running elsewhere (requires app restart)
|
||||
</p>
|
||||
</div>
|
||||
<div className="flex items-center">
|
||||
<Switch
|
||||
checked={config.enabled}
|
||||
onCheckedChange={(checked) => saveConfig(updateField('enabled', checked))}
|
||||
disabled={isSaving}
|
||||
className={urlError ? 'border-red-500' : ''}
|
||||
variant="mono"
|
||||
/>
|
||||
{urlError && (
|
||||
<p className="text-xs text-red-500 flex items-center gap-1">
|
||||
<AlertCircle size={12} />
|
||||
{urlError}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{config.enabled && (
|
||||
<>
|
||||
<div className="space-y-2">
|
||||
<label htmlFor="external-url" className="text-text-default text-xs">
|
||||
Server URL
|
||||
</label>
|
||||
<Input
|
||||
id="external-url"
|
||||
type="url"
|
||||
placeholder="http://127.0.0.1:3000"
|
||||
value={config.url}
|
||||
onChange={(e) => handleUrlChange(e.target.value)}
|
||||
onBlur={handleUrlBlur}
|
||||
disabled={isSaving}
|
||||
className={urlError ? 'border-red-500' : ''}
|
||||
/>
|
||||
{urlError && (
|
||||
<p className="text-xs text-red-500 flex items-center gap-1">
|
||||
<AlertCircle size={12} />
|
||||
{urlError}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="space-y-2">
|
||||
<label htmlFor="external-secret" className="text-text-default text-xs">
|
||||
Secret Key
|
||||
</label>
|
||||
<Input
|
||||
id="external-secret"
|
||||
type="password"
|
||||
placeholder="Enter the server's secret key"
|
||||
value={config.secret}
|
||||
onChange={(e) => updateField('secret', e.target.value)}
|
||||
onBlur={() => saveConfig(config)}
|
||||
disabled={isSaving}
|
||||
/>
|
||||
<p className="text-xs text-text-muted">
|
||||
The secret key configured on the goosed server (GOOSE_SERVER__SECRET_KEY)
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="space-y-2">
|
||||
<label htmlFor="external-secret" className="text-text-default text-xs">
|
||||
Secret Key
|
||||
</label>
|
||||
<Input
|
||||
id="external-secret"
|
||||
type="password"
|
||||
placeholder="Enter the server's secret key"
|
||||
value={config.secret}
|
||||
onChange={(e) => updateField('secret', e.target.value)}
|
||||
onBlur={() => saveConfig(config)}
|
||||
disabled={isSaving}
|
||||
/>
|
||||
<p className="text-xs text-text-muted">
|
||||
The secret key configured on the goosed server (GOOSE_SERVER__SECRET_KEY)
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div className="bg-amber-50 dark:bg-amber-950 border border-amber-200 dark:border-amber-800 rounded-md p-3">
|
||||
<p className="text-xs text-amber-800 dark:text-amber-200">
|
||||
<strong>Note:</strong> Changes require restarting Goose to take effect. New chat
|
||||
windows will connect to the external server.
|
||||
</p>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</CardContent>
|
||||
</Card>
|
||||
<div className="bg-amber-50 dark:bg-amber-950 border border-amber-200 dark:border-amber-800 rounded-md p-3">
|
||||
<p className="text-xs text-amber-800 dark:text-amber-200">
|
||||
<strong>Note:</strong> Changes require restarting Goose to take effect. New chat
|
||||
windows will connect to the external server.
|
||||
</p>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</CardContent>
|
||||
</Card>
|
||||
</section>
|
||||
);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user