Skip to content

Commit 416804b

Browse files
simplify
1 parent 1def7a4 commit 416804b

4 files changed

Lines changed: 34 additions & 133 deletions

File tree

crates/client-api/src/routes/internal.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -158,11 +158,11 @@ mod jemalloc_profiling {
158158
}
159159

160160
// The internal router is for things that are not meant to be exposed to the public API.
161-
pub fn router<S>(task_dumps: TaskDumpRegistry) -> axum::Router<S>
161+
pub fn router<S>() -> axum::Router<S>
162162
where
163163
S: NodeDelegate + Clone + 'static,
164164
{
165165
axum::Router::new()
166166
.nest("/heap", jemalloc_profiling::jemalloc_router())
167-
.nest("/task-dump", task_dump::router(task_dumps))
167+
.nest("/task-dump", task_dump::router())
168168
}
Lines changed: 26 additions & 107 deletions
Original file line numberDiff line numberDiff line change
@@ -1,124 +1,73 @@
1-
use axum::{extract::Extension, routing::get};
1+
use axum::routing::get;
22

33
#[cfg(all(
44
tokio_unstable,
55
target_os = "linux",
66
any(target_arch = "aarch64", target_arch = "x86", target_arch = "x86_64")
77
))]
88
mod imp {
9-
use std::{collections::BTreeMap, fmt::Write as _, sync::Arc, time::Duration};
9+
use std::{collections::BTreeMap, fmt::Write as _, num::NonZeroU64, sync::Arc, time::Duration};
1010

1111
use axum::{
1212
extract::{Extension, Query},
1313
response::Response,
1414
};
1515
use http::{header::CONTENT_TYPE, StatusCode};
1616
use serde::Deserialize;
17-
use tokio::{runtime::Handle, sync::Mutex};
18-
19-
const MAX_TIMEOUT_MS: u64 = 30_000;
17+
use tokio::runtime::Handle;
2018

21-
#[derive(Clone)]
22-
struct Runtime {
23-
handle: Handle,
24-
dump_lock: Arc<Mutex<()>>,
25-
}
19+
const DEFAULT_TIMEOUT_MS: u64 = 2_000;
2620

2721
/// The Tokio runtimes which can be inspected by the internal task dump endpoint.
2822
#[derive(Clone, Default)]
2923
pub struct TaskDumpRegistry {
30-
runtimes: Arc<BTreeMap<&'static str, Runtime>>,
24+
runtimes: Arc<BTreeMap<&'static str, Handle>>,
3125
}
3226

3327
impl TaskDumpRegistry {
3428
pub fn new(runtimes: impl IntoIterator<Item = (&'static str, Handle)>) -> Self {
35-
let runtimes = runtimes
36-
.into_iter()
37-
.map(|(name, handle)| {
38-
(
39-
name,
40-
Runtime {
41-
handle,
42-
dump_lock: Arc::new(Mutex::new(())),
43-
},
44-
)
45-
})
46-
.collect();
4729
Self {
48-
runtimes: Arc::new(runtimes),
30+
runtimes: Arc::new(runtimes.into_iter().collect()),
4931
}
5032
}
51-
52-
fn get(&self, name: &str) -> Option<&Runtime> {
53-
self.runtimes.get(name)
54-
}
55-
56-
fn names(&self) -> impl Iterator<Item = &'static str> + '_ {
57-
self.runtimes.keys().copied()
58-
}
5933
}
6034

6135
#[derive(Deserialize)]
6236
pub(super) struct TaskDumpQuery {
6337
runtime: String,
64-
timeout_ms: u64,
65-
}
66-
67-
impl TaskDumpQuery {
68-
fn validate(&self) -> Result<(), (StatusCode, String)> {
69-
if !(1..=MAX_TIMEOUT_MS).contains(&self.timeout_ms) {
70-
return Err((
71-
StatusCode::BAD_REQUEST,
72-
format!("timeout_ms must be between 1 and {MAX_TIMEOUT_MS}"),
73-
));
74-
}
75-
Ok(())
76-
}
38+
timeout_ms: Option<NonZeroU64>,
7739
}
7840

7941
pub(super) async fn handle_get_task_dump(
80-
Extension(registry): Extension<TaskDumpRegistry>,
42+
registry: Option<Extension<TaskDumpRegistry>>,
8143
Query(query): Query<TaskDumpQuery>,
8244
) -> Result<Response, (StatusCode, String)> {
83-
query.validate()?;
45+
let Some(Extension(registry)) = registry else {
46+
return Err((
47+
StatusCode::NOT_IMPLEMENTED,
48+
"Tokio task dumps are not configured for this server".into(),
49+
));
50+
};
8451

85-
let Some(runtime) = registry.get(&query.runtime).cloned() else {
86-
let valid = registry.names().collect::<Vec<_>>().join(", ");
52+
let Some(runtime) = registry.runtimes.get(query.runtime.as_str()).cloned() else {
53+
let valid = registry.runtimes.keys().copied().collect::<Vec<_>>().join(", ");
8754
return Err((
8855
StatusCode::BAD_REQUEST,
8956
format!("unknown Tokio runtime {:?}; valid runtimes: {valid}", query.runtime),
9057
));
9158
};
9259

93-
let permit = runtime.dump_lock.try_lock_owned().map_err(|_| {
94-
(
95-
StatusCode::CONFLICT,
96-
format!("a task dump is already in progress for runtime {:?}", query.runtime),
97-
)
98-
})?;
99-
100-
// Keep the permit in the spawned task so that a timed-out dump continues
101-
// to exclude new requests until Tokio's dump future actually finishes.
102-
let dump_task = tokio::spawn(async move {
103-
let dump = runtime.handle.dump().await;
104-
(permit, dump)
105-
});
106-
let (_permit, dump) = tokio::time::timeout(Duration::from_millis(query.timeout_ms), dump_task)
60+
let timeout_ms = query.timeout_ms.map_or(DEFAULT_TIMEOUT_MS, NonZeroU64::get);
61+
let dump = tokio::time::timeout(Duration::from_millis(timeout_ms), runtime.dump())
10762
.await
10863
.map_err(|_| {
10964
(
11065
StatusCode::GATEWAY_TIMEOUT,
11166
format!(
112-
"timed out after {}ms while dumping Tokio runtime {:?}",
113-
query.timeout_ms, query.runtime
67+
"timed out after {timeout_ms}ms while dumping Tokio runtime {:?}",
68+
query.runtime
11469
),
11570
)
116-
})?
117-
.map_err(|err| {
118-
(
119-
StatusCode::INTERNAL_SERVER_ERROR,
120-
format!("task dump worker failed for runtime {:?}: {err}", query.runtime),
121-
)
12271
})?;
12372

12473
let runtime_name = query.runtime;
@@ -155,44 +104,16 @@ mod imp {
155104
mod tests {
156105
use super::*;
157106
use http_body_util::BodyExt as _;
158-
159-
#[tokio::test]
160-
async fn registry_names_are_sorted() {
161-
let handle = Handle::current();
162-
let registry = TaskDumpRegistry::new([("replication", handle.clone()), ("main", handle)]);
163-
164-
assert_eq!(registry.names().collect::<Vec<_>>(), ["main", "replication"]);
165-
}
166-
167-
#[test]
168-
fn timeout_must_be_in_range() {
169-
for timeout_ms in [1, MAX_TIMEOUT_MS] {
170-
assert!(TaskDumpQuery {
171-
runtime: "main".into(),
172-
timeout_ms,
173-
}
174-
.validate()
175-
.is_ok());
176-
}
177-
178-
for timeout_ms in [0, MAX_TIMEOUT_MS + 1] {
179-
assert!(TaskDumpQuery {
180-
runtime: "main".into(),
181-
timeout_ms,
182-
}
183-
.validate()
184-
.is_err());
185-
}
186-
}
107+
use tokio::runtime::Handle;
187108

188109
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
189-
async fn dumps_registered_runtime() {
110+
async fn dumps_registered_runtime_with_default_timeout() {
190111
let registry = TaskDumpRegistry::new([("main", Handle::current())]);
191112
let response = handle_get_task_dump(
192-
Extension(registry),
113+
Some(Extension(registry)),
193114
Query(TaskDumpQuery {
194115
runtime: "main".into(),
195-
timeout_ms: 10_000,
116+
timeout_ms: None,
196117
}),
197118
)
198119
.await
@@ -248,8 +169,6 @@ mod imp {
248169
use imp::handle_get_task_dump;
249170
pub use imp::TaskDumpRegistry;
250171

251-
pub fn router<S: Clone + Send + Sync + 'static>(registry: TaskDumpRegistry) -> axum::Router<S> {
252-
axum::Router::new()
253-
.route("/", get(handle_get_task_dump))
254-
.layer(Extension(registry))
172+
pub fn router<S: Clone + Send + Sync + 'static>() -> axum::Router<S> {
173+
axum::Router::new().route("/", get(handle_get_task_dump))
255174
}

crates/client-api/src/routes/mod.rs

Lines changed: 1 addition & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -27,25 +27,6 @@ pub fn router<S>(
2727
identity_routes: IdentityRoutes<S>,
2828
extra: axum::Router<S>,
2929
) -> axum::Router<S>
30-
where
31-
S: NodeDelegate + ControlStateDelegate + Authorization + Clone + 'static,
32-
{
33-
router_with_task_dumps(
34-
ctx,
35-
database_routes,
36-
identity_routes,
37-
extra,
38-
TaskDumpRegistry::default(),
39-
)
40-
}
41-
42-
pub fn router_with_task_dumps<S>(
43-
ctx: &S,
44-
database_routes: DatabaseRoutes<S>,
45-
identity_routes: IdentityRoutes<S>,
46-
extra: axum::Router<S>,
47-
task_dumps: TaskDumpRegistry,
48-
) -> axum::Router<S>
4930
where
5031
S: NodeDelegate + ControlStateDelegate + Authorization + Clone + 'static,
5132
{
@@ -66,5 +47,5 @@ where
6647

6748
axum::Router::new()
6849
.nest("/v1", router.layer(cors))
69-
.nest("/internal", internal::router(task_dumps))
50+
.nest("/internal", internal::router())
7051
}

crates/standalone/src/subcommands/start.rs

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@ use std::sync::Arc;
77

88
use crate::{StandaloneEnv, StandaloneOptions};
99
use anyhow::Context;
10-
use axum::extract::DefaultBodyLimit;
10+
use axum::extract::{DefaultBodyLimit, Extension};
1111
use clap::ArgAction::SetTrue;
1212
use clap::{Arg, ArgMatches};
1313
use spacetimedb::config::{parse_config, CertificateAuthority};
@@ -17,7 +17,7 @@ use spacetimedb::startup::{self, TracingOptions};
1717
use spacetimedb::util::jobs::JobCores;
1818
use spacetimedb::worker_metrics;
1919
use spacetimedb_client_api::routes::database::DatabaseRoutes;
20-
use spacetimedb_client_api::routes::router_with_task_dumps;
20+
use spacetimedb_client_api::routes::router;
2121
use spacetimedb_client_api::routes::subscribe::WebSocketOptions;
2222
use spacetimedb_client_api::routes::TaskDumpRegistry;
2323
use spacetimedb_paths::cli::{PrivKeyPath, PubKeyPath};
@@ -208,8 +208,9 @@ pub async fn exec(args: &ArgMatches, db_cores: JobCores) -> anyhow::Result<()> {
208208
db_routes.pre_publish = db_routes.pre_publish.layer(DefaultBodyLimit::disable());
209209
let extra = axum::Router::new().nest("/health", spacetimedb_client_api::routes::health::router());
210210
let task_dumps = TaskDumpRegistry::new([("main", main_rt)]);
211-
let service =
212-
router_with_task_dumps(&ctx, db_routes, IdentityRoutes::default(), extra, task_dumps).with_state(ctx.clone());
211+
let service = router(&ctx, db_routes, IdentityRoutes::default(), extra)
212+
.layer(Extension(task_dumps))
213+
.with_state(ctx.clone());
213214

214215
// Check if the requested port is available on both IPv4 and IPv6.
215216
// If not, offer to find an available port by incrementing (unless non-interactive).

0 commit comments

Comments
 (0)