|
1 | 1 | use std::net::TcpListener; |
2 | 2 | use std::sync::Arc; |
| 3 | +use std::sync::atomic::{AtomicUsize, Ordering}; |
3 | 4 | use std::time::Duration; |
4 | 5 |
|
5 | 6 | use async_trait::async_trait; |
6 | | -use github_copilot_sdk::generated::session_events::SessionEventType; |
| 7 | +use github_copilot_sdk::generated::session_events::{ |
| 8 | + PermissionCompletedData, PermissionResult as EventPermissionResult, SessionEventType, |
| 9 | +}; |
7 | 10 | use github_copilot_sdk::handler::{PermissionResult, SessionHandler}; |
8 | 11 | use github_copilot_sdk::{ |
9 | | - Client, PermissionRequestData, RequestId, ResumeSessionConfig, SessionConfig, SessionId, Tool, |
10 | | - ToolInvocation, ToolResult, Transport, |
| 12 | + Client, PermissionRequestData, RequestId, ResumeSessionConfig, SessionConfig, SessionEvent, |
| 13 | + SessionId, Tool, ToolInvocation, ToolResult, Transport, |
11 | 14 | }; |
12 | 15 | use serde_json::json; |
13 | 16 |
|
@@ -101,16 +104,169 @@ async fn both_clients_see_tool_request_and_completion_events() { |
101 | 104 |
|
102 | 105 | #[tokio::test] |
103 | 106 | async fn one_client_approves_permission_and_both_see_the_result() { |
104 | | - let result = PermissionResult::Approved; |
| 107 | + with_e2e_context( |
| 108 | + "rust_multi_client", |
| 109 | + "one_client_approves_permission_and_both_see_the_result", |
| 110 | + |ctx| { |
| 111 | + Box::pin(async move { |
| 112 | + ctx.set_default_copilot_user(); |
| 113 | + let port = free_tcp_port(); |
| 114 | + let server = start_tcp_server(ctx, port).await; |
| 115 | + let permission_requests = Arc::new(AtomicUsize::new(0)); |
| 116 | + let session1 = server |
| 117 | + .create_session( |
| 118 | + SessionConfig::default() |
| 119 | + .with_github_token(DEFAULT_TEST_TOKEN) |
| 120 | + .with_handler(permission_handler_with_counter( |
| 121 | + PermissionResult::Approved, |
| 122 | + Arc::clone(&permission_requests), |
| 123 | + )), |
| 124 | + ) |
| 125 | + .await |
| 126 | + .expect("create session"); |
| 127 | + let client2 = start_external_client(ctx, port).await; |
| 128 | + let session2 = client2 |
| 129 | + .resume_session( |
| 130 | + resume_config(session1.id().clone()) |
| 131 | + .with_request_permission(false) |
| 132 | + .with_handler(permission_handler(PermissionResult::NoResult)), |
| 133 | + ) |
| 134 | + .await |
| 135 | + .expect("resume session"); |
| 136 | + |
| 137 | + let client1_requested = wait_for_event( |
| 138 | + session1.subscribe(), |
| 139 | + "client1 permission request", |
| 140 | + |event| event.parsed_type() == SessionEventType::PermissionRequested, |
| 141 | + ); |
| 142 | + let client2_requested = wait_for_event( |
| 143 | + session2.subscribe(), |
| 144 | + "client2 permission request", |
| 145 | + |event| event.parsed_type() == SessionEventType::PermissionRequested, |
| 146 | + ); |
| 147 | + let client1_completed = wait_for_event( |
| 148 | + session1.subscribe(), |
| 149 | + "client1 permission approved", |
| 150 | + |event| is_permission_approved(event), |
| 151 | + ); |
| 152 | + let client2_completed = wait_for_event( |
| 153 | + session2.subscribe(), |
| 154 | + "client2 permission approved", |
| 155 | + |event| is_permission_approved(event), |
| 156 | + ); |
| 157 | + |
| 158 | + let answer = session1 |
| 159 | + .send_and_wait( |
| 160 | + "Create a file called hello.txt containing the text 'hello world'", |
| 161 | + ) |
| 162 | + .await |
| 163 | + .expect("send") |
| 164 | + .expect("assistant message"); |
| 165 | + assert!(!assistant_message_content(&answer).is_empty()); |
| 166 | + assert!( |
| 167 | + permission_requests.load(Ordering::SeqCst) > 0, |
| 168 | + "expected client 1 to handle at least one permission request" |
| 169 | + ); |
| 170 | + let _ = tokio::join!( |
| 171 | + client1_requested, |
| 172 | + client2_requested, |
| 173 | + client1_completed, |
| 174 | + client2_completed |
| 175 | + ); |
105 | 176 |
|
106 | | - assert!(matches!(result, PermissionResult::Approved)); |
| 177 | + session2 |
| 178 | + .disconnect() |
| 179 | + .await |
| 180 | + .expect("disconnect second session"); |
| 181 | + client2.force_stop(); |
| 182 | + session1 |
| 183 | + .disconnect() |
| 184 | + .await |
| 185 | + .expect("disconnect first session"); |
| 186 | + server.stop().await.expect("stop server client"); |
| 187 | + }) |
| 188 | + }, |
| 189 | + ) |
| 190 | + .await; |
107 | 191 | } |
108 | 192 |
|
109 | 193 | #[tokio::test] |
110 | 194 | async fn one_client_rejects_permission_and_both_see_the_result() { |
111 | | - let result = PermissionResult::Denied; |
| 195 | + with_e2e_context( |
| 196 | + "rust_multi_client", |
| 197 | + "one_client_rejects_permission_and_both_see_the_result", |
| 198 | + |ctx| { |
| 199 | + Box::pin(async move { |
| 200 | + ctx.set_default_copilot_user(); |
| 201 | + let protected_file = ctx.work_dir().join("protected.txt"); |
| 202 | + std::fs::write(&protected_file, "protected content").expect("write protected file"); |
| 203 | + let port = free_tcp_port(); |
| 204 | + let server = start_tcp_server(ctx, port).await; |
| 205 | + let session1 = server |
| 206 | + .create_session( |
| 207 | + SessionConfig::default() |
| 208 | + .with_github_token(DEFAULT_TEST_TOKEN) |
| 209 | + .with_handler(permission_handler(PermissionResult::Denied)), |
| 210 | + ) |
| 211 | + .await |
| 212 | + .expect("create session"); |
| 213 | + let client2 = start_external_client(ctx, port).await; |
| 214 | + let session2 = client2 |
| 215 | + .resume_session( |
| 216 | + resume_config(session1.id().clone()) |
| 217 | + .with_request_permission(false) |
| 218 | + .with_handler(permission_handler(PermissionResult::NoResult)), |
| 219 | + ) |
| 220 | + .await |
| 221 | + .expect("resume session"); |
112 | 222 |
|
113 | | - assert!(matches!(result, PermissionResult::Denied)); |
| 223 | + let client1_requested = wait_for_event( |
| 224 | + session1.subscribe(), |
| 225 | + "client1 permission request", |
| 226 | + |event| event.parsed_type() == SessionEventType::PermissionRequested, |
| 227 | + ); |
| 228 | + let client2_requested = wait_for_event( |
| 229 | + session2.subscribe(), |
| 230 | + "client2 permission request", |
| 231 | + |event| event.parsed_type() == SessionEventType::PermissionRequested, |
| 232 | + ); |
| 233 | + let client1_completed = |
| 234 | + wait_for_event(session1.subscribe(), "client1 permission denied", |event| { |
| 235 | + is_permission_denied(event) |
| 236 | + }); |
| 237 | + let client2_completed = |
| 238 | + wait_for_event(session2.subscribe(), "client2 permission denied", |event| { |
| 239 | + is_permission_denied(event) |
| 240 | + }); |
| 241 | + |
| 242 | + session1 |
| 243 | + .send_and_wait("Edit protected.txt and replace 'protected' with 'hacked'.") |
| 244 | + .await |
| 245 | + .expect("send"); |
| 246 | + let content = |
| 247 | + std::fs::read_to_string(&protected_file).expect("read protected file"); |
| 248 | + assert_eq!(content, "protected content"); |
| 249 | + let _ = tokio::join!( |
| 250 | + client1_requested, |
| 251 | + client2_requested, |
| 252 | + client1_completed, |
| 253 | + client2_completed |
| 254 | + ); |
| 255 | + |
| 256 | + session2 |
| 257 | + .disconnect() |
| 258 | + .await |
| 259 | + .expect("disconnect second session"); |
| 260 | + client2.force_stop(); |
| 261 | + session1 |
| 262 | + .disconnect() |
| 263 | + .await |
| 264 | + .expect("disconnect first session"); |
| 265 | + server.stop().await.expect("stop server client"); |
| 266 | + }) |
| 267 | + }, |
| 268 | + ) |
| 269 | + .await; |
114 | 270 | } |
115 | 271 |
|
116 | 272 | #[tokio::test] |
@@ -300,6 +456,62 @@ fn selective_handler(tools: Vec<EchoTool>) -> Arc<SelectiveToolHandler> { |
300 | 456 | Arc::new(SelectiveToolHandler { tools }) |
301 | 457 | } |
302 | 458 |
|
| 459 | +fn permission_handler(result: PermissionResult) -> Arc<PermissionDecisionHandler> { |
| 460 | + Arc::new(PermissionDecisionHandler { |
| 461 | + result, |
| 462 | + request_count: None, |
| 463 | + }) |
| 464 | +} |
| 465 | + |
| 466 | +fn permission_handler_with_counter( |
| 467 | + result: PermissionResult, |
| 468 | + request_count: Arc<AtomicUsize>, |
| 469 | +) -> Arc<PermissionDecisionHandler> { |
| 470 | + Arc::new(PermissionDecisionHandler { |
| 471 | + result, |
| 472 | + request_count: Some(request_count), |
| 473 | + }) |
| 474 | +} |
| 475 | + |
| 476 | +fn is_permission_approved(event: &SessionEvent) -> bool { |
| 477 | + event.parsed_type() == SessionEventType::PermissionCompleted |
| 478 | + && event |
| 479 | + .typed_data::<PermissionCompletedData>() |
| 480 | + .is_some_and(|data| matches!(data.result, EventPermissionResult::Approved(_))) |
| 481 | +} |
| 482 | + |
| 483 | +fn is_permission_denied(event: &SessionEvent) -> bool { |
| 484 | + event.parsed_type() == SessionEventType::PermissionCompleted |
| 485 | + && event |
| 486 | + .typed_data::<PermissionCompletedData>() |
| 487 | + .is_some_and(|data| { |
| 488 | + matches!( |
| 489 | + data.result, |
| 490 | + EventPermissionResult::DeniedInteractivelyByUser(_) |
| 491 | + ) |
| 492 | + }) |
| 493 | +} |
| 494 | + |
| 495 | +struct PermissionDecisionHandler { |
| 496 | + result: PermissionResult, |
| 497 | + request_count: Option<Arc<AtomicUsize>>, |
| 498 | +} |
| 499 | + |
| 500 | +#[async_trait] |
| 501 | +impl SessionHandler for PermissionDecisionHandler { |
| 502 | + async fn on_permission_request( |
| 503 | + &self, |
| 504 | + _session_id: SessionId, |
| 505 | + _request_id: RequestId, |
| 506 | + _data: PermissionRequestData, |
| 507 | + ) -> PermissionResult { |
| 508 | + if let Some(request_count) = &self.request_count { |
| 509 | + request_count.fetch_add(1, Ordering::SeqCst); |
| 510 | + } |
| 511 | + self.result.clone() |
| 512 | + } |
| 513 | +} |
| 514 | + |
303 | 515 | struct SelectiveToolHandler { |
304 | 516 | tools: Vec<EchoTool>, |
305 | 517 | } |
|
0 commit comments