fix(goose): only send agent-session-id when a session exists (#6657)

Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
Adrian Cole
2026-01-27 13:43:05 +09:00
committed by GitHub
parent 3dae12765f
commit f5b402bbdf
37 changed files with 398 additions and 383 deletions
+9 -8
View File
@@ -116,7 +116,11 @@ impl AnthropicProvider {
headers
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<ApiResponse, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<ApiResponse, ProviderError> {
let mut request = self.api_client.request(session_id, "v1/messages");
for (key, value) in self.get_conditional_headers() {
@@ -198,7 +202,7 @@ impl Provider for AnthropicProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -228,11 +232,8 @@ impl Provider for AnthropicProvider {
Ok((message, provider_usage))
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
let response = self.api_client.api_get(session_id, "v1/models").await?;
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = self.api_client.request(None, "v1/models").api_get().await?;
if response.status != StatusCode::OK {
return Err(map_http_error_to_provider_error(
@@ -268,7 +269,7 @@ impl Provider for AnthropicProvider {
.unwrap()
.insert("stream".to_string(), Value::Bool(true));
let mut request = self.api_client.request(session_id, "v1/messages");
let mut request = self.api_client.request(Some(session_id), "v1/messages");
let mut log = RequestLog::start(&self.model, &payload)?;
for (key, value) in self.get_conditional_headers() {
+51 -69
View File
@@ -196,7 +196,7 @@ pub struct ApiRequestBuilder<'a> {
client: &'a ApiClient,
path: &'a str,
headers: HeaderMap,
session_id: &'a str,
session_id: Option<&'a str>,
}
impl ApiClient {
@@ -273,10 +273,15 @@ impl ApiClient {
Ok(self)
}
pub fn request<'a>(&'a self, session_id: &'a str, path: &'a str) -> ApiRequestBuilder<'a> {
/// - `session_id`: Use `None` only for configuration or pre-session tasks.
pub fn request<'a>(
&'a self,
session_id: Option<&'a str>,
path: &'a str,
) -> ApiRequestBuilder<'a> {
ApiRequestBuilder {
client: self,
session_id,
session_id: session_id.filter(|id| !id.is_empty()),
path,
headers: HeaderMap::new(),
}
@@ -284,7 +289,7 @@ impl ApiClient {
pub async fn api_post(
&self,
session_id: &str,
session_id: Option<&str>,
path: &str,
payload: &Value,
) -> Result<ApiResponse> {
@@ -293,18 +298,18 @@ impl ApiClient {
pub async fn response_post(
&self,
session_id: &str,
session_id: Option<&str>,
path: &str,
payload: &Value,
) -> Result<Response> {
self.request(session_id, path).response_post(payload).await
}
pub async fn api_get(&self, session_id: &str, path: &str) -> Result<ApiResponse> {
pub async fn api_get(&self, session_id: Option<&str>, path: &str) -> Result<ApiResponse> {
self.request(session_id, path).api_get().await
}
pub async fn response_get(&self, session_id: &str, path: &str) -> Result<Response> {
pub async fn response_get(&self, session_id: Option<&str>, path: &str) -> Result<Response> {
self.request(session_id, path).response_get().await
}
@@ -373,10 +378,16 @@ impl<'a> ApiRequestBuilder<'a> {
F: FnOnce(url::Url, &Client) -> reqwest::RequestBuilder,
{
let url = self.client.build_url(self.path)?;
let mut request = request_builder(url, &self.client.client);
request = request.headers(self.headers.clone());
let mut headers = self.headers.clone();
headers.remove(SESSION_ID_HEADER);
if let Some(session_id) = self.session_id {
let header_name = HeaderName::from_static(SESSION_ID_HEADER);
let header_value = HeaderValue::from_str(session_id)?;
headers.insert(header_name, header_value);
}
request = request.header(SESSION_ID_HEADER, self.session_id);
let mut request = request_builder(url, &self.client.client);
request = request.headers(headers);
request = match &self.client.auth {
AuthMethod::BearerToken(token) => {
@@ -411,69 +422,40 @@ impl fmt::Debug for ApiClient {
#[cfg(test)]
mod tests {
use super::*;
use test_case::test_case;
#[tokio::test]
async fn test_session_id_header_injection() {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
let builder = client.request("test-session_id-456", "/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
#[test_case(Some("test-session_id-456"), None, Some("test-session_id-456"); "header set")]
#[test_case(Some("new-session"), Some(("Agent-Session-Id", "old-session")), Some("new-session"); "replaces existing")]
#[test_case(None, Some(("Agent-Session-Id", "old-session")), None; "removes existing on none")]
#[test_case(Some(""), Some(("agent-session-id", "old-session")), None; "removes existing on empty")]
fn test_session_id_header(
session_id: Option<&str>,
existing_header: Option<(&str, &str)>,
expected: Option<&str>,
) {
let runtime = tokio::runtime::Runtime::new().unwrap();
runtime.block_on(async {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
let headers = request.build().unwrap().headers().clone();
let mut builder = client.request(session_id, "/test");
if let Some((key, value)) = existing_header {
builder = builder.header(key, value).unwrap();
}
let request = builder
.send_request(|url, client| client.get(url))
.await
.unwrap();
assert!(headers.contains_key(SESSION_ID_HEADER));
assert_eq!(
headers.get(SESSION_ID_HEADER).unwrap().to_str().unwrap(),
"test-session_id-456"
);
}
let headers = request.build().unwrap().headers().clone();
#[tokio::test]
async fn test_session_id_header_with_different_id() {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
let builder = client.request("another-session_id-789", "/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
.unwrap();
let headers = request.build().unwrap().headers().clone();
assert!(headers.contains_key(SESSION_ID_HEADER));
assert_eq!(
headers.get(SESSION_ID_HEADER).unwrap().to_str().unwrap(),
"another-session_id-789"
);
}
#[tokio::test]
async fn test_session_id_header_always_present() {
let client = ApiClient::new(
"http://localhost:8080".to_string(),
AuthMethod::BearerToken("test-token".to_string()),
)
.unwrap();
let builder = client.request("required-session_id", "/test");
let request = builder
.send_request(|url, client| client.get(url))
.await
.unwrap();
let headers = request.build().unwrap().headers().clone();
assert!(headers.contains_key(SESSION_ID_HEADER));
let actual = headers
.get(SESSION_ID_HEADER)
.and_then(|value| value.to_str().ok());
assert_eq!(actual, expected);
});
}
}
+2 -6
View File
@@ -1,10 +1,7 @@
use crate::model::ModelConfig;
use crate::providers::retry::{retry_operation, RetryConfig};
pub async fn detect_provider_from_api_key(
session_id: &str,
api_key: &str,
) -> Option<(String, Vec<String>)> {
pub async fn detect_provider_from_api_key(api_key: &str) -> Option<(String, Vec<String>)> {
let provider_tests = vec![
("anthropic", "ANTHROPIC_API_KEY"),
("openai", "OPENAI_API_KEY"),
@@ -18,7 +15,6 @@ pub async fn detect_provider_from_api_key(
.into_iter()
.map(|(provider_name, env_key)| {
let api_key = api_key.to_string();
let session_id = session_id.to_string();
tokio::spawn(async move {
let original_value = std::env::var(env_key).ok();
std::env::set_var(env_key, &api_key);
@@ -31,7 +27,7 @@ pub async fn detect_provider_from_api_key(
{
Ok(provider) => {
match retry_operation(&RetryConfig::default(), || async {
provider.fetch_supported_models(&session_id).await
provider.fetch_supported_models().await
})
.await
{
+6 -2
View File
@@ -99,7 +99,11 @@ impl AzureProvider {
})
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
// Build the path for Azure OpenAI
let path = format!(
"openai/deployments/{}/chat/completions?api-version={}",
@@ -147,7 +151,7 @@ impl Provider for AzureProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
+18 -15
View File
@@ -368,9 +368,12 @@ pub trait Provider: Send + Sync {
// Internal implementation of complete, used by complete_fast and complete
// Providers should override this to implement their actual completion logic
//
/// # Parameters
/// - `session_id`: Use `None` only for configuration or pre-session tasks.
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -386,7 +389,7 @@ pub trait Provider: Send + Sync {
tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
let model_config = self.get_model_config();
self.complete_with_model(session_id, &model_config, system, messages, tools)
self.complete_with_model(Some(session_id), &model_config, system, messages, tools)
.await
}
@@ -402,7 +405,7 @@ pub trait Provider: Send + Sync {
let fast_config = model_config.use_fast_model();
match self
.complete_with_model(session_id, &fast_config, system, messages, tools)
.complete_with_model(Some(session_id), &fast_config, system, messages, tools)
.await
{
Ok(result) => Ok(result),
@@ -414,8 +417,14 @@ pub trait Provider: Send + Sync {
e,
model_config.model_name
);
self.complete_with_model(session_id, &model_config, system, messages, tools)
.await
self.complete_with_model(
Some(session_id),
&model_config,
system,
messages,
tools,
)
.await
} else {
Err(e)
}
@@ -430,19 +439,13 @@ pub trait Provider: Send + Sync {
RetryConfig::default()
}
async fn fetch_supported_models(
&self,
_session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
Ok(None)
}
/// Fetch models filtered by canonical registry and usability
async fn fetch_recommended_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
let all_models = match self.fetch_supported_models(session_id).await? {
async fn fetch_recommended_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let all_models = match self.fetch_supported_models().await? {
Some(models) => models,
None => return Ok(None),
};
@@ -490,7 +493,7 @@ pub trait Provider: Send + Sync {
false
}
async fn supports_cache_control(&self, _session_id: &str) -> bool {
async fn supports_cache_control(&self) -> bool {
false
}
+16 -2
View File
@@ -11,6 +11,7 @@ use async_trait::async_trait;
use aws_sdk_bedrockruntime::config::ProvideCredentials;
use aws_sdk_bedrockruntime::operation::converse::ConverseError;
use aws_sdk_bedrockruntime::{types as bedrock, Client};
use reqwest::header::HeaderValue;
use rmcp::model::Tool;
use serde_json::Value;
@@ -18,6 +19,7 @@ use serde_json::Value;
use super::formats::bedrock::{
from_bedrock_message, from_bedrock_usage, to_bedrock_message, to_bedrock_tool_config,
};
use crate::session_context::SESSION_ID_HEADER;
pub const BEDROCK_DOC_LINK: &str =
"https://docs.aws.amazon.com/bedrock/latest/userguide/models-supported.html";
@@ -130,6 +132,7 @@ impl BedrockProvider {
async fn converse(
&self,
session_id: Option<&str>,
system: &str,
messages: &[Message],
tools: &[Tool],
@@ -153,6 +156,17 @@ impl BedrockProvider {
request = request.tool_config(to_bedrock_tool_config(tools)?);
}
let mut request = request.customize();
if let Some(session_id) = session_id.filter(|id| !id.is_empty()) {
let session_id = session_id.to_string();
request = request.mutate_request(move |req| {
if let Ok(value) = HeaderValue::from_str(&session_id) {
req.headers_mut().insert(SESSION_ID_HEADER, value);
}
});
}
let response = request
.send()
.await
@@ -227,7 +241,7 @@ impl Provider for BedrockProvider {
)]
async fn complete_with_model(
&self,
_session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -236,7 +250,7 @@ impl Provider for BedrockProvider {
let model_name = model_config.model_name.clone();
let (bedrock_message, bedrock_usage) = self
.with_retry(|| self.converse(system, messages, tools))
.with_retry(|| self.converse(session_id, system, messages, tools))
.await?;
let usage = bedrock_usage
@@ -550,9 +550,7 @@ async fn check_provider(
}
};
// Provider probe runs outside any user session; use an ephemeral id.
let session_id = uuid::Uuid::new_v4().to_string();
let fetched_models = match provider.fetch_supported_models(&session_id).await {
let fetched_models = match provider.fetch_supported_models().await {
Ok(Some(models)) => {
println!(" ✓ Fetched {} models", models.len());
models
+20 -9
View File
@@ -7,6 +7,7 @@ use crate::providers::errors::ProviderError;
use crate::providers::formats::openai_responses::responses_api_to_streaming_message;
use crate::providers::retry::ProviderRetry;
use crate::providers::utils::handle_status_openai_compat;
use crate::session_context::SESSION_ID_HEADER;
use anyhow::{anyhow, Result};
use async_stream::try_stream;
use async_trait::async_trait;
@@ -16,6 +17,7 @@ use chrono::{DateTime, Utc};
use futures::{StreamExt, TryStreamExt};
use jsonwebtoken::jwk::JwkSet;
use jsonwebtoken::{decode, decode_header, DecodingKey, Validation};
use reqwest::header::{HeaderName, HeaderValue};
use rmcp::model::{RawContent, Role, Tool};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
@@ -789,7 +791,11 @@ impl ChatGptCodexProvider {
})
}
async fn post_streaming(&self, payload: &Value) -> Result<reqwest::Response, ProviderError> {
async fn post_streaming(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<reqwest::Response, ProviderError> {
let token_data = self
.auth_provider
.get_valid_token()
@@ -805,6 +811,14 @@ impl ChatGptCodexProvider {
);
}
if let Some(session_id) = session_id.filter(|id| !id.is_empty()) {
headers.insert(
HeaderName::from_static(SESSION_ID_HEADER),
HeaderValue::from_str(session_id)
.map_err(|e| ProviderError::ExecutionError(e.to_string()))?,
);
}
let client = reqwest::Client::new();
let response = client
.post(format!("{}/responses", CODEX_API_ENDPOINT))
@@ -856,7 +870,7 @@ impl Provider for ChatGptCodexProvider {
)]
async fn complete_with_model(
&self,
_session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -870,7 +884,7 @@ impl Provider for ChatGptCodexProvider {
let response = self
.with_retry(|| async {
let payload_clone = payload.clone();
self.post_streaming(&payload_clone).await
self.post_streaming(session_id, &payload_clone).await
})
.await?;
@@ -914,7 +928,7 @@ impl Provider for ChatGptCodexProvider {
async fn stream(
&self,
_session_id: &str,
session_id: &str,
system: &str,
messages: &[Message],
tools: &[Tool],
@@ -926,7 +940,7 @@ impl Provider for ChatGptCodexProvider {
let response = self
.with_retry(|| async {
let payload_clone = payload.clone();
self.post_streaming(&payload_clone).await
self.post_streaming(Some(session_id), &payload_clone).await
})
.await?;
@@ -953,10 +967,7 @@ impl Provider for ChatGptCodexProvider {
Ok(())
}
async fn fetch_supported_models(
&self,
_session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
Ok(Some(
CHATGPT_CODEX_KNOWN_MODELS
.iter()
+1 -1
View File
@@ -417,7 +417,7 @@ impl Provider for ClaudeCodeProvider {
)]
async fn complete_with_model(
&self,
_session_id: &str,
_session_id: Option<&str>, // create_session == YYYYMMDD_N, but --session-id requires a UUID
model_config: &ModelConfig,
system: &str,
messages: &[Message],
+1 -1
View File
@@ -507,7 +507,7 @@ impl Provider for CodexProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
_session_id: Option<&str>, // CLI has no external session-id flag to propagate.
model_config: &ModelConfig,
system: &str,
messages: &[Message],
+1 -1
View File
@@ -352,7 +352,7 @@ impl Provider for CursorAgentProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
_session_id: Option<&str>, // CLI has no external session-id flag to propagate.
model_config: &ModelConfig,
system: &str,
messages: &[Message],
+7 -9
View File
@@ -206,7 +206,7 @@ impl DatabricksProvider {
async fn post(
&self,
session_id: &str,
session_id: Option<&str>,
payload: Value,
model_name: Option<&str>,
) -> Result<Value, ProviderError> {
@@ -257,7 +257,7 @@ impl Provider for DatabricksProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -314,7 +314,7 @@ impl Provider for DatabricksProvider {
.with_retry(|| async {
let resp = self
.api_client
.response_post(session_id, &path, &payload)
.response_post(Some(session_id), &path, &payload)
.await?;
if !resp.status().is_success() {
let status = resp.status();
@@ -352,13 +352,11 @@ impl Provider for DatabricksProvider {
.map_err(|e| ProviderError::ExecutionError(e.to_string()))
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = match self
.api_client
.response_get(session_id, "api/2.0/serving-endpoints")
.request(None, "api/2.0/serving-endpoints")
.response_get()
.await
{
Ok(resp) => resp,
@@ -434,7 +432,7 @@ impl EmbeddingCapable for DatabricksProvider {
});
let response = self
.with_retry(|| self.post(session_id, request.clone(), None))
.with_retry(|| self.post(Some(session_id), request.clone(), None))
.await?;
let embeddings = response["data"]
+27 -16
View File
@@ -25,6 +25,7 @@ use crate::providers::formats::gcpvertexai::{
use crate::providers::gcpauth::GcpAuth;
use crate::providers::retry::RetryConfig;
use crate::providers::utils::RequestLog;
use crate::session_context::SESSION_ID_HEADER;
use rmcp::model::Tool;
/// Base URL for GCP Vertex AI documentation
@@ -262,6 +263,7 @@ impl GcpVertexAIProvider {
async fn send_request_with_retry(
&self,
session_id: Option<&str>,
url: Url,
payload: &Value,
) -> Result<reqwest::Response, ProviderError> {
@@ -285,11 +287,17 @@ impl GcpVertexAIProvider {
.await
.map_err(|e| ProviderError::Authentication(e.to_string()))?;
let response = self
let mut request = self
.client
.post(url.clone())
.json(payload)
.header("Authorization", auth_header)
.header("Authorization", auth_header);
if let Some(session_id) = session_id.filter(|id| !id.is_empty()) {
request = request.header(SESSION_ID_HEADER, session_id);
}
let response = request
.send()
.await
.map_err(|e| ProviderError::RequestFailed(e.to_string()))?;
@@ -348,6 +356,7 @@ impl GcpVertexAIProvider {
async fn post_with_location(
&self,
session_id: Option<&str>,
payload: &Value,
context: &RequestContext,
location: &str,
@@ -356,7 +365,9 @@ impl GcpVertexAIProvider {
.build_request_url(context.provider(), location, false)
.map_err(|e| ProviderError::RequestFailed(e.to_string()))?;
let response = self.send_request_with_retry(url, payload).await?;
let response = self
.send_request_with_retry(session_id, url, payload)
.await?;
response
.json::<Value>()
@@ -366,11 +377,12 @@ impl GcpVertexAIProvider {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
context: &RequestContext,
) -> Result<Value, ProviderError> {
let result = self
.post_with_location(payload, context, &self.location)
.post_with_location(session_id, payload, context, &self.location)
.await;
if self.location == context.model.known_location().to_string() || result.is_ok() {
@@ -387,7 +399,7 @@ impl GcpVertexAIProvider {
"Trying known location {known_location} for {model_name} instead of {configured_location}: {msg}"
);
self.post_with_location(payload, context, &known_location)
self.post_with_location(session_id, payload, context, &known_location)
.await
}
_ => result,
@@ -396,6 +408,7 @@ impl GcpVertexAIProvider {
async fn post_stream_with_location(
&self,
session_id: Option<&str>,
payload: &Value,
context: &RequestContext,
location: &str,
@@ -404,16 +417,17 @@ impl GcpVertexAIProvider {
.build_request_url(context.provider(), location, true)
.map_err(|e| ProviderError::RequestFailed(e.to_string()))?;
self.send_request_with_retry(url, payload).await
self.send_request_with_retry(session_id, url, payload).await
}
async fn post_stream(
&self,
session_id: Option<&str>,
payload: &Value,
context: &RequestContext,
) -> Result<reqwest::Response, ProviderError> {
let result = self
.post_stream_with_location(payload, context, &self.location)
.post_stream_with_location(session_id, payload, context, &self.location)
.await;
if self.location == context.model.known_location().to_string() || result.is_ok() {
@@ -430,7 +444,7 @@ impl GcpVertexAIProvider {
"Trying known location {known_location} for {model_name} instead of {configured_location}: {msg}"
);
self.post_stream_with_location(payload, context, &known_location)
self.post_stream_with_location(session_id, payload, context, &known_location)
.await
}
_ => result,
@@ -593,7 +607,7 @@ impl Provider for GcpVertexAIProvider {
)]
async fn complete_with_model(
&self,
_session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -603,7 +617,7 @@ impl Provider for GcpVertexAIProvider {
let (request, context) = create_request(model_config, system, messages, tools)?;
// Send request and process response
let response = self.post(&request, &context).await?;
let response = self.post(session_id, &request, &context).await?;
let usage = get_usage(&response, &context)?;
let mut log = RequestLog::start(model_config, &request)?;
@@ -627,7 +641,7 @@ impl Provider for GcpVertexAIProvider {
async fn stream(
&self,
_session_id: &str,
session_id: &str,
system: &str,
messages: &[Message],
tools: &[Tool],
@@ -644,7 +658,7 @@ impl Provider for GcpVertexAIProvider {
let mut log = RequestLog::start(&model_config, &request)?;
let response = self
.post_stream(&request, &context)
.post_stream(Some(session_id), &request, &context)
.await
.inspect_err(|e| {
let _ = log.error(e);
@@ -672,10 +686,7 @@ impl Provider for GcpVertexAIProvider {
}))
}
async fn fetch_supported_models(
&self,
_session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let models: Vec<String> = KNOWN_MODELS.iter().map(|s| s.to_string()).collect();
let filtered = self.filter_by_org_policy(models).await;
Ok(Some(filtered))
+1 -1
View File
@@ -264,7 +264,7 @@ impl Provider for GeminiCliProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
_session_id: Option<&str>, // CLI has no external session-id flag to propagate.
_model_config: &ModelConfig,
system: &str,
messages: &[Message],
+8 -7
View File
@@ -169,7 +169,11 @@ impl GithubCopilotProvider {
})
}
async fn post(&self, session_id: &str, payload: &mut Value) -> Result<Response, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &mut Value,
) -> Result<Response, ProviderError> {
let (endpoint, token) = self.get_api_info().await?;
let auth = AuthMethod::BearerToken(token);
let mut headers = self.get_github_headers();
@@ -411,7 +415,7 @@ impl Provider for GithubCopilotProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -469,7 +473,7 @@ impl Provider for GithubCopilotProvider {
let response = self
.with_retry(|| async {
let mut payload_clone = payload.clone();
let resp = self.post(session_id, &mut payload_clone).await?;
let resp = self.post(Some(session_id), &mut payload_clone).await?;
handle_status_openai_compat(resp).await
})
.await
@@ -480,10 +484,7 @@ impl Provider for GithubCopilotProvider {
stream_openai_compat(response, log)
}
async fn fetch_supported_models(
&self,
_session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let (endpoint, token) = self.get_api_info().await?;
let url = format!("{}/models", endpoint);
+7 -9
View File
@@ -93,7 +93,7 @@ impl GoogleProvider {
async fn post(
&self,
session_id: &str,
session_id: Option<&str>,
model_name: &str,
payload: &Value,
) -> Result<Value, ProviderError> {
@@ -107,7 +107,7 @@ impl GoogleProvider {
async fn post_stream(
&self,
session_id: &str,
session_id: Option<&str>,
model_name: &str,
payload: &Value,
) -> Result<reqwest::Response, ProviderError> {
@@ -151,7 +151,7 @@ impl Provider for GoogleProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -178,13 +178,11 @@ impl Provider for GoogleProvider {
Ok((message, provider_usage))
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = self
.api_client
.response_get(session_id, "v1beta/models")
.request(None, "v1beta/models")
.response_get()
.await?;
let json: serde_json::Value = response.json().await?;
let arr = match json.get("models").and_then(|v| v.as_array()) {
@@ -216,7 +214,7 @@ impl Provider for GoogleProvider {
let response = self
.with_retry(|| async {
self.post_stream(session_id, &self.model.model_name, &payload)
self.post_stream(Some(session_id), &self.model.model_name, &payload)
.await
})
.await
+12 -17
View File
@@ -342,7 +342,7 @@ impl Provider for LeadWorkerProvider {
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
_model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -393,7 +393,10 @@ impl Provider for LeadWorkerProvider {
}
// Make the completion request
let result = provider.complete(session_id, system, messages, tools).await;
let model_config = provider.get_model_config();
let result = provider
.complete_with_model(session_id, &model_config, system, messages, tools)
.await;
// For technical failures, try with default model (lead provider) instead
let final_result = match &result {
@@ -401,9 +404,10 @@ impl Provider for LeadWorkerProvider {
tracing::warn!("Technical failure with {} provider, retrying with default model (lead provider)", provider_type);
// Try with lead provider as the default/fallback for technical failures
let model_config = self.lead_provider.get_model_config();
let default_result = self
.lead_provider
.complete(session_id, system, messages, tools)
.complete_with_model(session_id, &model_config, system, messages, tools)
.await;
match &default_result {
@@ -428,19 +432,10 @@ impl Provider for LeadWorkerProvider {
final_result
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
// Combine models from both providers
let lead_models = self
.lead_provider
.fetch_supported_models(session_id)
.await?;
let worker_models = self
.worker_provider
.fetch_supported_models(session_id)
.await?;
let lead_models = self.lead_provider.fetch_supported_models().await?;
let worker_models = self.worker_provider.fetch_supported_models().await?;
match (lead_models, worker_models) {
(Some(lead), Some(worker)) => {
@@ -517,7 +512,7 @@ mod tests {
async fn complete_with_model(
&self,
_session_id: &str,
_session_id: Option<&str>,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
@@ -703,7 +698,7 @@ mod tests {
async fn complete_with_model(
&self,
_session_id: &str,
_session_id: Option<&str>,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
+18 -13
View File
@@ -73,8 +73,12 @@ impl LiteLLMProvider {
})
}
async fn fetch_models(&self, session: &str) -> Result<Vec<ModelInfo>, ProviderError> {
let response = self.api_client.response_get(session, "model/info").await?;
async fn fetch_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
let response = self
.api_client
.request(None, "model/info")
.response_get()
.await?;
if !response.status().is_success() {
return Err(ProviderError::RequestFailed(format!(
@@ -112,7 +116,11 @@ impl LiteLLMProvider {
Ok(models)
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, &self.base_path, payload)
@@ -168,7 +176,7 @@ impl Provider for LiteLLMProvider {
#[tracing::instrument(skip_all, name = "provider_complete")]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -183,7 +191,7 @@ impl Provider for LiteLLMProvider {
false,
)?;
if self.supports_cache_control(session_id).await {
if self.supports_cache_control().await {
payload = update_request_for_cache_control(&payload);
}
@@ -206,8 +214,8 @@ impl Provider for LiteLLMProvider {
true
}
async fn supports_cache_control(&self, session_id: &str) -> bool {
if let Ok(models) = self.fetch_models(session_id).await {
async fn supports_cache_control(&self) -> bool {
if let Ok(models) = self.fetch_models().await {
if let Some(model_info) = models.iter().find(|m| m.name == self.model.model_name) {
return model_info.supports_cache_control.unwrap_or(false);
}
@@ -216,11 +224,8 @@ impl Provider for LiteLLMProvider {
self.model.model_name.to_lowercase().contains("claude")
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
match self.fetch_models(session_id).await {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
match self.fetch_models().await {
Ok(models) => {
let model_names: Vec<String> = models.into_iter().map(|m| m.name).collect();
Ok(Some(model_names))
@@ -251,7 +256,7 @@ impl EmbeddingCapable for LiteLLMProvider {
let response = self
.api_client
.response_post(session_id, "v1/embeddings", &payload)
.response_post(Some(session_id), "v1/embeddings", &payload)
.await?;
let response_text = response.text().await?;
let response_json: Value = serde_json::from_str(&response_text)?;
+10 -8
View File
@@ -119,7 +119,11 @@ impl OllamaProvider {
})
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, "v1/chat/completions", payload)
@@ -173,7 +177,7 @@ impl Provider for OllamaProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -273,7 +277,7 @@ impl Provider for OllamaProvider {
.with_retry(|| async {
let resp = self
.api_client
.response_post(session_id, "v1/chat/completions", &payload)
.response_post(Some(session_id), "v1/chat/completions", &payload)
.await?;
handle_status_openai_compat(resp).await
})
@@ -284,13 +288,11 @@ impl Provider for OllamaProvider {
stream_openai_compat(response, log)
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = self
.api_client
.response_get(session_id, "api/tags")
.request(None, "api/tags")
.response_get()
.await
.map_err(|e| ProviderError::RequestFailed(format!("Failed to fetch models: {}", e)))?;
+13 -11
View File
@@ -192,7 +192,11 @@ impl OpenAiProvider {
model_name.starts_with("gpt-5-codex") || model_name.starts_with("gpt-5.1-codex")
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, &self.base_path, payload)
@@ -202,7 +206,7 @@ impl OpenAiProvider {
async fn post_responses(
&self,
session_id: &str,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
@@ -253,7 +257,7 @@ impl Provider for OpenAiProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -323,14 +327,12 @@ impl Provider for OpenAiProvider {
}
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let models_path = self.base_path.replace("v1/chat/completions", "v1/models");
let response = self
.api_client
.response_get(session_id, &models_path)
.request(None, &models_path)
.response_get()
.await?;
let json = handle_response_openai_compat(response).await?;
if let Some(err_obj) = json.get("error") {
@@ -388,7 +390,7 @@ impl Provider for OpenAiProvider {
let payload_clone = payload.clone();
let resp = self
.api_client
.response_post(session_id, "v1/responses", &payload_clone)
.response_post(Some(session_id), "v1/responses", &payload_clone)
.await?;
handle_status_openai_compat(resp).await
})
@@ -426,7 +428,7 @@ impl Provider for OpenAiProvider {
.with_retry(|| async {
let resp = self
.api_client
.response_post(session_id, &self.base_path, &payload)
.response_post(Some(session_id), &self.base_path, &payload)
.await?;
handle_status_openai_compat(resp).await
})
@@ -479,7 +481,7 @@ impl EmbeddingCapable for OpenAiProvider {
let request_value = serde_json::to_value(request_clone)
.map_err(|e| ProviderError::ExecutionError(e.to_string()))?;
self.api_client
.api_post(session_id, "v1/embeddings", &request_value)
.api_post(Some(session_id), "v1/embeddings", &request_value)
.await
.map_err(|e| ProviderError::ExecutionError(e.to_string()))
})
+20 -12
View File
@@ -69,7 +69,11 @@ impl OpenRouterProvider {
})
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, "api/v1/chat/completions", payload)
@@ -189,7 +193,7 @@ fn is_gemini_model(model_name: &str) -> bool {
async fn create_request_based_on_model(
provider: &OpenRouterProvider,
session_id: &str,
session_id: Option<&str>,
system: &str,
messages: &[Message],
tools: &[Tool],
@@ -203,7 +207,13 @@ async fn create_request_based_on_model(
false,
)?;
if provider.supports_cache_control(session_id).await {
if let Some(session_id) = session_id.filter(|id| !id.is_empty()) {
if let Some(obj) = payload.as_object_mut() {
obj.insert("user".to_string(), Value::String(session_id.to_string()));
}
}
if provider.supports_cache_control().await {
payload = update_request_for_anthropic(&payload);
}
@@ -254,7 +264,7 @@ impl Provider for OpenRouterProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -287,15 +297,13 @@ impl Provider for OpenRouterProvider {
}
/// Fetch supported models from OpenRouter API (only models with tool support)
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
// Handle request failures gracefully
// If the request fails, fall back to manual entry
let response = match self
.api_client
.response_get(session_id, "api/v1/models")
.request(None, "api/v1/models")
.response_get()
.await
{
Ok(response) => response,
@@ -370,7 +378,7 @@ impl Provider for OpenRouterProvider {
Ok(Some(models))
}
async fn supports_cache_control(&self, _session_id: &str) -> bool {
async fn supports_cache_control(&self) -> bool {
self.model
.model_name
.starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC)
@@ -396,7 +404,7 @@ impl Provider for OpenRouterProvider {
true,
)?;
if self.supports_cache_control(session_id).await {
if self.supports_cache_control().await {
payload = update_request_for_anthropic(&payload);
}
@@ -414,7 +422,7 @@ impl Provider for OpenRouterProvider {
.with_retry(|| async {
let resp = self
.api_client
.response_post(session_id, "api/v1/chat/completions", &payload)
.response_post(Some(session_id), "api/v1/chat/completions", &payload)
.await?;
handle_status_openai_compat(resp).await
})
+16 -5
View File
@@ -14,6 +14,7 @@ use super::errors::ProviderError;
use super::retry::ProviderRetry;
use super::utils::RequestLog;
use crate::conversation::message::{Message, MessageContent};
use crate::session_context::SESSION_ID_HEADER;
use crate::model::ModelConfig;
use chrono::Utc;
@@ -154,17 +155,27 @@ impl SageMakerTgiProvider {
Ok(request)
}
async fn invoke_endpoint(&self, payload: Value) -> Result<Value, ProviderError> {
async fn invoke_endpoint(
&self,
session_id: Option<&str>,
payload: Value,
) -> Result<Value, ProviderError> {
let body = serde_json::to_string(&payload).map_err(|e| {
ProviderError::RequestFailed(format!("Failed to serialize request: {}", e))
})?;
let response = self
let mut request = self
.sagemaker_client
.invoke_endpoint()
.endpoint_name(&self.endpoint_name)
.content_type("application/json")
.body(body.into_bytes().into())
.body(body.into_bytes().into());
if let Some(session_id) = session_id.filter(|id| !id.is_empty()) {
request = request.custom_attributes(format!("{SESSION_ID_HEADER}={session_id}"));
}
let response = request
.send()
.await
.map_err(|e| ProviderError::RequestFailed(format!("SageMaker invoke failed: {}", e)))?;
@@ -289,7 +300,7 @@ impl Provider for SageMakerTgiProvider {
)]
async fn complete_with_model(
&self,
_session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -302,7 +313,7 @@ impl Provider for SageMakerTgiProvider {
})?;
let response = self
.with_retry(|| self.invoke_endpoint(request_payload.clone()))
.with_retry(|| self.invoke_endpoint(session_id, request_payload.clone()))
.await?;
let message = self.parse_tgi_response(response)?;
+6 -2
View File
@@ -107,7 +107,11 @@ impl SnowflakeProvider {
})
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, "api/v2/cortex/inference:complete", payload)
@@ -319,7 +323,7 @@ impl Provider for SnowflakeProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
+6 -3
View File
@@ -121,7 +121,7 @@ impl Provider for TestProvider {
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
_model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -130,7 +130,10 @@ impl Provider for TestProvider {
let hash = Self::hash_input(messages);
if let Some(inner) = &self.inner {
let (message, usage) = inner.complete(session_id, system, messages, tools).await?;
let model_config = inner.get_model_config();
let (message, usage) = inner
.complete_with_model(session_id, &model_config, system, messages, tools)
.await?;
let record = TestRecord {
input: TestInput {
@@ -203,7 +206,7 @@ mod tests {
async fn complete_with_model(
&self,
_session_id: &str,
_session_id: Option<&str>,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
+14 -8
View File
@@ -63,7 +63,11 @@ impl TetrateProvider {
})
}
async fn post(&self, session_id: &str, payload: &Value) -> Result<Value, ProviderError> {
async fn post(
&self,
session_id: Option<&str>,
payload: &Value,
) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, "v1/chat/completions", payload)
@@ -158,7 +162,7 @@ impl Provider for TetrateProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -215,7 +219,7 @@ impl Provider for TetrateProvider {
.with_retry(|| async {
let resp = self
.api_client
.response_post(session_id, "v1/chat/completions", &payload)
.response_post(Some(session_id), "v1/chat/completions", &payload)
.await?;
handle_status_openai_compat(resp).await
})
@@ -228,12 +232,14 @@ impl Provider for TetrateProvider {
}
/// Fetch supported models from Tetrate Agent Router Service API (only models with tool support)
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
// Use the existing api_client which already has authentication configured
let response = match self.api_client.response_get(session_id, "v1/models").await {
let response = match self
.api_client
.request(None, "v1/models")
.response_get()
.await
{
Ok(response) => response,
Err(e) => {
tracing::warn!("Failed to fetch models from Tetrate Agent Router Service API: {}, falling back to manual model entry", e);
+5 -7
View File
@@ -115,7 +115,7 @@ impl VeniceProvider {
async fn post(
&self,
session_id: &str,
session_id: Option<&str>,
path: &str,
payload: &Value,
) -> Result<Value, ProviderError> {
@@ -229,13 +229,11 @@ impl Provider for VeniceProvider {
self.model.clone()
}
async fn fetch_supported_models(
&self,
session_id: &str,
) -> Result<Option<Vec<String>>, ProviderError> {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = self
.api_client
.response_get(session_id, &self.models_path)
.request(None, &self.models_path)
.response_get()
.await?;
let json: serde_json::Value = response.json().await?;
@@ -265,7 +263,7 @@ impl Provider for VeniceProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
+3 -3
View File
@@ -69,7 +69,7 @@ impl XaiProvider {
})
}
async fn post(&self, session_id: &str, payload: Value) -> Result<Value, ProviderError> {
async fn post(&self, session_id: Option<&str>, payload: Value) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post(session_id, "chat/completions", &payload)
@@ -110,7 +110,7 @@ impl Provider for XaiProvider {
)]
async fn complete_with_model(
&self,
session_id: &str,
session_id: Option<&str>,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
@@ -165,7 +165,7 @@ impl Provider for XaiProvider {
.with_retry(|| async {
let resp = self
.api_client
.response_post(session_id, "chat/completions", &payload)
.response_post(Some(session_id), "chat/completions", &payload)
.await?;
handle_status_openai_compat(resp).await
})