malvin 0.2.7

Non-interactive research and coding agent
use std::collections::HashSet;
use std::path::Path;
use std::str::FromStr;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;

use pi::sdk::{SessionOptions, ThinkingLevel};

use crate::acp::AgentError;
use crate::agent_backend::SdkSession;
use crate::bridge_sdk::{BridgeSpawnArgs, StreamLog};

use super::isolated_bash::isolated_tool_factory;
use super::openrouter_pricing;
use super::runtime::PiRuntime;
use super::session::PiEmbeddedSession;
use super::session_spawn_local::local_session_overrides;

type SandboxBaseline = HashSet<u32>;

fn sandbox_note_or_error(
    ticket: crate::malvin_sandbox::SandboxSpawnTicket,
    cwd: &Path,
) -> Result<SandboxBaseline, AgentError> {
    let baseline = crate::malvin_sandbox::malvin_spawn_baseline();
    crate::malvin_sandbox::note_active_sandbox_session(ticket, None, baseline.clone(), cwd)
        .map_err(AgentError)?;
    Ok(baseline)
}

fn prewarm_openrouter_pricing(provider: &str) {
    if provider.eq_ignore_ascii_case("openrouter") {
        openrouter_pricing::warm_openrouter_pricing_cache(false);
    }
}

fn spawn_live_pi_bridge(
    ticket: crate::malvin_sandbox::SandboxSpawnTicket,
    args: &BridgeSpawnArgs<'_>,
    provider: &str,
    model: &str,
) -> Result<SdkSession, AgentError> {
    prewarm_openrouter_pricing(provider);
    let options = build_session_options(args, provider, model)?;
    let runtime = PiRuntime::start(options).map_err(AgentError)?;
    let session = embedded_session(ticket, args, runtime, (provider, model))?;
    super::session_spawn_watch::start_embedded_mem_watch(&session);
    Ok(SdkSession::Pi(Box::new(session)))
}

pub(crate) async fn pi_spawn_bridge(args: BridgeSpawnArgs<'_>) -> Result<SdkSession, AgentError> {
    let ticket = crate::malvin_sandbox::take_sandbox_spawn_ticket().map_err(AgentError)?;
    let (provider, model) = args.model.pi_provider_and_model().ok_or_else(|| {
        AgentError(format!(
            "rpi model id must be `rpi:<provider>/<model>` (got `{}`)",
            args.model.canonical()
        ))
    })?;
    if test_no_real_agent() {
        return Ok(SdkSession::Pi(Box::new(fake_embedded_session(
            ticket, &args, provider, model,
        ))));
    }
    spawn_live_pi_bridge(ticket, &args, provider, model)
}

fn test_no_real_agent() -> bool {
    crate::acp::test_no_real_agent_enabled()
}

fn pi_thinking_level(thinking: &str) -> Result<ThinkingLevel, String> {
    let mapped = match thinking {
        "ultra" => "max",
        other => other,
    };
    ThinkingLevel::from_str(mapped)
}

fn ensure_local_catalog(
    cwd: &std::path::Path,
    provider: &str,
    model: &str,
) -> Result<(), AgentError> {
    let context_size = super::local_context::context_size_for_workdir(cwd);
    super::local_context::ensure_capped_local_model_catalog(provider, model, context_size)
        .map_err(AgentError)?;
    if pi::provider_metadata::provider_is_keyless_local(provider) {
        super::local_lifecycle::ensure_local_llm(provider, model).map_err(AgentError)?;
    }
    Ok(())
}

fn build_session_options(
    args: &BridgeSpawnArgs<'_>,
    provider: &str,
    model: &str,
) -> Result<SessionOptions, AgentError> {
    let thinking = args
        .thinking
        .map(pi_thinking_level)
        .transpose()
        .map_err(AgentError)?;
    ensure_local_catalog(args.cwd, provider, model)?;
    let keyless = pi::provider_metadata::provider_is_keyless_local(provider);
    let (append_system_prompt, enabled_tools, max_tool_iterations) =
        local_session_overrides(keyless, provider, model);
    Ok(SessionOptions {
        provider: Some(provider.to_string()),
        model: Some(model.to_string()),
        api_key: keyless.then(|| super::local_context::KEYLESS_LOCAL_API_KEY.to_string()),
        thinking,
        append_system_prompt,
        enabled_tools,
        working_directory: Some(args.cwd.to_path_buf()),
        no_session: true,
        extension_paths: Vec::new(),
        tool_factory: Some(isolated_tool_factory()),
        max_tool_iterations,
        ..SessionOptions::default()
    })
}

fn fake_embedded_session(
    ticket: crate::malvin_sandbox::SandboxSpawnTicket,
    args: &BridgeSpawnArgs<'_>,
    provider: &str,
    model: &str,
) -> PiEmbeddedSession {
    let baseline = crate::malvin_sandbox::malvin_spawn_baseline();
    note_sandbox_baseline(ticket, None, baseline.clone(), args.cwd);
    PiEmbeddedSession {
        runtime: None,
        log: StreamLog::from_spawn(args),
        work_dir: args.cwd.to_path_buf(),
        reader_dead: Arc::new(AtomicBool::new(false)),
        spawn_pid_baseline: baseline,
        pi_provider: provider.to_string(),
        pi_model: model.to_string(),
        local_hold: take_local_hold(provider).unwrap_or(false),
    }
}

fn embedded_session(
    ticket: crate::malvin_sandbox::SandboxSpawnTicket,
    args: &BridgeSpawnArgs<'_>,
    runtime: PiRuntime,
    model_id: (&str, &str),
) -> Result<PiEmbeddedSession, AgentError> {
    let (provider, model) = model_id;
    let baseline = sandbox_note_or_error(ticket, args.cwd)?;
    let local_hold = take_local_hold(provider)?;
    Ok(PiEmbeddedSession {
        runtime: Some(runtime),
        log: StreamLog::from_spawn(args),
        work_dir: args.cwd.to_path_buf(),
        reader_dead: Arc::new(AtomicBool::new(false)),
        spawn_pid_baseline: baseline,
        pi_provider: provider.to_string(),
        pi_model: model.to_string(),
        local_hold,
    })
}

fn take_local_hold(provider: &str) -> Result<bool, AgentError> {
    if !pi::provider_metadata::provider_is_keyless_local(provider) {
        return Ok(false);
    }
    super::local_lifecycle::hold_local_llm().map_err(AgentError)?;
    Ok(true)
}

fn note_sandbox_baseline(
    ticket: crate::malvin_sandbox::SandboxSpawnTicket,
    pgid: Option<u32>,
    baseline: SandboxBaseline,
    cwd: &Path,
) {
    let _ = crate::malvin_sandbox::note_active_sandbox_session(ticket, pgid, baseline.clone(), cwd);
}