Repository navigation
Expand file tree
/
Copy pathexternal_tool_cancellation.rs
More file actions
123 lines (109 loc) · 4.59 KB
/
Copy pathexternal_tool_cancellation.rs
File metadata and controls
123 lines (109 loc) · 4.59 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
use std::sync::Arc;
use async_trait::async_trait;
use github_copilot_sdk::handler::ApproveAllHandler;
use github_copilot_sdk::tool::ToolHandler;
use github_copilot_sdk::{Error, SessionConfig, Tool, ToolInvocation, ToolResult};
use serde_json::json;
use tokio::sync::{Mutex, mpsc, oneshot};
use tokio::time::{Duration, timeout};
use super::support::DEFAULT_TEST_TOKEN;
#[tokio::test]
async fn should_cancel_tool_handler_when_session_disconnects() {
super::support::with_dedicated_e2e_context(
"external_tool_cancellation",
"should_cancel_tool_handler_when_session_disconnects",
|ctx| {
Box::pin(async move {
ctx.set_default_copilot_user();
let client = ctx.start_client().await;
let (started_tx, mut started_rx) = mpsc::unbounded_channel();
let (release_tx, release_rx) = oneshot::channel();
let (cancelled_tx, cancelled_rx) = oneshot::channel();
let tool = Arc::new(CancelAwareSlowTool {
started_tx,
release_rx: Mutex::new(Some(release_rx)),
cancelled_tx: Mutex::new(Some(cancelled_tx)),
});
let session = client
.create_session(
SessionConfig::default()
.with_github_token(DEFAULT_TEST_TOKEN)
.with_permission_handler(Arc::new(ApproveAllHandler))
.with_tools(vec![
Tool::new("slow_analysis")
.with_description(
"A slow analysis tool that blocks until released",
)
.with_parameters(json!({
"type": "object",
"properties": {
"value": {
"type": "string",
"description": "Value to analyze"
}
},
"required": ["value"]
}))
.with_handler(tool),
]),
)
.await
.expect("create session");
session
.send("Use slow_analysis with value 'test_abort'. Wait for the result.")
.await
.expect("send tool turn");
let started_value = timeout(Duration::from_secs(60), started_rx.recv())
.await
.expect("tool start wait timed out")
.expect("tool start channel closed");
assert_eq!(started_value, "test_abort");
session.disconnect().await.expect("disconnect session");
timeout(Duration::from_secs(60), cancelled_rx)
.await
.expect("tool cancellation wait timed out")
.expect("tool cancellation sender dropped");
let _ = release_tx.send("RELEASED".to_string());
client.stop().await.expect("stop client");
})
},
)
.await;
}
struct CancelAwareSlowTool {
started_tx: mpsc::UnboundedSender<String>,
release_rx: Mutex<Option<oneshot::Receiver<String>>>,
cancelled_tx: Mutex<Option<oneshot::Sender<()>>>,
}
struct CancelSignalGuard {
cancelled_tx: Option<oneshot::Sender<()>>,
}
impl Drop for CancelSignalGuard {
fn drop(&mut self) {
if let Some(sender) = self.cancelled_tx.take() {
let _ = sender.send(());
}
}
}
#[async_trait]
impl ToolHandler for CancelAwareSlowTool {
async fn call(&self, invocation: ToolInvocation) -> Result<ToolResult, Error> {
let value = invocation
.arguments
.get("value")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.to_string();
let _ = self.started_tx.send(value);
let cancelled_tx = self.cancelled_tx.lock().await.take();
let _guard = CancelSignalGuard { cancelled_tx };
let release_rx = self
.release_rx
.lock()
.await
.take()
.expect("slow tool called once");
let released = release_rx.await.unwrap_or_else(|_| "released".to_string());
Ok(ToolResult::Text(released))
}
}