1use std::sync::Arc;
11use std::time::Duration;
12
13use async_trait::async_trait;
14use serde_json::{Value, json};
15use tracing::{debug, instrument};
16
17use adk_core::ToolContext;
18
19use crate::error::CodeError;
20use crate::rust_executor::RustExecutor;
21
22const DEFAULT_TIMEOUT_SECS: u64 = 30;
24
25const MIN_TIMEOUT_SECS: u64 = 1;
27
28const MAX_TIMEOUT_SECS: u64 = 300;
30
31const REQUIRED_SCOPES: &[&str] = &["code:execute", "code:execute:rust"];
33
34pub struct CodeTool {
59 executor: RustExecutor,
60}
61
62impl CodeTool {
63 pub fn new(executor: RustExecutor) -> Self {
65 Self { executor }
66 }
67}
68
69fn code_error_to_json(err: &CodeError) -> Value {
74 match err {
75 CodeError::CompileError { diagnostics, stderr } => {
76 let diag_json: Vec<Value> = diagnostics
77 .iter()
78 .map(|d| {
79 json!({
80 "level": d.level,
81 "message": d.message,
82 "spans": d.spans.iter().map(|s| json!({
83 "file_name": s.file_name,
84 "line_start": s.line_start,
85 "line_end": s.line_end,
86 "column_start": s.column_start,
87 "column_end": s.column_end,
88 })).collect::<Vec<_>>(),
89 "code": d.code,
90 })
91 })
92 .collect();
93 json!({
94 "status": "compile_error",
95 "diagnostics": diag_json,
96 "stderr": stderr,
97 })
98 }
99 CodeError::DependencyNotFound { name, searched } => json!({
100 "status": "error",
101 "stderr": format!("dependency not found: {name} (searched: {searched:?})"),
102 }),
103 CodeError::Sandbox(sandbox_err) => {
104 use adk_sandbox::SandboxError;
105 match sandbox_err {
106 SandboxError::Timeout { timeout } => json!({
107 "status": "timeout",
108 "stderr": format!("execution timed out after {timeout:?}"),
109 "duration_ms": timeout.as_millis() as u64,
110 }),
111 SandboxError::MemoryExceeded { limit_mb } => json!({
112 "status": "memory_exceeded",
113 "stderr": format!("memory limit exceeded: {limit_mb} MB"),
114 }),
115 SandboxError::ExecutionFailed(msg) => json!({
116 "status": "error",
117 "stderr": msg,
118 }),
119 SandboxError::InvalidRequest(msg) => json!({
120 "status": "error",
121 "stderr": msg,
122 }),
123 SandboxError::BackendUnavailable(msg) => json!({
124 "status": "error",
125 "stderr": msg,
126 }),
127 SandboxError::EnforcerFailed { enforcer, message } => json!({
128 "status": "error",
129 "stderr": format!("sandbox enforcer '{enforcer}' failed: {message}"),
130 }),
131 SandboxError::EnforcerUnavailable { enforcer, message } => json!({
132 "status": "error",
133 "stderr": format!("sandbox enforcer '{enforcer}' unavailable: {message}"),
134 }),
135 SandboxError::PolicyViolation(msg) => json!({
136 "status": "error",
137 "stderr": msg,
138 }),
139 }
140 }
141 CodeError::InvalidCode(msg) => json!({
142 "status": "error",
143 "stderr": msg,
144 }),
145 }
146}
147
148#[async_trait]
149impl adk_core::Tool for CodeTool {
150 fn name(&self) -> &str {
151 "code_exec"
152 }
153
154 fn description(&self) -> &str {
155 "Execute Rust code through a check → build → execute pipeline. \
156 The code must provide a `fn run(input: serde_json::Value) -> serde_json::Value` \
157 entry point. Compile errors are returned as structured diagnostics."
158 }
159
160 fn required_scopes(&self) -> &[&str] {
161 REQUIRED_SCOPES
162 }
163
164 fn parameters_schema(&self) -> Option<Value> {
165 Some(json!({
166 "type": "object",
167 "properties": {
168 "language": {
169 "type": "string",
170 "enum": ["rust"],
171 "description": "The programming language. Currently only \"rust\" is supported.",
172 "default": "rust"
173 },
174 "code": {
175 "type": "string",
176 "description": "The Rust source code to execute. Must provide `fn run(input: serde_json::Value) -> serde_json::Value`."
177 },
178 "input": {
179 "type": "object",
180 "description": "Optional JSON input passed to the `run()` function via stdin."
181 },
182 "timeout_secs": {
183 "type": "integer",
184 "description": "Maximum execution time in seconds.",
185 "default": DEFAULT_TIMEOUT_SECS,
186 "minimum": MIN_TIMEOUT_SECS,
187 "maximum": MAX_TIMEOUT_SECS
188 }
189 },
190 "required": ["code"]
191 }))
192 }
193
194 #[instrument(skip_all, fields(tool = "code_exec"))]
195 async fn execute(&self, _ctx: Arc<dyn ToolContext>, args: Value) -> adk_core::Result<Value> {
196 let language = args.get("language").and_then(|v| v.as_str()).unwrap_or("rust");
198
199 if language != "rust" {
200 return Ok(json!({
201 "status": "error",
202 "stderr": format!(
203 "unsupported language \"{language}\". Only \"rust\" is currently supported."
204 ),
205 }));
206 }
207
208 let code = match args.get("code").and_then(|v| v.as_str()) {
210 Some(c) => c,
211 None => {
212 return Ok(json!({
213 "status": "error",
214 "stderr": "missing required field \"code\"",
215 }));
216 }
217 };
218
219 let input = args.get("input").cloned();
221
222 let timeout_secs = args
224 .get("timeout_secs")
225 .and_then(|v| v.as_u64())
226 .unwrap_or(DEFAULT_TIMEOUT_SECS)
227 .clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS);
228
229 let timeout = Duration::from_secs(timeout_secs);
230
231 debug!(language, timeout_secs, has_input = input.is_some(), "dispatching to RustExecutor");
232
233 match self.executor.execute(code, input.as_ref(), timeout).await {
234 Ok(result) => Ok(json!({
235 "status": "success",
236 "stdout": result.display_stdout,
237 "stderr": result.exec_result.stderr,
238 "exit_code": result.exec_result.exit_code,
239 "duration_ms": result.exec_result.duration.as_millis() as u64,
240 "output": result.output,
241 "diagnostics": result.diagnostics.iter().map(|d| json!({
242 "level": d.level,
243 "message": d.message,
244 "spans": d.spans.iter().map(|s| json!({
245 "file_name": s.file_name,
246 "line_start": s.line_start,
247 "line_end": s.line_end,
248 "column_start": s.column_start,
249 "column_end": s.column_end,
250 })).collect::<Vec<_>>(),
251 "code": d.code,
252 })).collect::<Vec<_>>(),
253 })),
254 Err(err) => Ok(code_error_to_json(&err)),
255 }
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262 use crate::diagnostics::RustDiagnostic;
263 use crate::rust_executor::RustExecutorConfig;
264 use adk_core::{CallbackContext, Content, EventActions, ReadonlyContext, Tool};
265 use adk_sandbox::SandboxBackend;
266 use adk_sandbox::backend::{BackendCapabilities, EnforcedLimits};
267 use adk_sandbox::error::SandboxError;
268 use adk_sandbox::types::{ExecRequest, ExecResult, Language};
269 use std::sync::Mutex;
270 use std::time::Duration;
271
272 struct MockBackend {
275 response: Mutex<Option<Result<ExecResult, SandboxError>>>,
276 }
277
278 impl MockBackend {
279 fn success(stdout: &str) -> Self {
280 Self {
281 response: Mutex::new(Some(Ok(ExecResult {
282 stdout: stdout.to_string(),
283 stderr: String::new(),
284 exit_code: 0,
285 duration: Duration::from_millis(10),
286 }))),
287 }
288 }
289 }
290
291 #[async_trait]
292 impl SandboxBackend for MockBackend {
293 fn name(&self) -> &str {
294 "mock"
295 }
296
297 fn capabilities(&self) -> BackendCapabilities {
298 BackendCapabilities {
299 supported_languages: vec![Language::Command],
300 isolation_class: "mock".to_string(),
301 enforced_limits: EnforcedLimits {
302 timeout: true,
303 memory: false,
304 network_isolation: false,
305 filesystem_write_isolation: false,
306 filesystem_read_isolation: false,
307 environment_isolation: false,
308 },
309 }
310 }
311
312 async fn execute(&self, _request: ExecRequest) -> Result<ExecResult, SandboxError> {
313 self.response
314 .lock()
315 .unwrap()
316 .take()
317 .unwrap_or(Err(SandboxError::ExecutionFailed("no canned response".to_string())))
318 }
319 }
320
321 struct MockToolContext {
324 content: Content,
325 actions: Mutex<EventActions>,
326 }
327
328 impl MockToolContext {
329 fn new() -> Self {
330 Self { content: Content::new("user"), actions: Mutex::new(EventActions::default()) }
331 }
332 }
333
334 #[async_trait]
335 impl ReadonlyContext for MockToolContext {
336 fn invocation_id(&self) -> &str {
337 "inv-1"
338 }
339 fn agent_name(&self) -> &str {
340 "test-agent"
341 }
342 fn user_id(&self) -> &str {
343 "user"
344 }
345 fn app_name(&self) -> &str {
346 "app"
347 }
348 fn session_id(&self) -> &str {
349 "session"
350 }
351 fn branch(&self) -> &str {
352 ""
353 }
354 fn user_content(&self) -> &Content {
355 &self.content
356 }
357 }
358
359 #[async_trait]
360 impl CallbackContext for MockToolContext {
361 fn artifacts(&self) -> Option<Arc<dyn adk_core::Artifacts>> {
362 None
363 }
364 }
365
366 #[async_trait]
367 impl ToolContext for MockToolContext {
368 fn function_call_id(&self) -> &str {
369 "call-1"
370 }
371 fn actions(&self) -> EventActions {
372 self.actions.lock().unwrap().clone()
373 }
374 fn set_actions(&self, actions: EventActions) {
375 *self.actions.lock().unwrap() = actions;
376 }
377 async fn search_memory(
378 &self,
379 _query: &str,
380 ) -> adk_core::Result<Vec<adk_core::MemoryEntry>> {
381 Ok(vec![])
382 }
383 }
384
385 fn ctx() -> Arc<dyn ToolContext> {
386 Arc::new(MockToolContext::new())
387 }
388
389 fn make_tool() -> CodeTool {
390 let backend = Arc::new(MockBackend::success(""));
391 let executor = RustExecutor::new(backend, RustExecutorConfig::default());
392 CodeTool::new(executor)
393 }
394
395 #[test]
398 fn test_name() {
399 let tool = make_tool();
400 assert_eq!(tool.name(), "code_exec");
401 }
402
403 #[test]
404 fn test_description_is_nonempty() {
405 let tool = make_tool();
406 assert!(!tool.description().is_empty());
407 }
408
409 #[test]
410 fn test_required_scopes() {
411 let tool = make_tool();
412 assert_eq!(tool.required_scopes(), &["code:execute", "code:execute:rust"]);
413 }
414
415 #[test]
416 fn test_parameters_schema_is_valid() {
417 let tool = make_tool();
418 let schema = tool.parameters_schema().expect("schema should be Some");
419 assert_eq!(schema["type"], "object");
420 assert!(schema["properties"]["language"].is_object());
421 assert!(schema["properties"]["code"].is_object());
422 assert!(schema["properties"]["input"].is_object());
423 assert!(schema["properties"]["timeout_secs"].is_object());
424
425 let required = schema["required"].as_array().unwrap();
426 let required_strs: Vec<&str> = required.iter().map(|v| v.as_str().unwrap()).collect();
427 assert!(required_strs.contains(&"code"));
428 assert!(!required_strs.contains(&"language"));
430 }
431
432 #[tokio::test]
433 async fn test_missing_code_field() {
434 let tool = make_tool();
435 let args = json!({ "language": "rust" });
436 let result = tool.execute(ctx(), args).await.unwrap();
437 assert_eq!(result["status"], "error");
438 assert!(result["stderr"].as_str().unwrap().contains("code"));
439 }
440
441 #[tokio::test]
442 async fn test_unsupported_language() {
443 let tool = make_tool();
444 let args = json!({ "language": "python", "code": "print('hi')" });
445 let result = tool.execute(ctx(), args).await.unwrap();
446 assert_eq!(result["status"], "error");
447 assert!(result["stderr"].as_str().unwrap().contains("python"));
448 assert!(result["stderr"].as_str().unwrap().contains("unsupported"));
449 }
450
451 #[tokio::test]
452 async fn test_missing_language_defaults_to_rust() {
453 let tool = make_tool();
457 let args =
458 json!({ "code": "fn run(input: serde_json::Value) -> serde_json::Value { input }" });
459 let result = tool.execute(ctx(), args).await.unwrap();
460 let status = result["status"].as_str().unwrap();
463 assert_ne!(status, "error_unsupported_language");
464 if status == "error" {
466 let stderr = result["stderr"].as_str().unwrap_or("");
467 assert!(!stderr.contains("unsupported language"));
468 }
469 }
470
471 #[test]
472 fn test_code_error_to_json_compile_error() {
473 let err = CodeError::CompileError {
474 diagnostics: vec![RustDiagnostic {
475 level: "error".to_string(),
476 message: "expected `;`".to_string(),
477 spans: vec![],
478 code: Some("E0308".to_string()),
479 }],
480 stderr: "error: expected `;`".to_string(),
481 };
482 let json = code_error_to_json(&err);
483 assert_eq!(json["status"], "compile_error");
484 assert!(json["diagnostics"].is_array());
485 assert_eq!(json["diagnostics"][0]["level"], "error");
486 assert_eq!(json["diagnostics"][0]["message"], "expected `;`");
487 assert_eq!(json["diagnostics"][0]["code"], "E0308");
488 assert_eq!(json["stderr"], "error: expected `;`");
489 }
490
491 #[test]
492 fn test_code_error_to_json_dependency_not_found() {
493 let err = CodeError::DependencyNotFound {
494 name: "serde_json".to_string(),
495 searched: vec!["config: /fake/path".to_string()],
496 };
497 let json = code_error_to_json(&err);
498 assert_eq!(json["status"], "error");
499 assert!(json["stderr"].as_str().unwrap().contains("serde_json"));
500 }
501
502 #[test]
503 fn test_code_error_to_json_sandbox_timeout() {
504 let err = CodeError::Sandbox(SandboxError::Timeout { timeout: Duration::from_secs(5) });
505 let json = code_error_to_json(&err);
506 assert_eq!(json["status"], "timeout");
507 assert!(json["stderr"].as_str().unwrap().contains("timed out"));
508 }
509
510 #[test]
511 fn test_code_error_to_json_invalid_code() {
512 let err = CodeError::InvalidCode("missing `fn run()` entry point".to_string());
513 let json = code_error_to_json(&err);
514 assert_eq!(json["status"], "error");
515 assert!(json["stderr"].as_str().unwrap().contains("fn run()"));
516 }
517
518 #[test]
519 fn test_code_error_to_json_sandbox_memory() {
520 let err = CodeError::Sandbox(SandboxError::MemoryExceeded { limit_mb: 128 });
521 let json = code_error_to_json(&err);
522 assert_eq!(json["status"], "memory_exceeded");
523 }
524
525 #[test]
526 fn test_code_error_to_json_sandbox_execution_failed() {
527 let err = CodeError::Sandbox(SandboxError::ExecutionFailed("boom".into()));
528 let json = code_error_to_json(&err);
529 assert_eq!(json["status"], "error");
530 assert_eq!(json["stderr"], "boom");
531 }
532}