fix(acp): forward ACP server context window size to clients (#9455)
Signed-off-by: Matt Toohey <contact@matttoohey.com>
This commit is contained in:
@@ -20,7 +20,7 @@ use std::future::Future;
|
|||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
use std::process::Stdio;
|
use std::process::Stdio;
|
||||||
use std::sync::{
|
use std::sync::{
|
||||||
atomic::{AtomicBool, Ordering},
|
atomic::{AtomicBool, AtomicU64, Ordering},
|
||||||
Arc, Mutex,
|
Arc, Mutex,
|
||||||
};
|
};
|
||||||
use std::thread::JoinHandle;
|
use std::thread::JoinHandle;
|
||||||
@@ -144,6 +144,11 @@ pub struct AcpProvider {
|
|||||||
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
|
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
|
||||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||||
handoff_context_sent: AtomicBool,
|
handoff_context_sent: AtomicBool,
|
||||||
|
/// Latest `size` reported by the ACP server in a `session/update` →
|
||||||
|
/// `usage_update` notification. 0 means no real update has arrived yet,
|
||||||
|
/// in which case `get_model_config()` falls back to the static model
|
||||||
|
/// configuration's context limit.
|
||||||
|
context_size: Arc<AtomicU64>,
|
||||||
|
|
||||||
tx: Option<mpsc::Sender<ClientRequest>>,
|
tx: Option<mpsc::Sender<ClientRequest>>,
|
||||||
loop_thread: Option<JoinHandle<()>>,
|
loop_thread: Option<JoinHandle<()>>,
|
||||||
@@ -222,10 +227,12 @@ impl AcpProvider {
|
|||||||
let goose_mode_shared = Arc::new(Mutex::new(goose_mode));
|
let goose_mode_shared = Arc::new(Mutex::new(goose_mode));
|
||||||
let pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>> =
|
let pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>> =
|
||||||
Arc::new(Mutex::new(HashMap::new()));
|
Arc::new(Mutex::new(HashMap::new()));
|
||||||
|
let context_size = Arc::new(AtomicU64::new(0));
|
||||||
let client_loop = AcpClientLoop::new(
|
let client_loop = AcpClientLoop::new(
|
||||||
config,
|
config,
|
||||||
goose_mode_shared.clone(),
|
goose_mode_shared.clone(),
|
||||||
pending_tool_updates.clone(),
|
pending_tool_updates.clone(),
|
||||||
|
context_size.clone(),
|
||||||
);
|
);
|
||||||
let loop_thread = spawn_client_loop(run(client_loop, rx, init_tx));
|
let loop_thread = spawn_client_loop(run(client_loop, rx, init_tx));
|
||||||
|
|
||||||
@@ -273,6 +280,7 @@ impl AcpProvider {
|
|||||||
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
||||||
pending_tool_updates,
|
pending_tool_updates,
|
||||||
handoff_context_sent: AtomicBool::new(false),
|
handoff_context_sent: AtomicBool::new(false),
|
||||||
|
context_size,
|
||||||
tx: Some(tx),
|
tx: Some(tx),
|
||||||
loop_thread: Some(loop_thread),
|
loop_thread: Some(loop_thread),
|
||||||
})
|
})
|
||||||
@@ -370,7 +378,12 @@ impl Provider for AcpProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn get_model_config(&self) -> ModelConfig {
|
fn get_model_config(&self) -> ModelConfig {
|
||||||
self.model.clone()
|
let mut model = self.model.clone();
|
||||||
|
let size = self.context_size.load(Ordering::Relaxed);
|
||||||
|
if size > 0 {
|
||||||
|
model.context_limit = Some(size as usize);
|
||||||
|
}
|
||||||
|
model
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> {
|
async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> {
|
||||||
@@ -629,6 +642,7 @@ struct AcpClientLoop {
|
|||||||
goose_mode: Arc<Mutex<GooseMode>>,
|
goose_mode: Arc<Mutex<GooseMode>>,
|
||||||
prompt_response_tx: Arc<Mutex<Option<mpsc::Sender<AcpUpdate>>>>,
|
prompt_response_tx: Arc<Mutex<Option<mpsc::Sender<AcpUpdate>>>>,
|
||||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||||
|
context_size: Arc<AtomicU64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl AcpClientLoop {
|
impl AcpClientLoop {
|
||||||
@@ -636,12 +650,14 @@ impl AcpClientLoop {
|
|||||||
config: AcpProviderConfig,
|
config: AcpProviderConfig,
|
||||||
goose_mode: Arc<Mutex<GooseMode>>,
|
goose_mode: Arc<Mutex<GooseMode>>,
|
||||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||||
|
context_size: Arc<AtomicU64>,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
config,
|
config,
|
||||||
goose_mode,
|
goose_mode,
|
||||||
prompt_response_tx: Arc::new(Mutex::new(None)),
|
prompt_response_tx: Arc::new(Mutex::new(None)),
|
||||||
pending_tool_updates,
|
pending_tool_updates,
|
||||||
|
context_size,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -695,6 +711,7 @@ impl AcpClientLoop {
|
|||||||
goose_mode,
|
goose_mode,
|
||||||
prompt_response_tx,
|
prompt_response_tx,
|
||||||
pending_tool_updates,
|
pending_tool_updates,
|
||||||
|
context_size,
|
||||||
} = self;
|
} = self;
|
||||||
let notification_callback = config.notification_callback.clone();
|
let notification_callback = config.notification_callback.clone();
|
||||||
let reverse_modes = reverse_mode_mapping(&config.mode_mapping);
|
let reverse_modes = reverse_mode_mapping(&config.mode_mapping);
|
||||||
@@ -707,6 +724,7 @@ impl AcpClientLoop {
|
|||||||
let reverse_modes = reverse_modes.clone();
|
let reverse_modes = reverse_modes.clone();
|
||||||
let goose_mode = goose_mode.clone();
|
let goose_mode = goose_mode.clone();
|
||||||
let pending_tool_updates = pending_tool_updates.clone();
|
let pending_tool_updates = pending_tool_updates.clone();
|
||||||
|
let context_size = context_size.clone();
|
||||||
async move |notification: SessionNotification, _cx| {
|
async move |notification: SessionNotification, _cx| {
|
||||||
if let Some(ref cb) = notification_callback {
|
if let Some(ref cb) = notification_callback {
|
||||||
cb(notification.clone());
|
cb(notification.clone());
|
||||||
@@ -740,6 +758,9 @@ impl AcpClientLoop {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
SessionUpdate::UsageUpdate(usage) => {
|
||||||
|
context_size.store(usage.size, Ordering::Relaxed);
|
||||||
|
}
|
||||||
_ => {}
|
_ => {}
|
||||||
}
|
}
|
||||||
if let Some(tx) = prompt_response_tx
|
if let Some(tx) = prompt_response_tx
|
||||||
@@ -1529,6 +1550,7 @@ mod tests {
|
|||||||
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
||||||
pending_tool_updates: Arc::new(Mutex::new(HashMap::new())),
|
pending_tool_updates: Arc::new(Mutex::new(HashMap::new())),
|
||||||
handoff_context_sent: AtomicBool::new(false),
|
handoff_context_sent: AtomicBool::new(false),
|
||||||
|
context_size: Arc::new(AtomicU64::new(0)),
|
||||||
tx,
|
tx,
|
||||||
loop_thread: None,
|
loop_thread: None,
|
||||||
}
|
}
|
||||||
@@ -1633,6 +1655,18 @@ mod tests {
|
|||||||
assert!(!later_claim.include_context);
|
assert!(!later_claim.include_context);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn get_model_config_surfaces_captured_context_size() {
|
||||||
|
let provider = test_provider();
|
||||||
|
assert_eq!(
|
||||||
|
provider.get_model_config().context_limit(),
|
||||||
|
crate::model::DEFAULT_CONTEXT_LIMIT
|
||||||
|
);
|
||||||
|
|
||||||
|
provider.context_size.store(200_000, Ordering::Relaxed);
|
||||||
|
assert_eq!(provider.get_model_config().context_limit(), 200_000);
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn failed_first_prompt_send_rolls_back_handoff_context_claim() {
|
async fn failed_first_prompt_send_rolls_back_handoff_context_claim() {
|
||||||
let (tx, rx) = mpsc::channel(1);
|
let (tx, rx) = mpsc::channel(1);
|
||||||
|
|||||||
Reference in New Issue
Block a user