1use std::{fmt::Debug, future::Future, hash::Hash};
2
3#[cfg(feature = "unstable_protocol_v2")]
4use futures::{StreamExt as _, future};
5#[cfg(feature = "unstable_protocol_v2")]
6use serde::{Serialize, de::DeserializeOwned};
7
8#[cfg(feature = "unstable_protocol_v2")]
9use crate::DynConnectTo;
10use crate::jsonrpc::{Builder, handlers::NullHandler, run::NullRun};
11#[cfg(feature = "unstable_protocol_v2")]
12use crate::jsonrpc::{
13 TransportBatch, TransportBatchEntry, TransportFrame, V2Builder, is_response_only_shape,
14 raw_is_response_only_shape,
15};
16use crate::role::{HasPeer, RemoteStyle};
17#[cfg(not(feature = "unstable_protocol_v2"))]
18use crate::schema::InitializeProxyRequest;
19use crate::schema::METHOD_INITIALIZE_PROXY;
20#[cfg(feature = "unstable_protocol_v2")]
21use crate::schema::v1::RequestId;
22use crate::schema::v1::{InitializeRequest, SessionId};
23#[cfg(not(feature = "unstable_protocol_v2"))]
24use crate::schema::v1::{NewSessionRequest, NewSessionResponse};
25#[cfg(feature = "unstable_protocol_v2")]
26use crate::schema::{ProtocolVersion, v2};
27use crate::util::MatchDispatchFrom;
28#[cfg(feature = "unstable_protocol_v2")]
29use crate::{
30 Channel, RawJsonRpcError, RawJsonRpcMessage, RawJsonRpcParams,
31 RawJsonRpcResponse as RpcResponse,
32};
33use crate::{ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, Role, RoleId};
34
35#[cfg(feature = "unstable_protocol_v2")]
36#[derive(serde::Deserialize)]
37struct NewSessionResponseEnvelope {
38 #[serde(rename = "sessionId")]
39 session_id: SessionId,
40}
41
42#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
46pub struct Client;
47
48impl Role for Client {
49 type Counterpart = Agent;
50
51 fn builder(self) -> Builder<Self> {
52 Builder::new(self).v1_client()
53 }
54
55 fn default_handle_dispatch_from(
56 &self,
57 message: Dispatch,
58 _connection: ConnectionTo<Client>,
59 ) -> impl Future<Output = Result<Handled<Dispatch>, crate::Error>> + Send {
60 std::future::ready(Ok(Handled::No {
61 message,
62 retry: false,
63 }))
64 }
65
66 fn role_id(&self) -> RoleId {
67 RoleId::from_singleton(self)
68 }
69
70 fn counterpart(&self) -> Self::Counterpart {
71 Agent
72 }
73}
74
75impl Client {
76 pub fn builder(self) -> Builder<Client, NullHandler, NullRun> {
78 <Self as Role>::builder(self)
79 }
80
81 #[cfg(feature = "unstable_protocol_v2")]
89 pub fn v2(self) -> V2Builder<Client, NullHandler, NullRun> {
90 self.builder().v2_client()
91 }
92
93 #[cfg(feature = "unstable_protocol_v2")]
106 #[must_use]
107 pub fn protocol_connector(self) -> ClientProtocolConnector {
108 ClientProtocolConnector::new()
109 }
110
111 pub async fn connect_with<R>(
116 self,
117 agent: impl ConnectTo<Client>,
118 main_fn: impl AsyncFnOnce(ConnectionTo<Agent>) -> Result<R, crate::Error>,
119 ) -> Result<R, crate::Error> {
120 self.builder().connect_with(agent, main_fn).await
121 }
122}
123
124#[cfg(feature = "unstable_protocol_v2")]
131#[derive(Debug, Default)]
132pub struct ClientProtocolConnector {
133 v1: Option<DynConnectToFactory<Agent>>,
134 v2: Option<DynConnectToFactory<Agent>>,
135}
136
137#[cfg(feature = "unstable_protocol_v2")]
138impl ClientProtocolConnector {
139 #[must_use]
141 pub fn new() -> Self {
142 Self::default()
143 }
144
145 #[must_use]
147 pub fn with_v1<C>(mut self, client: impl FnMut() -> C + Send + 'static) -> Self
148 where
149 C: ConnectTo<Agent>,
150 {
151 self.v1 = Some(DynConnectToFactory::new(client));
152 self
153 }
154
155 #[must_use]
157 pub fn with_v2<C>(mut self, client: impl FnMut() -> C + Send + 'static) -> Self
158 where
159 C: ConnectTo<Agent>,
160 {
161 self.v2 = Some(DynConnectToFactory::new(client));
162 self
163 }
164
165 pub async fn connect_to<C>(
168 mut self,
169 mut agent: impl FnMut() -> C + Send + 'static,
170 ) -> Result<(), crate::Error>
171 where
172 C: ConnectTo<Client>,
173 {
174 let supported = SupportedClientProtocols {
175 v1: self.v1.is_some(),
176 v2: self.v2.is_some(),
177 };
178 let Some(selected) = supported.highest_configured() else {
179 return Err(crate::Error::invalid_request()
180 .data("client protocol connector has no configured ACP protocol implementations"));
181 };
182
183 match selected {
184 ClientProtocol::V1 => {
185 let client = self
186 .v1
187 .as_mut()
188 .expect("selected protocol is configured")
189 .create();
190 connect_client_protocol(ClientProtocol::V1, client, agent()).await
191 }
192 ClientProtocol::V2 => {
193 let client = self
194 .v2
195 .as_mut()
196 .expect("selected protocol is configured")
197 .create();
198 let agent_connection = RunningProtocolPeer::new(agent());
199 let (client, initialize) =
200 start_client_protocol(ClientProtocol::V2, client).await?;
201 let v2_initialize_as_v1 = normalize_v2_initialize_params_for_reuse(&initialize);
205 let (client, agent_connection, initialize_response) =
206 send_initialize_and_receive(client, agent_connection, initialize).await?;
207
208 if initialize_response_negotiated_v1(&initialize_response)
209 && let Some(v1) = self.v1.as_mut()
210 {
211 let fallback_client = v1.create();
212 let (fallback_client, fallback_initialize) =
213 start_client_protocol(ClientProtocol::V1, fallback_client).await?;
214 let v1_initialize =
215 validated_initialize_params::<InitializeRequest>(&fallback_initialize)?;
216
217 if v2_initialize_as_v1
218 .as_ref()
219 .is_ok_and(|v2_initialize| v2_initialize == &v1_initialize)
220 {
221 let fallback_response = initialize_response.with_id(
222 initialize_request_id(&fallback_initialize)
223 .expect("validated initialize request has an id"),
224 );
225 drop(client);
230 fallback_client.send(fallback_response)?;
231 return pipe_protocol_peers_until_done(fallback_client, agent_connection)
232 .await;
233 }
234
235 drop((
239 client,
240 fallback_client,
241 agent_connection,
242 initialize_response,
243 ));
244 return connect_client_protocol(ClientProtocol::V1, v1.create(), agent()).await;
245 }
246
247 client.send(initialize_response.into_message())?;
248 pipe_protocol_peers_until_done(client, agent_connection).await
249 }
250 }
251 }
252}
253
254#[cfg(feature = "unstable_protocol_v2")]
255struct DynConnectToFactory<R: Role> {
256 inner: Box<dyn FnMut() -> DynConnectTo<R> + Send>,
257}
258
259#[cfg(feature = "unstable_protocol_v2")]
260impl<R: Role> DynConnectToFactory<R> {
261 fn new<C>(mut factory: impl FnMut() -> C + Send + 'static) -> Self
262 where
263 C: ConnectTo<R>,
264 {
265 Self {
266 inner: Box::new(move || DynConnectTo::new(factory())),
267 }
268 }
269
270 fn create(&mut self) -> DynConnectTo<R> {
271 (self.inner)()
272 }
273}
274
275#[cfg(feature = "unstable_protocol_v2")]
276impl<R: Role> Debug for DynConnectToFactory<R> {
277 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
278 f.debug_struct("DynConnectToFactory")
279 .finish_non_exhaustive()
280 }
281}
282
283impl HasPeer<Client> for Client {
284 fn remote_style(&self, _peer: Client) -> RemoteStyle {
285 RemoteStyle::Counterpart
286 }
287}
288
289#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
294pub struct Agent;
295
296impl Role for Agent {
297 type Counterpart = Client;
298
299 fn builder(self) -> Builder<Self> {
300 Builder::new(self).v1_agent()
301 }
302
303 fn role_id(&self) -> RoleId {
304 RoleId::from_singleton(self)
305 }
306
307 fn counterpart(&self) -> Self::Counterpart {
308 Client
309 }
310
311 async fn default_handle_dispatch_from(
312 &self,
313 message: Dispatch,
314 connection: ConnectionTo<Agent>,
315 ) -> Result<Handled<Dispatch>, crate::Error> {
316 MatchDispatchFrom::new(message, &connection)
317 .if_dispatch_from(Agent, async |message: Dispatch| {
318 #[cfg(feature = "unstable_protocol_v2")]
326 let retry = message.has_session_id()
327 && connection.acp_protocol_version()
328 != Some(crate::schema::ProtocolVersion::V2);
329 #[cfg(not(feature = "unstable_protocol_v2"))]
330 let retry = message.has_session_id();
331 Ok(Handled::No { message, retry })
332 })
333 .await
334 .done()
335 }
336}
337
338impl Agent {
339 pub fn builder(self) -> Builder<Agent, NullHandler, NullRun> {
341 <Self as Role>::builder(self)
342 }
343
344 #[cfg(feature = "unstable_protocol_v2")]
351 pub fn v2(self) -> V2Builder<Agent, NullHandler, NullRun> {
352 self.builder().v2_agent()
353 }
354
355 #[cfg(feature = "unstable_protocol_v2")]
368 #[must_use]
369 pub fn protocol_router(self) -> AgentProtocolRouter {
370 AgentProtocolRouter::new()
371 }
372}
373
374#[cfg(feature = "unstable_protocol_v2")]
380#[derive(Debug, Default)]
381pub struct AgentProtocolRouter {
382 v1: Option<DynConnectTo<Client>>,
383 v2: Option<DynConnectTo<Client>>,
384}
385
386#[cfg(feature = "unstable_protocol_v2")]
387impl AgentProtocolRouter {
388 #[must_use]
390 pub fn new() -> Self {
391 Self::default()
392 }
393
394 #[must_use]
396 pub fn with_v1(mut self, agent: impl ConnectTo<Client>) -> Self {
397 self.v1 = Some(DynConnectTo::new(agent));
398 self
399 }
400
401 #[must_use]
403 pub fn with_v2(mut self, agent: impl ConnectTo<Client>) -> Self {
404 self.v2 = Some(DynConnectTo::new(agent));
405 self
406 }
407}
408
409#[cfg(feature = "unstable_protocol_v2")]
410impl ConnectTo<Client> for AgentProtocolRouter {
411 async fn connect_to(self, client: impl ConnectTo<Agent>) -> Result<(), crate::Error> {
412 let supported = SupportedProtocols {
413 v1: self.v1.is_some(),
414 v2: self.v2.is_some(),
415 };
416 let mut client = RunningProtocolPeer::new(client);
417 let (first_frame, client, selected) = loop {
418 let Some((mut frame, next_client)) = client.next_frame().await? else {
419 return Ok(());
420 };
421 let message = match initialize_message_mut(&mut frame) {
422 Ok(Some(message)) => message,
423 Ok(None) => {
424 client = next_client;
425 continue;
426 }
427 Err(error) => return reject_initialize(next_client, &frame, error).await,
428 };
429 let selected = match select_agent_protocol(message, supported) {
430 Ok(selected) => selected,
431 Err(error) => return reject_initialize(next_client, &frame, error).await,
432 };
433 break (frame, next_client, selected);
434 };
435 let Some(agent) = selected.take_agent(self) else {
436 let error = selected.unsupported_error(supported);
437 return reject_initialize(client, &first_frame, error).await;
438 };
439
440 let agent = RunningProtocolPeer::new(agent);
441 agent.send_frame(first_frame)?;
442 pipe_protocol_peers_until_done(client, agent).await
443 }
444}
445
446#[cfg(feature = "unstable_protocol_v2")]
447#[derive(Debug, Clone, Copy, PartialEq, Eq)]
448enum SelectedProtocol {
449 V1,
450 V2,
451}
452
453#[cfg(feature = "unstable_protocol_v2")]
454impl SelectedProtocol {
455 fn take_agent(self, agent: AgentProtocolRouter) -> Option<DynConnectTo<Client>> {
456 match self {
457 Self::V1 => agent.v1,
458 Self::V2 => agent.v2,
459 }
460 }
461
462 fn version(self) -> ProtocolVersion {
463 match self {
464 Self::V1 => ProtocolVersion::V1,
465 Self::V2 => ProtocolVersion::V2,
466 }
467 }
468
469 fn name(self) -> &'static str {
470 match self {
471 Self::V1 => "1",
472 Self::V2 => "2",
473 }
474 }
475
476 fn unsupported_error(self, supported: SupportedProtocols) -> crate::Error {
477 crate::Error::invalid_request().data(format!(
478 "ACP protocol version {} is not configured; this endpoint supports {}",
479 self.name(),
480 supported.description()
481 ))
482 }
483}
484
485#[cfg(feature = "unstable_protocol_v2")]
486#[derive(Debug, Clone, Copy, PartialEq, Eq)]
487struct SupportedProtocols {
488 v1: bool,
489 v2: bool,
490}
491
492#[cfg(feature = "unstable_protocol_v2")]
493impl SupportedProtocols {
494 fn highest_compatible(self, requested: ProtocolVersion) -> Option<SelectedProtocol> {
495 if self.v2 && requested >= ProtocolVersion::V2 {
496 return Some(SelectedProtocol::V2);
497 }
498
499 if self.v1 && requested >= ProtocolVersion::V1 {
500 return Some(SelectedProtocol::V1);
501 }
502
503 None
504 }
505
506 fn exact(self, requested: ProtocolVersion) -> Option<SelectedProtocol> {
507 if self.v1 && requested == ProtocolVersion::V1 {
508 Some(SelectedProtocol::V1)
509 } else if self.v2 && requested == ProtocolVersion::V2 {
510 Some(SelectedProtocol::V2)
511 } else {
512 None
513 }
514 }
515
516 fn description(self) -> String {
517 match (self.v1, self.v2) {
518 (true, true) => "ACP protocol versions 1 and 2".into(),
519 (true, false) => "ACP protocol version 1".into(),
520 (false, true) => "ACP protocol version 2".into(),
521 (false, false) => "no ACP protocol versions".into(),
522 }
523 }
524}
525
526#[cfg(feature = "unstable_protocol_v2")]
527fn select_agent_protocol(
528 message: &mut RawJsonRpcMessage,
529 supported: SupportedProtocols,
530) -> Result<SelectedProtocol, crate::Error> {
531 let RawJsonRpcMessage::Request(request) = message else {
532 return Err(
533 crate::Error::invalid_request().data("first ACP message must be an initialize request")
534 );
535 };
536
537 if request.method.as_ref() != "initialize" {
538 return Err(crate::Error::invalid_request().data("first ACP request must be initialize"));
539 }
540
541 let Some(RawJsonRpcParams::Object(params)) = &mut request.params else {
542 return Err(invalid_initialize_protocol_version());
543 };
544 let Some(protocol_version) = params.get("protocolVersion") else {
545 return Err(invalid_initialize_protocol_version());
546 };
547
548 let requested = serde_json::from_value::<ProtocolVersion>(protocol_version.clone())
549 .map_err(|_| invalid_initialize_protocol_version())?;
550 let selected = highest_compatible_agent_protocol(requested, supported)?;
551 rewrite_initialize_params(params, requested, selected)?;
552
553 Ok(selected)
554}
555
556#[cfg(feature = "unstable_protocol_v2")]
557fn initialize_request_params(
558 message: &RawJsonRpcMessage,
559) -> Result<&serde_json::Map<String, serde_json::Value>, crate::Error> {
560 let RawJsonRpcMessage::Request(request) = message else {
561 return Err(
562 crate::Error::invalid_request().data("first ACP message must be an initialize request")
563 );
564 };
565
566 if request.method.as_ref() != "initialize" {
567 return Err(crate::Error::invalid_request().data("first ACP request must be initialize"));
568 }
569
570 let Some(RawJsonRpcParams::Object(params)) = &request.params else {
571 return Err(invalid_initialize_protocol_version());
572 };
573 if !params.contains_key("protocolVersion") {
574 return Err(invalid_initialize_protocol_version());
575 }
576 Ok(params)
577}
578
579#[cfg(feature = "unstable_protocol_v2")]
580fn validated_initialize_params<T: DeserializeOwned>(
581 message: &RawJsonRpcMessage,
582) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
583 let params = initialize_request_params(message)?;
584 parse_initialize_params::<T>(params)?;
585 Ok(params.clone())
586}
587
588#[cfg(feature = "unstable_protocol_v2")]
589fn normalize_v2_initialize_params_for_reuse(
590 message: &RawJsonRpcMessage,
591) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
592 let params = initialize_request_params(message)?;
593 let requested = params
594 .get("protocolVersion")
595 .cloned()
596 .ok_or_else(invalid_initialize_protocol_version)
597 .and_then(|version| {
598 serde_json::from_value::<ProtocolVersion>(version)
599 .map_err(|_| invalid_initialize_protocol_version())
600 })?;
601 if requested == ProtocolVersion::V1 {
602 parse_initialize_params::<InitializeRequest>(params)?;
603 return Ok(params.clone());
604 }
605 normalize_v2_initialize_params_for_v1(params, true)
606}
607
608#[cfg(feature = "unstable_protocol_v2")]
609fn rewrite_initialize_params(
610 params: &mut serde_json::Map<String, serde_json::Value>,
611 requested: ProtocolVersion,
612 selected: SelectedProtocol,
613) -> Result<(), crate::Error> {
614 if requested == selected.version() {
618 match selected {
619 SelectedProtocol::V1 => {
620 parse_initialize_params::<InitializeRequest>(params)?;
621 }
622 SelectedProtocol::V2 => {
623 parse_initialize_params::<v2::InitializeRequest>(params)?;
624 }
625 }
626 return Ok(());
627 }
628
629 match selected {
630 SelectedProtocol::V1 => {
631 debug_assert!(requested >= ProtocolVersion::V2);
632 *params = normalize_v2_initialize_params_for_v1(params, false)?;
633 Ok(())
634 }
635 SelectedProtocol::V2 => {
636 let mut initialize = parse_initialize_params::<v2::InitializeRequest>(params)?;
637 initialize.protocol_version = ProtocolVersion::V2;
638 *params = serialize_initialize_params(initialize)?;
639 Ok(())
640 }
641 }
642}
643
644#[cfg(feature = "unstable_protocol_v2")]
645fn normalize_v2_initialize_params_for_v1(
646 params: &serde_json::Map<String, serde_json::Value>,
647 require_lossless: bool,
648) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
649 let initialize = parse_initialize_params::<v2::InitializeRequest>(params)?;
654 let mut target = serialize_initialize_params(initialize)?;
655 if require_lossless && target != *params {
656 return Err(invalid_initialize_params(
657 "v2 initialize parameters are not losslessly representable in v1",
658 ));
659 }
660
661 target.insert(
662 "protocolVersion".into(),
663 serde_json::to_value(ProtocolVersion::V1).map_err(crate::Error::into_internal_error)?,
664 );
665 let info = target
666 .remove("info")
667 .ok_or_else(|| invalid_initialize_params("v2 InitializeRequest.info is required"))?;
668 target.insert("clientInfo".into(), info);
669 let capabilities = target
670 .remove("capabilities")
671 .and_then(|capabilities| capabilities.as_object().cloned())
672 .ok_or_else(|| {
673 crate::util::internal_error("v2 initialize capabilities did not serialize as an object")
674 })?;
675 let mut capabilities = capabilities;
676 if let Some(auth) = capabilities
677 .get_mut("auth")
678 .and_then(serde_json::Value::as_object_mut)
679 {
680 let terminal = auth.remove("terminal");
681 if require_lossless
682 && terminal
683 .as_ref()
684 .and_then(serde_json::Value::as_object)
685 .is_some_and(|terminal| terminal.contains_key("_meta"))
686 {
687 return Err(invalid_initialize_params(
688 "v2 terminal authentication metadata is not representable in v1",
689 ));
690 }
691 auth.insert("terminal".into(), terminal.is_some().into());
692 }
693 capabilities.insert(
694 "session".into(),
695 serde_json::json!({ "configOptions": { "boolean": {} } }),
696 );
697 target.insert("clientCapabilities".into(), capabilities.into());
698
699 let initialize = parse_initialize_params::<InitializeRequest>(&target)?;
700 let normalized = serialize_initialize_params(initialize)?;
701 if require_lossless && !json_object_contains(&normalized, &target) {
702 return Err(invalid_initialize_params(
703 "v2 initialize parameters are not losslessly representable in v1",
704 ));
705 }
706 Ok(normalized)
707}
708
709#[cfg(all(test, feature = "unstable_protocol_v2"))]
710mod initialize_normalization_tests {
711 use super::*;
712
713 fn v2_initialize_params() -> serde_json::Map<String, serde_json::Value> {
714 let value = serde_json::to_value(v2::InitializeRequest::new(
715 ProtocolVersion::V2,
716 v2::Implementation::new("test-client", "1.0.0"),
717 ))
718 .expect("serialize v2 initialize request");
719 value
720 .as_object()
721 .expect("initialize params serialize as an object")
722 .clone()
723 }
724
725 #[test]
726 fn v2_tolerant_fields_are_canonicalized_before_v1_normalization() {
727 let mut params = v2_initialize_params();
728 params.insert(
729 "capabilities".into(),
730 serde_json::Value::String("malformed".into()),
731 );
732 params.insert(
733 "_meta".into(),
734 serde_json::Value::String("malformed".into()),
735 );
736
737 let normalized = normalize_v2_initialize_params_for_v1(¶ms, false)
738 .expect("tolerant v2 fields should normalize through their defaults");
739 let normalized = serde_json::Value::Object(normalized);
740
741 assert!(normalized.get("_meta").is_none());
742 assert_eq!(
743 normalized.pointer("/clientCapabilities/session/configOptions/boolean"),
744 Some(&serde_json::json!({}))
745 );
746 }
747
748 #[test]
749 fn noncanonical_v2_fields_disable_reuse_but_not_v1_routing() {
750 let mut params = v2_initialize_params();
751 params
752 .get_mut("info")
753 .and_then(serde_json::Value::as_object_mut)
754 .expect("v2 initialize info is an object")
755 .insert("buildCommit".into(), serde_json::json!("abc123"));
756
757 normalize_v2_initialize_params_for_v1(¶ms, false)
758 .expect("v1 routing may ignore parameters unavailable in v1");
759 normalize_v2_initialize_params_for_v1(¶ms, true)
760 .expect_err("connection reuse requires lossless normalization");
761 }
762
763 #[test]
764 fn v1_reuse_probe_preserves_raw_initialize_params() {
765 let mut params = v2_initialize_params();
766 params.insert(
767 "protocolVersion".into(),
768 serde_json::json!(ProtocolVersion::V1),
769 );
770 let message = RawJsonRpcMessage::request(
771 "initialize".into(),
772 serde_json::Value::Object(params.clone()),
773 RequestId::Number(1),
774 )
775 .expect("build initialize request");
776
777 let normalized = normalize_v2_initialize_params_for_reuse(&message)
778 .expect("v1-shaped initialize request should be valid");
779
780 assert_eq!(normalized, params);
781 }
782
783 #[test]
784 fn null_v2_terminal_marker_meta_is_omitted_before_v1_normalization() {
785 let mut params = v2_initialize_params();
786 params.insert(
787 "capabilities".into(),
788 serde_json::json!({
789 "auth": {
790 "terminal": { "_meta": null }
791 }
792 }),
793 );
794
795 let normalized = normalize_v2_initialize_params_for_v1(¶ms, false)
796 .expect("null marker metadata is equivalent to omission");
797 let normalized = serde_json::Value::Object(normalized);
798
799 assert_eq!(
800 normalized.pointer("/clientCapabilities/auth/terminal"),
801 Some(&serde_json::Value::Bool(true))
802 );
803 }
804
805 #[test]
806 fn terminal_marker_metadata_disables_reuse_but_not_v1_routing() {
807 let mut params = v2_initialize_params();
808 params.insert(
809 "capabilities".into(),
810 serde_json::json!({
811 "auth": {
812 "terminal": {
813 "_meta": { "source": "test" }
814 }
815 }
816 }),
817 );
818
819 normalize_v2_initialize_params_for_v1(¶ms, false)
820 .expect("v1 routing may discard terminal marker metadata");
821 normalize_v2_initialize_params_for_v1(¶ms, true)
822 .expect_err("connection reuse must preserve terminal marker metadata");
823 }
824}
825
826#[cfg(feature = "unstable_protocol_v2")]
827fn parse_initialize_params<T: DeserializeOwned>(
828 params: &serde_json::Map<String, serde_json::Value>,
829) -> Result<T, crate::Error> {
830 serde_json::from_value(serde_json::Value::Object(params.clone()))
831 .map_err(invalid_initialize_params)
832}
833
834#[cfg(feature = "unstable_protocol_v2")]
835fn serialize_initialize_params(
836 initialize: impl Serialize,
837) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
838 let value = serde_json::to_value(initialize).map_err(crate::Error::into_internal_error)?;
839 let serde_json::Value::Object(object) = value else {
840 return Err(crate::util::internal_error(
841 "initialize params did not serialize to an object",
842 ));
843 };
844 Ok(object)
845}
846
847#[cfg(feature = "unstable_protocol_v2")]
848fn json_object_contains(
849 actual: &serde_json::Map<String, serde_json::Value>,
850 expected: &serde_json::Map<String, serde_json::Value>,
851) -> bool {
852 fn contains(actual: &serde_json::Value, expected: &serde_json::Value) -> bool {
853 match (actual, expected) {
854 (serde_json::Value::Object(actual), serde_json::Value::Object(expected)) => expected
855 .iter()
856 .all(|(key, value)| actual.get(key).is_some_and(|item| contains(item, value))),
857 _ => actual == expected,
858 }
859 }
860
861 expected
862 .iter()
863 .all(|(key, value)| actual.get(key).is_some_and(|item| contains(item, value)))
864}
865
866#[cfg(feature = "unstable_protocol_v2")]
867fn highest_compatible_agent_protocol(
868 requested: ProtocolVersion,
869 supported: SupportedProtocols,
870) -> Result<SelectedProtocol, crate::Error> {
871 supported.highest_compatible(requested).ok_or_else(|| {
872 crate::Error::invalid_request().data(format!(
873 "unsupported ACP protocol version {requested}; this endpoint supports {}",
874 supported.description()
875 ))
876 })
877}
878
879#[cfg(feature = "unstable_protocol_v2")]
880fn invalid_initialize_protocol_version() -> crate::Error {
881 crate::Error::invalid_params()
882 .data("initialize.protocolVersion must be a valid ACP protocol version")
883}
884
885#[cfg(feature = "unstable_protocol_v2")]
886fn invalid_initialize_params(error: impl ToString) -> crate::Error {
887 crate::Error::invalid_params().data(format!("invalid initialize params: {}", error.to_string()))
888}
889
890#[cfg(feature = "unstable_protocol_v2")]
891fn send_initialize_error(
892 tx: &futures::channel::mpsc::UnboundedSender<TransportFrame>,
893 frame: &TransportFrame,
894 error: crate::Error,
895) -> Result<(), crate::Error> {
896 fn response_for_message(
897 entry: &RawJsonRpcMessage,
898 initialize_error: &crate::Error,
899 ) -> Option<RawJsonRpcMessage> {
900 match entry {
901 RawJsonRpcMessage::Request(request) => Some(RawJsonRpcMessage::response(
902 request.id.clone(),
903 Err(initialize_error.clone()),
904 )),
905 RawJsonRpcMessage::Notification(_) | RawJsonRpcMessage::Response(_) => None,
906 }
907 }
908
909 fn response_for_entry(
910 entry: &TransportBatchEntry,
911 initialize_error: &crate::Error,
912 ) -> Option<RawJsonRpcMessage> {
913 match entry {
914 TransportBatchEntry::Message(message) => {
915 response_for_message(message, initialize_error)
916 }
917 TransportBatchEntry::Malformed { raw, error } if !is_response_only_shape(raw) => Some(
918 RawJsonRpcMessage::response(RequestId::Null, Err(error.clone())),
919 ),
920 TransportBatchEntry::Malformed { .. } => None,
921 }
922 }
923
924 let response = match frame {
925 TransportFrame::Single(entry) => {
926 let Some(response) = response_for_message(entry, &error) else {
927 return Ok(());
928 };
929 TransportFrame::Single(response)
930 }
931 TransportFrame::Malformed { raw, error } if !raw_is_response_only_shape(raw) => {
932 TransportFrame::Single(RawJsonRpcMessage::response(
933 RequestId::Null,
934 Err(error.clone()),
935 ))
936 }
937 TransportFrame::Malformed { .. } => return Ok(()),
938 TransportFrame::Batch(batch) => {
939 let responses = batch
940 .entries()
941 .filter_map(|entry| response_for_entry(entry, &error))
942 .collect::<Vec<_>>();
943 let Some(responses) = TransportBatch::from_messages(responses) else {
944 return Ok(());
945 };
946 TransportFrame::Batch(responses)
947 }
948 };
949
950 tx.unbounded_send(response)
951 .map_err(crate::util::internal_error)
952}
953
954#[cfg(feature = "unstable_protocol_v2")]
955async fn reject_initialize(
956 client: RunningProtocolPeer,
957 frame: &TransportFrame,
958 error: crate::Error,
959) -> Result<(), crate::Error> {
960 let RunningProtocolPeer { mut rx, tx, driver } = client;
961 send_initialize_error(&tx, frame, error)?;
962 drop(tx);
963
964 let Some(mut driver) = driver.into_driver() else {
965 return Ok(());
968 };
969 if !driver.request_finish() {
970 return crate::util::run_until(driver, future::ready(Ok(()))).await;
973 }
974
975 let drain_incoming = async move {
976 while rx.next().await.is_some() {}
980 Ok::<_, crate::Error>(())
981 };
982
983 match future::select(driver, Box::pin(drain_incoming)).await {
984 future::Either::Left((result, _)) => result,
985 future::Either::Right((result, driver)) => {
986 result?;
987 driver.await
988 }
989 }
990}
991
992#[cfg(feature = "unstable_protocol_v2")]
993struct RunningProtocolPeer {
994 rx: futures::channel::mpsc::UnboundedReceiver<TransportFrame>,
995 tx: futures::channel::mpsc::UnboundedSender<TransportFrame>,
996 driver: ProtocolPeerDriver,
997}
998
999#[cfg(feature = "unstable_protocol_v2")]
1000enum ProtocolPeerDriver {
1001 Passive,
1002 Active(crate::ConnectionDriver),
1003 Completed {
1004 finish: Option<crate::component::FinishControl>,
1005 },
1006}
1007
1008#[cfg(feature = "unstable_protocol_v2")]
1009impl ProtocolPeerDriver {
1010 fn into_driver(self) -> Option<crate::ConnectionDriver> {
1011 match self {
1012 Self::Passive => None,
1013 Self::Active(driver) => Some(driver),
1014 Self::Completed { finish } => Some(match finish {
1018 Some(mut finish) => {
1019 crate::ConnectionDriver::with_finish(future::ready(Ok(())), move || {
1020 finish.request();
1021 })
1022 }
1023 None => crate::ConnectionDriver::new(future::ready(Ok(()))),
1024 }),
1025 }
1026 }
1027}
1028
1029#[cfg(feature = "unstable_protocol_v2")]
1030impl RunningProtocolPeer {
1031 fn new<R: Role>(component: impl ConnectTo<R>) -> Self {
1032 let (Channel { rx, tx }, future) = component.into_channel_and_future();
1033 let driver = match future {
1034 None => ProtocolPeerDriver::Passive,
1035 Some(future) => ProtocolPeerDriver::Active(future),
1036 };
1037 Self { rx, tx, driver }
1038 }
1039
1040 async fn next_frame(self) -> Result<Option<(TransportFrame, Self)>, crate::Error> {
1041 let Self { mut rx, tx, driver } = self;
1042 let ProtocolPeerDriver::Active(mut future) = driver else {
1043 return Ok(rx
1044 .next()
1045 .await
1046 .map(|frame| (frame, Self { rx, tx, driver })));
1047 };
1048
1049 match future::select(&mut future, Box::pin(rx.next())).await {
1052 future::Either::Right((Some(frame), _)) => Ok(Some((
1053 frame,
1054 Self {
1055 rx,
1056 tx,
1057 driver: ProtocolPeerDriver::Active(future),
1058 },
1059 ))),
1060 future::Either::Right((None, _)) => {
1061 drop(tx);
1062 future.await?;
1063 Ok(None)
1064 }
1065 future::Either::Left((result, next_message)) => {
1066 result?;
1067 drop(next_message);
1068 rx.close();
1071 let Some(frame) = rx.next().await else {
1072 return Ok(None);
1073 };
1074 Ok(Some((
1075 frame,
1076 Self {
1077 rx,
1078 tx,
1079 driver: ProtocolPeerDriver::Completed {
1080 finish: future.take_finish(),
1081 },
1082 },
1083 )))
1084 }
1085 }
1086 }
1087
1088 async fn next_message(self) -> Result<Option<(RawJsonRpcMessage, Self)>, crate::Error> {
1089 let Some((frame, peer)) = self.next_frame().await? else {
1090 return Ok(None);
1091 };
1092 Ok(Some((initialize_message(frame)?, peer)))
1093 }
1094
1095 fn send(&self, message: RawJsonRpcMessage) -> Result<(), crate::Error> {
1096 self.send_frame(TransportFrame::Single(message))
1097 }
1098
1099 fn send_frame(&self, frame: TransportFrame) -> Result<(), crate::Error> {
1100 self.tx
1101 .unbounded_send(frame)
1102 .map_err(crate::util::internal_error)
1103 }
1104}
1105
1106#[cfg(feature = "unstable_protocol_v2")]
1107fn initialize_message(frame: TransportFrame) -> Result<RawJsonRpcMessage, crate::Error> {
1108 match frame {
1109 TransportFrame::Single(message) => Ok(message),
1110 TransportFrame::Malformed { error, .. } => Err(error),
1111 TransportFrame::Batch(_) => Err(crate::Error::invalid_request()
1112 .data("ACP initialize request and response messages must be sent individually")),
1113 }
1114}
1115
1116#[cfg(feature = "unstable_protocol_v2")]
1117fn initialize_message_mut(
1118 frame: &mut TransportFrame,
1119) -> Result<Option<&mut RawJsonRpcMessage>, crate::Error> {
1120 match frame {
1121 TransportFrame::Single(RawJsonRpcMessage::Response(_)) => Ok(None),
1122 TransportFrame::Single(entry) => Ok(Some(entry)),
1123 TransportFrame::Malformed { raw, .. } if raw_is_response_only_shape(raw) => Ok(None),
1124 TransportFrame::Malformed { error, .. } => Err(error.clone()),
1125 TransportFrame::Batch(batch) => {
1126 for entry in batch.entries_mut() {
1127 match entry {
1128 TransportBatchEntry::Message(RawJsonRpcMessage::Response(_)) => {}
1129 TransportBatchEntry::Message(message) => return Ok(Some(message)),
1130 TransportBatchEntry::Malformed { raw, .. } if is_response_only_shape(raw) => {}
1131 TransportBatchEntry::Malformed { error, .. } => return Err(error.clone()),
1132 }
1133 }
1134 Ok(None)
1135 }
1136 }
1137}
1138
1139#[cfg(feature = "unstable_protocol_v2")]
1143async fn pipe_protocol_peers_until_done(
1144 left: RunningProtocolPeer,
1145 right: RunningProtocolPeer,
1146) -> Result<(), crate::Error> {
1147 let mut left_driver = left.driver.into_driver();
1148 let mut right_driver = right.driver.into_driver();
1149 let left_passive = left_driver.is_none();
1150 let right_passive = right_driver.is_none();
1151 let left_finish = left_driver
1152 .as_mut()
1153 .and_then(crate::ConnectionDriver::take_finish);
1154 let right_finish = right_driver
1155 .as_mut()
1156 .and_then(crate::ConnectionDriver::take_finish);
1157 let (stop_left_tx, stop_left_rx) = futures::channel::oneshot::channel();
1158 let (stop_right_tx, stop_right_rx) = futures::channel::oneshot::channel();
1159 let stop = async |rx: futures::channel::oneshot::Receiver<()>| {
1160 if rx.await.is_err() {
1161 future::pending::<()>().await;
1162 }
1163 };
1164 let left_to_right = Box::pin(
1165 Channel {
1166 rx: left.rx,
1167 tx: right.tx,
1168 }
1169 .copy_with_driver_until(left_driver, stop(stop_left_rx)),
1170 );
1171 let right_to_left = Box::pin(
1172 Channel {
1173 rx: right.rx,
1174 tx: left.tx,
1175 }
1176 .copy_with_driver_until(right_driver, stop(stop_right_rx)),
1177 );
1178
1179 match future::select(left_to_right, right_to_left).await {
1180 future::Either::Left((result, right_to_left)) => {
1181 result?;
1182 if !left_passive {
1183 let _ = stop_right_tx.send(());
1184 }
1185 if left_passive || right_finish.is_some() {
1186 if !left_passive && let Some(mut finish) = right_finish {
1187 finish.request();
1188 }
1189 right_to_left.await
1190 } else {
1191 crate::util::run_until(right_to_left, future::ready(Ok(()))).await
1194 }
1195 }
1196 future::Either::Right((result, left_to_right)) => {
1197 result?;
1198 if !right_passive {
1199 let _ = stop_left_tx.send(());
1200 }
1201 if right_passive || left_finish.is_some() {
1202 if !right_passive && let Some(mut finish) = left_finish {
1203 finish.request();
1204 }
1205 left_to_right.await
1206 } else {
1207 crate::util::run_until(left_to_right, future::ready(Ok(()))).await
1208 }
1209 }
1210 }
1211}
1212
1213#[cfg(feature = "unstable_protocol_v2")]
1214#[derive(Debug)]
1215struct InitializeResponse {
1216 id: RequestId,
1217 result: Result<serde_json::Value, Box<RawJsonRpcError>>,
1218}
1219
1220#[cfg(feature = "unstable_protocol_v2")]
1221impl InitializeResponse {
1222 fn from_message(message: RawJsonRpcMessage) -> Result<Self, crate::Error> {
1223 match message {
1224 RawJsonRpcMessage::Response(RpcResponse::Result { id, result }) => Ok(Self {
1225 id,
1226 result: Ok(result),
1227 }),
1228 RawJsonRpcMessage::Response(RpcResponse::Error { id, error }) => Ok(Self {
1229 id,
1230 result: Err(error),
1231 }),
1232 message => Err(crate::Error::invalid_request().data(format!(
1233 "first ACP response must be an initialize response, got {message:?}",
1234 ))),
1235 }
1236 }
1237
1238 fn into_message(self) -> RawJsonRpcMessage {
1239 RawJsonRpcMessage::Response(RpcResponse::new(self.id, self.result))
1240 }
1241
1242 fn with_id(self, id: RequestId) -> RawJsonRpcMessage {
1243 RawJsonRpcMessage::Response(RpcResponse::new(id, self.result))
1244 }
1245
1246 fn protocol_version(&self) -> Option<ProtocolVersion> {
1247 serde_json::from_value(self.result.as_ref().ok()?.get("protocolVersion")?.clone()).ok()
1248 }
1249}
1250
1251#[cfg(all(test, feature = "unstable_protocol_v2"))]
1252mod raw_initialize_tests {
1253 use super::*;
1254 use serde_json::json;
1255
1256 #[test]
1257 fn initialize_error_forwarding_preserves_raw_fields() {
1258 for data in [
1259 None,
1260 Some(serde_json::Value::Null),
1261 Some(json!({"detail":"kept"})),
1262 ] {
1263 let mut error = json!({
1264 "code":-32000, "message":"peer", "extension":{"retry":true}
1265 });
1266 if let Some(data) = data {
1267 error["data"] = data;
1268 }
1269 let wire = json!({"jsonrpc":"2.0", "id":"original", "error":error});
1270 let response =
1271 InitializeResponse::from_message(serde_json::from_value(wire.clone()).unwrap())
1272 .unwrap();
1273 assert_eq!(serde_json::to_value(response.into_message()).unwrap(), wire);
1274 let response =
1275 InitializeResponse::from_message(serde_json::from_value(wire).unwrap()).unwrap();
1276 let forwarded = response.with_id(RequestId::Str("replacement".into()));
1277 assert_eq!(
1278 serde_json::to_value(forwarded).unwrap(),
1279 json!({"jsonrpc":"2.0", "id":"replacement", "error":error})
1280 );
1281 }
1282 }
1283}
1284
1285#[cfg(feature = "unstable_protocol_v2")]
1286#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1287enum ClientProtocol {
1288 V1,
1289 V2,
1290}
1291
1292#[cfg(feature = "unstable_protocol_v2")]
1293impl ClientProtocol {
1294 fn name(self) -> &'static str {
1295 match self {
1296 Self::V1 => "1",
1297 Self::V2 => "2",
1298 }
1299 }
1300}
1301
1302#[cfg(feature = "unstable_protocol_v2")]
1303#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1304struct SupportedClientProtocols {
1305 v1: bool,
1306 v2: bool,
1307}
1308
1309#[cfg(feature = "unstable_protocol_v2")]
1310impl SupportedClientProtocols {
1311 fn highest_configured(self) -> Option<ClientProtocol> {
1312 if self.v2 {
1313 return Some(ClientProtocol::V2);
1314 }
1315
1316 if self.v1 {
1317 return Some(ClientProtocol::V1);
1318 }
1319
1320 None
1321 }
1322}
1323
1324#[cfg(feature = "unstable_protocol_v2")]
1325async fn start_client_protocol(
1326 protocol: ClientProtocol,
1327 client: DynConnectTo<Agent>,
1328) -> Result<(RunningProtocolPeer, RawJsonRpcMessage), crate::Error> {
1329 let client = RunningProtocolPeer::new(client);
1330 let Some((initialize, client)) = client.next_message().await? else {
1331 return Err(crate::Error::invalid_request().data(format!(
1332 "ACP protocol version {} client implementation ended before initialize",
1333 protocol.name()
1334 )));
1335 };
1336 ensure_client_initialize_request(protocol, &initialize)?;
1337 Ok((client, initialize))
1338}
1339
1340#[cfg(feature = "unstable_protocol_v2")]
1341async fn send_initialize_and_receive(
1342 client: RunningProtocolPeer,
1343 agent: RunningProtocolPeer,
1344 initialize: RawJsonRpcMessage,
1345) -> Result<(RunningProtocolPeer, RunningProtocolPeer, InitializeResponse), crate::Error> {
1346 agent.send(initialize)?;
1347 let Some((response, agent)) = agent.next_message().await? else {
1348 return Err(crate::Error::internal_error().data("agent closed before initialize response"));
1349 };
1350 let response = InitializeResponse::from_message(response)?;
1351 Ok((client, agent, response))
1352}
1353
1354#[cfg(feature = "unstable_protocol_v2")]
1355async fn initialize_client_protocol(
1356 protocol: ClientProtocol,
1357 client: DynConnectTo<Agent>,
1358 agent: impl ConnectTo<Client>,
1359) -> Result<(RunningProtocolPeer, RunningProtocolPeer, InitializeResponse), crate::Error> {
1360 let agent = RunningProtocolPeer::new(agent);
1361 let (client, initialize) = start_client_protocol(protocol, client).await?;
1362 send_initialize_and_receive(client, agent, initialize).await
1363}
1364
1365#[cfg(feature = "unstable_protocol_v2")]
1366async fn connect_client_protocol(
1367 protocol: ClientProtocol,
1368 client: DynConnectTo<Agent>,
1369 agent: impl ConnectTo<Client>,
1370) -> Result<(), crate::Error> {
1371 let (client, agent, initialize_response) =
1372 initialize_client_protocol(protocol, client, agent).await?;
1373 client.send(initialize_response.into_message())?;
1374 pipe_protocol_peers_until_done(client, agent).await
1375}
1376
1377#[cfg(feature = "unstable_protocol_v2")]
1378fn ensure_client_initialize_request(
1379 protocol: ClientProtocol,
1380 message: &RawJsonRpcMessage,
1381) -> Result<(), crate::Error> {
1382 let RawJsonRpcMessage::Request(request) = message else {
1383 return Err(crate::Error::invalid_request().data(format!(
1384 "ACP protocol version {} client implementation must send initialize first",
1385 protocol.name()
1386 )));
1387 };
1388
1389 if request.method.as_ref() != "initialize" {
1390 return Err(crate::Error::invalid_request().data(format!(
1391 "ACP protocol version {} client implementation must send initialize first",
1392 protocol.name()
1393 )));
1394 }
1395
1396 Ok(())
1397}
1398
1399#[cfg(feature = "unstable_protocol_v2")]
1400fn initialize_request_id(message: &RawJsonRpcMessage) -> Option<RequestId> {
1401 let RawJsonRpcMessage::Request(request) = message else {
1402 return None;
1403 };
1404 Some(request.id.clone())
1405}
1406
1407#[cfg(feature = "unstable_protocol_v2")]
1408fn initialize_response_negotiated_v1(response: &InitializeResponse) -> bool {
1409 response.protocol_version() == Some(ProtocolVersion::V1)
1410}
1411
1412impl HasPeer<Agent> for Agent {
1413 fn remote_style(&self, _peer: Agent) -> RemoteStyle {
1414 RemoteStyle::Counterpart
1415 }
1416}
1417
1418#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
1427pub struct Proxy;
1428
1429impl Role for Proxy {
1430 type Counterpart = Conductor;
1431
1432 fn default_handle_dispatch_from(
1433 &self,
1434 message: crate::Dispatch,
1435 _connection: crate::ConnectionTo<Self>,
1436 ) -> impl Future<Output = Result<crate::Handled<crate::Dispatch>, crate::Error>> + Send {
1437 std::future::ready(Ok(Handled::No {
1438 message,
1439 retry: false,
1440 }))
1441 }
1442
1443 fn role_id(&self) -> RoleId {
1444 RoleId::from_singleton(self)
1445 }
1446
1447 fn counterpart(&self) -> Self::Counterpart {
1448 Conductor
1449 }
1450}
1451
1452impl Proxy {
1453 pub fn builder(self) -> Builder<Proxy, NullHandler, NullRun> {
1460 Builder::new(self)
1461 }
1462
1463 #[cfg(feature = "unstable_protocol_v2")]
1472 pub fn v2(self) -> V2Builder<Proxy, NullHandler, NullRun> {
1473 self.builder().v2_proxy()
1474 }
1475
1476 #[cfg(feature = "unstable_protocol_v2")]
1487 #[must_use]
1488 pub fn protocol_router(self) -> ProxyProtocolRouter {
1489 ProxyProtocolRouter::new()
1490 }
1491}
1492
1493#[cfg(feature = "unstable_protocol_v2")]
1501#[derive(Debug, Default)]
1502pub struct ProxyProtocolRouter {
1503 v1: Option<DynConnectTo<Conductor>>,
1504 v2: Option<DynConnectTo<Conductor>>,
1505}
1506
1507#[cfg(feature = "unstable_protocol_v2")]
1508impl ProxyProtocolRouter {
1509 #[must_use]
1511 pub fn new() -> Self {
1512 Self::default()
1513 }
1514
1515 #[must_use]
1517 pub fn with_v1(mut self, proxy: impl ConnectTo<Conductor>) -> Self {
1518 self.v1 = Some(DynConnectTo::new(proxy));
1519 self
1520 }
1521
1522 #[must_use]
1524 pub fn with_v2(mut self, proxy: impl ConnectTo<Conductor>) -> Self {
1525 self.v2 = Some(DynConnectTo::new(proxy));
1526 self
1527 }
1528}
1529
1530#[cfg(feature = "unstable_protocol_v2")]
1531impl ConnectTo<Conductor> for ProxyProtocolRouter {
1532 async fn connect_to(self, conductor: impl ConnectTo<Proxy>) -> Result<(), crate::Error> {
1533 let supported = SupportedProtocols {
1534 v1: self.v1.is_some(),
1535 v2: self.v2.is_some(),
1536 };
1537 let mut conductor = RunningProtocolPeer::new(conductor);
1538 let (first_frame, conductor, selected) = loop {
1539 let Some((mut frame, next_conductor)) = conductor.next_frame().await? else {
1540 return Ok(());
1541 };
1542 let message = match initialize_message_mut(&mut frame) {
1543 Ok(Some(message)) => message,
1544 Ok(None) => {
1545 conductor = next_conductor;
1546 continue;
1547 }
1548 Err(error) => return reject_initialize(next_conductor, &frame, error).await,
1549 };
1550 let selected = match select_proxy_protocol(message, supported) {
1551 Ok(selected) => selected,
1552 Err(error) => return reject_initialize(next_conductor, &frame, error).await,
1553 };
1554 break (frame, next_conductor, selected);
1555 };
1556 let Some(proxy) = selected.take_proxy(self) else {
1557 let error = selected.unsupported_error(supported);
1558 return reject_initialize(conductor, &first_frame, error).await;
1559 };
1560
1561 let proxy = RunningProtocolPeer::new(proxy);
1562 proxy.send_frame(first_frame)?;
1563 pipe_protocol_peers_until_done(conductor, proxy).await
1564 }
1565}
1566
1567#[cfg(feature = "unstable_protocol_v2")]
1568impl SelectedProtocol {
1569 fn take_proxy(self, proxy: ProxyProtocolRouter) -> Option<DynConnectTo<Conductor>> {
1570 match self {
1571 Self::V1 => proxy.v1,
1572 Self::V2 => proxy.v2,
1573 }
1574 }
1575}
1576
1577#[cfg(feature = "unstable_protocol_v2")]
1578fn select_proxy_protocol(
1579 message: &RawJsonRpcMessage,
1580 supported: SupportedProtocols,
1581) -> Result<SelectedProtocol, crate::Error> {
1582 let RawJsonRpcMessage::Request(request) = message else {
1583 return Err(crate::Error::invalid_request()
1584 .data("first ACP proxy message must be an `_proxy/initialize` request"));
1585 };
1586
1587 if request.method.as_ref() != METHOD_INITIALIZE_PROXY {
1588 return Err(crate::Error::invalid_request()
1589 .data("first ACP proxy request must be `_proxy/initialize`"));
1590 }
1591
1592 let Some(RawJsonRpcParams::Object(params)) = &request.params else {
1593 return Err(invalid_initialize_protocol_version());
1594 };
1595 let Some(protocol_version) = params.get("protocolVersion") else {
1596 return Err(invalid_initialize_protocol_version());
1597 };
1598 let requested = serde_json::from_value::<ProtocolVersion>(protocol_version.clone())
1599 .map_err(|_| invalid_initialize_protocol_version())?;
1600 let selected = supported.exact(requested).ok_or_else(|| {
1601 crate::Error::invalid_request().data(format!(
1602 "unsupported ACP protocol version {requested}; this proxy supports {}",
1603 supported.description()
1604 ))
1605 })?;
1606
1607 match selected {
1608 SelectedProtocol::V1 => {
1609 parse_initialize_params::<crate::schema::InitializeProxyRequest>(params)?;
1610 }
1611 SelectedProtocol::V2 => {
1612 parse_initialize_params::<v2::InitializeProxyRequest>(params)?;
1613 }
1614 }
1615 Ok(selected)
1616}
1617
1618impl HasPeer<Proxy> for Proxy {
1619 fn remote_style(&self, _peer: Proxy) -> RemoteStyle {
1620 RemoteStyle::Counterpart
1621 }
1622}
1623
1624#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
1629pub struct Conductor;
1630
1631impl Role for Conductor {
1632 type Counterpart = Proxy;
1633
1634 fn role_id(&self) -> RoleId {
1635 RoleId::from_singleton(self)
1636 }
1637
1638 fn counterpart(&self) -> Self::Counterpart {
1639 Proxy
1640 }
1641
1642 async fn default_handle_dispatch_from(
1643 &self,
1644 message: Dispatch,
1645 cx: ConnectionTo<Conductor>,
1646 ) -> Result<Handled<Dispatch>, crate::Error> {
1647 #[cfg(not(feature = "unstable_protocol_v2"))]
1648 {
1649 MatchDispatchFrom::new(message, &cx)
1650 .if_request_from(Client, async |_req: InitializeRequest, responder| {
1651 responder.respond_with_error(crate::Error::invalid_request().data(format!(
1652 "proxies must be initialized with `{METHOD_INITIALIZE_PROXY}`"
1653 )))
1654 })
1655 .await
1656 .if_request_from(
1657 Client,
1658 async |request: InitializeProxyRequest, responder| {
1659 let InitializeProxyRequest { initialize } = request;
1660 cx.send_ordered_request_to(Agent, initialize)
1661 .forward_response_to(responder)
1662 },
1663 )
1664 .await
1665 .if_request_from(Client, async |request: NewSessionRequest, responder| {
1666 let sent = cx.send_ordered_request_to(Agent, request);
1667 let sent = sent.forward_cancellation_from(responder.cancellation());
1668 sent.on_receiving_result({
1669 let cx = cx.clone();
1670 async move |result| {
1671 if let Ok(NewSessionResponse { session_id, .. }) = &result {
1672 cx.add_dynamic_handler(ProxySessionMessages::new(
1673 session_id.clone(),
1674 ))?
1675 .detach();
1676 }
1677 responder.respond_with_result(result)
1678 }
1679 })
1680 })
1681 .await
1682 .if_dispatch_from(Client, async |message: Dispatch| {
1683 cx.send_proxied_message_to(Agent, message)
1684 })
1685 .await
1686 .if_dispatch_from(Agent, async |message: Dispatch| {
1687 cx.send_proxied_message_to(Client, message)
1688 })
1689 .await
1690 .done()
1691 }
1692
1693 #[cfg(feature = "unstable_protocol_v2")]
1694 {
1695 let message = match message {
1696 Dispatch::Request(request, responder) if request.method() == "initialize" => {
1697 responder.respond_with_error(crate::Error::invalid_request().data(format!(
1698 "proxies must be initialized with `{METHOD_INITIALIZE_PROXY}`"
1699 )))?;
1700 return Ok(Handled::Yes);
1701 }
1702 Dispatch::Request(mut request, responder)
1703 if request.method() == METHOD_INITIALIZE_PROXY =>
1704 {
1705 request.method = "initialize".to_string();
1706 cx.send_ordered_request_to(Agent, request)
1707 .forward_response_to(responder)?;
1708 return Ok(Handled::Yes);
1709 }
1710 Dispatch::Request(request, responder) if request.method() == "session/new" => {
1711 let sent = cx.send_ordered_request_to(Agent, request);
1712 let sent = sent.forward_cancellation_from(responder.cancellation());
1717 sent.on_receiving_result({
1718 let cx = cx.clone();
1719 async move |result| {
1720 let result = result.and_then(|response| {
1721 let envelope: NewSessionResponseEnvelope =
1722 crate::util::json_cast(response.clone())?;
1723 cx.add_dynamic_handler(ProxySessionMessages::new(
1724 envelope.session_id,
1725 ))?
1726 .detach();
1727 Ok(response)
1728 });
1729 responder.respond_with_result(result)
1730 }
1731 })?;
1732 return Ok(Handled::Yes);
1733 }
1734 message => message,
1735 };
1736
1737 MatchDispatchFrom::new(message, &cx)
1738 .if_dispatch_from(Client, async |message: Dispatch| {
1739 cx.send_proxied_message_to(Agent, message)
1740 })
1741 .await
1742 .if_dispatch_from(Agent, async |message: Dispatch| {
1743 cx.send_proxied_message_to(Client, message)
1744 })
1745 .await
1746 .done()
1747 }
1748 }
1749}
1750
1751impl Conductor {
1752 pub fn builder(self) -> Builder<Conductor, NullHandler, NullRun> {
1754 Builder::new(self)
1755 }
1756}
1757
1758impl HasPeer<Client> for Conductor {
1759 fn remote_style(&self, _peer: Client) -> RemoteStyle {
1760 RemoteStyle::Predecessor
1761 }
1762}
1763
1764impl HasPeer<Agent> for Conductor {
1765 fn remote_style(&self, _peer: Agent) -> RemoteStyle {
1766 RemoteStyle::Successor
1767 }
1768}
1769
1770pub(crate) struct ProxySessionMessages {
1775 session_id: SessionId,
1776}
1777
1778impl ProxySessionMessages {
1779 pub fn new(session_id: SessionId) -> Self {
1781 Self { session_id }
1782 }
1783}
1784
1785impl<Counterpart: Role> HandleDispatchFrom<Counterpart> for ProxySessionMessages
1786where
1787 Counterpart: HasPeer<Agent> + HasPeer<Client>,
1788{
1789 async fn handle_dispatch_from(
1790 &mut self,
1791 message: Dispatch,
1792 connection: ConnectionTo<Counterpart>,
1793 ) -> Result<Handled<Dispatch>, crate::Error> {
1794 MatchDispatchFrom::new(message, &connection)
1795 .if_dispatch_from(Agent, async |message| {
1796 if let Some(session_id) = message.get_session_id()?
1798 && session_id == self.session_id
1799 {
1800 connection.send_proxied_message_to(Client, message)?;
1801 return Ok(Handled::Yes);
1802 }
1803
1804 Ok(Handled::No {
1806 message,
1807 retry: false,
1808 })
1809 })
1810 .await
1811 .done()
1812 }
1813
1814 fn describe_chain(&self) -> impl std::fmt::Debug {
1815 format!("ProxySessionMessages({})", self.session_id)
1816 }
1817}
1818
1819#[cfg(all(test, feature = "unstable_protocol_v2"))]
1820mod lifetime_tests {
1821 use super::*;
1822 use crate::{ConnectionDriver, UntypedRole};
1823 use futures::FutureExt as _;
1824
1825 fn frame() -> TransportFrame {
1826 TransportFrame::parse_json(r#"{"jsonrpc":"2.0","method":"test/queued","params":{}}"#)
1827 }
1828
1829 #[tokio::test]
1830 async fn passive_initialization_waits_for_a_frame_not_driver_readiness() {
1831 let (channel, remote) = Channel::duplex();
1832 let peer = RunningProtocolPeer::new::<UntypedRole>(channel);
1833 let mut next = Box::pin(peer.next_frame());
1834 assert!(next.as_mut().now_or_never().is_none());
1835
1836 remote.tx.unbounded_send(frame()).unwrap();
1837 let (_, peer) = next.await.unwrap().expect("passive peer remains connected");
1838 assert!(matches!(peer.driver, ProtocolPeerDriver::Passive));
1839 }
1840
1841 #[tokio::test]
1842 async fn owned_peer_preserves_finish_metadata_through_active_and_completed_states() {
1843 let (Channel { rx, tx }, remote) = Channel::duplex();
1844 let (done_tx, done_rx) = futures::channel::oneshot::channel();
1845 let (finish_tx, finish_rx) = futures::channel::oneshot::channel();
1846 let peer = RunningProtocolPeer {
1847 rx,
1848 tx,
1849 driver: ProtocolPeerDriver::Active(ConnectionDriver::with_finish(
1850 async move {
1851 done_rx.await.unwrap();
1852 Ok(())
1853 },
1854 move || {
1855 let _ = finish_tx.send(());
1856 },
1857 )),
1858 };
1859
1860 remote.tx.unbounded_send(frame()).unwrap();
1861 let (_, peer) = peer.next_frame().await.unwrap().unwrap();
1862 assert!(matches!(&peer.driver, ProtocolPeerDriver::Active(_)));
1863
1864 remote.tx.unbounded_send(frame()).unwrap();
1865 done_tx.send(()).unwrap();
1866 let (_, peer) = peer.next_frame().await.unwrap().unwrap();
1867 assert!(matches!(
1868 &peer.driver,
1869 ProtocolPeerDriver::Completed { finish: Some(_) }
1870 ));
1871 let mut driver = peer
1872 .driver
1873 .into_driver()
1874 .expect("completed owned work must not become passive");
1875 assert!(driver.request_finish());
1876 finish_rx.await.unwrap();
1877 driver.await.unwrap();
1878 }
1879
1880 #[tokio::test]
1881 async fn active_initialization_drains_accepted_frames_without_escaped_sender_eof() {
1882 let (Channel { rx, tx }, remote) = Channel::duplex();
1883 remote.tx.unbounded_send(frame()).unwrap();
1884 remote.tx.unbounded_send(frame()).unwrap();
1885 let mut polls = 0;
1886 let driver = ConnectionDriver::new(future::poll_fn(move |_| {
1887 polls += 1;
1888 assert_eq!(polls, 1, "the completed driver must never be re-polled");
1889 std::task::Poll::Ready(Ok(()))
1890 }));
1891 let peer = RunningProtocolPeer {
1892 rx,
1893 tx,
1894 driver: ProtocolPeerDriver::Active(driver),
1895 };
1896
1897 let (_, peer) = peer.next_frame().await.unwrap().unwrap();
1898 assert!(
1899 remote.tx.unbounded_send(frame()).is_err(),
1900 "active completion must reject new output from escaped handles"
1901 );
1902 let (_, peer) = peer.next_frame().await.unwrap().unwrap();
1903 assert!(peer.next_frame().await.unwrap().is_none());
1904 }
1905
1906 #[tokio::test]
1907 async fn ready_initialization_driver_error_beats_queued_frames() {
1908 let (Channel { rx, tx }, remote) = Channel::duplex();
1909 remote.tx.unbounded_send(frame()).unwrap();
1910 let error = crate::Error::internal_error().data("owned initialization failed");
1911 let peer = RunningProtocolPeer {
1912 rx,
1913 tx,
1914 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Err(
1915 error.clone()
1916 )))),
1917 };
1918
1919 match peer.next_frame().await {
1920 Err(actual) => assert_eq!(actual, error),
1921 Ok(_) => panic!("ready driver error must not be hidden by a queued frame"),
1922 }
1923 }
1924
1925 struct QueuedFinalClient;
1926
1927 impl ConnectTo<Agent> for QueuedFinalClient {
1928 async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), crate::Error> {
1929 let (mut channel, driver) = agent.into_channel_and_future();
1930 let foreground = async move {
1931 channel
1932 .tx
1933 .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request(
1934 "initialize".into(),
1935 serde_json::json!({ "protocolVersion": 1, "clientCapabilities": {} }),
1936 RequestId::Number(1),
1937 )?))
1938 .unwrap();
1939 assert!(channel.rx.next().await.is_some(), "initialize response");
1940 for index in 0..3 {
1943 channel
1944 .tx
1945 .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification(
1946 "test/final".into(),
1947 serde_json::json!({ "index": index, "payload": "x".repeat(1024) }),
1948 )?))
1949 .unwrap();
1950 }
1951 Ok(())
1952 };
1953 match driver {
1954 Some(driver) => crate::util::run_until(driver, foreground).await,
1955 None => foreground.await,
1956 }
1957 }
1958 }
1959
1960 #[tokio::test]
1961 async fn connector_completion_flushes_byte_streams_without_remote_read_eof() {
1962 use tokio::io::{AsyncBufReadExt as _, AsyncWriteExt as _, BufReader};
1963 use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _};
1964
1965 let (writer, remote_reader) = tokio::io::duplex(64);
1966 let (mut remote_input, reader) = tokio::io::duplex(64);
1967 let physical = crate::ByteStreams::new(writer.compat_write(), reader.compat());
1968 let mut physical = Some(crate::DynConnectTo::<Client>::new(physical));
1969 let connector = tokio::spawn(
1970 ClientProtocolConnector::new()
1971 .with_v1(|| QueuedFinalClient)
1972 .connect_to(move || physical.take().expect("one physical connection")),
1973 );
1974 let mut lines = BufReader::new(remote_reader).lines();
1975
1976 tokio::time::timeout(std::time::Duration::from_secs(5), async {
1977 let initialize = lines
1978 .next_line()
1979 .await
1980 .unwrap()
1981 .expect("initialize request");
1982 let value: serde_json::Value = serde_json::from_str(&initialize).unwrap();
1983 assert_eq!(value["method"], "initialize");
1984 remote_input
1985 .write_all(b"{\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{\"protocolVersion\":1}}\n")
1986 .await
1987 .unwrap();
1988 for index in 0..3 {
1989 let line = lines
1990 .next_line()
1991 .await
1992 .unwrap()
1993 .expect("accepted final output");
1994 let value: serde_json::Value = serde_json::from_str(&line).unwrap();
1995 assert_eq!(value["params"]["index"], index);
1996 assert_eq!(value["params"]["payload"].as_str().unwrap().len(), 1024);
1997 }
1998 connector.await.unwrap().unwrap();
1999 assert!(lines.next_line().await.unwrap().is_none());
2000 })
2001 .await
2002 .expect("physical flush must not wait for independent remote input EOF");
2003 drop(remote_input);
2004 }
2005
2006 #[derive(Default, Debug)]
2007 struct GatedLineSinkState {
2008 pending: Vec<String>,
2009 flushed: Vec<String>,
2010 closed: bool,
2011 dropped: bool,
2012 }
2013
2014 struct GatedLineSink {
2015 state: std::sync::Arc<std::sync::Mutex<GatedLineSinkState>>,
2016 release: futures::channel::oneshot::Receiver<()>,
2017 released: bool,
2018 }
2019
2020 impl futures::Sink<String> for GatedLineSink {
2021 type Error = std::io::Error;
2022
2023 fn poll_ready(
2024 self: std::pin::Pin<&mut Self>,
2025 _: &mut std::task::Context<'_>,
2026 ) -> std::task::Poll<Result<(), Self::Error>> {
2027 std::task::Poll::Ready(Ok(()))
2028 }
2029
2030 fn start_send(self: std::pin::Pin<&mut Self>, line: String) -> Result<(), Self::Error> {
2031 self.state.lock().unwrap().pending.push(line);
2032 Ok(())
2033 }
2034
2035 fn poll_flush(
2036 mut self: std::pin::Pin<&mut Self>,
2037 cx: &mut std::task::Context<'_>,
2038 ) -> std::task::Poll<Result<(), Self::Error>> {
2039 if !self.released {
2040 if std::pin::Pin::new(&mut self.release).poll(cx).is_pending() {
2041 return std::task::Poll::Pending;
2042 }
2043 self.released = true;
2044 }
2045 let mut state = self.state.lock().unwrap();
2046 let pending = std::mem::take(&mut state.pending);
2047 state.flushed.extend(pending);
2048 std::task::Poll::Ready(Ok(()))
2049 }
2050
2051 fn poll_close(
2052 mut self: std::pin::Pin<&mut Self>,
2053 cx: &mut std::task::Context<'_>,
2054 ) -> std::task::Poll<Result<(), Self::Error>> {
2055 match self.as_mut().poll_flush(cx) {
2056 std::task::Poll::Ready(Ok(())) => {
2057 self.state.lock().unwrap().closed = true;
2058 std::task::Poll::Ready(Ok(()))
2059 }
2060 result => result,
2061 }
2062 }
2063 }
2064
2065 impl Drop for GatedLineSink {
2066 fn drop(&mut self) {
2067 self.state.lock().unwrap().dropped = true;
2068 }
2069 }
2070
2071 #[tokio::test]
2072 async fn foreground_completion_flushes_lines_with_already_normalized_remote_input() {
2073 let state = std::sync::Arc::new(std::sync::Mutex::new(GatedLineSinkState::default()));
2074 let (release_tx, release_rx) = futures::channel::oneshot::channel();
2075 let sink = GatedLineSink {
2076 state: state.clone(),
2077 release: release_rx,
2078 released: false,
2079 };
2080 let (remote_input, incoming) = futures::channel::mpsc::unbounded();
2081 remote_input
2082 .unbounded_send(Ok(
2083 r#"{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":1}}"#.to_string(),
2084 ))
2085 .unwrap();
2086 remote_input
2087 .unbounded_send(Ok(
2088 r#"{"jsonrpc":"2.0","method":"test/queued","params":{}}"#.to_string(),
2089 ))
2090 .unwrap();
2091 let physical = RunningProtocolPeer::new::<Client>(crate::Lines::new(sink, incoming));
2092 let (initialize, physical) = physical.next_frame().await.unwrap().unwrap();
2094 assert_eq!(
2095 futures::Stream::size_hint(&physical.rx).0,
2096 1,
2097 "trailing input must already be in the original normalized queue"
2098 );
2099
2100 let (Channel { rx, tx }, mut local) = Channel::duplex();
2101 tx.unbounded_send(initialize).unwrap();
2102 let foreground = RunningProtocolPeer {
2103 rx,
2104 tx,
2105 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(async move {
2106 assert!(local.rx.next().await.is_some(), "initialize response");
2107 for index in 0..3 {
2108 local
2109 .tx
2110 .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification(
2111 "test/final".into(),
2112 serde_json::json!({ "index": index, "payload": "x".repeat(1024) }),
2113 )?))
2114 .unwrap();
2115 }
2116 drop(local);
2118 Ok(())
2119 })),
2120 };
2121 let mut bridge = Box::pin(pipe_protocol_peers_until_done(foreground, physical));
2122 let early_result = bridge.as_mut().now_or_never();
2123 assert!(
2124 !state.lock().unwrap().pending.is_empty(),
2125 "the real physical writer must accept output before the gated flush"
2126 );
2127 let _released = release_tx.send(());
2129 let result = match early_result {
2130 Some(result) => result,
2131 None => tokio::time::timeout(std::time::Duration::from_secs(1), bridge)
2132 .await
2133 .expect("physical drain must not require remote input EOF"),
2134 };
2135 let state = state.lock().unwrap();
2136 assert!(
2137 result.is_ok() && state.flushed.len() == 3 && state.closed,
2138 "accepted output must drain cleanly despite queued remote input: result={result:?}, pending={}, flushed={}, closed={}, dropped={}",
2139 state.pending.len(),
2140 state.flushed.len(),
2141 state.closed,
2142 state.dropped,
2143 );
2144 for (index, line) in state.flushed.iter().enumerate() {
2145 let value: serde_json::Value = serde_json::from_str(line).unwrap();
2146 assert_eq!(value["method"], "test/final");
2147 assert_eq!(value["params"]["index"], index);
2148 assert_eq!(value["params"]["payload"].as_str().unwrap().len(), 1024);
2149 }
2150 drop(remote_input);
2151 }
2152
2153 #[tokio::test]
2154 async fn foreground_completion_keeps_read_errors_during_lines_drain() {
2155 let state = std::sync::Arc::new(std::sync::Mutex::new(GatedLineSinkState::default()));
2156 let (_release_tx, release_rx) = futures::channel::oneshot::channel();
2157 let sink = GatedLineSink {
2158 state: state.clone(),
2159 release: release_rx,
2160 released: false,
2161 };
2162 let (remote_input, incoming) = futures::channel::mpsc::unbounded();
2163 let physical = RunningProtocolPeer::new::<Client>(crate::Lines::new(sink, incoming));
2164 let (Channel { rx, tx }, local) = Channel::duplex();
2165 local.tx.unbounded_send(frame()).unwrap();
2166 drop(local);
2167 let foreground = RunningProtocolPeer {
2168 rx,
2169 tx,
2170 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))),
2171 };
2172 let mut bridge = Box::pin(pipe_protocol_peers_until_done(foreground, physical));
2173 assert!(bridge.as_mut().now_or_never().is_none());
2174 assert_eq!(state.lock().unwrap().pending.len(), 1);
2175
2176 remote_input
2179 .unbounded_send(Ok(r#"{"jsonrpc":"2.0","method":"late"}"#.to_string()))
2180 .unwrap();
2181 remote_input
2182 .unbounded_send(Err(std::io::Error::other("read failed after foreground")))
2183 .unwrap();
2184 let error = tokio::time::timeout(std::time::Duration::from_secs(1), bridge)
2185 .await
2186 .expect("read failure must not wait for the sink gate")
2187 .expect_err("real read error must win over foreground success");
2188 assert!(
2189 error
2190 .data
2191 .unwrap()
2192 .to_string()
2193 .contains("read failed after foreground"),
2194 "the read error must not become a receiver-gone forwarding error"
2195 );
2196 assert_eq!(state.lock().unwrap().flushed, Vec::<String>::new());
2197 }
2198
2199 #[tokio::test]
2200 async fn foreground_completion_cancels_opposed_opaque_work_in_both_directions() {
2201 for foreground_on_left in [true, false] {
2202 let (Channel { rx, tx }, foreground_remote) = Channel::duplex();
2203 foreground_remote.tx.unbounded_send(frame()).unwrap();
2204 let foreground = RunningProtocolPeer {
2205 rx,
2206 tx,
2207 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))),
2208 };
2209 let (Channel { rx, tx }, mut opposed_remote) = Channel::duplex();
2210 let (work_tx, work_rx) = futures::channel::oneshot::channel::<()>();
2211 let opposed = RunningProtocolPeer {
2212 rx,
2213 tx,
2214 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(async move {
2215 work_rx.await.map_err(crate::util::internal_error)?;
2216 Ok(())
2217 })),
2218 };
2219 let mut bridge = Box::pin(if foreground_on_left {
2220 pipe_protocol_peers_until_done(foreground, opposed)
2221 } else {
2222 pipe_protocol_peers_until_done(opposed, foreground)
2223 });
2224
2225 assert_eq!(
2226 bridge.as_mut().now_or_never(),
2227 Some(Ok(())),
2228 "finite foreground must not join an opaque pending peer"
2229 );
2230 assert!(work_tx.is_canceled(), "opaque work must be dropped");
2231 assert!(opposed_remote.rx.next().await.is_some());
2232 assert!(opposed_remote.rx.next().await.is_none());
2233 }
2234 }
2235
2236 #[tokio::test]
2237 async fn foreground_completion_does_not_hide_opposed_ready_driver_error() {
2238 let (Channel { rx, tx }, _foreground_remote) = Channel::duplex();
2239 let foreground = RunningProtocolPeer {
2240 rx,
2241 tx,
2242 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))),
2243 };
2244 let (Channel { rx, tx }, _opposed_remote) = Channel::duplex();
2245 let error = crate::Error::internal_error().data("opposed driver failed");
2246 let opposed = RunningProtocolPeer {
2247 rx,
2248 tx,
2249 driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Err(
2250 error.clone()
2251 )))),
2252 };
2253
2254 assert_eq!(
2255 pipe_protocol_peers_until_done(foreground, opposed).await,
2256 Err(error)
2257 );
2258 }
2259
2260 #[tokio::test]
2261 async fn passive_protocol_bridge_preserves_the_reverse_half_after_eof() {
2262 let (left, mut remote_left) = Channel::duplex();
2263 let (right, mut remote_right) = Channel::duplex();
2264 let mut bridge = Box::pin(pipe_protocol_peers_until_done(
2265 RunningProtocolPeer::new::<UntypedRole>(left),
2266 RunningProtocolPeer::new::<UntypedRole>(right),
2267 ));
2268 remote_left.tx.close_channel();
2269 assert!(bridge.as_mut().now_or_never().is_none());
2270 assert!(remote_right.rx.next().await.is_none());
2271
2272 remote_right.tx.unbounded_send(frame()).unwrap();
2273 remote_right.tx.close_channel();
2274 bridge.await.unwrap();
2275 assert!(remote_left.rx.next().await.is_some());
2276 assert!(remote_left.rx.next().await.is_none());
2277 }
2278}