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;

use super::super::tool_args_normalize::bind_shared_job_scope;

pub(super) struct CompleteWrite {
    inner: Arc<dyn Tool>,
}

impl CompleteWrite {
    pub(super) fn from_builtin(inner: Arc<dyn Tool>) -> Self {
        Self { inner }
    }
}

fn write_field_str<'a>(input: &'a Value, keys: &[&str]) -> &'a str {
    keys.iter()
        .find_map(|k| input.get(*k).and_then(Value::as_str))
        .unwrap_or("")
}

fn path_looks_like_bin(path: &str) -> bool {
    let normalized = path.replace('\\', "/");
    normalized.contains("/bin/") || normalized.starts_with("bin/")
}

pub(super) fn stubby_write_error(input: &Value) -> Option<String> {
    let path = write_field_str(input, &["path", "file"]);
    if !path_looks_like_bin(path) {
        return None;
    }
    let trimmed = write_field_str(input, &["content", "text", "contents"]).trim();
    if trimmed.len() < 120 {
        return Some(
            "write of a bin/ script is too short; include the complete file body, not a stub"
                .to_string(),
        );
    }
    if trimmed.contains("...") && trimmed.lines().count() < 12 {
        return Some(
            "write of a bin/ script looks stubbed; include the complete file body".to_string(),
        );
    }
    None
}

fn push_escaped(out: &mut String, next: Option<char>) -> bool {
    match next {
        Some('n') => out.push('\n'),
        Some('t') => out.push('\t'),
        Some('\\') => out.push('\\'),
        Some('"') => out.push('"'),
        _ => {
            out.push('\\');
            return false;
        }
    }
    true
}

pub(super) fn unescape_local_write_content(content: &str) -> String {
    if content.contains('\n') || content.contains('\r') {
        return content.to_string();
    }
    if !(content.contains("\\n") || content.contains("\\t")) {
        return content.to_string();
    }
    let mut out = String::with_capacity(content.len());
    let mut chars = content.chars().peekable();
    while let Some(c) = chars.next() {
        if c == '\\' {
            let next = chars.peek().copied();
            if push_escaped(&mut out, next) {
                chars.next();
            }
        } else {
            out.push(c);
        }
    }
    out
}

pub(super) fn normalize_write_input(mut input: Value) -> Value {
    let Some(obj) = input.as_object_mut() else {
        return input;
    };
    for key in ["content", "text", "contents"] {
        if let Some(Value::String(s)) = obj.get(key).cloned() {
            let fixed = unescape_local_write_content(&s);
            if fixed != s {
                obj.insert(key.to_string(), Value::String(fixed));
            }
            break;
        }
    }
    Value::Object(obj.clone())
}

#[async_trait]
impl Tool for CompleteWrite {
    fn name(&self) -> &'static str {
        "write"
    }

    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 = normalize_write_input(input);
        if let Some(msg) = stubby_write_error(&input) {
            return Err(pi::error::Error::validation(msg));
        }
        self.inner.execute(tool_call_id, input, on_update).await
    }
}