Skip to main content

agent_client_protocol/role/
acp.rs

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/// The client role - typically an IDE or CLI that controls an agent.
43///
44/// Clients send prompts and receive responses from agents.
45#[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    /// Create a connection builder for a client.
77    pub fn builder(self) -> Builder<Client, NullHandler, NullRun> {
78        <Self as Role>::builder(self)
79    }
80
81    /// Create a client builder that requires an ACP protocol v2 agent.
82    ///
83    /// If the agent negotiates v1 during initialization, the initialize
84    /// request resolves with an error so callers can choose an explicit v1
85    /// fallback path.
86    ///
87    /// Requires the `unstable_protocol_v2` crate feature.
88    #[cfg(feature = "unstable_protocol_v2")]
89    pub fn v2(self) -> V2Builder<Client, NullHandler, NullRun> {
90        self.builder().v2_client()
91    }
92
93    /// Create a connector that chooses between configured protocol implementations.
94    ///
95    /// Add implementation factories with [`ClientProtocolConnector::with_v1`]
96    /// and [`ClientProtocolConnector::with_v2`]. The resulting connector starts
97    /// the highest configured protocol implementation. If a v2 implementation
98    /// successfully negotiates v1 and a v1 implementation is configured, the
99    /// connector reuses the connection only when the v1 implementation's
100    /// complete `initialize` parameters match what the agent already saw;
101    /// otherwise it opens a fresh agent connection and restarts with v1.
102    ///
103    /// Requires the `unstable_protocol_v2` crate feature while protocol v2
104    /// stabilizes.
105    #[cfg(feature = "unstable_protocol_v2")]
106    #[must_use]
107    pub fn protocol_connector(self) -> ClientProtocolConnector {
108        ClientProtocolConnector::new()
109    }
110
111    /// Connect to `agent` and run `main_fn` with the [`ConnectionTo`].
112    /// Returns the result of `main_fn` (or an error if something goes wrong).
113    ///
114    /// Equivalent to `self.builder().connect_with(agent, main_fn)`.
115    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/// Client connector that opens an agent connection with a configured protocol implementation.
125///
126/// Use [`Client::protocol_connector`] to start the builder, then add each
127/// supported protocol version independently. Implementations and the agent
128/// connection are provided as factories because fallback from v2 to v1 may
129/// require a fresh connection initialized by the v1 implementation.
130#[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    /// Create an empty client protocol connector.
140    #[must_use]
141    pub fn new() -> Self {
142        Self::default()
143    }
144
145    /// Return this connector with an ACP v1 implementation factory configured.
146    #[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    /// Return this connector with an ACP v2 implementation factory configured.
156    #[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    /// Connect to an agent produced by `agent` using the highest configured
166    /// compatible protocol implementation.
167    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                // This normalization is only a probe for the connection-reuse
202                // optimization. A request that cannot be represented in v1
203                // can still be valid v2 traffic and must reach the agent.
204                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                        // The v2 implementation will never receive its initialize response once
226                        // the matching v1 implementation takes over this connection. Drop its
227                        // future and channels before running the v1 session so any resources it
228                        // owns are released promptly.
229                        drop(client);
230                        fallback_client.send(fallback_response)?;
231                        return pipe_protocol_peers_until_done(fallback_client, agent_connection)
232                            .await;
233                    }
234
235                    // Neither probe can continue on the replacement connection. Release both
236                    // client implementations and the original agent connection before starting
237                    // the real v1 session so they cannot retain resources for its lifetime.
238                    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/// The agent role - typically an LLM that responds to prompts.
290///
291/// Agents receive prompts from clients and respond with answers,
292/// potentially invoking tools along the way.
293#[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                // Stable v1 session helpers install a dynamic handler after
319                // `session/new`. Retry session messages to close the race
320                // between the response and that registration.
321                //
322                // V2 uses typed handlers installed before the connection
323                // starts. Retrying an unhandled v2 message would retain it
324                // forever because no per-session dynamic handler is expected.
325                #[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    /// Create a connection builder for an agent.
340    pub fn builder(self) -> Builder<Agent, NullHandler, NullRun> {
341        <Self as Role>::builder(self)
342    }
343
344    /// Create an agent builder that uses the ACP protocol v2 API.
345    ///
346    /// This builder requires clients to negotiate protocol v2 during
347    /// initialization. Use a v1 builder for v1 clients.
348    ///
349    /// Requires the `unstable_protocol_v2` crate feature.
350    #[cfg(feature = "unstable_protocol_v2")]
351    pub fn v2(self) -> V2Builder<Agent, NullHandler, NullRun> {
352        self.builder().v2_agent()
353    }
354
355    /// Create a router that chooses between configured protocol implementations.
356    ///
357    /// Add implementations with [`AgentProtocolRouter::with_v1`] and
358    /// [`AgentProtocolRouter::with_v2`].
359    /// The resulting router reads the initial
360    /// `initialize` request, selects the highest configured implementation
361    /// compatible with the client's requested protocol version, then forwards
362    /// the connection to that implementation. It does not convert traffic
363    /// between protocol versions after routing.
364    ///
365    /// Requires the `unstable_protocol_v2` crate feature while protocol v2
366    /// stabilizes.
367    #[cfg(feature = "unstable_protocol_v2")]
368    #[must_use]
369    pub fn protocol_router(self) -> AgentProtocolRouter {
370        AgentProtocolRouter::new()
371    }
372}
373
374/// Agent component that routes each connection to a configured protocol implementation.
375///
376/// Use [`Agent::protocol_router`] to start the builder, then add each supported
377/// protocol version independently. The selected implementation owns the
378/// connection after the initial `initialize` negotiation.
379#[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    /// Create an empty agent protocol router.
389    #[must_use]
390    pub fn new() -> Self {
391        Self::default()
392    }
393
394    /// Return this router with an ACP v1 implementation configured.
395    #[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    /// Return this router with an ACP v2 implementation configured.
402    #[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    // Validate exact-version initialization without replacing its raw
615    // parameters. Reserializing through the SDK's pinned schema would discard
616    // fields added by newer compatible peers.
617    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    // Canonicalize through v2 first so tolerant field semantics are applied.
650    // A lossless result is only required by the connection-reuse probe. Normal
651    // v1 routing may discard fields that have no meaning in the selected
652    // protocol version.
653    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(&params, 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(&params, false)
758            .expect("v1 routing may ignore parameters unavailable in v1");
759        normalize_v2_initialize_params_for_v1(&params, 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(&params, 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(&params, false)
820            .expect("v1 routing may discard terminal marker metadata");
821        normalize_v2_initialize_params_for_v1(&params, 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        // The rejection has already been handed to the raw channel. There is
966        // no owned transport work or physical drain to await.
967        return Ok(());
968    };
969    if !driver.request_finish() {
970        // An opaque driver has no finite physical-finish contract. Preserve a
971        // ready error before cancelling it rather than wait for remote EOF.
972        return crate::util::run_until(driver, future::ready(Ok(()))).await;
973    }
974
975    let drain_incoming = async move {
976        // Later input has no protocol meaning once initialization is rejected.
977        // Keep draining it only so the transport can flush the queued rejection;
978        // treating a malformed trailing frame as fatal would cancel that flush.
979        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            // Conversion happens only when handing the peer to its final
1015            // bridge, never while reading its remaining queued frames. This
1016            // records actual owned completion, not a passive ready sentinel.
1017            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        // Poll the owned driver first: a ready error must not be hidden by
1050        // an equally ready frame or clean channel EOF.
1051        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                // No more output may be accepted from an owned endpoint once
1069                // its driver completes, even if a sender escaped the component.
1070                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// Every protocol router uses the same ownership rule. Passive halves keep
1140// independent lifetimes; owned completion drains output and then either joins
1141// an opposed cooperative driver or cancels opaque work after polling errors.
1142#[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                // Without a cooperative finish hook, opposed work may remain
1192                // open indefinitely. Poll ready errors before cancelling it.
1193                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/// The proxy role - an intermediary that can intercept and modify messages.
1419///
1420/// Proxies sit between a client and an agent (or another proxy), and can:
1421/// - Add tools via MCP servers
1422/// - Filter or transform messages
1423/// - Inject additional context
1424///
1425/// Proxies connect to a [`Conductor`] which orchestrates the proxy chain.
1426#[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    /// Create a stable protocol v1 connection builder for a proxy.
1454    ///
1455    /// Use `Proxy::v2` for a protocol-v2-only proxy with typed callbacks and
1456    /// wire validation. Protocol-routing infrastructure that deliberately
1457    /// selects a version itself can disable the guard with
1458    /// `Builder::without_acp_version_guard`.
1459    pub fn builder(self) -> Builder<Proxy, NullHandler, NullRun> {
1460        Builder::new(self)
1461    }
1462
1463    /// Create a proxy builder that uses the ACP protocol v2 API.
1464    ///
1465    /// This builder requires `_proxy/initialize` to select protocol v2.
1466    /// Fluent callbacks receive [`crate::V2ConnectionTo<Conductor>`], while
1467    /// low-level custom handlers and runners retain the protocol-neutral
1468    /// [`ConnectionTo`] interface.
1469    ///
1470    /// Requires the `unstable_protocol_v2` crate feature.
1471    #[cfg(feature = "unstable_protocol_v2")]
1472    pub fn v2(self) -> V2Builder<Proxy, NullHandler, NullRun> {
1473        self.builder().v2_proxy()
1474    }
1475
1476    /// Create a router that chooses between configured proxy implementations.
1477    ///
1478    /// Add implementations with [`ProxyProtocolRouter::with_v1`] and
1479    /// [`ProxyProtocolRouter::with_v2`]. The router reads the initial
1480    /// `_proxy/initialize` request, selects the implementation for that exact
1481    /// protocol version, and hands over the complete initial transport frame.
1482    /// It does not downgrade proxy traffic or convert later messages.
1483    ///
1484    /// Requires the `unstable_protocol_v2` crate feature while protocol v2
1485    /// stabilizes.
1486    #[cfg(feature = "unstable_protocol_v2")]
1487    #[must_use]
1488    pub fn protocol_router(self) -> ProxyProtocolRouter {
1489        ProxyProtocolRouter::new()
1490    }
1491}
1492
1493/// Proxy component that routes each connection to a configured protocol implementation.
1494///
1495/// Use [`Proxy::protocol_router`] to start the builder, then add stable-v1 and
1496/// draft-v2 proxy implementations independently. Unlike
1497/// [`AgentProtocolRouter`], this router requires an exact version match: the
1498/// conductor has already selected and canonicalized the wire protocol before
1499/// sending `_proxy/initialize` to a proxy.
1500#[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    /// Create an empty proxy protocol router.
1510    #[must_use]
1511    pub fn new() -> Self {
1512        Self::default()
1513    }
1514
1515    /// Return this router with a stable ACP v1 proxy implementation.
1516    #[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    /// Return this router with a draft ACP v2 proxy implementation.
1523    #[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/// The conductor role - orchestrates proxy chains.
1625///
1626/// Conductors manage connections between clients, proxies, and agents,
1627/// routing messages through the appropriate proxy chain.
1628#[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                    // The dynamic-handler hook below means we cannot use
1713                    // `forward_response_to`, so wire up cancellation forwarding
1714                    // explicitly to keep `session/new` cancellable like every
1715                    // other proxied request.
1716                    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    /// Create a connection builder for a conductor.
1753    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
1770/// Dynamic handler that proxies session messages from Agent to Client.
1771///
1772/// This is used internally to handle session message routing after a
1773/// `session.new` request has been forwarded.
1774pub(crate) struct ProxySessionMessages {
1775    session_id: SessionId,
1776}
1777
1778impl ProxySessionMessages {
1779    /// Create a new proxy handler for the given session.
1780    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 this is for our session-id, proxy it to the client.
1797                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                // Otherwise, leave it alone.
1805                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                // Exceed the physical writer capacity so only concurrent polling
1941                // of the sink can make this bridge finish.
1942                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        // The real Lines driver reads both ready lines before this returns.
2093        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                // This owns the foreground input receiver: completion closes it.
2117                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        // Never use time to establish the race: only the sink gate controls drain.
2128        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        // Newly read successful input is irrelevant to the completed foreground,
2177        // but a genuine read failure must still cancel the blocked sink drain.
2178        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}