Agent loop defensive (#3554)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-07-22 18:21:47 +02:00
committed by GitHub
parent b3cd03ef61
commit 9f356e7009
19 changed files with 349 additions and 773 deletions
+2 -4
View File
@@ -209,7 +209,7 @@ async fn serve_static(axum::extract::Path(path): axum::extract::Path<String>) ->
include_bytes!("../../../../documentation/static/img/logo_light.png").to_vec(),
)
.into_response(),
_ => (axum::http::StatusCode::NOT_FOUND, "Not found").into_response(),
_ => (http::StatusCode::NOT_FOUND, "Not found").into_response(),
}
}
@@ -484,7 +484,6 @@ async fn process_message_streaming(
)
.await?;
// Create a session config
let session_config = SessionConfig {
id: session::Identifier::Path(session_file.clone()),
working_dir: std::env::current_dir()?,
@@ -494,8 +493,7 @@ async fn process_message_streaming(
retry_config: None,
};
// Get response from agent
match agent.reply(&messages, Some(session_config)).await {
match agent.reply(&messages, Some(session_config), None).await {
Ok(mut stream) => {
while let Some(result) = stream.next().await {
match result {
+17 -18
View File
@@ -48,6 +48,7 @@ use std::path::PathBuf;
use std::sync::Arc;
use std::time::Instant;
use tokio;
use tokio_util::sync::CancellationToken;
pub enum RunMode {
Normal,
@@ -132,13 +133,10 @@ impl Session {
retry_config: Option<RetryConfig>,
) -> Self {
let messages = if let Some(session_file) = &session_file {
match session::read_messages(session_file) {
Ok(msgs) => msgs,
Err(e) => {
eprintln!("Warning: Failed to load message history: {}", e);
Vec::new()
}
}
session::read_messages(session_file).unwrap_or_else(|e| {
eprintln!("Warning: Failed to load message history: {}", e);
Vec::new()
})
} else {
// Don't try to read messages if we're not saving sessions
Vec::new()
@@ -180,7 +178,7 @@ impl Session {
/// Format: "ENV1=val1 ENV2=val2 command args..."
pub async fn add_extension(&mut self, extension_command: String) -> Result<()> {
let mut parts: Vec<&str> = extension_command.split_whitespace().collect();
let mut envs = std::collections::HashMap::new();
let mut envs = HashMap::new();
// Parse environment variables (format: KEY=value)
while let Some(part) = parts.first() {
@@ -473,7 +471,7 @@ impl Session {
self.display_context_usage().await?;
match input::get_input(&mut editor)? {
input::InputResult::Message(content) => {
InputResult::Message(content) => {
match self.run_mode {
RunMode::Normal => {
save_history(&mut editor);
@@ -495,15 +493,11 @@ impl Session {
eprintln!("Warning: Failed to update project tracker with instruction: {}", e);
}
// Get the provider from the agent for description generation
let provider = self.agent.provider().await?;
// Persist messages with provider for automatic description generation
if let Some(session_file) = &self.session_file {
let working_dir = Some(
std::env::current_dir()
.expect("failed to get current session working directory"),
);
let working_dir = Some(std::env::current_dir().unwrap_or_default());
session::persist_messages_with_schedule_id(
session_file,
@@ -847,12 +841,14 @@ impl Session {
}
async fn process_agent_response(&mut self, interactive: bool) -> Result<()> {
let cancel_token = CancellationToken::new();
let cancel_token_clone = cancel_token.clone();
let session_config = self.session_file.as_ref().map(|s| {
let session_id = session::Identifier::Path(s.clone());
SessionConfig {
id: session_id.clone(),
working_dir: std::env::current_dir()
.expect("failed to get current session working directory"),
working_dir: std::env::current_dir().unwrap_or_default(),
schedule_id: self.scheduled_job_id.clone(),
execution_mode: None,
max_turns: self.max_turns,
@@ -861,7 +857,7 @@ impl Session {
});
let mut stream = self
.agent
.reply(&self.messages, session_config.clone())
.reply(&self.messages, session_config.clone(), Some(cancel_token))
.await?;
let mut progress_bars = output::McpSpinners::new();
@@ -919,7 +915,7 @@ impl Session {
)
.await?;
}
cancel_token_clone.cancel();
drop(stream);
break;
} else {
@@ -1001,6 +997,7 @@ impl Session {
.reply(
&self.messages,
session_config.clone(),
None
)
.await?;
}
@@ -1157,6 +1154,7 @@ impl Session {
Some(Err(e)) => {
eprintln!("Error: {}", e);
cancel_token_clone.cancel();
drop(stream);
if let Err(e) = self.handle_interrupted_messages(false).await {
eprintln!("Error handling interruption: {}", e);
@@ -1173,6 +1171,7 @@ impl Session {
}
}
_ = tokio::signal::ctrl_c() => {
cancel_token_clone.cancel();
drop(stream);
if let Err(e) = self.handle_interrupted_messages(true).await {
eprintln!("Error handling interruption: {}", e);