1use agent_client_protocol_schema::v1::{
4 JsonRpcMessage as VersionedJsonRpcMessage, Notification as RpcNotification,
5 Request as RpcRequest, RequestId, SessionId,
6};
7
8use serde::ser::SerializeSeq as _;
10use serde::{Deserialize, Serialize};
11use std::any::TypeId;
12use std::collections::HashMap;
13use std::fmt::Debug;
14use std::marker::PhantomData;
15use std::panic::Location;
16use std::pin::pin;
17use std::sync::{
18 Arc, Mutex, Weak,
19 atomic::{AtomicBool, Ordering},
20};
21use uuid::Uuid;
22
23use futures::FutureExt;
24use futures::channel::{mpsc, oneshot};
25use futures::future::{self, BoxFuture, Either};
26use futures::{AsyncRead, AsyncWrite, StreamExt};
27
28pub(crate) mod close;
29mod dynamic_handler;
30pub(crate) mod handlers;
31mod incoming_actor;
32mod outgoing_actor;
33mod protocol_compat;
34mod raw_error;
35pub(crate) mod run;
36mod task_actor;
37mod transport_actor;
38
39use crate::jsonrpc::close::{ChainedClose, CloseCallback};
40pub use crate::jsonrpc::close::{HandleConnectionClose, NullClose};
41use crate::jsonrpc::dynamic_handler::DynamicHandlerMessage;
42pub use crate::jsonrpc::handlers::NullHandler;
43use crate::jsonrpc::handlers::{ChainedHandler, NamedHandler};
44use crate::jsonrpc::handlers::{MessageHandler, NotificationHandler, RequestHandler};
45use crate::jsonrpc::outgoing_actor::{OutgoingMessageTx, send_raw_message};
46use crate::jsonrpc::protocol_compat::{ProtocolCompat, ProtocolMode};
47pub use crate::jsonrpc::raw_error::{RawJsonRpcError, RawJsonRpcResponse};
48use crate::jsonrpc::run::SpawnedRun;
49use crate::jsonrpc::run::{ChainRun, NullRun, RunWithConnectionTo};
50use crate::jsonrpc::task_actor::{Task, TaskTx};
51#[cfg(feature = "unstable_mcp_over_acp")]
52use crate::mcp_server::McpServer;
53use crate::role::HasPeer;
54use crate::role::Role;
55use crate::{Agent, Client, ConnectTo, Proxy, RoleId};
56
57#[derive(Debug, Clone)]
63pub enum RawJsonRpcMessage {
64 Request(RpcRequest<RawJsonRpcParams>),
66 Notification(RpcNotification<RawJsonRpcParams>),
68 Response(RawJsonRpcResponse),
70}
71
72#[derive(Clone, Debug)]
79pub enum TransportFrame {
80 Single(RawJsonRpcMessage),
82 Malformed {
84 raw: String,
86 error: crate::Error,
88 },
89 Batch(TransportBatch),
91}
92
93#[derive(Clone, Debug)]
95pub struct TransportBatch {
96 first: TransportBatchEntry,
97 rest: Vec<TransportBatchEntry>,
98}
99
100#[derive(Clone, Debug)]
102pub enum TransportBatchEntry {
103 Message(RawJsonRpcMessage),
105 Malformed {
107 raw: serde_json::Value,
109 error: crate::Error,
111 },
112}
113
114pub(crate) fn is_response_only_shape(value: &serde_json::Value) -> bool {
115 value.as_object().is_some_and(|object| {
116 !object.contains_key("method")
117 && (object.contains_key("result") || object.contains_key("error"))
118 })
119}
120
121pub(crate) fn raw_is_response_only_shape(raw: &str) -> bool {
122 serde_json::from_str(raw).is_ok_and(|value| is_response_only_shape(&value))
123}
124
125impl TransportBatchEntry {
126 #[must_use]
128 pub fn message(message: RawJsonRpcMessage) -> Self {
129 Self::Message(message)
130 }
131
132 #[must_use]
134 pub fn malformed(raw: serde_json::Value, error: crate::Error) -> Self {
135 Self::Malformed { raw, error }
136 }
137
138 #[cfg(test)]
139 fn as_result(&self) -> Result<&RawJsonRpcMessage, &crate::Error> {
140 match self {
141 Self::Message(message) => Ok(message),
142 Self::Malformed { error, .. } => Err(error),
143 }
144 }
145
146 fn message_ref(&self) -> Option<&RawJsonRpcMessage> {
147 match self {
148 Self::Message(message) => Some(message),
149 Self::Malformed { .. } => None,
150 }
151 }
152}
153
154impl Serialize for TransportBatchEntry {
155 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
156 where
157 S: serde::Serializer,
158 {
159 match self {
160 Self::Message(message) => message.serialize(serializer),
161 Self::Malformed { raw, .. } => raw.serialize(serializer),
162 }
163 }
164}
165
166impl TransportBatch {
167 pub fn from_entries(entries: impl IntoIterator<Item = TransportBatchEntry>) -> Option<Self> {
171 let mut entries = entries.into_iter();
172 Some(Self {
173 first: entries.next()?,
174 rest: entries.collect(),
175 })
176 }
177
178 pub fn from_messages(messages: impl IntoIterator<Item = RawJsonRpcMessage>) -> Option<Self> {
182 Self::from_entries(messages.into_iter().map(TransportBatchEntry::message))
183 }
184
185 pub fn entries(&self) -> impl Iterator<Item = &TransportBatchEntry> {
187 std::iter::once(&self.first).chain(&self.rest)
188 }
189
190 pub fn entries_mut(&mut self) -> impl Iterator<Item = &mut TransportBatchEntry> {
192 std::iter::once(&mut self.first).chain(&mut self.rest)
193 }
194
195 pub fn into_entries(self) -> impl Iterator<Item = TransportBatchEntry> {
197 std::iter::once(self.first).chain(self.rest)
198 }
199
200 #[must_use]
202 pub fn len(&self) -> usize {
203 1 + self.rest.len()
204 }
205
206 #[must_use]
211 pub const fn is_empty(&self) -> bool {
212 false
213 }
214
215 #[cfg(test)]
216 pub(crate) fn iter_results(
217 &self,
218 ) -> impl Iterator<Item = Result<&RawJsonRpcMessage, &crate::Error>> {
219 self.entries().map(TransportBatchEntry::as_result)
220 }
221
222 fn messages(&self) -> impl Iterator<Item = &RawJsonRpcMessage> {
223 self.entries().filter_map(TransportBatchEntry::message_ref)
224 }
225}
226
227impl Serialize for TransportBatch {
228 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
229 where
230 S: serde::Serializer,
231 {
232 let mut sequence = serializer.serialize_seq(Some(1 + self.rest.len()))?;
233 sequence.serialize_element(&self.first)?;
234 for entry in &self.rest {
235 sequence.serialize_element(entry)?;
236 }
237 sequence.end()
238 }
239}
240
241impl TransportFrame {
242 fn inspect_messages(
243 &self,
244 observer: &mut impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error>,
245 ) -> Result<(), crate::Error> {
246 match self {
247 Self::Single(message) => observer(message),
248 Self::Malformed { .. } => Ok(()),
249 Self::Batch(batch) => {
250 for message in batch.messages() {
251 observer(message)?;
252 }
253 Ok(())
254 }
255 }
256 }
257}
258
259#[derive(Debug, Clone, PartialEq)]
263pub enum RawJsonRpcParams {
264 Array(Vec<serde_json::Value>),
266 Object(serde_json::Map<String, serde_json::Value>),
268}
269
270impl RawJsonRpcParams {
271 pub fn from_value(value: serde_json::Value) -> Result<Option<Self>, crate::Error> {
273 match value {
274 serde_json::Value::Null => Ok(None),
275 serde_json::Value::Array(array) => Ok(Some(Self::Array(array))),
276 serde_json::Value::Object(object) => Ok(Some(Self::Object(object))),
277 _ => {
278 Err(crate::Error::invalid_params()
279 .data("JSON-RPC params must be an object or array"))
280 }
281 }
282 }
283
284 #[must_use]
286 pub fn into_value(self) -> serde_json::Value {
287 match self {
288 Self::Array(array) => serde_json::Value::Array(array),
289 Self::Object(object) => serde_json::Value::Object(object),
290 }
291 }
292}
293
294impl Serialize for RawJsonRpcParams {
295 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
296 where
297 S: serde::Serializer,
298 {
299 match self {
300 Self::Array(array) => array.serialize(serializer),
301 Self::Object(object) => object.serialize(serializer),
302 }
303 }
304}
305
306impl<'de> Deserialize<'de> for RawJsonRpcParams {
307 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
308 where
309 D: serde::Deserializer<'de>,
310 {
311 let value = serde_json::Value::deserialize(deserializer)?;
312 match value {
313 serde_json::Value::Array(array) => Ok(Self::Array(array)),
314 serde_json::Value::Object(object) => Ok(Self::Object(object)),
315 _ => Err(serde::de::Error::custom(
316 "JSON-RPC params must be an object or array",
317 )),
318 }
319 }
320}
321
322impl RawJsonRpcMessage {
323 pub fn request(
325 method: String,
326 params: serde_json::Value,
327 id: RequestId,
328 ) -> Result<Self, crate::Error> {
329 Ok(Self::Request(RpcRequest {
330 id,
331 method: Arc::from(method),
332 params: RawJsonRpcParams::from_value(params)?,
333 }))
334 }
335
336 pub fn notification(method: String, params: serde_json::Value) -> Result<Self, crate::Error> {
338 Ok(Self::Notification(RpcNotification {
339 method: Arc::from(method),
340 params: RawJsonRpcParams::from_value(params)?,
341 }))
342 }
343
344 #[must_use]
349 pub fn response(id: RequestId, response: Result<serde_json::Value, crate::Error>) -> Self {
350 Self::Response(RawJsonRpcResponse::new(
351 id,
352 response.map_err(|error| Box::new(error.into())),
353 ))
354 }
355
356 #[must_use]
358 pub fn response_id(&self) -> Option<&RequestId> {
359 match self {
360 Self::Response(
361 RawJsonRpcResponse::Result { id, .. } | RawJsonRpcResponse::Error { id, .. },
362 ) => Some(id),
363 Self::Request(_) | Self::Notification(_) => None,
364 }
365 }
366}
367
368impl Serialize for RawJsonRpcMessage {
369 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
370 where
371 S: serde::Serializer,
372 {
373 match self {
374 Self::Request(request) => {
375 VersionedJsonRpcMessage::wrap(request.clone()).serialize(serializer)
376 }
377 Self::Notification(notification) => {
378 VersionedJsonRpcMessage::wrap(notification.clone()).serialize(serializer)
379 }
380 Self::Response(response) => {
381 VersionedJsonRpcMessage::wrap(response.clone()).serialize(serializer)
382 }
383 }
384 }
385}
386
387impl<'de> Deserialize<'de> for RawJsonRpcMessage {
388 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
389 where
390 D: serde::Deserializer<'de>,
391 {
392 let value = serde_json::Value::deserialize(deserializer)?;
393 let Some(object) = value.as_object() else {
394 return Err(serde::de::Error::custom("invalid JSON-RPC message"));
395 };
396
397 let has_method = object.contains_key("method");
398 let has_id = object.contains_key("id");
399 let has_result = object.contains_key("result");
400 let has_error = object.contains_key("error");
401
402 if has_method && !has_result && !has_error {
403 if has_id {
404 let request = serde_json::from_value::<
405 VersionedJsonRpcMessage<RpcRequest<RawJsonRpcParams>>,
406 >(value)
407 .map_err(serde::de::Error::custom)?
408 .into_inner();
409 Ok(Self::Request(request))
410 } else {
411 let notification = serde_json::from_value::<
412 VersionedJsonRpcMessage<RpcNotification<RawJsonRpcParams>>,
413 >(value)
414 .map_err(serde::de::Error::custom)?
415 .into_inner();
416 Ok(Self::Notification(notification))
417 }
418 } else if !has_method && has_id && has_result != has_error {
419 let response =
420 serde_json::from_value::<VersionedJsonRpcMessage<RawJsonRpcResponse>>(value)
421 .map_err(serde::de::Error::custom)?
422 .into_inner();
423 Ok(Self::Response(response))
424 } else {
425 Err(serde::de::Error::custom("invalid JSON-RPC message"))
426 }
427 }
428}
429
430fn params_from_transport(params: Option<RawJsonRpcParams>) -> serde_json::Value {
431 params.map_or(serde_json::Value::Null, RawJsonRpcParams::into_value)
432}
433
434#[allow(async_fn_in_trait)]
585pub trait HandleDispatchFrom<Counterpart: Role>: Send {
596 fn handle_dispatch_from(
617 &mut self,
618 message: Dispatch,
619 connection: ConnectionTo<Counterpart>,
620 ) -> impl Future<Output = Result<Handled<Dispatch>, crate::Error>> + Send;
621
622 fn describe_chain(&self) -> impl std::fmt::Debug;
624}
625
626impl<Counterpart: Role, H> HandleDispatchFrom<Counterpart> for &mut H
627where
628 H: HandleDispatchFrom<Counterpart>,
629{
630 fn handle_dispatch_from(
631 &mut self,
632 message: Dispatch,
633 cx: ConnectionTo<Counterpart>,
634 ) -> impl Future<Output = Result<Handled<Dispatch>, crate::Error>> + Send {
635 H::handle_dispatch_from(self, message, cx)
636 }
637
638 fn describe_chain(&self) -> impl std::fmt::Debug {
639 H::describe_chain(self)
640 }
641}
642
643#[doc(hidden)]
649#[allow(private_bounds)]
650pub trait ConnectionContext: connection_context::Sealed + Send + Sync + 'static {
651 type Connection<Counterpart: Role>: Clone + Send + Sync + 'static;
653}
654
655mod connection_context {
656 use super::{ConnectionContext, ConnectionTo, Role};
657
658 pub trait Sealed {
659 fn from_raw<Counterpart: Role>(
660 connection: ConnectionTo<Counterpart>,
661 ) -> <Self as ConnectionContext>::Connection<Counterpart>
662 where
663 Self: ConnectionContext;
664 }
665
666 pub(crate) fn from_raw<Context: ConnectionContext, Counterpart: Role>(
667 connection: ConnectionTo<Counterpart>,
668 ) -> Context::Connection<Counterpart> {
669 <Context as Sealed>::from_raw(connection)
670 }
671}
672
673#[doc(hidden)]
675#[derive(Copy, Clone, Debug, Default)]
676pub struct RawConnectionContext;
677
678impl connection_context::Sealed for RawConnectionContext {
679 fn from_raw<Counterpart: Role>(
680 connection: ConnectionTo<Counterpart>,
681 ) -> <Self as ConnectionContext>::Connection<Counterpart> {
682 connection
683 }
684}
685
686impl ConnectionContext for RawConnectionContext {
687 type Connection<Counterpart: Role> = ConnectionTo<Counterpart>;
688}
689
690#[cfg(feature = "unstable_protocol_v2")]
692#[doc(hidden)]
693#[derive(Copy, Clone, Debug, Default)]
694pub struct V2ConnectionContext;
695
696#[cfg(feature = "unstable_protocol_v2")]
697impl connection_context::Sealed for V2ConnectionContext {
698 fn from_raw<Counterpart: Role>(
699 connection: ConnectionTo<Counterpart>,
700 ) -> <Self as ConnectionContext>::Connection<Counterpart> {
701 V2ConnectionTo { inner: connection }
702 }
703}
704
705#[cfg(feature = "unstable_protocol_v2")]
706impl ConnectionContext for V2ConnectionContext {
707 type Connection<Counterpart: Role> = V2ConnectionTo<Counterpart>;
708}
709
710#[cfg(feature = "unstable_protocol_v2")]
712pub type V2Builder<Host, Handler = NullHandler, Runner = NullRun, Close = NullClose> =
713 Builder<Host, Handler, Runner, Close, V2ConnectionContext>;
714
715#[must_use]
1009#[derive(Debug)]
1010pub struct Builder<
1011 Host: Role,
1012 Handler = NullHandler,
1013 Runner = NullRun,
1014 Close = NullClose,
1015 Context = RawConnectionContext,
1016> where
1017 Handler: HandleDispatchFrom<Host::Counterpart>,
1018 Runner: RunWithConnectionTo<Host::Counterpart>,
1019 Close: HandleConnectionClose<Host::Counterpart>,
1020 Context: ConnectionContext,
1021{
1022 host: Host,
1024
1025 name: Option<String>,
1027
1028 handler: Handler,
1030
1031 runner: Runner,
1033
1034 protocol_mode: ProtocolMode,
1036
1037 on_close: Close,
1039
1040 context: PhantomData<fn() -> Context>,
1042}
1043
1044fn default_protocol_mode<Host: Role>() -> ProtocolMode {
1045 let role = TypeId::of::<Host>();
1046
1047 if role == TypeId::of::<Agent>() {
1048 ProtocolMode::v1_agent()
1049 } else if role == TypeId::of::<Client>() {
1050 ProtocolMode::v1_client()
1051 } else if role == TypeId::of::<Proxy>() {
1052 ProtocolMode::v1_proxy()
1053 } else {
1054 ProtocolMode::disabled()
1055 }
1056}
1057
1058impl<Host: Role> Builder<Host, NullHandler, NullRun, NullClose> {
1059 pub fn new(role: Host) -> Self {
1063 Self {
1064 host: role,
1065 name: None,
1066 handler: NullHandler,
1067 runner: NullRun,
1068 protocol_mode: default_protocol_mode::<Host>(),
1069 on_close: NullClose,
1070 context: PhantomData,
1071 }
1072 }
1073}
1074
1075impl<Host: Role, Handler> Builder<Host, Handler, NullRun, NullClose>
1076where
1077 Handler: HandleDispatchFrom<Host::Counterpart>,
1078{
1079 pub fn new_with(role: Host, handler: Handler) -> Self {
1081 Self {
1082 host: role,
1083 name: None,
1084 handler,
1085 runner: NullRun,
1086 protocol_mode: default_protocol_mode::<Host>(),
1087 on_close: NullClose,
1088 context: PhantomData,
1089 }
1090 }
1091}
1092
1093#[cfg(feature = "unstable_protocol_v2")]
1094impl<
1095 Host: Role,
1096 Handler: HandleDispatchFrom<Host::Counterpart>,
1097 Runner: RunWithConnectionTo<Host::Counterpart>,
1098 Close: HandleConnectionClose<Host::Counterpart>,
1099> Builder<Host, Handler, Runner, Close>
1100{
1101 pub(crate) fn v2_agent(self) -> V2Builder<Host, Handler, Runner, Close> {
1102 Builder {
1103 host: self.host,
1104 name: self.name,
1105 handler: self.handler,
1106 runner: self.runner,
1107 protocol_mode: ProtocolMode::v2_agent(),
1108 on_close: self.on_close,
1109 context: PhantomData,
1110 }
1111 }
1112
1113 pub(crate) fn v2_client(self) -> V2Builder<Host, Handler, Runner, Close> {
1114 Builder {
1115 host: self.host,
1116 name: self.name,
1117 handler: self.handler,
1118 runner: self.runner,
1119 protocol_mode: ProtocolMode::v2_client(),
1120 on_close: self.on_close,
1121 context: PhantomData,
1122 }
1123 }
1124
1125 pub(crate) fn v2_proxy(self) -> V2Builder<Host, Handler, Runner, Close> {
1126 Builder {
1127 host: self.host,
1128 name: self.name,
1129 handler: self.handler,
1130 runner: self.runner,
1131 protocol_mode: ProtocolMode::v2_proxy(),
1132 on_close: self.on_close,
1133 context: PhantomData,
1134 }
1135 }
1136
1137 pub fn without_acp_version_guard(mut self) -> Self {
1157 self.protocol_mode = ProtocolMode::disabled();
1158 self
1159 }
1160}
1161
1162#[cfg(feature = "unstable_protocol_v2")]
1163impl<
1164 Handler: HandleDispatchFrom<Agent>,
1165 Runner: RunWithConnectionTo<Agent>,
1166 Close: HandleConnectionClose<Agent>,
1167> Builder<Client, Handler, Runner, Close>
1168{
1169 pub fn with_v2_protocol_guard(mut self) -> Self {
1179 self.protocol_mode = ProtocolMode::v2_client();
1180 self
1181 }
1182}
1183
1184#[cfg(feature = "unstable_protocol_v2")]
1185impl<
1186 Handler: HandleDispatchFrom<Client>,
1187 Runner: RunWithConnectionTo<Client>,
1188 Close: HandleConnectionClose<Client>,
1189> Builder<Agent, Handler, Runner, Close>
1190{
1191 pub fn with_v2_protocol_guard(mut self) -> Self {
1201 self.protocol_mode = ProtocolMode::v2_agent();
1202 self
1203 }
1204}
1205
1206impl<
1207 Host: Role,
1208 Handler: HandleDispatchFrom<Host::Counterpart>,
1209 Runner: RunWithConnectionTo<Host::Counterpart>,
1210 Close: HandleConnectionClose<Host::Counterpart>,
1211 Context: ConnectionContext,
1212> Builder<Host, Handler, Runner, Close, Context>
1213{
1214 pub fn name(mut self, name: impl ToString) -> Self {
1216 self.name = Some(name.to_string());
1217 self
1218 }
1219
1220 pub(crate) fn v1_agent(mut self) -> Self {
1221 self.protocol_mode = ProtocolMode::v1_agent();
1222 self
1223 }
1224
1225 pub(crate) fn v1_client(mut self) -> Self {
1226 self.protocol_mode = ProtocolMode::v1_client();
1227 self
1228 }
1229
1230 pub fn with_connection_builder(
1235 self,
1236 other: Builder<
1237 Host,
1238 impl HandleDispatchFrom<Host::Counterpart>,
1239 impl RunWithConnectionTo<Host::Counterpart>,
1240 impl HandleConnectionClose<Host::Counterpart>,
1241 Context,
1242 >,
1243 ) -> Builder<
1244 Host,
1245 impl HandleDispatchFrom<Host::Counterpart>,
1246 impl RunWithConnectionTo<Host::Counterpart>,
1247 impl HandleConnectionClose<Host::Counterpart>,
1248 Context,
1249 > {
1250 let Builder {
1251 name: other_name,
1252 handler: other_handler,
1253 runner: other_runner,
1254 protocol_mode: other_protocol_mode,
1255 on_close: other_on_close,
1256 context: _,
1257 host: _,
1258 } = other;
1259 Builder {
1260 host: self.host,
1261 name: self.name,
1262 handler: ChainedHandler::new(
1263 self.handler,
1264 NamedHandler::new(other_name, other_handler),
1265 ),
1266 runner: ChainRun::new(self.runner, other_runner),
1267 protocol_mode: self.protocol_mode.merge(other_protocol_mode),
1268 on_close: ChainedClose::new(self.on_close, other_on_close),
1269 context: PhantomData,
1270 }
1271 }
1272
1273 pub fn with_handler(
1278 self,
1279 handler: impl HandleDispatchFrom<Host::Counterpart>,
1280 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context> {
1281 Builder {
1282 host: self.host,
1283 name: self.name,
1284 handler: ChainedHandler::new(self.handler, handler),
1285 runner: self.runner,
1286 protocol_mode: self.protocol_mode,
1287 on_close: self.on_close,
1288 context: PhantomData,
1289 }
1290 }
1291
1292 pub fn with_runner<Run1>(
1294 self,
1295 runner: Run1,
1296 ) -> Builder<Host, Handler, impl RunWithConnectionTo<Host::Counterpart>, Close, Context>
1297 where
1298 Run1: RunWithConnectionTo<Host::Counterpart>,
1299 {
1300 Builder {
1301 host: self.host,
1302 name: self.name,
1303 handler: self.handler,
1304 runner: ChainRun::new(self.runner, runner),
1305 protocol_mode: self.protocol_mode,
1306 on_close: self.on_close,
1307 context: PhantomData,
1308 }
1309 }
1310
1311 #[track_caller]
1313 pub fn with_spawned<F, Fut>(
1314 self,
1315 task: F,
1316 ) -> Builder<Host, Handler, impl RunWithConnectionTo<Host::Counterpart>, Close, Context>
1317 where
1318 F: FnOnce(Context::Connection<Host::Counterpart>) -> Fut + Send,
1319 Fut: Future<Output = Result<(), crate::Error>> + Send,
1320 {
1321 let location = Location::caller();
1322 self.with_runner(SpawnedRun::<_, Context>::new(location, task))
1323 }
1324
1325 pub fn on_close<F, Fut>(
1359 self,
1360 callback: F,
1361 ) -> Builder<Host, Handler, Runner, impl HandleConnectionClose<Host::Counterpart>, Context>
1362 where
1363 F: FnOnce(Context::Connection<Host::Counterpart>) -> Fut + Send,
1364 Fut: Future<Output = Result<(), crate::Error>> + Send,
1365 {
1366 Builder {
1367 host: self.host,
1368 name: self.name,
1369 handler: self.handler,
1370 runner: self.runner,
1371 protocol_mode: self.protocol_mode,
1372 on_close: ChainedClose::new(self.on_close, CloseCallback::<_, Context>::new(callback)),
1373 context: PhantomData,
1374 }
1375 }
1376
1377 pub fn on_receive_dispatch<Req, Notif, F, T, ToFut>(
1424 self,
1425 op: F,
1426 to_future_hack: ToFut,
1427 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context>
1428 where
1429 Host::Counterpart: HasPeer<Host::Counterpart>,
1430 Req: JsonRpcRequest,
1431 Notif: JsonRpcNotification,
1432 F: AsyncFnMut(
1433 Dispatch<Req, Notif>,
1434 Context::Connection<Host::Counterpart>,
1435 ) -> Result<T, crate::Error>
1436 + Send,
1437 T: IntoHandled<Dispatch<Req, Notif>>,
1438 ToFut: Fn(
1439 &mut F,
1440 Dispatch<Req, Notif>,
1441 Context::Connection<Host::Counterpart>,
1442 ) -> crate::BoxFuture<'_, Result<T, crate::Error>>
1443 + Send
1444 + Sync,
1445 {
1446 let handler = MessageHandler::<_, _, _, _, _, _, Context>::new(
1447 self.host.counterpart(),
1448 self.host.counterpart(),
1449 op,
1450 to_future_hack,
1451 );
1452 self.with_handler(handler)
1453 }
1454
1455 pub fn on_receive_request<Req: JsonRpcRequest, F, T, ToFut>(
1501 self,
1502 op: F,
1503 to_future_hack: ToFut,
1504 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context>
1505 where
1506 Host::Counterpart: HasPeer<Host::Counterpart>,
1507 F: AsyncFnMut(
1508 Req,
1509 Responder<Req::Response>,
1510 Context::Connection<Host::Counterpart>,
1511 ) -> Result<T, crate::Error>
1512 + Send,
1513 T: IntoHandled<(Req, Responder<Req::Response>)>,
1514 ToFut: Fn(
1515 &mut F,
1516 Req,
1517 Responder<Req::Response>,
1518 Context::Connection<Host::Counterpart>,
1519 ) -> crate::BoxFuture<'_, Result<T, crate::Error>>
1520 + Send
1521 + Sync,
1522 {
1523 let handler = RequestHandler::<_, _, _, _, _, Context>::new(
1524 self.host.counterpart(),
1525 self.host.counterpart(),
1526 op,
1527 to_future_hack,
1528 );
1529 self.with_handler(handler)
1530 }
1531
1532 pub fn on_receive_notification<Notif, F, T, ToFut>(
1576 self,
1577 op: F,
1578 to_future_hack: ToFut,
1579 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context>
1580 where
1581 Host::Counterpart: HasPeer<Host::Counterpart>,
1582 Notif: JsonRpcNotification,
1583 F: AsyncFnMut(Notif, Context::Connection<Host::Counterpart>) -> Result<T, crate::Error>
1584 + Send,
1585 T: IntoHandled<(Notif, Context::Connection<Host::Counterpart>)>,
1586 ToFut: Fn(
1587 &mut F,
1588 Notif,
1589 Context::Connection<Host::Counterpart>,
1590 ) -> crate::BoxFuture<'_, Result<T, crate::Error>>
1591 + Send
1592 + Sync,
1593 {
1594 let handler = NotificationHandler::<_, _, _, _, _, Context>::new(
1595 self.host.counterpart(),
1596 self.host.counterpart(),
1597 op,
1598 to_future_hack,
1599 );
1600 self.with_handler(handler)
1601 }
1602
1603 pub fn on_receive_dispatch_from<
1619 Req: JsonRpcRequest,
1620 Notif: JsonRpcNotification,
1621 Peer: Role,
1622 F,
1623 T,
1624 ToFut,
1625 >(
1626 self,
1627 peer: Peer,
1628 op: F,
1629 to_future_hack: ToFut,
1630 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context>
1631 where
1632 Host::Counterpart: HasPeer<Peer>,
1633 F: AsyncFnMut(
1634 Dispatch<Req, Notif>,
1635 Context::Connection<Host::Counterpart>,
1636 ) -> Result<T, crate::Error>
1637 + Send,
1638 T: IntoHandled<Dispatch<Req, Notif>>,
1639 ToFut: Fn(
1640 &mut F,
1641 Dispatch<Req, Notif>,
1642 Context::Connection<Host::Counterpart>,
1643 ) -> crate::BoxFuture<'_, Result<T, crate::Error>>
1644 + Send
1645 + Sync,
1646 {
1647 let handler = MessageHandler::<_, _, _, _, _, _, Context>::new(
1648 self.host.counterpart(),
1649 peer,
1650 op,
1651 to_future_hack,
1652 );
1653 self.with_handler(handler)
1654 }
1655
1656 pub fn on_receive_request_from<Req: JsonRpcRequest, Peer: Role, F, T, ToFut>(
1685 self,
1686 peer: Peer,
1687 op: F,
1688 to_future_hack: ToFut,
1689 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context>
1690 where
1691 Host::Counterpart: HasPeer<Peer>,
1692 F: AsyncFnMut(
1693 Req,
1694 Responder<Req::Response>,
1695 Context::Connection<Host::Counterpart>,
1696 ) -> Result<T, crate::Error>
1697 + Send,
1698 T: IntoHandled<(Req, Responder<Req::Response>)>,
1699 ToFut: Fn(
1700 &mut F,
1701 Req,
1702 Responder<Req::Response>,
1703 Context::Connection<Host::Counterpart>,
1704 ) -> crate::BoxFuture<'_, Result<T, crate::Error>>
1705 + Send
1706 + Sync,
1707 {
1708 let handler = RequestHandler::<_, _, _, _, _, Context>::new(
1709 self.host.counterpart(),
1710 peer,
1711 op,
1712 to_future_hack,
1713 );
1714 self.with_handler(handler)
1715 }
1716
1717 pub fn on_receive_notification_from<Notif: JsonRpcNotification, Peer: Role, F, T, ToFut>(
1733 self,
1734 peer: Peer,
1735 op: F,
1736 to_future_hack: ToFut,
1737 ) -> Builder<Host, impl HandleDispatchFrom<Host::Counterpart>, Runner, Close, Context>
1738 where
1739 Host::Counterpart: HasPeer<Peer>,
1740 F: AsyncFnMut(Notif, Context::Connection<Host::Counterpart>) -> Result<T, crate::Error>
1741 + Send,
1742 T: IntoHandled<(Notif, Context::Connection<Host::Counterpart>)>,
1743 ToFut: Fn(
1744 &mut F,
1745 Notif,
1746 Context::Connection<Host::Counterpart>,
1747 ) -> crate::BoxFuture<'_, Result<T, crate::Error>>
1748 + Send
1749 + Sync,
1750 {
1751 let handler = NotificationHandler::<_, _, _, _, _, Context>::new(
1752 self.host.counterpart(),
1753 peer,
1754 op,
1755 to_future_hack,
1756 );
1757 self.with_handler(handler)
1758 }
1759
1760 pub async fn connect_to(
1800 self,
1801 transport: impl ConnectTo<Host> + 'static,
1802 ) -> Result<(), crate::Error> {
1803 let (_, future) = self.into_connection_and_future(transport, true, async move |cx| {
1804 cx.incoming_closed().await;
1805 Ok(())
1806 });
1807 future.await
1808 }
1809
1810 pub async fn connect_with<R>(
1874 self,
1875 transport: impl ConnectTo<Host> + 'static,
1876 main_fn: impl AsyncFnOnce(Context::Connection<Host::Counterpart>) -> Result<R, crate::Error>,
1877 ) -> Result<R, crate::Error> {
1878 let (_, future) =
1879 self.into_connection_and_future(transport, false, async move |connection| {
1880 main_fn(connection_context::from_raw::<Context, _>(connection)).await
1881 });
1882 future.await
1883 }
1884
1885 fn into_connection_and_future<R>(
1887 self,
1888 transport: impl ConnectTo<Host> + 'static,
1889 wait_owned_transport: bool,
1890 main_fn: impl AsyncFnOnce(ConnectionTo<Host::Counterpart>) -> Result<R, crate::Error>,
1891 ) -> (
1892 ConnectionTo<Host::Counterpart>,
1893 impl Future<Output = Result<R, crate::Error>>,
1894 ) {
1895 let Self {
1896 name,
1897 handler,
1898 runner,
1899 host: me,
1900 protocol_mode,
1901 on_close,
1902 context: _,
1903 } = self;
1904
1905 let (outgoing_tx, outgoing_rx) = mpsc::unbounded();
1906 let (new_task_tx, new_task_rx) = mpsc::unbounded();
1907 let (dynamic_handler_tx, dynamic_handler_rx) = mpsc::unbounded();
1908 let (foreground_succeeded_tx, foreground_succeeded) = completion_signal();
1909 let (foreground_done_tx, foreground_done) = completion_signal();
1910 let pending_replies = PendingReplies::default();
1911
1912 let transport_component = crate::DynConnectTo::new(transport);
1914 let (transport_channel, mut transport_future) =
1915 transport_component.into_channel_and_future();
1916 let owned_transport = transport_future.is_some();
1917 let transport_finish = transport_future
1918 .as_mut()
1919 .and_then(crate::ConnectionDriver::take_finish);
1920 let (transport_completion_tx, transport_completion_rx) = oneshot::channel();
1921 let transport_completion = transport_completion_rx
1922 .map(|result| {
1923 result.unwrap_or_else(|error| {
1924 Err(crate::util::internal_error(format!(
1925 "transport task dropped before reporting completion: {error}"
1926 )))
1927 })
1928 })
1929 .boxed()
1930 .shared();
1931
1932 let connection = ConnectionTo::new(
1933 me.counterpart(),
1934 outgoing_tx,
1935 new_task_tx,
1936 dynamic_handler_tx,
1937 transport_completion,
1938 pending_replies.registrar(),
1939 protocol_mode,
1940 );
1941 let transport_driver = if let Some(driver) = transport_future {
1945 async move {
1946 let result = driver.await;
1947 drop(transport_completion_tx.send(result.clone()));
1948 result
1949 }
1950 .boxed()
1951 } else {
1952 drop(transport_completion_tx.send(Ok(())));
1955 future::ready(Ok(())).boxed()
1956 };
1957
1958 let Channel {
1960 rx: mut transport_incoming_rx,
1961 tx: transport_outgoing_tx,
1962 } = transport_channel;
1963
1964 let transport_incoming = futures::stream::poll_fn({
1965 let mut completion = connection.transport_completion.clone();
1966 let mut completed = false;
1967 move |cx| {
1968 if owned_transport
1969 && !completed
1970 && let std::task::Poll::Ready(Ok(())) =
1971 std::pin::Pin::new(&mut completion).poll(cx)
1972 {
1973 transport_incoming_rx.close();
1977 completed = true;
1978 }
1979 transport_incoming_rx.poll_next_unpin(cx)
1980 }
1981 });
1982
1983 let protocol_compat = ProtocolCompat::new(protocol_mode);
1984
1985 let future = crate::util::instrument_with_connection_name(name, {
1986 let connection = connection.clone();
1987 async move {
1988 let background = async {
1989 let incoming = {
1990 let pending_replies = pending_replies.clone();
1991 let protocol_compat = protocol_compat.clone();
1992 async {
1993 let mut transport_incoming = std::pin::pin!(transport_incoming);
1994 let incoming = incoming_actor::incoming_protocol_actor(
1995 me.counterpart(),
1996 &connection,
1997 transport_incoming.as_mut(),
1998 dynamic_handler_rx,
1999 pending_replies,
2000 incoming_actor::IncomingHandlers::new(
2001 handler,
2002 on_close,
2003 foreground_succeeded.clone(),
2004 ),
2005 protocol_compat,
2006 );
2007 let result = run_incoming_until_foreground_succeeds(
2010 incoming,
2011 foreground_succeeded,
2012 connection.incoming_closed.clone(),
2013 )
2014 .await;
2015 if result.is_err() {
2016 connection.request_shutdown();
2017 }
2018 result?;
2019 while transport_incoming.next().await.is_some() {}
2023 Ok(())
2024 }
2025 };
2026 let other_actors = async {
2027 let result = futures::try_join!(
2028 transport_driver,
2031 outgoing_actor::outgoing_protocol_actor(
2033 outgoing_rx,
2034 pending_replies,
2035 transport_outgoing_tx,
2036 protocol_compat,
2037 foreground_done,
2038 ),
2039 );
2040 if result.is_err() {
2042 connection.request_shutdown();
2043 }
2044 result?;
2045 Ok(())
2046 };
2047
2048 run_until_connection_close(
2052 incoming,
2053 other_actors,
2054 connection.incoming_closed.clone(),
2055 )
2056 .await
2057 };
2058
2059 run_until_connection_close(
2060 finish_actor_error(background, &connection),
2061 async {
2062 let application = async {
2063 futures::try_join!(
2064 finish_actor_error(
2065 task_actor::task_actor(new_task_rx, &connection),
2066 &connection,
2067 ),
2068 finish_actor_error(
2069 runner.run_with_connection_to(connection.clone()),
2070 &connection,
2071 ),
2072 )?;
2073 Ok(())
2074 };
2075 let result = run_until_connection_close(
2076 application,
2077 async {
2078 let result = main_fn(connection.clone()).await;
2079 connection.request_shutdown();
2080 if result.is_ok() {
2081 let _ = foreground_succeeded_tx.send(());
2084 }
2085 connection.pending_replies.disarm_cancellations();
2089 connection.wait_protected_operations().await;
2093 result
2094 },
2095 connection.incoming_closed.clone(),
2096 )
2097 .await?;
2098 let _ = foreground_done_tx.send(());
2099 connection
2100 .drain_outgoing(transport_finish, wait_owned_transport)
2101 .await?;
2102 Ok(result)
2103 },
2104 connection.incoming_closed.clone(),
2105 )
2106 .await
2107 }
2108 });
2109
2110 (connection, future)
2111 }
2112}
2113
2114async fn finish_actor_error<R: Role>(
2117 actor: impl Future<Output = Result<(), crate::Error>>,
2118 connection: &ConnectionTo<R>,
2119) -> Result<(), crate::Error> {
2120 let result = actor.await;
2121 if result.is_err() {
2122 connection.request_shutdown();
2123 if connection.incoming_closed.is_closing() {
2124 connection.incoming_closed.closed().await;
2125 }
2126 connection.wait_protected_operations().await;
2127 }
2128 result
2129}
2130
2131#[cfg(feature = "unstable_mcp_over_acp")]
2132impl<
2133 Host: Role,
2134 Handler: HandleDispatchFrom<Host::Counterpart>,
2135 Runner: RunWithConnectionTo<Host::Counterpart>,
2136 Close: HandleConnectionClose<Host::Counterpart>,
2137> Builder<Host, Handler, Runner, Close, RawConnectionContext>
2138{
2139 pub fn with_mcp_server(
2148 self,
2149 mcp_server: McpServer<Host::Counterpart, impl RunWithConnectionTo<Host::Counterpart>>,
2150 ) -> Builder<
2151 Host,
2152 impl HandleDispatchFrom<Host::Counterpart>,
2153 impl RunWithConnectionTo<Host::Counterpart>,
2154 Close,
2155 RawConnectionContext,
2156 >
2157 where
2158 Host::Counterpart: HasPeer<Agent> + HasPeer<Client>,
2159 {
2160 let (handler, runner) = mcp_server.into_handler_and_runner();
2161 self.with_handler(handler).with_runner(runner)
2162 }
2163}
2164
2165#[cfg(all(feature = "unstable_mcp_over_acp", feature = "unstable_protocol_v2"))]
2166impl<
2167 Host: Role,
2168 Handler: HandleDispatchFrom<Host::Counterpart>,
2169 Runner: RunWithConnectionTo<Host::Counterpart>,
2170 Close: HandleConnectionClose<Host::Counterpart>,
2171> Builder<Host, Handler, Runner, Close, V2ConnectionContext>
2172{
2173 pub fn with_mcp_server(
2182 self,
2183 mcp_server: McpServer<Host::Counterpart, impl RunWithConnectionTo<Host::Counterpart>>,
2184 ) -> Builder<
2185 Host,
2186 impl HandleDispatchFrom<Host::Counterpart>,
2187 impl RunWithConnectionTo<Host::Counterpart>,
2188 Close,
2189 V2ConnectionContext,
2190 >
2191 where
2192 Host::Counterpart: HasPeer<Agent> + HasPeer<Client>,
2193 {
2194 let (handler, runner) = mcp_server.into_v2_handler_and_runner();
2195 self.with_handler(handler).with_runner(runner)
2196 }
2197}
2198
2199impl<R, H, Run, Close, Context> ConnectTo<R::Counterpart> for Builder<R, H, Run, Close, Context>
2200where
2201 R: Role,
2202 H: HandleDispatchFrom<R::Counterpart> + 'static,
2203 Run: RunWithConnectionTo<R::Counterpart> + 'static,
2204 Close: HandleConnectionClose<R::Counterpart> + 'static,
2205 Context: ConnectionContext,
2206{
2207 async fn connect_to(self, client: impl ConnectTo<R>) -> Result<(), crate::Error> {
2208 Builder::connect_to(self, client).await
2209 }
2210}
2211
2212pub(crate) struct ResponsePayload {
2217 pub(crate) result: Result<serde_json::Value, crate::Error>,
2219
2220 pub(crate) ack_tx: Option<oneshot::Sender<()>>,
2232}
2233
2234type ResponseRouteHook =
2235 Box<dyn FnOnce(&str, &serde_json::Value) -> Result<(), crate::Error> + Send>;
2236
2237struct RequestReadiness {
2240 future: BoxFuture<'static, Result<(), crate::Error>>,
2241}
2242
2243impl RequestReadiness {
2244 fn new(future: impl Future<Output = Result<(), crate::Error>> + Send + 'static) -> Self {
2245 Self {
2246 future: future.boxed(),
2247 }
2248 }
2249}
2250
2251impl Future for RequestReadiness {
2252 type Output = Result<(), crate::Error>;
2253
2254 fn poll(
2255 mut self: std::pin::Pin<&mut Self>,
2256 cx: &mut std::task::Context<'_>,
2257 ) -> std::task::Poll<Self::Output> {
2258 self.future.as_mut().poll(cx)
2259 }
2260}
2261
2262impl Debug for RequestReadiness {
2263 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2264 formatter
2265 .debug_struct("RequestReadiness")
2266 .finish_non_exhaustive()
2267 }
2268}
2269
2270impl std::fmt::Debug for ResponsePayload {
2271 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2272 f.debug_struct("ResponsePayload")
2273 .field("result", &self.result)
2274 .field("ack_tx", &self.ack_tx.as_ref().map(|_| "..."))
2275 .finish()
2276 }
2277}
2278
2279#[derive(Clone, Debug, Default)]
2280struct ResponseOrdering {
2281 ordered: Arc<AtomicBool>,
2282}
2283
2284impl ResponseOrdering {
2285 fn mark_ordered(&self) {
2286 self.ordered.store(true, Ordering::Release);
2287 }
2288
2289 fn is_ordered(&self) -> bool {
2290 self.ordered.load(Ordering::Acquire)
2291 }
2292}
2293
2294struct PendingReply {
2295 method: String,
2296 role_id: RoleId,
2297 sender: oneshot::Sender<ResponsePayload>,
2298 cancellation_disarm: SentRequestCancellationDisarm,
2299 ordering: ResponseOrdering,
2300 response_route_hook: Option<ResponseRouteHook>,
2301}
2302
2303impl PendingReply {
2304 fn fail(self, error: crate::Error) {
2305 self.cancellation_disarm.disarm();
2306 if self
2307 .sender
2308 .send(ResponsePayload {
2309 result: Err(error),
2310 ack_tx: None,
2311 })
2312 .is_err()
2313 {
2314 tracing::trace!(method = %self.method, "Pending request was already dropped");
2315 }
2316 }
2317
2318 fn fail_incoming_closed(self) {
2319 let error = incoming_transport_closed_error(&self.method);
2320 self.fail(error);
2321 }
2322}
2323
2324#[derive(Default)]
2325struct PendingRepliesInner {
2326 incoming_closed: bool,
2327 replies: HashMap<RequestId, PendingReply>,
2328}
2329
2330#[derive(Clone, Default)]
2331struct PendingReplies {
2332 inner: Arc<Mutex<PendingRepliesInner>>,
2333}
2334
2335impl PendingReplies {
2336 fn registrar(&self) -> PendingRepliesRegistrar {
2337 PendingRepliesRegistrar {
2338 inner: Arc::downgrade(&self.inner),
2339 }
2340 }
2341
2342 fn contains(&self, id: &RequestId) -> bool {
2343 self.inner
2344 .lock()
2345 .expect("pending replies mutex poisoned")
2346 .replies
2347 .contains_key(id)
2348 }
2349
2350 fn remove(&self, id: &RequestId) -> Option<PendingReply> {
2351 self.inner
2352 .lock()
2353 .expect("pending replies mutex poisoned")
2354 .replies
2355 .remove(id)
2356 }
2357
2358 fn close_incoming(&self) -> usize {
2360 let replies = {
2361 let mut inner = self.inner.lock().expect("pending replies mutex poisoned");
2362 inner.incoming_closed = true;
2363 std::mem::take(&mut inner.replies)
2364 };
2365 let count = replies.len();
2366 for (_, reply) in replies {
2367 reply.fail_incoming_closed();
2368 }
2369 count
2370 }
2371}
2372
2373#[derive(Clone)]
2377struct PendingRepliesRegistrar {
2378 inner: Weak<Mutex<PendingRepliesInner>>,
2379}
2380
2381impl PendingRepliesRegistrar {
2382 fn disarm_cancellations(&self) {
2383 if let Some(inner) = self.inner.upgrade() {
2384 let inner = inner.lock().expect("pending replies mutex poisoned");
2385 for reply in inner.replies.values() {
2386 reply.cancellation_disarm.disarm();
2387 }
2388 }
2389 }
2390
2391 fn subscribe(
2396 &self,
2397 id: RequestId,
2398 reply: PendingReply,
2399 incoming_closed: &IncomingClosed,
2400 ) -> bool {
2401 let Some(inner) = self.inner.upgrade() else {
2402 if incoming_closed.is_closing() {
2403 reply.fail_incoming_closed();
2404 } else {
2405 let method = reply.method.clone();
2406 reply.fail(crate::util::internal_error(format!(
2407 "failed to send outgoing request `{method}`: connection is no longer running"
2408 )));
2409 }
2410 return false;
2411 };
2412
2413 let result = {
2414 let mut inner = inner.lock().expect("pending replies mutex poisoned");
2415 if inner.incoming_closed {
2416 Err(reply)
2417 } else {
2418 Ok(inner.replies.insert(id, reply))
2419 }
2420 };
2421
2422 match result {
2423 Err(rejected) => {
2424 rejected.fail_incoming_closed();
2425 false
2426 }
2427 Ok(replaced) => {
2428 if let Some(replaced) = replaced {
2429 replaced.fail(
2430 crate::Error::internal_error()
2431 .data("outgoing request ID was reused before its response arrived"),
2432 );
2433 }
2434 true
2435 }
2436 }
2437 }
2438
2439 fn remove(&self, id: &RequestId) -> Option<PendingReply> {
2440 self.inner
2441 .upgrade()?
2442 .lock()
2443 .expect("pending replies mutex poisoned")
2444 .replies
2445 .remove(id)
2446 }
2447}
2448
2449impl Debug for PendingRepliesRegistrar {
2450 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2451 formatter
2452 .debug_struct("PendingRepliesRegistrar")
2453 .field("is_connected", &(self.inner.strong_count() > 0))
2454 .finish()
2455 }
2456}
2457
2458#[derive(Clone)]
2464pub struct RequestCancellation {
2465 state: Arc<RequestCancellationState>,
2466}
2467
2468struct RequestCancellationState {
2469 cancelled: AtomicBool,
2470 signal_tx: Mutex<Option<oneshot::Sender<()>>>,
2471 signal_rx: future::Shared<BoxFuture<'static, ()>>,
2472}
2473
2474impl RequestCancellation {
2475 fn new() -> Self {
2476 let (signal_tx, signal_rx) = oneshot::channel();
2477 let signal_rx = signal_rx.map(|_| ()).boxed().shared();
2478 Self {
2479 state: Arc::new(RequestCancellationState {
2480 cancelled: AtomicBool::new(false),
2481 signal_tx: Mutex::new(Some(signal_tx)),
2482 signal_rx,
2483 }),
2484 }
2485 }
2486
2487 pub async fn cancelled(&self) {
2491 self.state.signal_rx.clone().await;
2492 }
2493
2494 pub async fn run_until_cancelled<T>(
2509 &self,
2510 future: impl std::future::Future<Output = Result<T, crate::Error>>,
2511 ) -> Result<T, crate::Error> {
2512 if self.is_cancelled() {
2513 return Err(crate::Error::request_cancelled());
2514 }
2515
2516 match future::select(pin!(future), pin!(self.cancelled())).await {
2517 Either::Left((result, _)) => result,
2518 Either::Right(((), _)) => Err(crate::Error::request_cancelled()),
2519 }
2520 }
2521
2522 #[must_use]
2524 pub fn is_cancelled(&self) -> bool {
2525 self.state.cancelled.load(Ordering::Acquire)
2526 }
2527
2528 fn cancel(&self) {
2529 if self.state.cancelled.swap(true, Ordering::AcqRel) {
2530 return;
2531 }
2532
2533 let signal_tx = self
2534 .state
2535 .signal_tx
2536 .lock()
2537 .expect("request cancellation signal mutex poisoned")
2538 .take();
2539
2540 if let Some(signal_tx) = signal_tx {
2543 let _ = signal_tx.send(());
2544 }
2545 }
2546}
2547
2548impl Debug for RequestCancellation {
2549 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2550 formatter
2551 .debug_struct("RequestCancellation")
2552 .field("is_cancelled", &self.is_cancelled())
2553 .finish_non_exhaustive()
2554 }
2555}
2556
2557#[derive(Debug)]
2564enum RequestCancellationEntry {
2565 Armed,
2567 Cancelled,
2569 Marker(RequestCancellation),
2571}
2572
2573#[derive(Debug)]
2581struct RequestCancellationSlot {
2582 generation: u64,
2583 entry: RequestCancellationEntry,
2584}
2585
2586#[derive(Debug, Default)]
2587struct RequestCancellationRegistryInner {
2588 slots: HashMap<RequestId, RequestCancellationSlot>,
2589 next_generation: u64,
2590}
2591
2592#[derive(Clone, Debug, Default)]
2593struct RequestCancellationRegistry {
2594 inner: Arc<Mutex<RequestCancellationRegistryInner>>,
2595}
2596
2597#[derive(Debug)]
2598struct ResponderCancellation {
2599 id: RequestId,
2600 generation: u64,
2601 registry: RequestCancellationRegistry,
2602}
2603
2604impl RequestCancellationRegistry {
2605 fn new() -> Self {
2606 Self::default()
2607 }
2608
2609 fn register(&self, id: &RequestId) -> ResponderCancellation {
2610 let generation = {
2611 let mut inner = self
2612 .inner
2613 .lock()
2614 .expect("request cancellation registry mutex poisoned");
2615 let generation = inner.next_generation;
2616 inner.next_generation += 1;
2617 if inner
2618 .slots
2619 .insert(
2620 id.clone(),
2621 RequestCancellationSlot {
2622 generation,
2623 entry: RequestCancellationEntry::Armed,
2624 },
2625 )
2626 .is_some()
2627 {
2628 tracing::debug!(
2629 ?id,
2630 "peer reused the ID of a request that is still in flight"
2631 );
2632 }
2633 generation
2634 };
2635 ResponderCancellation {
2636 id: id.clone(),
2637 generation,
2638 registry: self.clone(),
2639 }
2640 }
2641
2642 fn marker(&self, id: &RequestId, generation: u64) -> RequestCancellation {
2651 let mut inner = self
2652 .inner
2653 .lock()
2654 .expect("request cancellation registry mutex poisoned");
2655 let Some(slot) = inner.slots.get_mut(id) else {
2656 return RequestCancellation::new();
2661 };
2662 if slot.generation != generation {
2663 return RequestCancellation::new();
2668 }
2669 let entry = &mut slot.entry;
2670 match entry {
2671 RequestCancellationEntry::Marker(marker) => marker.clone(),
2672 RequestCancellationEntry::Armed => {
2673 let marker = RequestCancellation::new();
2674 *entry = RequestCancellationEntry::Marker(marker.clone());
2675 marker
2676 }
2677 RequestCancellationEntry::Cancelled => {
2678 let marker = RequestCancellation::new();
2681 marker.cancel();
2682 *entry = RequestCancellationEntry::Marker(marker.clone());
2683 marker
2684 }
2685 }
2686 }
2687
2688 fn cancel_if_requested(&self, dispatch: &Dispatch) -> Result<bool, crate::Error> {
2689 let Some(request_id) = cancellation_request_id(dispatch)? else {
2690 return Ok(false);
2691 };
2692 Ok(self.cancel(&request_id))
2693 }
2694
2695 fn cancel(&self, request_id: &RequestId) -> bool {
2697 let marker = {
2698 let mut inner = self
2699 .inner
2700 .lock()
2701 .expect("request cancellation registry mutex poisoned");
2702 let Some(slot) = inner.slots.get_mut(request_id) else {
2703 return false;
2704 };
2705 let entry = &mut slot.entry;
2706 match entry {
2707 RequestCancellationEntry::Marker(marker) => marker.clone(),
2708 RequestCancellationEntry::Cancelled => return true,
2709 RequestCancellationEntry::Armed => {
2710 *entry = RequestCancellationEntry::Cancelled;
2711 return true;
2712 }
2713 }
2714 };
2715
2716 marker.cancel();
2719 true
2720 }
2721
2722 fn remove(&self, request_id: &RequestId, generation: u64) {
2725 let mut inner = self
2726 .inner
2727 .lock()
2728 .expect("request cancellation registry mutex poisoned");
2729 if inner
2730 .slots
2731 .get(request_id)
2732 .is_some_and(|slot| slot.generation == generation)
2733 {
2734 inner.slots.remove(request_id);
2735 }
2736 }
2737}
2738
2739impl ResponderCancellation {
2740 fn cancellation(&self) -> RequestCancellation {
2741 self.registry.marker(&self.id, self.generation)
2742 }
2743}
2744
2745impl Drop for ResponderCancellation {
2746 fn drop(&mut self) {
2747 self.registry.remove(&self.id, self.generation);
2748 }
2749}
2750
2751fn cancellation_request_id(dispatch: &Dispatch) -> Result<Option<RequestId>, crate::Error> {
2752 let Dispatch::Notification(message) = dispatch else {
2753 return Ok(None);
2754 };
2755 cancellation_request_id_from_message(message)
2756}
2757
2758fn cancellation_request_id_from_message(
2759 message: &UntypedMessage,
2760) -> Result<Option<RequestId>, crate::Error> {
2761 let (method, params) = peel_successor_envelopes(&message.method, &message.params);
2762 if !crate::schema::v1::CancelRequestNotification::matches_method(method) {
2763 return Ok(None);
2764 }
2765
2766 let notification = crate::schema::v1::CancelRequestNotification::parse_message(method, params)?;
2767 Ok(Some(notification.request_id))
2768}
2769
2770fn peel_successor_envelopes<'message>(
2783 mut method: &'message str,
2784 mut params: &'message serde_json::Value,
2785) -> (&'message str, &'message serde_json::Value) {
2786 while crate::schema::SuccessorMessage::<UntypedMessage>::matches_method(method) {
2787 let Some(inner_method) = params.get("method").and_then(serde_json::Value::as_str) else {
2788 break;
2789 };
2790 method = inner_method;
2791 params = params.get("params").unwrap_or(&serde_json::Value::Null);
2792 }
2793 (method, params)
2794}
2795
2796#[must_use]
2811pub fn is_cancel_request_notification<N: JsonRpcNotification>(notification: &N) -> bool {
2812 let method = notification.method();
2813 if crate::schema::v1::CancelRequestNotification::matches_method(method) {
2814 return true;
2815 }
2816 if !crate::schema::SuccessorMessage::<UntypedMessage>::matches_method(method) {
2817 return false;
2818 }
2819
2820 match notification.to_untyped_message() {
2821 Ok(untyped) => {
2822 let (method, _params) = peel_successor_envelopes(&untyped.method, &untyped.params);
2823 crate::schema::v1::CancelRequestNotification::matches_method(method)
2824 }
2825 Err(error) => {
2826 tracing::debug!(
2827 ?error,
2828 "failed to inspect successor-wrapped notification for cancellation"
2829 );
2830 false
2831 }
2832 }
2833}
2834
2835#[derive(Clone)]
2837enum ResponseDestination {
2838 Individual(IndividualResponseSlot),
2839 Batch(BatchResponseSlot),
2840}
2841
2842impl std::fmt::Debug for ResponseDestination {
2843 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2844 match self {
2845 Self::Individual(slot) => formatter.debug_tuple("Individual").field(slot).finish(),
2846 Self::Batch(slot) => formatter.debug_tuple("Batch").field(slot).finish(),
2847 }
2848 }
2849}
2850
2851impl ResponseDestination {
2852 fn individual() -> Self {
2853 Self::Individual(IndividualResponseSlot::default())
2854 }
2855
2856 fn batch(slot_count: usize) -> (impl Iterator<Item = Self>, BatchDispatchCompletion) {
2857 let state = Arc::new(Mutex::new(BatchResponseState {
2858 remaining: slot_count,
2859 responses: (0..slot_count).map(|_| None).collect(),
2860 abandoned: (0..slot_count).map(|_| None).collect(),
2861 active_handler_attempts: (0..slot_count).map(|_| 0).collect(),
2862 dispatch_complete: false,
2863 emitted: false,
2864 }));
2865
2866 (
2867 (0..slot_count).map({
2868 let state = state.clone();
2869 move |index| {
2870 Self::Batch(BatchResponseSlot {
2871 state: state.clone(),
2872 index,
2873 })
2874 }
2875 }),
2876 BatchDispatchCompletion { state },
2877 )
2878 }
2879
2880 fn complete(self, response: RawJsonRpcMessage) -> Option<TransportFrame> {
2881 match self {
2882 Self::Individual(slot) => slot.complete(response),
2883 Self::Batch(slot) => slot.complete(response).map(batch_response_frame),
2884 }
2885 }
2886
2887 fn abandon(self, fallback: RawJsonRpcMessage) -> Option<TransportFrame> {
2888 match self {
2889 Self::Individual(_) => None,
2890 Self::Batch(slot) => slot.abandon(fallback).map(batch_response_frame),
2891 }
2892 }
2893
2894 fn is_batch(&self) -> bool {
2895 matches!(self, Self::Batch(_))
2896 }
2897
2898 fn begin_handler_attempt(
2899 &self,
2900 message_tx: OutgoingMessageTx,
2901 ) -> Option<ResponderHandlerAttempt> {
2902 let Self::Batch(slot) = self else {
2903 return None;
2904 };
2905 slot.begin_handler_attempt();
2906 Some(ResponderHandlerAttempt {
2907 message_tx,
2908 destination: self.clone(),
2909 })
2910 }
2911
2912 fn finish_handler_attempt(self) -> Option<TransportFrame> {
2913 match self {
2914 Self::Individual(_) => None,
2915 Self::Batch(slot) => slot.finish_handler_attempt().map(batch_response_frame),
2916 }
2917 }
2918}
2919
2920#[derive(Clone, Debug, Default)]
2921struct IndividualResponseSlot {
2922 completed: Arc<AtomicBool>,
2923}
2924
2925impl IndividualResponseSlot {
2926 fn complete(self, response: RawJsonRpcMessage) -> Option<TransportFrame> {
2927 if self.completed.swap(true, Ordering::AcqRel) {
2928 tracing::warn!("Ignoring duplicate completion of JSON-RPC request");
2929 return None;
2930 }
2931
2932 Some(TransportFrame::Single(response))
2933 }
2934}
2935
2936fn batch_response_frame(responses: Vec<RawJsonRpcMessage>) -> TransportFrame {
2937 TransportFrame::Batch(
2938 TransportBatch::from_messages(responses)
2939 .expect("a completed JSON-RPC response batch is non-empty"),
2940 )
2941}
2942
2943#[derive(Clone)]
2944struct BatchDispatchCompletion {
2945 state: Arc<Mutex<BatchResponseState>>,
2946}
2947
2948impl std::fmt::Debug for BatchDispatchCompletion {
2949 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2950 formatter
2951 .debug_struct("BatchDispatchCompletion")
2952 .finish_non_exhaustive()
2953 }
2954}
2955
2956impl BatchDispatchCompletion {
2957 fn complete(self) -> Option<TransportFrame> {
2958 let mut state = self
2959 .state
2960 .lock()
2961 .expect("batch response accumulator mutex poisoned");
2962 if state.dispatch_complete {
2963 tracing::warn!("Ignoring duplicate JSON-RPC batch dispatch completion");
2964 return None;
2965 }
2966 state.dispatch_complete = true;
2967 for index in 0..state.responses.len() {
2968 promote_abandoned_response(&mut state, index);
2969 }
2970 take_completed_batch(&mut state).map(batch_response_frame)
2971 }
2972}
2973
2974fn promote_abandoned_response(state: &mut BatchResponseState, index: usize) {
2975 if state.active_handler_attempts[index] == 0
2976 && state.responses[index].is_none()
2977 && let Some(fallback) = state.abandoned[index].take()
2978 {
2979 state.responses[index] = Some(fallback);
2980 state.remaining -= 1;
2981 }
2982}
2983
2984fn take_completed_batch(state: &mut BatchResponseState) -> Option<Vec<RawJsonRpcMessage>> {
2985 if !state.dispatch_complete || state.remaining != 0 || state.emitted {
2986 return None;
2987 }
2988
2989 state.emitted = true;
2990 Some(
2991 state
2992 .responses
2993 .iter_mut()
2994 .map(|response| {
2995 response
2996 .take()
2997 .expect("completed JSON-RPC batch has every response slot")
2998 })
2999 .collect(),
3000 )
3001}
3002
3003#[derive(Clone)]
3004struct BatchResponseSlot {
3005 state: Arc<Mutex<BatchResponseState>>,
3006 index: usize,
3007}
3008
3009impl std::fmt::Debug for BatchResponseSlot {
3010 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
3011 formatter
3012 .debug_struct("BatchResponseSlot")
3013 .field("index", &self.index)
3014 .finish_non_exhaustive()
3015 }
3016}
3017
3018impl BatchResponseSlot {
3019 fn begin_handler_attempt(&self) {
3020 let mut state = self
3021 .state
3022 .lock()
3023 .expect("batch response accumulator mutex poisoned");
3024 state.active_handler_attempts[self.index] += 1;
3025 }
3026
3027 fn finish_handler_attempt(self) -> Option<Vec<RawJsonRpcMessage>> {
3028 let mut state = self
3029 .state
3030 .lock()
3031 .expect("batch response accumulator mutex poisoned");
3032 state.active_handler_attempts[self.index] = state.active_handler_attempts[self.index]
3033 .checked_sub(1)
3034 .expect("handler attempt completion without a matching start");
3035 if state.dispatch_complete {
3036 promote_abandoned_response(&mut state, self.index);
3037 }
3038 take_completed_batch(&mut state)
3039 }
3040
3041 fn complete(self, response: RawJsonRpcMessage) -> Option<Vec<RawJsonRpcMessage>> {
3042 let mut state = self
3043 .state
3044 .lock()
3045 .expect("batch response accumulator mutex poisoned");
3046 if state.emitted {
3047 tracing::warn!(
3048 index = self.index,
3049 "Ignoring response after JSON-RPC batch was already completed"
3050 );
3051 return None;
3052 }
3053 if self.index >= state.responses.len() {
3054 tracing::error!(index = self.index, "Invalid JSON-RPC batch response slot");
3055 return None;
3056 }
3057 if state.responses[self.index].is_some() {
3058 tracing::warn!(
3059 index = self.index,
3060 "Ignoring duplicate completion of JSON-RPC batch response slot"
3061 );
3062 return None;
3063 }
3064
3065 state.abandoned[self.index] = None;
3066 state.responses[self.index] = Some(response);
3067 state.remaining -= 1;
3068 take_completed_batch(&mut state)
3069 }
3070
3071 fn abandon(self, fallback: RawJsonRpcMessage) -> Option<Vec<RawJsonRpcMessage>> {
3072 let mut state = self
3073 .state
3074 .lock()
3075 .expect("batch response accumulator mutex poisoned");
3076 if state.emitted || state.responses[self.index].is_some() {
3077 return None;
3078 }
3079 if state.abandoned[self.index].is_some() {
3080 tracing::warn!(
3081 index = self.index,
3082 "Ignoring duplicate abandonment of JSON-RPC batch response slot"
3083 );
3084 return None;
3085 }
3086
3087 if state.dispatch_complete && state.active_handler_attempts[self.index] == 0 {
3088 state.responses[self.index] = Some(fallback);
3089 state.remaining -= 1;
3090 } else {
3091 state.abandoned[self.index] = Some(fallback);
3092 }
3093 take_completed_batch(&mut state)
3094 }
3095}
3096
3097struct BatchResponseState {
3098 remaining: usize,
3099 responses: Vec<Option<RawJsonRpcMessage>>,
3100 abandoned: Vec<Option<RawJsonRpcMessage>>,
3101 active_handler_attempts: Vec<usize>,
3102 dispatch_complete: bool,
3103 emitted: bool,
3104}
3105
3106#[derive(Clone, Debug)]
3107struct RequestReplyTarget {
3108 id: RequestId,
3109 method: String,
3110 destination: ResponseDestination,
3111}
3112
3113struct ResponderHandlerAttempt {
3114 message_tx: OutgoingMessageTx,
3115 destination: ResponseDestination,
3116}
3117
3118impl Drop for ResponderHandlerAttempt {
3119 fn drop(&mut self) {
3120 if let Err(error) = send_raw_message(
3121 &self.message_tx,
3122 OutgoingMessage::BatchHandlerAttemptComplete {
3123 destination: self.destination.clone(),
3124 },
3125 ) {
3126 tracing::debug!(?error, "could not complete JSON-RPC batch handler attempt");
3127 }
3128 }
3129}
3130
3131#[derive(Clone)]
3132struct ResponseReplyTarget {
3133 id: RequestId,
3134 method: String,
3135 sender: Arc<Mutex<Option<oneshot::Sender<ResponsePayload>>>>,
3136 ordering: ResponseOrdering,
3137 dispatch: ResponseDispatch,
3138}
3139
3140impl ResponseReplyTarget {
3141 fn route(self, result: Result<serde_json::Value, crate::Error>) {
3142 let sender = self
3143 .sender
3144 .lock()
3145 .expect("response reply mutex poisoned")
3146 .take();
3147 let Some(sender) = sender else {
3148 tracing::debug!(
3149 method = %self.method,
3150 id = ?self.id,
3151 "response was already routed to its local awaiter"
3152 );
3153 return;
3154 };
3155
3156 let ack_tx = self.dispatch.acknowledgment(&self.ordering);
3157 if sender.send(ResponsePayload { result, ack_tx }).is_err() {
3158 tracing::debug!(
3159 method = %self.method,
3160 id = ?self.id,
3161 "dropped response because local receiver was gone"
3162 );
3163 }
3164 }
3165}
3166
3167#[derive(Clone, Default)]
3168struct ResponseDispatch {
3169 state: Arc<Mutex<ResponseDispatchState>>,
3170}
3171
3172#[derive(Default)]
3173struct ResponseDispatchState {
3174 complete: bool,
3175 ack_rx: Option<oneshot::Receiver<()>>,
3176}
3177
3178impl ResponseDispatch {
3179 fn acknowledgment(&self, ordering: &ResponseOrdering) -> Option<oneshot::Sender<()>> {
3180 if !ordering.is_ordered() {
3181 return None;
3182 }
3183
3184 let mut state = self.state.lock().expect("response dispatch mutex poisoned");
3185 if state.complete {
3186 return None;
3187 }
3188
3189 let (ack_tx, ack_rx) = oneshot::channel();
3190 let previous_ack = state.ack_rx.replace(ack_rx);
3191 debug_assert!(
3192 previous_ack.is_none(),
3193 "a response dispatch can only be routed once"
3194 );
3195 Some(ack_tx)
3196 }
3197
3198 fn complete(&self) -> Option<oneshot::Receiver<()>> {
3199 let mut state = self.state.lock().expect("response dispatch mutex poisoned");
3200 state.complete = true;
3201 state.ack_rx.take()
3202 }
3203}
3204
3205enum HandlerErrorTarget {
3206 Request(RequestReplyTarget),
3207 Response(ResponseReplyTarget),
3208}
3209
3210impl HandlerErrorTarget {
3211 fn begin_handler_attempt(
3212 &self,
3213 message_tx: &OutgoingMessageTx,
3214 ) -> Option<ResponderHandlerAttempt> {
3215 match self {
3216 Self::Request(target) => target.destination.begin_handler_attempt(message_tx.clone()),
3217 Self::Response(_) => None,
3218 }
3219 }
3220}
3221
3222#[derive(Debug)]
3223enum OutgoingMessage {
3224 CloseAfterDraining { done: oneshot::Sender<()> },
3227
3228 BatchDispatchComplete { completion: BatchDispatchCompletion },
3231
3232 BatchHandlerAttemptComplete { destination: ResponseDestination },
3235
3236 AbandonedBatchResponse {
3240 id: RequestId,
3241 method: String,
3242 destination: ResponseDestination,
3243 },
3244
3245 Request {
3247 id: RequestId,
3249
3250 method: String,
3252
3253 untyped: UntypedMessage,
3255
3256 remote_style: crate::role::RemoteStyle,
3258
3259 readiness: Option<RequestReadiness>,
3262 },
3263
3264 Notification {
3266 untyped: UntypedMessage,
3269 },
3270
3271 Response {
3273 id: RequestId,
3274
3275 method: String,
3277
3278 response: Result<serde_json::Value, crate::Error>,
3279
3280 destination: ResponseDestination,
3281 },
3282
3283 UncorrelatedErrorResponse {
3285 error: crate::Error,
3286 destination: ResponseDestination,
3287 },
3288}
3289
3290#[must_use]
3292#[derive(Debug)]
3293pub enum Handled<T> {
3294 Yes,
3296
3297 No {
3301 message: T,
3305
3306 retry: bool,
3314 },
3315}
3316
3317pub trait IntoHandled<T> {
3322 fn into_handled(self) -> Handled<T>;
3324}
3325
3326impl<T> IntoHandled<T> for () {
3327 fn into_handled(self) -> Handled<T> {
3328 Handled::Yes
3329 }
3330}
3331
3332impl<T> IntoHandled<T> for Handled<T> {
3333 fn into_handled(self) -> Handled<T> {
3334 self
3335 }
3336}
3337
3338#[cfg(feature = "unstable_protocol_v2")]
3350#[derive(Clone, Debug)]
3351pub struct V2ConnectionTo<Counterpart: Role> {
3352 inner: ConnectionTo<Counterpart>,
3353}
3354
3355#[cfg(feature = "unstable_protocol_v2")]
3356impl<Counterpart: Role> V2ConnectionTo<Counterpart> {
3357 pub(crate) fn raw_connection(&self) -> &ConnectionTo<Counterpart> {
3359 &self.inner
3360 }
3361
3362 pub fn counterpart(&self) -> Counterpart {
3364 self.inner.counterpart()
3365 }
3366
3367 pub async fn incoming_closed(&self) {
3369 self.inner.incoming_closed().await;
3370 }
3371
3372 #[must_use]
3374 pub fn is_incoming_closed(&self) -> bool {
3375 self.inner.is_incoming_closed()
3376 }
3377
3378 #[track_caller]
3380 pub fn spawn(
3381 &self,
3382 task: impl IntoFuture<Output = Result<(), crate::Error>, IntoFuture: Send + 'static>,
3383 ) -> Result<(), crate::Error> {
3384 self.inner.spawn(task)
3385 }
3386
3387 #[track_caller]
3407 pub fn spawn_connection<R: Role, Context: ConnectionContext>(
3408 &self,
3409 builder: Builder<
3410 R,
3411 impl HandleDispatchFrom<R::Counterpart> + 'static,
3412 impl RunWithConnectionTo<R::Counterpart> + 'static,
3413 impl HandleConnectionClose<R::Counterpart> + 'static,
3414 Context,
3415 >,
3416 transport: impl ConnectTo<R> + 'static,
3417 ) -> Result<Context::Connection<R::Counterpart>, crate::Error> {
3418 self.inner.spawn_connection_with_context(builder, transport)
3419 }
3420
3421 pub fn send_proxied_message<Req: JsonRpcRequest<Response: Send>, Notif: JsonRpcNotification>(
3423 &self,
3424 message: Dispatch<Req, Notif>,
3425 ) -> Result<(), crate::Error>
3426 where
3427 Counterpart: HasPeer<Counterpart>,
3428 {
3429 self.inner.send_proxied_message(message)
3430 }
3431
3432 pub fn send_proxied_message_to<
3435 Peer: Role,
3436 Req: JsonRpcRequest<Response: Send>,
3437 Notif: JsonRpcNotification,
3438 >(
3439 &self,
3440 peer: Peer,
3441 message: Dispatch<Req, Notif>,
3442 ) -> Result<(), crate::Error>
3443 where
3444 Counterpart: HasPeer<Peer>,
3445 {
3446 self.inner.send_proxied_message_to(peer, message)
3447 }
3448
3449 pub fn send_request<Req: JsonRpcRequest>(&self, request: Req) -> SentRequest<Req::Response>
3451 where
3452 Counterpart: HasPeer<Counterpart>,
3453 {
3454 self.inner.send_request(request)
3455 }
3456
3457 pub fn send_request_to<Peer: Role, Req: JsonRpcRequest>(
3459 &self,
3460 peer: Peer,
3461 request: Req,
3462 ) -> SentRequest<Req::Response>
3463 where
3464 Counterpart: HasPeer<Peer>,
3465 {
3466 self.inner.send_request_to(peer, request)
3467 }
3468
3469 pub fn send_notification<N: JsonRpcNotification>(
3471 &self,
3472 notification: N,
3473 ) -> Result<(), crate::Error>
3474 where
3475 Counterpart: HasPeer<Counterpart>,
3476 {
3477 self.inner.send_notification(notification)
3478 }
3479
3480 pub fn send_notification_to<Peer: Role, N: JsonRpcNotification>(
3482 &self,
3483 peer: Peer,
3484 notification: N,
3485 ) -> Result<(), crate::Error>
3486 where
3487 Counterpart: HasPeer<Peer>,
3488 {
3489 self.inner.send_notification_to(peer, notification)
3490 }
3491
3492 pub fn send_cancel_request(
3494 &self,
3495 request_id: impl Into<crate::schema::v1::RequestId>,
3496 ) -> Result<(), crate::Error>
3497 where
3498 Counterpart: HasPeer<Counterpart>,
3499 {
3500 self.inner.send_cancel_request(request_id)
3501 }
3502
3503 pub fn send_cancel_request_to<Peer: Role>(
3505 &self,
3506 peer: Peer,
3507 request_id: impl Into<crate::schema::v1::RequestId>,
3508 ) -> Result<(), crate::Error>
3509 where
3510 Counterpart: HasPeer<Peer>,
3511 {
3512 self.inner.send_cancel_request_to(peer, request_id)
3513 }
3514
3515 pub fn add_dynamic_handler(
3540 &self,
3541 handler: impl HandleDispatchFrom<Counterpart> + 'static,
3542 ) -> Result<DynamicHandlerGuard<Counterpart>, crate::Error> {
3543 self.inner.add_dynamic_handler(handler)
3544 }
3545}
3546
3547#[derive(Clone, Debug)]
3570pub struct ConnectionTo<Counterpart: Role> {
3571 counterpart: Counterpart,
3572 message_tx: OutgoingMessageTx,
3573 task_tx: TaskTx,
3574 dynamic_handler_tx: mpsc::UnboundedSender<DynamicHandlerMessage<Counterpart>>,
3575 transport_completion: SharedTransportCompletion,
3576 pending_replies: PendingRepliesRegistrar,
3577 #[cfg_attr(
3578 not(feature = "unstable_protocol_v2"),
3579 allow(
3580 dead_code,
3581 reason = "retained so ConnectionTo has one constructor shape"
3582 )
3583 )]
3584 protocol_mode: ProtocolMode,
3585 incoming_closed: IncomingClosed,
3586 protected_operations: Arc<Mutex<ProtectedOperations>>,
3587 runner_error_scope: Option<run::RunnerErrorScope>,
3588}
3589
3590type SharedTransportCompletion = future::Shared<BoxFuture<'static, Result<(), crate::Error>>>;
3591
3592type SharedCompletionSignal = future::Shared<BoxFuture<'static, ()>>;
3593
3594#[derive(Default)]
3595struct ProtectedOperations {
3596 pending: Vec<oneshot::Receiver<()>>,
3597 joining: Option<SharedCompletionSignal>,
3598}
3599
3600impl Debug for ProtectedOperations {
3601 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
3602 formatter
3603 .debug_struct("ProtectedOperations")
3604 .field("pending", &self.pending.len())
3605 .field("joining", &self.joining.is_some())
3606 .finish_non_exhaustive()
3607 }
3608}
3609
3610fn completion_signal() -> (oneshot::Sender<()>, SharedCompletionSignal) {
3611 let (tx, rx) = oneshot::channel();
3612 let signal = async move {
3613 if rx.await.is_err() {
3615 future::pending::<()>().await;
3616 }
3617 }
3618 .boxed()
3619 .shared();
3620 (tx, signal)
3621}
3622
3623#[derive(Clone)]
3624struct IncomingClosed {
3625 state: Arc<IncomingClosedState>,
3626}
3627
3628struct IncomingClosedState {
3629 closing: AtomicBool,
3630 closed: AtomicBool,
3631 signal_tx: Mutex<Option<oneshot::Sender<()>>>,
3632 signal_rx: future::Shared<BoxFuture<'static, ()>>,
3633 shutdown_tx: Mutex<Option<oneshot::Sender<()>>>,
3634 #[cfg(any(feature = "unstable_mcp_over_acp", test))]
3635 shutdown_rx: SharedCompletionSignal,
3636}
3637
3638impl IncomingClosed {
3639 fn new() -> Self {
3640 let (signal_tx, signal_rx) = oneshot::channel();
3641 let (shutdown_tx, shutdown_rx) = oneshot::channel();
3642 #[cfg(not(any(feature = "unstable_mcp_over_acp", test)))]
3643 drop(shutdown_rx);
3644 Self {
3645 state: Arc::new(IncomingClosedState {
3646 closing: AtomicBool::new(false),
3647 closed: AtomicBool::new(false),
3648 signal_tx: Mutex::new(Some(signal_tx)),
3649 signal_rx: signal_rx.map(|_| ()).boxed().shared(),
3650 shutdown_tx: Mutex::new(Some(shutdown_tx)),
3651 #[cfg(any(feature = "unstable_mcp_over_acp", test))]
3652 shutdown_rx: shutdown_rx.map(|_| ()).boxed().shared(),
3653 }),
3654 }
3655 }
3656
3657 fn begin_close(&self) {
3658 self.state.closing.store(true, Ordering::Release);
3659 self.request_shutdown();
3660 }
3661
3662 fn request_shutdown(&self) {
3663 if let Some(tx) = self
3664 .state
3665 .shutdown_tx
3666 .lock()
3667 .expect("shutdown signal mutex poisoned")
3668 .take()
3669 {
3670 let _ = tx.send(());
3671 }
3672 }
3673
3674 fn finish_close(&self) {
3675 self.state.closed.store(true, Ordering::Release);
3676 let signal_tx = self
3677 .state
3678 .signal_tx
3679 .lock()
3680 .expect("incoming-close signal mutex poisoned")
3681 .take();
3682
3683 if let Some(signal_tx) = signal_tx {
3684 let _ = signal_tx.send(());
3685 }
3686 }
3687
3688 async fn closed(&self) {
3689 self.state.signal_rx.clone().await;
3690 }
3691
3692 fn is_closed(&self) -> bool {
3693 self.state.closed.load(Ordering::Acquire)
3694 }
3695
3696 fn is_closing(&self) -> bool {
3697 self.state.closing.load(Ordering::Acquire)
3698 }
3699}
3700
3701impl Debug for IncomingClosed {
3702 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
3703 formatter
3704 .debug_struct("IncomingClosed")
3705 .field("is_closing", &self.is_closing())
3706 .field("is_closed", &self.is_closed())
3707 .finish_non_exhaustive()
3708 }
3709}
3710
3711pub const INCOMING_TRANSPORT_CLOSED_REASON: &str = "incoming_transport_closed";
3715
3716#[must_use]
3719pub fn is_incoming_transport_closed(error: &crate::Error) -> bool {
3720 error
3721 .data
3722 .as_ref()
3723 .and_then(|data| data.get("reason"))
3724 .and_then(serde_json::Value::as_str)
3725 == Some(INCOMING_TRANSPORT_CLOSED_REASON)
3726}
3727
3728fn incoming_transport_closed_error(method: &str) -> crate::Error {
3729 let mut error = crate::Error::internal_error();
3730 error.message = "Incoming transport closed".to_string();
3731 error.data(serde_json::json!({
3732 "reason": INCOMING_TRANSPORT_CLOSED_REASON,
3733 "method": method,
3734 }))
3735}
3736
3737fn run_incoming_until_foreground_succeeds(
3741 incoming: impl Future<Output = Result<(), crate::Error>>,
3742 foreground_succeeded: SharedCompletionSignal,
3743 incoming_closed: IncomingClosed,
3744) -> impl Future<Output = Result<(), crate::Error>> {
3745 let mut incoming = Box::pin(incoming);
3746 future::poll_fn(move |cx| {
3747 if foreground_succeeded.clone().poll_unpin(cx).is_ready() && !incoming_closed.is_closing() {
3748 return std::task::Poll::Ready(Ok(()));
3749 }
3750 incoming.as_mut().poll(cx)
3753 })
3754}
3755
3756fn run_until_connection_close<R>(
3759 background: impl Future<Output = Result<(), crate::Error>>,
3760 foreground: impl Future<Output = Result<R, crate::Error>>,
3761 incoming_closed: IncomingClosed,
3762) -> impl Future<Output = Result<R, crate::Error>> {
3763 let background = Box::pin(background);
3767 let foreground = Box::pin(foreground);
3768
3769 async move {
3770 match future::select(background, foreground).await {
3771 Either::Left((background_result, foreground)) => {
3772 background_result?;
3773 foreground.await
3774 }
3775 Either::Right((foreground_result, background)) => {
3776 if !incoming_closed.is_closing() {
3777 return foreground_result;
3778 }
3779
3780 match future::select(background, Box::pin(incoming_closed.closed())).await {
3781 Either::Left((background_result, _)) => {
3782 background_result?;
3783 foreground_result
3784 }
3785 Either::Right(((), background)) => {
3786 crate::util::run_until(background, future::ready(foreground_result)).await
3790 }
3791 }
3792 }
3793 }
3794 }
3795}
3796
3797impl<Counterpart: Role> ConnectionTo<Counterpart> {
3798 fn new(
3799 counterpart: Counterpart,
3800 message_tx: mpsc::UnboundedSender<OutgoingMessage>,
3801 task_tx: mpsc::UnboundedSender<Task>,
3802 dynamic_handler_tx: mpsc::UnboundedSender<DynamicHandlerMessage<Counterpart>>,
3803 transport_completion: SharedTransportCompletion,
3804 pending_replies: PendingRepliesRegistrar,
3805 protocol_mode: ProtocolMode,
3806 ) -> Self {
3807 Self {
3808 counterpart,
3809 message_tx,
3810 task_tx,
3811 dynamic_handler_tx,
3812 transport_completion,
3813 pending_replies,
3814 protocol_mode,
3815 incoming_closed: IncomingClosed::new(),
3816 protected_operations: Arc::default(),
3817 runner_error_scope: None,
3818 }
3819 }
3820
3821 pub(crate) fn with_runner_error_scope(mut self, scope: run::RunnerErrorScope) -> Self {
3822 self.runner_error_scope = Some(scope);
3823 self
3824 }
3825
3826 pub(crate) fn finish_runner_error(
3827 &self,
3828 error: crate::Error,
3829 ) -> impl Future<Output = ()> + Send + '_ {
3830 if let Some(scope) = &self.runner_error_scope {
3831 Either::Left(scope.finish(error))
3832 } else {
3833 self.request_shutdown();
3835 Either::Right(self.wait_protected_operations())
3836 }
3837 }
3838
3839 #[cfg(any(feature = "unstable_mcp_over_acp", test))]
3842 #[track_caller]
3843 pub(crate) fn spawn_protected(
3844 &self,
3845 task: impl IntoFuture<Output = Result<(), crate::Error>, IntoFuture: Send + 'static>,
3846 ) -> Result<(), crate::Error> {
3847 let mut state = self
3848 .protected_operations
3849 .lock()
3850 .expect("protected operations mutex poisoned");
3851 if state.joining.is_some() {
3852 return Err(crate::Error::request_cancelled());
3853 }
3854 state
3857 .pending
3858 .retain_mut(|done| matches!(done.try_recv(), Ok(None)));
3859 let (done_tx, done_rx) = oneshot::channel();
3860 let task = task.into_future();
3861 self.spawn(async move {
3862 let result = task.await;
3863 let _ = done_tx.send(());
3864 result
3865 })?;
3866 state.pending.push(done_rx);
3867 Ok(())
3868 }
3869
3870 pub(crate) async fn wait_protected_operations(&self) {
3871 let joining = {
3872 let mut state = self
3873 .protected_operations
3874 .lock()
3875 .expect("protected operations mutex poisoned");
3876 if state.joining.is_none() {
3877 let operations = std::mem::take(&mut state.pending);
3878 state.joining = Some(
3879 async move {
3880 for operation in operations {
3881 let _ = operation.await;
3882 }
3883 }
3884 .boxed()
3885 .shared(),
3886 );
3887 }
3888 state.joining.as_ref().expect("join initialized").clone()
3889 };
3890 joining.await;
3891 }
3892
3893 pub(crate) fn request_shutdown(&self) {
3894 self.incoming_closed.request_shutdown();
3895 }
3896
3897 #[cfg(any(feature = "unstable_mcp_over_acp", test))]
3899 pub(crate) async fn shutdown_requested(&self) {
3900 self.incoming_closed.state.shutdown_rx.clone().await;
3901 }
3902
3903 #[cfg(feature = "unstable_protocol_v2")]
3904 pub(crate) fn acp_protocol_version(&self) -> Option<crate::schema::ProtocolVersion> {
3905 self.protocol_mode.api_protocol_version()
3906 }
3907
3908 pub fn counterpart(&self) -> Counterpart {
3910 self.counterpart.clone()
3911 }
3912
3913 pub async fn incoming_closed(&self) {
3922 self.incoming_closed.closed().await;
3923 }
3924
3925 #[must_use]
3929 pub fn is_incoming_closed(&self) -> bool {
3930 self.incoming_closed.is_closed()
3931 }
3932
3933 async fn drain_outgoing(
3937 &self,
3938 finish: Option<crate::component::FinishControl>,
3939 wait_owned_transport: bool,
3940 ) -> Result<(), crate::Error> {
3941 let (done_tx, done_rx) = oneshot::channel();
3942 let marker_result = send_raw_message(
3943 &self.message_tx,
3944 OutgoingMessage::CloseAfterDraining { done: done_tx },
3945 );
3946 let marker_result = match marker_result {
3947 Ok(()) => done_rx.await.map_err(|error| {
3948 crate::util::internal_error(format!(
3949 "outgoing drain marker was dropped before completion: {error}"
3950 ))
3951 }),
3952 Err(error) => Err(error),
3953 };
3954
3955 let physical_finish = finish.is_some();
3956 if let Some(mut finish) = finish {
3957 finish.request();
3960 }
3961 if physical_finish || wait_owned_transport {
3962 self.transport_completion.clone().await?;
3965 }
3966 marker_result
3970 }
3971
3972 fn is_incoming_closing(&self) -> bool {
3973 self.incoming_closed.is_closing()
3974 }
3975
3976 pub(super) fn begin_incoming_close(&self) {
3977 self.incoming_closed.begin_close();
3978 }
3979
3980 pub(super) fn finish_incoming_close(&self) {
3981 self.incoming_closed.finish_close();
3982 }
3983
3984 #[track_caller]
4022 pub fn spawn(
4023 &self,
4024 task: impl IntoFuture<Output = Result<(), crate::Error>, IntoFuture: Send + 'static>,
4025 ) -> Result<(), crate::Error> {
4026 let location = std::panic::Location::caller();
4027 let task = task.into_future();
4028 Task::new(location, task).spawn(&self.task_tx)
4029 }
4030
4031 #[track_caller]
4074 pub fn spawn_connection<R: Role>(
4075 &self,
4076 builder: Builder<
4077 R,
4078 impl HandleDispatchFrom<R::Counterpart> + 'static,
4079 impl RunWithConnectionTo<R::Counterpart> + 'static,
4080 impl HandleConnectionClose<R::Counterpart> + 'static,
4081 impl ConnectionContext,
4082 >,
4083 transport: impl ConnectTo<R> + 'static,
4084 ) -> Result<ConnectionTo<R::Counterpart>, crate::Error> {
4085 self.spawn_connection_raw(builder, transport)
4086 }
4087
4088 #[cfg(feature = "unstable_protocol_v2")]
4113 #[track_caller]
4114 pub fn spawn_connection_with_context<R: Role, Context: ConnectionContext>(
4115 &self,
4116 builder: Builder<
4117 R,
4118 impl HandleDispatchFrom<R::Counterpart> + 'static,
4119 impl RunWithConnectionTo<R::Counterpart> + 'static,
4120 impl HandleConnectionClose<R::Counterpart> + 'static,
4121 Context,
4122 >,
4123 transport: impl ConnectTo<R> + 'static,
4124 ) -> Result<Context::Connection<R::Counterpart>, crate::Error> {
4125 let connection = self.spawn_connection_raw(builder, transport)?;
4126 Ok(connection_context::from_raw::<Context, _>(connection))
4127 }
4128
4129 #[track_caller]
4130 fn spawn_connection_raw<R: Role, Context: ConnectionContext>(
4131 &self,
4132 builder: Builder<
4133 R,
4134 impl HandleDispatchFrom<R::Counterpart> + 'static,
4135 impl RunWithConnectionTo<R::Counterpart> + 'static,
4136 impl HandleConnectionClose<R::Counterpart> + 'static,
4137 Context,
4138 >,
4139 transport: impl ConnectTo<R> + 'static,
4140 ) -> Result<ConnectionTo<R::Counterpart>, crate::Error> {
4141 let (connection, future) =
4142 builder.into_connection_and_future(transport, false, |_| std::future::pending());
4143 Task::new(std::panic::Location::caller(), future).spawn(&self.task_tx)?;
4144 Ok(connection)
4145 }
4146
4147 pub fn send_proxied_message<Req: JsonRpcRequest<Response: Send>, Notif: JsonRpcNotification>(
4152 &self,
4153 message: Dispatch<Req, Notif>,
4154 ) -> Result<(), crate::Error>
4155 where
4156 Counterpart: HasPeer<Counterpart>,
4157 {
4158 self.send_proxied_message_to(self.counterpart(), message)
4159 }
4160
4161 pub fn send_proxied_message_to<
4173 Peer: Role,
4174 Req: JsonRpcRequest<Response: Send>,
4175 Notif: JsonRpcNotification,
4176 >(
4177 &self,
4178 peer: Peer,
4179 message: Dispatch<Req, Notif>,
4180 ) -> Result<(), crate::Error>
4181 where
4182 Counterpart: HasPeer<Peer>,
4183 {
4184 match message {
4185 Dispatch::Request(request, responder) => self
4186 .send_ordered_request_to(peer, request)
4187 .forward_response_to(responder),
4188 Dispatch::Notification(notification) => {
4189 if is_cancel_request_notification(¬ification) {
4197 tracing::debug!(
4198 "not forwarding hop-scoped `$/cancel_request` notification across proxy hop"
4199 );
4200 return Ok(());
4201 }
4202 self.send_notification_to(peer, notification)
4203 }
4204 Dispatch::Response(result, router) => {
4205 router.route_with_result(result)
4207 }
4208 }
4209 }
4210
4211 pub fn send_request<Req: JsonRpcRequest>(&self, request: Req) -> SentRequest<Req::Response>
4264 where
4265 Counterpart: HasPeer<Counterpart>,
4266 {
4267 self.send_request_to(self.counterpart.clone(), request)
4268 }
4269
4270 pub fn send_request_to<Peer: Role, Req: JsonRpcRequest>(
4275 &self,
4276 peer: Peer,
4277 request: Req,
4278 ) -> SentRequest<Req::Response>
4279 where
4280 Counterpart: HasPeer<Peer>,
4281 {
4282 self.send_request_to_with_options(peer, request, false, None, None)
4283 }
4284
4285 #[cfg(feature = "unstable_protocol_v2")]
4289 pub(crate) fn send_request_to_with_response_hook_after<
4290 Peer: Role,
4291 Req: JsonRpcRequest,
4292 BeforeSend: Future<Output = Result<(), crate::Error>> + Send + 'static,
4293 >(
4294 &self,
4295 peer: Peer,
4296 request: Req,
4297 before_send: BeforeSend,
4298 response_hook: impl FnOnce(&Req::Response) -> Result<(), crate::Error> + Send + 'static,
4299 ) -> SentRequest<Req::Response>
4300 where
4301 Counterpart: HasPeer<Peer>,
4302 {
4303 let hook: ResponseRouteHook = Box::new(move |method, value| {
4304 let response = Req::Response::from_value(method, value.clone())?;
4305 response_hook(&response)
4306 });
4307 self.send_request_to_with_options(
4308 peer,
4309 request,
4310 false,
4311 Some(RequestReadiness::new(before_send)),
4312 Some(hook),
4313 )
4314 }
4315
4316 #[cfg(feature = "unstable_protocol_v2")]
4318 pub(crate) fn send_ordered_request_to_with_response_hook_after<
4319 Peer: Role,
4320 Req: JsonRpcRequest,
4321 BeforeSend: Future<Output = Result<(), crate::Error>> + Send + 'static,
4322 >(
4323 &self,
4324 peer: Peer,
4325 request: Req,
4326 before_send: BeforeSend,
4327 response_hook: impl FnOnce(&Req::Response) -> Result<(), crate::Error> + Send + 'static,
4328 ) -> SentRequest<Req::Response>
4329 where
4330 Counterpart: HasPeer<Peer>,
4331 {
4332 let hook: ResponseRouteHook = Box::new(move |method, value| {
4333 let response = Req::Response::from_value(method, value.clone())?;
4334 response_hook(&response)
4335 });
4336 self.send_request_to_with_options(
4337 peer,
4338 request,
4339 true,
4340 Some(RequestReadiness::new(before_send)),
4341 Some(hook),
4342 )
4343 }
4344
4345 pub(crate) fn send_ordered_request_to<Peer: Role, Req: JsonRpcRequest>(
4352 &self,
4353 peer: Peer,
4354 request: Req,
4355 ) -> SentRequest<Req::Response>
4356 where
4357 Counterpart: HasPeer<Peer>,
4358 {
4359 self.send_request_to_with_options(peer, request, true, None, None)
4360 }
4361
4362 pub(crate) fn send_ordered_request_to_after<
4369 Peer: Role,
4370 Req: JsonRpcRequest,
4371 BeforeSend: Future<Output = Result<(), crate::Error>> + Send + 'static,
4372 >(
4373 &self,
4374 peer: Peer,
4375 request: Req,
4376 before_send: BeforeSend,
4377 ) -> SentRequest<Req::Response>
4378 where
4379 Counterpart: HasPeer<Peer>,
4380 {
4381 self.send_request_to_with_options(
4382 peer,
4383 request,
4384 true,
4385 Some(RequestReadiness::new(before_send)),
4386 None,
4387 )
4388 }
4389
4390 fn send_request_to_with_options<Peer: Role, Req: JsonRpcRequest>(
4391 &self,
4392 peer: Peer,
4393 request: Req,
4394 ordered: bool,
4395 readiness: Option<RequestReadiness>,
4396 response_route_hook: Option<ResponseRouteHook>,
4397 ) -> SentRequest<Req::Response>
4398 where
4399 Counterpart: HasPeer<Peer>,
4400 {
4401 let method = request.method().to_string();
4402 let id = RequestId::Str(uuid::Uuid::new_v4().to_string());
4403 let (response_tx, response_rx) = oneshot::channel();
4404 let response_ordering = ResponseOrdering::default();
4405 if ordered {
4406 response_ordering.mark_ordered();
4407 }
4408 let role_id = peer.role_id();
4409 let remote_style = self.counterpart.remote_style(peer);
4410 let cancellation =
4411 SentRequestCancellation::new(self.message_tx.clone(), remote_style, id.clone());
4412 if self.is_incoming_closing() {
4413 cancellation.disarm();
4414 drop(response_tx.send(ResponsePayload {
4415 result: Err(incoming_transport_closed_error(&method)),
4416 ack_tx: None,
4417 }));
4418 return SentRequest::new(
4419 id,
4420 method.clone(),
4421 self.task_tx.clone(),
4422 response_rx,
4423 cancellation,
4424 response_ordering,
4425 )
4426 .map(move |json| <Req::Response>::from_value(&method, json));
4427 }
4428
4429 match request.to_untyped_message() {
4430 Ok(untyped) => {
4431 let pending_reply = PendingReply {
4436 method: method.clone(),
4437 role_id,
4438 sender: response_tx,
4439 cancellation_disarm: cancellation.disarm_handle(),
4440 ordering: response_ordering.clone(),
4441 response_route_hook,
4442 };
4443
4444 if self
4445 .pending_replies
4446 .subscribe(id.clone(), pending_reply, &self.incoming_closed)
4447 {
4448 let message = OutgoingMessage::Request {
4449 id: id.clone(),
4450 method: method.clone(),
4451 untyped,
4452 remote_style,
4453 readiness,
4454 };
4455
4456 if let Err(error) = self.message_tx.unbounded_send(message) {
4457 cancellation.disarm();
4458
4459 let OutgoingMessage::Request { id, method, .. } = error.into_inner() else {
4460 unreachable!();
4461 };
4462
4463 if let Some(pending_reply) = self.pending_replies.remove(&id) {
4464 if self.is_incoming_closing() {
4465 pending_reply.fail_incoming_closed();
4466 } else {
4467 pending_reply.fail(crate::util::internal_error(format!(
4468 "failed to send outgoing request `{method}`"
4469 )));
4470 }
4471 }
4472 }
4473 }
4474 }
4475
4476 Err(err) => {
4477 cancellation.disarm();
4478
4479 response_tx
4480 .send(ResponsePayload {
4481 result: Err(crate::util::internal_error(format!(
4482 "failed to create untyped request for `{method}`: {err}"
4483 ))),
4484 ack_tx: None,
4485 })
4486 .unwrap();
4487 }
4488 }
4489
4490 SentRequest::new(
4491 id,
4492 method.clone(),
4493 self.task_tx.clone(),
4494 response_rx,
4495 cancellation,
4496 response_ordering,
4497 )
4498 .map(move |json| <Req::Response>::from_value(&method, json))
4499 }
4500
4501 pub fn send_notification<N: JsonRpcNotification>(
4519 &self,
4520 notification: N,
4521 ) -> Result<(), crate::Error>
4522 where
4523 Counterpart: HasPeer<Counterpart>,
4524 {
4525 self.send_notification_to(self.counterpart.clone(), notification)
4526 }
4527
4528 pub fn send_notification_to<Peer: Role, N: JsonRpcNotification>(
4533 &self,
4534 peer: Peer,
4535 notification: N,
4536 ) -> Result<(), crate::Error>
4537 where
4538 Counterpart: HasPeer<Peer>,
4539 {
4540 let remote_style = self.counterpart.remote_style(peer);
4541 tracing::debug!(
4542 role = std::any::type_name::<Counterpart>(),
4543 peer = std::any::type_name::<Peer>(),
4544 notification_type = std::any::type_name::<N>(),
4545 ?remote_style,
4546 original_method = notification.method(),
4547 "send_notification_to"
4548 );
4549 let transformed = remote_style.transform_outgoing_message(notification)?;
4550 tracing::debug!(
4551 transformed_method = %transformed.method,
4552 "send_notification_to transformed"
4553 );
4554 send_raw_message(
4555 &self.message_tx,
4556 OutgoingMessage::Notification {
4557 untyped: transformed,
4558 },
4559 )
4560 }
4561
4562 pub fn send_cancel_request(
4570 &self,
4571 request_id: impl Into<crate::schema::v1::RequestId>,
4572 ) -> Result<(), crate::Error>
4573 where
4574 Counterpart: HasPeer<Counterpart>,
4575 {
4576 self.send_cancel_request_to(self.counterpart.clone(), request_id)
4577 }
4578
4579 pub fn send_cancel_request_to<Peer: Role>(
4587 &self,
4588 peer: Peer,
4589 request_id: impl Into<crate::schema::v1::RequestId>,
4590 ) -> Result<(), crate::Error>
4591 where
4592 Counterpart: HasPeer<Peer>,
4593 {
4594 self.send_notification_to(
4595 peer,
4596 crate::schema::v1::CancelRequestNotification::new(request_id),
4597 )
4598 }
4599
4600 pub fn add_dynamic_handler(
4608 &self,
4609 handler: impl HandleDispatchFrom<Counterpart> + 'static,
4610 ) -> Result<DynamicHandlerGuard<Counterpart>, crate::Error> {
4611 let uuid = Uuid::new_v4();
4612 let active = Arc::new(AtomicBool::new(true));
4613 self.dynamic_handler_tx
4614 .unbounded_send(DynamicHandlerMessage::AddDynamicHandler(
4615 uuid,
4616 Box::new(GuardedDynamicHandler {
4617 active: active.clone(),
4618 handler,
4619 }),
4620 ))
4621 .map_err(crate::util::internal_error)?;
4622
4623 Ok(DynamicHandlerGuard::new(uuid, active, self.clone()))
4624 }
4625
4626 pub(crate) fn dynamic_handler_barrier(&self) -> BoxFuture<'static, Result<(), crate::Error>> {
4629 let (acknowledgment_tx, acknowledgment_rx) = oneshot::channel();
4630 if let Err(error) =
4631 self.dynamic_handler_tx
4632 .unbounded_send(DynamicHandlerMessage::AcknowledgedBarrier(
4633 acknowledgment_tx,
4634 ))
4635 {
4636 return future::ready(Err(crate::Error::into_internal_error(error))).boxed();
4637 }
4638
4639 async move {
4640 acknowledgment_rx.await.map_err(|error| {
4641 crate::util::internal_error(format!(
4642 "dynamic-handler barrier was dropped before acknowledgment: {error}"
4643 ))
4644 })
4645 }
4646 .boxed()
4647 }
4648
4649 fn remove_dynamic_handler(&self, uuid: Uuid) {
4650 drop(
4652 self.dynamic_handler_tx
4653 .unbounded_send(DynamicHandlerMessage::RemoveDynamicHandler(uuid)),
4654 );
4655 }
4656}
4657
4658struct GuardedDynamicHandler<Handler> {
4659 active: Arc<AtomicBool>,
4660 handler: Handler,
4661}
4662
4663impl<Counterpart, Handler> HandleDispatchFrom<Counterpart> for GuardedDynamicHandler<Handler>
4664where
4665 Counterpart: Role,
4666 Handler: HandleDispatchFrom<Counterpart>,
4667{
4668 async fn handle_dispatch_from(
4669 &mut self,
4670 message: Dispatch,
4671 connection: ConnectionTo<Counterpart>,
4672 ) -> Result<Handled<Dispatch>, crate::Error> {
4673 if !self.active.load(Ordering::Acquire) {
4674 return Ok(Handled::No {
4675 message,
4676 retry: false,
4677 });
4678 }
4679 self.handler.handle_dispatch_from(message, connection).await
4680 }
4681
4682 fn describe_chain(&self) -> impl Debug {
4683 self.handler.describe_chain()
4684 }
4685}
4686
4687#[must_use = "dropping this guard unregisters the dynamic handler"]
4693#[derive(Debug)]
4694pub struct DynamicHandlerGuard<R: Role> {
4695 uuid: Option<Uuid>,
4696 active: Arc<AtomicBool>,
4697 cx: ConnectionTo<R>,
4698 cleanup: Option<Arc<dyn DynamicHandlerCleanup>>,
4699}
4700
4701pub(crate) trait DynamicHandlerCleanup: std::fmt::Debug + Send + Sync {
4703 fn close(&self);
4704 fn wait(&self) -> futures::future::BoxFuture<'static, ()>;
4705}
4706
4707impl<R: Role> DynamicHandlerGuard<R> {
4708 fn new(uuid: Uuid, active: Arc<AtomicBool>, cx: ConnectionTo<R>) -> Self {
4709 Self {
4710 uuid: Some(uuid),
4711 active,
4712 cx,
4713 cleanup: None,
4714 }
4715 }
4716
4717 #[cfg(feature = "unstable_mcp_over_acp")]
4718 pub(crate) fn with_cleanup(mut self, cleanup: Arc<dyn DynamicHandlerCleanup>) -> Self {
4719 self.cleanup = Some(cleanup);
4720 self
4721 }
4722
4723 pub(crate) fn cleanup(&self) -> Option<Arc<dyn DynamicHandlerCleanup>> {
4724 self.cleanup.clone()
4725 }
4726
4727 pub fn detach(mut self) {
4733 self.uuid = None;
4734 }
4735}
4736
4737impl<R: Role> Drop for DynamicHandlerGuard<R> {
4738 fn drop(&mut self) {
4739 if let Some(uuid) = self.uuid {
4740 self.active.store(false, Ordering::Release);
4741 if let Some(cleanup) = &self.cleanup {
4742 cleanup.close();
4743 }
4744 self.cx.remove_dynamic_handler(uuid);
4745 }
4746 }
4747}
4748
4749#[must_use]
4799pub struct Responder<T: JsonRpcResponse = serde_json::Value> {
4800 method: String,
4802
4803 id: RequestId,
4805
4806 cancellation: ResponderCancellation,
4808
4809 destination: ResponseDestination,
4811
4812 send_fn: Box<dyn FnOnce(Result<T, crate::Error>) -> Result<(), crate::Error> + Send>,
4817
4818 drop_guard: ResponderDropGuard,
4820}
4821
4822struct ResponderDropGuard {
4823 message_tx: OutgoingMessageTx,
4824 id: RequestId,
4825 method: String,
4826 destination: ResponseDestination,
4827 armed: bool,
4828}
4829
4830impl ResponderDropGuard {
4831 fn disarm(&mut self) {
4832 self.armed = false;
4833 }
4834}
4835
4836impl Drop for ResponderDropGuard {
4837 fn drop(&mut self) {
4838 if !self.armed || !self.destination.is_batch() {
4839 return;
4840 }
4841
4842 if let Err(error) = send_raw_message(
4843 &self.message_tx,
4844 OutgoingMessage::AbandonedBatchResponse {
4845 id: self.id.clone(),
4846 method: self.method.clone(),
4847 destination: self.destination.clone(),
4848 },
4849 ) {
4850 tracing::debug!(
4851 id = ?self.id,
4852 method = %self.method,
4853 ?error,
4854 "could not complete abandoned JSON-RPC batch response slot"
4855 );
4856 }
4857 }
4858}
4859
4860impl<T: JsonRpcResponse> std::fmt::Debug for Responder<T> {
4861 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
4862 f.debug_struct("Responder")
4863 .field("method", &self.method)
4864 .field("id", &self.id)
4865 .field("response_type", &std::any::type_name::<T>())
4866 .finish_non_exhaustive()
4867 }
4868}
4869
4870impl Responder<serde_json::Value> {
4871 fn new(
4875 message_tx: OutgoingMessageTx,
4876 method: String,
4877 id: RequestId,
4878 cancellation_registry: &RequestCancellationRegistry,
4879 destination: ResponseDestination,
4880 ) -> Self {
4881 let id_clone = id.clone();
4882 let method_clone = method.clone();
4883 let cancellation = cancellation_registry.register(&id);
4884 let send_destination = destination.clone();
4885 let drop_guard = ResponderDropGuard {
4886 message_tx: message_tx.clone(),
4887 id: id.clone(),
4888 method: method.clone(),
4889 destination: destination.clone(),
4890 armed: true,
4891 };
4892 Self {
4893 method,
4894 id,
4895 cancellation,
4896 destination,
4897 send_fn: Box::new(move |response: Result<serde_json::Value, crate::Error>| {
4898 send_raw_message(
4899 &message_tx,
4900 OutgoingMessage::Response {
4901 id: id_clone,
4902 method: method_clone,
4903 response,
4904 destination: send_destination,
4905 },
4906 )
4907 }),
4908 drop_guard,
4909 }
4910 }
4911
4912 pub fn cast<T: JsonRpcResponse>(self) -> Responder<T> {
4916 self.wrap_params(move |method, value| match value {
4917 Ok(value) => T::into_json(value, method),
4918 Err(e) => Err(e),
4919 })
4920 }
4921}
4922
4923impl<T: JsonRpcResponse> Responder<T> {
4924 #[must_use]
4926 pub fn method(&self) -> &str {
4927 &self.method
4928 }
4929
4930 #[must_use]
4932 pub fn id(&self) -> &RequestId {
4933 &self.id
4934 }
4935
4936 #[must_use]
4945 pub fn cancellation(&self) -> RequestCancellation {
4946 self.cancellation.cancellation()
4947 }
4948
4949 pub fn erase_to_json(self) -> Responder<serde_json::Value> {
4953 self.wrap_params(|method, value| T::from_value(method, value?))
4954 }
4955
4956 pub fn wrap_method(mut self, method: String) -> Responder<T> {
4958 self.drop_guard.method.clone_from(&method);
4959 Responder {
4960 method,
4961 id: self.id,
4962 cancellation: self.cancellation,
4963 destination: self.destination,
4964 send_fn: self.send_fn,
4965 drop_guard: self.drop_guard,
4966 }
4967 }
4968
4969 pub fn wrap_params<U: JsonRpcResponse>(
4974 self,
4975 wrap_fn: impl FnOnce(&str, Result<U, crate::Error>) -> Result<T, crate::Error> + Send + 'static,
4976 ) -> Responder<U> {
4977 let method = self.method.clone();
4978 Responder {
4979 method: self.method,
4980 id: self.id,
4981 cancellation: self.cancellation,
4982 destination: self.destination,
4983 send_fn: Box::new(move |input: Result<U, crate::Error>| {
4984 let t_value = wrap_fn(&method, input);
4985 (self.send_fn)(t_value)
4986 }),
4987 drop_guard: self.drop_guard,
4988 }
4989 }
4990
4991 pub fn respond_with_result(
4993 mut self,
4994 response: Result<T, crate::Error>,
4995 ) -> Result<(), crate::Error> {
4996 tracing::debug!(id = ?self.id, "respond called");
4997 self.drop_guard.disarm();
4998 (self.send_fn)(response)
4999 }
5000
5001 pub fn respond(self, response: T) -> Result<(), crate::Error> {
5003 self.respond_with_result(Ok(response))
5004 }
5005
5006 pub fn respond_with_internal_error(self, message: impl ToString) -> Result<(), crate::Error> {
5008 self.respond_with_error(crate::util::internal_error(message))
5009 }
5010
5011 pub fn respond_with_error(self, error: crate::Error) -> Result<(), crate::Error> {
5013 tracing::debug!(id = ?self.id, ?error, "respond_with_error called");
5014 self.respond_with_result(Err(error))
5015 }
5016
5017 fn reply_target(&self) -> RequestReplyTarget {
5018 RequestReplyTarget {
5019 id: self.id.clone(),
5020 method: self.method.clone(),
5021 destination: self.destination.clone(),
5022 }
5023 }
5024}
5025
5026#[must_use]
5044pub struct ResponseRouter<T: JsonRpcResponse = serde_json::Value> {
5045 method: String,
5047
5048 id: RequestId,
5050
5051 role_id: RoleId,
5054
5055 send_fn: Box<dyn FnOnce(Result<T, crate::Error>) -> Result<(), crate::Error> + Send>,
5057
5058 reply_target: ResponseReplyTarget,
5060}
5061
5062impl<T: JsonRpcResponse> std::fmt::Debug for ResponseRouter<T> {
5063 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
5064 f.debug_struct("ResponseRouter")
5065 .field("method", &self.method)
5066 .field("id", &self.id)
5067 .field("response_type", &std::any::type_name::<T>())
5068 .finish_non_exhaustive()
5069 }
5070}
5071
5072impl ResponseRouter<serde_json::Value> {
5073 fn new(id: RequestId, pending_reply: PendingReply, dispatch: ResponseDispatch) -> Self {
5079 let PendingReply {
5080 method,
5081 role_id,
5082 sender,
5083 cancellation_disarm,
5084 ordering,
5085 response_route_hook,
5086 } = pending_reply;
5087 let reply_target = ResponseReplyTarget {
5088 id: id.clone(),
5089 method: method.clone(),
5090 sender: Arc::new(Mutex::new(Some(sender))),
5091 ordering,
5092 dispatch,
5093 };
5094 let send_target = reply_target.clone();
5095 cancellation_disarm.disarm();
5100 let hook_method = method.clone();
5101 Self {
5102 method,
5103 id,
5104 role_id,
5105 send_fn: Box::new(move |response: Result<serde_json::Value, crate::Error>| {
5106 let response = match response {
5107 Ok(value) => match response_route_hook {
5108 Some(hook) => hook(&hook_method, &value).map(|()| value),
5109 None => Ok(value),
5110 },
5111 Err(error) => Err(error),
5112 };
5113 send_target.route(response);
5114 Ok(())
5115 }),
5116 reply_target,
5117 }
5118 }
5119
5120 pub fn cast<T: JsonRpcResponse>(self) -> ResponseRouter<T> {
5124 self.wrap_params(move |method, value| match value {
5125 Ok(value) => T::into_json(value, method),
5126 Err(e) => Err(e),
5127 })
5128 }
5129}
5130
5131impl<T: JsonRpcResponse> ResponseRouter<T> {
5132 #[must_use]
5134 pub fn method(&self) -> &str {
5135 &self.method
5136 }
5137
5138 #[must_use]
5140 pub fn id(&self) -> &RequestId {
5141 &self.id
5142 }
5143
5144 #[must_use]
5148 pub fn role_id(&self) -> RoleId {
5149 self.role_id.clone()
5150 }
5151
5152 pub fn erase_to_json(self) -> ResponseRouter<serde_json::Value> {
5156 self.wrap_params(|method, value| T::from_value(method, value?))
5157 }
5158
5159 fn wrap_params<U: JsonRpcResponse>(
5164 self,
5165 wrap_fn: impl FnOnce(&str, Result<U, crate::Error>) -> Result<T, crate::Error> + Send + 'static,
5166 ) -> ResponseRouter<U> {
5167 let method = self.method.clone();
5168 ResponseRouter {
5169 method: self.method,
5170 id: self.id,
5171 role_id: self.role_id,
5172 send_fn: Box::new(move |input: Result<U, crate::Error>| {
5173 let t_value = wrap_fn(&method, input);
5174 (self.send_fn)(t_value)
5175 }),
5176 reply_target: self.reply_target,
5177 }
5178 }
5179
5180 pub fn route_with_result(self, response: Result<T, crate::Error>) -> Result<(), crate::Error> {
5182 tracing::debug!(id = ?self.id, "response routed to awaiter");
5183 (self.send_fn)(response)
5184 }
5185
5186 pub fn route(self, response: T) -> Result<(), crate::Error> {
5188 self.route_with_result(Ok(response))
5189 }
5190
5191 pub fn route_with_internal_error(self, message: impl ToString) -> Result<(), crate::Error> {
5193 self.route_with_error(crate::util::internal_error(message))
5194 }
5195
5196 pub fn route_with_error(self, error: crate::Error) -> Result<(), crate::Error> {
5198 tracing::debug!(id = ?self.id, ?error, "error routed to awaiter");
5199 self.route_with_result(Err(error))
5200 }
5201}
5202
5203pub trait JsonRpcMessage: 'static + Debug + Sized + Send + Clone {
5211 fn matches_method(method: &str) -> bool;
5213
5214 fn method(&self) -> &str;
5216
5217 fn to_untyped_message(&self) -> Result<UntypedMessage, crate::Error>;
5219
5220 fn parse_message(method: &str, params: &impl Serialize) -> Result<Self, crate::Error>;
5225}
5226
5227pub trait JsonRpcResponse: 'static + Debug + Sized + Send + Clone {
5243 fn into_json(self, method: &str) -> Result<serde_json::Value, crate::Error>;
5245
5246 fn from_value(method: &str, value: serde_json::Value) -> Result<Self, crate::Error>;
5248}
5249
5250impl JsonRpcResponse for serde_json::Value {
5251 fn from_value(_method: &str, value: serde_json::Value) -> Result<Self, crate::Error> {
5252 Ok(value)
5253 }
5254
5255 fn into_json(self, _method: &str) -> Result<serde_json::Value, crate::Error> {
5256 Ok(self)
5257 }
5258}
5259
5260pub trait JsonRpcNotification: JsonRpcMessage {}
5277
5278pub trait JsonRpcRequest: JsonRpcMessage {
5300 type Response: JsonRpcResponse;
5302}
5303
5304#[derive(Debug)]
5312pub enum Dispatch<Req: JsonRpcRequest = UntypedMessage, Notif: JsonRpcNotification = UntypedMessage>
5313{
5314 Request(Req, Responder<Req::Response>),
5316
5317 Notification(Notif),
5319
5320 Response(
5326 Result<Req::Response, crate::Error>,
5327 ResponseRouter<Req::Response>,
5328 ),
5329}
5330
5331impl<Req: JsonRpcRequest, Notif: JsonRpcNotification> Dispatch<Req, Notif> {
5332 pub fn map<Req1, Notif1>(
5337 self,
5338 map_request: impl FnOnce(Req, Responder<Req::Response>) -> (Req1, Responder<Req1::Response>),
5339 map_notification: impl FnOnce(Notif) -> Notif1,
5340 ) -> Dispatch<Req1, Notif1>
5341 where
5342 Req1: JsonRpcRequest<Response = Req::Response>,
5343 Notif1: JsonRpcNotification,
5344 {
5345 match self {
5346 Dispatch::Request(request, responder) => {
5347 let (new_request, new_responder) = map_request(request, responder);
5348 Dispatch::Request(new_request, new_responder)
5349 }
5350 Dispatch::Notification(notification) => {
5351 let new_notification = map_notification(notification);
5352 Dispatch::Notification(new_notification)
5353 }
5354 Dispatch::Response(result, router) => Dispatch::Response(result, router),
5355 }
5356 }
5357
5358 pub fn to_untyped_message(&self) -> Result<UntypedMessage, crate::Error> {
5363 match self {
5364 Dispatch::Request(request, _) => request.to_untyped_message(),
5365 Dispatch::Notification(notification) => notification.to_untyped_message(),
5366 Dispatch::Response(_, _) => Err(crate::util::internal_error(
5367 "Response variant has no untyped message representation",
5368 )),
5369 }
5370 }
5371
5372 pub fn into_untyped_dispatch(self) -> Result<Dispatch, crate::Error> {
5376 match self {
5377 Dispatch::Request(request, responder) => Ok(Dispatch::Request(
5378 request.to_untyped_message()?,
5379 responder.erase_to_json(),
5380 )),
5381 Dispatch::Notification(notification) => {
5382 Ok(Dispatch::Notification(notification.to_untyped_message()?))
5383 }
5384 Dispatch::Response(_, _) => Err(crate::util::internal_error(
5385 "cannot convert Response variant to untyped message context",
5386 )),
5387 }
5388 }
5389
5390 pub fn id(&self) -> Option<&RequestId> {
5392 match self {
5393 Dispatch::Request(_, cx) => Some(cx.id()),
5394 Dispatch::Notification(_) => None,
5395 Dispatch::Response(_, cx) => Some(cx.id()),
5396 }
5397 }
5398
5399 fn handler_error_target(&self) -> Option<HandlerErrorTarget> {
5400 match self {
5401 Dispatch::Request(_, responder) => {
5402 Some(HandlerErrorTarget::Request(responder.reply_target()))
5403 }
5404 Dispatch::Notification(_) => None,
5405 Dispatch::Response(_, router) => {
5406 Some(HandlerErrorTarget::Response(router.reply_target.clone()))
5407 }
5408 }
5409 }
5410
5411 pub fn method(&self) -> &str {
5416 match self {
5417 Dispatch::Request(msg, _) => msg.method(),
5418 Dispatch::Notification(msg) => msg.method(),
5419 Dispatch::Response(_, cx) => cx.method(),
5420 }
5421 }
5422}
5423
5424impl Dispatch {
5425 #[tracing::instrument(skip(self), fields(Request = ?std::any::type_name::<Req>(), Notif = ?std::any::type_name::<Notif>()), level = "trace", ret)]
5433 pub(crate) fn into_typed_dispatch<Req: JsonRpcRequest, Notif: JsonRpcNotification>(
5434 self,
5435 ) -> Result<Result<Dispatch<Req, Notif>, Dispatch>, crate::Error> {
5436 tracing::debug!(
5437 message = ?self,
5438 "into_typed_dispatch"
5439 );
5440 match self {
5441 Dispatch::Request(message, responder) => {
5442 if Req::matches_method(&message.method) {
5443 match Req::parse_message(&message.method, &message.params) {
5444 Ok(req) => {
5445 tracing::trace!(?req, "parsed ok");
5446 Ok(Ok(Dispatch::Request(req, responder.cast())))
5447 }
5448 Err(err) => {
5449 tracing::trace!(?err, "parse error");
5450 Err(err)
5451 }
5452 }
5453 } else {
5454 tracing::trace!("method doesn't match");
5455 Ok(Err(Dispatch::Request(message, responder)))
5456 }
5457 }
5458
5459 Dispatch::Notification(message) => {
5460 if Notif::matches_method(&message.method) {
5461 match Notif::parse_message(&message.method, &message.params) {
5462 Ok(notif) => {
5463 tracing::trace!(?notif, "parse ok");
5464 Ok(Ok(Dispatch::Notification(notif)))
5465 }
5466 Err(err) => {
5467 tracing::trace!(?err, "parse error");
5468 Err(err)
5469 }
5470 }
5471 } else {
5472 tracing::trace!("method doesn't match");
5473 Ok(Err(Dispatch::Notification(message)))
5474 }
5475 }
5476
5477 Dispatch::Response(result, cx) => {
5478 let method = cx.method();
5479 if Req::matches_method(method) {
5480 let typed_result = match result {
5482 Ok(value) => {
5483 match <Req::Response as JsonRpcResponse>::from_value(method, value) {
5484 Ok(parsed) => {
5485 tracing::trace!(?parsed, "parse ok");
5486 Ok(parsed)
5487 }
5488 Err(err) => {
5489 tracing::trace!(?err, "parse error");
5490 return Err(err);
5491 }
5492 }
5493 }
5494 Err(err) => {
5495 tracing::trace!("error, passthrough");
5496 Err(err)
5497 }
5498 };
5499 Ok(Ok(Dispatch::Response(typed_result, cx.cast())))
5500 } else {
5501 tracing::trace!("method doesn't match");
5502 Ok(Err(Dispatch::Response(result, cx)))
5503 }
5504 }
5505 }
5506 }
5507
5508 #[must_use]
5512 pub fn has_field(&self, field_name: &str) -> bool {
5513 self.message()
5514 .and_then(|m| m.params().get(field_name))
5515 .is_some()
5516 }
5517
5518 pub(crate) fn has_session_id(&self) -> bool {
5522 self.has_field("sessionId")
5523 }
5524
5525 pub(crate) fn get_session_id(&self) -> Result<Option<SessionId>, crate::Error> {
5529 let Some(message) = self.message() else {
5530 return Ok(None);
5531 };
5532 let Some(value) = message.params().get("sessionId") else {
5533 return Ok(None);
5534 };
5535 let session_id = serde_json::from_value(value.clone())?;
5536 Ok(Some(session_id))
5537 }
5538
5539 pub fn into_notification<N: JsonRpcNotification>(
5547 self,
5548 ) -> Result<Result<N, Dispatch>, crate::Error> {
5549 match self {
5550 Dispatch::Notification(msg) => {
5551 if !N::matches_method(&msg.method) {
5552 return Ok(Err(Dispatch::Notification(msg)));
5553 }
5554 match N::parse_message(&msg.method, &msg.params) {
5555 Ok(n) => Ok(Ok(n)),
5556 Err(err) => Err(err),
5557 }
5558 }
5559 Dispatch::Request(..) | Dispatch::Response(..) => Ok(Err(self)),
5560 }
5561 }
5562
5563 pub fn into_request<Req: JsonRpcRequest>(
5571 self,
5572 ) -> Result<Result<(Req, Responder<Req::Response>), Dispatch>, crate::Error> {
5573 match self {
5574 Dispatch::Request(msg, responder) => {
5575 if !Req::matches_method(&msg.method) {
5576 return Ok(Err(Dispatch::Request(msg, responder)));
5577 }
5578 match Req::parse_message(&msg.method, &msg.params) {
5579 Ok(req) => Ok(Ok((req, responder.cast()))),
5580 Err(err) => Err(err),
5581 }
5582 }
5583 Dispatch::Notification(..) | Dispatch::Response(..) => Ok(Err(self)),
5584 }
5585 }
5586}
5587
5588impl<M: JsonRpcRequest + JsonRpcNotification> Dispatch<M, M> {
5589 pub fn message(&self) -> Option<&M> {
5593 match self {
5594 Dispatch::Request(msg, _) | Dispatch::Notification(msg) => Some(msg),
5595 Dispatch::Response(_, _) => None,
5596 }
5597 }
5598
5599 pub(crate) fn try_map_message(
5603 self,
5604 map_message: impl FnOnce(M) -> Result<M, crate::Error>,
5605 ) -> Result<Dispatch<M, M>, crate::Error> {
5606 match self {
5607 Dispatch::Request(request, cx) => Ok(Dispatch::Request(map_message(request)?, cx)),
5608 Dispatch::Notification(notification) => {
5609 Ok(Dispatch::<M, M>::Notification(map_message(notification)?))
5610 }
5611 Dispatch::Response(result, cx) => Ok(Dispatch::Response(result, cx)),
5612 }
5613 }
5614}
5615
5616#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
5618pub struct UntypedMessage {
5619 pub method: String,
5621 pub params: serde_json::Value,
5623}
5624
5625impl UntypedMessage {
5626 pub fn new(method: &str, params: impl Serialize) -> Result<Self, crate::Error> {
5628 let params = serde_json::to_value(params)?;
5629 Ok(Self {
5630 method: method.to_string(),
5631 params,
5632 })
5633 }
5634
5635 #[must_use]
5637 pub fn method(&self) -> &str {
5638 &self.method
5639 }
5640
5641 #[must_use]
5643 pub fn params(&self) -> &serde_json::Value {
5644 &self.params
5645 }
5646
5647 #[must_use]
5649 pub fn into_parts(self) -> (String, serde_json::Value) {
5650 (self.method, self.params)
5651 }
5652
5653 pub(crate) fn into_raw_jsonrpc_message(
5655 self,
5656 id: Option<RequestId>,
5657 ) -> Result<RawJsonRpcMessage, crate::Error> {
5658 let Self { method, params } = self;
5659 match id {
5660 Some(id) => RawJsonRpcMessage::request(method, params, id),
5661 None => RawJsonRpcMessage::notification(method, params),
5662 }
5663 }
5664}
5665
5666impl JsonRpcMessage for UntypedMessage {
5667 fn matches_method(_method: &str) -> bool {
5668 true
5670 }
5671
5672 fn method(&self) -> &str {
5673 &self.method
5674 }
5675
5676 fn to_untyped_message(&self) -> Result<UntypedMessage, crate::Error> {
5677 Ok(self.clone())
5678 }
5679
5680 fn parse_message(method: &str, params: &impl Serialize) -> Result<Self, crate::Error> {
5681 UntypedMessage::new(method, params)
5682 }
5683}
5684
5685impl JsonRpcRequest for UntypedMessage {
5686 type Response = serde_json::Value;
5687}
5688
5689impl JsonRpcNotification for UntypedMessage {}
5690
5691#[must_use = "dropping a SentRequest asks the peer to cancel the request and \
5792 discards the response; consume it with `block_task`, \
5793 `on_receiving_result`, `forward_response_to`, or `detach`"]
5794pub struct SentRequest<T> {
5795 id: RequestId,
5796 method: String,
5797 task_tx: TaskTx,
5798 response_rx: oneshot::Receiver<ResponsePayload>,
5799 to_result: Box<dyn FnOnce(serde_json::Value) -> Result<T, crate::Error> + Send>,
5800 cancellation: SentRequestCancellation,
5801 response_ordering: ResponseOrdering,
5802 cancellation_sources: Vec<RequestCancellation>,
5806}
5807
5808#[derive(Clone, Debug)]
5809pub(crate) struct SentRequestCancellationDisarm {
5810 armed: Arc<AtomicBool>,
5811}
5812
5813impl SentRequestCancellationDisarm {
5814 fn new() -> Self {
5815 Self {
5816 armed: Arc::new(AtomicBool::new(true)),
5817 }
5818 }
5819
5820 fn disarm(&self) {
5821 self.armed.store(false, Ordering::Release);
5822 }
5823}
5824
5825struct SentRequestCancellation {
5826 message_tx: OutgoingMessageTx,
5827 remote_style: crate::role::RemoteStyle,
5828 request_id: RequestId,
5829 disarm: SentRequestCancellationDisarm,
5830}
5831
5832impl SentRequestCancellation {
5833 fn new(
5834 message_tx: OutgoingMessageTx,
5835 remote_style: crate::role::RemoteStyle,
5836 request_id: RequestId,
5837 ) -> Self {
5838 Self {
5839 message_tx,
5840 remote_style,
5841 request_id,
5842 disarm: SentRequestCancellationDisarm::new(),
5843 }
5844 }
5845
5846 fn disarm(&self) {
5847 self.disarm.disarm();
5848 }
5849
5850 fn disarm_handle(&self) -> SentRequestCancellationDisarm {
5851 self.disarm.clone()
5852 }
5853
5854 fn send(&self) -> Result<(), crate::Error> {
5855 if !self.disarm.armed.swap(false, Ordering::AcqRel) {
5856 return Ok(());
5857 }
5858
5859 let untyped = self.remote_style.transform_outgoing_message(
5862 crate::schema::v1::CancelRequestNotification::new(self.request_id.clone()),
5863 )?;
5864
5865 send_raw_message(&self.message_tx, OutgoingMessage::Notification { untyped })
5866 }
5867}
5868
5869impl Drop for SentRequestCancellation {
5870 fn drop(&mut self) {
5871 if let Err(error) = self.send() {
5872 tracing::debug!(?error, "failed to auto-cancel dropped request");
5873 }
5874 }
5875}
5876
5877impl Debug for SentRequestCancellation {
5878 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
5879 f.debug_struct("SentRequestCancellation")
5880 .field("request_id", &self.request_id)
5881 .field("remote_style", &self.remote_style)
5882 .field("armed", &self.disarm.armed.load(Ordering::Acquire))
5883 .finish_non_exhaustive()
5884 }
5885}
5886
5887async fn await_response_forwarding_cancellation(
5898 response_rx: oneshot::Receiver<ResponsePayload>,
5899 cancellation: &SentRequestCancellation,
5900 sources: &[RequestCancellation],
5901) -> Result<ResponsePayload, oneshot::Canceled> {
5902 let forward_cancellation = || {
5906 if let Err(error) = cancellation.send() {
5907 tracing::debug!(
5908 ?error,
5909 "failed to forward cancellation to downstream request"
5910 );
5911 }
5912 };
5913
5914 let response = if sources.is_empty() {
5915 response_rx.await
5916 } else if sources.iter().any(RequestCancellation::is_cancelled) {
5917 forward_cancellation();
5918 response_rx.await
5919 } else {
5920 let cancelled = sources.iter().map(|source| source.state.signal_rx.clone());
5921 match future::select(future::select_all(cancelled), response_rx).await {
5922 Either::Left((_, response_rx)) => {
5923 forward_cancellation();
5924 response_rx.await
5925 }
5926 Either::Right((response, _)) => response,
5927 }
5928 };
5929
5930 cancellation.disarm();
5931 response
5932}
5933
5934impl<T: Debug> Debug for SentRequest<T> {
5935 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
5936 let mut debug = f.debug_struct("SentRequest");
5937 debug
5938 .field("id", &self.id)
5939 .field("method", &self.method)
5940 .field("task_tx", &self.task_tx)
5941 .field("response_rx", &self.response_rx);
5942 debug
5943 .field("cancellation", &self.cancellation)
5944 .field("cancellation_sources", &self.cancellation_sources);
5945 debug.finish_non_exhaustive()
5946 }
5947}
5948
5949impl SentRequest<serde_json::Value> {
5950 fn new(
5951 id: RequestId,
5952 method: String,
5953 task_tx: mpsc::UnboundedSender<Task>,
5954 response_rx: oneshot::Receiver<ResponsePayload>,
5955 cancellation: SentRequestCancellation,
5956 response_ordering: ResponseOrdering,
5957 ) -> Self {
5958 Self {
5959 id,
5960 method,
5961 response_rx,
5962 task_tx,
5963 to_result: Box::new(Ok),
5964 cancellation,
5965 response_ordering,
5966 cancellation_sources: Vec::new(),
5967 }
5968 }
5969}
5970
5971impl<T> SentRequest<T> {
5972 pub fn detach(self) {
5985 self.cancellation.disarm();
5986 }
5987
5988 pub fn cancel(&self) -> Result<(), crate::Error> {
6004 self.cancellation.send()
6005 }
6006
6007 pub fn forward_cancellation_from(mut self, source: RequestCancellation) -> Self {
6046 self.cancellation_sources.push(source);
6047 self
6048 }
6049}
6050
6051impl<T> SentRequest<T> {
6052 #[must_use]
6054 pub fn id(&self) -> &RequestId {
6055 &self.id
6056 }
6057
6058 #[must_use]
6060 pub fn method(&self) -> &str {
6061 &self.method
6062 }
6063
6064 pub fn map<U>(
6073 self,
6074 map_fn: impl FnOnce(T) -> Result<U, crate::Error> + 'static + Send,
6075 ) -> SentRequest<U>
6076 where
6077 T: 'static,
6078 {
6079 SentRequest {
6080 id: self.id,
6081 method: self.method,
6082 response_rx: self.response_rx,
6083 task_tx: self.task_tx,
6084 to_result: Box::new(move |value| map_fn((self.to_result)(value)?)),
6085 cancellation: self.cancellation,
6086 response_ordering: self.response_ordering,
6087 cancellation_sources: self.cancellation_sources,
6088 }
6089 }
6090
6091 #[track_caller]
6154 pub fn forward_response_to(self, responder: Responder<T>) -> Result<(), crate::Error>
6155 where
6156 T: JsonRpcResponse,
6157 {
6158 let this = self.forward_cancellation_from(responder.cancellation());
6159
6160 this.consume_with(async move |response| {
6161 responder.respond_with_result(response.unwrap_or_else(Err))
6164 })
6165 }
6166
6167 #[track_caller]
6181 fn consume_with<F>(
6182 self,
6183 handle: impl FnOnce(Result<Result<T, crate::Error>, crate::Error>) -> F + 'static + Send,
6184 ) -> Result<(), crate::Error>
6185 where
6186 T: 'static,
6187 F: Future<Output = Result<(), crate::Error>> + 'static + Send,
6188 {
6189 self.response_ordering.mark_ordered();
6190 let task_tx = self.task_tx.clone();
6191 let method = self.method;
6192 let response_rx = self.response_rx;
6193 let to_result = self.to_result;
6194 let cancellation = self.cancellation;
6195 let cancellation_sources = self.cancellation_sources;
6196 let location = Location::caller();
6197
6198 Task::new(location, async move {
6199 let response = await_response_forwarding_cancellation(
6200 response_rx,
6201 &cancellation,
6202 &cancellation_sources,
6203 )
6204 .await;
6205
6206 match response {
6207 Ok(ResponsePayload { result, ack_tx }) => {
6208 let typed_result = match result {
6210 Ok(json_value) => to_result(json_value),
6211 Err(err) => Err(err),
6212 };
6213
6214 let outcome = handle(Ok(typed_result)).await;
6215
6216 if let Some(tx) = ack_tx {
6220 let _ = tx.send(());
6221 }
6222
6223 outcome
6224 }
6225 Err(err) => {
6226 handle(Err(crate::util::internal_error(format!(
6227 "response to `{method}` never received: {err}"
6228 ))))
6229 .await
6230 }
6231 }
6232 })
6233 .spawn(&task_tx)
6234 }
6235
6236 pub async fn block_task(self) -> Result<T, crate::Error> {
6300 let response = await_response_forwarding_cancellation(
6301 self.response_rx,
6302 &self.cancellation,
6303 &self.cancellation_sources,
6304 )
6305 .await;
6306
6307 match response {
6308 Ok(ResponsePayload {
6309 result: Ok(json_value),
6310 ack_tx,
6311 }) => {
6312 if let Some(tx) = ack_tx {
6315 let _ = tx.send(());
6316 }
6317 match (self.to_result)(json_value) {
6318 Ok(value) => Ok(value),
6319 Err(err) => Err(err),
6320 }
6321 }
6322 Ok(ResponsePayload {
6323 result: Err(err),
6324 ack_tx,
6325 }) => {
6326 if let Some(tx) = ack_tx {
6327 let _ = tx.send(());
6328 }
6329 Err(err)
6330 }
6331 Err(err) => Err(crate::util::internal_error(format!(
6332 "response to `{}` never received: {}",
6333 self.method, err
6334 ))),
6335 }
6336 }
6337
6338 pub(crate) async fn block_task_with_ordered_result<U>(
6346 self,
6347 transform: impl FnOnce(Result<T, crate::Error>) -> Result<U, crate::Error>,
6348 ) -> Result<U, crate::Error> {
6349 let response = await_response_forwarding_cancellation(
6350 self.response_rx,
6351 &self.cancellation,
6352 &self.cancellation_sources,
6353 )
6354 .await;
6355
6356 let (result, ack_tx) = match response {
6357 Ok(ResponsePayload { result, ack_tx }) => {
6358 let typed_result = match result {
6359 Ok(json_value) => (self.to_result)(json_value),
6360 Err(error) => Err(error),
6361 };
6362 (typed_result, ack_tx)
6363 }
6364 Err(error) => (
6365 Err(crate::util::internal_error(format!(
6366 "response to `{}` never received: {error}",
6367 self.method
6368 ))),
6369 None,
6370 ),
6371 };
6372
6373 let outcome = transform(result);
6374 if let Some(acknowledgment) = ack_tx {
6375 let _ = acknowledgment.send(());
6376 }
6377 outcome
6378 }
6379
6380 #[track_caller]
6436 pub fn on_receiving_ok_result<F>(
6437 self,
6438 responder: Responder<T>,
6439 task: impl FnOnce(T, Responder<T>) -> F + 'static + Send,
6440 ) -> Result<(), crate::Error>
6441 where
6442 F: Future<Output = Result<(), crate::Error>> + 'static + Send,
6443 T: JsonRpcResponse,
6444 {
6445 self.on_receiving_result(async move |result| match result {
6446 Ok(value) => task(value, responder).await,
6447 Err(err) => responder.respond_with_error(err),
6448 })
6449 }
6450
6451 #[track_caller]
6530 pub fn on_receiving_result<F>(
6531 self,
6532 task: impl FnOnce(Result<T, crate::Error>) -> F + 'static + Send,
6533 ) -> Result<(), crate::Error>
6534 where
6535 T: 'static,
6536 F: Future<Output = Result<(), crate::Error>> + 'static + Send,
6537 {
6538 self.consume_with(move |response| match response {
6539 Ok(result) => Either::Left(task(result)),
6542 Err(err) => Either::Right(future::ready(Err(err))),
6545 })
6546 }
6547}
6548
6549#[derive(Debug)]
6580pub struct Lines<OutgoingSink, IncomingStream> {
6581 outgoing: OutgoingSink,
6582 incoming: IncomingStream,
6583}
6584
6585impl<OutgoingSink, IncomingStream> Lines<OutgoingSink, IncomingStream>
6586where
6587 OutgoingSink: futures::Sink<String, Error = std::io::Error> + Send + 'static,
6588 IncomingStream: futures::Stream<Item = std::io::Result<String>> + Send + 'static,
6589{
6590 pub fn new(outgoing: OutgoingSink, incoming: IncomingStream) -> Self {
6592 Self { outgoing, incoming }
6593 }
6594
6595 fn into_channel_transport(self) -> (Channel, crate::ConnectionDriver) {
6596 let Self { outgoing, incoming } = self;
6597 let (channel_for_caller, channel_for_lines) = Channel::duplex();
6598 let Channel { mut rx, tx } = channel_for_lines;
6599 let (finish_tx, finish_rx) = oneshot::channel();
6600 let finish = async move {
6601 if finish_rx.await.is_err() {
6603 future::pending::<()>().await;
6604 }
6605 }
6606 .boxed()
6607 .shared();
6608 let outgoing_frames = futures::stream::poll_fn({
6609 let mut finish = finish.clone();
6610 let mut finishing = false;
6611 move |cx| {
6612 if !finishing && std::pin::Pin::new(&mut finish).poll(cx).is_ready() {
6613 rx.close();
6614 finishing = true;
6615 }
6616 rx.poll_next_unpin(cx)
6617 }
6618 });
6619 let discard_incoming = Arc::new(AtomicBool::new(false));
6620 let incoming = incoming.filter_map({
6621 let discard_incoming = discard_incoming.clone();
6622 move |item| {
6623 let discard = discard_incoming.load(Ordering::Acquire);
6624 future::ready((!discard || item.is_err()).then_some(item))
6625 }
6626 });
6627 let outgoing = transport_actor::transport_outgoing_lines_actor(outgoing_frames, outgoing)
6628 .boxed()
6629 .shared();
6630 let serve_self = Box::pin({
6631 let outgoing = outgoing.clone();
6632 async move {
6633 futures::try_join!(
6634 outgoing,
6635 transport_actor::transport_incoming_lines_actor(incoming, tx),
6636 )?;
6637 Ok(())
6638 }
6639 });
6640 let server_future = crate::ConnectionDriver::with_finish(
6641 async move {
6642 match future::select(finish, serve_self).await {
6643 Either::Left(((), serve_self)) => {
6644 discard_incoming.store(true, Ordering::Release);
6645 match future::select(serve_self, outgoing).await {
6648 Either::Left((result, _)) | Either::Right((result, _)) => result,
6649 }
6650 }
6651 Either::Right((result, _)) => result,
6652 }
6653 },
6654 move || {
6655 let _ = finish_tx.send(());
6656 },
6657 );
6658
6659 (channel_for_caller, server_future)
6660 }
6661}
6662
6663impl<OutgoingSink, IncomingStream, R: Role> ConnectTo<R> for Lines<OutgoingSink, IncomingStream>
6664where
6665 OutgoingSink: futures::Sink<String, Error = std::io::Error> + Send + 'static,
6666 IncomingStream: futures::Stream<Item = std::io::Result<String>> + Send + 'static,
6667{
6668 async fn connect_to(self, client: impl ConnectTo<R::Counterpart>) -> Result<(), crate::Error> {
6669 let (channel, mut serve_self) = self.into_channel_transport();
6670 let mut finish = serve_self
6671 .take_finish()
6672 .expect("built-in Lines transport supports explicit finishing");
6673 let client_future = Box::pin(ConnectTo::<R>::connect_to(channel, client));
6674
6675 match futures::future::select(client_future, serve_self).await {
6676 Either::Left((result, serve_self)) => {
6677 result?;
6678 finish.request();
6681 serve_self.await
6682 }
6683 Either::Right((result, _)) => result,
6684 }
6685 }
6686
6687 fn into_channel_and_future(self) -> (Channel, Option<crate::ConnectionDriver>) {
6688 let (channel, driver) = self.into_channel_transport();
6689 (channel, Some(driver))
6690 }
6691}
6692
6693#[derive(Debug)]
6732pub struct ByteStreams<OB, IB> {
6733 outgoing: OB,
6734 incoming: IB,
6735}
6736
6737impl<OB, IB> ByteStreams<OB, IB>
6738where
6739 OB: AsyncWrite + Send + 'static,
6740 IB: AsyncRead + Send + 'static,
6741{
6742 pub fn new(outgoing: OB, incoming: IB) -> Self {
6744 Self { outgoing, incoming }
6745 }
6746
6747 fn into_lines(
6748 self,
6749 ) -> Lines<
6750 impl futures::Sink<String, Error = std::io::Error> + Send + 'static,
6751 impl futures::Stream<Item = std::io::Result<String>> + Send + 'static,
6752 > {
6753 use futures::AsyncBufReadExt;
6754 use futures::io::BufReader;
6755 let Self { outgoing, incoming } = self;
6756
6757 let incoming_lines = Box::pin(BufReader::new(incoming).lines());
6758 let outgoing_lines = transport_actor::LineWriter::new(outgoing);
6759
6760 Lines::new(outgoing_lines, incoming_lines)
6761 }
6762}
6763
6764#[cfg(any(
6765 all(
6766 any(feature = "process", feature = "stdio"),
6767 not(target_family = "wasm")
6768 ),
6769 test
6770))]
6771pub(crate) async fn write_line<W>(writer: &mut W, line: String) -> std::io::Result<()>
6772where
6773 W: AsyncWrite + Unpin + ?Sized,
6774{
6775 use futures::AsyncWriteExt as _;
6776
6777 let mut bytes = line.into_bytes();
6778 bytes.push(b'\n');
6779 writer.write_all(&bytes).await?;
6780 writer.flush().await
6781}
6782
6783impl<OB, IB, R: Role> ConnectTo<R> for ByteStreams<OB, IB>
6784where
6785 OB: AsyncWrite + Send + 'static,
6786 IB: AsyncRead + Send + 'static,
6787{
6788 async fn connect_to(self, client: impl ConnectTo<R::Counterpart>) -> Result<(), crate::Error> {
6789 ConnectTo::<R>::connect_to(self.into_lines(), client).await
6790 }
6791
6792 fn into_channel_and_future(self) -> (Channel, Option<crate::ConnectionDriver>) {
6793 ConnectTo::<R>::into_channel_and_future(self.into_lines())
6794 }
6795}
6796
6797#[derive(Debug)]
6820pub struct Channel {
6821 pub rx: mpsc::UnboundedReceiver<TransportFrame>,
6823 pub tx: mpsc::UnboundedSender<TransportFrame>,
6825}
6826
6827impl Channel {
6828 #[must_use]
6832 pub fn duplex() -> (Self, Self) {
6833 let (a_tx, b_rx) = mpsc::unbounded();
6834 let (b_tx, a_rx) = mpsc::unbounded();
6835
6836 (Self { rx: a_rx, tx: a_tx }, Self { rx: b_rx, tx: b_tx })
6837 }
6838
6839 pub(crate) async fn copy(mut self) -> Result<(), crate::Error> {
6845 while let Some(frame) = self.rx.next().await {
6846 self.tx
6847 .unbounded_send(frame)
6848 .map_err(crate::util::internal_error)?;
6849 }
6850 Ok(())
6851 }
6852
6853 pub(crate) async fn copy_with_driver(
6856 self,
6857 driver: Option<crate::ConnectionDriver>,
6858 ) -> Result<(), crate::Error> {
6859 self.copy_with_driver_until(driver, future::pending()).await
6860 }
6861
6862 pub(crate) async fn copy_with_driver_until(
6865 mut self,
6866 mut driver: Option<crate::ConnectionDriver>,
6867 stop_delivery: impl Future<Output = ()>,
6868 ) -> Result<(), crate::Error> {
6869 let mut stop_delivery = pin!(stop_delivery);
6870 let mut delivering = true;
6871 let mut done = false;
6872 loop {
6873 let event = future::poll_fn(|cx| {
6874 if delivering && stop_delivery.as_mut().poll(cx).is_ready() {
6875 delivering = false;
6876 }
6877 if !done
6879 && let Some(driver) = driver.as_mut()
6880 && let std::task::Poll::Ready(result) = std::pin::Pin::new(driver).poll(cx)
6881 {
6882 return std::task::Poll::Ready(Either::Left(result));
6883 }
6884 if !delivering && driver.is_none() {
6885 return std::task::Poll::Ready(Either::Right(None));
6886 }
6887 self.rx.poll_next_unpin(cx).map(Either::Right)
6888 })
6889 .await;
6890 let frame = match event {
6891 Either::Left(result) => {
6892 result?;
6893 done = true;
6894 self.rx.close();
6895 continue;
6896 }
6897 Either::Right(frame) => frame,
6898 };
6899 let Some(frame) = frame else {
6900 break;
6901 };
6902 if delivering {
6903 self.tx
6904 .unbounded_send(frame)
6905 .map_err(crate::util::internal_error)?;
6906 }
6907 }
6908 drop(self);
6910 if !done && let Some(driver) = driver {
6911 driver.await?;
6912 }
6913 Ok(())
6914 }
6915
6916 pub async fn bridge_with_inspection(
6926 left: Self,
6927 right: Self,
6928 mut left_to_right: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send,
6929 mut right_to_left: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send,
6930 ) -> Result<(), crate::Error> {
6931 let Self {
6932 rx: mut left_rx,
6933 tx: left_tx,
6934 } = left;
6935 let Self {
6936 rx: mut right_rx,
6937 tx: right_tx,
6938 } = right;
6939
6940 let left_to_right = async move {
6941 while let Some(frame) = left_rx.next().await {
6942 frame.inspect_messages(&mut left_to_right)?;
6943 right_tx
6944 .unbounded_send(frame)
6945 .map_err(crate::util::internal_error)?;
6946 }
6947 Ok::<(), crate::Error>(())
6948 };
6949 let right_to_left = async move {
6950 while let Some(frame) = right_rx.next().await {
6951 frame.inspect_messages(&mut right_to_left)?;
6952 left_tx
6953 .unbounded_send(frame)
6954 .map_err(crate::util::internal_error)?;
6955 }
6956 Ok::<(), crate::Error>(())
6957 };
6958
6959 futures::try_join!(left_to_right, right_to_left)?;
6960 Ok(())
6961 }
6962}
6963
6964impl<R: Role> ConnectTo<R> for Channel {
6965 async fn connect_to(self, client: impl ConnectTo<R::Counterpart>) -> Result<(), crate::Error> {
6966 let (client_channel, client_future) = client.into_channel_and_future();
6967
6968 let passive = client_future.is_none();
6969 let outgoing = Box::pin(
6970 Channel {
6971 rx: client_channel.rx,
6972 tx: self.tx,
6973 }
6974 .copy_with_driver(client_future),
6975 );
6976 let incoming = Box::pin(
6977 Channel {
6978 rx: self.rx,
6979 tx: client_channel.tx,
6980 }
6981 .copy(),
6982 );
6983 if passive {
6984 futures::try_join!(outgoing, incoming)?;
6985 return Ok(());
6986 }
6987
6988 match future::select(outgoing, incoming).await {
6989 Either::Left((result, _)) => result,
6990 Either::Right((result, outgoing)) => {
6991 result?;
6992 outgoing.await
6993 }
6994 }
6995 }
6996
6997 fn into_channel_and_future(self) -> (Channel, Option<crate::ConnectionDriver>) {
6998 (self, None)
6999 }
7000}
7001
7002#[cfg(test)]
7003mod tests {
7004 use super::*;
7005
7006 #[test]
7007 fn protected_cleanup_keeps_scoped_runners_polled_on_every_shutdown_path() {
7008 #[derive(Clone, Copy, Debug)]
7009 enum Stop {
7010 ForegroundSuccess,
7011 ForegroundError,
7012 InputEof,
7013 TransportError,
7014 TaskError,
7015 RunnerError,
7016 SupervisorError,
7017 }
7018
7019 struct Dropped(Arc<AtomicBool>);
7020 impl Drop for Dropped {
7021 fn drop(&mut self) {
7022 self.0.store(true, Ordering::Release);
7023 }
7024 }
7025
7026 for stop in [
7027 Stop::ForegroundSuccess,
7028 Stop::ForegroundError,
7029 Stop::InputEof,
7030 Stop::TransportError,
7031 Stop::TaskError,
7032 Stop::RunnerError,
7033 Stop::SupervisorError,
7034 ] {
7035 let cleaned = Arc::new(AtomicBool::new(false));
7036 let disposable_dropped = Arc::new(AtomicBool::new(false));
7037 let close_finished = Arc::new(AtomicBool::new(false));
7038 let (cleanup_tx, cleanup_rx) = oneshot::channel::<()>();
7039 let (scoped_done_tx, scoped_done_rx) = completion_signal();
7040 let (stop_tx, stop_rx) = oneshot::channel::<()>();
7041 let stop_signal = stop_rx.map(|_| ()).boxed().shared();
7042 let (incoming_tx, incoming_rx) = mpsc::unbounded();
7043 let outgoing = futures::sink::unfold((), |(), _line: String| {
7044 future::ready(Ok::<_, std::io::Error>(()))
7045 });
7046 let builder = Client
7047 .builder()
7048 .with_spawned({
7049 let cleaned = cleaned.clone();
7050 async move |cx: ConnectionTo<Agent>| {
7051 cx.shutdown_requested().await;
7052 cleanup_rx.await.unwrap();
7055 cleaned.store(true, Ordering::Release);
7056 let _ = scoped_done_tx.send(());
7057 Ok(())
7058 }
7059 })
7060 .with_spawned({
7061 let stop_signal = stop_signal.clone();
7062 async move |_cx| {
7063 stop_signal.await;
7064 if matches!(stop, Stop::RunnerError) {
7065 Err(crate::Error::internal_error().data("runner failure"))
7066 } else {
7067 future::pending().await
7068 }
7069 }
7070 })
7071 .on_close({
7072 let close_finished = close_finished.clone();
7073 let scoped_done = scoped_done_rx.clone();
7074 async move |cx: ConnectionTo<Agent>| {
7075 cx.shutdown_requested().await;
7077 assert!(!cx.is_incoming_closed());
7078 scoped_done.await;
7079 close_finished.store(true, Ordering::Release);
7080 Ok(())
7081 }
7082 });
7083 let (connection, driver) =
7084 builder.into_connection_and_future(Lines::new(outgoing, incoming_rx), false, {
7085 let stop_signal = stop_signal.clone();
7086 async move |cx| {
7087 if matches!(stop, Stop::InputEof) {
7088 cx.incoming_closed().await;
7089 return Ok(());
7090 }
7091 stop_signal.await;
7092 match stop {
7093 Stop::ForegroundSuccess | Stop::SupervisorError => Ok(()),
7094 Stop::ForegroundError => {
7095 Err(crate::Error::internal_error().data("foreground failure"))
7096 }
7097 _ => future::pending().await,
7098 }
7099 }
7100 });
7101 let disposable = Dropped(disposable_dropped.clone());
7102 connection
7103 .spawn(async move {
7104 let _disposable = disposable;
7105 future::pending().await
7106 })
7107 .unwrap();
7108 connection
7109 .spawn({
7110 let stop_signal = stop_signal.clone();
7111 async move {
7112 stop_signal.await;
7113 if matches!(stop, Stop::TaskError) {
7114 Err(crate::Error::internal_error().data("task failure"))
7115 } else {
7116 future::pending().await
7117 }
7118 }
7119 })
7120 .unwrap();
7121 connection
7122 .spawn_protected({
7123 let connection = connection.clone();
7124 async move {
7125 connection.shutdown_requested().await;
7126 scoped_done_rx.await;
7127 if matches!(stop, Stop::SupervisorError) {
7128 Err(crate::Error::internal_error().data("supervisor failure"))
7129 } else {
7130 Ok(())
7131 }
7132 }
7133 })
7134 .unwrap();
7135 let mut driver = Box::pin(driver);
7136 assert!(driver.as_mut().now_or_never().is_none(), "{stop:?}");
7137 assert!(connection.shutdown_requested().now_or_never().is_none());
7138 let _ = stop_tx.send(());
7139 let incoming_tx = match stop {
7140 Stop::InputEof => {
7141 drop(incoming_tx);
7142 None
7143 }
7144 Stop::TransportError => {
7145 incoming_tx
7146 .unbounded_send(Err(std::io::Error::other("transport failure")))
7147 .unwrap();
7148 Some(incoming_tx)
7149 }
7150 _ => Some(incoming_tx),
7151 };
7152 for _ in 0..10 {
7153 assert!(driver.as_mut().now_or_never().is_none(), "{stop:?}");
7154 if connection.shutdown_requested().now_or_never().is_some() {
7155 break;
7156 }
7157 }
7158 assert!(
7159 connection.shutdown_requested().now_or_never().is_some(),
7160 "{stop:?}"
7161 );
7162 assert!(!cleaned.load(Ordering::Acquire), "{stop:?}");
7163 assert!(!disposable_dropped.load(Ordering::Acquire), "{stop:?}");
7164 cleanup_tx.send(()).unwrap();
7165 let mut result = None;
7168 for _ in 0..10 {
7169 result = driver.as_mut().now_or_never();
7170 if result.is_some() {
7171 break;
7172 }
7173 }
7174 let result =
7175 result.unwrap_or_else(|| panic!("driver did not finish owned cleanup: {stop:?}"));
7176 match stop {
7177 Stop::ForegroundSuccess | Stop::InputEof => result.unwrap(),
7178 _ => {
7179 let error = result.expect_err("shutdown must preserve the first error");
7180 let expected = match stop {
7181 Stop::ForegroundError => "foreground failure",
7182 Stop::TransportError => "transport failure",
7183 Stop::TaskError => "task failure",
7184 Stop::RunnerError => "runner failure",
7185 Stop::SupervisorError => "supervisor failure",
7186 _ => unreachable!(),
7187 };
7188 assert!(
7189 error.data.unwrap().to_string().contains(expected),
7190 "{stop:?}"
7191 );
7192 }
7193 }
7194 assert!(cleaned.load(Ordering::Acquire), "{stop:?}");
7195 assert!(disposable_dropped.load(Ordering::Acquire), "{stop:?}");
7196 assert_eq!(
7197 close_finished.load(Ordering::Acquire),
7198 matches!(stop, Stop::InputEof),
7199 "{stop:?}",
7200 );
7201 assert!(connection.spawn_protected(async { Ok(()) }).is_err());
7202 drop(incoming_tx);
7203 }
7204 }
7205
7206 #[test]
7207 fn protected_operation_acknowledgments_are_reaped_and_join_seals_registration() {
7208 let (connection, _message_rx, _pending_replies) = connection_for_response_hook_tests();
7209 let (task_tx, mut task_rx) = mpsc::unbounded();
7211 let connection = ConnectionTo {
7212 task_tx,
7213 ..connection
7214 };
7215 for _ in 0..100 {
7216 connection.spawn_protected(async { Ok(()) }).unwrap();
7217 assert_eq!(
7218 connection
7219 .protected_operations
7220 .lock()
7221 .unwrap()
7222 .pending
7223 .len(),
7224 1
7225 );
7226 let task = task_rx.next().now_or_never().unwrap().unwrap();
7227 futures::executor::block_on(task.run_for_test()).unwrap();
7228 }
7229 assert!(
7230 connection
7231 .wait_protected_operations()
7232 .now_or_never()
7233 .is_some()
7234 );
7235 assert!(
7236 connection
7237 .wait_protected_operations()
7238 .now_or_never()
7239 .is_some()
7240 );
7241 assert!(
7242 connection
7243 .protected_operations
7244 .lock()
7245 .unwrap()
7246 .pending
7247 .is_empty()
7248 );
7249 assert!(connection.spawn_protected(async { Ok(()) }).is_err());
7250 assert!(task_rx.next().now_or_never().is_none());
7251 }
7252
7253 #[test]
7254 fn dropping_unused_finish_signal_preserves_physical_half_closes() {
7255 let outgoing = futures::sink::unfold((), |(), _line: String| {
7256 future::ready(Ok::<_, std::io::Error>(()))
7257 });
7258 let (incoming_tx, incoming_rx) = mpsc::unbounded();
7259 let (Channel { mut rx, tx }, mut driver) =
7260 Lines::new(outgoing, incoming_rx).into_channel_transport();
7261
7262 drop(
7263 driver
7264 .take_finish()
7265 .expect("built-in Lines driver is finishable"),
7266 );
7267 drop(tx);
7268 assert!((&mut driver).now_or_never().is_none());
7269 incoming_tx
7270 .unbounded_send(Ok(
7271 r#"{"jsonrpc":"2.0","method":"test/after-output-eof"}"#.into()
7272 ))
7273 .unwrap();
7274 assert!((&mut driver).now_or_never().is_none());
7275 assert!(rx.next().now_or_never().unwrap().is_some());
7276
7277 drop(incoming_tx);
7278 futures::executor::block_on(driver).unwrap();
7279 assert!(rx.next().now_or_never().unwrap().is_none());
7280 }
7281
7282 #[test]
7283 fn explicit_physical_finish_does_not_hide_a_ready_read_error() {
7284 let outgoing = futures::sink::unfold((), |(), _line: String| {
7285 future::ready(Ok::<_, std::io::Error>(()))
7286 });
7287 let incoming = futures::stream::iter([Err(std::io::Error::other("finish read failed"))]);
7288 let (_channel, mut driver) = Lines::new(outgoing, incoming).into_channel_transport();
7289 assert!(driver.request_finish());
7290
7291 let error = futures::executor::block_on(driver).unwrap_err();
7292 assert_eq!(
7293 error
7294 .data
7295 .and_then(|value| value.as_str().map(str::to_owned)),
7296 Some("finish read failed".into())
7297 );
7298 }
7299
7300 #[cfg(feature = "unstable_protocol_v2")]
7301 fn connection_with_task_receiver() -> (
7302 ConnectionTo<crate::role::UntypedRole>,
7303 mpsc::UnboundedReceiver<Task>,
7304 ) {
7305 let (message_tx, _message_rx) = mpsc::unbounded();
7306 let (task_tx, task_rx) = mpsc::unbounded();
7307 let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded();
7308 let transport_completion: SharedTransportCompletion =
7309 future::ready(Ok::<(), crate::Error>(())).boxed().shared();
7310 let pending_replies = PendingReplies::default();
7311
7312 (
7313 ConnectionTo::new(
7314 crate::role::UntypedRole,
7315 message_tx,
7316 task_tx,
7317 dynamic_handler_tx,
7318 transport_completion,
7319 pending_replies.registrar(),
7320 ProtocolMode::disabled(),
7321 ),
7322 task_rx,
7323 )
7324 }
7325
7326 #[cfg(feature = "unstable_protocol_v2")]
7327 #[test]
7328 fn v2_builder_exposes_typed_context_to_user_callbacks() {
7329 fn assert_v2_context(_connection: &V2ConnectionTo<Agent>) {}
7330
7331 let _builder = Client
7332 .v2()
7333 .on_receive_request(
7334 async |_request: UntypedMessage, _responder, connection| {
7335 assert_v2_context(&connection);
7336 Ok(())
7337 },
7338 crate::on_receive_request!(),
7339 )
7340 .on_receive_notification(
7341 async |_notification: UntypedMessage, connection| {
7342 assert_v2_context(&connection);
7343 Ok(())
7344 },
7345 crate::on_receive_notification!(),
7346 )
7347 .on_receive_dispatch(
7348 async |_dispatch: Dispatch<UntypedMessage, UntypedMessage>, connection| {
7349 assert_v2_context(&connection);
7350 Ok(())
7351 },
7352 crate::on_receive_dispatch!(),
7353 )
7354 .on_receive_request_from(
7355 Agent,
7356 async |_request: UntypedMessage, _responder, connection| {
7357 assert_v2_context(&connection);
7358 Ok(())
7359 },
7360 crate::on_receive_request!(),
7361 )
7362 .on_receive_notification_from(
7363 Agent,
7364 async |_notification: UntypedMessage, connection| {
7365 assert_v2_context(&connection);
7366 Ok(())
7367 },
7368 crate::on_receive_notification!(),
7369 )
7370 .on_receive_dispatch_from(
7371 Agent,
7372 async |_dispatch: Dispatch<UntypedMessage, UntypedMessage>, connection| {
7373 assert_v2_context(&connection);
7374 Ok(())
7375 },
7376 crate::on_receive_dispatch!(),
7377 )
7378 .with_spawned(async |connection| {
7379 assert_v2_context(&connection);
7380 Ok(())
7381 })
7382 .on_close(async |connection| {
7383 assert_v2_context(&connection);
7384 Ok(())
7385 });
7386 }
7387
7388 #[cfg(feature = "unstable_protocol_v2")]
7389 #[test]
7390 fn proxy_builders_select_exact_proxy_protocol_guards() -> Result<(), crate::Error> {
7391 use crate::schema::ProtocolVersion;
7392
7393 for (mode, selected, unsupported) in [
7394 (
7395 Proxy.builder().protocol_mode,
7396 ProtocolVersion::V1,
7397 ProtocolVersion::V2,
7398 ),
7399 (
7400 Proxy.v2().protocol_mode,
7401 ProtocolVersion::V2,
7402 ProtocolVersion::V1,
7403 ),
7404 ] {
7405 assert_eq!(mode.api_protocol_version(), Some(selected));
7406
7407 let error = ProtocolCompat::new(mode)
7408 .incoming_message(UntypedMessage::new(
7409 "_proxy/initialize",
7410 serde_json::json!({ "protocolVersion": unsupported }),
7411 )?)
7412 .expect_err("a proxy builder must reject the other protocol version");
7413 let data = error
7414 .data
7415 .as_ref()
7416 .and_then(|data| data.as_str())
7417 .unwrap_or_default();
7418 assert!(
7419 data.contains(&format!("only supports ACP protocol version {selected}")),
7420 "{error:?}"
7421 );
7422 }
7423
7424 Ok(())
7425 }
7426
7427 #[cfg(feature = "unstable_protocol_v2")]
7428 #[test]
7429 fn v2_proxy_rejects_explicitly_prewrapped_initialize_request() {
7430 let (message_tx, message_rx) = mpsc::unbounded();
7431 let (task_tx, _task_rx) = mpsc::unbounded();
7432 let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded();
7433 let transport_completion: SharedTransportCompletion =
7434 future::ready(Ok::<(), crate::Error>(())).boxed().shared();
7435 let pending_replies = PendingReplies::default();
7436 let connection = ConnectionTo::new(
7437 crate::Conductor,
7438 message_tx,
7439 task_tx,
7440 dynamic_handler_tx,
7441 transport_completion,
7442 pending_replies.registrar(),
7443 ProtocolMode::v2_proxy(),
7444 );
7445
7446 let request = crate::schema::SuccessorMessage {
7447 message: UntypedMessage::new(
7448 "initialize",
7449 serde_json::json!({ "protocolVersion": crate::schema::ProtocolVersion::V1 }),
7450 )
7451 .expect("test initialize request should serialize"),
7452 meta: None,
7453 };
7454 let sent = connection.send_request_to(Agent, request);
7455
7456 let (transport_tx, mut transport_rx) = mpsc::unbounded();
7457 let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor(
7458 message_rx,
7459 pending_replies,
7460 transport_tx,
7461 ProtocolCompat::new(ProtocolMode::v2_proxy()),
7462 future::pending::<()>().boxed().shared(),
7463 ));
7464 assert!(
7465 actor.as_mut().now_or_never().is_none(),
7466 "the outgoing actor should continue after rejecting the request"
7467 );
7468 assert!(
7469 transport_rx.next().now_or_never().is_none(),
7470 "an explicitly prewrapped initialize must not reach the transport"
7471 );
7472
7473 let error = futures::executor::block_on(sent.block_task())
7474 .expect_err("connection routing must own successor wrapping");
7475 let data = error
7476 .data
7477 .as_ref()
7478 .and_then(|data| data.as_str())
7479 .unwrap_or_default();
7480 assert!(data.contains("logical `initialize`"), "{error:?}");
7481 assert!(data.contains("_proxy/successor"), "{error:?}");
7482 }
7483
7484 #[cfg(feature = "unstable_protocol_v2")]
7485 #[test]
7486 fn v2_proxy_builder_exposes_typed_context_to_user_callbacks() {
7487 fn assert_v2_context(_connection: &V2ConnectionTo<crate::Conductor>) {}
7488
7489 let _builder = Proxy
7490 .v2()
7491 .on_receive_request_from(
7492 Client,
7493 async |_request: UntypedMessage, _responder, connection| {
7494 assert_v2_context(&connection);
7495 Ok(())
7496 },
7497 crate::on_receive_request!(),
7498 )
7499 .on_receive_notification_from(
7500 Agent,
7501 async |_notification: UntypedMessage, connection| {
7502 assert_v2_context(&connection);
7503 Ok(())
7504 },
7505 crate::on_receive_notification!(),
7506 )
7507 .on_receive_dispatch_from(
7508 Client,
7509 async |_dispatch: Dispatch<UntypedMessage, UntypedMessage>, connection| {
7510 assert_v2_context(&connection);
7511 Ok(())
7512 },
7513 crate::on_receive_dispatch!(),
7514 )
7515 .with_spawned(async |connection| {
7516 assert_v2_context(&connection);
7517 Ok(())
7518 })
7519 .on_close(async |connection| {
7520 assert_v2_context(&connection);
7521 Ok(())
7522 });
7523 }
7524
7525 #[cfg(feature = "unstable_protocol_v2")]
7526 #[test]
7527 fn raw_connection_spawns_v2_builder_with_typed_child_callback() {
7528 let (parent, mut task_rx) = connection_with_task_receiver();
7529 let (transport, _peer) = Channel::duplex();
7530 let (callback_tx, callback_rx) = oneshot::channel();
7531
7532 let child: ConnectionTo<Agent> = parent
7533 .spawn_connection::<Client>(
7534 Client
7535 .v2()
7536 .with_spawned(async move |_connection: V2ConnectionTo<Agent>| {
7537 callback_tx.send(()).map_err(|()| {
7538 crate::util::internal_error("typed child callback receiver was dropped")
7539 })
7540 }),
7541 transport,
7542 )
7543 .expect("v2 child connection should be spawned");
7544
7545 let task = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut task_rx))
7546 .expect("child connection task should already be queued")
7547 .expect("parent task queue should remain open");
7548 futures::executor::block_on(async {
7549 match future::select(Box::pin(task.run_for_test()), Box::pin(callback_rx)).await {
7550 Either::Right((Ok(()), child_task)) => drop(child_task),
7551 Either::Right((Err(error), _)) => {
7552 panic!("typed child callback sender was dropped: {error}")
7553 }
7554 Either::Left((result, _)) => {
7555 panic!("child connection stopped before its typed callback ran: {result:?}")
7556 }
7557 }
7558 });
7559
7560 drop(child);
7561 }
7562
7563 #[cfg(feature = "unstable_protocol_v2")]
7564 #[test]
7565 fn raw_connection_can_return_v2_context_for_spawned_builder() {
7566 let (parent, mut task_rx) = connection_with_task_receiver();
7567 let (transport, _peer) = Channel::duplex();
7568
7569 let child: V2ConnectionTo<Agent> = parent
7570 .spawn_connection_with_context(Client.v2(), transport)
7571 .expect("v2 child connection should be spawned");
7572
7573 let child_task = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut task_rx))
7574 .expect("child connection task should already be queued")
7575 .expect("parent task queue should remain open");
7576
7577 drop((child, child_task));
7578 }
7579
7580 fn connection_with_dynamic_handler_receiver() -> (
7581 ConnectionTo<crate::role::UntypedRole>,
7582 mpsc::UnboundedReceiver<DynamicHandlerMessage<crate::role::UntypedRole>>,
7583 ) {
7584 let (message_tx, _message_rx) = mpsc::unbounded();
7585 let (task_tx, _task_rx) = mpsc::unbounded();
7586 let (dynamic_handler_tx, dynamic_handler_rx) = mpsc::unbounded();
7587 let transport_completion: SharedTransportCompletion =
7588 future::ready(Ok::<(), crate::Error>(())).boxed().shared();
7589 let pending_replies = PendingReplies::default();
7590
7591 (
7592 ConnectionTo::new(
7593 crate::role::UntypedRole,
7594 message_tx,
7595 task_tx,
7596 dynamic_handler_tx,
7597 transport_completion,
7598 pending_replies.registrar(),
7599 ProtocolMode::disabled(),
7600 ),
7601 dynamic_handler_rx,
7602 )
7603 }
7604
7605 struct ClaimingDynamicHandler;
7606
7607 impl HandleDispatchFrom<crate::role::UntypedRole> for ClaimingDynamicHandler {
7608 fn handle_dispatch_from(
7609 &mut self,
7610 _message: Dispatch,
7611 _connection: ConnectionTo<crate::role::UntypedRole>,
7612 ) -> impl Future<Output = Result<Handled<Dispatch>, crate::Error>> + Send {
7613 future::ready(Ok(Handled::Yes))
7614 }
7615
7616 fn describe_chain(&self) -> impl Debug {
7617 "ClaimingDynamicHandler"
7618 }
7619 }
7620
7621 fn connection_for_response_hook_tests() -> (
7622 ConnectionTo<crate::role::UntypedRole>,
7623 mpsc::UnboundedReceiver<OutgoingMessage>,
7624 PendingReplies,
7625 ) {
7626 let (message_tx, message_rx) = mpsc::unbounded();
7627 let (task_tx, _task_rx) = mpsc::unbounded();
7628 let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded();
7629 let transport_completion: SharedTransportCompletion =
7630 future::ready(Ok::<(), crate::Error>(())).boxed().shared();
7631 let pending_replies = PendingReplies::default();
7632
7633 (
7634 ConnectionTo::new(
7635 crate::role::UntypedRole,
7636 message_tx,
7637 task_tx,
7638 dynamic_handler_tx,
7639 transport_completion,
7640 pending_replies.registrar(),
7641 ProtocolMode::disabled(),
7642 ),
7643 message_rx,
7644 pending_replies,
7645 )
7646 }
7647
7648 #[cfg(feature = "unstable_protocol_v2")]
7649 fn route_test_response(
7650 request_id: RequestId,
7651 pending_replies: &PendingReplies,
7652 result: Result<serde_json::Value, crate::Error>,
7653 ) {
7654 let pending_reply = pending_replies
7655 .remove(&request_id)
7656 .expect("the request should have a pending reply");
7657 let (dispatch, _) =
7658 incoming_actor::dispatch_from_response(request_id, pending_reply, result);
7659 let Dispatch::Response(result, router) = dispatch else {
7660 panic!("expected a response dispatch");
7661 };
7662 router
7663 .route_with_result(result)
7664 .expect("response should route to the pending request");
7665 }
7666
7667 #[cfg(feature = "unstable_protocol_v2")]
7668 #[test]
7669 fn response_hook_runs_when_success_is_routed_before_consumption() {
7670 let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests();
7671 let hook_ran = Arc::new(AtomicBool::new(false));
7672 let sent = connection.send_request_to_with_response_hook_after(
7673 crate::role::UntypedRole,
7674 UntypedMessage::new("hooked", serde_json::json!({}))
7675 .expect("test request should serialize"),
7676 future::ready(Ok(())),
7677 {
7678 let hook_ran = hook_ran.clone();
7679 move |response| {
7680 assert_eq!(response, &serde_json::json!({"ok": true}));
7681 hook_ran.store(true, Ordering::Release);
7682 Ok(())
7683 }
7684 },
7685 );
7686 let request_id = sent.id().clone();
7687
7688 route_test_response(
7689 request_id,
7690 &pending_replies,
7691 Ok(serde_json::json!({"ok": true})),
7692 );
7693
7694 assert!(hook_ran.load(Ordering::Acquire));
7695 assert_eq!(
7696 futures::executor::block_on(sent.block_task())
7697 .expect("routed response should remain consumable"),
7698 serde_json::json!({"ok": true})
7699 );
7700 }
7701
7702 #[cfg(feature = "unstable_protocol_v2")]
7703 #[test]
7704 fn response_hook_skips_errors_but_outlives_a_dropped_consumer() {
7705 let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests();
7706 let peer_error_hook_ran = Arc::new(AtomicBool::new(false));
7707 let peer_error = connection.send_request_to_with_response_hook_after(
7708 crate::role::UntypedRole,
7709 UntypedMessage::new("peer-error", serde_json::json!({}))
7710 .expect("test request should serialize"),
7711 future::ready(Ok(())),
7712 {
7713 let hook_ran = peer_error_hook_ran.clone();
7714 move |_| {
7715 hook_ran.store(true, Ordering::Release);
7716 Ok(())
7717 }
7718 },
7719 );
7720 let peer_error_id = peer_error.id().clone();
7721 route_test_response(
7722 peer_error_id,
7723 &pending_replies,
7724 Err(crate::Error::invalid_request()),
7725 );
7726 assert!(
7727 futures::executor::block_on(peer_error.block_task()).is_err(),
7728 "the peer error should reach the consumer"
7729 );
7730 assert!(!peer_error_hook_ran.load(Ordering::Acquire));
7731
7732 let dropped_hook_ran = Arc::new(AtomicBool::new(false));
7733 let dropped = connection.send_request_to_with_response_hook_after(
7734 crate::role::UntypedRole,
7735 UntypedMessage::new("dropped", serde_json::json!({}))
7736 .expect("test request should serialize"),
7737 future::ready(Ok(())),
7738 {
7739 let hook_ran = dropped_hook_ran.clone();
7740 move |_| {
7741 hook_ran.store(true, Ordering::Release);
7742 Ok(())
7743 }
7744 },
7745 );
7746 let dropped_id = dropped.id().clone();
7747 drop(dropped);
7748 route_test_response(
7749 dropped_id,
7750 &pending_replies,
7751 Ok(serde_json::json!({"ok": true})),
7752 );
7753 assert!(dropped_hook_ran.load(Ordering::Acquire));
7754 }
7755
7756 #[cfg(feature = "unstable_protocol_v2")]
7757 #[test]
7758 fn response_hook_failure_replaces_the_success_result() {
7759 let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests();
7760 let sent = connection.send_request_to_with_response_hook_after(
7761 crate::role::UntypedRole,
7762 UntypedMessage::new("hook-failure", serde_json::json!({}))
7763 .expect("test request should serialize"),
7764 future::ready(Ok(())),
7765 |_| Err(crate::Error::internal_error().data("response hook failed")),
7766 );
7767 let request_id = sent.id().clone();
7768 route_test_response(
7769 request_id,
7770 &pending_replies,
7771 Ok(serde_json::json!({"ok": true})),
7772 );
7773
7774 let error = futures::executor::block_on(sent.block_task())
7775 .expect_err("the hook failure should replace the successful response");
7776 assert_eq!(error.code, crate::ErrorCode::InternalError);
7777 assert_eq!(error.data, Some(serde_json::json!("response hook failed")));
7778 }
7779
7780 #[test]
7781 fn ordered_request_waits_for_readiness_before_publication() {
7782 let (connection, message_rx, pending_replies) = connection_for_response_hook_tests();
7783 let (ready_tx, ready_rx) = oneshot::channel();
7784 let sent = connection.send_ordered_request_to_after(
7785 crate::role::UntypedRole,
7786 UntypedMessage::new("after-ready", serde_json::json!({}))
7787 .expect("test request should serialize"),
7788 async move { ready_rx.await.map_err(crate::Error::into_internal_error) },
7789 );
7790
7791 let (transport_tx, mut transport_rx) = mpsc::unbounded();
7792 let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor(
7793 message_rx,
7794 pending_replies,
7795 transport_tx,
7796 ProtocolCompat::new(ProtocolMode::disabled()),
7797 future::pending::<()>().boxed().shared(),
7798 ));
7799
7800 assert!(
7801 actor.as_mut().now_or_never().is_none(),
7802 "the outgoing actor should wait for readiness"
7803 );
7804 assert!(
7805 transport_rx.next().now_or_never().is_none(),
7806 "the request must not be published before readiness"
7807 );
7808
7809 ready_tx
7810 .send(())
7811 .expect("the readiness receiver should remain active");
7812 assert!(
7813 actor.as_mut().now_or_never().is_none(),
7814 "the outgoing actor should continue serving after publication"
7815 );
7816 let frame = transport_rx
7817 .next()
7818 .now_or_never()
7819 .expect("the ready request should be published")
7820 .expect("the transport queue should remain open");
7821 assert!(matches!(
7822 frame,
7823 TransportFrame::Single(RawJsonRpcMessage::Request(_))
7824 ));
7825
7826 drop(sent);
7827 }
7828
7829 #[test]
7830 fn foreground_finish_settles_unready_requests_and_preserves_ready_output_fifo() {
7831 let (connection, message_rx, pending_replies) = connection_for_response_hook_tests();
7832 let unready = connection.send_ordered_request_to_after(
7833 crate::role::UntypedRole,
7834 UntypedMessage::new("unready", serde_json::json!({})).unwrap(),
7835 future::pending(),
7836 );
7837 let unready_id = unready.id().clone();
7838 send_raw_message(
7839 &connection.message_tx,
7840 OutgoingMessage::Notification {
7841 untyped: UntypedMessage::new("first", serde_json::json!({})).unwrap(),
7842 },
7843 )
7844 .unwrap();
7845 let ready = connection.send_ordered_request_to_after(
7846 crate::role::UntypedRole,
7847 UntypedMessage::new("ready", serde_json::json!({})).unwrap(),
7848 future::ready(Ok(())),
7849 );
7850 let unready_after = connection.send_ordered_request_to_after(
7851 crate::role::UntypedRole,
7852 UntypedMessage::new("unready-after", serde_json::json!({})).unwrap(),
7853 future::pending(),
7854 );
7855 let unready_after_id = unready_after.id().clone();
7856 send_raw_message(
7857 &connection.message_tx,
7858 OutgoingMessage::Notification {
7859 untyped: UntypedMessage::new("last", serde_json::json!({})).unwrap(),
7860 },
7861 )
7862 .unwrap();
7863 let (done_tx, done_rx) = oneshot::channel();
7864 send_raw_message(
7865 &connection.message_tx,
7866 OutgoingMessage::CloseAfterDraining { done: done_tx },
7867 )
7868 .unwrap();
7869 let (transport_tx, transport_rx) = mpsc::unbounded();
7870 futures::executor::block_on(outgoing_actor::outgoing_protocol_actor(
7871 message_rx,
7872 pending_replies.clone(),
7873 transport_tx,
7874 ProtocolCompat::new(ProtocolMode::disabled()),
7875 future::ready(()).boxed().shared(),
7876 ))
7877 .unwrap();
7878 futures::executor::block_on(done_rx).unwrap();
7879 let error = futures::executor::block_on(unready.block_task())
7880 .expect_err("an unresolved gate must explicitly fail its consumer");
7881 assert!(
7882 error
7883 .data
7884 .unwrap()
7885 .to_string()
7886 .contains("foreground completed before outgoing request readiness")
7887 );
7888 assert!(!pending_replies.contains(&unready_id));
7889 let error = futures::executor::block_on(unready_after.block_task())
7890 .expect_err("each unresolved gate must fail without repolling a consumed signal");
7891 assert!(
7892 error
7893 .data
7894 .unwrap()
7895 .to_string()
7896 .contains("foreground completed before outgoing request readiness")
7897 );
7898 assert!(!pending_replies.contains(&unready_after_id));
7899 assert!(pending_replies.contains(ready.id()));
7900 let frames = futures::executor::block_on(transport_rx.collect::<Vec<_>>());
7901 let methods = frames
7902 .into_iter()
7903 .map(|frame| match frame {
7904 TransportFrame::Single(RawJsonRpcMessage::Notification(message)) => {
7905 message.method.to_string()
7906 }
7907 TransportFrame::Single(RawJsonRpcMessage::Request(message)) => {
7908 message.method.to_string()
7909 }
7910 _ => panic!("expected ready request/notification output"),
7911 })
7912 .collect::<Vec<_>>();
7913 assert_eq!(methods, ["first", "ready", "last"]);
7914 }
7915
7916 #[test]
7917 fn ordered_blocking_transform_precedes_response_acknowledgment() {
7918 let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests();
7919 let sent = connection.send_ordered_request_to(
7920 crate::role::UntypedRole,
7921 UntypedMessage::new("ordered-transform", serde_json::json!({}))
7922 .expect("test request should serialize"),
7923 );
7924 let request_id = sent.id().clone();
7925 let pending_reply = pending_replies
7926 .remove(&request_id)
7927 .expect("the request should have a pending reply");
7928 let (dispatch, response_dispatch) = incoming_actor::dispatch_from_response(
7929 request_id,
7930 pending_reply,
7931 Err(crate::Error::invalid_params()),
7932 );
7933 let Dispatch::Response(result, router) = dispatch else {
7934 panic!("expected a response dispatch");
7935 };
7936 router
7937 .route_with_result(result)
7938 .expect("response should route to the pending request");
7939 let acknowledgment = response_dispatch
7940 .complete()
7941 .expect("an ordered response should wait for acknowledgment");
7942 let acknowledgment = Arc::new(Mutex::new(Some(acknowledgment)));
7943 let acknowledgment_probe = acknowledgment.clone();
7944
7945 let error =
7946 futures::executor::block_on(sent.block_task_with_ordered_result(move |result| {
7947 assert_eq!(
7948 acknowledgment_probe
7949 .lock()
7950 .expect("acknowledgment mutex poisoned")
7951 .as_mut()
7952 .expect("acknowledgment receiver should remain available")
7953 .try_recv()
7954 .expect("acknowledgment sender should remain open"),
7955 None,
7956 "the ordered response was acknowledged before its transform"
7957 );
7958 result
7959 }))
7960 .expect_err("the peer error should survive the ordered transform");
7961 assert_eq!(error.code, crate::ErrorCode::InvalidParams);
7962
7963 let acknowledgment = acknowledgment
7964 .lock()
7965 .expect("acknowledgment mutex poisoned")
7966 .take()
7967 .expect("acknowledgment receiver should remain available");
7968 futures::executor::block_on(acknowledgment)
7969 .expect("the transform should release the ordered response");
7970 }
7971
7972 #[cfg(feature = "unstable_protocol_v2")]
7973 #[test]
7974 fn outgoing_request_readiness_failure_rejects_without_publication() {
7975 let (connection, message_rx, pending_replies) = connection_for_response_hook_tests();
7976 let hook_ran = Arc::new(AtomicBool::new(false));
7977 let sent = connection.send_request_to_with_response_hook_after(
7978 crate::role::UntypedRole,
7979 UntypedMessage::new("never-published", serde_json::json!({}))
7980 .expect("test request should serialize"),
7981 future::ready(Err(crate::Error::internal_error().data("readiness failed"))),
7982 {
7983 let hook_ran = hook_ran.clone();
7984 move |_| {
7985 hook_ran.store(true, Ordering::Release);
7986 Ok(())
7987 }
7988 },
7989 );
7990
7991 let (transport_tx, mut transport_rx) = mpsc::unbounded();
7992 let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor(
7993 message_rx,
7994 pending_replies,
7995 transport_tx,
7996 ProtocolCompat::new(ProtocolMode::disabled()),
7997 future::pending::<()>().boxed().shared(),
7998 ));
7999
8000 assert!(
8001 actor.as_mut().now_or_never().is_none(),
8002 "the outgoing actor should continue serving after rejecting the request"
8003 );
8004 assert!(
8005 transport_rx.next().now_or_never().is_none(),
8006 "a request whose readiness failed must not be published"
8007 );
8008 let error = futures::executor::block_on(sent.block_task())
8009 .expect_err("the readiness error should reach the request consumer");
8010 assert_eq!(error.code, crate::ErrorCode::InternalError);
8011 assert_eq!(error.data, Some(serde_json::json!("readiness failed")));
8012 assert!(!hook_ran.load(Ordering::Acquire));
8013 }
8014
8015 #[test]
8016 fn ordered_request_is_marked_before_entering_outgoing_queue() {
8017 let (message_tx, mut message_rx) = mpsc::unbounded();
8018 let (task_tx, mut task_rx) = mpsc::unbounded();
8019 let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded();
8020 let transport_completion: SharedTransportCompletion =
8021 future::ready(Ok::<(), crate::Error>(())).boxed().shared();
8022 let pending_replies = PendingReplies::default();
8023 let connection = ConnectionTo::new(
8024 crate::role::UntypedRole,
8025 message_tx,
8026 task_tx,
8027 dynamic_handler_tx,
8028 transport_completion,
8029 pending_replies.registrar(),
8030 ProtocolMode::disabled(),
8031 );
8032
8033 let sent = connection.send_ordered_request_to(
8034 crate::role::UntypedRole,
8035 UntypedMessage::new("ordered", serde_json::json!({}))
8036 .expect("test request should serialize"),
8037 );
8038 let request_id = sent.id().clone();
8039 let message = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut message_rx))
8040 .expect("outgoing request should already be queued")
8041 .expect("outgoing request queue should remain open");
8042 let OutgoingMessage::Request { id, .. } = message else {
8043 panic!("expected an outgoing request");
8044 };
8045 assert_eq!(id, request_id);
8046
8047 let pending_reply = pending_replies
8048 .remove(&request_id)
8049 .expect("the request should have a pending reply");
8050 assert!(
8051 pending_reply.ordering.is_ordered(),
8052 "the response ordering barrier must be installed before publication"
8053 );
8054
8055 let (dispatch, response_dispatch) = incoming_actor::dispatch_from_response(
8059 request_id,
8060 pending_reply,
8061 Ok(serde_json::json!({"ok": true})),
8062 );
8063 let Dispatch::Response(result, router) = dispatch else {
8064 panic!("expected a response dispatch");
8065 };
8066 router
8067 .route_with_result(result)
8068 .expect("response should route to the pending request");
8069 let acknowledgment = response_dispatch
8070 .complete()
8071 .expect("an ordered response should require acknowledgment");
8072
8073 let callback_ran = Arc::new(AtomicBool::new(false));
8074 sent.on_receiving_result({
8075 let callback_ran = callback_ran.clone();
8076 async move |result| {
8077 assert_eq!(result?, serde_json::json!({"ok": true}));
8078 callback_ran.store(true, Ordering::Release);
8079 Ok(())
8080 }
8081 })
8082 .expect("ordered callback should be scheduled");
8083
8084 let task = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut task_rx))
8085 .expect("callback task should already be queued")
8086 .expect("callback task queue should remain open");
8087 futures::executor::block_on(task.run_for_test()).expect("callback task should succeed");
8088 futures::executor::block_on(acknowledgment)
8089 .expect("callback completion should acknowledge dispatch");
8090 assert!(callback_ran.load(Ordering::Acquire));
8091 }
8092
8093 fn next_dynamic_handler_message<Counterpart: Role>(
8094 receiver: &mut mpsc::UnboundedReceiver<DynamicHandlerMessage<Counterpart>>,
8095 ) -> Option<DynamicHandlerMessage<Counterpart>> {
8096 futures::FutureExt::now_or_never(futures::StreamExt::next(receiver))
8097 .expect("dynamic-handler receiver should be ready")
8098 }
8099
8100 #[cfg(feature = "unstable_protocol_v2")]
8101 #[test]
8102 fn v2_dynamic_handler_guard_registers_and_removes_handler() {
8103 let (message_tx, _message_rx) = mpsc::unbounded();
8104 let (task_tx, _task_rx) = mpsc::unbounded();
8105 let (dynamic_handler_tx, mut dynamic_handler_rx) = mpsc::unbounded();
8106 let transport_completion: SharedTransportCompletion =
8107 future::ready(Ok::<(), crate::Error>(())).boxed().shared();
8108 let pending_replies = PendingReplies::default();
8109 let connection = V2ConnectionTo {
8110 inner: ConnectionTo::new(
8111 Agent,
8112 message_tx,
8113 task_tx,
8114 dynamic_handler_tx,
8115 transport_completion,
8116 pending_replies.registrar(),
8117 ProtocolMode::v2_client(),
8118 ),
8119 };
8120
8121 let guard = connection
8122 .add_dynamic_handler(NullHandler)
8123 .expect("v2 dynamic handler should register");
8124 let added_uuid = match next_dynamic_handler_message(&mut dynamic_handler_rx) {
8125 Some(DynamicHandlerMessage::AddDynamicHandler(uuid, _)) => uuid,
8126 other => panic!("expected v2 handler registration, got {other:?}"),
8127 };
8128
8129 drop(guard);
8130
8131 match next_dynamic_handler_message(&mut dynamic_handler_rx) {
8132 Some(DynamicHandlerMessage::RemoveDynamicHandler(uuid)) => {
8133 assert_eq!(uuid, added_uuid);
8134 }
8135 other => panic!("expected v2 handler removal, got {other:?}"),
8136 }
8137 }
8138
8139 #[test]
8140 fn dropping_dynamic_handler_guard_unregisters_handler() {
8141 let (connection, mut receiver) = connection_with_dynamic_handler_receiver();
8142 let guard = connection.add_dynamic_handler(NullHandler).unwrap();
8143
8144 let added_uuid = match next_dynamic_handler_message(&mut receiver) {
8145 Some(DynamicHandlerMessage::AddDynamicHandler(uuid, _)) => uuid,
8146 other => panic!("expected handler registration, got {other:?}"),
8147 };
8148
8149 drop(guard);
8150
8151 match next_dynamic_handler_message(&mut receiver) {
8152 Some(DynamicHandlerMessage::RemoveDynamicHandler(uuid)) => {
8153 assert_eq!(uuid, added_uuid);
8154 }
8155 other => panic!("expected handler removal, got {other:?}"),
8156 }
8157 }
8158
8159 #[test]
8160 fn dropping_dynamic_handler_guard_deactivates_queued_handler_immediately() {
8161 let (connection, mut receiver) = connection_with_dynamic_handler_receiver();
8162 let guard = connection
8163 .add_dynamic_handler(ClaimingDynamicHandler)
8164 .expect("dynamic handler should register");
8165 let mut handler = match next_dynamic_handler_message(&mut receiver) {
8166 Some(DynamicHandlerMessage::AddDynamicHandler(_, handler)) => handler,
8167 other => panic!("expected handler registration, got {other:?}"),
8168 };
8169
8170 drop(guard);
8171
8172 let message = Dispatch::Notification(
8173 UntypedMessage::new("stale", serde_json::json!({}))
8174 .expect("test notification should serialize"),
8175 );
8176 let handled =
8177 futures::executor::block_on(handler.dyn_handle_dispatch_from(message, connection))
8178 .expect("inactive handler should decline cleanly");
8179 assert!(matches!(handled, Handled::No { retry: false, .. }));
8180 }
8181
8182 #[test]
8183 fn dynamic_handler_barrier_acknowledges_prior_messages() {
8184 let (connection, mut receiver) = connection_with_dynamic_handler_receiver();
8185 let _guard = connection.add_dynamic_handler(NullHandler).unwrap();
8186 let mut barrier = Box::pin(connection.dynamic_handler_barrier());
8187
8188 assert!(matches!(
8189 next_dynamic_handler_message(&mut receiver),
8190 Some(DynamicHandlerMessage::AddDynamicHandler(_, _))
8191 ));
8192 assert!(
8193 barrier.as_mut().now_or_never().is_none(),
8194 "the barrier must wait for the incoming actor"
8195 );
8196
8197 let acknowledgment = match next_dynamic_handler_message(&mut receiver) {
8198 Some(DynamicHandlerMessage::AcknowledgedBarrier(acknowledgment)) => acknowledgment,
8199 other => panic!("expected acknowledged barrier, got {other:?}"),
8200 };
8201 acknowledgment
8202 .send(())
8203 .expect("the barrier receiver should remain active");
8204 futures::executor::block_on(barrier)
8205 .expect("the acknowledged dynamic-handler barrier should complete");
8206 }
8207
8208 #[test]
8209 fn detaching_dynamic_handler_guard_does_not_leak_connection() {
8210 let (connection, mut receiver) = connection_with_dynamic_handler_receiver();
8211 let guard = connection.add_dynamic_handler(NullHandler).unwrap();
8212
8213 assert!(matches!(
8214 next_dynamic_handler_message(&mut receiver),
8215 Some(DynamicHandlerMessage::AddDynamicHandler(_, _))
8216 ));
8217
8218 drop(connection);
8219 guard.detach();
8220
8221 assert!(
8222 next_dynamic_handler_message(&mut receiver).is_none(),
8223 "detach should retain the handler without retaining a connection sender"
8224 );
8225 }
8226
8227 #[tokio::test]
8228 async fn write_line_flushes_buffered_writers() {
8229 let mut writer =
8230 futures::io::BufWriter::with_capacity(4096, futures::io::Cursor::new(Vec::new()));
8231
8232 write_line(&mut writer, "message".into()).await.unwrap();
8233
8234 assert_eq!(writer.into_inner().into_inner(), b"message\n");
8235 }
8236
8237 #[test]
8238 fn peel_successor_envelopes_returns_plain_messages_unchanged() {
8239 let params = serde_json::json!({ "key": "value" });
8240 let (method, peeled) = peel_successor_envelopes("session/update", ¶ms);
8241 assert_eq!(method, "session/update");
8242 assert_eq!(peeled, ¶ms);
8243 }
8244
8245 #[test]
8246 fn peel_successor_envelopes_unwraps_nested_envelopes() {
8247 let params = serde_json::json!({
8248 "method": "_proxy/successor",
8249 "params": {
8250 "method": "$/cancel_request",
8251 "params": { "requestId": "req-1" }
8252 }
8253 });
8254 let (method, peeled) = peel_successor_envelopes("_proxy/successor", ¶ms);
8255 assert_eq!(method, "$/cancel_request");
8256 assert_eq!(peeled, &serde_json::json!({ "requestId": "req-1" }));
8257 }
8258
8259 #[test]
8260 fn peel_successor_envelopes_leaves_malformed_envelopes_intact() {
8261 let params = serde_json::json!({ "unexpected": true });
8264 let (method, peeled) = peel_successor_envelopes("_proxy/successor", ¶ms);
8265 assert_eq!(method, "_proxy/successor");
8266 assert_eq!(peeled, ¶ms);
8267 }
8268
8269 mod cancel_request {
8270 use super::super::*;
8271
8272 fn notification(method: &str, params: serde_json::Value) -> UntypedMessage {
8273 UntypedMessage::new(method, params).expect("well-formed JSON")
8274 }
8275
8276 #[test]
8277 fn cancellation_request_id_is_extracted_from_wrapped_notifications() {
8278 let message = notification(
8279 "_proxy/successor",
8280 serde_json::json!({
8281 "method": "$/cancel_request",
8282 "params": { "requestId": "req-1" }
8283 }),
8284 );
8285 let request_id = cancellation_request_id_from_message(&message)
8286 .expect("wrapped cancel should parse");
8287 assert_eq!(request_id, Some(RequestId::Str("req-1".into())));
8288 }
8289
8290 #[test]
8291 fn malformed_successor_envelope_is_not_treated_as_cancellation() {
8292 let message = notification("_proxy/successor", serde_json::json!({ "bogus": true }));
8295 let request_id = cancellation_request_id_from_message(&message)
8296 .expect("malformed envelope should be left to the handler chain");
8297 assert_eq!(request_id, None);
8298 }
8299
8300 #[test]
8301 fn cancel_request_notifications_are_detected_even_when_wrapped() {
8302 let plain = notification("$/cancel_request", serde_json::json!({ "requestId": 1 }));
8303 assert!(is_cancel_request_notification(&plain));
8304
8305 let wrapped = notification(
8306 "_proxy/successor",
8307 serde_json::json!({
8308 "method": "$/cancel_request",
8309 "params": { "requestId": 1 }
8310 }),
8311 );
8312 assert!(is_cancel_request_notification(&wrapped));
8313
8314 let other_wrapped = notification(
8315 "_proxy/successor",
8316 serde_json::json!({
8317 "method": "session/update",
8318 "params": {}
8319 }),
8320 );
8321 assert!(!is_cancel_request_notification(&other_wrapped));
8322
8323 let malformed_envelope =
8324 notification("_proxy/successor", serde_json::json!({ "bogus": true }));
8325 assert!(!is_cancel_request_notification(&malformed_envelope));
8326 }
8327
8328 #[test]
8329 fn malformed_cancel_request_params_error() {
8330 let message = notification(
8331 "$/cancel_request",
8332 serde_json::json!({ "requestId": { "not": "an id" } }),
8333 );
8334 cancellation_request_id_from_message(&message)
8335 .expect_err("malformed cancel params should error");
8336 }
8337
8338 #[test]
8339 fn registry_marks_and_removes_requests() {
8340 let registry = RequestCancellationRegistry::new();
8341 let id = RequestId::Str("req-1".into());
8342
8343 let responder_cancellation = registry.register(&id);
8344 let marker = responder_cancellation.cancellation();
8345 assert!(!marker.is_cancelled());
8346
8347 assert!(registry.cancel(&id));
8348 assert!(marker.is_cancelled());
8349 assert!(responder_cancellation.cancellation().is_cancelled());
8350
8351 drop(responder_cancellation);
8352 assert!(!registry.cancel(&id), "slot should be removed on drop");
8353 }
8354
8355 #[test]
8356 fn reused_request_id_does_not_cross_wire_cancellation_state() {
8357 let registry = RequestCancellationRegistry::new();
8358 let id = RequestId::Str("dup".into());
8359
8360 let first = registry.register(&id);
8362 let first_marker = first.cancellation();
8363 let second = registry.register(&id);
8364 let second_marker = second.cancellation();
8365
8366 assert!(registry.cancel(&id));
8368 assert!(second_marker.is_cancelled());
8369 assert!(
8370 !first_marker.is_cancelled(),
8371 "the stale request must not observe the newer request's cancellation"
8372 );
8373
8374 assert!(!first.cancellation().is_cancelled());
8377
8378 drop(first);
8381 assert!(registry.cancel(&id), "newer slot should still be present");
8382
8383 drop(second);
8384 assert!(!registry.cancel(&id), "slot should be removed on drop");
8385 }
8386 }
8387}