refactor: Modified the naming of several generalized supervisor values
This commit is contained in:
@@ -454,8 +454,6 @@ impl Agent {
|
||||
output.push_str(DEFAULT_SPAWN_INSTRUCTIONS);
|
||||
}
|
||||
|
||||
// Job declarations are appended at init iff jobs are enabled for this
|
||||
// agent, so their presence doubles as the jobs_enabled predicate.
|
||||
if self
|
||||
.functions
|
||||
.declarations()
|
||||
|
||||
+13
-13
@@ -844,8 +844,8 @@ mod tests {
|
||||
|
||||
unsafe {
|
||||
match prev {
|
||||
Some(v) => std::env::set_var(&env_name, v),
|
||||
None => std::env::remove_var(&env_name),
|
||||
Some(v) => env::set_var(&env_name, v),
|
||||
None => env::remove_var(&env_name),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -868,35 +868,35 @@ mod tests {
|
||||
#[serial_test::serial]
|
||||
fn load_envs_overrides_max_concurrent_jobs() {
|
||||
let env_name = get_env_name("max_concurrent_jobs");
|
||||
let prev = std::env::var_os(&env_name);
|
||||
let prev = env::var_os(&env_name);
|
||||
|
||||
let mut app = AppConfig::default();
|
||||
|
||||
unsafe { std::env::set_var(&env_name, "7") };
|
||||
unsafe { env::set_var(&env_name, "7") };
|
||||
app.load_envs();
|
||||
assert_eq!(app.max_concurrent_jobs, Some(7));
|
||||
|
||||
unsafe { std::env::set_var(&env_name, "0") };
|
||||
unsafe { env::set_var(&env_name, "0") };
|
||||
app.load_envs();
|
||||
assert_eq!(app.max_concurrent_jobs, Some(0));
|
||||
|
||||
unsafe { std::env::remove_var(&env_name) };
|
||||
unsafe { env::remove_var(&env_name) };
|
||||
app.max_concurrent_jobs = Some(2);
|
||||
app.load_envs();
|
||||
assert_eq!(app.max_concurrent_jobs, Some(2));
|
||||
|
||||
unsafe {
|
||||
match prev {
|
||||
Some(v) => std::env::set_var(&env_name, v),
|
||||
None => std::env::remove_var(&env_name),
|
||||
Some(v) => env::set_var(&env_name, v),
|
||||
None => env::remove_var(&env_name),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn editor_returns_configured_value() {
|
||||
let configured = cached_editor()
|
||||
.unwrap_or_else(|| std::env::current_exe().unwrap().display().to_string());
|
||||
let configured =
|
||||
cached_editor().unwrap_or_else(|| env::current_exe().unwrap().display().to_string());
|
||||
let app = AppConfig {
|
||||
editor: Some(configured.clone()),
|
||||
..AppConfig::default()
|
||||
@@ -913,9 +913,9 @@ mod tests {
|
||||
return;
|
||||
}
|
||||
|
||||
let expected = std::env::current_exe().unwrap().display().to_string();
|
||||
let expected = env::current_exe().unwrap().display().to_string();
|
||||
unsafe {
|
||||
std::env::set_var("VISUAL", &expected);
|
||||
env::set_var("VISUAL", &expected);
|
||||
}
|
||||
|
||||
let app = AppConfig::default();
|
||||
@@ -983,7 +983,7 @@ mod tests {
|
||||
let app = AppConfig::from_config(cfg).unwrap();
|
||||
|
||||
let ua = app.user_agent.as_deref().unwrap();
|
||||
assert!(ua != "auto", "user_agent should have been resolved");
|
||||
assert_ne!(ua, "auto", "user_agent should have been resolved");
|
||||
assert!(ua.contains('/'), "user_agent should be '<name>/<version>'");
|
||||
}
|
||||
|
||||
|
||||
@@ -73,6 +73,7 @@ impl AppState {
|
||||
if !mcp_registry.is_empty() && config.mcp_server_support {
|
||||
functions.append_mcp_meta_functions(mcp_registry.server_features());
|
||||
}
|
||||
|
||||
if jobs_enabled(None, &config) {
|
||||
functions.append_job_functions();
|
||||
}
|
||||
|
||||
@@ -163,7 +163,6 @@ impl Input {
|
||||
self.data_urls.clone()
|
||||
}
|
||||
|
||||
/// Names of the function declarations this request will send to the model.
|
||||
pub fn declared_function_names(&self) -> HashSet<String> {
|
||||
self.functions
|
||||
.as_ref()
|
||||
|
||||
@@ -334,7 +334,7 @@ pub struct RequestContext {
|
||||
pub notification_queue: Arc<NotificationQueue>,
|
||||
pub current_depth: usize,
|
||||
pub auto_continue_count: usize,
|
||||
pub pending_agents_guardrail_count: u32,
|
||||
pub pending_tasks_guardrail_count: u32,
|
||||
pub todo_list: TodoList,
|
||||
pub skill_registry: SkillRegistry,
|
||||
pub last_continuation_response: Option<String>,
|
||||
@@ -369,7 +369,7 @@ impl RequestContext {
|
||||
notification_queue: Arc::new(NotificationQueue::new()),
|
||||
current_depth: 0,
|
||||
auto_continue_count: 0,
|
||||
pending_agents_guardrail_count: 0,
|
||||
pending_tasks_guardrail_count: 0,
|
||||
todo_list: TodoList::default(),
|
||||
skill_registry: SkillRegistry::default(),
|
||||
last_continuation_response: None,
|
||||
@@ -430,7 +430,7 @@ impl RequestContext {
|
||||
notification_queue: Arc::new(NotificationQueue::new()),
|
||||
current_depth: 0,
|
||||
auto_continue_count: 0,
|
||||
pending_agents_guardrail_count: 0,
|
||||
pending_tasks_guardrail_count: 0,
|
||||
todo_list: TodoList::default(),
|
||||
skill_registry: SkillRegistry::default(),
|
||||
last_continuation_response: None,
|
||||
@@ -478,7 +478,7 @@ impl RequestContext {
|
||||
notification_queue: self.notification_queue.clone(),
|
||||
current_depth: self.current_depth,
|
||||
auto_continue_count: 0,
|
||||
pending_agents_guardrail_count: 0,
|
||||
pending_tasks_guardrail_count: 0,
|
||||
todo_list: self.todo_list.clone(),
|
||||
skill_registry: self.skill_registry.clone(),
|
||||
last_continuation_response: None,
|
||||
@@ -524,7 +524,7 @@ impl RequestContext {
|
||||
notification_queue: Arc::new(NotificationQueue::new()),
|
||||
current_depth,
|
||||
auto_continue_count: 0,
|
||||
pending_agents_guardrail_count: 0,
|
||||
pending_tasks_guardrail_count: 0,
|
||||
todo_list: TodoList::default(),
|
||||
skill_registry: SkillRegistry::default(),
|
||||
last_continuation_response: None,
|
||||
@@ -923,6 +923,12 @@ impl RequestContext {
|
||||
pub fn before_chat_completion(&mut self, input: &Input) -> Result<()> {
|
||||
// `job__start` validates against exactly what was declared to the
|
||||
// model for THIS request; refresh it every time.
|
||||
//
|
||||
// This is necessary to prevent the model from invoking functions it
|
||||
// otherwise wouldn't have access to by going through the free `tool`
|
||||
// argument of `job__start`. If a function is disabled, the model
|
||||
// shouldn't be able to invoke it at all in any way. This prevents
|
||||
// that backdoor.
|
||||
self.declared_function_names = input.declared_function_names();
|
||||
self.last_message = Some(LastMessage::new(input.clone(), String::new()));
|
||||
Ok(())
|
||||
@@ -4082,11 +4088,6 @@ impl RequestContext {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn jobs_enabled(&self) -> bool {
|
||||
jobs_enabled(self.agent.as_ref(), &self.app.config)
|
||||
}
|
||||
|
||||
pub async fn use_agent(
|
||||
&mut self,
|
||||
app: &AppConfig,
|
||||
@@ -4179,7 +4180,7 @@ impl RequestContext {
|
||||
|
||||
let jobs_enabled = jobs_enabled(Some(&agent), app);
|
||||
let should_init_supervisor = agent.can_spawn_agents() || jobs_enabled;
|
||||
let max_concurrent = if agent.can_spawn_agents() {
|
||||
let max_concurrent_agents = if agent.can_spawn_agents() {
|
||||
agent.max_concurrent_agents()
|
||||
} else {
|
||||
0
|
||||
@@ -4188,7 +4189,8 @@ impl RequestContext {
|
||||
let max_jobs = effective_max_concurrent_jobs(Some(&agent), app);
|
||||
let supervisor = should_init_supervisor.then(|| {
|
||||
Arc::new(RwLock::new(
|
||||
Supervisor::new(max_concurrent, max_depth).with_max_concurrent_jobs(max_jobs),
|
||||
Supervisor::new(max_concurrent_agents, max_depth)
|
||||
.with_max_concurrent_jobs(max_jobs),
|
||||
))
|
||||
});
|
||||
|
||||
@@ -4254,7 +4256,7 @@ impl RequestContext {
|
||||
self.notification_queue = Arc::new(NotificationQueue::new());
|
||||
self.current_depth = 0;
|
||||
self.auto_continue_count = 0;
|
||||
self.pending_agents_guardrail_count = 0;
|
||||
self.pending_tasks_guardrail_count = 0;
|
||||
self.todo_list = TodoList::default();
|
||||
self.rag.take();
|
||||
// Cleared alongside `rag` so the pair never disagrees: an agent RAG is
|
||||
@@ -4750,18 +4752,22 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::config::AppState;
|
||||
use crate::config::agent::AgentConfig;
|
||||
use crate::function::jobs::RingBuf;
|
||||
use crate::function::{ToolCall, skill};
|
||||
use crate::mcp::{McpServer, McpServerFeatures, McpServersConfig, McpTransportType};
|
||||
use crate::supervisor::{
|
||||
AgentExitStatus, AgentHandle, AgentResult, JobHandle, JobResult, JobState, JobStatus,
|
||||
};
|
||||
use crate::utils;
|
||||
use crate::utils::get_env_name;
|
||||
use crate::vault::Vault;
|
||||
use rmcp::model::PromptArgument;
|
||||
use serde_json::json;
|
||||
use serial_test::serial;
|
||||
use std::env;
|
||||
use std::fs::{create_dir_all, remove_dir_all, write};
|
||||
use std::path::PathBuf;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use std::time::{Instant, SystemTime, UNIX_EPOCH};
|
||||
use std::{env, mem};
|
||||
|
||||
struct TestConfigDirGuard {
|
||||
key: String,
|
||||
@@ -5339,14 +5345,14 @@ mod tests {
|
||||
.unwrap()
|
||||
.block_on(async {
|
||||
let join_handle = tokio::spawn(async {
|
||||
Ok(crate::supervisor::AgentResult {
|
||||
Ok(AgentResult {
|
||||
id: "a1".into(),
|
||||
agent_name: "explore".into(),
|
||||
output: String::new(),
|
||||
exit_status: crate::supervisor::AgentExitStatus::Completed,
|
||||
exit_status: AgentExitStatus::Completed,
|
||||
})
|
||||
});
|
||||
let handle = crate::supervisor::AgentHandle {
|
||||
let handle = AgentHandle {
|
||||
id: "a1".to_string(),
|
||||
agent_name: "explore".to_string(),
|
||||
depth: 1,
|
||||
@@ -7696,7 +7702,7 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn make_running_job(abort_signal: utils::AbortSignal) -> crate::supervisor::JobHandle {
|
||||
fn make_running_job(abort_signal: utils::AbortSignal) -> JobHandle {
|
||||
// Leak the runtime so the spawned task is never polled and the job
|
||||
// stays running for the duration of the test.
|
||||
let rt = tokio::runtime::Builder::new_current_thread()
|
||||
@@ -7704,26 +7710,24 @@ mod tests {
|
||||
.build()
|
||||
.unwrap();
|
||||
let join_handle = rt.spawn(async {
|
||||
Ok(crate::supervisor::JobResult {
|
||||
Ok(JobResult {
|
||||
output: serde_json::Value::Null,
|
||||
exit_code: Some(0),
|
||||
output_bytes_captured: 0,
|
||||
})
|
||||
});
|
||||
std::mem::forget(rt);
|
||||
crate::supervisor::JobHandle {
|
||||
mem::forget(rt);
|
||||
JobHandle {
|
||||
id: "j1".to_string(),
|
||||
tool: "execute_command".to_string(),
|
||||
started_at: std::time::Instant::now(),
|
||||
started_at: Instant::now(),
|
||||
join_handle,
|
||||
abort_signal,
|
||||
state: Arc::new(parking_lot::Mutex::new(crate::supervisor::JobState {
|
||||
status: crate::supervisor::JobStatus::Running,
|
||||
state: Arc::new(parking_lot::Mutex::new(JobState {
|
||||
status: JobStatus::Running,
|
||||
pgid: None,
|
||||
})),
|
||||
output_buf: Arc::new(parking_lot::Mutex::new(
|
||||
crate::function::jobs::RingBuf::default(),
|
||||
)),
|
||||
output_buf: Arc::new(parking_lot::Mutex::new(RingBuf::default())),
|
||||
no_change_checks: 0,
|
||||
last_check_state: None,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user