diff --git a/bin/pudox-run/src/main.rs b/bin/pudox-run/src/main.rs index 8f1f1af..e04dfaa 100644 --- a/bin/pudox-run/src/main.rs +++ b/bin/pudox-run/src/main.rs @@ -2,8 +2,9 @@ use async_openai::config::OpenAIConfig; use async_openai::types::chat::{ ChatCompletionMessageToolCalls, ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage, ChatCompletionRequestToolMessage, - ChatCompletionRequestUserMessageArgs, ChatCompletionResponseMessage, ChatCompletionTool, - ChatCompletionTools, CreateChatCompletionRequestArgs, FunctionObjectArgs, + ChatCompletionRequestToolMessageArgs, ChatCompletionRequestUserMessageArgs, + ChatCompletionResponseMessage, ChatCompletionTool, ChatCompletionTools, + CreateChatCompletionRequestArgs, FunctionObjectArgs, }; use async_openai::Client; @@ -75,6 +76,44 @@ async fn process_task(task: &Task) -> Result = Vec::new(); + tools.push(ChatCompletionTools::Function(ChatCompletionTool { + function: FunctionObjectArgs::default() + .name("read") + .description("Read the contents of a file.") + .parameters(serde_json::json!({ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the file to read, relative to the working directory." + } + }, + "required": ["path"] + })) + .build()?, + })); + + tools.push(ChatCompletionTools::Function(ChatCompletionTool { + function: FunctionObjectArgs::default() + .name("write") + .description("Write content to a file, creating it if it does not exist.") + .parameters(serde_json::json!({ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the file to write, relative to the working directory." + }, + "content": { + "type": "string", + "description": "The content to write to the file." + } + }, + "required": ["path", "content"] + })) + .build()?, + })); + tools.push(ChatCompletionTools::Function(ChatCompletionTool { function: FunctionObjectArgs::default() .name("finish") @@ -175,25 +214,102 @@ enum ToolCallResult { async fn handle_tool_calls( message: &ChatCompletionResponseMessage, ) -> Result> { - if let Some(tool_calls) = &message.tool_calls { - for tool_call in tool_calls { - match tool_call { - ChatCompletionMessageToolCalls::Function(f) => { - tracing::info!( - "Dispatching tool call: {} args={}", - f.function.name, - f.function.arguments - ); - if f.function.name == "finish" { - return Ok(ToolCallResult::Finish); + let Some(tool_calls) = &message.tool_calls else { + return Ok(ToolCallResult::Continue); + }; + + let mut results = Vec::new(); + + for tool_call in tool_calls { + match tool_call { + ChatCompletionMessageToolCalls::Function(f) => { + tracing::info!( + "Dispatching tool: {} args={}", + f.function.name, + f.function.arguments + ); + + match f.function.name.as_str() { + "finish" => return Ok(ToolCallResult::Finish), + + "read" => { + let output = tool_read(&f.function.arguments); + tracing::info!("read: {} chars returned", output.len()); + results.push( + ChatCompletionRequestToolMessageArgs::default() + .tool_call_id(f.id.clone()) + .content(output) + .build()?, + ); + } + + "write" => { + let output = tool_write(&f.function.arguments); + tracing::info!("write: {}", output); + results.push( + ChatCompletionRequestToolMessageArgs::default() + .tool_call_id(f.id.clone()) + .content(output) + .build()?, + ); + } + + name => { + tracing::warn!("Unknown tool: {}", name); + results.push( + ChatCompletionRequestToolMessageArgs::default() + .tool_call_id(f.id.clone()) + .content(format!("Unknown tool: {}", name)) + .build()?, + ); } } - _ => {} } + _ => {} } } - Ok(ToolCallResult::Continue) + if results.is_empty() { + Ok(ToolCallResult::Continue) + } else { + Ok(ToolCallResult::Results(results)) + } +} + +fn tool_read(arguments: &str) -> String { + let args: serde_json::Value = match serde_json::from_str(arguments) { + Ok(v) => v, + Err(e) => return format!("Error parsing arguments: {}", e), + }; + let path = match args["path"].as_str() { + Some(p) => p, + None => return "Missing required argument: path".to_string(), + }; + tracing::info!("read path={}", path); + match std::fs::read_to_string(path) { + Ok(content) => content, + Err(e) => format!("Error reading {}: {}", path, e), + } +} + +fn tool_write(arguments: &str) -> String { + let args: serde_json::Value = match serde_json::from_str(arguments) { + Ok(v) => v, + Err(e) => return format!("Error parsing arguments: {}", e), + }; + let path = match args["path"].as_str() { + Some(p) => p, + None => return "Missing required argument: path".to_string(), + }; + let content = match args["content"].as_str() { + Some(c) => c, + None => return "Missing required argument: content".to_string(), + }; + tracing::info!("write path={} ({} bytes)", path, content.len()); + match std::fs::write(path, content) { + Ok(_) => format!("Wrote {} bytes to {}", content.len(), path), + Err(e) => format!("Error writing {}: {}", path, e), + } } async fn report_task_result(result: &TaskResult) -> Result<(), Box> {