Skip to main content

rivetkit_core/
serverless_http.rs

1use 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	/// Host to bind; accepts numeric IPs or DNS names. Defaults to `0.0.0.0`.
25	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
36/// Bind a TCP listener and serve `runtime` over HTTP until `shutdown` fires.
37pub 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	// Repeated header names get comma-joined per RFC 9110 ยง5.3.
114	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	// Cancel the runtime task when the response body is dropped.
167	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}