File size: 3,521 Bytes
17f328f | 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 | use super::*;
use codex_rmcp_client::InProcessTransportFactory;
use futures::FutureExt;
use futures::future::BoxFuture;
use pretty_assertions::assert_eq;
use rmcp::ServiceExt;
use rmcp::model::ClientCapabilities;
use rmcp::model::CustomNotification;
use rmcp::model::Implementation;
use rmcp::model::InitializeRequestParams;
use tokio::sync::mpsc;
use tokio::time::timeout;
#[derive(Clone)]
struct NotificationServer(mpsc::Sender<serde_json::Value>);
impl rmcp::ServerHandler for NotificationServer {
async fn on_custom_notification(
&self,
notification: CustomNotification,
_context: rmcp::service::NotificationContext<rmcp::RoleServer>,
) {
self.0.send(json!(notification)).await.unwrap();
}
}
impl InProcessTransportFactory for NotificationServer {
fn open(&self) -> BoxFuture<'static, std::io::Result<tokio::io::DuplexStream>> {
let server = self.clone();
async move {
let (client, transport) = tokio::io::duplex(/*max_buf_size*/ 4096);
tokio::spawn(async move {
let service = server.serve(transport).await.unwrap();
service.waiting().await.unwrap();
});
Ok(client)
}
.boxed()
}
}
#[tokio::test]
async fn auth_notifications_require_opt_in_and_follow_client_lifetime() -> Result<()> {
let (notifications, mut received) = mpsc::channel(/*buffer*/ 8);
let client = Arc::new(
RmcpClient::new_in_process_client(Arc::new(NotificationServer(notifications))).await?,
);
client
.initialize(
InitializeRequestParams::new(
ClientCapabilities::default(),
Implementation::new("test", "1"),
),
Some(SEND_TIMEOUT),
Box::new(|_, _| async { anyhow::bail!("unexpected elicitation") }.boxed()),
)
.await?;
let (changes, receiver) = watch::channel(AuthChangeState::default());
let mut capabilities = ServerCapabilities::default();
assert!(
start(Arc::clone(&client), &capabilities, Some(receiver.clone()))
.await?
.is_none()
);
capabilities.experimental = Some([(CAPABILITY.to_string(), Default::default())].into());
assert!(
start(Arc::clone(&client), &capabilities, /*changes*/ None)
.await?
.is_none()
);
assert_eq!(received.try_recv(), Err(mpsc::error::TryRecvError::Empty));
let watcher = start(Arc::clone(&client), &capabilities, Some(receiver))
.await?
.unwrap();
assert_eq!(
timeout(SEND_TIMEOUT, received.recv()).await?,
Some(
json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 0, "ownerGeneration": 0}})
),
);
changes.send_modify(|state| state.generation += 1);
assert_eq!(
timeout(SEND_TIMEOUT, received.recv()).await?,
Some(
json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 1, "ownerGeneration": 0}})
),
);
for _ in 0..2 {
changes.send_modify(|state| {
state.generation += 1;
state.owner_generation += 1;
});
}
assert_eq!(
timeout(SEND_TIMEOUT, received.recv()).await?,
Some(
json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 3, "ownerGeneration": 2}})
),
);
drop(watcher);
timeout(SEND_TIMEOUT, changes.closed()).await?;
client.shutdown().await;
Ok(())
}
|