use std::{fmt::Debug, future::Future, hash::Hash};
#[cfg(feature = "unstable_protocol_v2")]
use futures::{StreamExt as _, future};
#[cfg(feature = "unstable_protocol_v2")]
use serde::{Serialize, de::DeserializeOwned};
#[cfg(feature = "unstable_protocol_v2")]
use crate::DynConnectTo;
use crate::jsonrpc::{Builder, handlers::NullHandler, run::NullRun};
#[cfg(feature = "unstable_protocol_v2")]
use crate::jsonrpc::{
TransportBatch, TransportBatchEntry, TransportFrame, V2Builder, is_response_only_shape,
raw_is_response_only_shape,
};
use crate::role::{HasPeer, RemoteStyle};
#[cfg(not(feature = "unstable_protocol_v2"))]
use crate::schema::InitializeProxyRequest;
use crate::schema::METHOD_INITIALIZE_PROXY;
#[cfg(feature = "unstable_protocol_v2")]
use crate::schema::v1::RequestId;
use crate::schema::v1::{InitializeRequest, SessionId};
#[cfg(not(feature = "unstable_protocol_v2"))]
use crate::schema::v1::{NewSessionRequest, NewSessionResponse};
#[cfg(feature = "unstable_protocol_v2")]
use crate::schema::{ProtocolVersion, v2};
use crate::util::MatchDispatchFrom;
#[cfg(feature = "unstable_protocol_v2")]
use crate::{
Channel, RawJsonRpcError, RawJsonRpcMessage, RawJsonRpcParams,
RawJsonRpcResponse as RpcResponse,
};
use crate::{ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, Role, RoleId};
#[cfg(feature = "unstable_protocol_v2")]
#[derive(serde::Deserialize)]
struct NewSessionResponseEnvelope {
#[serde(rename = "sessionId")]
session_id: SessionId,
}
#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Client;
impl Role for Client {
type Counterpart = Agent;
fn builder(self) -> Builder<Self> {
Builder::new(self).v1_client()
}
fn default_handle_dispatch_from(
&self,
message: Dispatch,
_connection: ConnectionTo<Client>,
) -> impl Future<Output = Result<Handled<Dispatch>, crate::Error>> + Send {
std::future::ready(Ok(Handled::No {
message,
retry: false,
}))
}
fn role_id(&self) -> RoleId {
RoleId::from_singleton(self)
}
fn counterpart(&self) -> Self::Counterpart {
Agent
}
}
impl Client {
pub fn builder(self) -> Builder<Client, NullHandler, NullRun> {
<Self as Role>::builder(self)
}
#[cfg(feature = "unstable_protocol_v2")]
pub fn v2(self) -> V2Builder<Client, NullHandler, NullRun> {
self.builder().v2_client()
}
#[cfg(feature = "unstable_protocol_v2")]
#[must_use]
pub fn protocol_connector(self) -> ClientProtocolConnector {
ClientProtocolConnector::new()
}
pub async fn connect_with<R>(
self,
agent: impl ConnectTo<Client>,
main_fn: impl AsyncFnOnce(ConnectionTo<Agent>) -> Result<R, crate::Error>,
) -> Result<R, crate::Error> {
self.builder().connect_with(agent, main_fn).await
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Default)]
pub struct ClientProtocolConnector {
v1: Option<DynConnectToFactory<Agent>>,
v2: Option<DynConnectToFactory<Agent>>,
}
#[cfg(feature = "unstable_protocol_v2")]
impl ClientProtocolConnector {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_v1<C>(mut self, client: impl FnMut() -> C + Send + 'static) -> Self
where
C: ConnectTo<Agent>,
{
self.v1 = Some(DynConnectToFactory::new(client));
self
}
#[must_use]
pub fn with_v2<C>(mut self, client: impl FnMut() -> C + Send + 'static) -> Self
where
C: ConnectTo<Agent>,
{
self.v2 = Some(DynConnectToFactory::new(client));
self
}
pub async fn connect_to<C>(
mut self,
mut agent: impl FnMut() -> C + Send + 'static,
) -> Result<(), crate::Error>
where
C: ConnectTo<Client>,
{
let supported = SupportedClientProtocols {
v1: self.v1.is_some(),
v2: self.v2.is_some(),
};
let Some(selected) = supported.highest_configured() else {
return Err(crate::Error::invalid_request()
.data("client protocol connector has no configured ACP protocol implementations"));
};
match selected {
ClientProtocol::V1 => {
let client = self
.v1
.as_mut()
.expect("selected protocol is configured")
.create();
connect_client_protocol(ClientProtocol::V1, client, agent()).await
}
ClientProtocol::V2 => {
let client = self
.v2
.as_mut()
.expect("selected protocol is configured")
.create();
let agent_connection = RunningProtocolPeer::new(agent());
let (client, initialize) =
start_client_protocol(ClientProtocol::V2, client).await?;
let v2_initialize_as_v1 = normalize_v2_initialize_params_for_reuse(&initialize);
let (client, agent_connection, initialize_response) =
send_initialize_and_receive(client, agent_connection, initialize).await?;
if initialize_response_negotiated_v1(&initialize_response)
&& let Some(v1) = self.v1.as_mut()
{
let fallback_client = v1.create();
let (fallback_client, fallback_initialize) =
start_client_protocol(ClientProtocol::V1, fallback_client).await?;
let v1_initialize =
validated_initialize_params::<InitializeRequest>(&fallback_initialize)?;
if v2_initialize_as_v1
.as_ref()
.is_ok_and(|v2_initialize| v2_initialize == &v1_initialize)
{
let fallback_response = initialize_response.with_id(
initialize_request_id(&fallback_initialize)
.expect("validated initialize request has an id"),
);
drop(client);
fallback_client.send(fallback_response)?;
return pipe_protocol_peers_until_done(fallback_client, agent_connection)
.await;
}
drop((
client,
fallback_client,
agent_connection,
initialize_response,
));
return connect_client_protocol(ClientProtocol::V1, v1.create(), agent()).await;
}
client.send(initialize_response.into_message())?;
pipe_protocol_peers_until_done(client, agent_connection).await
}
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
struct DynConnectToFactory<R: Role> {
inner: Box<dyn FnMut() -> DynConnectTo<R> + Send>,
}
#[cfg(feature = "unstable_protocol_v2")]
impl<R: Role> DynConnectToFactory<R> {
fn new<C>(mut factory: impl FnMut() -> C + Send + 'static) -> Self
where
C: ConnectTo<R>,
{
Self {
inner: Box::new(move || DynConnectTo::new(factory())),
}
}
fn create(&mut self) -> DynConnectTo<R> {
(self.inner)()
}
}
#[cfg(feature = "unstable_protocol_v2")]
impl<R: Role> Debug for DynConnectToFactory<R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DynConnectToFactory")
.finish_non_exhaustive()
}
}
impl HasPeer<Client> for Client {
fn remote_style(&self, _peer: Client) -> RemoteStyle {
RemoteStyle::Counterpart
}
}
#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Agent;
impl Role for Agent {
type Counterpart = Client;
fn builder(self) -> Builder<Self> {
Builder::new(self).v1_agent()
}
fn role_id(&self) -> RoleId {
RoleId::from_singleton(self)
}
fn counterpart(&self) -> Self::Counterpart {
Client
}
async fn default_handle_dispatch_from(
&self,
message: Dispatch,
connection: ConnectionTo<Agent>,
) -> Result<Handled<Dispatch>, crate::Error> {
MatchDispatchFrom::new(message, &connection)
.if_dispatch_from(Agent, async |message: Dispatch| {
#[cfg(feature = "unstable_protocol_v2")]
let retry = message.has_session_id()
&& connection.acp_protocol_version()
!= Some(crate::schema::ProtocolVersion::V2);
#[cfg(not(feature = "unstable_protocol_v2"))]
let retry = message.has_session_id();
Ok(Handled::No { message, retry })
})
.await
.done()
}
}
impl Agent {
pub fn builder(self) -> Builder<Agent, NullHandler, NullRun> {
<Self as Role>::builder(self)
}
#[cfg(feature = "unstable_protocol_v2")]
pub fn v2(self) -> V2Builder<Agent, NullHandler, NullRun> {
self.builder().v2_agent()
}
#[cfg(feature = "unstable_protocol_v2")]
#[must_use]
pub fn protocol_router(self) -> AgentProtocolRouter {
AgentProtocolRouter::new()
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Default)]
pub struct AgentProtocolRouter {
v1: Option<DynConnectTo<Client>>,
v2: Option<DynConnectTo<Client>>,
}
#[cfg(feature = "unstable_protocol_v2")]
impl AgentProtocolRouter {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_v1(mut self, agent: impl ConnectTo<Client>) -> Self {
self.v1 = Some(DynConnectTo::new(agent));
self
}
#[must_use]
pub fn with_v2(mut self, agent: impl ConnectTo<Client>) -> Self {
self.v2 = Some(DynConnectTo::new(agent));
self
}
}
#[cfg(feature = "unstable_protocol_v2")]
impl ConnectTo<Client> for AgentProtocolRouter {
async fn connect_to(self, client: impl ConnectTo<Agent>) -> Result<(), crate::Error> {
let supported = SupportedProtocols {
v1: self.v1.is_some(),
v2: self.v2.is_some(),
};
let mut client = RunningProtocolPeer::new(client);
let (first_frame, client, selected) = loop {
let Some((mut frame, next_client)) = client.next_frame().await? else {
return Ok(());
};
let message = match initialize_message_mut(&mut frame) {
Ok(Some(message)) => message,
Ok(None) => {
client = next_client;
continue;
}
Err(error) => return reject_initialize(next_client, &frame, error).await,
};
let selected = match select_agent_protocol(message, supported) {
Ok(selected) => selected,
Err(error) => return reject_initialize(next_client, &frame, error).await,
};
break (frame, next_client, selected);
};
let Some(agent) = selected.take_agent(self) else {
let error = selected.unsupported_error(supported);
return reject_initialize(client, &first_frame, error).await;
};
let agent = RunningProtocolPeer::new(agent);
agent.send_frame(first_frame)?;
pipe_protocol_peers_until_done(client, agent).await
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SelectedProtocol {
V1,
V2,
}
#[cfg(feature = "unstable_protocol_v2")]
impl SelectedProtocol {
fn take_agent(self, agent: AgentProtocolRouter) -> Option<DynConnectTo<Client>> {
match self {
Self::V1 => agent.v1,
Self::V2 => agent.v2,
}
}
fn version(self) -> ProtocolVersion {
match self {
Self::V1 => ProtocolVersion::V1,
Self::V2 => ProtocolVersion::V2,
}
}
fn name(self) -> &'static str {
match self {
Self::V1 => "1",
Self::V2 => "2",
}
}
fn unsupported_error(self, supported: SupportedProtocols) -> crate::Error {
crate::Error::invalid_request().data(format!(
"ACP protocol version {} is not configured; this endpoint supports {}",
self.name(),
supported.description()
))
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct SupportedProtocols {
v1: bool,
v2: bool,
}
#[cfg(feature = "unstable_protocol_v2")]
impl SupportedProtocols {
fn highest_compatible(self, requested: ProtocolVersion) -> Option<SelectedProtocol> {
if self.v2 && requested >= ProtocolVersion::V2 {
return Some(SelectedProtocol::V2);
}
if self.v1 && requested >= ProtocolVersion::V1 {
return Some(SelectedProtocol::V1);
}
None
}
fn exact(self, requested: ProtocolVersion) -> Option<SelectedProtocol> {
if self.v1 && requested == ProtocolVersion::V1 {
Some(SelectedProtocol::V1)
} else if self.v2 && requested == ProtocolVersion::V2 {
Some(SelectedProtocol::V2)
} else {
None
}
}
fn description(self) -> String {
match (self.v1, self.v2) {
(true, true) => "ACP protocol versions 1 and 2".into(),
(true, false) => "ACP protocol version 1".into(),
(false, true) => "ACP protocol version 2".into(),
(false, false) => "no ACP protocol versions".into(),
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
fn select_agent_protocol(
message: &mut RawJsonRpcMessage,
supported: SupportedProtocols,
) -> Result<SelectedProtocol, crate::Error> {
let RawJsonRpcMessage::Request(request) = message else {
return Err(
crate::Error::invalid_request().data("first ACP message must be an initialize request")
);
};
if request.method.as_ref() != "initialize" {
return Err(crate::Error::invalid_request().data("first ACP request must be initialize"));
}
let Some(RawJsonRpcParams::Object(params)) = &mut request.params else {
return Err(invalid_initialize_protocol_version());
};
let Some(protocol_version) = params.get("protocolVersion") else {
return Err(invalid_initialize_protocol_version());
};
let requested = serde_json::from_value::<ProtocolVersion>(protocol_version.clone())
.map_err(|_| invalid_initialize_protocol_version())?;
let selected = highest_compatible_agent_protocol(requested, supported)?;
rewrite_initialize_params(params, requested, selected)?;
Ok(selected)
}
#[cfg(feature = "unstable_protocol_v2")]
fn initialize_request_params(
message: &RawJsonRpcMessage,
) -> Result<&serde_json::Map<String, serde_json::Value>, crate::Error> {
let RawJsonRpcMessage::Request(request) = message else {
return Err(
crate::Error::invalid_request().data("first ACP message must be an initialize request")
);
};
if request.method.as_ref() != "initialize" {
return Err(crate::Error::invalid_request().data("first ACP request must be initialize"));
}
let Some(RawJsonRpcParams::Object(params)) = &request.params else {
return Err(invalid_initialize_protocol_version());
};
if !params.contains_key("protocolVersion") {
return Err(invalid_initialize_protocol_version());
}
Ok(params)
}
#[cfg(feature = "unstable_protocol_v2")]
fn validated_initialize_params<T: DeserializeOwned>(
message: &RawJsonRpcMessage,
) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
let params = initialize_request_params(message)?;
parse_initialize_params::<T>(params)?;
Ok(params.clone())
}
#[cfg(feature = "unstable_protocol_v2")]
fn normalize_v2_initialize_params_for_reuse(
message: &RawJsonRpcMessage,
) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
let params = initialize_request_params(message)?;
let requested = params
.get("protocolVersion")
.cloned()
.ok_or_else(invalid_initialize_protocol_version)
.and_then(|version| {
serde_json::from_value::<ProtocolVersion>(version)
.map_err(|_| invalid_initialize_protocol_version())
})?;
if requested == ProtocolVersion::V1 {
parse_initialize_params::<InitializeRequest>(params)?;
return Ok(params.clone());
}
normalize_v2_initialize_params_for_v1(params, true)
}
#[cfg(feature = "unstable_protocol_v2")]
fn rewrite_initialize_params(
params: &mut serde_json::Map<String, serde_json::Value>,
requested: ProtocolVersion,
selected: SelectedProtocol,
) -> Result<(), crate::Error> {
if requested == selected.version() {
match selected {
SelectedProtocol::V1 => {
parse_initialize_params::<InitializeRequest>(params)?;
}
SelectedProtocol::V2 => {
parse_initialize_params::<v2::InitializeRequest>(params)?;
}
}
return Ok(());
}
match selected {
SelectedProtocol::V1 => {
debug_assert!(requested >= ProtocolVersion::V2);
*params = normalize_v2_initialize_params_for_v1(params, false)?;
Ok(())
}
SelectedProtocol::V2 => {
let mut initialize = parse_initialize_params::<v2::InitializeRequest>(params)?;
initialize.protocol_version = ProtocolVersion::V2;
*params = serialize_initialize_params(initialize)?;
Ok(())
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
fn normalize_v2_initialize_params_for_v1(
params: &serde_json::Map<String, serde_json::Value>,
require_lossless: bool,
) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
let initialize = parse_initialize_params::<v2::InitializeRequest>(params)?;
let mut target = serialize_initialize_params(initialize)?;
if require_lossless && target != *params {
return Err(invalid_initialize_params(
"v2 initialize parameters are not losslessly representable in v1",
));
}
target.insert(
"protocolVersion".into(),
serde_json::to_value(ProtocolVersion::V1).map_err(crate::Error::into_internal_error)?,
);
let info = target
.remove("info")
.ok_or_else(|| invalid_initialize_params("v2 InitializeRequest.info is required"))?;
target.insert("clientInfo".into(), info);
let capabilities = target
.remove("capabilities")
.and_then(|capabilities| capabilities.as_object().cloned())
.ok_or_else(|| {
crate::util::internal_error("v2 initialize capabilities did not serialize as an object")
})?;
let mut capabilities = capabilities;
if let Some(auth) = capabilities
.get_mut("auth")
.and_then(serde_json::Value::as_object_mut)
{
let terminal = auth.remove("terminal");
if require_lossless
&& terminal
.as_ref()
.and_then(serde_json::Value::as_object)
.is_some_and(|terminal| terminal.contains_key("_meta"))
{
return Err(invalid_initialize_params(
"v2 terminal authentication metadata is not representable in v1",
));
}
auth.insert("terminal".into(), terminal.is_some().into());
}
capabilities.insert(
"session".into(),
serde_json::json!({ "configOptions": { "boolean": {} } }),
);
target.insert("clientCapabilities".into(), capabilities.into());
let initialize = parse_initialize_params::<InitializeRequest>(&target)?;
let normalized = serialize_initialize_params(initialize)?;
if require_lossless && !json_object_contains(&normalized, &target) {
return Err(invalid_initialize_params(
"v2 initialize parameters are not losslessly representable in v1",
));
}
Ok(normalized)
}
#[cfg(all(test, feature = "unstable_protocol_v2"))]
mod initialize_normalization_tests {
use super::*;
fn v2_initialize_params() -> serde_json::Map<String, serde_json::Value> {
let value = serde_json::to_value(v2::InitializeRequest::new(
ProtocolVersion::V2,
v2::Implementation::new("test-client", "1.0.0"),
))
.expect("serialize v2 initialize request");
value
.as_object()
.expect("initialize params serialize as an object")
.clone()
}
#[test]
fn v2_tolerant_fields_are_canonicalized_before_v1_normalization() {
let mut params = v2_initialize_params();
params.insert(
"capabilities".into(),
serde_json::Value::String("malformed".into()),
);
params.insert(
"_meta".into(),
serde_json::Value::String("malformed".into()),
);
let normalized = normalize_v2_initialize_params_for_v1(¶ms, false)
.expect("tolerant v2 fields should normalize through their defaults");
let normalized = serde_json::Value::Object(normalized);
assert!(normalized.get("_meta").is_none());
assert_eq!(
normalized.pointer("/clientCapabilities/session/configOptions/boolean"),
Some(&serde_json::json!({}))
);
}
#[test]
fn noncanonical_v2_fields_disable_reuse_but_not_v1_routing() {
let mut params = v2_initialize_params();
params
.get_mut("info")
.and_then(serde_json::Value::as_object_mut)
.expect("v2 initialize info is an object")
.insert("buildCommit".into(), serde_json::json!("abc123"));
normalize_v2_initialize_params_for_v1(¶ms, false)
.expect("v1 routing may ignore parameters unavailable in v1");
normalize_v2_initialize_params_for_v1(¶ms, true)
.expect_err("connection reuse requires lossless normalization");
}
#[test]
fn v1_reuse_probe_preserves_raw_initialize_params() {
let mut params = v2_initialize_params();
params.insert(
"protocolVersion".into(),
serde_json::json!(ProtocolVersion::V1),
);
let message = RawJsonRpcMessage::request(
"initialize".into(),
serde_json::Value::Object(params.clone()),
RequestId::Number(1),
)
.expect("build initialize request");
let normalized = normalize_v2_initialize_params_for_reuse(&message)
.expect("v1-shaped initialize request should be valid");
assert_eq!(normalized, params);
}
#[test]
fn null_v2_terminal_marker_meta_is_omitted_before_v1_normalization() {
let mut params = v2_initialize_params();
params.insert(
"capabilities".into(),
serde_json::json!({
"auth": {
"terminal": { "_meta": null }
}
}),
);
let normalized = normalize_v2_initialize_params_for_v1(¶ms, false)
.expect("null marker metadata is equivalent to omission");
let normalized = serde_json::Value::Object(normalized);
assert_eq!(
normalized.pointer("/clientCapabilities/auth/terminal"),
Some(&serde_json::Value::Bool(true))
);
}
#[test]
fn terminal_marker_metadata_disables_reuse_but_not_v1_routing() {
let mut params = v2_initialize_params();
params.insert(
"capabilities".into(),
serde_json::json!({
"auth": {
"terminal": {
"_meta": { "source": "test" }
}
}
}),
);
normalize_v2_initialize_params_for_v1(¶ms, false)
.expect("v1 routing may discard terminal marker metadata");
normalize_v2_initialize_params_for_v1(¶ms, true)
.expect_err("connection reuse must preserve terminal marker metadata");
}
}
#[cfg(feature = "unstable_protocol_v2")]
fn parse_initialize_params<T: DeserializeOwned>(
params: &serde_json::Map<String, serde_json::Value>,
) -> Result<T, crate::Error> {
serde_json::from_value(serde_json::Value::Object(params.clone()))
.map_err(invalid_initialize_params)
}
#[cfg(feature = "unstable_protocol_v2")]
fn serialize_initialize_params(
initialize: impl Serialize,
) -> Result<serde_json::Map<String, serde_json::Value>, crate::Error> {
let value = serde_json::to_value(initialize).map_err(crate::Error::into_internal_error)?;
let serde_json::Value::Object(object) = value else {
return Err(crate::util::internal_error(
"initialize params did not serialize to an object",
));
};
Ok(object)
}
#[cfg(feature = "unstable_protocol_v2")]
fn json_object_contains(
actual: &serde_json::Map<String, serde_json::Value>,
expected: &serde_json::Map<String, serde_json::Value>,
) -> bool {
fn contains(actual: &serde_json::Value, expected: &serde_json::Value) -> bool {
match (actual, expected) {
(serde_json::Value::Object(actual), serde_json::Value::Object(expected)) => expected
.iter()
.all(|(key, value)| actual.get(key).is_some_and(|item| contains(item, value))),
_ => actual == expected,
}
}
expected
.iter()
.all(|(key, value)| actual.get(key).is_some_and(|item| contains(item, value)))
}
#[cfg(feature = "unstable_protocol_v2")]
fn highest_compatible_agent_protocol(
requested: ProtocolVersion,
supported: SupportedProtocols,
) -> Result<SelectedProtocol, crate::Error> {
supported.highest_compatible(requested).ok_or_else(|| {
crate::Error::invalid_request().data(format!(
"unsupported ACP protocol version {requested}; this endpoint supports {}",
supported.description()
))
})
}
#[cfg(feature = "unstable_protocol_v2")]
fn invalid_initialize_protocol_version() -> crate::Error {
crate::Error::invalid_params()
.data("initialize.protocolVersion must be a valid ACP protocol version")
}
#[cfg(feature = "unstable_protocol_v2")]
fn invalid_initialize_params(error: impl ToString) -> crate::Error {
crate::Error::invalid_params().data(format!("invalid initialize params: {}", error.to_string()))
}
#[cfg(feature = "unstable_protocol_v2")]
fn send_initialize_error(
tx: &futures::channel::mpsc::UnboundedSender<TransportFrame>,
frame: &TransportFrame,
error: crate::Error,
) -> Result<(), crate::Error> {
fn response_for_message(
entry: &RawJsonRpcMessage,
initialize_error: &crate::Error,
) -> Option<RawJsonRpcMessage> {
match entry {
RawJsonRpcMessage::Request(request) => Some(RawJsonRpcMessage::response(
request.id.clone(),
Err(initialize_error.clone()),
)),
RawJsonRpcMessage::Notification(_) | RawJsonRpcMessage::Response(_) => None,
}
}
fn response_for_entry(
entry: &TransportBatchEntry,
initialize_error: &crate::Error,
) -> Option<RawJsonRpcMessage> {
match entry {
TransportBatchEntry::Message(message) => {
response_for_message(message, initialize_error)
}
TransportBatchEntry::Malformed { raw, error } if !is_response_only_shape(raw) => Some(
RawJsonRpcMessage::response(RequestId::Null, Err(error.clone())),
),
TransportBatchEntry::Malformed { .. } => None,
}
}
let response = match frame {
TransportFrame::Single(entry) => {
let Some(response) = response_for_message(entry, &error) else {
return Ok(());
};
TransportFrame::Single(response)
}
TransportFrame::Malformed { raw, error } if !raw_is_response_only_shape(raw) => {
TransportFrame::Single(RawJsonRpcMessage::response(
RequestId::Null,
Err(error.clone()),
))
}
TransportFrame::Malformed { .. } => return Ok(()),
TransportFrame::Batch(batch) => {
let responses = batch
.entries()
.filter_map(|entry| response_for_entry(entry, &error))
.collect::<Vec<_>>();
let Some(responses) = TransportBatch::from_messages(responses) else {
return Ok(());
};
TransportFrame::Batch(responses)
}
};
tx.unbounded_send(response)
.map_err(crate::util::internal_error)
}
#[cfg(feature = "unstable_protocol_v2")]
async fn reject_initialize(
client: RunningProtocolPeer,
frame: &TransportFrame,
error: crate::Error,
) -> Result<(), crate::Error> {
let RunningProtocolPeer { mut rx, tx, driver } = client;
send_initialize_error(&tx, frame, error)?;
drop(tx);
let Some(mut driver) = driver.into_driver() else {
return Ok(());
};
if !driver.request_finish() {
return crate::util::run_until(driver, future::ready(Ok(()))).await;
}
let drain_incoming = async move {
while rx.next().await.is_some() {}
Ok::<_, crate::Error>(())
};
match future::select(driver, Box::pin(drain_incoming)).await {
future::Either::Left((result, _)) => result,
future::Either::Right((result, driver)) => {
result?;
driver.await
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
struct RunningProtocolPeer {
rx: futures::channel::mpsc::UnboundedReceiver<TransportFrame>,
tx: futures::channel::mpsc::UnboundedSender<TransportFrame>,
driver: ProtocolPeerDriver,
}
#[cfg(feature = "unstable_protocol_v2")]
enum ProtocolPeerDriver {
Passive,
Active(crate::ConnectionDriver),
Completed {
finish: Option<crate::component::FinishControl>,
},
}
#[cfg(feature = "unstable_protocol_v2")]
impl ProtocolPeerDriver {
fn into_driver(self) -> Option<crate::ConnectionDriver> {
match self {
Self::Passive => None,
Self::Active(driver) => Some(driver),
Self::Completed { finish } => Some(match finish {
Some(mut finish) => {
crate::ConnectionDriver::with_finish(future::ready(Ok(())), move || {
finish.request();
})
}
None => crate::ConnectionDriver::new(future::ready(Ok(()))),
}),
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
impl RunningProtocolPeer {
fn new<R: Role>(component: impl ConnectTo<R>) -> Self {
let (Channel { rx, tx }, future) = component.into_channel_and_future();
let driver = match future {
None => ProtocolPeerDriver::Passive,
Some(future) => ProtocolPeerDriver::Active(future),
};
Self { rx, tx, driver }
}
async fn next_frame(self) -> Result<Option<(TransportFrame, Self)>, crate::Error> {
let Self { mut rx, tx, driver } = self;
let ProtocolPeerDriver::Active(mut future) = driver else {
return Ok(rx
.next()
.await
.map(|frame| (frame, Self { rx, tx, driver })));
};
match future::select(&mut future, Box::pin(rx.next())).await {
future::Either::Right((Some(frame), _)) => Ok(Some((
frame,
Self {
rx,
tx,
driver: ProtocolPeerDriver::Active(future),
},
))),
future::Either::Right((None, _)) => {
drop(tx);
future.await?;
Ok(None)
}
future::Either::Left((result, next_message)) => {
result?;
drop(next_message);
rx.close();
let Some(frame) = rx.next().await else {
return Ok(None);
};
Ok(Some((
frame,
Self {
rx,
tx,
driver: ProtocolPeerDriver::Completed {
finish: future.take_finish(),
},
},
)))
}
}
}
async fn next_message(self) -> Result<Option<(RawJsonRpcMessage, Self)>, crate::Error> {
let Some((frame, peer)) = self.next_frame().await? else {
return Ok(None);
};
Ok(Some((initialize_message(frame)?, peer)))
}
fn send(&self, message: RawJsonRpcMessage) -> Result<(), crate::Error> {
self.send_frame(TransportFrame::Single(message))
}
fn send_frame(&self, frame: TransportFrame) -> Result<(), crate::Error> {
self.tx
.unbounded_send(frame)
.map_err(crate::util::internal_error)
}
}
#[cfg(feature = "unstable_protocol_v2")]
fn initialize_message(frame: TransportFrame) -> Result<RawJsonRpcMessage, crate::Error> {
match frame {
TransportFrame::Single(message) => Ok(message),
TransportFrame::Malformed { error, .. } => Err(error),
TransportFrame::Batch(_) => Err(crate::Error::invalid_request()
.data("ACP initialize request and response messages must be sent individually")),
}
}
#[cfg(feature = "unstable_protocol_v2")]
fn initialize_message_mut(
frame: &mut TransportFrame,
) -> Result<Option<&mut RawJsonRpcMessage>, crate::Error> {
match frame {
TransportFrame::Single(RawJsonRpcMessage::Response(_)) => Ok(None),
TransportFrame::Single(entry) => Ok(Some(entry)),
TransportFrame::Malformed { raw, .. } if raw_is_response_only_shape(raw) => Ok(None),
TransportFrame::Malformed { error, .. } => Err(error.clone()),
TransportFrame::Batch(batch) => {
for entry in batch.entries_mut() {
match entry {
TransportBatchEntry::Message(RawJsonRpcMessage::Response(_)) => {}
TransportBatchEntry::Message(message) => return Ok(Some(message)),
TransportBatchEntry::Malformed { raw, .. } if is_response_only_shape(raw) => {}
TransportBatchEntry::Malformed { error, .. } => return Err(error.clone()),
}
}
Ok(None)
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
async fn pipe_protocol_peers_until_done(
left: RunningProtocolPeer,
right: RunningProtocolPeer,
) -> Result<(), crate::Error> {
let mut left_driver = left.driver.into_driver();
let mut right_driver = right.driver.into_driver();
let left_passive = left_driver.is_none();
let right_passive = right_driver.is_none();
let left_finish = left_driver
.as_mut()
.and_then(crate::ConnectionDriver::take_finish);
let right_finish = right_driver
.as_mut()
.and_then(crate::ConnectionDriver::take_finish);
let (stop_left_tx, stop_left_rx) = futures::channel::oneshot::channel();
let (stop_right_tx, stop_right_rx) = futures::channel::oneshot::channel();
let stop = async |rx: futures::channel::oneshot::Receiver<()>| {
if rx.await.is_err() {
future::pending::<()>().await;
}
};
let left_to_right = Box::pin(
Channel {
rx: left.rx,
tx: right.tx,
}
.copy_with_driver_until(left_driver, stop(stop_left_rx)),
);
let right_to_left = Box::pin(
Channel {
rx: right.rx,
tx: left.tx,
}
.copy_with_driver_until(right_driver, stop(stop_right_rx)),
);
match future::select(left_to_right, right_to_left).await {
future::Either::Left((result, right_to_left)) => {
result?;
if !left_passive {
let _ = stop_right_tx.send(());
}
if left_passive || right_finish.is_some() {
if !left_passive && let Some(mut finish) = right_finish {
finish.request();
}
right_to_left.await
} else {
crate::util::run_until(right_to_left, future::ready(Ok(()))).await
}
}
future::Either::Right((result, left_to_right)) => {
result?;
if !right_passive {
let _ = stop_left_tx.send(());
}
if right_passive || left_finish.is_some() {
if !right_passive && let Some(mut finish) = left_finish {
finish.request();
}
left_to_right.await
} else {
crate::util::run_until(left_to_right, future::ready(Ok(()))).await
}
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug)]
struct InitializeResponse {
id: RequestId,
result: Result<serde_json::Value, Box<RawJsonRpcError>>,
}
#[cfg(feature = "unstable_protocol_v2")]
impl InitializeResponse {
fn from_message(message: RawJsonRpcMessage) -> Result<Self, crate::Error> {
match message {
RawJsonRpcMessage::Response(RpcResponse::Result { id, result }) => Ok(Self {
id,
result: Ok(result),
}),
RawJsonRpcMessage::Response(RpcResponse::Error { id, error }) => Ok(Self {
id,
result: Err(error),
}),
message => Err(crate::Error::invalid_request().data(format!(
"first ACP response must be an initialize response, got {message:?}",
))),
}
}
fn into_message(self) -> RawJsonRpcMessage {
RawJsonRpcMessage::Response(RpcResponse::new(self.id, self.result))
}
fn with_id(self, id: RequestId) -> RawJsonRpcMessage {
RawJsonRpcMessage::Response(RpcResponse::new(id, self.result))
}
fn protocol_version(&self) -> Option<ProtocolVersion> {
serde_json::from_value(self.result.as_ref().ok()?.get("protocolVersion")?.clone()).ok()
}
}
#[cfg(all(test, feature = "unstable_protocol_v2"))]
mod raw_initialize_tests {
use super::*;
use serde_json::json;
#[test]
fn initialize_error_forwarding_preserves_raw_fields() {
for data in [
None,
Some(serde_json::Value::Null),
Some(json!({"detail":"kept"})),
] {
let mut error = json!({
"code":-32000, "message":"peer", "extension":{"retry":true}
});
if let Some(data) = data {
error["data"] = data;
}
let wire = json!({"jsonrpc":"2.0", "id":"original", "error":error});
let response =
InitializeResponse::from_message(serde_json::from_value(wire.clone()).unwrap())
.unwrap();
assert_eq!(serde_json::to_value(response.into_message()).unwrap(), wire);
let response =
InitializeResponse::from_message(serde_json::from_value(wire).unwrap()).unwrap();
let forwarded = response.with_id(RequestId::Str("replacement".into()));
assert_eq!(
serde_json::to_value(forwarded).unwrap(),
json!({"jsonrpc":"2.0", "id":"replacement", "error":error})
);
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ClientProtocol {
V1,
V2,
}
#[cfg(feature = "unstable_protocol_v2")]
impl ClientProtocol {
fn name(self) -> &'static str {
match self {
Self::V1 => "1",
Self::V2 => "2",
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct SupportedClientProtocols {
v1: bool,
v2: bool,
}
#[cfg(feature = "unstable_protocol_v2")]
impl SupportedClientProtocols {
fn highest_configured(self) -> Option<ClientProtocol> {
if self.v2 {
return Some(ClientProtocol::V2);
}
if self.v1 {
return Some(ClientProtocol::V1);
}
None
}
}
#[cfg(feature = "unstable_protocol_v2")]
async fn start_client_protocol(
protocol: ClientProtocol,
client: DynConnectTo<Agent>,
) -> Result<(RunningProtocolPeer, RawJsonRpcMessage), crate::Error> {
let client = RunningProtocolPeer::new(client);
let Some((initialize, client)) = client.next_message().await? else {
return Err(crate::Error::invalid_request().data(format!(
"ACP protocol version {} client implementation ended before initialize",
protocol.name()
)));
};
ensure_client_initialize_request(protocol, &initialize)?;
Ok((client, initialize))
}
#[cfg(feature = "unstable_protocol_v2")]
async fn send_initialize_and_receive(
client: RunningProtocolPeer,
agent: RunningProtocolPeer,
initialize: RawJsonRpcMessage,
) -> Result<(RunningProtocolPeer, RunningProtocolPeer, InitializeResponse), crate::Error> {
agent.send(initialize)?;
let Some((response, agent)) = agent.next_message().await? else {
return Err(crate::Error::internal_error().data("agent closed before initialize response"));
};
let response = InitializeResponse::from_message(response)?;
Ok((client, agent, response))
}
#[cfg(feature = "unstable_protocol_v2")]
async fn initialize_client_protocol(
protocol: ClientProtocol,
client: DynConnectTo<Agent>,
agent: impl ConnectTo<Client>,
) -> Result<(RunningProtocolPeer, RunningProtocolPeer, InitializeResponse), crate::Error> {
let agent = RunningProtocolPeer::new(agent);
let (client, initialize) = start_client_protocol(protocol, client).await?;
send_initialize_and_receive(client, agent, initialize).await
}
#[cfg(feature = "unstable_protocol_v2")]
async fn connect_client_protocol(
protocol: ClientProtocol,
client: DynConnectTo<Agent>,
agent: impl ConnectTo<Client>,
) -> Result<(), crate::Error> {
let (client, agent, initialize_response) =
initialize_client_protocol(protocol, client, agent).await?;
client.send(initialize_response.into_message())?;
pipe_protocol_peers_until_done(client, agent).await
}
#[cfg(feature = "unstable_protocol_v2")]
fn ensure_client_initialize_request(
protocol: ClientProtocol,
message: &RawJsonRpcMessage,
) -> Result<(), crate::Error> {
let RawJsonRpcMessage::Request(request) = message else {
return Err(crate::Error::invalid_request().data(format!(
"ACP protocol version {} client implementation must send initialize first",
protocol.name()
)));
};
if request.method.as_ref() != "initialize" {
return Err(crate::Error::invalid_request().data(format!(
"ACP protocol version {} client implementation must send initialize first",
protocol.name()
)));
}
Ok(())
}
#[cfg(feature = "unstable_protocol_v2")]
fn initialize_request_id(message: &RawJsonRpcMessage) -> Option<RequestId> {
let RawJsonRpcMessage::Request(request) = message else {
return None;
};
Some(request.id.clone())
}
#[cfg(feature = "unstable_protocol_v2")]
fn initialize_response_negotiated_v1(response: &InitializeResponse) -> bool {
response.protocol_version() == Some(ProtocolVersion::V1)
}
impl HasPeer<Agent> for Agent {
fn remote_style(&self, _peer: Agent) -> RemoteStyle {
RemoteStyle::Counterpart
}
}
#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Proxy;
impl Role for Proxy {
type Counterpart = Conductor;
fn default_handle_dispatch_from(
&self,
message: crate::Dispatch,
_connection: crate::ConnectionTo<Self>,
) -> impl Future<Output = Result<crate::Handled<crate::Dispatch>, crate::Error>> + Send {
std::future::ready(Ok(Handled::No {
message,
retry: false,
}))
}
fn role_id(&self) -> RoleId {
RoleId::from_singleton(self)
}
fn counterpart(&self) -> Self::Counterpart {
Conductor
}
}
impl Proxy {
pub fn builder(self) -> Builder<Proxy, NullHandler, NullRun> {
Builder::new(self)
}
#[cfg(feature = "unstable_protocol_v2")]
pub fn v2(self) -> V2Builder<Proxy, NullHandler, NullRun> {
self.builder().v2_proxy()
}
#[cfg(feature = "unstable_protocol_v2")]
#[must_use]
pub fn protocol_router(self) -> ProxyProtocolRouter {
ProxyProtocolRouter::new()
}
}
#[cfg(feature = "unstable_protocol_v2")]
#[derive(Debug, Default)]
pub struct ProxyProtocolRouter {
v1: Option<DynConnectTo<Conductor>>,
v2: Option<DynConnectTo<Conductor>>,
}
#[cfg(feature = "unstable_protocol_v2")]
impl ProxyProtocolRouter {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_v1(mut self, proxy: impl ConnectTo<Conductor>) -> Self {
self.v1 = Some(DynConnectTo::new(proxy));
self
}
#[must_use]
pub fn with_v2(mut self, proxy: impl ConnectTo<Conductor>) -> Self {
self.v2 = Some(DynConnectTo::new(proxy));
self
}
}
#[cfg(feature = "unstable_protocol_v2")]
impl ConnectTo<Conductor> for ProxyProtocolRouter {
async fn connect_to(self, conductor: impl ConnectTo<Proxy>) -> Result<(), crate::Error> {
let supported = SupportedProtocols {
v1: self.v1.is_some(),
v2: self.v2.is_some(),
};
let mut conductor = RunningProtocolPeer::new(conductor);
let (first_frame, conductor, selected) = loop {
let Some((mut frame, next_conductor)) = conductor.next_frame().await? else {
return Ok(());
};
let message = match initialize_message_mut(&mut frame) {
Ok(Some(message)) => message,
Ok(None) => {
conductor = next_conductor;
continue;
}
Err(error) => return reject_initialize(next_conductor, &frame, error).await,
};
let selected = match select_proxy_protocol(message, supported) {
Ok(selected) => selected,
Err(error) => return reject_initialize(next_conductor, &frame, error).await,
};
break (frame, next_conductor, selected);
};
let Some(proxy) = selected.take_proxy(self) else {
let error = selected.unsupported_error(supported);
return reject_initialize(conductor, &first_frame, error).await;
};
let proxy = RunningProtocolPeer::new(proxy);
proxy.send_frame(first_frame)?;
pipe_protocol_peers_until_done(conductor, proxy).await
}
}
#[cfg(feature = "unstable_protocol_v2")]
impl SelectedProtocol {
fn take_proxy(self, proxy: ProxyProtocolRouter) -> Option<DynConnectTo<Conductor>> {
match self {
Self::V1 => proxy.v1,
Self::V2 => proxy.v2,
}
}
}
#[cfg(feature = "unstable_protocol_v2")]
fn select_proxy_protocol(
message: &RawJsonRpcMessage,
supported: SupportedProtocols,
) -> Result<SelectedProtocol, crate::Error> {
let RawJsonRpcMessage::Request(request) = message else {
return Err(crate::Error::invalid_request()
.data("first ACP proxy message must be an `_proxy/initialize` request"));
};
if request.method.as_ref() != METHOD_INITIALIZE_PROXY {
return Err(crate::Error::invalid_request()
.data("first ACP proxy request must be `_proxy/initialize`"));
}
let Some(RawJsonRpcParams::Object(params)) = &request.params else {
return Err(invalid_initialize_protocol_version());
};
let Some(protocol_version) = params.get("protocolVersion") else {
return Err(invalid_initialize_protocol_version());
};
let requested = serde_json::from_value::<ProtocolVersion>(protocol_version.clone())
.map_err(|_| invalid_initialize_protocol_version())?;
let selected = supported.exact(requested).ok_or_else(|| {
crate::Error::invalid_request().data(format!(
"unsupported ACP protocol version {requested}; this proxy supports {}",
supported.description()
))
})?;
match selected {
SelectedProtocol::V1 => {
parse_initialize_params::<crate::schema::InitializeProxyRequest>(params)?;
}
SelectedProtocol::V2 => {
parse_initialize_params::<v2::InitializeProxyRequest>(params)?;
}
}
Ok(selected)
}
impl HasPeer<Proxy> for Proxy {
fn remote_style(&self, _peer: Proxy) -> RemoteStyle {
RemoteStyle::Counterpart
}
}
#[derive(Debug, Default, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Conductor;
impl Role for Conductor {
type Counterpart = Proxy;
fn role_id(&self) -> RoleId {
RoleId::from_singleton(self)
}
fn counterpart(&self) -> Self::Counterpart {
Proxy
}
async fn default_handle_dispatch_from(
&self,
message: Dispatch,
cx: ConnectionTo<Conductor>,
) -> Result<Handled<Dispatch>, crate::Error> {
#[cfg(not(feature = "unstable_protocol_v2"))]
{
MatchDispatchFrom::new(message, &cx)
.if_request_from(Client, async |_req: InitializeRequest, responder| {
responder.respond_with_error(crate::Error::invalid_request().data(format!(
"proxies must be initialized with `{METHOD_INITIALIZE_PROXY}`"
)))
})
.await
.if_request_from(
Client,
async |request: InitializeProxyRequest, responder| {
let InitializeProxyRequest { initialize } = request;
cx.send_ordered_request_to(Agent, initialize)
.forward_response_to(responder)
},
)
.await
.if_request_from(Client, async |request: NewSessionRequest, responder| {
let sent = cx.send_ordered_request_to(Agent, request);
let sent = sent.forward_cancellation_from(responder.cancellation());
sent.on_receiving_result({
let cx = cx.clone();
async move |result| {
if let Ok(NewSessionResponse { session_id, .. }) = &result {
cx.add_dynamic_handler(ProxySessionMessages::new(
session_id.clone(),
))?
.detach();
}
responder.respond_with_result(result)
}
})
})
.await
.if_dispatch_from(Client, async |message: Dispatch| {
cx.send_proxied_message_to(Agent, message)
})
.await
.if_dispatch_from(Agent, async |message: Dispatch| {
cx.send_proxied_message_to(Client, message)
})
.await
.done()
}
#[cfg(feature = "unstable_protocol_v2")]
{
let message = match message {
Dispatch::Request(request, responder) if request.method() == "initialize" => {
responder.respond_with_error(crate::Error::invalid_request().data(format!(
"proxies must be initialized with `{METHOD_INITIALIZE_PROXY}`"
)))?;
return Ok(Handled::Yes);
}
Dispatch::Request(mut request, responder)
if request.method() == METHOD_INITIALIZE_PROXY =>
{
request.method = "initialize".to_string();
cx.send_ordered_request_to(Agent, request)
.forward_response_to(responder)?;
return Ok(Handled::Yes);
}
Dispatch::Request(request, responder) if request.method() == "session/new" => {
let sent = cx.send_ordered_request_to(Agent, request);
let sent = sent.forward_cancellation_from(responder.cancellation());
sent.on_receiving_result({
let cx = cx.clone();
async move |result| {
let result = result.and_then(|response| {
let envelope: NewSessionResponseEnvelope =
crate::util::json_cast(response.clone())?;
cx.add_dynamic_handler(ProxySessionMessages::new(
envelope.session_id,
))?
.detach();
Ok(response)
});
responder.respond_with_result(result)
}
})?;
return Ok(Handled::Yes);
}
message => message,
};
MatchDispatchFrom::new(message, &cx)
.if_dispatch_from(Client, async |message: Dispatch| {
cx.send_proxied_message_to(Agent, message)
})
.await
.if_dispatch_from(Agent, async |message: Dispatch| {
cx.send_proxied_message_to(Client, message)
})
.await
.done()
}
}
}
impl Conductor {
pub fn builder(self) -> Builder<Conductor, NullHandler, NullRun> {
Builder::new(self)
}
}
impl HasPeer<Client> for Conductor {
fn remote_style(&self, _peer: Client) -> RemoteStyle {
RemoteStyle::Predecessor
}
}
impl HasPeer<Agent> for Conductor {
fn remote_style(&self, _peer: Agent) -> RemoteStyle {
RemoteStyle::Successor
}
}
pub(crate) struct ProxySessionMessages {
session_id: SessionId,
}
impl ProxySessionMessages {
pub fn new(session_id: SessionId) -> Self {
Self { session_id }
}
}
impl<Counterpart: Role> HandleDispatchFrom<Counterpart> for ProxySessionMessages
where
Counterpart: HasPeer<Agent> + HasPeer<Client>,
{
async fn handle_dispatch_from(
&mut self,
message: Dispatch,
connection: ConnectionTo<Counterpart>,
) -> Result<Handled<Dispatch>, crate::Error> {
MatchDispatchFrom::new(message, &connection)
.if_dispatch_from(Agent, async |message| {
if let Some(session_id) = message.get_session_id()?
&& session_id == self.session_id
{
connection.send_proxied_message_to(Client, message)?;
return Ok(Handled::Yes);
}
Ok(Handled::No {
message,
retry: false,
})
})
.await
.done()
}
fn describe_chain(&self) -> impl std::fmt::Debug {
format!("ProxySessionMessages({})", self.session_id)
}
}
#[cfg(all(test, feature = "unstable_protocol_v2"))]
mod lifetime_tests {
use super::*;
use crate::{ConnectionDriver, UntypedRole};
use futures::FutureExt as _;
fn frame() -> TransportFrame {
TransportFrame::parse_json(r#"{"jsonrpc":"2.0","method":"test/queued","params":{}}"#)
}
#[tokio::test]
async fn passive_initialization_waits_for_a_frame_not_driver_readiness() {
let (channel, remote) = Channel::duplex();
let peer = RunningProtocolPeer::new::<UntypedRole>(channel);
let mut next = Box::pin(peer.next_frame());
assert!(next.as_mut().now_or_never().is_none());
remote.tx.unbounded_send(frame()).unwrap();
let (_, peer) = next.await.unwrap().expect("passive peer remains connected");
assert!(matches!(peer.driver, ProtocolPeerDriver::Passive));
}
#[tokio::test]
async fn owned_peer_preserves_finish_metadata_through_active_and_completed_states() {
let (Channel { rx, tx }, remote) = Channel::duplex();
let (done_tx, done_rx) = futures::channel::oneshot::channel();
let (finish_tx, finish_rx) = futures::channel::oneshot::channel();
let peer = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::with_finish(
async move {
done_rx.await.unwrap();
Ok(())
},
move || {
let _ = finish_tx.send(());
},
)),
};
remote.tx.unbounded_send(frame()).unwrap();
let (_, peer) = peer.next_frame().await.unwrap().unwrap();
assert!(matches!(&peer.driver, ProtocolPeerDriver::Active(_)));
remote.tx.unbounded_send(frame()).unwrap();
done_tx.send(()).unwrap();
let (_, peer) = peer.next_frame().await.unwrap().unwrap();
assert!(matches!(
&peer.driver,
ProtocolPeerDriver::Completed { finish: Some(_) }
));
let mut driver = peer
.driver
.into_driver()
.expect("completed owned work must not become passive");
assert!(driver.request_finish());
finish_rx.await.unwrap();
driver.await.unwrap();
}
#[tokio::test]
async fn active_initialization_drains_accepted_frames_without_escaped_sender_eof() {
let (Channel { rx, tx }, remote) = Channel::duplex();
remote.tx.unbounded_send(frame()).unwrap();
remote.tx.unbounded_send(frame()).unwrap();
let mut polls = 0;
let driver = ConnectionDriver::new(future::poll_fn(move |_| {
polls += 1;
assert_eq!(polls, 1, "the completed driver must never be re-polled");
std::task::Poll::Ready(Ok(()))
}));
let peer = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(driver),
};
let (_, peer) = peer.next_frame().await.unwrap().unwrap();
assert!(
remote.tx.unbounded_send(frame()).is_err(),
"active completion must reject new output from escaped handles"
);
let (_, peer) = peer.next_frame().await.unwrap().unwrap();
assert!(peer.next_frame().await.unwrap().is_none());
}
#[tokio::test]
async fn ready_initialization_driver_error_beats_queued_frames() {
let (Channel { rx, tx }, remote) = Channel::duplex();
remote.tx.unbounded_send(frame()).unwrap();
let error = crate::Error::internal_error().data("owned initialization failed");
let peer = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Err(
error.clone()
)))),
};
match peer.next_frame().await {
Err(actual) => assert_eq!(actual, error),
Ok(_) => panic!("ready driver error must not be hidden by a queued frame"),
}
}
struct QueuedFinalClient;
impl ConnectTo<Agent> for QueuedFinalClient {
async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), crate::Error> {
let (mut channel, driver) = agent.into_channel_and_future();
let foreground = async move {
channel
.tx
.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request(
"initialize".into(),
serde_json::json!({ "protocolVersion": 1, "clientCapabilities": {} }),
RequestId::Number(1),
)?))
.unwrap();
assert!(channel.rx.next().await.is_some(), "initialize response");
for index in 0..3 {
channel
.tx
.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification(
"test/final".into(),
serde_json::json!({ "index": index, "payload": "x".repeat(1024) }),
)?))
.unwrap();
}
Ok(())
};
match driver {
Some(driver) => crate::util::run_until(driver, foreground).await,
None => foreground.await,
}
}
}
#[tokio::test]
async fn connector_completion_flushes_byte_streams_without_remote_read_eof() {
use tokio::io::{AsyncBufReadExt as _, AsyncWriteExt as _, BufReader};
use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _};
let (writer, remote_reader) = tokio::io::duplex(64);
let (mut remote_input, reader) = tokio::io::duplex(64);
let physical = crate::ByteStreams::new(writer.compat_write(), reader.compat());
let mut physical = Some(crate::DynConnectTo::<Client>::new(physical));
let connector = tokio::spawn(
ClientProtocolConnector::new()
.with_v1(|| QueuedFinalClient)
.connect_to(move || physical.take().expect("one physical connection")),
);
let mut lines = BufReader::new(remote_reader).lines();
tokio::time::timeout(std::time::Duration::from_secs(5), async {
let initialize = lines
.next_line()
.await
.unwrap()
.expect("initialize request");
let value: serde_json::Value = serde_json::from_str(&initialize).unwrap();
assert_eq!(value["method"], "initialize");
remote_input
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{\"protocolVersion\":1}}\n")
.await
.unwrap();
for index in 0..3 {
let line = lines
.next_line()
.await
.unwrap()
.expect("accepted final output");
let value: serde_json::Value = serde_json::from_str(&line).unwrap();
assert_eq!(value["params"]["index"], index);
assert_eq!(value["params"]["payload"].as_str().unwrap().len(), 1024);
}
connector.await.unwrap().unwrap();
assert!(lines.next_line().await.unwrap().is_none());
})
.await
.expect("physical flush must not wait for independent remote input EOF");
drop(remote_input);
}
#[derive(Default, Debug)]
struct GatedLineSinkState {
pending: Vec<String>,
flushed: Vec<String>,
closed: bool,
dropped: bool,
}
struct GatedLineSink {
state: std::sync::Arc<std::sync::Mutex<GatedLineSinkState>>,
release: futures::channel::oneshot::Receiver<()>,
released: bool,
}
impl futures::Sink<String> for GatedLineSink {
type Error = std::io::Error;
fn poll_ready(
self: std::pin::Pin<&mut Self>,
_: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Ready(Ok(()))
}
fn start_send(self: std::pin::Pin<&mut Self>, line: String) -> Result<(), Self::Error> {
self.state.lock().unwrap().pending.push(line);
Ok(())
}
fn poll_flush(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
if !self.released {
if std::pin::Pin::new(&mut self.release).poll(cx).is_pending() {
return std::task::Poll::Pending;
}
self.released = true;
}
let mut state = self.state.lock().unwrap();
let pending = std::mem::take(&mut state.pending);
state.flushed.extend(pending);
std::task::Poll::Ready(Ok(()))
}
fn poll_close(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
match self.as_mut().poll_flush(cx) {
std::task::Poll::Ready(Ok(())) => {
self.state.lock().unwrap().closed = true;
std::task::Poll::Ready(Ok(()))
}
result => result,
}
}
}
impl Drop for GatedLineSink {
fn drop(&mut self) {
self.state.lock().unwrap().dropped = true;
}
}
#[tokio::test]
async fn foreground_completion_flushes_lines_with_already_normalized_remote_input() {
let state = std::sync::Arc::new(std::sync::Mutex::new(GatedLineSinkState::default()));
let (release_tx, release_rx) = futures::channel::oneshot::channel();
let sink = GatedLineSink {
state: state.clone(),
release: release_rx,
released: false,
};
let (remote_input, incoming) = futures::channel::mpsc::unbounded();
remote_input
.unbounded_send(Ok(
r#"{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":1}}"#.to_string(),
))
.unwrap();
remote_input
.unbounded_send(Ok(
r#"{"jsonrpc":"2.0","method":"test/queued","params":{}}"#.to_string(),
))
.unwrap();
let physical = RunningProtocolPeer::new::<Client>(crate::Lines::new(sink, incoming));
let (initialize, physical) = physical.next_frame().await.unwrap().unwrap();
assert_eq!(
futures::Stream::size_hint(&physical.rx).0,
1,
"trailing input must already be in the original normalized queue"
);
let (Channel { rx, tx }, mut local) = Channel::duplex();
tx.unbounded_send(initialize).unwrap();
let foreground = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(async move {
assert!(local.rx.next().await.is_some(), "initialize response");
for index in 0..3 {
local
.tx
.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification(
"test/final".into(),
serde_json::json!({ "index": index, "payload": "x".repeat(1024) }),
)?))
.unwrap();
}
drop(local);
Ok(())
})),
};
let mut bridge = Box::pin(pipe_protocol_peers_until_done(foreground, physical));
let early_result = bridge.as_mut().now_or_never();
assert!(
!state.lock().unwrap().pending.is_empty(),
"the real physical writer must accept output before the gated flush"
);
let _released = release_tx.send(());
let result = match early_result {
Some(result) => result,
None => tokio::time::timeout(std::time::Duration::from_secs(1), bridge)
.await
.expect("physical drain must not require remote input EOF"),
};
let state = state.lock().unwrap();
assert!(
result.is_ok() && state.flushed.len() == 3 && state.closed,
"accepted output must drain cleanly despite queued remote input: result={result:?}, pending={}, flushed={}, closed={}, dropped={}",
state.pending.len(),
state.flushed.len(),
state.closed,
state.dropped,
);
for (index, line) in state.flushed.iter().enumerate() {
let value: serde_json::Value = serde_json::from_str(line).unwrap();
assert_eq!(value["method"], "test/final");
assert_eq!(value["params"]["index"], index);
assert_eq!(value["params"]["payload"].as_str().unwrap().len(), 1024);
}
drop(remote_input);
}
#[tokio::test]
async fn foreground_completion_keeps_read_errors_during_lines_drain() {
let state = std::sync::Arc::new(std::sync::Mutex::new(GatedLineSinkState::default()));
let (_release_tx, release_rx) = futures::channel::oneshot::channel();
let sink = GatedLineSink {
state: state.clone(),
release: release_rx,
released: false,
};
let (remote_input, incoming) = futures::channel::mpsc::unbounded();
let physical = RunningProtocolPeer::new::<Client>(crate::Lines::new(sink, incoming));
let (Channel { rx, tx }, local) = Channel::duplex();
local.tx.unbounded_send(frame()).unwrap();
drop(local);
let foreground = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))),
};
let mut bridge = Box::pin(pipe_protocol_peers_until_done(foreground, physical));
assert!(bridge.as_mut().now_or_never().is_none());
assert_eq!(state.lock().unwrap().pending.len(), 1);
remote_input
.unbounded_send(Ok(r#"{"jsonrpc":"2.0","method":"late"}"#.to_string()))
.unwrap();
remote_input
.unbounded_send(Err(std::io::Error::other("read failed after foreground")))
.unwrap();
let error = tokio::time::timeout(std::time::Duration::from_secs(1), bridge)
.await
.expect("read failure must not wait for the sink gate")
.expect_err("real read error must win over foreground success");
assert!(
error
.data
.unwrap()
.to_string()
.contains("read failed after foreground"),
"the read error must not become a receiver-gone forwarding error"
);
assert_eq!(state.lock().unwrap().flushed, Vec::<String>::new());
}
#[tokio::test]
async fn foreground_completion_cancels_opposed_opaque_work_in_both_directions() {
for foreground_on_left in [true, false] {
let (Channel { rx, tx }, foreground_remote) = Channel::duplex();
foreground_remote.tx.unbounded_send(frame()).unwrap();
let foreground = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))),
};
let (Channel { rx, tx }, mut opposed_remote) = Channel::duplex();
let (work_tx, work_rx) = futures::channel::oneshot::channel::<()>();
let opposed = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(async move {
work_rx.await.map_err(crate::util::internal_error)?;
Ok(())
})),
};
let mut bridge = Box::pin(if foreground_on_left {
pipe_protocol_peers_until_done(foreground, opposed)
} else {
pipe_protocol_peers_until_done(opposed, foreground)
});
assert_eq!(
bridge.as_mut().now_or_never(),
Some(Ok(())),
"finite foreground must not join an opaque pending peer"
);
assert!(work_tx.is_canceled(), "opaque work must be dropped");
assert!(opposed_remote.rx.next().await.is_some());
assert!(opposed_remote.rx.next().await.is_none());
}
}
#[tokio::test]
async fn foreground_completion_does_not_hide_opposed_ready_driver_error() {
let (Channel { rx, tx }, _foreground_remote) = Channel::duplex();
let foreground = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))),
};
let (Channel { rx, tx }, _opposed_remote) = Channel::duplex();
let error = crate::Error::internal_error().data("opposed driver failed");
let opposed = RunningProtocolPeer {
rx,
tx,
driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Err(
error.clone()
)))),
};
assert_eq!(
pipe_protocol_peers_until_done(foreground, opposed).await,
Err(error)
);
}
#[tokio::test]
async fn passive_protocol_bridge_preserves_the_reverse_half_after_eof() {
let (left, mut remote_left) = Channel::duplex();
let (right, mut remote_right) = Channel::duplex();
let mut bridge = Box::pin(pipe_protocol_peers_until_done(
RunningProtocolPeer::new::<UntypedRole>(left),
RunningProtocolPeer::new::<UntypedRole>(right),
));
remote_left.tx.close_channel();
assert!(bridge.as_mut().now_or_never().is_none());
assert!(remote_right.rx.next().await.is_none());
remote_right.tx.unbounded_send(frame()).unwrap();
remote_right.tx.close_channel();
bridge.await.unwrap();
assert!(remote_left.rx.next().await.is_some());
assert!(remote_left.rx.next().await.is_none());
}
}