File size: 6,131 Bytes
afa0cbf | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | use std::io;
use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient;
use codex_code_mode_protocol::host::MAX_FRAME_BYTES;
use codex_http_client::ClientRouteClass;
use codex_http_client::HttpClientFactory;
use http_body_util::BodyExt;
use tonic::body::Body;
use tonic::codegen::http::Request;
use tonic::codegen::http::Response;
use tonic::codegen::http::Uri;
use tonic::transport::Channel;
use tonic::transport::Endpoint;
use tower::ServiceExt;
use tower::service_fn;
use tower::util::BoxCloneSyncService;
use super::GrpcClient;
pub(super) type GrpcTransport = BoxCloneSyncService<Request<Body>, Response<Body>, io::Error>;
pub(super) struct SharedTransport {
endpoint: TransportEndpoint,
client: tokio::sync::OnceCell<GrpcClient>,
}
enum TransportEndpoint {
Url {
endpoint: String,
http_client_factory: HttpClientFactory,
},
Connected(Channel),
}
impl SharedTransport {
pub(super) fn new(endpoint: String, http_client_factory: HttpClientFactory) -> Self {
Self {
endpoint: TransportEndpoint::Url {
endpoint,
http_client_factory,
},
client: tokio::sync::OnceCell::new(),
}
}
pub(super) fn with_channel(channel: Channel) -> Self {
Self {
endpoint: TransportEndpoint::Connected(channel),
client: tokio::sync::OnceCell::new(),
}
}
pub(super) async fn client(&self) -> Result<GrpcClient, String> {
self.client
.get_or_try_init(|| async {
let client = match &self.endpoint {
TransportEndpoint::Url { endpoint, .. } if endpoint.starts_with("unix:") => {
let channel = Endpoint::from_shared(endpoint.clone())
.map_err(|error| {
format!("invalid gRPC code-mode Unix socket endpoint: {error}")
})?
.connect_lazy();
let transport = channel.map_err(io::Error::other);
CodeModeHostClient::new(BoxCloneSyncService::new(transport))
}
TransportEndpoint::Url {
endpoint,
http_client_factory,
} => {
let target = reqwest::Url::parse(endpoint)
.map_err(|error| format!("invalid gRPC code-mode host URL: {error}"))?;
if !matches!(target.scheme(), "http" | "https") {
return Err("gRPC code-mode host URL must use http or https".to_string());
}
if !target.username().is_empty() || target.password().is_some() {
return Err(
"gRPC code-mode host URL must not include credentials".to_string(),
);
}
if target.path() != "/"
|| target.query().is_some()
|| target.fragment().is_some()
{
return Err("gRPC code-mode host URL must not include a path, query, or fragment".to_string());
}
let origin: Uri = endpoint
.parse()
.map_err(|error| format!("invalid gRPC code-mode host origin: {error}"))?;
let endpoint = endpoint.clone();
let http_client_factory = http_client_factory.clone();
let client = tokio::task::spawn_blocking(move || {
http_client_factory
.build_reqwest_client(
reqwest::Client::builder()
.http2_prior_knowledge()
.redirect(reqwest::redirect::Policy::none()),
&endpoint,
ClientRouteClass::Other,
)
.map_err(|error| {
format!(
"failed to configure gRPC code-mode host transport: {error}"
)
})
})
.await
.map_err(|error| {
format!("gRPC code-mode host transport task failed: {error}")
})??;
let transport = service_fn(move |request: Request<Body>| {
let client = client.clone();
async move {
let request = request.map(|body| {
reqwest::Body::wrap_stream(body.into_data_stream())
});
let request =
reqwest::Request::try_from(request).map_err(io::Error::other)?;
let response: Response<reqwest::Body> =
client.execute(request).await.map_err(io::Error::other)?.into();
Ok::<_, io::Error>(response.map(Body::new))
}
});
CodeModeHostClient::with_origin(BoxCloneSyncService::new(transport), origin)
}
TransportEndpoint::Connected(channel) => {
let transport = channel.clone().map_err(io::Error::other);
CodeModeHostClient::new(BoxCloneSyncService::new(transport))
}
};
Ok(client
.max_decoding_message_size(MAX_FRAME_BYTES)
.max_encoding_message_size(MAX_FRAME_BYTES))
})
.await
.cloned()
}
}
|