mirror of
https://github.com/aaif-goose/goose.git
synced 2026-07-03 14:10:03 +02:00
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::process::Stdio;
|
||||
use std::sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
atomic::{AtomicBool, AtomicU64, Ordering},
|
||||
Arc, Mutex,
|
||||
};
|
||||
use std::thread::JoinHandle;
|
||||
@@ -144,6 +144,11 @@ pub struct AcpProvider {
|
||||
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
|
||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||
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>>,
|
||||
loop_thread: Option<JoinHandle<()>>,
|
||||
@@ -222,10 +227,12 @@ impl AcpProvider {
|
||||
let goose_mode_shared = Arc::new(Mutex::new(goose_mode));
|
||||
let pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>> =
|
||||
Arc::new(Mutex::new(HashMap::new()));
|
||||
let context_size = Arc::new(AtomicU64::new(0));
|
||||
let client_loop = AcpClientLoop::new(
|
||||
config,
|
||||
goose_mode_shared.clone(),
|
||||
pending_tool_updates.clone(),
|
||||
context_size.clone(),
|
||||
);
|
||||
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_tool_updates,
|
||||
handoff_context_sent: AtomicBool::new(false),
|
||||
context_size,
|
||||
tx: Some(tx),
|
||||
loop_thread: Some(loop_thread),
|
||||
})
|
||||
@@ -370,7 +378,12 @@ impl Provider for AcpProvider {
|
||||
}
|
||||
|
||||
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> {
|
||||
@@ -629,6 +642,7 @@ struct AcpClientLoop {
|
||||
goose_mode: Arc<Mutex<GooseMode>>,
|
||||
prompt_response_tx: Arc<Mutex<Option<mpsc::Sender<AcpUpdate>>>>,
|
||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||
context_size: Arc<AtomicU64>,
|
||||
}
|
||||
|
||||
impl AcpClientLoop {
|
||||
@@ -636,12 +650,14 @@ impl AcpClientLoop {
|
||||
config: AcpProviderConfig,
|
||||
goose_mode: Arc<Mutex<GooseMode>>,
|
||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||
context_size: Arc<AtomicU64>,
|
||||
) -> Self {
|
||||
Self {
|
||||
config,
|
||||
goose_mode,
|
||||
prompt_response_tx: Arc::new(Mutex::new(None)),
|
||||
pending_tool_updates,
|
||||
context_size,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -695,6 +711,7 @@ impl AcpClientLoop {
|
||||
goose_mode,
|
||||
prompt_response_tx,
|
||||
pending_tool_updates,
|
||||
context_size,
|
||||
} = self;
|
||||
let notification_callback = config.notification_callback.clone();
|
||||
let reverse_modes = reverse_mode_mapping(&config.mode_mapping);
|
||||
@@ -707,6 +724,7 @@ impl AcpClientLoop {
|
||||
let reverse_modes = reverse_modes.clone();
|
||||
let goose_mode = goose_mode.clone();
|
||||
let pending_tool_updates = pending_tool_updates.clone();
|
||||
let context_size = context_size.clone();
|
||||
async move |notification: SessionNotification, _cx| {
|
||||
if let Some(ref cb) = notification_callback {
|
||||
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
|
||||
@@ -1529,6 +1550,7 @@ mod tests {
|
||||
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
||||
pending_tool_updates: Arc::new(Mutex::new(HashMap::new())),
|
||||
handoff_context_sent: AtomicBool::new(false),
|
||||
context_size: Arc::new(AtomicU64::new(0)),
|
||||
tx,
|
||||
loop_thread: None,
|
||||
}
|
||||
@@ -1633,6 +1655,18 @@ mod tests {
|
||||
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]
|
||||
async fn failed_first_prompt_send_rolls_back_handoff_context_claim() {
|
||||
let (tx, rx) = mpsc::channel(1);
|
||||
|
||||
Reference in New Issue
Block a user