fix(acp): resume provider-native sessions (#10379)

Co-authored-by: John Tennant <jtennant@block.xyz>
Co-authored-by: Douwe M Osinga <douwe@sidewalklabs.com>
This commit is contained in:
John Matthew Tennant
2026-08-10 06:08:34 -04:00
committed by GitHub
parent d41b7dad89
commit 5e7430fefa
8 changed files with 257 additions and 26 deletions
+8
View File
@@ -426,6 +426,14 @@ pub trait Provider: Send + Sync {
/// Get the name of this provider instance
fn get_name(&self) -> &str;
fn provider_session_id(&self) -> Option<String> {
None
}
async fn resume(&self, _session_id: &str) -> Result<(), ProviderError> {
Ok(())
}
/// Primary streaming method that all providers must implement.
async fn stream(
&self,
@@ -1856,6 +1856,7 @@ mod tests {
provider: "test-provider".to_string(),
requested_model: "test-model".to_string(),
resolved_model: None,
provider_session_id: None,
};
let mut limited = Message::assistant()
.with_id("turn-1")
@@ -673,6 +673,8 @@ pub struct InferenceMetadata {
pub requested_model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub resolved_model: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_session_id: Option<String>,
}
#[derive(Clone, PartialEq, Serialize, Deserialize, Debug, Default)]
+162 -16
View File
@@ -1,8 +1,8 @@
use agent_client_protocol::schema::v1::{
Annotations as AcpAnnotations, ClientCapabilities, CloseSessionRequest, ContentBlock,
ContentChunk, EnvVariable, HttpHeader, ImageContent, InitializeRequest, InitializeResponse,
McpCapabilities, McpServer, McpServerHttp, McpServerStdio, NewSessionRequest,
NewSessionResponse, PromptRequest, PromptResponse, RequestPermissionOutcome,
LoadSessionRequest, McpCapabilities, McpServer, McpServerHttp, McpServerStdio,
NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, RequestPermissionOutcome,
RequestPermissionRequest, RequestPermissionResponse, Role as AcpRole, SessionConfigKind,
SessionConfigOption, SessionConfigOptionCategory, SessionConfigSelectOptions, SessionId,
SessionModeState, SessionNotification, SessionUpdate, SetSessionConfigOptionRequest,
@@ -69,6 +69,13 @@ enum ClientRequest {
NewSession {
response_tx: oneshot::Sender<Result<NewSessionResponse>>,
},
LoadSession {
session_id: SessionId,
response_tx: oneshot::Sender<Result<NewSessionResponse>>,
},
CloseSession {
session_id: SessionId,
},
SetMode {
session_id: SessionId,
mode_id: String,
@@ -171,7 +178,7 @@ pub struct AcpProvider {
goose_mode: Arc<Mutex<GooseMode>>,
mode_mapping: HashMap<GooseMode, Vec<String>>,
session: AcpSession,
session: Mutex<AcpSession>,
pending_confirmations:
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
@@ -304,7 +311,7 @@ impl AcpProvider {
name,
goose_mode: goose_mode_shared,
mode_mapping,
session,
session: Mutex::new(session),
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
pending_tool_updates,
handoff_context_sent: AtomicBool::new(false),
@@ -318,7 +325,25 @@ impl AcpProvider {
}
fn acp_session_id(&self) -> SessionId {
self.session.id.clone()
self.session.lock().unwrap().id.clone()
}
async fn load_session(&self, session_id: SessionId) -> Result<AcpSession> {
let (response_tx, response_rx) = oneshot::channel();
self.tx
.as_ref()
.unwrap()
.send(ClientRequest::LoadSession {
session_id,
response_tx,
})
.await
.context("ACP client is unavailable")?;
let response = response_rx.await.context("ACP session load cancelled")??;
Ok(AcpSession {
id: response.session_id.clone(),
response,
})
}
pub(crate) async fn send_set_mode(&self, _goose_id: &str, mode_id: String) -> Result<()> {
@@ -413,6 +438,8 @@ impl AcpProvider {
fn session_has_config_option(&self, category: SessionConfigOptionCategory) -> bool {
self.session
.lock()
.unwrap()
.response
.config_options
.as_ref()
@@ -441,6 +468,33 @@ impl Provider for AcpProvider {
&self.name
}
fn provider_session_id(&self) -> Option<String> {
Some(self.acp_session_id().to_string())
}
async fn resume(&self, session_id: &str) -> Result<(), ProviderError> {
if self.acp_session_id().0.as_ref() == session_id {
return Ok(());
}
let previous_session_id = self.acp_session_id();
let loaded = self
.load_session(SessionId::new(session_id))
.await
.map_err(|error| ProviderError::RequestFailed(error.to_string()))?;
*self.session.lock().unwrap() = loaded;
self.handoff_context_sent.store(true, Ordering::Release);
let _ = self
.tx
.as_ref()
.unwrap()
.send(ClientRequest::CloseSession {
session_id: previous_session_id,
})
.await;
Ok(())
}
async fn get_context_limit(&self, model_config: &ModelConfig) -> Result<usize, ProviderError> {
let size = self.context_size.load(Ordering::Relaxed);
if size > 0 {
@@ -451,8 +505,9 @@ impl Provider for AcpProvider {
async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> {
if let Some(candidates) = self.mode_mapping.get(&mode) {
let mode_str = select_mode_id(candidates, self.session.response.modes.as_ref())
.ok_or_else(|| {
let session = self.session.lock().unwrap().clone();
let mode_str =
select_mode_id(candidates, session.response.modes.as_ref()).ok_or_else(|| {
ProviderError::RequestFailed(format!(
"None of the mode ids [{}] are offered by the agent",
candidates.join(", ")
@@ -708,7 +763,8 @@ impl Provider for AcpProvider {
}
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
let (_, available) = resolve_model_info(&self.name, &self.session.response)?;
let session = self.session.lock().unwrap().clone();
let (_, available) = resolve_model_info(&self.name, &session.response)?;
Ok(available)
}
}
@@ -1125,6 +1181,7 @@ async fn handle_requests(
.session_capabilities
.close
.is_some();
let supports_load = init_response.agent_capabilities.load_session;
let mcp_capabilities = init_response.agent_capabilities.mcp_capabilities.clone();
if let Some(tx) = init_tx.take() {
log_undelivered(tx.send(Ok(init_response)), AGENT_METHOD_NAMES.initialize);
@@ -1156,6 +1213,52 @@ async fn handle_requests(
};
log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_new);
}
ClientRequest::LoadSession {
session_id,
response_tx,
} => {
let result = if supports_load {
let mcp_servers =
filter_supported_servers(&config.mcp_servers, &mcp_capabilities);
cx.send_request(
LoadSessionRequest::new(session_id.clone(), config.work_dir.clone())
.mcp_servers(mcp_servers),
)
.block_task()
.await
.map(|response| {
NewSessionResponse::new(session_id.clone())
.modes(response.modes)
.config_options(response.config_options)
.meta(response.meta)
})
.map_err(anyhow::Error::from)
} else {
Err(anyhow::anyhow!("ACP agent does not support session/load"))
};
let result = match result {
Ok(session) => {
session_ids.push(session.session_id.clone());
apply_session_config_options(&config, &cx, session.session_id.clone())
.await?;
apply_session_mode(&config, &goose_mode, &cx, session).await
}
Err(error) => Err(error),
};
log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_load);
}
ClientRequest::CloseSession { session_id } => {
if supports_close {
if let Err(error) = cx
.send_request(CloseSessionRequest::new(session_id.clone()))
.block_task()
.await
{
tracing::debug!(method = AGENT_METHOD_NAMES.session_close, session_id = %session_id, %error, "failed to close replaced ACP session");
}
}
session_ids.retain(|id| id != &session_id);
}
ClientRequest::SetMode {
session_id,
mode_id,
@@ -1223,7 +1326,7 @@ async fn handle_requests(
}
}
if supports_close {
if supports_close && !supports_load {
for session_id in session_ids {
if let Err(e) = cx
.send_request(CloseSessionRequest::new(session_id.clone()))
@@ -1752,10 +1855,10 @@ mod tests {
name: "acp-test".to_string(),
goose_mode: Arc::new(Mutex::new(GooseMode::Auto)),
mode_mapping: HashMap::new(),
session: AcpSession {
session: Mutex::new(AcpSession {
id: SessionId::new("test-session"),
response: NewSessionResponse::new("test-session"),
},
}),
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
pending_tool_updates: Arc::new(Mutex::new(HashMap::new())),
handoff_context_sent: AtomicBool::new(false),
@@ -2081,6 +2184,50 @@ mod tests {
assert!(!later_claim.include_context);
}
#[tokio::test]
async fn resume_replaces_session_and_skips_handoff() {
let (tx, mut rx) = mpsc::channel(2);
let (provider, _) = test_provider_with_tx(Some(tx));
let handle = tokio::spawn(async move {
provider.resume("saved-session").await.unwrap();
provider
});
let ClientRequest::LoadSession {
session_id,
response_tx,
} = rx.recv().await.expect("expected session/load")
else {
panic!("expected session/load");
};
assert_eq!(session_id.to_string(), "saved-session");
response_tx
.send(Ok(NewSessionResponse::new("saved-session")))
.unwrap();
let ClientRequest::CloseSession { session_id } =
rx.recv().await.expect("expected temporary session close")
else {
panic!("expected temporary session close");
};
assert_eq!(session_id.to_string(), "test-session");
let provider = handle.await.unwrap();
assert_eq!(
provider.provider_session_id().as_deref(),
Some("saved-session")
);
let messages = vec![
Message::assistant().with_text("prior answer"),
Message::user().with_text("current request"),
];
let claim = provider.claim_handoff_context(&messages);
assert!(!claim.first_prompt);
assert!(!claim.include_context);
}
#[tokio::test]
async fn get_context_limit_surfaces_captured_context_size() {
let (provider, model) = test_provider();
@@ -2324,10 +2471,9 @@ mod tests {
GooseMode::Auto,
vec!["full-access".to_string(), "agent-full-access".to_string()],
)]);
provider.session.response = NewSessionResponse::new("test-session").modes(mode_state(
"read-only",
&["read-only", "agent", "agent-full-access"],
));
provider.session.lock().unwrap().response = NewSessionResponse::new("test-session").modes(
mode_state("read-only", &["read-only", "agent", "agent-full-access"]),
);
let handle = tokio::spawn(async move {
provider
@@ -2358,7 +2504,7 @@ mod tests {
let (tx, mut rx) = mpsc::channel(1);
let (mut provider, _) = test_provider_with_tx(Some(tx));
provider.mode_mapping = HashMap::from([(GooseMode::Chat, vec!["read-only".to_string()])]);
provider.session.response = NewSessionResponse::new("test-session")
provider.session.lock().unwrap().response = NewSessionResponse::new("test-session")
.modes(mode_state("agent", &["agent", "agent-full-access"]));
let result = provider.update_mode("session", GooseMode::Chat).await;
+48 -5
View File
@@ -2163,17 +2163,33 @@ impl Agent {
let provider = self.provider().await?;
let provider_name = provider.get_name().to_string();
let saved_provider_session_id =
super::latest_provider_session_id(conversation.messages(), &provider_name);
if let Some(saved_provider_session_id) = saved_provider_session_id {
if let Err(error) = provider.resume(saved_provider_session_id).await {
warn!(
provider = provider_name,
%error,
"Could not resume provider session; continuing with a handoff"
);
}
}
let requested_model = model_config.model_name.clone();
let inference = provider
let resolved_model = provider
.fetch_model_info(&requested_model)
.await
.ok()
.and_then(|model_info| model_info.resolved_model)
.map(|resolved_model| InferenceMetadata {
.and_then(|model_info| model_info.resolved_model);
let provider_session_id = provider.provider_session_id();
let inference = (resolved_model.is_some() || provider_session_id.is_some()).then(|| {
InferenceMetadata {
provider: provider_name.clone(),
requested_model,
resolved_model: Some(resolved_model),
});
resolved_model,
provider_session_id,
}
});
let session_manager = self.config.session_manager.clone();
let session_id = session_config.id.clone();
if !self.config.disable_session_naming {
@@ -3941,6 +3957,33 @@ mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use tempfile::TempDir;
#[test]
fn provider_session_id_comes_from_latest_inference() {
let messages = vec![
Message::assistant().with_inference(InferenceMetadata {
provider: "codex-acp".to_string(),
requested_model: "current".to_string(),
resolved_model: None,
provider_session_id: Some("codex-session".to_string()),
}),
Message::assistant().with_inference(InferenceMetadata {
provider: "claude-acp".to_string(),
requested_model: "current".to_string(),
resolved_model: None,
provider_session_id: Some("claude-session".to_string()),
}),
];
assert_eq!(
super::super::latest_provider_session_id(&messages, "claude-acp"),
Some("claude-session")
);
assert_eq!(
super::super::latest_provider_session_id(&messages, "codex-acp"),
None
);
}
#[test]
fn recipe_history_excludes_turn_context_events() {
use crate::conversation::message::MessageMetadata;
+13
View File
@@ -36,3 +36,16 @@ pub use subagent_handler::SUBAGENT_TOOL_REQUEST_TYPE;
pub use subagent_task_config::TaskConfig;
pub use tool_execution::ToolCallContext;
pub use types::{FrontendTool, RetryConfig, SessionConfig, SuccessCheck};
fn latest_provider_session_id<'a>(
messages: &'a [crate::conversation::message::Message],
provider: &str,
) -> Option<&'a str> {
let inference = messages
.iter()
.rev()
.find_map(|message| message.metadata.inference.as_ref())?;
(inference.provider == provider)
.then_some(inference.provider_session_id.as_deref())
.flatten()
}
@@ -434,6 +434,19 @@ impl Inference for InferenceRunner<'_> {
.get_context_limit(&self.model_config)
.await
.unwrap_or_else(|_| self.model_config.context_limit());
let provider_name = self.provider.get_name();
if let Some(session_id) = super::super::latest_provider_session_id(
conversation.messages(),
provider_name,
) {
if let Err(error) = self.provider.resume(session_id).await {
tracing::warn!(
provider = provider_name,
%error,
"Could not resume provider session; continuing with a handoff"
);
}
}
let turn = messages_since_kickoff(conversation)?;
let turn_start = turn
.first()
@@ -481,17 +494,21 @@ impl Inference for InferenceRunner<'_> {
};
let requested_model = self.model_config.model_name.clone();
let inference = self
let resolved_model = self
.provider
.fetch_model_info(&requested_model)
.await
.ok()
.and_then(|model_info| model_info.resolved_model)
.map(|resolved_model| InferenceMetadata {
.and_then(|model_info| model_info.resolved_model);
let provider_session_id = self.provider.provider_session_id();
let inference = (resolved_model.is_some() || provider_session_id.is_some()).then(|| {
InferenceMetadata {
provider: self.provider.get_name().to_string(),
requested_model,
resolved_model: Some(resolved_model),
});
resolved_model,
provider_session_id,
}
});
let mut accumulator = Conversation::empty();
let mut tool_request_ids = std::collections::HashSet::new();
+1
View File
@@ -168,6 +168,7 @@ export type InferenceMetadata = {
provider: string;
requestedModel: string;
resolvedModel?: string | null;
providerSessionId?: string | null;
};
/** Mirrors the backend `MessageUsage` schema (camelCase). */