1use std::collections::BTreeMap;
2
3use serde::{de::Error as _, Deserialize, Deserializer, Serialize};
4use serde_json::Value;
5
6use crate::model_settings::ModelSettings;
7use crate::prompt::PromptBundle;
8use crate::tools::common::trim_portable_whitespace;
9use crate::tools::{ToolPolicy, ToolSideEffect};
10
11use super::{
12 json_value_from_serializable, AgentStatus, CompletionReason, Message, Metadata, NoToolPolicy,
13};
14
15pub const INVALID_SUB_AGENT_MODEL_CODE: &str = "invalid_sub_agent_model";
16pub const INVALID_SUB_AGENT_MODEL_MESSAGE: &str = "sub-agent model cannot be empty";
17pub const INVALID_SUB_AGENT_SYSTEM_PROMPT_CODE: &str = "invalid_sub_agent_system_prompt";
18pub const INVALID_SUB_AGENT_SYSTEM_PROMPT_MESSAGE: &str =
19 "sub-agent system_prompt cannot be empty when provided";
20
21#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct SubAgentConfigValidationError {
23 code: &'static str,
24 message: &'static str,
25}
26
27impl SubAgentConfigValidationError {
28 fn new(code: &'static str, message: &'static str) -> Self {
29 Self { code, message }
30 }
31
32 pub fn code(&self) -> &'static str {
33 self.code
34 }
35
36 pub fn message(&self) -> &'static str {
37 self.message
38 }
39}
40
41impl std::fmt::Display for SubAgentConfigValidationError {
42 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
43 formatter.write_str(self.message)
44 }
45}
46
47impl std::error::Error for SubAgentConfigValidationError {}
48
49#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
50pub struct SubAgentConfig {
51 pub model: String,
52 pub description: String,
53 pub backend: Option<String>,
54 pub system_prompt: Option<String>,
55 pub max_cycles: u32,
56 pub session_memory_enabled: bool,
57 pub exclude_tools: Vec<String>,
58 pub metadata: Metadata,
59 pub denied_side_effects: Vec<ToolSideEffect>,
60 pub denied_capability_tags: Vec<String>,
61 pub deny_terminal_tools: bool,
62 pub denied_cost_dimensions: Vec<String>,
63}
64
65#[derive(Deserialize)]
66#[serde(deny_unknown_fields)]
67struct SubAgentConfigWire {
68 model: String,
69 #[serde(default)]
70 description: String,
71 #[serde(default)]
72 backend: Option<String>,
73 #[serde(default)]
74 system_prompt: Option<String>,
75 #[serde(default = "default_sub_agent_max_cycles")]
76 max_cycles: u32,
77 #[serde(default)]
78 session_memory_enabled: bool,
79 #[serde(default)]
80 exclude_tools: Vec<String>,
81 #[serde(default)]
82 metadata: Metadata,
83 #[serde(default)]
84 denied_side_effects: Vec<ToolSideEffect>,
85 #[serde(default)]
86 denied_capability_tags: Vec<String>,
87 #[serde(default)]
88 deny_terminal_tools: bool,
89 #[serde(default)]
90 denied_cost_dimensions: Vec<String>,
91}
92
93const fn default_sub_agent_max_cycles() -> u32 {
94 8
95}
96
97impl<'de> Deserialize<'de> for SubAgentConfig {
98 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
99 where
100 D: Deserializer<'de>,
101 {
102 let value = Value::deserialize(deserializer)?;
103 if !value.is_object() {
104 return Err(D::Error::custom("SubAgentConfig payload must be an object"));
105 }
106 let wire = serde_json::from_value::<SubAgentConfigWire>(value).map_err(D::Error::custom)?;
107 let mut config = Self {
108 model: trim_portable_whitespace(&wire.model).to_string(),
109 description: wire.description,
110 backend: wire.backend,
111 system_prompt: wire.system_prompt,
112 max_cycles: wire.max_cycles,
113 session_memory_enabled: wire.session_memory_enabled,
114 exclude_tools: wire.exclude_tools,
115 metadata: wire.metadata,
116 denied_side_effects: wire.denied_side_effects,
117 denied_capability_tags: wire.denied_capability_tags,
118 deny_terminal_tools: wire.deny_terminal_tools,
119 denied_cost_dimensions: wire.denied_cost_dimensions,
120 };
121 config
122 .normalize_policy_denials()
123 .map_err(D::Error::custom)?;
124 config.validate().map_err(D::Error::custom)?;
125 Ok(config)
126 }
127}
128
129impl SubAgentConfig {
130 pub fn new(model: impl Into<String>, description: impl Into<String>) -> Self {
131 let model = model.into();
132 Self {
133 model: trim_portable_whitespace(&model).to_string(),
134 description: description.into(),
135 backend: None,
136 system_prompt: None,
137 max_cycles: 8,
138 session_memory_enabled: false,
139 exclude_tools: Vec::new(),
140 metadata: Metadata::new(),
141 denied_side_effects: Vec::new(),
142 denied_capability_tags: Vec::new(),
143 deny_terminal_tools: false,
144 denied_cost_dimensions: Vec::new(),
145 }
146 }
147
148 pub fn validate(&self) -> Result<(), SubAgentConfigValidationError> {
149 if trim_portable_whitespace(&self.model).is_empty() {
150 return Err(SubAgentConfigValidationError::new(
151 INVALID_SUB_AGENT_MODEL_CODE,
152 INVALID_SUB_AGENT_MODEL_MESSAGE,
153 ));
154 }
155 if self
156 .system_prompt
157 .as_deref()
158 .is_some_and(|prompt| trim_portable_whitespace(prompt).is_empty())
159 {
160 return Err(SubAgentConfigValidationError::new(
161 INVALID_SUB_AGENT_SYSTEM_PROMPT_CODE,
162 INVALID_SUB_AGENT_SYSTEM_PROMPT_MESSAGE,
163 ));
164 }
165 self.declared_tool_policy().normalized().map_err(|_| {
166 SubAgentConfigValidationError::new(
167 "invalid_sub_agent_tool_policy",
168 "sub-agent tool policy is invalid",
169 )
170 })?;
171 Ok(())
172 }
173
174 pub fn declared_tool_policy(&self) -> ToolPolicy {
175 ToolPolicy {
176 denied_side_effects: self.denied_side_effects.clone(),
177 denied_capability_tags: self.denied_capability_tags.clone(),
178 deny_terminal_tools: self.deny_terminal_tools,
179 denied_cost_dimensions: self.denied_cost_dimensions.clone(),
180 ..ToolPolicy::default()
181 }
182 }
183
184 fn normalize_policy_denials(&mut self) -> Result<(), crate::tools::ToolMetadataError> {
185 let policy = self.declared_tool_policy().normalized()?;
186 self.denied_side_effects = policy.denied_side_effects;
187 self.denied_capability_tags = policy.denied_capability_tags;
188 self.deny_terminal_tools = policy.deny_terminal_tools;
189 self.denied_cost_dimensions = policy.denied_cost_dimensions;
190 Ok(())
191 }
192}
193
194#[derive(Debug, Clone, PartialEq, Serialize)]
195pub struct AgentTask {
196 pub task_id: String,
197 pub model: String,
198 pub prompt_bundle: PromptBundle,
199 pub user_prompt: String,
200 pub max_cycles: u32,
201 pub memory_compact_threshold: u64,
202 pub memory_threshold_percentage: u8,
203 pub no_tool_policy: NoToolPolicy,
204 pub allow_interruption: bool,
205 pub use_workspace: bool,
206 pub sub_agents: BTreeMap<String, SubAgentConfig>,
207 pub agent_type: Option<String>,
208 pub native_multimodal: bool,
209 pub extra_tool_names: Vec<String>,
210 pub exclude_tools: Vec<String>,
211 pub initial_messages: Vec<Message>,
212 pub initial_shared_state: Metadata,
213 pub model_settings: Option<ModelSettings>,
214 pub metadata: Metadata,
215}
216
217#[derive(Deserialize)]
218#[serde(deny_unknown_fields)]
219struct AgentTaskWire {
220 task_id: String,
221 model: String,
222 prompt_bundle: PromptBundle,
223 user_prompt: String,
224 #[serde(default = "default_agent_task_max_cycles")]
225 max_cycles: u32,
226 #[serde(default = "default_memory_compact_threshold")]
227 memory_compact_threshold: u64,
228 #[serde(default = "default_memory_threshold_percentage")]
229 memory_threshold_percentage: u8,
230 #[serde(default)]
231 no_tool_policy: NoToolPolicy,
232 #[serde(default = "default_true")]
233 allow_interruption: bool,
234 #[serde(default = "default_true")]
235 use_workspace: bool,
236 #[serde(default)]
237 sub_agents: BTreeMap<String, SubAgentConfig>,
238 #[serde(default)]
239 agent_type: Option<String>,
240 #[serde(default)]
241 native_multimodal: bool,
242 #[serde(default)]
243 extra_tool_names: Vec<String>,
244 #[serde(default)]
245 exclude_tools: Vec<String>,
246 #[serde(default, deserialize_with = "deserialize_agent_task_messages")]
247 initial_messages: Vec<Message>,
248 #[serde(default)]
249 initial_shared_state: Metadata,
250 #[serde(default, deserialize_with = "deserialize_agent_task_model_settings")]
251 model_settings: Option<ModelSettings>,
252 #[serde(default)]
253 metadata: Metadata,
254}
255
256impl<'de> Deserialize<'de> for AgentTask {
257 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
258 where
259 D: Deserializer<'de>,
260 {
261 let value = Value::deserialize(deserializer)?;
262 if !value.is_object() {
263 return Err(D::Error::custom("AgentTask payload must be an object"));
264 }
265 let wire = serde_json::from_value::<AgentTaskWire>(value).map_err(D::Error::custom)?;
266 Ok(Self {
267 task_id: wire.task_id,
268 model: wire.model,
269 prompt_bundle: wire.prompt_bundle,
270 user_prompt: wire.user_prompt,
271 max_cycles: wire.max_cycles,
272 memory_compact_threshold: wire.memory_compact_threshold,
273 memory_threshold_percentage: wire.memory_threshold_percentage,
274 no_tool_policy: wire.no_tool_policy,
275 allow_interruption: wire.allow_interruption,
276 use_workspace: wire.use_workspace,
277 sub_agents: wire.sub_agents,
278 agent_type: wire.agent_type,
279 native_multimodal: wire.native_multimodal,
280 extra_tool_names: wire.extra_tool_names,
281 exclude_tools: wire.exclude_tools,
282 initial_messages: wire.initial_messages,
283 initial_shared_state: wire.initial_shared_state,
284 model_settings: wire.model_settings,
285 metadata: wire.metadata,
286 })
287 }
288}
289
290const fn default_agent_task_max_cycles() -> u32 {
291 8
292}
293
294const fn default_memory_compact_threshold() -> u64 {
295 250_000
296}
297
298const fn default_memory_threshold_percentage() -> u8 {
299 90
300}
301
302const fn default_true() -> bool {
303 true
304}
305
306fn deserialize_agent_task_messages<'de, D>(deserializer: D) -> Result<Vec<Message>, D::Error>
307where
308 D: Deserializer<'de>,
309{
310 Vec::<Value>::deserialize(deserializer)?
311 .into_iter()
312 .enumerate()
313 .map(|(index, value)| {
314 validate_agent_task_message(&value, index).map_err(D::Error::custom)?;
315 Message::from_dict(&value).map_err(D::Error::custom)
316 })
317 .collect()
318}
319
320fn deserialize_agent_task_model_settings<'de, D>(
321 deserializer: D,
322) -> Result<Option<ModelSettings>, D::Error>
323where
324 D: Deserializer<'de>,
325{
326 let value = Option::<Value>::deserialize(deserializer)?;
327 value
328 .map(|value| {
329 if !value.is_object() {
330 return Err(D::Error::custom(
331 "AgentTask field 'model_settings' must be an object or null",
332 ));
333 }
334 serde_json::from_value(value).map_err(D::Error::custom)
335 })
336 .transpose()
337}
338
339fn validate_agent_task_message(value: &Value, index: usize) -> Result<(), String> {
340 let object = value
341 .as_object()
342 .ok_or_else(|| format!("AgentTask initial_messages[{index}] must be an object"))?;
343 let role = object
344 .get("role")
345 .and_then(Value::as_str)
346 .ok_or_else(|| format!("AgentTask initial_messages[{index}].role must be a string"))?;
347 if !matches!(role, "system" | "user" | "assistant" | "tool") {
348 return Err(format!(
349 "unknown AgentTask initial_messages[{index}].role: {role}"
350 ));
351 }
352 if object
353 .get("content")
354 .is_some_and(|value| !value.is_string())
355 {
356 return Err(format!(
357 "AgentTask initial_messages[{index}].content must be a string"
358 ));
359 }
360 for field_name in ["name", "tool_call_id", "reasoning_content", "image_url"] {
361 if object
362 .get(field_name)
363 .is_some_and(|value| !value.is_null() && !value.is_string())
364 {
365 return Err(format!(
366 "AgentTask initial_messages[{index}].{field_name} must be a string or null"
367 ));
368 }
369 }
370 if object.get("tool_calls").is_some_and(|value| {
371 !value
372 .as_array()
373 .is_some_and(|items| items.iter().all(Value::is_object))
374 }) {
375 return Err(format!(
376 "AgentTask initial_messages[{index}].tool_calls must be an array of objects"
377 ));
378 }
379 if object
380 .get("metadata")
381 .is_some_and(|value| !value.is_object())
382 {
383 return Err(format!(
384 "AgentTask initial_messages[{index}].metadata must be an object"
385 ));
386 }
387 Ok(())
388}
389
390impl AgentTask {
391 pub fn new(
392 task_id: impl Into<String>,
393 model: impl Into<String>,
394 prompt_bundle: PromptBundle,
395 user_prompt: impl Into<String>,
396 ) -> Self {
397 Self {
398 task_id: task_id.into(),
399 model: model.into(),
400 prompt_bundle,
401 user_prompt: user_prompt.into(),
402 max_cycles: 8,
403 memory_compact_threshold: default_memory_compact_threshold(),
404 memory_threshold_percentage: 90,
405 no_tool_policy: NoToolPolicy::Continue,
406 allow_interruption: true,
407 use_workspace: true,
408 sub_agents: BTreeMap::new(),
409 agent_type: None,
410 native_multimodal: false,
411 extra_tool_names: Vec::new(),
412 exclude_tools: Vec::new(),
413 initial_messages: Vec::new(),
414 initial_shared_state: Metadata::new(),
415 model_settings: None,
416 metadata: Metadata::new(),
417 }
418 }
419
420 pub fn sub_agents_enabled(&self) -> bool {
421 !self.sub_agents.is_empty()
422 }
423}
424
425#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
426pub struct SubTaskRequest {
427 pub agent_name: String,
428 pub task_description: String,
429 pub output_requirements: String,
430 pub include_main_summary: bool,
431 pub exclude_files_pattern: Option<String>,
432 pub metadata: Metadata,
433}
434
435impl SubTaskRequest {
436 pub fn new(agent_name: impl Into<String>, task_description: impl Into<String>) -> Self {
437 Self {
438 agent_name: agent_name.into(),
439 task_description: task_description.into(),
440 output_requirements: String::new(),
441 include_main_summary: false,
442 exclude_files_pattern: None,
443 metadata: Metadata::new(),
444 }
445 }
446}
447
448#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
449pub struct SubTaskOutcome {
450 pub task_id: String,
451 pub agent_name: String,
452 pub status: AgentStatus,
453 pub session_id: Option<String>,
454 pub final_answer: Option<String>,
455 pub wait_reason: Option<String>,
456 pub error: Option<String>,
457 #[serde(default, skip_serializing_if = "Option::is_none")]
458 pub error_code: Option<String>,
459 #[serde(default, skip_serializing_if = "Option::is_none")]
460 pub completion_reason: Option<CompletionReason>,
461 #[serde(default, skip_serializing_if = "Option::is_none")]
462 pub completion_tool_name: Option<String>,
463 #[serde(default, skip_serializing_if = "Option::is_none")]
464 pub partial_output: Option<String>,
465 pub cycles: u32,
466 pub todo_list: Vec<Value>,
467 pub resolved: BTreeMap<String, String>,
468}
469
470impl Default for SubTaskOutcome {
471 fn default() -> Self {
472 Self {
473 task_id: String::new(),
474 agent_name: String::new(),
475 status: AgentStatus::Pending,
476 session_id: None,
477 final_answer: None,
478 wait_reason: None,
479 error: None,
480 error_code: None,
481 completion_reason: None,
482 completion_tool_name: None,
483 partial_output: None,
484 cycles: 0,
485 todo_list: Vec::new(),
486 resolved: BTreeMap::new(),
487 }
488 }
489}
490
491impl SubTaskOutcome {
492 pub fn to_dict(&self) -> Value {
493 self.to_value()
494 }
495
496 pub fn to_value(&self) -> Value {
497 json_value_from_serializable(self)
498 }
499}