1use crate::error::Result;
4use crate::shared::sse_parser::SseParser;
5use crate::shared::{Transport, TransportMessage};
6use async_trait::async_trait;
7use bytes::Bytes;
8use http_body_util::{BodyExt, Full, LengthLimitError, Limited};
9use hyper::{Method, Request, StatusCode};
10use hyper_util::client::legacy::Client;
11use hyper_util::rt::TokioExecutor;
12use parking_lot::RwLock;
13use std::sync::Arc;
14use std::time::Duration;
15#[cfg(not(target_arch = "wasm32"))]
16use tokio::sync::mpsc;
17#[cfg(not(target_arch = "wasm32"))]
18use tokio::sync::Mutex as AsyncMutex;
19#[cfg(not(target_arch = "wasm32"))]
20use tokio::time::timeout;
21use tracing::{debug, error, info, warn};
22use url::Url;
23
24#[derive(Debug, Clone)]
26pub struct HttpConfig {
27 pub base_url: Url,
29 pub sse_endpoint: Option<String>,
31 pub timeout: Duration,
33 pub headers: Vec<(String, String)>,
35 pub enable_pooling: bool,
37 pub max_idle_per_host: usize,
39}
40
41impl Default for HttpConfig {
42 fn default() -> Self {
43 Self {
44 base_url: "http://localhost:8080".parse().expect("Valid default URL"),
45 sse_endpoint: Some("/events".to_string()),
46 timeout: Duration::from_secs(30),
47 headers: vec![],
48 enable_pooling: true,
49 max_idle_per_host: 10,
50 }
51 }
52}
53
54pub struct HttpTransport {
56 config: HttpConfig,
57 client: Client<hyper_util::client::legacy::connect::HttpConnector, Full<Bytes>>,
58 message_queue: Arc<AsyncMutex<mpsc::Receiver<TransportMessage>>>,
59 message_tx: mpsc::Sender<TransportMessage>,
60 connected: Arc<RwLock<bool>>,
61 sse_buffered_bytes: usize,
72 max_collected_body_bytes: usize,
80}
81
82impl std::fmt::Debug for HttpTransport {
83 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
84 f.debug_struct("HttpTransport")
85 .field("config", &self.config)
86 .field("connected", &self.connected)
87 .field("sse_buffered_bytes", &self.sse_buffered_bytes)
88 .field("max_collected_body_bytes", &self.max_collected_body_bytes)
89 .finish_non_exhaustive()
90 }
91}
92
93pub use crate::shared::http_constants::DEFAULT_HTTP_SSE_BUFFERED_BYTES;
100
101pub const DEFAULT_HTTP_COLLECTED_BODY_BYTES: usize = 16 * 1024 * 1024;
126
127fn sse_reader_parser(sse_buffered_bytes: usize) -> SseParser {
135 SseParser::with_max_buffer_size(sse_buffered_bytes)
136}
137
138fn report_sse_overflow(parser: &SseParser) -> bool {
161 if !parser.overflowed() {
162 return false;
163 }
164 error!(
165 "an SSE chunk pushed the buffered stream state past the {}-byte parser \
166 bound; the buffered bytes were discarded, so the stream is corrupt and \
167 the connection is being closed",
168 parser.max_buffer_size()
169 );
170 true
171}
172
173impl HttpTransport {
174 pub fn new(config: HttpConfig) -> Self {
176 let connector = hyper_util::client::legacy::connect::HttpConnector::new();
177 let client = Client::builder(TokioExecutor::new())
178 .pool_idle_timeout(Duration::from_secs(30))
179 .pool_max_idle_per_host(config.max_idle_per_host)
180 .build(connector);
181
182 let (tx, rx) = mpsc::channel(100);
183
184 Self {
185 config,
186 client,
187 message_queue: Arc::new(AsyncMutex::new(rx)),
188 message_tx: tx,
189 connected: Arc::new(RwLock::new(false)),
190 sse_buffered_bytes: DEFAULT_HTTP_SSE_BUFFERED_BYTES,
191 max_collected_body_bytes: DEFAULT_HTTP_COLLECTED_BODY_BYTES,
192 }
193 }
194
195 pub fn with_url(url: impl Into<Url>) -> Result<Self> {
197 Ok(Self::new(HttpConfig {
198 base_url: url.into(),
199 ..Default::default()
200 }))
201 }
202
203 #[must_use]
228 pub fn with_sse_buffered_bytes(mut self, sse_buffered_bytes: usize) -> Self {
229 self.sse_buffered_bytes = sse_buffered_bytes;
230 self
231 }
232
233 #[must_use]
252 pub fn with_max_collected_body_bytes(mut self, max_collected_body_bytes: usize) -> Self {
253 self.max_collected_body_bytes = max_collected_body_bytes;
254 self
255 }
256
257 async fn collect_body_within_cap(
267 response: hyper::Response<hyper::body::Incoming>,
268 max_bytes: usize,
269 ) -> Result<Bytes> {
270 let declared = response
271 .headers()
272 .get(hyper::header::CONTENT_LENGTH)
273 .and_then(|value| value.to_str().ok())
274 .and_then(|value| value.parse::<usize>().ok());
275 if let Some(declared) = declared {
276 if declared > max_bytes {
277 return Err(crate::error::Error::Transport(
278 crate::error::TransportError::Request(format!(
279 "response body declares Content-Length {declared}, over this transport's \
280 {max_bytes}-byte collected-body cap (DEFAULT_HTTP_COLLECTED_BODY_BYTES); \
281 raise it with HttpTransport::with_max_collected_body_bytes"
282 )),
283 ));
284 }
285 }
286 match Limited::new(response.into_body(), max_bytes)
287 .collect()
288 .await
289 {
290 Ok(collected) => Ok(collected.to_bytes()),
291 Err(error) if error.is::<LengthLimitError>() => Err(crate::error::Error::Transport(
292 crate::error::TransportError::Request(format!(
293 "response body delivered more than this transport's {max_bytes}-byte \
294 collected-body cap (Content-Length absent or understated); raise it with \
295 HttpTransport::with_max_collected_body_bytes"
296 )),
297 )),
298 Err(error) => Err(crate::error::Error::Transport(
299 crate::error::TransportError::Request(error.to_string()),
300 )),
301 }
302 }
303
304 pub async fn connect_sse(&self) -> Result<()> {
306 if let Some(sse_path) = &self.config.sse_endpoint {
307 let sse_url = self
308 .config
309 .base_url
310 .join(sse_path)
311 .map_err(|e| crate::error::TransportError::InvalidMessage(e.to_string()))?;
312 info!("Connecting to SSE endpoint: {}", sse_url);
313
314 let req = Request::builder()
315 .method(Method::GET)
316 .uri(sse_url.as_str())
317 .header("Accept", "text/event-stream")
318 .header("Cache-Control", "no-cache")
319 .body(Full::new(Bytes::new()))
320 .map_err(|e| crate::error::TransportError::InvalidMessage(e.to_string()))?;
321
322 let response = self
323 .client
324 .request(req)
325 .await
326 .map_err(|e| crate::error::TransportError::InvalidMessage(e.to_string()))?;
327
328 if response.status() != StatusCode::OK {
329 return Err(crate::error::Error::Transport(
330 crate::error::TransportError::InvalidMessage(format!(
331 "SSE connection failed with status: {}",
332 response.status()
333 )),
334 ));
335 }
336
337 let message_tx = self.message_tx.clone();
339 let connected = self.connected.clone();
340 let sse_buffered_bytes = self.sse_buffered_bytes;
341
342 tokio::spawn(async move {
343 *connected.write() = true;
344
345 let mut body = response.into_body();
346 let mut sse_parser = sse_reader_parser(sse_buffered_bytes);
347 let mut undecoded: Vec<u8> = Vec::new();
353
354 while let Some(chunk) = body.frame().await {
355 match chunk {
356 Ok(frame) => {
357 if let Some(data) = frame.data_ref() {
358 undecoded.extend_from_slice(data);
359 let text =
360 crate::shared::sse_parser::take_utf8_prefix(&mut undecoded);
361 let events = sse_parser.feed(&text);
362
363 let overflowed = report_sse_overflow(&sse_parser);
372
373 for event in events {
374 match crate::shared::stdio::StdioTransport::parse_message(
376 event.data.as_bytes(),
377 ) {
378 Ok(msg) => {
379 if message_tx.send(msg).await.is_err() {
380 error!("Failed to send SSE message");
381 break;
382 }
383 },
384 Err(e) => {
385 error!("Failed to parse SSE message: {}", e);
386 },
387 }
388 }
389
390 if overflowed {
391 break;
396 }
397 }
398 },
399 Err(e) => {
400 error!("SSE stream error: {}", e);
401 break;
402 },
403 }
404 }
405
406 *connected.write() = false;
407 warn!("SSE connection closed");
408 });
409 } else {
410 *self.connected.write() = true;
412 }
413 Ok(())
414 }
415
416 async fn send_request(&self, message: &TransportMessage) -> Result<()> {
417 let json_bytes = crate::shared::stdio::StdioTransport::serialize_message(message)?;
418 let json = String::from_utf8(json_bytes).map_err(|e| {
419 crate::error::Error::Transport(crate::error::TransportError::InvalidMessage(format!(
420 "Invalid UTF-8: {}",
421 e
422 )))
423 })?;
424
425 let req = Request::builder()
426 .method(Method::POST)
427 .uri(self.config.base_url.as_str())
428 .header("Content-Type", "application/json")
429 .body(Full::new(Bytes::from(json)))
430 .map_err(|e| crate::error::TransportError::InvalidMessage(e.to_string()))?;
431
432 let response = timeout(self.config.timeout, self.client.request(req))
433 .await
434 .map_err(|_| crate::error::Error::Timeout(self.config.timeout.as_secs() * 1000))?
435 .map_err(|e| {
436 crate::error::Error::Transport(crate::error::TransportError::InvalidMessage(
437 e.to_string(),
438 ))
439 })?;
440
441 if response.status() != StatusCode::OK {
442 return Err(crate::error::Error::Transport(
443 crate::error::TransportError::InvalidMessage(format!(
444 "HTTP request failed with status: {}",
445 response.status()
446 )),
447 ));
448 }
449
450 let body_bytes = Self::collect_body_within_cap(response, self.max_collected_body_bytes)
458 .await
459 .map_err(|e| {
460 crate::error::Error::Transport(crate::error::TransportError::InvalidMessage(
461 e.to_string(),
462 ))
463 })?;
464 let response_msg = crate::shared::stdio::StdioTransport::parse_message(&body_bytes)?;
465
466 self.message_tx.send(response_msg).await.map_err(|_| {
468 crate::error::Error::Transport(crate::error::TransportError::ConnectionClosed)
469 })?;
470
471 Ok(())
472 }
473}
474
475#[async_trait]
476impl Transport for HttpTransport {
477 async fn send(&mut self, message: TransportMessage) -> Result<()> {
478 debug!("Sending HTTP message: {:?}", message);
479 self.send_request(&message).await
480 }
481
482 async fn receive(&mut self) -> Result<TransportMessage> {
483 let mut rx = self.message_queue.lock().await;
484 rx.recv().await.ok_or_else(|| {
485 crate::error::Error::Transport(crate::error::TransportError::ConnectionClosed)
486 })
487 }
488
489 async fn close(&mut self) -> Result<()> {
490 *self.connected.write() = false;
491 info!("HTTP transport closed");
492 Ok(())
493 }
494
495 fn is_connected(&self) -> bool {
496 *self.connected.read()
497 }
498}
499
500#[cfg(test)]
501mod tests {
502 use super::*;
503 use crate::types::{ClientRequest, Request, RequestId};
504
505 #[test]
506 fn test_http_config_default() {
507 let config = HttpConfig::default();
508 assert!(config.enable_pooling);
509 assert_eq!(config.timeout, Duration::from_secs(30));
510 assert_eq!(config.sse_endpoint, Some("/events".to_string()));
511 assert_eq!(config.max_idle_per_host, 10);
512 assert_eq!(config.headers.len(), 0);
513 }
514
515 #[test]
516 fn test_http_config_custom() {
517 let config = HttpConfig {
518 base_url: "http://example.com:3000".parse().unwrap(),
519 sse_endpoint: None,
520 timeout: Duration::from_mins(1),
521 headers: vec![("X-Custom".to_string(), "value".to_string())],
522 enable_pooling: false,
523 max_idle_per_host: 5,
524 };
525 assert_eq!(config.base_url.as_str(), "http://example.com:3000/");
526 assert!(config.sse_endpoint.is_none());
527 assert_eq!(config.timeout, Duration::from_mins(1));
528 assert_eq!(config.headers.len(), 1);
529 assert!(!config.enable_pooling);
530 assert_eq!(config.max_idle_per_host, 5);
531 }
532
533 #[test]
534 fn test_http_transport_creation() {
535 let config = HttpConfig::default();
536 let transport = HttpTransport::new(config);
537 assert!(!transport.is_connected());
538 }
539
540 #[test]
541 fn test_http_transport_with_url() {
542 let transport =
543 HttpTransport::with_url("http://localhost:9000".parse::<Url>().unwrap()).unwrap();
544 assert!(!transport.is_connected());
545 assert_eq!(transport.config.base_url.as_str(), "http://localhost:9000/");
546 }
547
548 #[test]
549 fn test_http_transport_debug() {
550 let config = HttpConfig::default();
551 let transport = HttpTransport::new(config);
552 let debug_str = format!("{:?}", transport);
553 assert!(debug_str.contains("HttpTransport"));
554 assert!(debug_str.contains("config"));
555 assert!(debug_str.contains("connected"));
556 }
557
558 #[tokio::test]
559 async fn test_http_transport_close() {
560 let config = HttpConfig::default();
561 let mut transport = HttpTransport::new(config);
562
563 *transport.connected.write() = true;
565 assert!(transport.is_connected());
566
567 transport.close().await.unwrap();
569 assert!(!transport.is_connected());
570 }
571
572 #[tokio::test]
573 async fn test_connect_sse_no_endpoint() {
574 let config = HttpConfig {
575 base_url: "http://localhost:8080".parse().unwrap(),
576 sse_endpoint: None,
577 ..Default::default()
578 };
579 let transport = HttpTransport::new(config);
580
581 transport.connect_sse().await.unwrap();
583 assert!(transport.is_connected());
584 }
585
586 #[tokio::test]
587 async fn test_send_request_not_connected() {
588 let config = HttpConfig::default();
589 let mut transport = HttpTransport::new(config);
590
591 let message = TransportMessage::Request {
592 id: RequestId::from(1i64),
593 request: Request::Client(Box::new(ClientRequest::Ping)),
594 };
595
596 let result = transport.send(message).await;
598 assert!(result.is_err());
599 }
600
601 #[test]
602 fn test_http_config_with_headers() {
603 let config = HttpConfig {
604 base_url: "http://localhost:8080".parse().unwrap(),
605 headers: vec![
606 ("Authorization".to_string(), "Bearer token".to_string()),
607 ("X-API-Key".to_string(), "secret".to_string()),
608 ],
609 ..Default::default()
610 };
611 assert_eq!(config.headers.len(), 2);
612 assert_eq!(config.headers[0].0, "Authorization");
613 assert_eq!(config.headers[0].1, "Bearer token");
614 }
615
616 #[test]
617 fn test_http_config_clone() {
618 let config = HttpConfig::default();
619 let cloned = config.clone();
620 assert_eq!(config.base_url, cloned.base_url);
621 assert_eq!(config.timeout, cloned.timeout);
622 assert_eq!(config.enable_pooling, cloned.enable_pooling);
623 }
624
625 fn sse_frame_of_len(len: usize) -> String {
627 assert!(
632 len >= 8,
633 "an SSE frame cannot be shorter than its 8 bytes of framing (asked for {len})"
634 );
635 format!("data: {}\n\n", "A".repeat(len - 8))
636 }
637
638 #[test]
642 fn an_oversized_sse_line_ends_the_reader_task() {
643 let mut parser = SseParser::with_max_buffer_size(64);
644 assert!(
645 !report_sse_overflow(&parser),
646 "a fresh parser has lost nothing, so the task keeps reading"
647 );
648
649 assert!(
650 parser.feed(&"x".repeat(256)).is_empty(),
651 "an unterminated line completes no event"
652 );
653 assert!(
654 report_sse_overflow(&parser),
655 "the discarded bytes end the task instead of being silently swallowed"
656 );
657 }
658
659 #[test]
666 fn a_newline_carrying_flood_ends_the_reader_task_too() {
667 let mut parser = sse_reader_parser(64);
668 let mut ended = false;
669 for _ in 0..1_000 {
670 assert!(
671 parser.feed("data: AAAAAAAA\n").is_empty(),
672 "a `data:` line with no blank line after it completes no event"
673 );
674 if report_sse_overflow(&parser) {
675 ended = true;
676 break;
677 }
678 }
679 assert!(ended, "accumulated `data:` lines must end the reader task");
680 }
681
682 #[test]
690 fn connect_sse_uses_its_own_named_bound() {
691 let transport = HttpTransport::new(HttpConfig::default());
692 assert_eq!(
693 transport.sse_buffered_bytes, DEFAULT_HTTP_SSE_BUFFERED_BYTES,
694 "the transport defaults its ceiling from the named constant"
695 );
696
697 let mut parser = sse_reader_parser(transport.sse_buffered_bytes);
698 assert_eq!(parser.max_buffer_size(), DEFAULT_HTTP_SSE_BUFFERED_BYTES);
699 let _ = parser.feed(&"x".repeat(256));
700 assert!(
701 !report_sse_overflow(&parser),
702 "256 bytes is nowhere near the {DEFAULT_HTTP_SSE_BUFFERED_BYTES}-byte default"
703 );
704 }
705
706 #[test]
712 fn the_configured_ceiling_admits_up_to_and_including_itself() {
713 let ceiling = 256;
714
715 let mut under = sse_reader_parser(ceiling);
716 let events = under.feed(&sse_frame_of_len(ceiling - 1));
717 assert_eq!(events.len(), 1, "one byte under the ceiling parses");
718 assert!(!under.overflowed());
719
720 let mut exact = sse_reader_parser(ceiling);
721 let events = exact.feed(&sse_frame_of_len(ceiling));
722 assert_eq!(events.len(), 1, "exactly the ceiling parses");
723 assert!(!exact.overflowed(), "the comparison is `>`, not `>=`");
724
725 let mut over = sse_reader_parser(ceiling);
726 assert!(
727 over.feed(&sse_frame_of_len(ceiling + 1)).is_empty(),
728 "one byte over the ceiling is refused whole"
729 );
730 assert!(over.overflowed(), "and the refusal is observable");
731 assert!(report_sse_overflow(&over), "so the reader task ends");
732 }
733
734 #[test]
741 fn raising_the_ceiling_admits_a_payload_the_lower_one_refuses() {
742 let base = 256;
743 let payload = sse_frame_of_len(base + 1);
744
745 let mut low = sse_reader_parser(base);
746 assert!(low.feed(&payload).is_empty());
747 assert!(report_sse_overflow(&low), "refused at the lower ceiling");
748
749 let raised = HttpTransport::new(HttpConfig::default()).with_sse_buffered_bytes(base * 4);
750 assert_eq!(
751 raised.sse_buffered_bytes,
752 base * 4,
753 "the builder overrides the default"
754 );
755
756 let mut parser = sse_reader_parser(raised.sse_buffered_bytes);
757 let events = parser.feed(&payload);
758 assert_eq!(events.len(), 1, "the same bytes now parse");
759 assert!(!report_sse_overflow(&parser));
760 }
761
762 #[test]
770 fn base64_expansion_puts_a_12_to_16_binary_over_the_ceiling() {
771 use base64::Engine as _;
772
773 let raw_len = 12 * 1024;
774 let ceiling = 16 * 1024;
775
776 let encoded = base64::engine::general_purpose::STANDARD.encode(vec![0u8; raw_len]);
777
778 assert_eq!(
781 encoded.len(),
782 raw_len.div_ceil(3) * 4,
783 "base64 expands 3 raw bytes into 4"
784 );
785 assert_eq!(
786 encoded.len(),
787 ceiling,
788 "a '12 MiB' binary is ALREADY the whole '16 MiB' ceiling once encoded, \
789 with nothing left for JSON, the `data: ` prefix or the MIME type"
790 );
791
792 let frame = format!("data: {encoded}\n\n");
794 assert!(frame.len() > ceiling, "the envelope is what tips it");
795
796 let mut parser = sse_reader_parser(ceiling);
797 assert!(
798 parser.feed(&frame).is_empty(),
799 "so the payload is refused at a ceiling sized for its RAW bytes"
800 );
801 assert!(parser.overflowed());
802 }
803
804 #[tokio::test]
805 async fn test_message_queue_receive_closed() {
806 let config = HttpConfig::default();
807 let transport = HttpTransport::new(config);
808
809 let (_, rx) = mpsc::channel::<TransportMessage>(1);
811 let mut transport = HttpTransport {
812 config: transport.config,
813 client: transport.client,
814 message_queue: Arc::new(AsyncMutex::new(rx)),
815 message_tx: transport.message_tx,
816 connected: transport.connected,
817 sse_buffered_bytes: transport.sse_buffered_bytes,
818 max_collected_body_bytes: transport.max_collected_body_bytes,
819 };
820
821 let result = transport.receive().await;
823 assert!(result.is_err());
824 if let Err(crate::error::Error::Transport(e)) = result {
825 assert!(matches!(e, crate::error::TransportError::ConnectionClosed));
826 }
827 }
828}