1use crate::error::ProxyError as ProxyLifecycleError;
6use async_stream::try_stream;
7use axum::Json;
8use axum::body::Body;
9use axum::http::{HeaderMap, Response, StatusCode, header};
10use axum::response::IntoResponse;
11use bytes::Bytes;
12use futures_util::{FutureExt, Stream, StreamExt};
13use serde::{Deserialize, Serialize};
14use serde_json::Value;
15use std::env;
16use std::fmt;
17use std::future::Future;
18use std::sync::atomic::{AtomicUsize, Ordering};
19use std::time::Duration;
20use tokio::task::JoinHandle;
21
22pub fn run<F, Fut>(run_async: F) -> Result<(), ProxyLifecycleError>
24where
25 F: FnOnce() -> Fut,
26 Fut: Future<Output = Result<(), ProxyLifecycleError>>,
27{
28 let runtime = tokio::runtime::Builder::new_multi_thread()
29 .enable_all()
30 .build()
31 .map_err(|error| ProxyLifecycleError::Lifecycle {
32 message: format!("failed to create proxy tokio runtime: {error}"),
33 })?;
34 runtime.block_on(run_async())
35}
36
37#[derive(Serialize)]
39pub struct ProxyHealthcheckResponse {
40 pub ready: bool,
41 pub prefill_instances: usize,
42 pub decode_instances: usize,
43}
44
45pub(crate) fn healthcheck_response(
47 ready: bool,
48 prefill_instances: usize,
49 decode_instances: usize,
50) -> (StatusCode, Json<ProxyHealthcheckResponse>) {
51 let status = if ready {
52 StatusCode::OK
53 } else {
54 StatusCode::SERVICE_UNAVAILABLE
55 };
56 (
57 status,
58 Json(ProxyHealthcheckResponse {
59 ready,
60 prefill_instances,
61 decode_instances,
62 }),
63 )
64}
65
66pub(crate) fn require_endpoints(
68 proxy_name: &'static str,
69 prefill_is_empty: bool,
70 decode_is_empty: bool,
71) -> Result<(), ProxyLifecycleError> {
72 if prefill_is_empty {
73 return Err(ProxyLifecycleError::Invalid {
74 message: format!("{proxy_name} requires at least one prefill endpoint"),
75 });
76 }
77 if decode_is_empty {
78 return Err(ProxyLifecycleError::Invalid {
79 message: format!("{proxy_name} requires at least one decode endpoint"),
80 });
81 }
82 Ok(())
83}
84
85pub(crate) fn pooled_client(
87 proxy_name: &'static str,
88) -> Result<reqwest::Client, ProxyLifecycleError> {
89 build_pooled_client().map_err(|error| ProxyLifecycleError::Io {
90 message: format!("failed to create {proxy_name} HTTP client: {error}"),
91 })
92}
93
94pub(crate) async fn serve_router(
96 proxy_name: &'static str,
97 host: &str,
98 port: u16,
99 router: axum::Router,
100) -> Result<(), ProxyLifecycleError> {
101 let listener = tokio::net::TcpListener::bind((host, port))
102 .await
103 .map_err(|error| ProxyLifecycleError::Io {
104 message: format!("failed to bind {proxy_name} on {host}:{port}: {error}"),
105 })?;
106 axum::serve(listener, router)
107 .await
108 .map_err(|error| ProxyLifecycleError::Io {
109 message: format!("{proxy_name} server failed: {error}"),
110 })
111}
112
113pub(crate) async fn await_backends(client: reqwest::Client, urls: Vec<String>, path: &'static str) {
115 let waits = urls
116 .into_iter()
117 .map(|url| await_backend(client.clone(), url, path));
118 futures_util::future::join_all(waits).await;
119}
120
121async fn await_backend(client: reqwest::Client, url: String, path: &'static str) {
122 loop {
123 if client
124 .get(join_path(&url, path))
125 .send()
126 .await
127 .is_ok_and(|response| response.status().is_success())
128 {
129 return;
130 }
131 tokio::time::sleep(BACKEND_RETRY_INTERVAL).await;
132 }
133}
134
135pub(crate) const BACKEND_RETRY_INTERVAL: Duration = Duration::from_secs(1);
138
139pub(crate) fn fanout_target_urls<'a>(
142 prefill_urls: impl IntoIterator<Item = &'a str>,
143 decode_urls: impl IntoIterator<Item = &'a str>,
144) -> Vec<String> {
145 prefill_urls
146 .into_iter()
147 .chain(decode_urls)
148 .map(str::to_owned)
149 .collect()
150}
151
152pub(crate) fn upstream_response_builder(
155 response: &reqwest::Response,
156) -> Result<axum::http::response::Builder, ProxyHttpError> {
157 let mut builder = Response::builder().status(status_code(response.status())?);
158 if let Some(content_type) = response
159 .headers()
160 .get(reqwest::header::CONTENT_TYPE)
161 .and_then(|value| value.to_str().ok())
162 {
163 builder = builder.header(header::CONTENT_TYPE, content_type);
164 }
165 Ok(builder)
166}
167
168pub(crate) fn response_body(
170 builder: axum::http::response::Builder,
171 body: Body,
172) -> Result<Response<Body>, ProxyHttpError> {
173 builder.body(body).map_err(|error| {
174 ProxyHttpError::internal(format!("failed to build proxy response: {error}"))
175 })
176}
177
178pub async fn forward_response(
181 response: reqwest::Response,
182) -> Result<Response<Body>, ProxyHttpError> {
183 let builder = upstream_response_builder(&response)?;
184 let bytes = response
185 .bytes()
186 .await
187 .map_err(|error| ProxyHttpError::upstream("upstream response body read failed", error))?;
188 response_body(builder, Body::from(bytes))
189}
190
191pub async fn upstream_status_error(context: &str, response: reqwest::Response) -> ProxyHttpError {
194 let status = response.status();
195 let body = match response.text().await {
196 Ok(text) => text,
197 Err(error) => format!("<failed to read upstream error body: {error}>"),
198 };
199 ProxyHttpError::status(
200 StatusCode::BAD_GATEWAY,
201 format!("{context} returned HTTP {status}: {body}"),
202 )
203}
204
205pub fn outbound_authorization(headers: &HeaderMap) -> Option<String> {
208 headers
209 .get(header::AUTHORIZATION)
210 .and_then(|value| value.to_str().ok())
211 .map(str::to_owned)
212 .or_else(|| {
213 env::var("OPENAI_API_KEY")
214 .ok()
215 .map(|key| format!("Bearer {key}"))
216 })
217}
218
219pub fn join_path(base: &str, path: &str) -> String {
222 format!("{}{}", base.trim_end_matches('/'), path)
223}
224
225pub fn status_code(status: reqwest::StatusCode) -> Result<StatusCode, ProxyHttpError> {
227 StatusCode::from_u16(status.as_u16())
228 .map_err(|error| ProxyHttpError::internal(format!("invalid upstream status code: {error}")))
229}
230
231pub(crate) fn round_robin_index(cursor: &AtomicUsize, len: usize) -> usize {
235 cursor.fetch_add(1, Ordering::SeqCst) % len
236}
237
238pub(crate) async fn send_json_post(
244 client: reqwest::Client,
245 url: String,
246 body: &Value,
247 request_id: Option<&str>,
248 authorization: Option<&str>,
249 extra_headers: &[(&str, String)],
250 context: &'static str,
251) -> Result<reqwest::Response, ProxyHttpError> {
252 let response = send_json_post_status(
253 client,
254 url,
255 body,
256 request_id,
257 authorization,
258 extra_headers,
259 context,
260 )
261 .await?;
262 if !response.status().is_success() {
263 return Err(upstream_status_error(context, response).await);
264 }
265 Ok(response)
266}
267
268pub(crate) async fn send_json_post_status(
272 client: reqwest::Client,
273 url: String,
274 body: &Value,
275 request_id: Option<&str>,
276 authorization: Option<&str>,
277 extra_headers: &[(&str, String)],
278 context: &'static str,
279) -> Result<reqwest::Response, ProxyHttpError> {
280 let mut request = client.post(url).json(body);
281 if let Some(request_id) = request_id {
282 request = request.header("X-Request-Id", request_id);
283 }
284 for (name, value) in extra_headers {
287 request = request.header(*name, value);
288 }
289 if let Some(authorization) = authorization {
290 request = request.header(reqwest::header::AUTHORIZATION, authorization);
291 }
292 request
293 .send()
294 .await
295 .map_err(|error| ProxyHttpError::upstream(&format!("{context} failed"), error))
296}
297
298pub(crate) fn next_request_id(counter: &AtomicUsize) -> String {
301 let value = counter.fetch_add(1, Ordering::SeqCst);
302 format!("{}-{value}", std::process::id())
303}
304
305pub(crate) fn build_pooled_client() -> reqwest::Result<reqwest::Client> {
309 reqwest::Client::builder()
310 .pool_max_idle_per_host(usize::MAX)
311 .build()
312}
313
314#[derive(Debug, Deserialize, Serialize)]
316pub struct FanoutFailure {
317 pub url: String,
318 pub error: String,
319}
320
321#[derive(Debug, Deserialize, Serialize)]
325pub struct ResetPrefixCacheResponse {
326 pub successful: Vec<String>,
327 pub failed: Vec<FanoutFailure>,
328}
329
330#[derive(Debug, Deserialize, Serialize)]
334pub struct PrimePrefixCacheResponse {
335 pub targets: Vec<PrimePrefixCacheTarget>,
336}
337
338#[derive(Debug, Deserialize, Serialize)]
341pub struct PrimePrefixCacheTarget {
342 pub url: String,
343 pub rank: u32,
344 pub http_status: Option<u16>,
345 pub elapsed_ms: u64,
346 pub error: Option<String>,
347}
348
349pub(crate) struct PrimeFlowFailure {
352 pub http_status: Option<u16>,
353 pub error: String,
354}
355
356impl PrimeFlowFailure {
357 pub(crate) fn transport(error: ProxyHttpError) -> Self {
358 Self {
359 http_status: None,
360 error: error.to_string(),
361 }
362 }
363
364 pub(crate) fn status(status: u16, detail: String) -> Self {
365 Self {
366 http_status: Some(status),
367 error: detail,
368 }
369 }
370}
371
372pub(crate) async fn expect_2xx(
377 context: &'static str,
378 response: reqwest::Response,
379) -> Result<(u16, String), PrimeFlowFailure> {
380 let status = response.status().as_u16();
381 let text = response.text().await.map_err(|error| {
382 PrimeFlowFailure::transport(ProxyHttpError::upstream(
383 &format!("{context} response read failed"),
384 error,
385 ))
386 })?;
387 if !(200..300).contains(&status) {
388 return Err(PrimeFlowFailure::status(
389 status,
390 format!("{context} returned HTTP {status}: {text}"),
391 ));
392 }
393 Ok((status, text))
394}
395
396pub(crate) trait PrimeFanoutTarget {
399 fn url(&self) -> &str;
400 fn rank(&self) -> u32;
401}
402
403pub(crate) trait PrimeReplica {
406 fn url(&self) -> &str;
407 fn data_parallel_size(&self) -> u32;
408}
409
410pub(crate) struct RankedPrimeTarget<R> {
413 pub replica: R,
414 pub rank: u32,
415}
416
417impl<R: PrimeReplica> PrimeFanoutTarget for RankedPrimeTarget<R> {
418 fn url(&self) -> &str {
419 self.replica.url()
420 }
421
422 fn rank(&self) -> u32 {
423 self.rank
424 }
425}
426
427pub(crate) fn ranked_prime_targets<R: PrimeReplica + Clone>(
430 replicas: &[R],
431) -> Vec<RankedPrimeTarget<R>> {
432 let mut targets = Vec::new();
433 for replica in replicas {
434 for rank in 0..replica.data_parallel_size().max(1) {
435 targets.push(RankedPrimeTarget {
436 replica: replica.clone(),
437 rank,
438 });
439 }
440 }
441 targets
442}
443
444pub(crate) async fn run_sweep_fanout(
452 client: reqwest::Client,
453 operation: &'static str,
454 path: &'static str,
455 targets: Vec<String>,
456 authorization: Option<String>,
457) -> Response<Body> {
458 if targets.is_empty() {
459 return empty_fanout_failure(operation);
460 }
461 let attempts = targets
462 .into_iter()
463 .map(|url| sweep_target(client.clone(), operation, path, url, authorization.clone()));
464 let mut successful = Vec::new();
465 let mut failed = Vec::new();
466 for result in futures_util::future::join_all(attempts).await {
467 match result {
468 Ok(url) => successful.push(url),
469 Err(failure) => failed.push(failure),
470 }
471 }
472 let status = if failed.is_empty() {
473 StatusCode::OK
474 } else {
475 StatusCode::PARTIAL_CONTENT
476 };
477 (
478 status,
479 Json(ResetPrefixCacheResponse { successful, failed }),
480 )
481 .into_response()
482}
483
484async fn sweep_target(
485 client: reqwest::Client,
486 operation: &'static str,
487 path: &'static str,
488 url: String,
489 authorization: Option<String>,
490) -> Result<String, FanoutFailure> {
491 let endpoint = join_path(&url, path);
492 let mut request = client.post(endpoint);
493 if let Some(authorization) = authorization {
494 request = request.header(reqwest::header::AUTHORIZATION, authorization);
495 }
496 let response = request.send().await.map_err(|error| FanoutFailure {
497 url: url.clone(),
498 error: format!("{operation} request failed: {error}"),
499 })?;
500 if response.status().is_success() && response.status() != reqwest::StatusCode::PARTIAL_CONTENT {
503 Ok(url)
504 } else {
505 let status = response.status();
506 let detail = response
507 .text()
508 .await
509 .unwrap_or_else(|error| format!("failed to read response body: {error}"));
510 Err(FanoutFailure {
511 url,
512 error: format!("HTTP {status}: {detail}"),
513 })
514 }
515}
516
517pub(crate) async fn run_prime_fanout<T, F, Fut>(
524 operation: &'static str,
525 targets: Vec<T>,
526 mut execute: F,
527) -> Response<Body>
528where
529 T: PrimeFanoutTarget,
530 F: FnMut(T) -> Fut,
531 Fut: Future<Output = Result<u16, PrimeFlowFailure>>,
532{
533 if targets.is_empty() {
534 return empty_fanout_failure(operation);
535 }
536 let mut results = Vec::new();
537 for target in targets {
538 let url = target.url().to_owned();
539 let rank = target.rank();
540 let started = std::time::Instant::now();
541 let outcome = execute(target).await;
542 let elapsed_ms = u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX);
543 results.push(match outcome {
544 Ok(status) => PrimePrefixCacheTarget {
545 url,
546 rank,
547 http_status: Some(status),
548 elapsed_ms,
549 error: None,
550 },
551 Err(failure) => PrimePrefixCacheTarget {
552 url,
553 rank,
554 http_status: failure.http_status,
555 elapsed_ms,
556 error: Some(failure.error),
557 },
558 });
559 }
560 let status = if results.iter().all(|target| target.error.is_none()) {
561 StatusCode::OK
562 } else {
563 StatusCode::PARTIAL_CONTENT
564 };
565 (status, Json(PrimePrefixCacheResponse { targets: results })).into_response()
566}
567
568fn empty_fanout_failure(operation: &str) -> Response<Body> {
571 ProxyHttpError::status(
572 StatusCode::BAD_GATEWAY,
573 format!("{operation} fan-out has no targets: no prefill replica or data-parallel rank is available"),
574 )
575 .into_response()
576}
577
578#[derive(Debug)]
580pub struct ProxyHttpError {
581 status: StatusCode,
582 message: String,
583}
584
585impl ProxyHttpError {
586 pub fn status(status: StatusCode, message: impl Into<String>) -> Self {
587 Self {
588 status,
589 message: message.into(),
590 }
591 }
592
593 pub fn upstream(context: &str, error: reqwest::Error) -> Self {
594 Self::status(StatusCode::BAD_GATEWAY, format!("{context}: {error}"))
595 }
596
597 pub fn internal(message: impl Into<String>) -> Self {
598 Self::status(StatusCode::INTERNAL_SERVER_ERROR, message)
599 }
600}
601
602impl fmt::Display for ProxyHttpError {
603 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
604 write!(formatter, "{}", self.message)
605 }
606}
607
608impl std::error::Error for ProxyHttpError {}
609
610impl IntoResponse for ProxyHttpError {
611 fn into_response(self) -> axum::response::Response {
612 let body = Json(ProxyErrorResponse {
613 error: self.message,
614 });
615 (self.status, body).into_response()
616 }
617}
618
619#[derive(Serialize)]
620pub struct ProxyErrorResponse {
621 pub error: String,
622}
623
624#[derive(Clone, Copy, Debug)]
627pub(crate) enum OnClientDrop {
628 Abort,
631 Detach,
637}
638
639pub(crate) fn stream_decode_response(
647 response: reqwest::Response,
648 prefill_task: JoinHandle<Result<(), ProxyHttpError>>,
649 on_client_drop: OnClientDrop,
650) -> Result<Response<Body>, ProxyHttpError> {
651 let builder = upstream_response_builder(&response)?;
652 let stream = decode_response_stream(response.bytes_stream(), prefill_task, on_client_drop);
653 response_body(builder, Body::from_stream(stream))
654}
655
656pub(crate) fn stream_response(
658 response: reqwest::Response,
659) -> Result<Response<Body>, ProxyHttpError> {
660 let builder = upstream_response_builder(&response)?;
661 let stream = response
662 .bytes_stream()
663 .map(|chunk| chunk.map_err(|error| stream_error(format!("decode stream failed: {error}"))));
664 response_body(builder, Body::from_stream(stream))
665}
666
667pub(crate) fn decode_response_stream<S, E>(
674 decode_stream: S,
675 prefill_task: JoinHandle<Result<(), ProxyHttpError>>,
676 on_client_drop: OnClientDrop,
677) -> impl Stream<Item = std::result::Result<Bytes, std::io::Error>>
678where
679 S: Stream<Item = std::result::Result<Bytes, E>> + Unpin,
680 E: fmt::Display,
681{
682 let prefill_abort = prefill_task.abort_handle();
683 try_stream! {
684 let mut decode_stream = decode_stream;
685 let mut prefill_task = prefill_task;
686 let mut prefill_abort = match on_client_drop {
690 OnClientDrop::Abort => Some(AbortOnDrop::new(prefill_abort)),
691 OnClientDrop::Detach => None,
692 };
693 let mut prefill_done = false;
694 loop {
695 match next_stream_event(&mut prefill_task, &mut decode_stream, prefill_done).await {
696 StreamEvent::Prefill(prefill) => {
697 prefill_done = true;
698 match decode_stream.next().now_or_never() {
707 Some(Some(Ok(bytes))) => yield bytes,
708 Some(Some(Err(error))) => {
709 Err(stream_error(format!("decode stream failed: {error}")))?;
710 }
711 Some(None) | None => {}
712 }
713 prefill
714 .map_err(join_error)?
715 .map_err(|error| stream_error(error.to_string()))?;
716 if let Some(abort) = &mut prefill_abort {
717 abort.disarm();
718 }
719 }
720 StreamEvent::Decode(Some(Ok(bytes))) => yield bytes,
721 StreamEvent::Decode(Some(Err(error))) => {
722 Err(stream_error(format!("decode stream failed: {error}")))?;
723 }
724 StreamEvent::Decode(None) => break,
725 }
726 }
727 if !prefill_done {
728 prefill_task
729 .await
730 .map_err(join_error)?
731 .map_err(|error| stream_error(error.to_string()))?;
732 if let Some(abort) = &mut prefill_abort {
733 abort.disarm();
734 }
735 }
736 }
737}
738
739enum StreamEvent<E> {
740 Prefill(std::result::Result<Result<(), ProxyHttpError>, tokio::task::JoinError>),
741 Decode(Option<std::result::Result<Bytes, E>>),
742}
743
744async fn next_stream_event<S, E>(
745 prefill_task: &mut JoinHandle<Result<(), ProxyHttpError>>,
746 decode_stream: &mut S,
747 prefill_done: bool,
748) -> StreamEvent<E>
749where
750 S: Stream<Item = std::result::Result<Bytes, E>> + Unpin,
751{
752 if !prefill_done && prefill_task.is_finished() {
759 return StreamEvent::Prefill(prefill_task.await);
760 }
761 tokio::select! {
767 prefill = prefill_task, if !prefill_done => StreamEvent::Prefill(prefill),
768 chunk = decode_stream.next() => StreamEvent::Decode(chunk),
769 }
770}
771
772fn join_error(error: tokio::task::JoinError) -> std::io::Error {
773 stream_error(format!("prefill task failed: {error}"))
774}
775
776fn stream_error(message: String) -> std::io::Error {
777 std::io::Error::other(message)
778}
779
780struct AbortOnDrop {
782 handle: tokio::task::AbortHandle,
783 armed: bool,
784}
785
786impl AbortOnDrop {
787 fn new(handle: tokio::task::AbortHandle) -> Self {
788 Self {
789 handle,
790 armed: true,
791 }
792 }
793
794 fn disarm(&mut self) {
795 self.armed = false;
796 }
797}
798
799impl Drop for AbortOnDrop {
800 fn drop(&mut self) {
801 if self.armed {
802 self.handle.abort();
803 }
804 }
805}
806
807#[cfg(test)]
808mod tests {
809 use super::*;
810 use anyhow::{Context, Result};
811
812 #[test]
813 fn join_path_normalizes_single_trailing_slash() {
814 assert_eq!(
815 join_path("http://h:1/", "/v1/models"),
816 "http://h:1/v1/models"
817 );
818 assert_eq!(
819 join_path("http://h:1", "/v1/models"),
820 "http://h:1/v1/models"
821 );
822 }
823
824 #[test]
825 fn status_code_maps_reqwest_status() -> Result<()> {
826 let mapped = status_code(reqwest::StatusCode::OK)
827 .map_err(|error| anyhow::anyhow!(error.to_string()))?;
828 assert_eq!(mapped, StatusCode::OK);
829 Ok(())
830 }
831
832 #[test]
833 fn outbound_authorization_prefers_inbound_header() -> Result<()> {
834 let mut headers = HeaderMap::new();
835 headers.insert(header::AUTHORIZATION, "Bearer inbound".parse()?);
836 assert_eq!(
837 outbound_authorization(&headers),
838 Some("Bearer inbound".to_owned())
839 );
840 Ok(())
841 }
842
843 #[test]
844 fn proxy_error_internal_uses_500() {
845 let error = ProxyHttpError::internal("boom");
846 assert_eq!(error.status, StatusCode::INTERNAL_SERVER_ERROR);
847 assert_eq!(error.to_string(), "boom");
848 }
849
850 struct StaticPrimeTarget {
851 url: &'static str,
852 rank: u32,
853 }
854
855 impl PrimeFanoutTarget for StaticPrimeTarget {
856 fn url(&self) -> &str {
857 self.url
858 }
859
860 fn rank(&self) -> u32 {
861 self.rank
862 }
863 }
864
865 #[test]
868 fn prime_fanout_rejects_an_empty_target_set() -> Result<()> {
869 let runtime = proxy_test_runtime()?;
870 let response = runtime.block_on(run_prime_fanout(
871 "prefix cache conditioning",
872 Vec::<StaticPrimeTarget>::new(),
873 |_target| async { Ok::<u16, PrimeFlowFailure>(200) },
874 ));
875 assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
876 let body = runtime.block_on(axum::body::to_bytes(response.into_body(), usize::MAX))?;
877 let value: Value = serde_json::from_slice(&body)?;
878 assert!(
879 value["error"]
880 .as_str()
881 .is_some_and(|error| error.contains("no targets")),
882 "got {value}"
883 );
884 Ok(())
885 }
886
887 #[test]
890 fn sweep_fanout_rejects_an_empty_target_set() -> Result<()> {
891 let runtime = proxy_test_runtime()?;
892 let client = build_pooled_client().map_err(|error| anyhow::anyhow!(error.to_string()))?;
893 let response = runtime.block_on(run_sweep_fanout(
894 client,
895 "prefix cache reset",
896 "/reset_prefix_cache",
897 Vec::new(),
898 None,
899 ));
900 assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
901 let body = runtime.block_on(axum::body::to_bytes(response.into_body(), usize::MAX))?;
902 let value: Value = serde_json::from_slice(&body)?;
903 assert!(
904 value["error"]
905 .as_str()
906 .is_some_and(|error| error.contains("no targets")),
907 "got {value}"
908 );
909 Ok(())
910 }
911
912 #[test]
916 fn prime_fanout_is_cancelled_when_the_caller_gives_up() -> Result<()> {
917 let runtime = proxy_test_runtime()?;
918 let dropped = Arc::new(AtomicBool::new(false));
919 let flag = dropped.clone();
920 let observed = dropped.clone();
921 runtime.block_on(async move {
922 let app = axum::Router::new().route(
923 "/prime",
924 axum::routing::post(move || {
925 let flag = flag.clone();
926 async move {
927 run_prime_fanout(
928 "prefix cache conditioning",
929 vec![StaticPrimeTarget {
930 url: "http://127.0.0.1:1",
931 rank: 0,
932 }],
933 move |_target| {
934 let guard = SetOnDrop(flag.clone());
935 async move {
936 let _guard = guard;
937 futures_util::future::pending::<()>().await;
938 Ok::<u16, PrimeFlowFailure>(200)
939 }
940 },
941 )
942 .await
943 }
944 }),
945 );
946 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
947 let addr = listener.local_addr()?;
948 tokio::spawn(async move { axum::serve(listener, app).await });
949 let result = reqwest::Client::new()
950 .post(format!("http://{addr}/prime"))
951 .timeout(Duration::from_millis(200))
952 .send()
953 .await;
954 assert!(result.is_err(), "the pending target must hold the response");
955 for _ in 0..50 {
956 if observed.load(Ordering::SeqCst) {
957 break;
958 }
959 tokio::time::sleep(Duration::from_millis(20)).await;
960 }
961 anyhow::Ok(())
962 })?;
963 assert!(
964 dropped.load(Ordering::SeqCst),
965 "target work outlived the caller's deadline"
966 );
967 Ok(())
968 }
969
970 use std::sync::Arc;
971 use std::sync::atomic::{AtomicBool, Ordering};
972
973 struct SetOnDrop(Arc<AtomicBool>);
976
977 impl Drop for SetOnDrop {
978 fn drop(&mut self) {
979 self.0.store(true, Ordering::SeqCst);
980 }
981 }
982
983 fn proxy_test_runtime() -> Result<tokio::runtime::Runtime> {
984 tokio::runtime::Builder::new_multi_thread()
985 .enable_all()
986 .build()
987 .map_err(|error| anyhow::anyhow!(error.to_string()))
988 }
989
990 #[test]
991 fn streamed_decode_yields_bytes_in_order_when_prefill_succeeds() -> Result<()> {
992 let runtime = proxy_test_runtime()?;
993 let bytes = runtime.block_on(async {
994 let decode = Box::pin(futures_util::stream::iter(vec![
995 std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"hello")),
996 Ok(Bytes::from_static(b" world")),
997 ]));
998 let prefill = tokio::spawn(async { Ok::<(), ProxyHttpError>(()) });
999 let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1000 let mut out = Vec::new();
1001 while let Some(item) = stream.next().await {
1002 out.push(item.map_err(|error| anyhow::anyhow!(error.to_string()))?);
1003 }
1004 anyhow::Ok(out)
1005 })?;
1006 let joined: Vec<u8> = bytes.into_iter().flatten().collect();
1007 assert_eq!(joined, b"hello world");
1008 Ok(())
1009 }
1010
1011 #[test]
1012 fn streamed_decode_surfaces_prefill_error_after_decode_ends() -> Result<()> {
1013 let runtime = proxy_test_runtime()?;
1014 let (bytes, error) = runtime.block_on(async {
1015 let decode = Box::pin(futures_util::stream::iter(vec![std::result::Result::<
1016 Bytes,
1017 std::io::Error,
1018 >::Ok(
1019 Bytes::from_static(b"partial"),
1020 )]));
1021 let prefill = tokio::spawn(async {
1022 Err::<(), ProxyHttpError>(ProxyHttpError::internal("prefill boom"))
1023 });
1024 let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1025 let mut bytes = Vec::new();
1026 let mut error = None;
1027 while let Some(item) = stream.next().await {
1028 match item {
1029 Ok(chunk) => bytes.extend_from_slice(&chunk),
1030 Err(stream_error) => {
1031 error = Some(stream_error.to_string());
1032 break;
1033 }
1034 }
1035 }
1036 anyhow::Ok((bytes, error))
1037 })?;
1038 assert_eq!(bytes, b"partial");
1039 let error = error.context("expected a prefill error to surface after decode ended")?;
1040 assert!(error.contains("prefill boom"), "got {error}");
1041 Ok(())
1042 }
1043
1044 #[test]
1049 fn prefill_error_surfaces_even_while_decode_stays_ready() -> Result<()> {
1050 let runtime = proxy_test_runtime()?;
1051 let error = runtime.block_on(async {
1052 let decode = Box::pin(futures_util::stream::repeat_with(|| {
1054 std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"x"))
1055 }));
1056 let prefill = tokio::spawn(async {
1057 Err::<(), ProxyHttpError>(ProxyHttpError::internal("prefill boom"))
1058 });
1059 let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1060 let mut chunks = 0usize;
1061 let mut error = None;
1062 while let Some(item) = stream.next().await {
1063 match item {
1064 Ok(_) => {
1065 chunks += 1;
1066 assert!(
1070 chunks < 100_000,
1071 "prefill error was suppressed by a continuously-ready decode stream"
1072 );
1073 }
1074 Err(stream_error) => {
1075 error = Some(stream_error.to_string());
1076 break;
1077 }
1078 }
1079 }
1080 anyhow::Ok(error)
1081 })?;
1082 let error = error.context("a prefill error must surface even while decode stays ready")?;
1083 assert!(error.contains("prefill boom"), "got {error}");
1084 Ok(())
1085 }
1086
1087 #[test]
1093 fn decode_error_ready_at_tiebreak_is_not_swallowed() -> Result<()> {
1094 let runtime = proxy_test_runtime()?;
1095 let error = runtime.block_on(async {
1096 let prefill = tokio::spawn(async { Ok::<(), ProxyHttpError>(()) });
1099 while !prefill.is_finished() {
1100 tokio::task::yield_now().await;
1101 }
1102 let decode = Box::pin(futures_util::stream::iter(vec![std::result::Result::<
1104 Bytes,
1105 std::io::Error,
1106 >::Err(
1107 std::io::Error::other("decode boom"),
1108 )]));
1109 let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1110 let mut error = None;
1111 while let Some(item) = stream.next().await {
1112 if let Err(stream_error) = item {
1113 error = Some(stream_error.to_string());
1114 break;
1115 }
1116 }
1117 anyhow::Ok(error)
1118 })?;
1119 let error =
1120 error.context("a decode error ready at the tie-break must surface, not truncate")?;
1121 assert!(error.contains("decode boom"), "got {error}");
1122 Ok(())
1123 }
1124
1125 #[test]
1126 fn dropping_the_stream_before_prefill_finishes_aborts_prefill() -> Result<()> {
1127 let runtime = proxy_test_runtime()?;
1128 let aborted = Arc::new(AtomicBool::new(false));
1129 let flag = aborted.clone();
1130 let cancelled = runtime.block_on(async move {
1131 let (started_tx, started_rx) = tokio::sync::oneshot::channel::<()>();
1132 let prefill = tokio::spawn(async move {
1137 let _guard = SetOnDrop(flag);
1138 let _ = started_tx.send(());
1139 futures_util::future::pending::<()>().await;
1140 Ok::<(), ProxyHttpError>(())
1141 });
1142 let _ = started_rx.await;
1143 let decode = Box::pin(
1146 futures_util::stream::once(async {
1147 std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"a"))
1148 })
1149 .chain(futures_util::stream::pending::<
1150 std::result::Result<Bytes, std::io::Error>,
1151 >()),
1152 );
1153 let mut stream = Box::pin(decode_response_stream(decode, prefill, OnClientDrop::Abort));
1154 assert!(matches!(stream.next().await, Some(Ok(_))));
1155 drop(stream);
1156 for _ in 0..200 {
1157 if aborted.load(Ordering::SeqCst) {
1158 return true;
1159 }
1160 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1161 }
1162 false
1163 });
1164 assert!(
1165 cancelled,
1166 "prefill task was not aborted when the response stream was dropped"
1167 );
1168 Ok(())
1169 }
1170
1171 #[test]
1175 fn dropping_the_stream_before_prefill_finishes_detaches_prefill() -> Result<()> {
1176 let runtime = proxy_test_runtime()?;
1177 let completed = Arc::new(AtomicBool::new(false));
1178 let flag = completed.clone();
1179 let finished = runtime.block_on(async move {
1180 let (started_tx, started_rx) = tokio::sync::oneshot::channel::<()>();
1181 let prefill = tokio::spawn(async move {
1184 let _ = started_tx.send(());
1185 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1186 flag.store(true, Ordering::SeqCst);
1187 Ok::<(), ProxyHttpError>(())
1188 });
1189 let _ = started_rx.await;
1190 let decode = Box::pin(
1193 futures_util::stream::once(async {
1194 std::result::Result::<Bytes, std::io::Error>::Ok(Bytes::from_static(b"a"))
1195 })
1196 .chain(futures_util::stream::pending::<
1197 std::result::Result<Bytes, std::io::Error>,
1198 >()),
1199 );
1200 let mut stream = Box::pin(decode_response_stream(
1201 decode,
1202 prefill,
1203 OnClientDrop::Detach,
1204 ));
1205 assert!(matches!(stream.next().await, Some(Ok(_))));
1206 drop(stream);
1207 for _ in 0..200 {
1208 if completed.load(Ordering::SeqCst) {
1209 return true;
1210 }
1211 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1212 }
1213 false
1214 });
1215 assert!(
1216 finished,
1217 "prefill task was aborted instead of detached when the response stream was dropped"
1218 );
1219 Ok(())
1220 }
1221}