1use crate::{
6 DEFAULT_USER_AGENT, REQUEST_TIMEOUT,
7 provider::{
8 mpp::transport::{LazyMppHttpTransport, lazy_mpp_ws_connect},
9 redact_url,
10 },
11};
12use alloy_json_rpc::{RequestPacket, ResponsePacket};
13use alloy_pubsub::{PubSubConnect, PubSubFrontend};
14use alloy_rpc_types_engine::{Claims, JwtSecret};
15use alloy_transport::{
16 Authorization, BoxTransport, TransportError, TransportErrorKind, TransportFut,
17 utils::guess_local_url,
18};
19use alloy_transport_ipc::IpcConnect;
20use alloy_transport_ws::WsConnect;
21use regex::{Captures, Regex};
22use reqwest::header::{HeaderName, HeaderValue};
23use std::{
24 error::Error as StdError,
25 fmt,
26 path::PathBuf,
27 str::FromStr,
28 sync::{Arc, LazyLock},
29};
30use thiserror::Error;
31use tokio::{
32 runtime::{Handle, Id},
33 sync::RwLock,
34};
35use tower::Service;
36use url::Url;
37
38const KNOWN_MPP_HOSTS: &[&str] = &[".mpp.tempo.xyz", ".mpp.moderato.tempo.xyz"];
43
44static HTTP_URL_RE: LazyLock<Regex> =
45 LazyLock::new(|| Regex::new(r#"(?i)https?://[^\s<>"']+"#).expect("valid URL regex"));
46
47#[derive(Clone, Debug)]
50pub enum InnerTransport {
51 Http(LazyMppHttpTransport),
57 Ws(PubSubFrontend),
59 Ipc(PubSubFrontend),
61}
62
63#[derive(Error, Debug)]
65pub enum RuntimeTransportError {
66 #[error("Internal transport error: {0} with {1}")]
68 TransportError(TransportError, String),
69
70 #[error("URL scheme is not supported: {0}")]
72 BadScheme(String),
73
74 #[error("Invalid HTTP header: {0}")]
76 BadHeader(String),
77
78 #[error("Invalid IPC file path: {0}")]
80 BadPath(String),
81
82 #[error(transparent)]
84 HttpConstructionError(#[from] reqwest::Error),
85
86 #[error("Invalid JWT: {0}")]
88 InvalidJwt(String),
89}
90
91#[derive(Clone, Debug)]
99pub struct RuntimeTransport {
100 inner: Arc<RwLock<Option<(Id, InnerTransport)>>>,
102 url: Url,
104 headers: Vec<String>,
106 jwt: Option<String>,
108 timeout: std::time::Duration,
110 accept_invalid_certs: bool,
112 no_proxy: bool,
114}
115
116#[derive(Debug)]
118pub struct RuntimeTransportBuilder {
119 url: Url,
120 headers: Vec<String>,
121 jwt: Option<String>,
122 timeout: std::time::Duration,
123 accept_invalid_certs: bool,
124 no_proxy: bool,
125}
126
127impl RuntimeTransportBuilder {
128 pub const fn new(url: Url) -> Self {
130 Self {
131 url,
132 headers: vec![],
133 jwt: None,
134 timeout: REQUEST_TIMEOUT,
135 accept_invalid_certs: false,
136 no_proxy: false,
137 }
138 }
139
140 pub fn with_headers(mut self, headers: Vec<String>) -> Self {
142 self.headers = headers;
143 self
144 }
145
146 pub fn with_jwt(mut self, jwt: Option<String>) -> Self {
148 self.jwt = jwt;
149 self
150 }
151
152 pub const fn with_timeout(mut self, timeout: std::time::Duration) -> Self {
154 self.timeout = timeout;
155 self
156 }
157
158 pub const fn accept_invalid_certs(mut self, accept_invalid_certs: bool) -> Self {
160 self.accept_invalid_certs = accept_invalid_certs;
161 self
162 }
163
164 pub const fn no_proxy(mut self, no_proxy: bool) -> Self {
169 self.no_proxy = no_proxy;
170 self
171 }
172
173 pub fn build(self) -> RuntimeTransport {
176 RuntimeTransport {
177 inner: Arc::new(RwLock::new(None)),
178 url: self.url,
179 headers: self.headers,
180 jwt: self.jwt,
181 timeout: self.timeout,
182 accept_invalid_certs: self.accept_invalid_certs,
183 no_proxy: self.no_proxy,
184 }
185 }
186}
187
188impl fmt::Display for RuntimeTransport {
189 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
190 write!(f, "RuntimeTransport {}", redact_url(self.url.as_str()))
191 }
192}
193
194impl RuntimeTransport {
195 pub async fn connect(&self) -> Result<InnerTransport, RuntimeTransportError> {
197 match self.url.scheme() {
198 "http" | "https" => self.connect_http(),
199 "ws" | "wss" => self.connect_ws().await,
200 "file" => self.connect_ipc().await,
201 _ => Err(RuntimeTransportError::BadScheme(self.url.scheme().to_string())),
202 }
203 }
204
205 fn reqwest_headers(&self) -> Result<reqwest::header::HeaderMap, RuntimeTransportError> {
206 let mut headers = reqwest::header::HeaderMap::new();
207
208 if let Some(jwt) = self.jwt.clone() {
210 let auth =
211 build_auth(jwt).map_err(|e| RuntimeTransportError::InvalidJwt(e.to_string()))?;
212
213 let mut auth_value: HeaderValue =
214 HeaderValue::from_str(&auth.to_string()).expect("Header should be valid string");
215 auth_value.set_sensitive(true);
216
217 headers.insert(reqwest::header::AUTHORIZATION, auth_value);
218 };
219
220 for header in &self.headers {
222 let make_err = || RuntimeTransportError::BadHeader(header.clone());
223
224 let (key, val) = header.split_once(':').ok_or_else(make_err)?;
225
226 headers.insert(
227 HeaderName::from_str(key.trim()).map_err(|_| make_err())?,
228 HeaderValue::from_str(val.trim()).map_err(|_| make_err())?,
229 );
230 }
231
232 if !headers.contains_key(reqwest::header::USER_AGENT) {
233 headers.insert(
234 reqwest::header::USER_AGENT,
235 HeaderValue::from_str(DEFAULT_USER_AGENT)
236 .expect("User-Agent should be valid string"),
237 );
238 }
239
240 if !headers.contains_key(HeaderName::from_static("x-api-key"))
243 && let Ok(api_key) = std::env::var("MPP_API_KEY")
244 {
245 let api_key = api_key.trim();
246 if !api_key.is_empty() {
247 let mut value = HeaderValue::from_str(api_key)
248 .map_err(|_| RuntimeTransportError::BadHeader("MPP_API_KEY".to_string()))?;
249 value.set_sensitive(true);
250 headers.insert(HeaderName::from_static("x-api-key"), value);
251 }
252 }
253
254 Ok(headers)
255 }
256
257 fn reqwest_client_with_headers(
258 &self,
259 headers: reqwest::header::HeaderMap,
260 ) -> Result<reqwest::Client, RuntimeTransportError> {
261 let mut client_builder = reqwest::Client::builder()
262 .timeout(self.timeout)
263 .danger_accept_invalid_certs(self.accept_invalid_certs);
264
265 if self.no_proxy || guess_local_url(self.url.as_str()) {
269 client_builder = client_builder.no_proxy();
270 }
271
272 client_builder = client_builder.default_headers(headers);
273
274 Ok(client_builder.build()?)
275 }
276
277 pub fn reqwest_client(&self) -> Result<reqwest::Client, RuntimeTransportError> {
279 self.reqwest_client_with_headers(self.reqwest_headers()?)
280 }
281
282 fn connect_http(&self) -> Result<InnerTransport, RuntimeTransportError> {
284 let headers = self.reqwest_headers()?;
285 let client = self.reqwest_client_with_headers(headers.clone())?;
286 Ok(InnerTransport::Http(LazyMppHttpTransport::lazy(client, self.url.clone(), headers)))
287 }
288
289 async fn connect_ws(&self) -> Result<InnerTransport, RuntimeTransportError> {
295 let auth = self.jwt.as_ref().and_then(|jwt| build_auth(jwt.clone()).ok());
296
297 let service = if is_known_mpp_endpoint(&self.url) {
298 let mut ws = lazy_mpp_ws_connect(&self.url);
299 if let Some(auth) = auth {
300 ws = ws.with_auth(auth);
301 }
302 ws.into_service().await.map_err(|e| {
303 RuntimeTransportError::TransportError(e, redact_url(self.url.as_str()))
304 })?
305 } else {
306 let mut ws = WsConnect::new(self.url.to_string());
307 if let Some(auth) = auth {
308 ws = ws.with_auth(auth);
309 }
310 ws.into_service().await.map_err(|e| {
311 RuntimeTransportError::TransportError(e, redact_url(self.url.as_str()))
312 })?
313 };
314
315 Ok(InnerTransport::Ws(service))
316 }
317
318 async fn connect_ipc(&self) -> Result<InnerTransport, RuntimeTransportError> {
320 let path = url_to_file_path(&self.url)
321 .map_err(|_| RuntimeTransportError::BadPath(self.url.to_string()))?;
322 let ipc_connector = IpcConnect::new(path.clone());
323 let ipc = ipc_connector.into_service().await.map_err(|e| {
324 RuntimeTransportError::TransportError(e, path.clone().display().to_string())
325 })?;
326 Ok(InnerTransport::Ipc(ipc))
327 }
328
329 pub fn request(&self, req: RequestPacket) -> TransportFut<'static> {
337 let this = self.clone();
338 Box::pin(async move {
339 let runtime = Handle::current().id();
340 let inner = this.inner.read().await;
341 let transport = if let Some((id, transport)) = &*inner
342 && *id == runtime
343 {
344 transport.clone()
345 } else {
346 drop(inner);
347 let mut inner = this.inner.write().await;
348 if inner.as_ref().is_none_or(|(id, _)| *id != runtime) {
351 *inner =
352 Some((runtime, this.connect().await.map_err(TransportErrorKind::custom)?));
353 }
354 inner.as_ref().expect("transport is connected").1.clone()
355 };
356
357 match transport {
358 InnerTransport::Http(mut http) => http
359 .call(req)
360 .await
361 .map_err(|error| redact_http_transport_error(error, &this.url)),
362 InnerTransport::Ws(mut ws) => ws.call(req).await,
363 InnerTransport::Ipc(mut ipc) => ipc.call(req).await,
364 }
365 })
366 }
367
368 pub fn boxed(self) -> BoxTransport
370 where
371 Self: Sized + Clone + Send + Sync + 'static,
372 {
373 BoxTransport::new(self)
374 }
375}
376
377fn is_known_mpp_endpoint(url: &Url) -> bool {
379 url.host_str().is_some_and(|host| KNOWN_MPP_HOSTS.iter().any(|suffix| host.ends_with(suffix)))
380}
381
382fn redact_http_transport_error(error: TransportError, endpoint: &Url) -> TransportError {
383 let alloy_json_rpc::RpcError::Transport(TransportErrorKind::Custom(source)) = &error else {
384 return error;
385 };
386 let safe_endpoint = redact_url(endpoint.as_str());
387
388 let mut message = String::new();
389 let mut error: Option<&(dyn StdError + 'static)> = Some(source.as_ref());
390 while let Some(source) = error {
391 if !message.is_empty() {
392 message.push_str(": ");
393 }
394 message.push_str(&source.to_string());
395 error = source.source();
396 }
397 let message = HTTP_URL_RE.replace_all(&message, |captures: &Captures<'_>| {
398 let candidate = &captures[0];
399 let Ok(url) = Url::parse(candidate) else { return candidate.to_owned() };
400 if url.host() == endpoint.host()
401 && url.port_or_known_default() == endpoint.port_or_known_default()
402 {
403 safe_endpoint.clone()
404 } else {
405 candidate.to_owned()
406 }
407 });
408 TransportErrorKind::custom_str(&message)
409}
410
411impl tower::Service<RequestPacket> for RuntimeTransport {
412 type Response = ResponsePacket;
413 type Error = TransportError;
414 type Future = TransportFut<'static>;
415
416 #[inline]
417 fn poll_ready(
418 &mut self,
419 _cx: &mut std::task::Context<'_>,
420 ) -> std::task::Poll<Result<(), Self::Error>> {
421 std::task::Poll::Ready(Ok(()))
422 }
423
424 #[inline]
425 fn call(&mut self, req: RequestPacket) -> Self::Future {
426 self.request(req)
427 }
428}
429
430impl tower::Service<RequestPacket> for &RuntimeTransport {
431 type Response = ResponsePacket;
432 type Error = TransportError;
433 type Future = TransportFut<'static>;
434
435 #[inline]
436 fn poll_ready(
437 &mut self,
438 _cx: &mut std::task::Context<'_>,
439 ) -> std::task::Poll<Result<(), Self::Error>> {
440 std::task::Poll::Ready(Ok(()))
441 }
442
443 #[inline]
444 fn call(&mut self, req: RequestPacket) -> Self::Future {
445 self.request(req)
446 }
447}
448
449fn build_auth(jwt: String) -> eyre::Result<Authorization> {
450 let secret = JwtSecret::from_hex(jwt)?;
452 let claims = Claims::default();
453 let token = secret.encode(&claims)?;
454
455 let auth = Authorization::Bearer(token);
456
457 Ok(auth)
458}
459
460#[cfg(windows)]
461fn url_to_file_path(url: &Url) -> Result<PathBuf, ()> {
462 const PREFIX: &str = "file:///pipe/";
463
464 let url_str = url.as_str();
465
466 if let Some(pipe_name) = url_str.strip_prefix(PREFIX) {
467 let pipe_path = format!(r"\\.\pipe\{pipe_name}");
468 return Ok(PathBuf::from(pipe_path));
469 }
470
471 url.to_file_path()
472}
473
474#[cfg(not(windows))]
475fn url_to_file_path(url: &Url) -> Result<PathBuf, ()> {
476 url.to_file_path()
477}
478
479#[cfg(test)]
480mod tests {
481 use super::*;
482 use reqwest::header::HeaderMap;
483 use std::io;
484
485 #[derive(Debug, Error)]
486 #[error("request to https://example.com/private-api-key failed")]
487 struct ProviderError {
488 #[source]
489 source: io::Error,
490 }
491
492 #[test]
493 fn http_transport_errors_preserve_provider_guidance() {
494 let endpoint =
495 Url::parse("https://user:password@example.com/private-api-key?token=secret").unwrap();
496 let error = TransportErrorKind::custom(ProviderError {
497 source: io::Error::other(
498 "Authorize an access key with:\n cast tempo login --no-browser",
499 ),
500 });
501
502 let report = redact_http_transport_error(error, &endpoint).to_string();
503
504 assert!(report.contains("https://example.com/"));
505 assert!(report.contains("cast tempo login --no-browser"));
506 assert!(!report.contains("password"));
507 assert!(!report.contains("private-api-key"));
508 assert!(!report.contains("secret"));
509 }
510
511 #[test]
512 fn http_transport_errors_redact_endpoint_paths() {
513 let endpoint =
514 Url::parse("https://user:password@example.com/private-api-key?token=secret").unwrap();
515 let error = TransportErrorKind::custom_str(concat!(
516 "request to https://example.com/private-api-key failed: connection refused\n\n",
517 "Authorize an access key with:\n cast tempo login"
518 ));
519
520 let error = redact_http_transport_error(error, &endpoint);
521 let report = error.to_string();
522
523 assert!(report.contains("https://example.com/"));
524 assert!(!report.contains("password"));
525 assert!(!report.contains("private-api-key"));
526 assert!(!report.contains("secret"));
527 assert!(report.to_lowercase().contains("connection refused"));
528 assert!(report.contains("cast tempo login"));
529 }
530
531 #[test]
532 fn http_transport_errors_redact_normalized_endpoint_variants() {
533 let endpoint =
534 Url::parse("https://user:password@example.com/private-api-key?token=secret").unwrap();
535 let error = TransportErrorKind::custom_str(
536 "request to https://USER:normalized@example.com:443/different%2Fpath?key=other failed",
537 );
538
539 let report = redact_http_transport_error(error, &endpoint).to_string();
540
541 assert!(report.contains("https://example.com/"));
542 assert!(!report.contains("normalized"));
543 assert!(!report.contains("different"));
544 assert!(!report.contains("other"));
545 }
546
547 #[tokio::test]
548 async fn websocket_error_redacts_url_credentials() {
549 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
550 let address = listener.local_addr().unwrap();
551 drop(listener);
552 let url = Url::parse(&format!(
553 "ws://user:password@{address}/private-api-key?token=secret#fragment"
554 ))
555 .unwrap();
556 let transport = RuntimeTransportBuilder::new(url).build();
557
558 let error = transport.connect_ws().await.unwrap_err().to_string();
559
560 assert!(error.contains(&format!("ws://{address}/")));
561 assert!(!error.contains("user"));
562 assert!(!error.contains("password"));
563 assert!(!error.contains("private-api-key"));
564 assert!(!error.contains("secret"));
565 }
566
567 #[tokio::test]
568 async fn test_user_agent_header() {
569 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
570 let url = Url::parse(&format!("http://{}", listener.local_addr().unwrap())).unwrap();
571
572 let http_handler = axum::routing::get(|actual_headers: HeaderMap| {
573 let user_agent = HeaderName::from_str("User-Agent").unwrap();
574 assert_eq!(actual_headers[user_agent], HeaderValue::from_str("test-agent").unwrap());
575
576 async { "" }
577 });
578
579 let server_task = tokio::spawn(async move {
580 axum::serve(listener, http_handler.into_make_service()).await.unwrap()
581 });
582
583 let transport = RuntimeTransportBuilder::new(url.clone())
584 .with_headers(vec!["User-Agent: test-agent".to_string()])
585 .build();
586 let inner = transport.connect_http().unwrap();
587
588 match inner {
589 InnerTransport::Http(http) => {
590 let _ = http.client().get(url).send().await.unwrap();
591
592 }
594 _ => unreachable!(),
595 }
596
597 server_task.abort();
598 }
599}