Skip to content

Commit 708281d

Browse files
perf(cache): consolidate HotSnapshot atomic refcounts into single Arc<FrozenState> (#149)
1 parent b804b38 commit 708281d

3 files changed

Lines changed: 67 additions & 46 deletions

File tree

src/dispatch.rs

Lines changed: 16 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -316,6 +316,7 @@ fn ensure_compiled_snapshot(state: &Arc<RwLock<AppState>>) -> Arc<CompiledRouter
316316
let mut st = state.write();
317317
if st.compiled.is_none() {
318318
st.compiled = Some(Arc::new(st.snapshot_routers()));
319+
st.rebuild_snapshot();
319320
}
320321
Arc::clone(st.compiled.as_ref().expect("just populated"))
321322
}
@@ -472,20 +473,6 @@ pub async fn run_rsgi(
472473
let Some((method, path, query_string, is_head, snapshot)) = prelim else {
473474
return Ok(Python::with_gil(|py| py.None()));
474475
};
475-
type CfgClones = (
476-
Option<Py<PyAny>>,
477-
Option<Py<PyAny>>,
478-
Arc<Vec<Py<PyAny>>>,
479-
Arc<Vec<Py<PyAny>>>,
480-
);
481-
let (cors_cfg, security_cfg, req_mw, res_mw): CfgClones = Python::with_gil(|_py| {
482-
(
483-
snapshot.cors.clone(),
484-
snapshot.security_headers.clone(),
485-
snapshot.request_middleware.clone(),
486-
snapshot.response_middleware.clone(),
487-
)
488-
});
489476
if (method == "GET" || method == "HEAD") && path == "/openapi.json" && snapshot.include_openapi
490477
{
491478
let _ = Python::with_gil(|py| {
@@ -539,7 +526,7 @@ pub async fn run_rsgi(
539526
}
540527
}
541528
}
542-
for mw in req_mw.iter() {
529+
for mw in snapshot.request_middleware.iter() {
543530
let out: Py<PyAny> = match Python::with_gil(|py| {
544531
let f = mw.bind(py);
545532
f.call1((scope.bind(py), protocol.bind(py)))
@@ -556,10 +543,12 @@ pub async fn run_rsgi(
556543
return match Python::with_gil(|py| -> PyResult<()> {
557544
let mut mapped = map_handler_return(py, &out)?;
558545

559-
if !res_mw.is_empty() && !matches!(mapped, HandlerMap::AlreadySent) {
546+
if !snapshot.response_middleware.is_empty()
547+
&& !matches!(mapped, HandlerMap::AlreadySent)
548+
{
560549
let response_module = py.import("oxyroute.response")?;
561550
let response_class = response_module.getattr("Response")?;
562-
for res_m in res_mw.iter() {
551+
for res_m in snapshot.response_middleware.iter() {
563552
let kwargs = pyo3::types::PyDict::new(py);
564553
match &mapped {
565554
HandlerMap::Simple {
@@ -605,16 +594,16 @@ pub async fn run_rsgi(
605594
}
606595
}
607596

608-
let mapped = if security_cfg.is_some() || cors_cfg.is_some() {
597+
let mapped = if snapshot.security_headers.is_some() || snapshot.cors.is_some() {
609598
let scope_bound = scope.bind(py).clone();
610599
let mapped = merge_config_response_headers(
611600
py,
612-
&security_cfg,
601+
&snapshot.security_headers,
613602
scope_bound.clone(),
614603
mapped,
615604
true,
616605
)?;
617-
merge_config_response_headers(py, &cors_cfg, scope_bound, mapped, false)?
606+
merge_config_response_headers(py, &snapshot.cors, scope_bound, mapped, false)?
618607
} else {
619608
mapped
620609
};
@@ -1240,10 +1229,10 @@ pub async fn run_rsgi(
12401229
match Python::with_gil(|py| -> PyResult<()> {
12411230
let mut mapped = map_handler_return(py, &handler_out)?;
12421231

1243-
if !res_mw.is_empty() && !matches!(mapped, HandlerMap::AlreadySent) {
1232+
if !snapshot.response_middleware.is_empty() && !matches!(mapped, HandlerMap::AlreadySent) {
12441233
let response_module = py.import("oxyroute.response")?;
12451234
let response_class = response_module.getattr("Response")?;
1246-
for res_m in res_mw.iter() {
1235+
for res_m in snapshot.response_middleware.iter() {
12471236
let kwargs = pyo3::types::PyDict::new(py);
12481237
match &mapped {
12491238
HandlerMap::Simple {
@@ -1287,16 +1276,16 @@ pub async fn run_rsgi(
12871276
}
12881277
}
12891278

1290-
let mapped = if security_cfg.is_some() || cors_cfg.is_some() {
1279+
let mapped = if snapshot.security_headers.is_some() || snapshot.cors.is_some() {
12911280
let scope_bound = scope.bind(py).clone();
12921281
let mapped = merge_config_response_headers(
12931282
py,
1294-
&security_cfg,
1283+
&snapshot.security_headers,
12951284
scope_bound.clone(),
12961285
mapped,
12971286
true,
12981287
)?;
1299-
merge_config_response_headers(py, &cors_cfg, scope_bound, mapped, false)?
1288+
merge_config_response_headers(py, &snapshot.cors, scope_bound, mapped, false)?
13001289
} else {
13011290
mapped
13021291
};
@@ -1769,7 +1758,6 @@ async fn run_rsgi_websocket(
17691758
Some(c) => c,
17701759
None => ensure_compiled_snapshot(&state),
17711760
};
1772-
let ws_routes = Arc::clone(&snapshot.websocket_routes);
17731761
let route_match = match_ws_route_compiled(&compiled, &path);
17741762
let Some((route_idx, params)) = route_match else {
17751763
// No route → polite close. ``close`` is sync on RSGIWebsocketProtocol.
@@ -1781,7 +1769,8 @@ async fn run_rsgi_websocket(
17811769
return Ok(Python::with_gil(|py| py.None()));
17821770
};
17831771
let (handler, is_async) = Python::with_gil(|_py| -> PyResult<(Py<PyAny>, bool)> {
1784-
let e = ws_routes
1772+
let e = snapshot
1773+
.websocket_routes
17851774
.get(route_idx)
17861775
.ok_or_else(|| pyo3::exceptions::PyRuntimeError::new_err("ws route index"))?;
17871776
Ok((e.handler.clone(), e.is_async))

src/lib.rs

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -447,6 +447,7 @@ impl App {
447447
}
448448
// Keep auto-compiled routing snapshots fresh when routes are added before explicit freeze().
449449
st.compiled = None;
450+
st.rebuild_snapshot();
450451
Ok(())
451452
}
452453

@@ -480,6 +481,7 @@ impl App {
480481
.map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("{e}")))?;
481482
}
482483
st.compiled = None;
484+
st.rebuild_snapshot();
483485
Ok(())
484486
}
485487

@@ -490,12 +492,14 @@ impl App {
490492
if st.compiled.is_none() {
491493
st.compiled = Some(Arc::new(st.snapshot_routers()));
492494
}
495+
st.rebuild_snapshot();
493496
Ok(())
494497
}
495498

496499
fn set_openapi_served(&self, enabled: bool) -> PyResult<()> {
497500
let mut st = self.state.write();
498501
st.include_openapi = enabled;
502+
st.rebuild_snapshot();
499503
Ok(())
500504
}
501505

@@ -568,6 +572,7 @@ impl App {
568572
})
569573
.unwrap_or(false);
570574
Arc::make_mut(&mut st.exception_handlers).push((exc_type.unbind(), handler, is_async));
575+
st.rebuild_snapshot();
571576
Ok(())
572577
}
573578

@@ -578,6 +583,7 @@ impl App {
578583
} else {
579584
st.request_middleware = Arc::new(Vec::new());
580585
}
586+
st.rebuild_snapshot();
581587
Ok(())
582588
}
583589

@@ -596,13 +602,15 @@ impl App {
596602
"phase must be 'request', 'response', or 'both'",
597603
));
598604
}
605+
st.rebuild_snapshot();
599606
Ok(())
600607
}
601608

602609
/// Optional Python CORS config with ``response_header_pairs(scope)`` (see ``oxyroute.cors``).
603610
fn set_cors(&self, config: Option<Py<PyAny>>) -> PyResult<()> {
604611
let mut st = self.state.write();
605612
st.cors = config;
613+
st.rebuild_snapshot();
606614
Ok(())
607615
}
608616

@@ -611,6 +619,7 @@ impl App {
611619
fn set_security_headers(&self, config: Option<Py<PyAny>>) -> PyResult<()> {
612620
let mut st = self.state.write();
613621
st.security_headers = config;
622+
st.rebuild_snapshot();
614623
Ok(())
615624
}
616625

@@ -633,6 +642,7 @@ impl App {
633642
})?;
634643
let mut st = state.write();
635644
st.db_pool = Some(pool);
645+
st.rebuild_snapshot();
636646
Ok(())
637647
})
638648
}
@@ -643,7 +653,9 @@ impl App {
643653
pyo3_async_runtimes::tokio::future_into_py(py, async move {
644654
let pool = {
645655
let mut st = state.write();
646-
st.db_pool.take()
656+
let p = st.db_pool.take();
657+
st.rebuild_snapshot();
658+
p
647659
};
648660
if let Some(p) = pool {
649661
p.close().await;

src/state.rs

Lines changed: 38 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -169,6 +169,7 @@ pub struct AppState {
169169
pub db_pool: Option<sqlx::PgPool>,
170170
/// Bitmask of allowed HTTP methods per registered path template.
171171
pub path_method_masks: Mutex<std::collections::HashMap<String, MethodMask>>,
172+
pub snapshot: Arc<FrozenState>,
172173
}
173174

174175
impl AppState {
@@ -178,9 +179,26 @@ impl AppState {
178179
"info": { "title": "OxyRoute", "version": "0.5.0" },
179180
"paths": {}
180181
});
182+
let routes = Arc::new(Vec::new());
183+
let websocket_routes = Arc::new(Vec::new());
184+
let request_middleware = Arc::new(Vec::new());
185+
let response_middleware = Arc::new(Vec::new());
186+
let exception_handlers = Arc::new(Vec::new());
187+
let snapshot = Arc::new(FrozenState {
188+
routes: Arc::clone(&routes),
189+
websocket_routes: Arc::clone(&websocket_routes),
190+
compiled: None,
191+
cors: None,
192+
security_headers: None,
193+
request_middleware: Arc::clone(&request_middleware),
194+
response_middleware: Arc::clone(&response_middleware),
195+
exception_handlers: Arc::clone(&exception_handlers),
196+
include_openapi: true,
197+
db_pool: None,
198+
});
181199
Self {
182-
routes: Arc::new(Vec::new()),
183-
websocket_routes: Arc::new(Vec::new()),
200+
routes,
201+
websocket_routes,
184202
get: Mutex::new(Router::new()),
185203
post: Mutex::new(Router::new()),
186204
put: Mutex::new(Router::new()),
@@ -192,24 +210,19 @@ impl AppState {
192210
compiled: None,
193211
frozen: false,
194212
include_openapi: true,
195-
request_middleware: Arc::new(Vec::new()),
196-
response_middleware: Arc::new(Vec::new()),
197-
exception_handlers: Arc::new(Vec::new()),
213+
request_middleware,
214+
response_middleware,
215+
exception_handlers,
198216
cors: None,
199217
security_headers: None,
200218
db_pool: None,
201219
path_method_masks: Mutex::new(std::collections::HashMap::new()),
220+
snapshot,
202221
}
203222
}
204223

205-
/// Cheap read-side snapshot of the fields the request hot path touches: the
206-
/// returned [`HotSnapshot`] is built **inside one** `state.read()` so the request
207-
/// dispatch can release the `RwLock` immediately and avoid further reads.
208-
///
209-
/// Cheap because every cloned field is `Arc::clone` / `Option<Py<PyAny>>::clone`
210-
/// (both refcount bumps), not deep clones.
211-
pub fn hot_snapshot(&self) -> HotSnapshot {
212-
HotSnapshot {
224+
pub fn rebuild_snapshot(&mut self) {
225+
self.snapshot = Arc::new(FrozenState {
213226
routes: Arc::clone(&self.routes),
214227
websocket_routes: Arc::clone(&self.websocket_routes),
215228
compiled: self.compiled.as_ref().map(Arc::clone),
@@ -220,7 +233,13 @@ impl AppState {
220233
exception_handlers: Arc::clone(&self.exception_handlers),
221234
include_openapi: self.include_openapi,
222235
db_pool: self.db_pool.clone(),
223-
}
236+
});
237+
}
238+
239+
/// Read-side snapshot of the fields the request hot path touches: only 1 atomic
240+
/// pointer clone (`Arc::clone(&self.snapshot)`).
241+
pub fn hot_snapshot(&self) -> Arc<FrozenState> {
242+
Arc::clone(&self.snapshot)
224243
}
225244

226245
/// Clone current mutex-protected [`Router`]s into a snapshot (used at freeze / tests).
@@ -242,10 +261,9 @@ impl AppState {
242261
}
243262
}
244263

245-
/// One-shot read-side view of [`AppState`] for [`run_rsgi`]. All fields are cheap to clone
246-
/// (`Arc`/`Option<Py<PyAny>>` refcount bumps) so the hot path can drop the `RwLock` after a
247-
/// single `read()`. See [`AppState::hot_snapshot`].
248-
pub struct HotSnapshot {
264+
/// Consolidated immutable read-side view of [`AppState`] for request dispatching.
265+
/// Only 1 atomic refcount increment is needed per request.
266+
pub struct FrozenState {
249267
pub routes: Arc<Vec<RouteEntry>>,
250268
pub websocket_routes: Arc<Vec<WebsocketRoute>>,
251269
pub compiled: Option<Arc<CompiledRouters>>,
@@ -258,6 +276,8 @@ pub struct HotSnapshot {
258276
pub db_pool: Option<sqlx::PgPool>,
259277
}
260278

279+
pub type HotSnapshot = Arc<FrozenState>;
280+
261281
/// Lookup a WebSocket route in a precomputed [`CompiledRouters`] (lock-free).
262282
pub fn match_ws_route_compiled<'a, 'b>(
263283
compiled: &'a CompiledRouters,

0 commit comments

Comments
 (0)