use std::collections::BTreeMap;
use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
use anyhow::{Context, Result, bail};
use salvor_core::Effect;
use salvor_llm::{AuthKind, Config};
use salvor_runtime::{Agent, AgentBuildError, Budgets, Pricing};
use salvor_tools::mcp::{EffectOverrides, McpServer};
use salvor_wasm::{DirGrant, WasmEngine, WasmTool, WasmToolSpec};
use serde::Deserialize;
const RECORD_PROMPTS_ENV: &str = "SALVOR_RECORD_PROMPTS";
pub const MAX_NAME_LEN: usize = 64;
fn resolve_record_prompts(per_agent: Option<bool>, env_default: Option<bool>) -> bool {
per_agent.or(env_default).unwrap_or(false)
}
fn parse_record_prompts_env(raw: Option<&str>) -> Option<bool> {
match raw
.map(|value| value.trim().to_ascii_lowercase())
.as_deref()
{
Some("1" | "true" | "yes") => Some(true),
_ => None,
}
}
fn env_record_prompts_default() -> Option<bool> {
parse_record_prompts_env(std::env::var(RECORD_PROMPTS_ENV).ok().as_deref())
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct AgentConfig {
pub model: String,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub system_prompt: Option<String>,
#[serde(default)]
pub system_prompt_path: Option<String>,
#[serde(default)]
pub llm: LlmConfig,
#[serde(default)]
pub budgets: BudgetsConfig,
#[serde(default)]
pub pricing: Option<PricingConfig>,
#[serde(default)]
pub max_response_tokens: Option<u32>,
#[serde(default)]
pub mcp_servers: Vec<McpServerConfig>,
#[serde(default)]
pub wasm_tools: Vec<WasmToolConfig>,
#[serde(default)]
pub record_prompts: Option<bool>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ApiKeyKind {
#[default]
ApiKey,
Oauth,
}
impl ApiKeyKind {
fn auth_kind(self) -> AuthKind {
match self {
ApiKeyKind::ApiKey => AuthKind::ApiKey,
ApiKeyKind::Oauth => AuthKind::Bearer,
}
}
}
#[derive(Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct LlmConfig {
pub base_url: Option<String>,
pub base_url_env: Option<String>,
pub api_key_env: Option<String>,
#[serde(default)]
pub api_key_kind: ApiKeyKind,
pub max_retries: Option<u32>,
pub timeout_seconds: Option<u64>,
}
#[derive(Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BudgetsConfig {
pub steps: Option<u64>,
pub tokens: Option<u64>,
pub cost_usd: Option<f64>,
pub wall_time_seconds: Option<f64>,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct PricingConfig {
pub input_per_mtok: f64,
pub output_per_mtok: f64,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct McpServerConfig {
#[serde(default)]
pub command: Option<String>,
#[serde(default)]
pub args: Vec<String>,
#[serde(default)]
pub env: BTreeMap<String, String>,
#[serde(default)]
pub url: Option<String>,
#[serde(default)]
pub bearer_token_env: Option<String>,
#[serde(default)]
pub effect_overrides: BTreeMap<String, Effect>,
}
impl McpServerConfig {
fn validate(&self) -> Result<()> {
match (self.command.is_some(), self.url.is_some()) {
(false, false) => {
bail!(
"an [[mcp_servers]] entry needs exactly one of `command` or `url`; neither is set"
)
}
(true, true) => {
bail!("an [[mcp_servers]] entry sets both `command` and `url`; use exactly one")
}
(true, false) => {
if self.bearer_token_env.is_some() {
bail!("`bearer_token_env` applies only to a `url` server, not a `command` one");
}
}
(false, true) => {
if !self.args.is_empty() {
bail!("`args` applies only to a `command` server, not a `url` one");
}
if !self.env.is_empty() {
bail!("`env` applies only to a `command` server, not a `url` one");
}
}
}
Ok(())
}
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct WasmToolConfig {
pub path: String,
#[serde(default)]
pub sha256: Option<String>,
pub name: String,
pub description: String,
#[serde(default)]
pub effect: Option<Effect>,
#[serde(default)]
pub input_schema: Option<String>,
#[serde(default)]
pub input_schema_path: Option<String>,
#[serde(default)]
pub limits: WasmLimitsConfig,
#[serde(default)]
pub grants: WasmGrantsConfig,
}
#[derive(Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct WasmLimitsConfig {
pub wall_time_ms: Option<u64>,
pub memory_bytes: Option<u64>,
pub fuel: Option<u64>,
}
impl WasmLimitsConfig {
fn tool_limits(&self) -> salvor_wasm::ToolLimits {
let defaults = salvor_wasm::ToolLimits::default();
salvor_wasm::ToolLimits {
wall_time_ms: self.wall_time_ms.unwrap_or(defaults.wall_time_ms),
memory_bytes: self.memory_bytes.unwrap_or(defaults.memory_bytes),
fuel: self.fuel,
}
}
}
#[derive(Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct WasmGrantsConfig {
#[serde(default)]
pub preopen: Vec<PreopenConfig>,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct PreopenConfig {
pub host: String,
pub guest: String,
pub perms: PreopenPermsConfig,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum PreopenPermsConfig {
Read,
ReadWrite,
}
impl PreopenPermsConfig {
fn grant_perms(self) -> salvor_wasm::GrantPerms {
match self {
PreopenPermsConfig::Read => salvor_wasm::GrantPerms::Read,
PreopenPermsConfig::ReadWrite => salvor_wasm::GrantPerms::ReadWrite,
}
}
}
impl WasmToolConfig {
fn validate(&self) -> Result<()> {
if self.effect.is_none() {
bail!(
"wasm tool `{}`: `effect` is required (\"read\", \"idempotent\", or \"write\") \
and has no default. The sandboxed binary gets no say in its own side-effect \
class, so a missing effect is a missing operator decision, not something to \
guess",
self.name
);
}
match (
self.input_schema.is_some(),
self.input_schema_path.is_some(),
) {
(false, false) => bail!(
"wasm tool `{}`: set exactly one of `input_schema` or `input_schema_path`; \
neither is set",
self.name
),
(true, true) => bail!(
"wasm tool `{}`: set exactly one of `input_schema` or `input_schema_path`, \
not both",
self.name
),
_ => {}
}
if let Some(inline) = &self.input_schema {
serde_json::from_str::<serde_json::Value>(inline).with_context(|| {
format!(
"wasm tool `{}`: `input_schema` is not valid JSON",
self.name
)
})?;
}
Ok(())
}
fn resolved_input_schema(&self, agent_dir: &Path) -> Result<serde_json::Value> {
if let Some(inline) = &self.input_schema {
return serde_json::from_str(inline).with_context(|| {
format!(
"wasm tool `{}`: `input_schema` is not valid JSON",
self.name
)
});
}
let rel = self
.input_schema_path
.as_ref()
.expect("validate guarantees a schema source");
let path = agent_dir.join(rel);
let text = std::fs::read_to_string(&path).with_context(|| {
format!(
"wasm tool `{}`: reading input schema file {}",
self.name,
path.display()
)
})?;
serde_json::from_str(&text).with_context(|| {
format!(
"wasm tool `{}`: input schema file {} is not valid JSON",
self.name,
path.display()
)
})
}
}
impl AgentConfig {
pub fn load(path: &Path) -> Result<Self> {
let text = std::fs::read_to_string(path)
.with_context(|| format!("reading agent file {}", path.display()))?;
let config: AgentConfig = toml::from_str(&text)
.with_context(|| format!("parsing agent file {}", path.display()))?;
config.validate()?;
Ok(config)
}
pub fn from_toml_str(text: &str) -> Result<Self> {
let config: AgentConfig =
toml::from_str(text).context("parsing agent definition as TOML")?;
config.validate()?;
Ok(config)
}
pub fn from_json_str(text: &str) -> Result<Self> {
let config: AgentConfig =
serde_json::from_str(text).context("parsing agent definition as JSON")?;
config.validate()?;
Ok(config)
}
pub fn validate(&self) -> Result<()> {
if self.system_prompt.is_some() && self.system_prompt_path.is_some() {
bail!("set only one of `system_prompt` or `system_prompt_path`, not both");
}
if let Some(name) = &self.name {
if name.trim().is_empty() {
bail!("`name`, if set, must not be empty or all whitespace");
}
let len = name.chars().count();
if len > MAX_NAME_LEN {
bail!("`name` is {len} characters, over the {MAX_NAME_LEN}-character cap");
}
}
for server in &self.mcp_servers {
server.validate()?;
}
for tool in &self.wasm_tools {
tool.validate()?;
}
Ok(())
}
fn budgets(&self) -> Budgets {
Budgets {
max_steps: self.budgets.steps,
max_tokens: self.budgets.tokens,
max_cost_usd: self.budgets.cost_usd,
max_wall_time: self.budgets.wall_time_seconds.map(Duration::from_secs_f64),
}
}
#[must_use]
pub fn client_config(&self) -> Config {
let mut config = Config::new();
let override_url = self
.llm
.base_url_env
.as_deref()
.and_then(|name| std::env::var(name).ok())
.filter(|url| !url.is_empty());
if let Some(url) = override_url {
config = config.with_base_url(url);
} else if let Some(base_url) = &self.llm.base_url {
config = config.with_base_url(base_url);
}
let key_env = self
.llm
.api_key_env
.as_deref()
.unwrap_or("ANTHROPIC_API_KEY");
if let Ok(key) = std::env::var(key_env)
&& !key.is_empty()
{
config = config.with_api_key(key);
}
config = config.with_auth_kind(self.llm.api_key_kind.auth_kind());
if let Some(max_retries) = self.llm.max_retries {
config = config.with_max_retries(max_retries);
}
if let Some(timeout) = self.llm.timeout_seconds {
config = config.with_timeout(Duration::from_secs(timeout));
}
config
}
#[must_use]
pub fn record_prompts_enabled(&self) -> bool {
resolve_record_prompts(self.record_prompts, env_record_prompts_default())
}
fn system_prompt(&self, agent_dir: &Path) -> Result<Option<String>> {
if let Some(prompt) = &self.system_prompt {
return Ok(Some(prompt.clone()));
}
if let Some(rel) = &self.system_prompt_path {
let path = agent_dir.join(rel);
let text = std::fs::read_to_string(&path)
.with_context(|| format!("reading system prompt file {}", path.display()))?;
return Ok(Some(text));
}
Ok(None)
}
}
pub async fn build_agent(
config: &AgentConfig,
agent_path: &Path,
) -> Result<(Agent, Vec<McpServer>)> {
let agent_dir = agent_path.parent().unwrap_or_else(|| Path::new("."));
let mut builder = Agent::builder().model(config.client_config(), &config.model);
if let Some(name) = &config.name {
builder = builder.name(name.clone());
}
if let Some(prompt) = config.system_prompt(agent_dir)? {
builder = builder.system_prompt(prompt);
}
let budgets = config.budgets();
if budgets.any_declared() {
builder = builder.budgets(budgets);
}
if let Some(pricing) = &config.pricing {
builder = builder.pricing(Pricing {
input_per_mtok: pricing.input_per_mtok,
output_per_mtok: pricing.output_per_mtok,
});
}
if let Some(max_tokens) = config.max_response_tokens {
builder = builder.max_response_tokens(max_tokens);
}
builder = builder.record_prompts(config.record_prompts_enabled());
let mut servers = Vec::new();
for server_config in &config.mcp_servers {
let mut overrides = EffectOverrides::new();
for (name, effect) in &server_config.effect_overrides {
overrides.insert(name.clone(), *effect);
}
let mut server = if let Some(url) = &server_config.url {
let token = server_config
.bearer_token_env
.as_deref()
.and_then(|name| std::env::var(name).ok())
.filter(|t| !t.is_empty());
McpServer::connect_http(url, token.as_deref(), &overrides)
.await
.with_context(|| format!("connecting to MCP server at `{url}`"))?
} else {
let command_name = server_config
.command
.as_deref()
.expect("validate guarantees a command when there is no url");
let mut command = tokio::process::Command::new(command_name);
command.args(&server_config.args);
for (key, value) in &server_config.env {
command.env(key, value);
}
McpServer::connect(command, &overrides)
.await
.with_context(|| format!("connecting to MCP server `{command_name}`"))?
};
for tool in server.take_tools() {
builder = builder.tool_dyn(Box::new(tool));
}
servers.push(server);
}
if !config.wasm_tools.is_empty() {
let engine = WasmEngine::new().context("initializing the wasm sandbox engine")?;
for tool_config in &config.wasm_tools {
let component_path = agent_dir.join(&tool_config.path);
let spec = WasmToolSpec {
name: tool_config.name.clone(),
description: tool_config.description.clone(),
effect: tool_config
.effect
.expect("validate (run at load) guarantees an effect"),
input_schema: tool_config.resolved_input_schema(agent_dir)?,
limits: tool_config.limits.tool_limits(),
grants: tool_config
.grants
.preopen
.iter()
.map(|preopen| DirGrant {
host: agent_dir.join(&preopen.host),
guest: preopen.guest.clone(),
perms: preopen.perms.grant_perms(),
})
.collect(),
};
let tool = WasmTool::load(
Arc::clone(&engine),
&component_path,
tool_config.sha256.as_deref(),
spec,
)
.with_context(|| {
format!(
"loading wasm tool `{}` from {}",
tool_config.name,
component_path.display()
)
})?;
builder = builder.tool_dyn(Box::new(tool));
}
}
let agent = builder.build().map_err(build_error_context)?;
Ok((agent, servers))
}
fn build_error_context(error: AgentBuildError) -> anyhow::Error {
match error {
AgentBuildError::CostBudgetWithoutPricing => anyhow::anyhow!(
"budgets.cost_usd is set but there is no [pricing] table; add pricing with input_per_mtok and output_per_mtok, or remove the cost budget"
),
other => anyhow::Error::new(other),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn record_prompts_precedence() {
assert!(resolve_record_prompts(Some(true), None));
assert!(resolve_record_prompts(None, Some(true)));
assert!(!resolve_record_prompts(None, None));
assert!(!resolve_record_prompts(Some(false), Some(true)));
}
#[test]
fn record_prompts_env_parsing() {
for on in ["1", "true", "TRUE", "Yes", " yes "] {
assert_eq!(parse_record_prompts_env(Some(on)), Some(true), "{on:?}");
}
for unset in [
None,
Some(""),
Some("0"),
Some("false"),
Some("no"),
Some("x"),
] {
assert_eq!(parse_record_prompts_env(unset), None, "{unset:?}");
}
}
#[test]
fn record_prompts_parses_from_toml() {
let absent = AgentConfig::from_toml_str("model = \"m\"\n").expect("parses");
assert_eq!(absent.record_prompts, None);
let on =
AgentConfig::from_toml_str("model = \"m\"\nrecord_prompts = true\n").expect("parses");
assert_eq!(on.record_prompts, Some(true));
let off =
AgentConfig::from_toml_str("model = \"m\"\nrecord_prompts = false\n").expect("parses");
assert_eq!(off.record_prompts, Some(false));
}
#[test]
fn name_parses_from_toml() {
let absent = AgentConfig::from_toml_str("model = \"m\"\n").expect("parses");
assert_eq!(absent.name, None);
let named = AgentConfig::from_toml_str("model = \"m\"\nname = \"support-triage\"\n")
.expect("parses");
assert_eq!(named.name.as_deref(), Some("support-triage"));
}
#[test]
fn blank_name_is_rejected() {
for blank in ["", " ", "\t"] {
let error = AgentConfig::from_toml_str(&format!("model = \"m\"\nname = \"{blank}\"\n"))
.expect_err("blank name should be rejected");
assert!(format!("{error:#}").contains("empty or all whitespace"));
}
}
#[test]
fn oversized_name_is_rejected() {
let long_name = "a".repeat(MAX_NAME_LEN + 1);
let toml = format!("model = \"m\"\nname = \"{long_name}\"\n");
let error = AgentConfig::from_toml_str(&toml).expect_err("oversized name rejected");
let message = format!("{error:#}");
assert!(message.contains("65 characters"), "{message}");
assert!(
message.contains(&format!("{MAX_NAME_LEN}-character cap")),
"{message}"
);
}
#[test]
fn name_exactly_at_the_cap_is_valid() {
let name = "a".repeat(MAX_NAME_LEN);
let toml = format!("model = \"m\"\nname = \"{name}\"\n");
let config = AgentConfig::from_toml_str(&toml).expect("parses and validates");
assert_eq!(config.name.as_deref(), Some(name.as_str()));
}
}