rivetkit_core/
serverless_http.rs1use std::collections::HashMap;
2use std::path::PathBuf;
3use std::pin::Pin;
4use std::task::{Context as TaskContext, Poll};
5
6use anyhow::{Context, Result};
7use axum::Router;
8use axum::body::{Body, Bytes};
9use axum::extract::{Request, State};
10use axum::http::{HeaderMap, HeaderName, HeaderValue, StatusCode};
11use axum::response::IntoResponse;
12use axum::routing::any;
13use futures::Stream;
14use futures::StreamExt;
15use http_body_util::LengthLimitError;
16use tokio_stream::wrappers::UnboundedReceiverStream;
17use tokio_util::sync::CancellationToken;
18use tower_http::services::ServeDir;
19
20use crate::serverless::{CoreServerlessRuntime, ServerlessRequest, ServerlessResponse};
21
22#[derive(Debug, Clone)]
23pub struct ListenerConfig {
24 pub host: Option<String>,
26 pub port: u16,
27 pub public_dir: Option<PathBuf>,
28}
29
30#[derive(Clone)]
31struct AppState {
32 runtime: CoreServerlessRuntime,
33 shutdown_token: CancellationToken,
34}
35
36pub async fn serve(
38 runtime: CoreServerlessRuntime,
39 listener: ListenerConfig,
40 shutdown: CancellationToken,
41) -> Result<()> {
42 let host = listener.host.as_deref().unwrap_or("0.0.0.0");
43 let port = listener.port;
44
45 let state = AppState {
46 runtime,
47 shutdown_token: shutdown.clone(),
48 };
49
50 let forward_service = any(forward_request).with_state(state);
51
52 let router = match listener.public_dir.as_ref() {
53 Some(dir) => Router::new().fallback_service(
54 ServeDir::new(dir)
55 .call_fallback_on_method_not_allowed(true)
56 .fallback(forward_service),
57 ),
58 None => Router::new().fallback_service(forward_service),
59 };
60
61 let tcp = tokio::net::TcpListener::bind((host, port))
62 .await
63 .with_context(|| format!("bind tcp listener on {host}:{port}"))?;
64 let bound = tcp
65 .local_addr()
66 .context("read local address of bound listener")?;
67 tracing::info!(host = %bound.ip(), port = bound.port(), "rivetkit server listening");
68
69 let shutdown_fut = {
70 let shutdown = shutdown.clone();
71 async move { shutdown.cancelled().await }
72 };
73
74 axum::serve(tcp, router.into_make_service())
75 .with_graceful_shutdown(shutdown_fut)
76 .await
77 .context("axum::serve returned an error")?;
78
79 Ok(())
80}
81
82async fn forward_request(
83 State(state): State<AppState>,
84 request: Request,
85) -> axum::response::Response {
86 let (parts, body) = request.into_parts();
87 let body_limit = state.runtime.max_request_body_bytes();
88 let request_token = state.shutdown_token.child_token();
89 let body_bytes = match axum::body::to_bytes(body, body_limit).await {
90 Ok(bytes) => bytes,
91 Err(error) if is_length_limit_error(&error) => {
92 tracing::warn!(body_limit, "request body exceeded limit");
93 return into_axum_response(state.runtime.incoming_too_long_response(), request_token);
94 }
95 Err(error) => {
96 tracing::warn!(?error, "failed to read request body");
97 return into_axum_response(
98 state
99 .runtime
100 .invalid_request_response("failed to read request body"),
101 request_token,
102 );
103 }
104 };
105
106 let path_and_query = parts
107 .uri
108 .path_and_query()
109 .map(|pq| pq.as_str())
110 .unwrap_or("/");
111 let url = format!("http://internal{path_and_query}");
112
113 let mut headers: HashMap<String, String> = HashMap::new();
115 for (name, value) in parts.headers.iter() {
116 let Ok(value_str) = value.to_str() else {
117 continue;
118 };
119 let key = name.as_str().to_ascii_lowercase();
120 headers
121 .entry(key)
122 .and_modify(|existing| {
123 existing.push_str(", ");
124 existing.push_str(value_str);
125 })
126 .or_insert_with(|| value_str.to_owned());
127 }
128
129 let req = ServerlessRequest {
130 method: parts.method.as_str().to_owned(),
131 url,
132 headers,
133 body: body_bytes.to_vec(),
134 cancel_token: request_token.clone(),
135 };
136
137 into_axum_response(state.runtime.handle_request(req).await, request_token)
138}
139
140fn into_axum_response(
141 response: ServerlessResponse,
142 request_token: CancellationToken,
143) -> axum::response::Response {
144 let status = StatusCode::from_u16(response.status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
145 let mut header_map = HeaderMap::with_capacity(response.headers.len());
146 for (name, value) in response.headers {
147 if let (Ok(name), Ok(value)) = (
148 HeaderName::try_from(name.as_str()),
149 HeaderValue::from_str(&value),
150 ) {
151 header_map.append(name, value);
152 }
153 }
154
155 let stream = UnboundedReceiverStream::new(response.body).map(|chunk| match chunk {
156 Ok(bytes) => Ok::<Bytes, std::io::Error>(Bytes::from(bytes)),
157 Err(error) => {
158 tracing::warn!(?error, "serverless stream error");
159 Err(std::io::Error::other(format!(
160 "{}.{}: {}",
161 error.group, error.code, error.message
162 )))
163 }
164 });
165
166 let guarded = CancelOnDropStream {
168 inner: stream,
169 _guard: CancelOnDrop {
170 token: request_token,
171 },
172 };
173
174 (status, header_map, Body::from_stream(guarded)).into_response()
175}
176
177fn is_length_limit_error(error: &axum::Error) -> bool {
178 let mut source: Option<&dyn std::error::Error> = Some(error);
179 while let Some(err) = source {
180 if err.is::<LengthLimitError>() {
181 return true;
182 }
183 source = err.source();
184 }
185 false
186}
187
188struct CancelOnDrop {
189 token: CancellationToken,
190}
191
192impl Drop for CancelOnDrop {
193 fn drop(&mut self) {
194 self.token.cancel();
195 }
196}
197
198struct CancelOnDropStream<S> {
199 inner: S,
200 _guard: CancelOnDrop,
201}
202
203impl<S: Stream + Unpin> Stream for CancelOnDropStream<S> {
204 type Item = S::Item;
205
206 fn poll_next(mut self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll<Option<Self::Item>> {
207 Pin::new(&mut self.inner).poll_next(cx)
208 }
209}