malvin 0.2.9

Non-interactive research and coding agent
use std::sync::Arc;

use async_trait::async_trait;
use pi::jobs::JobSessionScope;
use pi::sdk::{Tool, ToolOutput, ToolUpdate};
use pi::tools::{ToolEffects, ToolOrigin};
use serde_json::Value;

const PATH_TOOLS: [&str; 5] = ["read", "write", "edit", "ls", "find"];
const PATH_ALIASES: [&str; 3] = ["file", "file_path", "filePath"];

pub(super) fn bind_shared_job_scope(inner: &mut Arc<dyn Tool>, scope: JobSessionScope) {
    if let Some(tool) = Arc::get_mut(inner) {
        tool.bind_job_session_scope(scope);
    }
}

pub(super) fn wrap_tool_args(tool: Arc<dyn Tool>) -> Box<dyn Tool> {
    let path_aliases = PATH_TOOLS.contains(&tool.name());
    let schema = tool.parameters();
    Box::new(NormalizedArgsTool {
        inner: tool,
        schema,
        path_aliases,
    })
}

pub(super) fn normalize_path_alias(mut input: Value) -> Value {
    let Some(obj) = input.as_object_mut() else {
        return input;
    };
    if obj.contains_key("path") {
        return input;
    }
    if let Some(v) = PATH_ALIASES.iter().find_map(|k| obj.remove(*k)) {
        obj.insert("path".to_string(), v);
    }
    input
}

fn coerce_scalar(declared: &str, s: &str) -> Option<Value> {
    let s = s.trim();
    match declared {
        "boolean" if s.eq_ignore_ascii_case("true") => Some(Value::Bool(true)),
        "boolean" if s.eq_ignore_ascii_case("false") => Some(Value::Bool(false)),
        "integer" => s.parse::<i64>().ok().map(Value::from),
        _ => None,
    }
}

pub(super) fn coerce_string_scalars(mut input: Value, schema: &Value) -> Value {
    let Some(obj) = input.as_object_mut() else {
        return input;
    };
    for (key, value) in obj.iter_mut() {
        let declared = schema["properties"][key.as_str()]["type"].as_str();
        let fixed = match (declared, value.as_str()) {
            (Some(d), Some(s)) => coerce_scalar(d, s),
            _ => None,
        };
        if let Some(v) = fixed {
            *value = v;
        }
    }
    input
}

struct NormalizedArgsTool {
    inner: Arc<dyn Tool>,
    schema: Value,
    path_aliases: bool,
}

#[async_trait]
impl Tool for NormalizedArgsTool {
    fn name(&self) -> &str {
        self.inner.name()
    }

    fn label(&self) -> &str {
        self.inner.label()
    }

    fn description(&self) -> &str {
        self.inner.description()
    }

    fn parameters(&self) -> Value {
        self.inner.parameters()
    }

    fn effects(&self) -> ToolEffects {
        self.inner.effects()
    }

    fn bind_job_session_scope(&mut self, scope: JobSessionScope) {
        bind_shared_job_scope(&mut self.inner, scope);
    }

    fn origin(&self) -> ToolOrigin {
        self.inner.origin()
    }

    async fn execute(
        &self,
        tool_call_id: &str,
        input: Value,
        on_update: Option<Box<dyn Fn(ToolUpdate) + Send + Sync>>,
    ) -> pi::sdk::Result<ToolOutput> {
        let input = if self.path_aliases {
            normalize_path_alias(input)
        } else {
            input
        };
        let input = coerce_string_scalars(input, &self.schema);
        self.inner.execute(tool_call_id, input, on_update).await
    }
}

#[cfg(test)]
#[path = "tool_args_normalize_tests.rs"]
mod tests;