1use serde::{Deserialize, Serialize};
2use std::collections::HashMap;
3
4#[derive(Debug, Clone, Serialize, Deserialize)]
5pub struct HITLConfig {
6 #[serde(default = "default_hitl_timeout")]
7 pub default_timeout_seconds: u64,
8
9 #[serde(default)]
10 pub on_timeout: TimeoutAction,
11
12 #[serde(default)]
13 pub message_language: MessageLanguageConfig,
14
15 #[serde(default)]
16 pub tools: HashMap<String, ToolApprovalConfig>,
17
18 #[serde(default)]
19 pub conditions: Vec<ApprovalCondition>,
20
21 #[serde(default)]
22 pub states: HashMap<String, StateApprovalConfig>,
23}
24
25fn default_hitl_timeout() -> u64 {
26 300
27}
28
29impl Default for HITLConfig {
30 fn default() -> Self {
31 Self {
32 default_timeout_seconds: default_hitl_timeout(),
33 on_timeout: TimeoutAction::default(),
34 message_language: MessageLanguageConfig::default(),
35 tools: HashMap::new(),
36 conditions: Vec::new(),
37 states: HashMap::new(),
38 }
39 }
40}
41
42#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
43#[serde(rename_all = "lowercase")]
44pub enum TimeoutAction {
45 #[default]
46 Reject,
47 Approve,
48 Error,
49}
50
51#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
52#[serde(rename_all = "snake_case")]
53pub enum MessageLanguageStrategy {
54 #[default]
55 Auto,
56 User,
57 Approver,
58 Explicit,
59 LlmGenerate,
60}
61
62#[derive(Debug, Clone, Serialize, Deserialize)]
63pub struct MessageLanguageConfig {
64 #[serde(default)]
65 pub strategy: MessageLanguageStrategy,
66
67 #[serde(default = "default_fallback_chain")]
68 pub fallback: Vec<MessageLanguageStrategy>,
69
70 #[serde(default)]
71 pub explicit: Option<String>,
72
73 #[serde(default)]
74 pub llm_generate: Option<LlmGenerateConfig>,
75}
76
77fn default_fallback_chain() -> Vec<MessageLanguageStrategy> {
78 vec![
79 MessageLanguageStrategy::Approver,
80 MessageLanguageStrategy::User,
81 MessageLanguageStrategy::Explicit,
82 MessageLanguageStrategy::LlmGenerate,
83 ]
84}
85
86impl Default for MessageLanguageConfig {
87 fn default() -> Self {
88 Self {
89 strategy: MessageLanguageStrategy::default(),
90 fallback: default_fallback_chain(),
91 explicit: None,
92 llm_generate: None,
93 }
94 }
95}
96
97#[derive(Debug, Clone, Serialize, Deserialize)]
98pub struct LlmGenerateConfig {
99 #[serde(default = "default_router")]
100 pub llm: String,
101
102 #[serde(default = "default_true")]
103 pub include_context: bool,
104}
105
106fn default_router() -> String {
107 "router".to_string()
108}
109
110fn default_true() -> bool {
111 true
112}
113
114impl Default for LlmGenerateConfig {
115 fn default() -> Self {
116 Self {
117 llm: default_router(),
118 include_context: default_true(),
119 }
120 }
121}
122
123#[derive(Debug, Clone, Serialize, Deserialize)]
124#[serde(untagged)]
125pub enum ApprovalMessage {
126 Simple(String),
127 MultiLanguage {
128 #[serde(flatten)]
129 messages: HashMap<String, String>,
130 #[serde(default)]
131 description: Option<String>,
132 },
133}
134
135impl Default for ApprovalMessage {
136 fn default() -> Self {
137 ApprovalMessage::Simple(String::new())
138 }
139}
140
141impl ApprovalMessage {
142 pub fn simple(message: impl Into<String>) -> Self {
143 ApprovalMessage::Simple(message.into())
144 }
145
146 pub fn multi_language(messages: HashMap<String, String>) -> Self {
147 ApprovalMessage::MultiLanguage {
148 messages,
149 description: None,
150 }
151 }
152
153 pub fn with_description(mut self, desc: impl Into<String>) -> Self {
154 if let ApprovalMessage::MultiLanguage {
155 ref mut description,
156 ..
157 } = self
158 {
159 *description = Some(desc.into());
160 }
161 self
162 }
163
164 pub fn get(&self, lang: &str) -> Option<String> {
165 match self {
166 ApprovalMessage::Simple(s) => Some(s.clone()),
167 ApprovalMessage::MultiLanguage { messages, .. } => messages.get(lang).cloned(),
168 }
169 }
170
171 pub fn get_any(&self) -> Option<String> {
172 match self {
173 ApprovalMessage::Simple(s) if !s.is_empty() => Some(s.clone()),
174 ApprovalMessage::Simple(_) => None,
175 ApprovalMessage::MultiLanguage { messages, .. } => messages
176 .get("en")
177 .cloned()
178 .or_else(|| messages.values().next().cloned()),
179 }
180 }
181
182 pub fn description(&self) -> Option<&str> {
183 match self {
184 ApprovalMessage::Simple(_) => None,
185 ApprovalMessage::MultiLanguage { description, .. } => description.as_deref(),
186 }
187 }
188
189 pub fn is_empty(&self) -> bool {
190 match self {
191 ApprovalMessage::Simple(s) => s.is_empty(),
192 ApprovalMessage::MultiLanguage {
193 messages,
194 description,
195 } => messages.is_empty() && description.is_none(),
196 }
197 }
198
199 pub fn available_languages(&self) -> Vec<String> {
200 match self {
201 ApprovalMessage::Simple(_) => vec![],
202 ApprovalMessage::MultiLanguage { messages, .. } => messages.keys().cloned().collect(),
203 }
204 }
205}
206
207#[derive(Debug, Clone, Serialize, Deserialize, Default)]
208pub struct ToolApprovalConfig {
209 #[serde(default)]
210 pub require_approval: bool,
211
212 #[serde(default)]
213 pub approval_context: Vec<String>,
214
215 #[serde(default)]
216 pub approval_message: ApprovalMessage,
217
218 #[serde(default)]
219 pub message_language: Option<MessageLanguageConfig>,
220
221 #[serde(default)]
222 pub timeout_seconds: Option<u64>,
223}
224
225#[derive(Debug, Clone, Serialize, Deserialize)]
226pub struct ApprovalCondition {
227 pub name: String,
228
229 pub when: String,
230
231 #[serde(default)]
232 pub require_approval: bool,
233
234 #[serde(default)]
235 pub approval_message: ApprovalMessage,
236
237 #[serde(default)]
238 pub message_language: Option<MessageLanguageConfig>,
239}
240
241#[derive(Debug, Clone, Serialize, Deserialize, Default)]
242pub struct StateApprovalConfig {
243 #[serde(default)]
244 pub on_enter: StateApprovalTrigger,
245
246 #[serde(default)]
247 pub approval_message: ApprovalMessage,
248
249 #[serde(default)]
250 pub message_language: Option<MessageLanguageConfig>,
251}
252
253#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
254#[serde(rename_all = "snake_case")]
255pub enum StateApprovalTrigger {
256 #[default]
257 None,
258 RequireApproval,
259}
260
261#[cfg(test)]
262mod tests {
263 use super::*;
264
265 #[test]
266 fn test_hitl_config_default() {
267 let config = HITLConfig::default();
268 assert_eq!(config.default_timeout_seconds, 300);
269 assert_eq!(config.on_timeout, TimeoutAction::Reject);
270 assert!(config.tools.is_empty());
271 assert!(config.conditions.is_empty());
272 assert!(config.states.is_empty());
273 assert_eq!(
274 config.message_language.strategy,
275 MessageLanguageStrategy::Auto
276 );
277 }
278
279 #[test]
280 fn test_hitl_config_from_yaml() {
281 let yaml = r#"
282default_timeout_seconds: 600
283on_timeout: approve
284message_language:
285 strategy: approver
286 fallback:
287 - user
288 - explicit
289 explicit: en
290tools:
291 send_payment:
292 require_approval: true
293 approval_context:
294 - amount
295 - recipient
296 approval_message: "Approve payment?"
297 timeout_seconds: 120
298 delete_record:
299 require_approval: true
300conditions:
301 - name: high_value
302 when: "amount > 1000"
303 require_approval: true
304 approval_message: "High value transaction"
305states:
306 escalation:
307 on_enter: require_approval
308 approval_message: "Escalate to human?"
309"#;
310 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
311 assert_eq!(config.default_timeout_seconds, 600);
312 assert_eq!(config.on_timeout, TimeoutAction::Approve);
313 assert_eq!(
314 config.message_language.strategy,
315 MessageLanguageStrategy::Approver
316 );
317 assert_eq!(config.message_language.explicit, Some("en".to_string()));
318 assert_eq!(config.tools.len(), 2);
319
320 let payment = config.tools.get("send_payment").unwrap();
321 assert!(payment.require_approval);
322 assert_eq!(payment.approval_context, vec!["amount", "recipient"]);
323 assert_eq!(payment.timeout_seconds, Some(120));
324
325 assert_eq!(config.conditions.len(), 1);
326 assert_eq!(config.conditions[0].name, "high_value");
327
328 let escalation = config.states.get("escalation").unwrap();
329 assert_eq!(escalation.on_enter, StateApprovalTrigger::RequireApproval);
330 }
331
332 #[test]
333 fn test_multi_language_approval_message() {
334 let yaml = r#"
335tools:
336 process_payment:
337 require_approval: true
338 approval_message:
339 en: "Approve payment of {{ amount }}?"
340 ko: "{{ amount }} 결제를 승인하시겠습니까?"
341 ja: "{{ amount }}の支払いを承認しますか?"
342 description: "Payment approval request"
343"#;
344 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
345 let payment = config.tools.get("process_payment").unwrap();
346
347 let msg = &payment.approval_message;
348 assert_eq!(
349 msg.get("en"),
350 Some("Approve payment of {{ amount }}?".to_string())
351 );
352 assert_eq!(
353 msg.get("ko"),
354 Some("{{ amount }} 결제를 승인하시겠습니까?".to_string())
355 );
356 assert_eq!(
357 msg.get("ja"),
358 Some("{{ amount }}の支払いを承認しますか?".to_string())
359 );
360 assert_eq!(msg.description(), Some("Payment approval request"));
361 }
362
363 #[test]
364 fn test_tool_level_message_language_override() {
365 let yaml = r#"
366message_language:
367 strategy: auto
368tools:
369 admin_action:
370 require_approval: true
371 message_language:
372 strategy: explicit
373 explicit: en
374 approval_message:
375 en: "Admin action required"
376"#;
377 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
378 assert_eq!(
379 config.message_language.strategy,
380 MessageLanguageStrategy::Auto
381 );
382
383 let admin = config.tools.get("admin_action").unwrap();
384 let tool_lang = admin.message_language.as_ref().unwrap();
385 assert_eq!(tool_lang.strategy, MessageLanguageStrategy::Explicit);
386 assert_eq!(tool_lang.explicit, Some("en".to_string()));
387 }
388
389 #[test]
390 fn test_llm_generate_config() {
391 let yaml = r#"
392message_language:
393 strategy: llm_generate
394 llm_generate:
395 llm: router
396 include_context: true
397tools:
398 dynamic_action:
399 require_approval: true
400 approval_message:
401 description: "Dynamic action requiring approval"
402"#;
403 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
404 assert_eq!(
405 config.message_language.strategy,
406 MessageLanguageStrategy::LlmGenerate
407 );
408
409 let llm_config = config.message_language.llm_generate.as_ref().unwrap();
410 assert_eq!(llm_config.llm, "router");
411 assert!(llm_config.include_context);
412 }
413
414 #[test]
415 fn test_timeout_action_variants() {
416 let yaml_reject = "reject";
417 let yaml_approve = "approve";
418 let yaml_error = "error";
419
420 let action: TimeoutAction = serde_yaml::from_str(yaml_reject).unwrap();
421 assert_eq!(action, TimeoutAction::Reject);
422
423 let action: TimeoutAction = serde_yaml::from_str(yaml_approve).unwrap();
424 assert_eq!(action, TimeoutAction::Approve);
425
426 let action: TimeoutAction = serde_yaml::from_str(yaml_error).unwrap();
427 assert_eq!(action, TimeoutAction::Error);
428 }
429
430 #[test]
431 fn test_message_language_strategy_variants() {
432 assert_eq!(
433 serde_yaml::from_str::<MessageLanguageStrategy>("auto").unwrap(),
434 MessageLanguageStrategy::Auto
435 );
436 assert_eq!(
437 serde_yaml::from_str::<MessageLanguageStrategy>("user").unwrap(),
438 MessageLanguageStrategy::User
439 );
440 assert_eq!(
441 serde_yaml::from_str::<MessageLanguageStrategy>("approver").unwrap(),
442 MessageLanguageStrategy::Approver
443 );
444 assert_eq!(
445 serde_yaml::from_str::<MessageLanguageStrategy>("explicit").unwrap(),
446 MessageLanguageStrategy::Explicit
447 );
448 assert_eq!(
449 serde_yaml::from_str::<MessageLanguageStrategy>("llm_generate").unwrap(),
450 MessageLanguageStrategy::LlmGenerate
451 );
452 }
453
454 #[test]
455 fn test_approval_message_simple() {
456 let msg = ApprovalMessage::simple("Approve?");
457 assert_eq!(msg.get("en"), Some("Approve?".to_string()));
458 assert_eq!(msg.get("ko"), Some("Approve?".to_string()));
459 assert!(msg.description().is_none());
460 assert!(!msg.is_empty());
461 }
462
463 #[test]
464 fn test_approval_message_multi_language() {
465 let mut messages = HashMap::new();
466 messages.insert("en".to_string(), "Approve?".to_string());
467 messages.insert("ko".to_string(), "승인?".to_string());
468
469 let msg = ApprovalMessage::multi_language(messages).with_description("Test approval");
470 assert_eq!(msg.get("en"), Some("Approve?".to_string()));
471 assert_eq!(msg.get("ko"), Some("승인?".to_string()));
472 assert_eq!(msg.get("ja"), None);
473 assert_eq!(msg.description(), Some("Test approval"));
474 assert!(!msg.is_empty());
475
476 let langs = msg.available_languages();
477 assert!(langs.contains(&"en".to_string()));
478 assert!(langs.contains(&"ko".to_string()));
479 }
480
481 #[test]
482 fn test_approval_message_get_any() {
483 let msg = ApprovalMessage::simple("Test");
484 assert_eq!(msg.get_any(), Some("Test".to_string()));
485
486 let empty = ApprovalMessage::simple("");
487 assert_eq!(empty.get_any(), None);
488
489 let mut messages = HashMap::new();
490 messages.insert("ko".to_string(), "한국어".to_string());
491 let msg = ApprovalMessage::multi_language(messages);
492 assert_eq!(msg.get_any(), Some("한국어".to_string()));
493
494 let mut messages_with_en = HashMap::new();
495 messages_with_en.insert("en".to_string(), "English".to_string());
496 messages_with_en.insert("ko".to_string(), "한국어".to_string());
497 let msg = ApprovalMessage::multi_language(messages_with_en);
498 assert_eq!(msg.get_any(), Some("English".to_string()));
499 }
500
501 #[test]
502 fn test_tool_approval_config_default() {
503 let config = ToolApprovalConfig::default();
504 assert!(!config.require_approval);
505 assert!(config.approval_context.is_empty());
506 assert!(config.approval_message.is_empty());
507 assert!(config.message_language.is_none());
508 assert!(config.timeout_seconds.is_none());
509 }
510
511 #[test]
512 fn test_backward_compatible_simple_message() {
513 let yaml = r#"
514tools:
515 old_tool:
516 require_approval: true
517 approval_message: "Simple approval message"
518"#;
519 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
520 let tool = config.tools.get("old_tool").unwrap();
521 assert_eq!(
522 tool.approval_message.get_any(),
523 Some("Simple approval message".to_string())
524 );
525 }
526
527 #[test]
528 fn test_condition_with_multi_language() {
529 let yaml = r#"
530conditions:
531 - name: high_value
532 when: "amount > 1000"
533 require_approval: true
534 approval_message:
535 en: "High value: {{ amount }}"
536 ko: "고액 거래: {{ amount }}"
537 message_language:
538 strategy: user
539"#;
540 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
541 let condition = &config.conditions[0];
542 assert_eq!(condition.name, "high_value");
543 assert_eq!(
544 condition.approval_message.get("en"),
545 Some("High value: {{ amount }}".to_string())
546 );
547 assert_eq!(
548 condition.message_language.as_ref().unwrap().strategy,
549 MessageLanguageStrategy::User
550 );
551 }
552
553 #[test]
554 fn test_state_with_multi_language() {
555 let yaml = r#"
556states:
557 escalation:
558 on_enter: require_approval
559 approval_message:
560 en: "Escalate to human agent?"
561 ko: "상담원에게 연결하시겠습니까?"
562 ja: "人間のエージェントにエスカレーションしますか?"
563 message_language:
564 strategy: approver
565"#;
566 let config: HITLConfig = serde_yaml::from_str(yaml).unwrap();
567 let state = config.states.get("escalation").unwrap();
568 assert_eq!(state.on_enter, StateApprovalTrigger::RequireApproval);
569 assert_eq!(
570 state.approval_message.get("ko"),
571 Some("상담원에게 연결하시겠습니까?".to_string())
572 );
573 }
574
575 #[test]
576 fn test_default_fallback_chain() {
577 let chain = default_fallback_chain();
578 assert_eq!(chain.len(), 4);
579 assert_eq!(chain[0], MessageLanguageStrategy::Approver);
580 assert_eq!(chain[1], MessageLanguageStrategy::User);
581 assert_eq!(chain[2], MessageLanguageStrategy::Explicit);
582 assert_eq!(chain[3], MessageLanguageStrategy::LlmGenerate);
583 }
584
585 #[test]
586 fn test_llm_generate_config_defaults() {
587 let config = LlmGenerateConfig::default();
588 assert_eq!(config.llm, "router");
589 assert!(config.include_context);
590 }
591}