Skip to main content

ai_agents_hitl/
config.rs

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}