systemprompt_api/services/server/
startup.rs1use std::convert::Infallible;
14use std::net::SocketAddr;
15use std::pin::Pin;
16use std::sync::{Arc, PoisonError, RwLock};
17use std::task::{Context, Poll};
18
19use anyhow::{Context as _, Result};
20use axum::body::Body;
21use axum::http::{Request, StatusCode};
22use axum::response::{IntoResponse, Response};
23use axum::routing::get;
24use axum::{Json, Router};
25use serde_json::json;
26use systemprompt_models::api::ApiError;
27use systemprompt_models::modules::ApiPaths;
28use systemprompt_runtime::ShutdownRequest;
29use systemprompt_traits::{OwnedTask, StartupEvent, StartupEventExt, StartupEventSender};
30use tower::ServiceExt;
31
32#[derive(Debug)]
33pub struct EarlyServer {
34 swap: Arc<RwLock<Router>>,
35 join: OwnedTask<Result<()>>,
36 local_addr: SocketAddr,
37}
38
39impl EarlyServer {
40 pub const fn local_addr(&self) -> SocketAddr {
41 self.local_addr
42 }
43
44 pub fn activate(&self, router: Router) {
45 *self.swap.write().unwrap_or_else(PoisonError::into_inner) = router;
46 tracing::info!("Full API router activated");
47 }
48
49 pub async fn join(self) -> Result<()> {
50 self.join.join().await.context("API serve task panicked")?
51 }
52}
53
54pub async fn bind_and_serve(
55 addr: &str,
56 events: Option<StartupEventSender>,
57 shutdown: ShutdownRequest,
58) -> Result<EarlyServer> {
59 if let Some(ref tx) = events
60 && tx
61 .unbounded_send(StartupEvent::ServerBinding {
62 address: addr.to_owned(),
63 })
64 .is_err()
65 {
66 tracing::debug!("Startup event receiver dropped");
67 }
68
69 let listener = tokio::net::TcpListener::bind(addr)
70 .await
71 .with_context(|| format!("Failed to bind to {addr}"))?;
72 let local_addr = listener
73 .local_addr()
74 .context("Failed to read bound address")?;
75
76 if let Some(ref tx) = events {
77 tx.server_listening(addr, std::process::id());
78 }
79
80 let swap = Arc::new(RwLock::new(starting_router()));
81 let outer = Router::new().fallback_service(SwapService {
82 swap: Arc::clone(&swap),
83 });
84
85 let join = OwnedTask::spawn("api_serve", async move {
86 axum::serve(
87 listener,
88 outer.into_make_service_with_connect_info::<SocketAddr>(),
89 )
90 .with_graceful_shutdown(super::shutdown::shutdown_signal(shutdown))
91 .await
92 .map_err(Into::into)
93 });
94
95 Ok(EarlyServer {
96 swap,
97 join,
98 local_addr,
99 })
100}
101
102pub fn starting_router() -> Router {
103 Router::new()
104 .route(ApiPaths::HEALTH, get(starting_health))
105 .route("/health", get(starting_health))
106 .route(ApiPaths::LIVEZ, get(starting_livez))
107 .route(ApiPaths::READYZ, get(starting_readyz))
108 .fallback(starting_fallback)
109}
110
111async fn starting_health() -> impl IntoResponse {
112 Json(json!({ "status": "starting" }))
113}
114
115async fn starting_livez() -> impl IntoResponse {
116 Json(json!({ "status": "alive", "version": env!("CARGO_PKG_VERSION") }))
117}
118
119async fn starting_readyz() -> impl IntoResponse {
120 (
121 StatusCode::SERVICE_UNAVAILABLE,
122 Json(json!({ "status": "starting", "version": env!("CARGO_PKG_VERSION") })),
123 )
124}
125
126async fn starting_fallback() -> Response {
127 tracing::debug!("request answered with 503 while the service is starting");
128 let body = ApiError::service_unavailable("service starting").with_error_key("service_starting");
129 (StatusCode::SERVICE_UNAVAILABLE, Json(body)).into_response()
130}
131
132#[derive(Clone)]
133struct SwapService {
134 swap: Arc<RwLock<Router>>,
135}
136
137impl tower::Service<Request<Body>> for SwapService {
138 type Response = Response;
139 type Error = Infallible;
140 type Future = Pin<Box<dyn Future<Output = Result<Response, Infallible>> + Send>>;
141
142 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Infallible>> {
143 Poll::Ready(Ok(()))
144 }
145
146 fn call(&mut self, req: Request<Body>) -> Self::Future {
147 let router = self
148 .swap
149 .read()
150 .unwrap_or_else(PoisonError::into_inner)
151 .clone();
152 Box::pin(router.oneshot(req))
153 }
154}