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:
committed by
GitHub
parent
d41b7dad89
commit
5e7430fefa
@@ -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)]
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -168,6 +168,7 @@ export type InferenceMetadata = {
|
||||
provider: string;
|
||||
requestedModel: string;
|
||||
resolvedModel?: string | null;
|
||||
providerSessionId?: string | null;
|
||||
};
|
||||
|
||||
/** Mirrors the backend `MessageUsage` schema (camelCase). */
|
||||
|
||||
Reference in New Issue
Block a user