fix: resolve deferred review items from PRs #883, #848, #788 (#915)

Address three deferred implementation items flagged during code review:

1. SIGHUP lock held across .await (#883): Split restart_with_addr into
   merged_router_clone() + install_listener() so the async TcpListener
   bind happens outside the mutex, eliminating lock contention risk.

2. Recursion depth limit for check_strings (#848): Cap JSON traversal
   at 32 levels to prevent stack overflow on pathological tool params.

3. Named error type for add_tokens (#788): Replace Result<(), String>
   with TokenBudgetExceeded { used, limit } for type-safe budget errors.

Co-authored-by: Claude Opus 4.6 <[email protected]>
This commit is contained in:
Illia Polosukhin
2026-03-11 07:12:45 +00:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 369741fc60
commit 8f513428f1
6 changed files with 174 additions and 75 deletions
+63 -47
View File
@@ -24,7 +24,7 @@ pub struct WebhookServerConfig {
pub struct WebhookServer {
config: WebhookServerConfig,
routes: Vec<Router>,
/// Merged router saved after start() for restart_with_addr().
/// Merged router saved after start() for restarts via `install_listener()`.
merged_router: Option<Router>,
shutdown_tx: Option<oneshot::Sender<()>>,
handle: Option<JoinHandle<()>>,
@@ -59,7 +59,7 @@ impl WebhookServer {
}
/// Bind a listener to the configured address and spawn the server task.
/// Private helper used by both start() and restart_with_addr().
/// Private helper used by `start()`.
async fn bind_and_spawn(&mut self, app: Router) -> Result<(), ChannelError> {
let listener = tokio::net::TcpListener::bind(self.config.addr)
.await
@@ -89,47 +89,49 @@ impl WebhookServer {
Ok(())
}
/// Gracefully shut down the current listener and rebind to a new address.
/// The merged router from the original `start()` call is reused.
///
/// If binding to the new address fails, the old listener remains active and
/// state is restored. This prevents a denial-of-service if the new address
/// is invalid or already in use.
pub async fn restart_with_addr(&mut self, new_addr: SocketAddr) -> Result<(), ChannelError> {
let app = self
.merged_router
.clone()
.ok_or_else(|| ChannelError::StartupFailed {
name: "webhook_server".to_string(),
reason: "restart_with_addr called before start()".to_string(),
})?;
/// Clone the merged router, if `start()` has been called.
pub fn merged_router_clone(&self) -> Option<Router> {
self.merged_router.clone()
}
// Save old state for rollback if new bind fails
let old_addr = self.config.addr;
/// Install a pre-bound listener, replacing the current one.
///
/// The caller is responsible for binding the `TcpListener` *outside* any
/// lock so that the async bind does not block other lock waiters. This
/// method only does synchronous bookkeeping plus spawning the (non-blocking)
/// server task, so it is safe to call while holding a mutex.
pub fn install_listener(
&mut self,
new_addr: SocketAddr,
listener: tokio::net::TcpListener,
app: Router,
) -> (Option<oneshot::Sender<()>>, Option<JoinHandle<()>>) {
// Capture old handles so the caller can shut them down outside the lock.
let old_shutdown_tx = self.shutdown_tx.take();
let old_handle = self.handle.take();
// Update config to new address and try to bind
self.config.addr = new_addr;
match self.bind_and_spawn(app).await {
Ok(()) => {
// New listener is running, gracefully shut down the old one
if let Some(tx) = old_shutdown_tx {
let _ = tx.send(());
}
if let Some(handle) = old_handle {
let _ = handle.await;
}
Ok(())
// Spawn the new server task (non-blocking).
let (shutdown_tx, shutdown_rx) = oneshot::channel();
self.shutdown_tx = Some(shutdown_tx);
let handle = tokio::spawn(async move {
if let Err(e) = axum::serve(listener, app)
.with_graceful_shutdown(async {
let _ = shutdown_rx.await;
tracing::debug!("Webhook server shutting down");
})
.await
{
tracing::error!("Webhook server error: {}", e);
}
Err(e) => {
// Restore old state; old listener remains active
self.config.addr = old_addr;
self.shutdown_tx = old_shutdown_tx;
self.handle = old_handle;
Err(e)
}
}
});
self.handle = Some(handle);
tracing::info!("Webhook server listening on {}", new_addr);
(old_shutdown_tx, old_handle)
}
/// Return the current bind address.
@@ -213,12 +215,21 @@ mod tests {
"First server should respond to health check"
);
// Restart on second port
let addr2 = format!("127.0.0.1:{}", port2).parse().unwrap();
server
.restart_with_addr(addr2)
// Restart on second port using two-phase approach
let addr2: SocketAddr = format!("127.0.0.1:{}", port2).parse().unwrap();
let app = server
.merged_router_clone()
.expect("Router should exist after start()");
let listener = tokio::net::TcpListener::bind(addr2)
.await
.expect("Failed to restart with new addr");
.expect("Failed to bind to new addr");
let (old_tx, old_handle) = server.install_listener(addr2, listener, app);
if let Some(tx) = old_tx {
let _ = tx.send(());
}
if let Some(handle) = old_handle {
let _ = handle.await;
}
// Assert the address changed
assert_eq!(
@@ -295,13 +306,18 @@ mod tests {
.expect("Failed to send request");
assert_eq!(response.status(), 200, "Server should be listening");
// Try to restart on an invalid address (port 0 is reserved, won't bind)
// Use port 1 which typically requires elevated privileges
// Try to restart on an invalid address (port 1 typically requires elevated privileges)
let invalid_addr: SocketAddr = "127.0.0.1:1".parse().unwrap();
// Attempt restart (should fail)
let result = server.restart_with_addr(invalid_addr).await;
assert!(result.is_err(), "Restart with invalid address should fail");
// Attempt bind (should fail); server state is untouched because we
// never call install_listener on failure.
let app = server
.merged_router_clone()
.expect("Router should exist after start()");
let result = tokio::net::TcpListener::bind(invalid_addr).await;
assert!(result.is_err(), "Bind to privileged port should fail");
// `app` is dropped — server state unchanged (rollback by construction)
drop(app);
// Verify the old address is still responding (rollback succeeded)
let response = client
+1 -1
View File
@@ -12,4 +12,4 @@ mod state;
pub use manager::ContextManager;
pub use memory::{ActionRecord, ConversationMemory, Memory};
pub use state::{JobContext, JobState, StateTransition};
pub use state::{JobContext, JobState, StateTransition, TokenBudgetExceeded};
+17 -7
View File
@@ -11,6 +11,16 @@ use uuid::Uuid;
use crate::llm::recording::HttpInterceptor;
/// Error returned when a job exceeds its token budget.
#[derive(Debug, thiserror::Error)]
#[error("Token budget exceeded: used {used} of {limit} allowed tokens")]
pub struct TokenBudgetExceeded {
/// Total tokens consumed (including the call that exceeded the budget).
pub used: u64,
/// Configured token limit for this job.
pub limit: u64,
}
/// State of a job.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
@@ -265,15 +275,15 @@ impl JobContext {
self.actual_cost += cost;
}
/// Record token usage from an LLM call. Returns an error string if the
/// token budget has been exceeded after this addition.
pub fn add_tokens(&mut self, tokens: u64) -> Result<(), String> {
/// Record token usage from an LLM call. Returns an error if the token
/// budget has been exceeded after this addition.
pub fn add_tokens(&mut self, tokens: u64) -> Result<(), TokenBudgetExceeded> {
self.total_tokens_used += tokens;
if self.max_tokens > 0 && self.total_tokens_used > self.max_tokens {
Err(format!(
"Token budget exceeded: used {} of {} allowed tokens",
self.total_tokens_used, self.max_tokens
))
Err(TokenBudgetExceeded {
used: self.total_tokens_used,
limit: self.max_tokens,
})
} else {
Ok(())
}
+41 -13
View File
@@ -799,12 +799,12 @@ async fn async_main() -> anyhow::Result<()> {
};
// Restart listener if addr changed.
// Minimize lock scope: acquire, read old addr, release, then restart.
// Two-phase approach: bind outside the lock, then swap under lock.
let mut restart_failed = false;
if let Some(ref ws_arc) = sighup_webhook_server {
let old_addr = {
let (old_addr, router) = {
let ws = ws_arc.lock().await;
ws.current_addr()
(ws.current_addr(), ws.merged_router_clone())
}; // Lock released here
if old_addr != new_addr {
@@ -813,17 +813,45 @@ async fn async_main() -> anyhow::Result<()> {
old_addr,
new_addr
);
// NOTE: Lock is held across restart_with_addr().await. This is
// acceptable because SIGHUP is infrequent and restart is fast. A full
// fix would require refactoring restart_with_addr to separate state
// mutation from async I/O.
let mut ws = ws_arc.lock().await;
match ws.restart_with_addr(new_addr).await {
Ok(()) => {
tracing::info!("SIGHUP: webhook server restarted on {}", new_addr);
match router {
Some(app) => {
// Phase 1: Bind new listener WITHOUT holding the lock.
match tokio::net::TcpListener::bind(new_addr).await {
Ok(listener) => {
// Phase 2: Swap state under lock (no await inside).
let (old_tx, old_handle) = {
let mut ws = ws_arc.lock().await;
ws.install_listener(new_addr, listener, app)
}; // Lock released here
// Phase 3: Shut down old listener outside the lock.
if let Some(tx) = old_tx {
let _ = tx.send(());
}
if let Some(handle) = old_handle {
let _ = handle.await;
}
tracing::info!(
"SIGHUP: webhook server restarted on {}",
new_addr
);
}
Err(e) => {
tracing::error!(
"SIGHUP: failed to bind to {}: {}",
new_addr,
e
);
restart_failed = true;
}
}
}
Err(e) => {
tracing::error!("SIGHUP: listener restart failed: {}", e);
None => {
tracing::error!(
"SIGHUP: cannot restart — server was never started"
);
restart_failed = true;
}
}
+49 -4
View File
@@ -197,13 +197,20 @@ impl Validator {
pub fn validate_tool_params(&self, params: &serde_json::Value) -> ValidationResult {
let mut result = ValidationResult::ok();
// Recursively check all string values in the JSON
// Recursively check all string values in the JSON.
// Depth is capped to prevent stack overflow on pathological input.
const MAX_DEPTH: usize = 32;
fn check_strings(
value: &serde_json::Value,
path: &str,
validator: &Validator,
result: &mut ValidationResult,
depth: usize,
) {
if depth > MAX_DEPTH {
return;
}
match value {
serde_json::Value::String(s) => {
let string_result = if s.is_empty() {
@@ -216,7 +223,7 @@ impl Validator {
serde_json::Value::Array(arr) => {
for (i, item) in arr.iter().enumerate() {
let child_path = format!("{path}[{i}]");
check_strings(item, &child_path, validator, result);
check_strings(item, &child_path, validator, result, depth + 1);
}
}
serde_json::Value::Object(obj) => {
@@ -226,14 +233,14 @@ impl Validator {
} else {
format!("{path}.{k}")
};
check_strings(v, &child_path, validator, result);
check_strings(v, &child_path, validator, result, depth + 1);
}
}
_ => {}
}
}
check_strings(params, "", self, &mut result);
check_strings(params, "", self, &mut result, 0);
result
}
}
@@ -423,4 +430,42 @@ mod tests {
.expect("expected forbidden content error");
assert_eq!(error.field, "metadata.tags[1]");
}
#[test]
fn test_tool_params_depth_limit_prevents_stack_overflow() {
let validator = Validator::new().forbid_pattern("evil");
// Build a deeply nested JSON object (depth > MAX_DEPTH of 32)
let mut value = serde_json::json!("evil payload");
for _ in 0..50 {
value = serde_json::json!({ "nested": value });
}
let result = validator.validate_tool_params(&value);
// The "evil payload" is beyond the depth limit so it should NOT be
// detected — the traversal stops before reaching it.
assert!(
result.is_valid,
"Strings beyond depth limit should be silently skipped, got errors: {:?}",
result.errors
);
}
#[test]
fn test_tool_params_within_depth_limit_still_validated() {
let validator = Validator::new().forbid_pattern("evil");
// Build a nested object within the depth limit
let mut value = serde_json::json!("evil payload");
for _ in 0..5 {
value = serde_json::json!({ "nested": value });
}
let result = validator.validate_tool_params(&value);
assert!(
!result.is_valid,
"Strings within depth limit should still be validated"
);
}
}
+3 -3
View File
@@ -1187,13 +1187,13 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
// TokenUsage; only respond_with_tools() usage is tracked here.
let total_tokens = output.usage.total() as u64;
if total_tokens > 0
&& let Err(msg) = self
&& let Err(err) = self
.worker
.context_manager()
.update_context(self.worker.job_id, |ctx| ctx.add_tokens(total_tokens))
.await?
{
self.worker.mark_failed(&msg).await?;
self.worker.mark_failed(&err.to_string()).await?;
}
Ok(output)
@@ -1796,7 +1796,7 @@ mod tests {
// Verify that mark_failed transitions job to Failed
worker
.mark_failed(&budget_result.unwrap_err())
.mark_failed(&budget_result.unwrap_err().to_string())
.await
.unwrap();
let ctx = worker