mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-09 00:05:48 +08:00
fix: normalize integer-valued shell timeouts
This commit is contained in:
@@ -39,6 +39,9 @@ pub(super) async fn start(
|
|||||||
context: &ExecContext,
|
context: &ExecContext,
|
||||||
store: Option<&Store>,
|
store: Option<&Store>,
|
||||||
) -> Result<ToolStart> {
|
) -> Result<ToolStart> {
|
||||||
|
let normalized_call = normalize_block_until_ms(call)?;
|
||||||
|
let call = normalized_call.as_ref().unwrap_or(call);
|
||||||
|
|
||||||
if let Some(definition) = dynamic_mcp.get(&call.name) {
|
if let Some(definition) = dynamic_mcp.get(&call.name) {
|
||||||
return exec::start_dynamic(runtime, call, definition, context).await;
|
return exec::start_dynamic(runtime, call, definition, context).await;
|
||||||
}
|
}
|
||||||
@@ -64,6 +67,55 @@ pub(super) async fn start(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn normalize_block_until_ms(call: &ToolCall) -> Result<Option<ToolCall>> {
|
||||||
|
if !matches!(normalized(&call.name).as_str(), "shell" | "awaitshell") {
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
|
let Some(value) = call.arguments.get("block_until_ms") else {
|
||||||
|
return Ok(None);
|
||||||
|
};
|
||||||
|
|
||||||
|
let integer = if let Some(value) = value.as_i64() {
|
||||||
|
value
|
||||||
|
} else {
|
||||||
|
let value = value.as_f64().ok_or_else(|| {
|
||||||
|
Error::Protocol(format!("{} block_until_ms must be an integer", call.name))
|
||||||
|
})?;
|
||||||
|
if !value.is_finite() || value.fract() != 0.0 {
|
||||||
|
return Err(Error::Protocol(format!(
|
||||||
|
"{} block_until_ms must be an integer",
|
||||||
|
call.name
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
if value < i64::MIN as f64 || value > i64::MAX as f64 {
|
||||||
|
return Err(Error::Protocol(format!(
|
||||||
|
"{} block_until_ms is out of range",
|
||||||
|
call.name
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
value as i64
|
||||||
|
};
|
||||||
|
|
||||||
|
if integer < 0 {
|
||||||
|
return Err(Error::Protocol(format!(
|
||||||
|
"{} block_until_ms is out of range",
|
||||||
|
call.name
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|
||||||
|
if value.as_i64().is_some() {
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut normalized_call = call.clone();
|
||||||
|
normalized_call
|
||||||
|
.arguments
|
||||||
|
.as_object_mut()
|
||||||
|
.ok_or_else(|| Error::Protocol(format!("{} arguments must be a JSON object", call.name)))?
|
||||||
|
.insert("block_until_ms".into(), serde_json::Value::from(integer));
|
||||||
|
Ok(Some(normalized_call))
|
||||||
|
}
|
||||||
|
|
||||||
fn is_mcp_auth(call: &ToolCall) -> bool {
|
fn is_mcp_auth(call: &ToolCall) -> bool {
|
||||||
normalized(&call.name) == "callmcptool"
|
normalized(&call.name) == "callmcptool"
|
||||||
&& call
|
&& call
|
||||||
@@ -89,3 +141,65 @@ pub(super) fn normalized(name: &str) -> String {
|
|||||||
.flat_map(char::to_lowercase)
|
.flat_map(char::to_lowercase)
|
||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
fn tool(name: &str, arguments: serde_json::Value) -> ToolCall {
|
||||||
|
ToolCall {
|
||||||
|
index: 0,
|
||||||
|
call_id: "call-1".into(),
|
||||||
|
model_call_id: "model-call-1".into(),
|
||||||
|
name: name.into(),
|
||||||
|
arguments_text: arguments.to_string(),
|
||||||
|
arguments,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn shell_accepts_integer_valued_float_timeout() {
|
||||||
|
let call = tool("Shell", serde_json::json!({
|
||||||
|
"command": "echo ok",
|
||||||
|
"block_until_ms": 45_000.0
|
||||||
|
}));
|
||||||
|
|
||||||
|
let call = normalize_block_until_ms(&call).unwrap().unwrap();
|
||||||
|
|
||||||
|
assert_eq!(call.arguments["block_until_ms"].as_i64(), Some(45_000));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn await_shell_accepts_integer_valued_scientific_timeout() {
|
||||||
|
let arguments = serde_json::from_str(r#"{"shell_id":"1","block_until_ms":3e4}"#).unwrap();
|
||||||
|
let call = tool("AwaitShell", arguments);
|
||||||
|
|
||||||
|
let call = normalize_block_until_ms(&call).unwrap().unwrap();
|
||||||
|
|
||||||
|
assert_eq!(call.arguments["block_until_ms"].as_u64(), Some(30_000));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn shell_rejects_fractional_timeout() {
|
||||||
|
let call = tool("Shell", serde_json::json!({
|
||||||
|
"command": "echo ok",
|
||||||
|
"block_until_ms": 30_000.5
|
||||||
|
}));
|
||||||
|
|
||||||
|
let error = normalize_block_until_ms(&call).unwrap_err();
|
||||||
|
|
||||||
|
assert_eq!(error.to_string(), "protocol error: Shell block_until_ms must be an integer");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn await_shell_rejects_negative_timeout_instead_of_defaulting() {
|
||||||
|
let call = tool("AwaitShell", serde_json::json!({
|
||||||
|
"shell_id": "1",
|
||||||
|
"block_until_ms": -1
|
||||||
|
}));
|
||||||
|
|
||||||
|
let error = normalize_block_until_ms(&call).unwrap_err();
|
||||||
|
|
||||||
|
assert_eq!(error.to_string(), "protocol error: AwaitShell block_until_ms is out of range");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user