use std::cmp::Ordering;
use std::collections::HashMap;
use convert_case::{Case, Casing};
use pyo3::prelude::*;
use pyo3::types::{IntoPyDict, PyDict};
use sapiens::tools::{
invoke_simple_from_toolbox, AdvancedTool, Describe, ProtoToolDescribe, ProtoToolInvoke,
ToolDescription, ToolUseError, Toolbox,
};
use sapiens_derive::{Describe, ProtoToolDescribe};
use serde::{Deserialize, Serialize};
use serde_yaml::Value;
pub(crate) mod utils;
use crate::python::utils::SimpleToolDescription;
const MAX_OUTPUT_SIZE: usize = 512;
#[derive(Debug, Default, ProtoToolDescribe)]
#[tool(
name = "SandboxedPython",
input = "PythonToolInput",
output = "PythonToolOutput"
)]
pub struct PythonTool {}
#[derive(Debug, Serialize, Deserialize, Describe)]
pub struct PythonToolInput {
pub code: String,
}
#[derive(Serialize, Deserialize, Describe)]
pub struct PythonToolOutput {
pub stdout: String,
pub stderr: String,
}
#[pyclass]
#[derive(Default)]
struct Logging {
output: String,
}
#[pymethods]
impl Logging {
fn write(&mut self, data: &str) {
self.output.push_str(data);
}
}
#[pyclass(unsendable)]
struct ToolsWrapper {
toolbox: Toolbox,
tool_list: Vec<SimpleToolDescription>,
}
impl ToolsWrapper {
async fn new(toolbox: Toolbox) -> Self {
let tools = toolbox.describe().await;
let tool_list = tools
.into_values()
.map(SimpleToolDescription::from)
.collect::<Vec<_>>();
ToolsWrapper { toolbox, tool_list }
}
}
#[pymethods]
impl ToolsWrapper {
#[pyo3(signature = ())]
fn list(&self, py: Python<'_>) -> PyResult<PyObject> {
let tools = self.tool_list.to_object(py);
Ok(tools)
}
#[pyo3(signature = (tool_name, input))]
fn invoke(
&self,
py: Python<'_>,
tool_name: &str,
input: Option<&PyDict>,
) -> PyResult<PyObject> {
let input = if let Some(input) = input {
let input: PyObject = input.into();
utils::to_yaml(py, &input).map_err(|e| {
pyo3::exceptions::PyException::new_err(format!("Invalid input: {}", e))
})?
} else {
Value::default()
};
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let (tx, mut rx) = tokio::sync::oneshot::channel::<Result<Value, ToolUseError>>();
let toolbox = self.toolbox.clone();
let tool_name = tool_name.to_string();
std::thread::spawn(move || {
rt.block_on(async move {
let output = invoke_simple_from_toolbox(toolbox, &tool_name, input).await;
match output {
Ok(output) => {
tx.send(Ok(output)).unwrap();
}
Err(e) => {
tx.send(Err(e)).unwrap();
}
}
});
});
loop {
if let Ok(output) = rx.try_recv() {
let output = output.map_err(|e| {
pyo3::exceptions::PyException::new_err(format!("Tool invocation failed: {}", e))
})?;
let output = utils::value_to_object(output, py);
return Ok(output);
}
}
}
}
impl PythonTool {
#[tracing::instrument(skip(self, toolbox))]
async fn invoke_typed(
&self,
toolbox: Toolbox,
input: &PythonToolInput,
) -> Result<PythonToolOutput, ToolUseError> {
let code = input.code.clone();
let tools = toolbox.describe().await;
let code = Self::transform_code(code, tools)?;
let toolwrapper = ToolsWrapper::new(toolbox).await;
let res: PyResult<(String, String)> = Python::with_gil(|py| {
let tools_cell = PyCell::new(py, toolwrapper)?;
let globals = [("toolbox", tools_cell)].into_py_dict(py);
let sys = py.import("sys")?;
let stdout = Logging::default();
let py_stdout_cell = PyCell::new(py, stdout)?;
let py_stdout = py_stdout_cell.borrow_mut();
sys.setattr("stdout", py_stdout.into_py(py))?;
let stderr = Logging::default();
let py_stderr_cell = PyCell::new(py, stderr)?;
let py_stderr = py_stderr_cell.borrow_mut();
sys.setattr("stderr", py_stderr.into_py(py))?;
Python::run(py, &code, globals.into(), None)?;
let stdout = py_stdout_cell.borrow().output.clone();
let stderr = py_stderr_cell.borrow().output.clone();
Ok((stdout, stderr))
});
let (stdout, stderr) = res.map_err(|e| {
ToolUseError::InvocationFailed(format!("Python code execution failed: {}", e))
})?;
Ok(PythonToolOutput { stdout, stderr })
}
fn transform_code(
code: String,
tools: HashMap<String, ToolDescription>,
) -> Result<String, ToolUseError> {
lazy_static::lazy_static! {
static ref EXEC_RE: regex::Regex = regex::Regex::new(r"(exec|pip)").unwrap();
static ref IMPORT_RE: regex::Regex = regex::Regex::new(r"(?x)import \s+ tools.*").unwrap();
static ref FROM_RE: regex::Regex =
regex::Regex::new(r"(?x)from \s+ tools \s+ import .*").unwrap();
}
if let Some(caps) = EXEC_RE.captures(code.as_ref()) {
return Err(ToolUseError::InvocationFailed(format!(
"Python code contains forbidden keywords such as {}",
caps.get(0).unwrap().as_str()
)));
}
let code = code
.lines()
.filter(|&l| !IMPORT_RE.is_match(l))
.filter(|&l| !FROM_RE.is_match(l))
.collect::<Vec<_>>()
.join("\n");
let mut tool_class_code = String::new();
tool_class_code.push_str("class Tools:\n");
tool_class_code.push_str(" def __init__(self, toolbox):\n");
tool_class_code.push_str(" self.toolbox = toolbox\n");
let mut binding_code = String::new();
for (name, description) in tools {
let inputs_parts = description.input_format.fields;
let mut inputs = inputs_parts.clone();
inputs.sort_by(|a, b| {
if a.optional && !b.optional {
Ordering::Greater
} else if !a.optional && b.optional {
Ordering::Less
} else {
Ordering::Equal
}
});
let inputs = inputs
.into_iter()
.map(|f| {
if f.optional {
format!("{}=None", f.name)
} else {
f.name
}
})
.collect::<Vec<_>>();
let inputs = inputs.join(", ");
let inputs = if inputs.is_empty() {
"(self)".to_string()
} else {
format!("(self, {})", inputs)
};
let dict = inputs_parts
.iter()
.map(|f| {
let name = &f.name;
format!("\"{}\": {}", name, name)
})
.collect::<Vec<_>>()
.join(", ");
tool_class_code.push_str(&format!(
" def {}{}:\n return self.toolbox.invoke(\"{}\", {{{}}})\n",
name.to_case(Case::Pascal),
inputs,
name,
dict
));
binding_code.push_str(&format!(
"def {}{}:\n return tools.toolbox.invoke(\"{}\", {{{}}})\n",
name, inputs, name, dict
));
}
tool_class_code.push_str(" def list(self):\n");
tool_class_code.push_str(" return self.toolbox.list()\n");
tool_class_code.push_str("tools = Tools(toolbox)\n");
let code_to_prepend = format!("{}{}", tool_class_code, binding_code);
let code = format!("{}\n# ======== user code\n{}", code_to_prepend, code);
Ok(code)
}
#[tracing::instrument(skip(self))]
fn invoke_sync_typed(&self, input: &PythonToolInput) -> Result<PythonToolOutput, ToolUseError> {
let code = input.code.clone();
let re = regex::Regex::new(r"(exec|pip)").unwrap();
if let Some(caps) = re.captures(&code) {
return Err(ToolUseError::InvocationFailed(format!(
"Python code contains forbidden keywords such as {}",
caps.get(0).unwrap().as_str()
)));
}
let res: PyResult<(String, String)> = Python::with_gil(|py| {
let globals = PyDict::new(py);
let sys = py.import("sys")?;
let stdout = Logging::default();
let py_stdout_cell = PyCell::new(py, stdout)?;
let py_stdout = py_stdout_cell.borrow_mut();
sys.setattr("stdout", py_stdout.into_py(py))?;
let stderr = Logging::default();
let py_stderr_cell = PyCell::new(py, stderr)?;
let py_stderr = py_stderr_cell.borrow_mut();
sys.setattr("stderr", py_stderr.into_py(py))?;
Python::run(py, &code, globals.into(), None)?;
let stdout = py_stdout_cell.borrow().output.clone();
let stderr = py_stderr_cell.borrow().output.clone();
Ok((stdout, stderr))
});
let (stdout, stderr) = res.map_err(|e| {
ToolUseError::InvocationFailed(format!("Python code execution failed: {}", e))
})?;
Ok(PythonToolOutput { stdout, stderr })
}
}
#[async_trait::async_trait]
impl ProtoToolInvoke for PythonTool {
async fn invoke(&self, input: serde_yaml::Value) -> Result<serde_yaml::Value, ToolUseError> {
let input = serde_yaml::from_value(input)?;
let output = self.invoke_sync_typed(&input)?;
let l = output.stdout.len() + output.stderr.len();
if l > MAX_OUTPUT_SIZE {
return Err(ToolUseError::InvocationFailed(format!(
"Python code produced too much output on stdout and stderr
combined ({} bytes) - max is {}",
l, MAX_OUTPUT_SIZE
)));
}
Ok(serde_yaml::to_value(output)?)
}
}
#[async_trait::async_trait]
impl AdvancedTool for PythonTool {
async fn invoke_with_toolbox(
&self,
toolbox: Toolbox,
input: Value,
) -> Result<Value, ToolUseError> {
let input = serde_yaml::from_value(input)?;
let output = self.invoke_typed(toolbox, &input).await?;
Ok(serde_yaml::to_value(output)?)
}
}
#[cfg(test)]
mod tests {
use indoc::indoc;
use insta::assert_display_snapshot;
use sapiens::tools::Toolbox;
use crate::conclude::ConcludeTool;
use crate::python::{PythonTool, PythonToolInput};
#[tokio::test]
async fn test_code_transformation() {
let input = PythonToolInput {
code: indoc! {
r#"
import tools
from tools import Arxiv
arxiv_results = Arxiv(
search_query='cat:cs.AI',
max_results=5,
sort_by='lastUpdatedDate',
sort_order='descending',
show_summary=True
)
formatted_results = []
for result in arxiv_results['result']:
formatted_results.append(f"{result['title']} : {result['pdf_url']}")
formatted_results = "\n".join(formatted_results)
print({'formatted_results': formatted_results})
"#}
.to_string(),
};
let mut toolbox = Toolbox::default();
toolbox.add_terminal_tool(ConcludeTool::default()).await;
let tools = toolbox.describe().await;
let code = PythonTool::transform_code(input.code, tools).unwrap();
assert_display_snapshot!(code);
}
}