feat: add streaming support for Claude Code CLI provider (#6833)
Signed-off-by: rabi <ramishra@redhat.com>
This commit is contained in:
@@ -1,16 +1,21 @@
|
|||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
|
use async_stream::try_stream;
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use futures::future::BoxFuture;
|
use futures::future::BoxFuture;
|
||||||
use rmcp::model::Role;
|
use rmcp::model::{Role, Tool};
|
||||||
use serde_json::{json, Value};
|
use serde_json::{json, Value};
|
||||||
use std::io::Write;
|
use std::io::Write;
|
||||||
use std::path::{Path, PathBuf};
|
use std::path::{Path, PathBuf};
|
||||||
use std::process::Stdio;
|
use std::process::Stdio;
|
||||||
|
use std::sync::Arc;
|
||||||
use tempfile::NamedTempFile;
|
use tempfile::NamedTempFile;
|
||||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||||
use tokio::process::Command;
|
use tokio::process::Command;
|
||||||
|
|
||||||
use super::base::{ConfigKey, Provider, ProviderDef, ProviderMetadata, ProviderUsage, Usage};
|
use super::base::{
|
||||||
|
stream_from_single_message, ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata,
|
||||||
|
ProviderUsage, Usage,
|
||||||
|
};
|
||||||
use super::errors::ProviderError;
|
use super::errors::ProviderError;
|
||||||
use super::utils::{filter_extensions_from_system_prompt, RequestLog};
|
use super::utils::{filter_extensions_from_system_prompt, RequestLog};
|
||||||
use crate::config::base::ClaudeCodeCommand;
|
use crate::config::base::ClaudeCodeCommand;
|
||||||
@@ -20,7 +25,8 @@ use crate::config::{Config, ExtensionConfig, GooseMode};
|
|||||||
use crate::conversation::message::{Message, MessageContent};
|
use crate::conversation::message::{Message, MessageContent};
|
||||||
use crate::model::ModelConfig;
|
use crate::model::ModelConfig;
|
||||||
use crate::subprocess::configure_subprocess;
|
use crate::subprocess::configure_subprocess;
|
||||||
use rmcp::model::Tool;
|
|
||||||
|
use super::cli_common::{error_from_event, extract_usage_tokens};
|
||||||
|
|
||||||
const CLAUDE_CODE_PROVIDER_NAME: &str = "claude-code";
|
const CLAUDE_CODE_PROVIDER_NAME: &str = "claude-code";
|
||||||
pub const CLAUDE_CODE_DEFAULT_MODEL: &str = "default";
|
pub const CLAUDE_CODE_DEFAULT_MODEL: &str = "default";
|
||||||
@@ -35,6 +41,7 @@ struct CliProcess {
|
|||||||
current_model: String,
|
current_model: String,
|
||||||
log_model_update: bool,
|
log_model_update: bool,
|
||||||
next_request_id: u64,
|
next_request_id: u64,
|
||||||
|
needs_drain: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl std::fmt::Debug for CliProcess {
|
impl std::fmt::Debug for CliProcess {
|
||||||
@@ -129,10 +136,58 @@ impl CliProcess {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn drain_pending_response(&mut self) {
|
||||||
|
if !self.needs_drain {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
tracing::debug!("Draining cancelled response from CLI process");
|
||||||
|
|
||||||
|
let drain = async {
|
||||||
|
let mut line = String::new();
|
||||||
|
loop {
|
||||||
|
line.clear();
|
||||||
|
match self.reader.read_line(&mut line).await {
|
||||||
|
Ok(0) => break,
|
||||||
|
Ok(_) => {
|
||||||
|
let trimmed = line.trim();
|
||||||
|
if trimmed.is_empty() {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if let Ok(parsed) = serde_json::from_str::<Value>(trimmed) {
|
||||||
|
match parsed.get("type").and_then(|t| t.as_str()) {
|
||||||
|
Some("result") | Some("error") => break,
|
||||||
|
_ => continue,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
tracing::trace!(line = trimmed, "Non-JSON line during drain");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(_) => break,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const DRAIN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
|
||||||
|
if tokio::time::timeout(DRAIN_TIMEOUT, drain).await.is_err() {
|
||||||
|
// CLI is still producing the old response. Leave needs_drain
|
||||||
|
// true so the next call retries — by then the old response
|
||||||
|
// likely completed and drain will succeed quickly.
|
||||||
|
tracing::warn!(
|
||||||
|
"Drain did not complete in {DRAIN_TIMEOUT:?}; \
|
||||||
|
will retry on next request"
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
self.needs_drain = false;
|
||||||
|
tracing::debug!("Drain complete, protocol re-synced");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Drop for CliProcess {
|
impl Drop for CliProcess {
|
||||||
fn drop(&mut self) {
|
fn drop(&mut self) {
|
||||||
|
self.stderr_handle.abort();
|
||||||
let _ = self.child.start_kill();
|
let _ = self.child.start_kill();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -152,7 +207,7 @@ pub struct ClaudeCodeProvider {
|
|||||||
#[serde(skip)]
|
#[serde(skip)]
|
||||||
mcp_config_file: Option<NamedTempFile>,
|
mcp_config_file: Option<NamedTempFile>,
|
||||||
#[serde(skip)]
|
#[serde(skip)]
|
||||||
cli_process: tokio::sync::OnceCell<tokio::sync::Mutex<CliProcess>>,
|
cli_process: tokio::sync::OnceCell<Arc<tokio::sync::Mutex<CliProcess>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ClaudeCodeProvider {
|
impl ClaudeCodeProvider {
|
||||||
@@ -245,14 +300,11 @@ impl ClaudeCodeProvider {
|
|||||||
.to_string(),
|
.to_string(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
GooseMode::Chat => {
|
GooseMode::Chat => {}
|
||||||
// Chat mode doesn't need permission flags
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Parse NDJSON stream-json response from Claude CLI
|
|
||||||
fn parse_claude_response(
|
fn parse_claude_response(
|
||||||
&self,
|
&self,
|
||||||
json_lines: &[String],
|
json_lines: &[String],
|
||||||
@@ -265,7 +317,6 @@ impl ClaudeCodeProvider {
|
|||||||
match parsed.get("type").and_then(|t| t.as_str()) {
|
match parsed.get("type").and_then(|t| t.as_str()) {
|
||||||
Some("assistant") => {
|
Some("assistant") => {
|
||||||
if let Some(message) = parsed.get("message") {
|
if let Some(message) = parsed.get("message") {
|
||||||
// Extract text content from this assistant message
|
|
||||||
if let Some(content) = message.get("content").and_then(|c| c.as_array())
|
if let Some(content) = message.get("content").and_then(|c| c.as_array())
|
||||||
{
|
{
|
||||||
for item in content {
|
for item in content {
|
||||||
@@ -276,65 +327,33 @@ impl ClaudeCodeProvider {
|
|||||||
all_text_content.push(text.to_string());
|
all_text_content.push(text.to_string());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Skip tool_use - those are claude CLI's internal tools
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Extract usage information
|
|
||||||
if let Some(usage_info) = message.get("usage") {
|
if let Some(usage_info) = message.get("usage") {
|
||||||
usage.input_tokens = usage_info
|
usage = extract_usage_tokens(usage_info);
|
||||||
.get("input_tokens")
|
|
||||||
.and_then(|v| v.as_i64())
|
|
||||||
.map(|v| v as i32);
|
|
||||||
usage.output_tokens = usage_info
|
|
||||||
.get("output_tokens")
|
|
||||||
.and_then(|v| v.as_i64())
|
|
||||||
.map(|v| v as i32);
|
|
||||||
|
|
||||||
if usage.total_tokens.is_none() {
|
|
||||||
if let (Some(input), Some(output)) =
|
|
||||||
(usage.input_tokens, usage.output_tokens)
|
|
||||||
{
|
|
||||||
usage.total_tokens = Some(input + output);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some("result") => {
|
Some("result") => {
|
||||||
// Extract additional usage info from result if available
|
|
||||||
if let Some(result_usage) = parsed.get("usage") {
|
if let Some(result_usage) = parsed.get("usage") {
|
||||||
if usage.input_tokens.is_none() {
|
let new = extract_usage_tokens(result_usage);
|
||||||
usage.input_tokens = result_usage
|
usage = Usage::new(
|
||||||
.get("input_tokens")
|
usage.input_tokens.or(new.input_tokens),
|
||||||
.and_then(|v| v.as_i64())
|
usage.output_tokens.or(new.output_tokens),
|
||||||
.map(|v| v as i32);
|
None,
|
||||||
}
|
);
|
||||||
if usage.output_tokens.is_none() {
|
|
||||||
usage.output_tokens = result_usage
|
|
||||||
.get("output_tokens")
|
|
||||||
.and_then(|v| v.as_i64())
|
|
||||||
.map(|v| v as i32);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some("error") => {
|
Some("error") => {
|
||||||
let error_msg = parsed
|
return Err(error_from_event("Claude CLI", &parsed));
|
||||||
.get("error")
|
|
||||||
.and_then(|e| e.as_str())
|
|
||||||
.unwrap_or("Unknown error");
|
|
||||||
return Err(ProviderError::RequestFailed(format!(
|
|
||||||
"Claude CLI error: {}",
|
|
||||||
error_msg
|
|
||||||
)));
|
|
||||||
}
|
}
|
||||||
Some("system") => {} // Ignore system init events
|
Some("system") => {}
|
||||||
_ => {} // Ignore other event types
|
_ => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Combine all text content into a single message
|
|
||||||
let combined_text = all_text_content.join("\n\n");
|
let combined_text = all_text_content.join("\n\n");
|
||||||
if combined_text.is_empty() {
|
if combined_text.is_empty() {
|
||||||
return Err(ProviderError::RequestFailed(
|
return Err(ProviderError::RequestFailed(
|
||||||
@@ -353,6 +372,73 @@ impl ClaudeCodeProvider {
|
|||||||
Ok((response_message, usage))
|
Ok((response_message, usage))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn spawn_process(&self, filtered_system: &str) -> Result<CliProcess, ProviderError> {
|
||||||
|
let mut cmd = self.build_stream_json_command();
|
||||||
|
|
||||||
|
if let Some(f) = &self.mcp_config_file {
|
||||||
|
cmd.arg("--mcp-config").arg(f.path());
|
||||||
|
cmd.arg("--strict-mcp-config");
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd.arg("--include-partial-messages")
|
||||||
|
.arg("--system-prompt")
|
||||||
|
.arg(filtered_system)
|
||||||
|
.arg("--model")
|
||||||
|
.arg(&self.model.model_name);
|
||||||
|
|
||||||
|
Self::apply_permission_flags(&mut cmd)?;
|
||||||
|
|
||||||
|
let mut child = cmd.spawn().map_err(|e| {
|
||||||
|
ProviderError::RequestFailed(format!(
|
||||||
|
"Failed to spawn Claude CLI command '{:?}': {}.",
|
||||||
|
self.command, e
|
||||||
|
))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let stdin = child
|
||||||
|
.stdin
|
||||||
|
.take()
|
||||||
|
.ok_or_else(|| ProviderError::RequestFailed("Failed to capture stdin".to_string()))?;
|
||||||
|
let stdout = child
|
||||||
|
.stdout
|
||||||
|
.take()
|
||||||
|
.ok_or_else(|| ProviderError::RequestFailed("Failed to capture stdout".to_string()))?;
|
||||||
|
|
||||||
|
let stderr = child.stderr.take();
|
||||||
|
let stderr_handle = tokio::spawn(async move {
|
||||||
|
let mut output = String::new();
|
||||||
|
if let Some(mut stderr) = stderr {
|
||||||
|
use tokio::io::AsyncReadExt;
|
||||||
|
let _ = stderr.read_to_string(&mut output).await;
|
||||||
|
}
|
||||||
|
output
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(CliProcess {
|
||||||
|
child,
|
||||||
|
stdin: Box::new(stdin),
|
||||||
|
reader: BufReader::new(Box::new(stdout)),
|
||||||
|
stderr_handle,
|
||||||
|
current_model: self.model.model_name.clone(),
|
||||||
|
log_model_update: false,
|
||||||
|
next_request_id: 0,
|
||||||
|
needs_drain: false,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn get_or_init_process(
|
||||||
|
&self,
|
||||||
|
filtered_system: &str,
|
||||||
|
) -> Result<&Arc<tokio::sync::Mutex<CliProcess>>, ProviderError> {
|
||||||
|
self.cli_process
|
||||||
|
.get_or_try_init(|| async {
|
||||||
|
Ok(Arc::new(tokio::sync::Mutex::new(
|
||||||
|
self.spawn_process(filtered_system)?,
|
||||||
|
)))
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
|
||||||
async fn execute_command(
|
async fn execute_command(
|
||||||
&self,
|
&self,
|
||||||
system: &str,
|
system: &str,
|
||||||
@@ -363,74 +449,19 @@ impl ClaudeCodeProvider {
|
|||||||
) -> Result<Vec<String>, ProviderError> {
|
) -> Result<Vec<String>, ProviderError> {
|
||||||
let filtered_system = filter_extensions_from_system_prompt(system);
|
let filtered_system = filter_extensions_from_system_prompt(system);
|
||||||
|
|
||||||
if std::env::var("GOOSE_CLAUDE_CODE_DEBUG").is_ok() {
|
tracing::debug!(
|
||||||
println!("=== CLAUDE CODE PROVIDER DEBUG ===");
|
command = ?self.command,
|
||||||
println!("Command: {:?}", self.command);
|
system_prompt_len = system.len(),
|
||||||
println!("Original system prompt length: {} chars", system.len());
|
filtered_system_prompt_len = filtered_system.len(),
|
||||||
println!(
|
"Executing Claude CLI command"
|
||||||
"Filtered system prompt length: {} chars",
|
);
|
||||||
filtered_system.len()
|
|
||||||
);
|
|
||||||
println!("Filtered system prompt: {}", filtered_system);
|
|
||||||
println!("================================");
|
|
||||||
}
|
|
||||||
|
|
||||||
// Spawn lazily on first call (OnceCell ensures exactly once)
|
|
||||||
let process_mutex = self
|
|
||||||
.cli_process
|
|
||||||
.get_or_try_init(|| async {
|
|
||||||
let mut cmd = self.build_stream_json_command();
|
|
||||||
if let Some(f) = &self.mcp_config_file {
|
|
||||||
cmd.arg("--mcp-config").arg(f.path());
|
|
||||||
cmd.arg("--strict-mcp-config");
|
|
||||||
}
|
|
||||||
// System prompt is set once at process start and cannot be updated at runtime.
|
|
||||||
cmd.arg("--system-prompt").arg(&filtered_system);
|
|
||||||
|
|
||||||
// The initial model can be updated later.
|
|
||||||
cmd.arg("--model").arg(&self.model.model_name);
|
|
||||||
|
|
||||||
Self::apply_permission_flags(&mut cmd)?;
|
|
||||||
|
|
||||||
let mut child = cmd.spawn().map_err(|e| {
|
|
||||||
ProviderError::RequestFailed(format!(
|
|
||||||
"Failed to spawn Claude CLI command '{:?}': {}.",
|
|
||||||
self.command, e
|
|
||||||
))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let stdin = child.stdin.take().ok_or_else(|| {
|
|
||||||
ProviderError::RequestFailed("Failed to capture stdin".to_string())
|
|
||||||
})?;
|
|
||||||
let stdout = child.stdout.take().ok_or_else(|| {
|
|
||||||
ProviderError::RequestFailed("Failed to capture stdout".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
// Drain stderr concurrently to prevent pipe buffer deadlock
|
|
||||||
let stderr = child.stderr.take();
|
|
||||||
let stderr_handle = tokio::spawn(async move {
|
|
||||||
let mut output = String::new();
|
|
||||||
if let Some(mut stderr) = stderr {
|
|
||||||
use tokio::io::AsyncReadExt;
|
|
||||||
let _ = stderr.read_to_string(&mut output).await;
|
|
||||||
}
|
|
||||||
output
|
|
||||||
});
|
|
||||||
|
|
||||||
Ok::<_, ProviderError>(tokio::sync::Mutex::new(CliProcess {
|
|
||||||
child,
|
|
||||||
stdin: Box::new(stdin),
|
|
||||||
reader: BufReader::new(Box::new(stdout)),
|
|
||||||
stderr_handle,
|
|
||||||
current_model: String::new(),
|
|
||||||
log_model_update: false,
|
|
||||||
next_request_id: 0,
|
|
||||||
}))
|
|
||||||
})
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
|
let process_mutex = self.get_or_init_process(&filtered_system).await?;
|
||||||
let mut process = process_mutex.lock().await;
|
let mut process = process_mutex.lock().await;
|
||||||
|
|
||||||
|
// Drain any pending response from a cancelled stream
|
||||||
|
process.drain_pending_response().await;
|
||||||
|
|
||||||
// Switch model if it differs from what the CLI is currently using.
|
// Switch model if it differs from what the CLI is currently using.
|
||||||
process.send_set_model(model).await?;
|
process.send_set_model(model).await?;
|
||||||
|
|
||||||
@@ -449,7 +480,7 @@ impl ClaudeCodeProvider {
|
|||||||
ProviderError::RequestFailed(format!("Failed to write newline to stdin: {}", e))
|
ProviderError::RequestFailed(format!("Failed to write newline to stdin: {}", e))
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
// Read lines until we see a "result" event
|
// Read lines until we see a "result" or "error" event
|
||||||
let mut lines = Vec::new();
|
let mut lines = Vec::new();
|
||||||
let mut line = String::new();
|
let mut line = String::new();
|
||||||
|
|
||||||
@@ -457,7 +488,6 @@ impl ClaudeCodeProvider {
|
|||||||
line.clear();
|
line.clear();
|
||||||
match process.reader.read_line(&mut line).await {
|
match process.reader.read_line(&mut line).await {
|
||||||
Ok(0) => {
|
Ok(0) => {
|
||||||
// EOF means the process died
|
|
||||||
return Err(ProviderError::RequestFailed(
|
return Err(ProviderError::RequestFailed(
|
||||||
"Claude CLI process terminated unexpectedly".to_string(),
|
"Claude CLI process terminated unexpectedly".to_string(),
|
||||||
));
|
));
|
||||||
@@ -467,31 +497,31 @@ impl ClaudeCodeProvider {
|
|||||||
if trimmed.is_empty() {
|
if trimmed.is_empty() {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
lines.push(trimmed.to_string());
|
|
||||||
|
|
||||||
if let Ok(parsed) = serde_json::from_str::<Value>(trimmed) {
|
if let Ok(parsed) = serde_json::from_str::<Value>(trimmed) {
|
||||||
match parsed.get("type").and_then(|t| t.as_str()) {
|
match parsed.get("type").and_then(|t| t.as_str()) {
|
||||||
Some("result") => break,
|
Some("stream_event") => continue,
|
||||||
Some("error") => break,
|
Some("result") | Some("error") => {
|
||||||
|
lines.push(trimmed.to_string());
|
||||||
|
break;
|
||||||
|
}
|
||||||
// The system init with the resolved model arrives here,
|
// The system init with the resolved model arrives here,
|
||||||
// not in send_set_model (which only sees control_response):
|
// not in send_set_model (which only sees control_response).
|
||||||
// send_set_model: {"type":"control_response",...}
|
|
||||||
// execute_command: {"type":"system",...,"model":"claude-sonnet-4-5-20250929",...}
|
|
||||||
Some("system") if process.log_model_update => {
|
Some("system") if process.log_model_update => {
|
||||||
if let Some(resolved) = parsed.get("model").and_then(|m| m.as_str())
|
if let Some(resolved) = parsed.get("model").and_then(|m| m.as_str())
|
||||||
{
|
{
|
||||||
if std::env::var("GOOSE_CLAUDE_CODE_DEBUG").is_ok() {
|
tracing::debug!(
|
||||||
println!(
|
from = %process.current_model,
|
||||||
"set_model: {} resolved to {}",
|
to = %resolved,
|
||||||
process.current_model, resolved
|
"set_model resolved"
|
||||||
);
|
);
|
||||||
}
|
|
||||||
}
|
}
|
||||||
process.log_model_update = false;
|
process.log_model_update = false;
|
||||||
}
|
}
|
||||||
_ => {}
|
_ => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
lines.push(trimmed.to_string());
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
return Err(ProviderError::RequestFailed(format!(
|
return Err(ProviderError::RequestFailed(format!(
|
||||||
@@ -503,57 +533,8 @@ impl ClaudeCodeProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
tracing::debug!("Command executed successfully, got {} lines", lines.len());
|
tracing::debug!("Command executed successfully, got {} lines", lines.len());
|
||||||
for (i, line) in lines.iter().enumerate() {
|
|
||||||
tracing::debug!("Line {}: {}", i, line);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(lines)
|
Ok(lines)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Generate a simple session description without calling subprocess
|
|
||||||
fn generate_simple_session_description(
|
|
||||||
&self,
|
|
||||||
messages: &[Message],
|
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
|
||||||
// Extract the first user message text
|
|
||||||
let description = messages
|
|
||||||
.iter()
|
|
||||||
.find(|m| m.role == Role::User)
|
|
||||||
.and_then(|m| {
|
|
||||||
m.content.iter().find_map(|c| match c {
|
|
||||||
MessageContent::Text(text_content) => Some(&text_content.text),
|
|
||||||
_ => None,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.map(|text| {
|
|
||||||
// Take first few words, limit to 4 words
|
|
||||||
text.split_whitespace()
|
|
||||||
.take(4)
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
.join(" ")
|
|
||||||
})
|
|
||||||
.unwrap_or_else(|| "Simple task".to_string());
|
|
||||||
|
|
||||||
if std::env::var("GOOSE_CLAUDE_CODE_DEBUG").is_ok() {
|
|
||||||
println!("=== CLAUDE CODE PROVIDER DEBUG ===");
|
|
||||||
println!("Generated simple session description: {}", description);
|
|
||||||
println!("Skipped subprocess call for session description");
|
|
||||||
println!("================================");
|
|
||||||
}
|
|
||||||
|
|
||||||
let message = Message::new(
|
|
||||||
Role::Assistant,
|
|
||||||
chrono::Utc::now().timestamp(),
|
|
||||||
vec![MessageContent::text(description.clone())],
|
|
||||||
);
|
|
||||||
|
|
||||||
let usage = Usage::default();
|
|
||||||
|
|
||||||
Ok((
|
|
||||||
message,
|
|
||||||
ProviderUsage::new(self.model.model_name.clone(), usage),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Extract model aliases from the CLI's initialize control_response.
|
/// Extract model aliases from the CLI's initialize control_response.
|
||||||
@@ -708,7 +689,6 @@ impl Provider for ClaudeCodeProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn get_model_config(&self) -> ModelConfig {
|
fn get_model_config(&self) -> ModelConfig {
|
||||||
// Return the model config with appropriate context limit for Claude models
|
|
||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -785,9 +765,11 @@ impl Provider for ClaudeCodeProvider {
|
|||||||
messages: &[Message],
|
messages: &[Message],
|
||||||
tools: &[Tool],
|
tools: &[Tool],
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||||
// Check if this is a session description request (short system prompt asking for 4 words or less)
|
if super::cli_common::is_session_description_request(system) {
|
||||||
if system.contains("four words or less") || system.contains("4 words or less") {
|
return super::cli_common::generate_simple_session_description(
|
||||||
return self.generate_simple_session_description(messages);
|
&model_config.model_name,
|
||||||
|
messages,
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
// session_id is None before a session is created (e.g. model listing).
|
// session_id is None before a session is created (e.g. model listing).
|
||||||
@@ -798,7 +780,6 @@ impl Provider for ClaudeCodeProvider {
|
|||||||
|
|
||||||
let (message, usage) = self.parse_claude_response(&json_lines)?;
|
let (message, usage) = self.parse_claude_response(&json_lines)?;
|
||||||
|
|
||||||
// Create a dummy payload for debug tracing
|
|
||||||
let payload = json!({
|
let payload = json!({
|
||||||
"command": self.command,
|
"command": self.command,
|
||||||
"model": model_config.model_name,
|
"model": model_config.model_name,
|
||||||
@@ -819,6 +800,172 @@ impl Provider for ClaudeCodeProvider {
|
|||||||
ProviderUsage::new(model_config.model_name.clone(), usage),
|
ProviderUsage::new(model_config.model_name.clone(), usage),
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn supports_streaming(&self) -> bool {
|
||||||
|
true
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn stream(
|
||||||
|
&self,
|
||||||
|
session_id: &str,
|
||||||
|
system: &str,
|
||||||
|
messages: &[Message],
|
||||||
|
_tools: &[Tool],
|
||||||
|
) -> Result<MessageStream, ProviderError> {
|
||||||
|
if super::cli_common::is_session_description_request(system) {
|
||||||
|
let (message, usage) = super::cli_common::generate_simple_session_description(
|
||||||
|
&self.model.model_name,
|
||||||
|
messages,
|
||||||
|
)?;
|
||||||
|
return Ok(stream_from_single_message(message, usage));
|
||||||
|
}
|
||||||
|
|
||||||
|
let filtered_system = filter_extensions_from_system_prompt(system);
|
||||||
|
let process_arc = Arc::clone(self.get_or_init_process(&filtered_system).await?);
|
||||||
|
|
||||||
|
// Prepare the payload outside the lock — these don't need the process.
|
||||||
|
let blocks = self.last_user_content_blocks(messages);
|
||||||
|
let ndjson_line = build_stream_json_input(&blocks, session_id);
|
||||||
|
let model_name = self.model.model_name.clone();
|
||||||
|
let message_id = uuid::Uuid::new_v4().to_string();
|
||||||
|
|
||||||
|
Ok(Box::pin(try_stream! {
|
||||||
|
// Single lock acquisition covers write-to-stdin and read-from-stdout,
|
||||||
|
// eliminating the race window between the two.
|
||||||
|
let mut process = process_arc.lock_owned().await;
|
||||||
|
process.drain_pending_response().await;
|
||||||
|
process.send_set_model(&model_name).await?;
|
||||||
|
|
||||||
|
process
|
||||||
|
.stdin
|
||||||
|
.write_all(ndjson_line.as_bytes())
|
||||||
|
.await
|
||||||
|
.map_err(|e| {
|
||||||
|
ProviderError::RequestFailed(format!("Failed to write to stdin: {}", e))
|
||||||
|
})?;
|
||||||
|
process.stdin.write_all(b"\n").await.map_err(|e| {
|
||||||
|
ProviderError::RequestFailed(format!("Failed to write newline to stdin: {}", e))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
process.needs_drain = true;
|
||||||
|
let mut line = String::new();
|
||||||
|
let mut accumulated_usage = Usage::default();
|
||||||
|
let mut stream_error: Option<ProviderError> = None;
|
||||||
|
let stream_timestamp = chrono::Utc::now().timestamp();
|
||||||
|
|
||||||
|
loop {
|
||||||
|
line.clear();
|
||||||
|
match process.reader.read_line(&mut line).await {
|
||||||
|
Ok(0) => {
|
||||||
|
process.needs_drain = false;
|
||||||
|
stream_error = Some(ProviderError::RequestFailed(
|
||||||
|
"Claude CLI process terminated unexpectedly".to_string(),
|
||||||
|
));
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
Ok(_) => {
|
||||||
|
let trimmed = line.trim();
|
||||||
|
if trimmed.is_empty() {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Ok(parsed) = serde_json::from_str::<Value>(trimmed) {
|
||||||
|
match parsed.get("type").and_then(|t| t.as_str()) {
|
||||||
|
Some("stream_event") => {
|
||||||
|
if let Some(event) = parsed.get("event") {
|
||||||
|
match event.get("type").and_then(|t| t.as_str()) {
|
||||||
|
Some("content_block_delta") => {
|
||||||
|
if let Some(text) = event
|
||||||
|
.get("delta")
|
||||||
|
.filter(|d| {
|
||||||
|
d.get("type").and_then(|t| t.as_str())
|
||||||
|
== Some("text_delta")
|
||||||
|
})
|
||||||
|
.and_then(|d| d.get("text"))
|
||||||
|
.and_then(|t| t.as_str())
|
||||||
|
{
|
||||||
|
let mut partial_message = Message::new(
|
||||||
|
Role::Assistant,
|
||||||
|
stream_timestamp,
|
||||||
|
vec![MessageContent::text(text)],
|
||||||
|
);
|
||||||
|
partial_message.id =
|
||||||
|
Some(message_id.clone());
|
||||||
|
yield (Some(partial_message), None);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Some("message_start") => {
|
||||||
|
if let Some(usage_info) = event
|
||||||
|
.get("message")
|
||||||
|
.and_then(|m| m.get("usage"))
|
||||||
|
{
|
||||||
|
let new = extract_usage_tokens(usage_info);
|
||||||
|
if let Some(i) = new.input_tokens {
|
||||||
|
accumulated_usage.input_tokens = Some(i);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Some("message_delta") => {
|
||||||
|
if let Some(usage_info) = event.get("usage") {
|
||||||
|
let new = extract_usage_tokens(usage_info);
|
||||||
|
if let Some(o) = new.output_tokens {
|
||||||
|
accumulated_usage.output_tokens = Some(o);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Some("result") => {
|
||||||
|
process.needs_drain = false;
|
||||||
|
if let Some(usage_info) = parsed.get("usage") {
|
||||||
|
let new = extract_usage_tokens(usage_info);
|
||||||
|
accumulated_usage = Usage::new(
|
||||||
|
new.input_tokens.or(accumulated_usage.input_tokens),
|
||||||
|
new.output_tokens.or(accumulated_usage.output_tokens),
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
Some("error") => {
|
||||||
|
process.needs_drain = false;
|
||||||
|
stream_error = Some(error_from_event("Claude CLI", &parsed));
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
Some("system") if process.log_model_update => {
|
||||||
|
if let Some(resolved) = parsed.get("model").and_then(|m| m.as_str()) {
|
||||||
|
tracing::debug!(
|
||||||
|
from = %process.current_model,
|
||||||
|
to = %resolved,
|
||||||
|
"set_model resolved"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
process.log_model_update = false;
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
process.needs_drain = false;
|
||||||
|
stream_error = Some(ProviderError::RequestFailed(format!(
|
||||||
|
"Failed to read streaming output: {e}"
|
||||||
|
)));
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(err) = stream_error {
|
||||||
|
Err(err)?;
|
||||||
|
}
|
||||||
|
|
||||||
|
let provider_usage = ProviderUsage::new(model_name, accumulated_usage);
|
||||||
|
yield (None, Some(provider_usage));
|
||||||
|
}))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -833,6 +980,44 @@ mod tests {
|
|||||||
use tempfile::tempdir;
|
use tempfile::tempdir;
|
||||||
use test_case::test_case;
|
use test_case::test_case;
|
||||||
|
|
||||||
|
#[test_case(
|
||||||
|
json!({"input_tokens": 100, "output_tokens": 50}),
|
||||||
|
Some(100), Some(50)
|
||||||
|
; "both_tokens"
|
||||||
|
)]
|
||||||
|
#[test_case(json!({"input_tokens": 100}), Some(100), None ; "input_only")]
|
||||||
|
#[test_case(json!({}), None, None ; "empty_usage")]
|
||||||
|
fn test_extract_usage_tokens(
|
||||||
|
usage_json: Value,
|
||||||
|
expected_input: Option<i32>,
|
||||||
|
expected_output: Option<i32>,
|
||||||
|
) {
|
||||||
|
let usage = extract_usage_tokens(&usage_json);
|
||||||
|
assert_eq!(usage.input_tokens, expected_input);
|
||||||
|
assert_eq!(usage.output_tokens, expected_output);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test_case(
|
||||||
|
r#"{"type":"error","error":"context window exceeded"}"#,
|
||||||
|
true
|
||||||
|
; "context_exceeded"
|
||||||
|
)]
|
||||||
|
#[test_case(
|
||||||
|
r#"{"type":"error","error":"Model not supported"}"#,
|
||||||
|
false
|
||||||
|
; "generic_error_from_event"
|
||||||
|
)]
|
||||||
|
#[test_case(r#"{"type":"error"}"#, false ; "missing_error_field")]
|
||||||
|
fn test_error_from_event(line: &str, is_context_exceeded: bool) {
|
||||||
|
let parsed: Value = serde_json::from_str(line).unwrap();
|
||||||
|
let err = error_from_event("Claude CLI", &parsed);
|
||||||
|
if is_context_exceeded {
|
||||||
|
assert!(matches!(err, ProviderError::ContextLengthExceeded(_)));
|
||||||
|
} else {
|
||||||
|
assert!(matches!(err, ProviderError::RequestFailed(_)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// (role, text, optional (image_data, mime_type))
|
/// (role, text, optional (image_data, mime_type))
|
||||||
type MsgSpec<'a> = (&'a str, &'a str, Option<(&'a str, &'a str)>);
|
type MsgSpec<'a> = (&'a str, &'a str, Option<(&'a str, &'a str)>);
|
||||||
|
|
||||||
@@ -958,6 +1143,17 @@ mod tests {
|
|||||||
Some(3), Some(3)
|
Some(3), Some(3)
|
||||||
; "system_init_filtered"
|
; "system_init_filtered"
|
||||||
)]
|
)]
|
||||||
|
#[test_case(
|
||||||
|
&[
|
||||||
|
r#"{"type":"stream_event","event":{"type":"content_block_delta","delta":{"type":"text_delta","text":"He"}}}"#,
|
||||||
|
r#"{"type":"stream_event","event":{"type":"content_block_delta","delta":{"type":"text_delta","text":"llo"}}}"#,
|
||||||
|
r#"{"type":"assistant","message":{"role":"assistant","content":[{"type":"text","text":"Hello"}],"usage":{"input_tokens":50,"output_tokens":10}}}"#,
|
||||||
|
r#"{"type":"result","subtype":"success","result":"Hello","session_id":"abc"}"#,
|
||||||
|
],
|
||||||
|
"Hello",
|
||||||
|
Some(50), Some(10)
|
||||||
|
; "streaming_events_ignored_by_parse"
|
||||||
|
)]
|
||||||
fn test_parse_claude_response_ok(
|
fn test_parse_claude_response_ok(
|
||||||
lines: &[&str],
|
lines: &[&str],
|
||||||
expected_text: &str,
|
expected_text: &str,
|
||||||
@@ -984,7 +1180,7 @@ mod tests {
|
|||||||
)]
|
)]
|
||||||
#[test_case(
|
#[test_case(
|
||||||
&[r#"{"type":"error","error":"context window exceeded"}"#],
|
&[r#"{"type":"error","error":"context window exceeded"}"#],
|
||||||
ProviderError::RequestFailed("Claude CLI error: context window exceeded".into())
|
ProviderError::ContextLengthExceeded("context window exceeded".into())
|
||||||
; "context_length"
|
; "context_length"
|
||||||
)]
|
)]
|
||||||
#[test_case(
|
#[test_case(
|
||||||
@@ -1160,6 +1356,7 @@ mod tests {
|
|||||||
current_model: String::new(),
|
current_model: String::new(),
|
||||||
log_model_update: false,
|
log_model_update: false,
|
||||||
next_request_id: 0,
|
next_request_id: 0,
|
||||||
|
needs_drain: false,
|
||||||
};
|
};
|
||||||
(process, stdin_reader)
|
(process, stdin_reader)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,75 @@
|
|||||||
|
use serde_json::Value;
|
||||||
|
|
||||||
|
use super::base::{ProviderUsage, Usage};
|
||||||
|
use super::errors::ProviderError;
|
||||||
|
use crate::conversation::message::{Message, MessageContent};
|
||||||
|
use rmcp::model::Role;
|
||||||
|
|
||||||
|
pub(crate) fn extract_usage_tokens(usage_info: &Value) -> Usage {
|
||||||
|
let get = |key: &str| {
|
||||||
|
usage_info
|
||||||
|
.get(key)
|
||||||
|
.and_then(|v| v.as_i64())
|
||||||
|
.and_then(|v| i32::try_from(v).ok())
|
||||||
|
};
|
||||||
|
Usage::new(
|
||||||
|
get("input_tokens"),
|
||||||
|
get("output_tokens"),
|
||||||
|
get("total_tokens"),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn error_from_event(provider_name: &str, parsed: &Value) -> ProviderError {
|
||||||
|
let error_msg = parsed
|
||||||
|
.get("error")
|
||||||
|
.and_then(|e| e.as_str())
|
||||||
|
.or_else(|| parsed.get("message").and_then(|m| m.as_str()))
|
||||||
|
.unwrap_or("Unknown error");
|
||||||
|
if error_msg.contains("context window exceeded") {
|
||||||
|
ProviderError::ContextLengthExceeded(error_msg.to_string())
|
||||||
|
} else {
|
||||||
|
ProviderError::RequestFailed(format!("{provider_name} error: {error_msg}"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn is_session_description_request(system: &str) -> bool {
|
||||||
|
system.contains("four words or less") || system.contains("4 words or less")
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn generate_simple_session_description(
|
||||||
|
model_name: &str,
|
||||||
|
messages: &[Message],
|
||||||
|
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||||
|
let description = messages
|
||||||
|
.iter()
|
||||||
|
.find(|m| m.role == Role::User)
|
||||||
|
.and_then(|m| {
|
||||||
|
m.content.iter().find_map(|c| match c {
|
||||||
|
MessageContent::Text(text_content) => Some(&text_content.text),
|
||||||
|
_ => None,
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.map(|text| {
|
||||||
|
text.split_whitespace()
|
||||||
|
.take(4)
|
||||||
|
.collect::<Vec<_>>()
|
||||||
|
.join(" ")
|
||||||
|
})
|
||||||
|
.unwrap_or_else(|| "Simple task".to_string());
|
||||||
|
|
||||||
|
tracing::debug!(
|
||||||
|
description = %description,
|
||||||
|
"Generated simple session description, skipped subprocess"
|
||||||
|
);
|
||||||
|
|
||||||
|
let message = Message::new(
|
||||||
|
Role::Assistant,
|
||||||
|
chrono::Utc::now().timestamp(),
|
||||||
|
vec![MessageContent::text(description)],
|
||||||
|
);
|
||||||
|
|
||||||
|
Ok((
|
||||||
|
message,
|
||||||
|
ProviderUsage::new(model_name.to_string(), Usage::default()),
|
||||||
|
))
|
||||||
|
}
|
||||||
@@ -409,51 +409,6 @@ impl CodexProvider {
|
|||||||
|
|
||||||
Ok((message, usage))
|
Ok((message, usage))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Generate a simple session description without calling subprocess
|
|
||||||
fn generate_simple_session_description(
|
|
||||||
&self,
|
|
||||||
messages: &[Message],
|
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
|
||||||
// Extract the first user message text
|
|
||||||
let description = messages
|
|
||||||
.iter()
|
|
||||||
.find(|m| m.role == Role::User)
|
|
||||||
.and_then(|m| {
|
|
||||||
m.content.iter().find_map(|c| match c {
|
|
||||||
MessageContent::Text(text_content) => Some(&text_content.text),
|
|
||||||
_ => None,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.map(|text| {
|
|
||||||
// Take first few words, limit to 4 words
|
|
||||||
text.split_whitespace()
|
|
||||||
.take(4)
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
.join(" ")
|
|
||||||
})
|
|
||||||
.unwrap_or_else(|| "Simple task".to_string());
|
|
||||||
|
|
||||||
if std::env::var("GOOSE_CODEX_DEBUG").is_ok() {
|
|
||||||
println!("=== CODEX PROVIDER DEBUG ===");
|
|
||||||
println!("Generated simple session description: {}", description);
|
|
||||||
println!("Skipped subprocess call for session description");
|
|
||||||
println!("============================");
|
|
||||||
}
|
|
||||||
|
|
||||||
let message = Message::new(
|
|
||||||
Role::Assistant,
|
|
||||||
chrono::Utc::now().timestamp(),
|
|
||||||
vec![MessageContent::text(description.clone())],
|
|
||||||
);
|
|
||||||
|
|
||||||
let usage = Usage::default();
|
|
||||||
|
|
||||||
Ok((
|
|
||||||
message,
|
|
||||||
ProviderUsage::new(self.model.model_name.clone(), usage),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Builds the text prompt and extracts images to temp files in a single pass.
|
/// Builds the text prompt and extracts images to temp files in a single pass.
|
||||||
@@ -724,9 +679,11 @@ impl Provider for CodexProvider {
|
|||||||
messages: &[Message],
|
messages: &[Message],
|
||||||
tools: &[Tool],
|
tools: &[Tool],
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||||
// Check if this is a session description request
|
if super::cli_common::is_session_description_request(system) {
|
||||||
if system.contains("four words or less") || system.contains("4 words or less") {
|
return super::cli_common::generate_simple_session_description(
|
||||||
return self.generate_simple_session_description(messages);
|
&model_config.model_name,
|
||||||
|
messages,
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let lines = self.execute_command(system, messages, tools).await?;
|
let lines = self.execute_command(system, messages, tools).await?;
|
||||||
@@ -1192,15 +1149,6 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_session_description_generation() {
|
fn test_session_description_generation() {
|
||||||
let provider = CodexProvider {
|
|
||||||
command: PathBuf::from("codex"),
|
|
||||||
model: ModelConfig::new("gpt-5.2-codex").unwrap(),
|
|
||||||
name: "codex".to_string(),
|
|
||||||
reasoning_effort: "high".to_string(),
|
|
||||||
skip_git_check: false,
|
|
||||||
mcp_config_overrides: Vec::new(),
|
|
||||||
};
|
|
||||||
|
|
||||||
let messages = vec![Message::new(
|
let messages = vec![Message::new(
|
||||||
Role::User,
|
Role::User,
|
||||||
chrono::Utc::now().timestamp(),
|
chrono::Utc::now().timestamp(),
|
||||||
@@ -1209,12 +1157,15 @@ mod tests {
|
|||||||
)],
|
)],
|
||||||
)];
|
)];
|
||||||
|
|
||||||
let result = provider.generate_simple_session_description(&messages);
|
let result = crate::providers::cli_common::generate_simple_session_description(
|
||||||
|
"gpt-5.2-codex",
|
||||||
|
&messages,
|
||||||
|
);
|
||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
|
|
||||||
let (message, _usage) = result.unwrap();
|
let (message, usage) = result.unwrap();
|
||||||
|
assert_eq!(usage.model, "gpt-5.2-codex");
|
||||||
if let MessageContent::Text(text) = &message.content[0] {
|
if let MessageContent::Text(text) = &message.content[0] {
|
||||||
// Should be truncated to 4 words
|
|
||||||
let word_count = text.text.split_whitespace().count();
|
let word_count = text.text.split_whitespace().count();
|
||||||
assert!(word_count <= 4);
|
assert!(word_count <= 4);
|
||||||
} else {
|
} else {
|
||||||
@@ -1224,18 +1175,12 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_session_description_empty_messages() {
|
fn test_session_description_empty_messages() {
|
||||||
let provider = CodexProvider {
|
|
||||||
command: PathBuf::from("codex"),
|
|
||||||
model: ModelConfig::new("gpt-5.2-codex").unwrap(),
|
|
||||||
name: "codex".to_string(),
|
|
||||||
reasoning_effort: "high".to_string(),
|
|
||||||
skip_git_check: false,
|
|
||||||
mcp_config_overrides: Vec::new(),
|
|
||||||
};
|
|
||||||
|
|
||||||
let messages: Vec<Message> = vec![];
|
let messages: Vec<Message> = vec![];
|
||||||
|
|
||||||
let result = provider.generate_simple_session_description(&messages);
|
let result = crate::providers::cli_common::generate_simple_session_description(
|
||||||
|
"gpt-5.2-codex",
|
||||||
|
&messages,
|
||||||
|
);
|
||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
|
|
||||||
let (message, _usage) = result.unwrap();
|
let (message, _usage) = result.unwrap();
|
||||||
|
|||||||
@@ -275,51 +275,6 @@ impl CursorAgentProvider {
|
|||||||
|
|
||||||
Ok(lines)
|
Ok(lines)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Generate a simple session description without calling subprocess
|
|
||||||
fn generate_simple_session_description(
|
|
||||||
&self,
|
|
||||||
messages: &[Message],
|
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
|
||||||
// Extract the first user message text
|
|
||||||
let description = messages
|
|
||||||
.iter()
|
|
||||||
.find(|m| m.role == Role::User)
|
|
||||||
.and_then(|m| {
|
|
||||||
m.content.iter().find_map(|c| match c {
|
|
||||||
MessageContent::Text(text_content) => Some(&text_content.text),
|
|
||||||
_ => None,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.map(|text| {
|
|
||||||
// Take first few words, limit to 4 words
|
|
||||||
text.split_whitespace()
|
|
||||||
.take(4)
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
.join(" ")
|
|
||||||
})
|
|
||||||
.unwrap_or_else(|| "Simple task".to_string());
|
|
||||||
|
|
||||||
if std::env::var("GOOSE_CURSOR_AGENT_DEBUG").is_ok() {
|
|
||||||
println!("=== CURSOR AGENT PROVIDER DEBUG ===");
|
|
||||||
println!("Generated simple session description: {}", description);
|
|
||||||
println!("Skipped subprocess call for session description");
|
|
||||||
println!("================================");
|
|
||||||
}
|
|
||||||
|
|
||||||
let message = Message::new(
|
|
||||||
Role::Assistant,
|
|
||||||
chrono::Utc::now().timestamp(),
|
|
||||||
vec![MessageContent::text(description.clone())],
|
|
||||||
);
|
|
||||||
|
|
||||||
let usage = Usage::default();
|
|
||||||
|
|
||||||
Ok((
|
|
||||||
message,
|
|
||||||
ProviderUsage::new(self.model.model_name.clone(), usage),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ProviderDef for CursorAgentProvider {
|
impl ProviderDef for CursorAgentProvider {
|
||||||
@@ -378,9 +333,11 @@ impl Provider for CursorAgentProvider {
|
|||||||
messages: &[Message],
|
messages: &[Message],
|
||||||
tools: &[Tool],
|
tools: &[Tool],
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||||
// Check if this is a session description request (short system prompt asking for 4 words or less)
|
if super::cli_common::is_session_description_request(system) {
|
||||||
if system.contains("four words or less") || system.contains("4 words or less") {
|
return super::cli_common::generate_simple_session_description(
|
||||||
return self.generate_simple_session_description(messages);
|
&model_config.model_name,
|
||||||
|
messages,
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let lines = self.execute_command(system, messages, tools).await?;
|
let lines = self.execute_command(system, messages, tools).await?;
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ use tokio::io::{AsyncBufReadExt, AsyncReadExt, BufReader};
|
|||||||
use tokio::process::Command;
|
use tokio::process::Command;
|
||||||
|
|
||||||
use super::base::{Provider, ProviderDef, ProviderMetadata, ProviderUsage, Usage};
|
use super::base::{Provider, ProviderDef, ProviderMetadata, ProviderUsage, Usage};
|
||||||
|
use super::cli_common::{error_from_event, extract_usage_tokens};
|
||||||
use super::errors::ProviderError;
|
use super::errors::ProviderError;
|
||||||
use super::utils::{filter_extensions_from_system_prompt, RequestLog};
|
use super::utils::{filter_extensions_from_system_prompt, RequestLog};
|
||||||
use crate::config::base::GeminiCliCommand;
|
use crate::config::base::GeminiCliCommand;
|
||||||
@@ -31,29 +32,6 @@ pub const GEMINI_CLI_KNOWN_MODELS: &[&str] = &[
|
|||||||
|
|
||||||
pub const GEMINI_CLI_DOC_URL: &str = "https://ai.google.dev/gemini-api/docs";
|
pub const GEMINI_CLI_DOC_URL: &str = "https://ai.google.dev/gemini-api/docs";
|
||||||
|
|
||||||
fn extract_usage_from_stats(stats: &Value) -> Usage {
|
|
||||||
let get = |key: &str| {
|
|
||||||
stats
|
|
||||||
.get(key)
|
|
||||||
.and_then(|v| v.as_i64())
|
|
||||||
.and_then(|v| i32::try_from(v).ok())
|
|
||||||
};
|
|
||||||
Usage::new(
|
|
||||||
get("input_tokens"),
|
|
||||||
get("output_tokens"),
|
|
||||||
get("total_tokens"),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn error_from_event(parsed: &Value) -> ProviderError {
|
|
||||||
let error_msg = parsed
|
|
||||||
.get("error")
|
|
||||||
.and_then(|e| e.as_str())
|
|
||||||
.or_else(|| parsed.get("message").and_then(|m| m.as_str()))
|
|
||||||
.unwrap_or("Unknown error");
|
|
||||||
ProviderError::RequestFailed(format!("Gemini CLI error: {error_msg}"))
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, serde::Serialize)]
|
#[derive(Debug, serde::Serialize)]
|
||||||
pub struct GeminiCliProvider {
|
pub struct GeminiCliProvider {
|
||||||
command: PathBuf,
|
command: PathBuf,
|
||||||
@@ -95,48 +73,6 @@ impl GeminiCliProvider {
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn is_session_description_request(system: &str) -> bool {
|
|
||||||
system.contains("four words or less") || system.contains("4 words or less")
|
|
||||||
}
|
|
||||||
|
|
||||||
fn generate_simple_session_description(
|
|
||||||
&self,
|
|
||||||
messages: &[Message],
|
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
|
||||||
let description = messages
|
|
||||||
.iter()
|
|
||||||
.find(|m| m.role == Role::User)
|
|
||||||
.and_then(|m| {
|
|
||||||
m.content.iter().find_map(|c| match c {
|
|
||||||
MessageContent::Text(text_content) => Some(&text_content.text),
|
|
||||||
_ => None,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.map(|text| {
|
|
||||||
text.split_whitespace()
|
|
||||||
.take(4)
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
.join(" ")
|
|
||||||
})
|
|
||||||
.unwrap_or_else(|| "Simple task".to_string());
|
|
||||||
|
|
||||||
tracing::debug!(
|
|
||||||
description = %description,
|
|
||||||
"Generated simple session description, skipped subprocess"
|
|
||||||
);
|
|
||||||
|
|
||||||
let message = Message::new(
|
|
||||||
Role::Assistant,
|
|
||||||
chrono::Utc::now().timestamp(),
|
|
||||||
vec![MessageContent::text(description)],
|
|
||||||
);
|
|
||||||
|
|
||||||
Ok((
|
|
||||||
message,
|
|
||||||
ProviderUsage::new(self.model.model_name.clone(), Usage::default()),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build the prompt for the CLI invocation. When resuming a session the CLI
|
/// Build the prompt for the CLI invocation. When resuming a session the CLI
|
||||||
/// maintains conversation context internally, so only the latest user
|
/// maintains conversation context internally, so only the latest user
|
||||||
/// message is needed. On the first turn (no session yet) the system prompt
|
/// message is needed. On the first turn (no session yet) the system prompt
|
||||||
@@ -300,11 +236,11 @@ impl GeminiCliProvider {
|
|||||||
}
|
}
|
||||||
Some("result") => {
|
Some("result") => {
|
||||||
if let Some(stats) = parsed.get("stats") {
|
if let Some(stats) = parsed.get("stats") {
|
||||||
usage = extract_usage_from_stats(stats);
|
usage = extract_usage_tokens(stats);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some("error") => {
|
Some("error") => {
|
||||||
return Err(error_from_event(parsed));
|
return Err(error_from_event("Gemini CLI", parsed));
|
||||||
}
|
}
|
||||||
_ => {}
|
_ => {}
|
||||||
}
|
}
|
||||||
@@ -380,8 +316,11 @@ impl Provider for GeminiCliProvider {
|
|||||||
messages: &[Message],
|
messages: &[Message],
|
||||||
tools: &[Tool],
|
tools: &[Tool],
|
||||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||||
if Self::is_session_description_request(system) {
|
if super::cli_common::is_session_description_request(system) {
|
||||||
return self.generate_simple_session_description(messages);
|
return super::cli_common::generate_simple_session_description(
|
||||||
|
&model_config.model_name,
|
||||||
|
messages,
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let payload = json!({
|
let payload = json!({
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ pub mod bedrock;
|
|||||||
pub mod canonical;
|
pub mod canonical;
|
||||||
pub mod chatgpt_codex;
|
pub mod chatgpt_codex;
|
||||||
pub mod claude_code;
|
pub mod claude_code;
|
||||||
|
pub(crate) mod cli_common;
|
||||||
pub mod codex;
|
pub mod codex;
|
||||||
pub mod cursor_agent;
|
pub mod cursor_agent;
|
||||||
pub mod databricks;
|
pub mod databricks;
|
||||||
|
|||||||
Reference in New Issue
Block a user