1use std::collections::HashSet;
2
3use serde::{Deserialize, Serialize};
4
5use crate::types::{ChatMessage, ImageAttachment, Message, MessageRole, ToolCallMessage};
6
7use crate::types::SessionId;
8
9#[derive(Clone, Debug, Default, Serialize, Deserialize)]
20pub struct RunState {
21 pub turn_tool_calls: usize,
25 pub run_has_tool_calls: bool,
29 pub reasoning_only_strikes: usize,
34 pub empty_response_strikes: usize,
39 pub nudge_count: usize,
43 #[serde(default)]
49 pub thinking_disabled_for_rest_of_run: bool,
50
51 #[serde(default)]
57 pub original_thinking_enabled: bool,
58
59 #[serde(default)]
65 pub truncation_strikes: usize,
66}
67
68impl RunState {
69 pub fn reset_for_new_run(&mut self) {
71 self.turn_tool_calls = 0;
72 self.run_has_tool_calls = false;
73 self.reasoning_only_strikes = 0;
74 self.empty_response_strikes = 0;
75 self.nudge_count = 0;
76 self.thinking_disabled_for_rest_of_run = false;
77 self.truncation_strikes = 0;
78 }
81
82 pub fn record_tool_calls(&mut self, n: usize) {
85 self.turn_tool_calls += n;
86 self.run_has_tool_calls = true;
87 self.reasoning_only_strikes = 0;
88 self.empty_response_strikes = 0;
89 self.truncation_strikes = 0;
90 }
91
92 pub fn record_reasoning_only(&mut self) -> usize {
97 self.empty_response_strikes = 0;
98 self.reasoning_only_strikes += 1;
99 if self.reasoning_only_strikes >= 3 {
100 tracing::info!(
101 strikes = self.reasoning_only_strikes,
102 "reasoning_only_strikes reached 3, disabling thinking for rest of run"
103 );
104 self.thinking_disabled_for_rest_of_run = true;
105 }
106 self.reasoning_only_strikes
107 }
108
109 pub fn record_empty_response(&mut self) -> usize {
113 self.reasoning_only_strikes = 0;
114 self.empty_response_strikes += 1;
115 self.empty_response_strikes
116 }
117
118 pub fn record_truncation(&mut self) -> usize {
121 self.truncation_strikes += 1;
122 self.truncation_strikes
123 }
124}
125
126#[derive(Deserialize)]
133struct RawAgentSession {
134 id: Option<SessionId>,
135 chat_messages: Vec<ChatMessage>,
136 always_allowed_actions: HashSet<String>,
137 total_tool_calls: usize,
138
139 run_state: Option<RunState>,
141
142 nudge_count: Option<usize>,
144 turn_tool_calls: Option<usize>,
145 reasoning_only_strikes: Option<usize>,
146 empty_response_strikes: Option<usize>,
147}
148
149impl From<RawAgentSession> for AgentSession {
150 fn from(raw: RawAgentSession) -> Self {
151 let run_state = raw.run_state.unwrap_or_else(|| RunState {
152 nudge_count: raw.nudge_count.unwrap_or(0),
153 turn_tool_calls: raw.turn_tool_calls.unwrap_or(0),
154 reasoning_only_strikes: raw.reasoning_only_strikes.unwrap_or(0),
155 empty_response_strikes: raw.empty_response_strikes.unwrap_or(0),
156 ..RunState::default()
157 });
158 Self {
159 id: raw.id,
160 chat_messages: raw.chat_messages,
161 always_allowed_actions: raw.always_allowed_actions,
162 total_tool_calls: raw.total_tool_calls,
163 run_state,
164 }
165 }
166}
167
168#[derive(Clone, Debug, Default, Serialize, Deserialize)]
173#[serde(from = "RawAgentSession")]
174pub struct AgentSession {
175 id: Option<SessionId>,
176 chat_messages: Vec<ChatMessage>,
179 always_allowed_actions: HashSet<String>,
180 pub total_tool_calls: usize,
183 pub run_state: RunState,
185}
186
187impl AgentSession {
188 pub fn new(id: SessionId) -> Self {
189 Self {
190 id: Some(id),
191 chat_messages: Vec::new(),
192 always_allowed_actions: HashSet::new(),
193 total_tool_calls: 0,
194 run_state: RunState::default(),
195 }
196 }
197
198 pub fn id(&self) -> Option<SessionId> {
199 self.id.clone()
200 }
201
202 pub fn simple_messages(&self) -> Vec<Message> {
206 self.chat_messages
207 .iter()
208 .filter_map(|cm| match cm {
209 ChatMessage::Assistant { content: None, .. } => None,
210 ChatMessage::Assistant {
211 content: Some(c),
212 tool_calls: Some(tc),
213 ..
214 } if c.is_empty() && !tc.is_empty() => None,
215 _ => Some(Message::from(cm)),
216 })
217 .collect()
218 }
219
220 pub fn chat_messages(&self) -> &[ChatMessage] {
221 &self.chat_messages
222 }
223
224 pub fn chat_messages_mut(&mut self) -> &mut Vec<ChatMessage> {
226 &mut self.chat_messages
227 }
228
229 pub fn is_action_allowed(&self, action_key: &str) -> bool {
230 self.always_allowed_actions.contains(action_key)
231 }
232
233 pub fn allow_action(&mut self, action_key: impl Into<String>) {
234 self.always_allowed_actions.insert(action_key.into());
235 }
236
237 pub fn push_message(&mut self, role: MessageRole, content: impl Into<String>) {
238 let content = content.into();
239 let chat_msg = match role {
240 MessageRole::System => ChatMessage::system(content),
241 MessageRole::User => ChatMessage::user(content),
242 MessageRole::Assistant => ChatMessage::assistant(content),
243 MessageRole::Tool => ChatMessage::tool(String::new(), content),
244 };
245 self.chat_messages.push(chat_msg);
246 }
247
248 pub fn push_message_ephemeral(&mut self, role: MessageRole, content: impl Into<String>) {
252 let content = content.into();
253 let chat_msg = match role {
254 MessageRole::System => ChatMessage::system_ephemeral(content),
255 MessageRole::User => ChatMessage::user_ephemeral(content),
256 MessageRole::Assistant | MessageRole::Tool => {
257 self.push_message(role, content);
260 return;
261 }
262 };
263 self.chat_messages.push(chat_msg);
264 }
265
266 pub fn set_system_prompt(&mut self, content: impl Into<String>) {
271 let content = content.into();
272 if let Some(idx) = self.chat_messages.iter().position(|m| {
273 matches!(
274 m,
275 ChatMessage::System {
276 ephemeral: false,
277 ..
278 }
279 )
280 }) {
281 self.chat_messages[idx] = ChatMessage::system(content);
282 } else {
283 self.chat_messages.insert(0, ChatMessage::system(content));
284 }
285 }
286
287 pub fn push_assistant_with_reasoning(
291 &mut self,
292 content: impl Into<String>,
293 reasoning: impl Into<String>,
294 ) {
295 self.chat_messages
296 .push(ChatMessage::assistant_with_reasoning(content, reasoning));
297 }
298
299 pub fn push_user_message_with_images(
300 &mut self,
301 content: impl Into<String>,
302 images: Vec<ImageAttachment>,
303 ) {
304 self.chat_messages
305 .push(ChatMessage::user_with_images(content, images));
306 }
307
308 pub fn push_assistant_tool_call(
309 &mut self,
310 tool_call_id: &str,
311 tool_name: &str,
312 arguments_json: &str,
313 ) {
314 self.chat_messages.push(ChatMessage::assistant_tool_call(
315 tool_call_id,
316 tool_name,
317 arguments_json,
318 ));
319 }
320
321 pub fn push_assistant_tool_calls(
322 &mut self,
323 tool_calls: &[(String, String, String)],
324 reasoning: Option<String>,
325 content: Option<String>,
326 ) {
327 let calls: Vec<ToolCallMessage> = tool_calls
328 .iter()
329 .map(|(id, name, args)| {
330 let valid_args = if serde_json::from_str::<serde_json::Value>(args).is_ok() {
340 args.clone()
341 } else {
342 tracing::warn!(
343 tool_name = %name,
344 args_len = args.len(),
345 "tool call arguments are not valid JSON (provider truncated them), \
346 sanitizing to empty object; re-issue instruction goes to the tool_result"
347 );
348 "{}".to_string()
349 };
350 ToolCallMessage {
351 id: id.clone(),
352 name: name.clone(),
353 arguments: valid_args,
354 }
355 })
356 .collect();
357 self.chat_messages.push(ChatMessage::Assistant {
358 content,
359 reasoning_content: reasoning,
360 tool_calls: Some(calls),
361 thinking_signature: None,
362 });
363 }
364
365 pub fn push_tool_result(&mut self, tool_call_id: &str, content: impl Into<String>) {
366 self.chat_messages
367 .push(ChatMessage::tool(tool_call_id, content));
368 }
369
370 pub fn remove_ephemeral_messages(&mut self) {
374 let before = self.chat_messages.len();
375 self.chat_messages.retain(|m| !m.is_ephemeral());
376 let removed = before - self.chat_messages.len();
377 if removed > 0 {
378 tracing::debug!(
379 removed,
380 remaining = self.chat_messages.len(),
381 "ephemeral messages cleaned up"
382 );
383 }
384 }
385
386 pub fn turn_count(&self) -> usize {
389 self.chat_messages
390 .iter()
391 .filter(|m| matches!(m, ChatMessage::User { .. }))
392 .count()
393 }
394
395 pub fn trim_oldest_turns(&mut self, max_turns: usize) {
398 let current_turns = self.turn_count();
399 if current_turns <= max_turns {
400 return;
401 }
402 let turns_to_remove = current_turns - max_turns;
403
404 let user_positions: Vec<usize> = self
406 .chat_messages
407 .iter()
408 .enumerate()
409 .filter_map(|(i, m)| {
410 if matches!(m, ChatMessage::User { .. }) {
411 Some(i)
412 } else {
413 None
414 }
415 })
416 .collect();
417
418 if user_positions.len() <= turns_to_remove {
419 return;
420 }
421
422 let system_prefix = self
424 .chat_messages
425 .iter()
426 .take_while(|m| matches!(m, ChatMessage::System { .. }))
427 .count();
428
429 let drain_end = user_positions[turns_to_remove];
431 if system_prefix >= drain_end {
432 return; }
434
435 self.chat_messages.drain(system_prefix..drain_end);
436 }
437
438 pub fn pop_last_message(&mut self) {
441 self.chat_messages.pop();
442 }
443
444 pub fn close_dangling_tool_calls(&mut self, error_summary: &str) {
445 let assistant_idx = self.chat_messages.iter().rposition(
446 |m| matches!(m, ChatMessage::Assistant { tool_calls: Some(tc), .. } if !tc.is_empty()),
447 );
448
449 let Some(assistant_idx) = assistant_idx else {
450 return;
451 };
452
453 let ChatMessage::Assistant {
454 tool_calls: Some(tc),
455 ..
456 } = &self.chat_messages[assistant_idx]
457 else {
458 return;
459 };
460
461 let all_ids: Vec<String> = tc.iter().map(|t| t.id.clone()).collect();
462
463 let answered_ids: Vec<String> = self.chat_messages[assistant_idx + 1..]
464 .iter()
465 .filter_map(|m| match m {
466 ChatMessage::Tool { tool_call_id, .. } => Some(tool_call_id.clone()),
467 _ => None,
468 })
469 .collect();
470
471 for id in &all_ids {
472 if !answered_ids.iter().any(|a| a == id) {
473 self.push_tool_result(id, error_summary);
474 }
475 }
476 }
477
478 pub fn set_chat_messages(&mut self, messages: Vec<ChatMessage>) -> Result<(), String> {
483 validate_message_sequence(&messages)?;
484 self.total_tool_calls = messages
487 .iter()
488 .filter_map(|m| match m {
489 ChatMessage::Assistant {
490 tool_calls: Some(tc),
491 ..
492 } => Some(tc.len()),
493 _ => None,
494 })
495 .sum();
496 self.chat_messages = messages;
497 Ok(())
498 }
499}
500
501pub fn validate_message_sequence(messages: &[ChatMessage]) -> Result<(), String> {
514 if !messages
515 .iter()
516 .any(|m| !matches!(m, ChatMessage::System { .. } | ChatMessage::Custom { .. }))
517 {
518 return Err(
519 "sequence contains no sendable message: System/Custom alone leave the \
520 provider `messages` array empty"
521 .to_string(),
522 );
523 }
524
525 let mut pending_tool_call_ids: HashSet<String> = HashSet::new();
526
527 for (i, msg) in messages.iter().enumerate() {
528 match msg {
529 ChatMessage::Tool { tool_call_id, .. } => {
530 if pending_tool_call_ids.is_empty() {
531 return Err(format!(
532 "message[{}]: Tool message with call_id '{}' has no preceding tool_call",
533 i, tool_call_id
534 ));
535 }
536 if !pending_tool_call_ids.remove(tool_call_id) {
538 return Err(format!(
539 "message[{}]: Tool message with call_id '{}' does not match any pending tool_call (already answered or unknown)",
540 i, tool_call_id
541 ));
542 }
543 }
544 ChatMessage::Assistant {
545 tool_calls: Some(tc),
546 ..
547 } => {
548 if !pending_tool_call_ids.is_empty() {
550 return Err(format!(
551 "message[{}]: Assistant message with new tool_calls appears before pending calls were answered: {:?}",
552 i, pending_tool_call_ids
553 ));
554 }
555 pending_tool_call_ids = tc.iter().map(|t| t.id.clone()).collect();
556 }
557 _ => {}
558 }
559 }
560
561 if !pending_tool_call_ids.is_empty() {
563 return Err(format!(
564 "message sequence ends with unanswered tool calls: {:?}",
565 pending_tool_call_ids
566 ));
567 }
568
569 Ok(())
570}
571
572#[cfg(test)]
573fn make_session() -> AgentSession {
574 AgentSession::new(SessionId::new(1))
575}
576
577#[cfg(test)]
578mod tests {
579 use super::*;
580
581 #[test]
582 fn test_turn_count_empty() {
583 let s = make_session();
584 assert_eq!(s.turn_count(), 0);
585 }
586
587 #[test]
588 fn test_turn_count_with_system_and_user() {
589 let mut s = make_session();
590 s.push_message(MessageRole::System, "system");
591 assert_eq!(s.turn_count(), 0);
592 s.push_message(MessageRole::User, "hello");
593 assert_eq!(s.turn_count(), 1);
594 s.push_message(MessageRole::Assistant, "hi");
595 assert_eq!(s.turn_count(), 1);
596 s.push_message(MessageRole::User, "bye");
597 assert_eq!(s.turn_count(), 2);
598 }
599
600 #[test]
601 fn test_turn_count_with_tool_calls() {
602 let mut s = make_session();
603 s.push_message(MessageRole::User, "do something");
604 s.push_assistant_tool_calls(&[("id1".into(), "tool".into(), "{}".into())], None, None);
605 s.push_tool_result("id1", "result");
606 s.push_message(MessageRole::Assistant, "done");
607 assert_eq!(s.turn_count(), 1);
609 }
610
611 #[test]
612 fn test_trim_oldest_turns_noop() {
613 let mut s = make_session();
614 s.push_message(MessageRole::User, "hello");
615 s.push_message(MessageRole::Assistant, "hi");
616 s.trim_oldest_turns(5);
617 assert_eq!(s.turn_count(), 1);
618 assert_eq!(s.chat_messages().len(), 2);
619 }
620
621 #[test]
622 fn test_trim_oldest_turns_removes_old() {
623 let mut s = make_session();
624 s.push_message(MessageRole::System, "sys");
625 s.push_message(MessageRole::User, "u1");
627 s.push_message(MessageRole::Assistant, "a1");
628 s.push_message(MessageRole::User, "u2");
630 s.push_message(MessageRole::Assistant, "a2");
631 s.push_message(MessageRole::User, "u3");
633 s.push_message(MessageRole::Assistant, "a3");
634
635 s.trim_oldest_turns(2);
636 assert_eq!(s.turn_count(), 2);
637 assert!(matches!(s.chat_messages()[0], ChatMessage::System { .. }));
639 assert!(
641 matches!(s.chat_messages()[1], ChatMessage::User { ref content, .. } if content == "u2")
642 );
643 }
644
645 #[test]
646 fn test_trim_oldest_turns_with_tool_calls() {
647 let mut s = make_session();
648 s.push_message(MessageRole::User, "u1");
650 s.push_assistant_tool_calls(&[("id1".into(), "t".into(), "{}".into())], None, None);
651 s.push_tool_result("id1", "r1");
652 s.push_message(MessageRole::Assistant, "a1");
653 s.push_message(MessageRole::User, "u2");
655 s.push_message(MessageRole::Assistant, "a2");
656
657 let msg_before = s.simple_messages().len();
658 let chat_before = s.chat_messages().len();
659 s.trim_oldest_turns(1);
660 assert_eq!(s.turn_count(), 1);
661 assert_eq!(s.chat_messages().len(), chat_before - 4);
663 assert_eq!(s.simple_messages().len(), msg_before - 3);
665 }
666
667 #[test]
668 fn test_pop_last_message_text() {
669 let mut s = make_session();
670 s.push_message(MessageRole::User, "hello");
671 s.push_message(MessageRole::Assistant, "hi");
672 assert_eq!(s.chat_messages().len(), 2);
673 s.pop_last_message();
674 assert_eq!(s.chat_messages().len(), 1);
675 assert_eq!(s.simple_messages().len(), 1);
676 }
677
678 #[test]
679 fn test_pop_last_message_tool_calls_only() {
680 let mut s = make_session();
681 s.push_message(MessageRole::User, "do it");
682 s.push_assistant_tool_calls(&[("id1".into(), "t".into(), "{}".into())], None, None);
683 assert_eq!(s.chat_messages().len(), 2);
684 assert_eq!(s.simple_messages().len(), 1); s.pop_last_message();
686 assert_eq!(s.chat_messages().len(), 1);
687 assert_eq!(s.simple_messages().len(), 1); }
689
690 #[test]
691 fn test_pop_last_message_empty_session() {
692 let mut s = make_session();
693 s.pop_last_message(); assert_eq!(s.chat_messages().len(), 0);
695 }
696
697 #[test]
700 fn test_id_and_action_allowlist() {
701 let mut s = make_session();
702 assert_eq!(s.id(), Some(SessionId::new(1)));
703 assert!(!s.is_action_allowed("approve:rm"));
704 s.allow_action("approve:rm");
705 assert!(s.is_action_allowed("approve:rm"));
706 assert!(!s.is_action_allowed("approve:shell"));
707 }
708
709 #[test]
710 fn test_chat_messages_mut() {
711 let mut s = make_session();
712 s.chat_messages_mut().push(ChatMessage::user("direct"));
713 assert_eq!(s.chat_messages().len(), 1);
714 }
715
716 #[test]
717 fn test_push_message_tool_role() {
718 let mut s = make_session();
719 s.push_message(MessageRole::Tool, "result");
720 assert!(matches!(s.chat_messages()[0], ChatMessage::Tool { .. }));
721 }
722
723 #[test]
724 fn test_push_assistant_with_reasoning() {
725 let mut s = make_session();
726 s.push_assistant_with_reasoning("answer", "thinking");
727 match &s.chat_messages()[0] {
728 ChatMessage::Assistant {
729 content,
730 reasoning_content,
731 ..
732 } => {
733 assert_eq!(content.as_deref(), Some("answer"));
734 assert_eq!(reasoning_content.as_deref(), Some("thinking"));
735 }
736 other => panic!("unexpected message: {other:?}"),
737 }
738 }
739
740 #[test]
741 fn test_push_user_message_with_images() {
742 let mut s = make_session();
743 s.push_user_message_with_images(
744 "look",
745 vec![ImageAttachment::Url {
746 url: "http://x".into(),
747 detail: None,
748 }],
749 );
750 match &s.chat_messages()[0] {
751 ChatMessage::User { images, .. } => assert_eq!(images.len(), 1),
752 other => panic!("unexpected message: {other:?}"),
753 }
754 }
755
756 #[test]
757 fn test_push_assistant_tool_call_singular() {
758 let mut s = make_session();
759 s.push_assistant_tool_call("call_1", "bash", "{}");
760 match &s.chat_messages()[0] {
761 ChatMessage::Assistant {
762 tool_calls: Some(tc),
763 ..
764 } => {
765 assert_eq!(tc.len(), 1);
766 assert_eq!(tc[0].id, "call_1");
767 assert_eq!(tc[0].name, "bash");
768 }
769 other => panic!("unexpected message: {other:?}"),
770 }
771 }
772
773 #[test]
774 fn test_simple_messages_filters_empty_content_tool_calls() {
775 let mut s = make_session();
776 s.chat_messages_mut().push(ChatMessage::Assistant {
777 content: Some(String::new()),
778 reasoning_content: None,
779 tool_calls: Some(vec![ToolCallMessage {
780 id: "c".into(),
781 name: "t".into(),
782 arguments: "{}".into(),
783 }]),
784 thinking_signature: None,
785 });
786 assert!(s.simple_messages().is_empty());
787 }
788
789 #[test]
790 fn test_remove_ephemeral_messages() {
791 let mut s = make_session();
792 s.push_message(MessageRole::System, "keep");
793 s.chat_messages_mut()
794 .push(ChatMessage::user_ephemeral("temp"));
795 s.chat_messages_mut()
796 .push(ChatMessage::system_ephemeral("temp2"));
797 s.push_message(MessageRole::User, "keep2");
798 assert_eq!(s.chat_messages().len(), 4);
799 s.remove_ephemeral_messages();
800 assert_eq!(s.chat_messages().len(), 2);
801 assert!(s.chat_messages().iter().all(|m| !m.is_ephemeral()));
802 }
803
804 #[test]
805 fn test_set_system_prompt_replaces_first_non_ephemeral_system() {
806 let mut s = make_session();
807 s.push_message(MessageRole::System, "old prompt");
808 s.push_message(MessageRole::User, "hi");
809 s.push_message(MessageRole::Assistant, "hello");
810
811 s.set_system_prompt("new prompt");
812
813 let msgs = s.chat_messages();
814 assert_eq!(msgs.len(), 3, "history length unchanged");
815 assert!(
816 matches!(&msgs[0], ChatMessage::System { content, ephemeral: false } if content == "new prompt")
817 );
818 assert!(matches!(&msgs[1], ChatMessage::User { content, .. } if content == "hi"));
819 assert!(
820 matches!(&msgs[2], ChatMessage::Assistant { content: Some(c), .. } if c == "hello")
821 );
822 }
823
824 #[test]
825 fn test_set_system_prompt_inserts_when_absent() {
826 let mut s = make_session();
827 s.push_message(MessageRole::User, "hi");
828
829 s.set_system_prompt("fresh prompt");
830
831 let msgs = s.chat_messages();
832 assert_eq!(msgs.len(), 2);
833 assert!(
834 matches!(&msgs[0], ChatMessage::System { content, ephemeral: false } if content == "fresh prompt")
835 );
836 assert!(matches!(&msgs[1], ChatMessage::User { .. }));
837 }
838
839 #[test]
840 fn test_set_system_prompt_skips_ephemeral_system_and_inserts() {
841 let mut s = make_session();
842 s.chat_messages_mut()
843 .push(ChatMessage::system_ephemeral("ephemeral nudge"));
844 s.push_message(MessageRole::User, "hi");
845
846 s.set_system_prompt("real prompt");
847
848 let msgs = s.chat_messages();
850 assert_eq!(msgs.len(), 3);
851 assert!(
852 matches!(&msgs[0], ChatMessage::System { content, ephemeral: false } if content == "real prompt")
853 );
854 assert!(matches!(
855 &msgs[1],
856 ChatMessage::System {
857 ephemeral: true,
858 ..
859 }
860 ));
861 assert!(matches!(&msgs[2], ChatMessage::User { .. }));
862 }
863
864 #[test]
865 fn test_close_dangling_tool_calls_noop_without_tool_call() {
866 let mut s = make_session();
867 s.push_message(MessageRole::User, "hi");
868 s.push_message(MessageRole::Assistant, "hi");
869 s.close_dangling_tool_calls("failed");
870 assert_eq!(s.chat_messages().len(), 2);
871 }
872
873 #[test]
874 fn test_close_dangling_tool_calls_adds_missing_results() {
875 let mut s = make_session();
876 s.push_message(MessageRole::User, "do");
877 s.push_assistant_tool_calls(
878 &[
879 ("c1".into(), "t".into(), "{}".into()),
880 ("c2".into(), "t".into(), "{}".into()),
881 ],
882 None,
883 None,
884 );
885 s.push_tool_result("c1", "ok"); s.close_dangling_tool_calls("failed");
887
888 let tool_results: Vec<(String, String)> = s
889 .chat_messages()
890 .iter()
891 .filter_map(|m| match m {
892 ChatMessage::Tool {
893 tool_call_id,
894 name: _,
895 content,
896 } => Some((tool_call_id.clone(), content.clone())),
897 _ => None,
898 })
899 .collect();
900 assert_eq!(tool_results.len(), 2);
901 assert!(
902 tool_results
903 .iter()
904 .any(|(id, c)| id == "c2" && c == "failed")
905 );
906 }
907
908 #[test]
909 fn test_set_chat_messages_recalculates_total_tool_calls() {
910 let mut s = make_session();
911 let msgs = vec![
912 ChatMessage::user("do"),
913 ChatMessage::assistant_tool_call("c1", "t", "{}"),
914 ChatMessage::tool("c1", "result"),
915 ];
916 s.set_chat_messages(msgs).unwrap();
917 assert_eq!(s.total_tool_calls, 1);
918 }
919}
920
921#[cfg(test)]
922mod validate_tests {
923 use super::*;
924
925 #[test]
926 fn test_valid_simple_sequence() {
927 let msgs = vec![ChatMessage::user("hello"), ChatMessage::assistant("hi")];
928 assert!(validate_message_sequence(&msgs).is_ok());
929 }
930
931 #[test]
932 fn test_valid_tool_call_sequence() {
933 let msgs = vec![
934 ChatMessage::user("run command"),
935 ChatMessage::assistant_tool_call("call_1", "bash", r#"{"cmd":"ls"}"#),
936 ChatMessage::tool("call_1", "file1 file2"),
937 ChatMessage::assistant("done"),
938 ];
939 assert!(validate_message_sequence(&msgs).is_ok());
940 }
941
942 #[test]
943 fn test_valid_multi_tool_call_sequence() {
944 let msgs = vec![
945 ChatMessage::user("run commands"),
946 ChatMessage::Assistant {
947 content: None,
948 reasoning_content: None,
949 tool_calls: Some(vec![
950 crate::types::ToolCallMessage {
951 id: "call_1".into(),
952 name: "bash".into(),
953 arguments: "{}".into(),
954 },
955 crate::types::ToolCallMessage {
956 id: "call_2".into(),
957 name: "read".into(),
958 arguments: "{}".into(),
959 },
960 ]),
961 thinking_signature: None,
962 },
963 ChatMessage::tool("call_1", "result1"),
964 ChatMessage::tool("call_2", "result2"),
965 ChatMessage::assistant("done"),
966 ];
967 assert!(validate_message_sequence(&msgs).is_ok());
968 }
969
970 #[test]
971 fn test_orphaned_tool_result() {
972 let msgs = vec![
973 ChatMessage::user("hello"),
974 ChatMessage::tool("call_1", "orphaned result"),
975 ];
976 let err = validate_message_sequence(&msgs).unwrap_err();
977 assert!(err.contains("no preceding tool_call"));
978 }
979
980 #[test]
981 fn test_system_only_sequence_rejected() {
982 let msgs = vec![
986 ChatMessage::system("prompt"),
987 ChatMessage::system_ephemeral("reminder"),
988 ChatMessage::Custom {
989 role: "artifact".into(),
990 data: serde_json::json!({"id": "x"}),
991 },
992 ];
993 let err = validate_message_sequence(&msgs).unwrap_err();
994 assert!(err.contains("no sendable message"));
995 }
996
997 #[test]
998 fn test_system_plus_user_ok() {
999 let msgs = vec![ChatMessage::system("prompt"), ChatMessage::user("hi")];
1000 assert!(validate_message_sequence(&msgs).is_ok());
1001 }
1002
1003 #[test]
1004 fn test_mismatched_tool_call_id() {
1005 let msgs = vec![
1006 ChatMessage::user("run"),
1007 ChatMessage::assistant_tool_call("call_1", "bash", "{}"),
1008 ChatMessage::tool("call_2", "wrong id"),
1009 ];
1010 let err = validate_message_sequence(&msgs).unwrap_err();
1011 assert!(err.contains("does not match"));
1012 }
1013
1014 #[test]
1015 fn test_set_chat_messages_valid() {
1016 let mut s = make_session();
1017 let msgs = vec![ChatMessage::user("hello"), ChatMessage::assistant("hi")];
1018 assert!(s.set_chat_messages(msgs.clone()).is_ok());
1019 assert_eq!(s.chat_messages().len(), 2);
1020 }
1021
1022 #[test]
1023 fn test_set_chat_messages_invalid() {
1024 let mut s = make_session();
1025 let msgs = vec![ChatMessage::tool("call_1", "orphaned")];
1026 assert!(s.set_chat_messages(msgs).is_err());
1027 }
1028
1029 #[test]
1032 fn run_state_default() {
1033 let rs = RunState::default();
1034 assert_eq!(rs.turn_tool_calls, 0);
1035 assert!(!rs.run_has_tool_calls);
1036 assert_eq!(rs.reasoning_only_strikes, 0);
1037 assert_eq!(rs.empty_response_strikes, 0);
1038 assert_eq!(rs.nudge_count, 0);
1039 }
1040
1041 #[test]
1042 fn run_state_reset_for_new_run() {
1043 let mut rs = RunState {
1044 turn_tool_calls: 5,
1045 run_has_tool_calls: true,
1046 reasoning_only_strikes: 2,
1047 empty_response_strikes: 1,
1048 nudge_count: 3,
1049 thinking_disabled_for_rest_of_run: true,
1050 original_thinking_enabled: true,
1051 truncation_strikes: 4,
1052 };
1053
1054 rs.reset_for_new_run();
1055
1056 assert_eq!(rs.turn_tool_calls, 0);
1057 assert!(!rs.run_has_tool_calls);
1058 assert_eq!(rs.reasoning_only_strikes, 0);
1059 assert_eq!(rs.empty_response_strikes, 0);
1060 assert_eq!(rs.nudge_count, 0);
1061 assert_eq!(rs.truncation_strikes, 0);
1062 assert!(!rs.thinking_disabled_for_rest_of_run);
1063 assert!(rs.original_thinking_enabled);
1065 }
1066
1067 #[test]
1068 fn run_state_record_tool_calls() {
1069 let mut rs = RunState {
1070 reasoning_only_strikes: 2,
1071 empty_response_strikes: 1,
1072 ..RunState::default()
1073 };
1074
1075 rs.record_tool_calls(3);
1076
1077 assert_eq!(rs.turn_tool_calls, 3);
1078 assert!(rs.run_has_tool_calls);
1079 assert_eq!(rs.reasoning_only_strikes, 0); assert_eq!(rs.empty_response_strikes, 0); }
1082
1083 #[test]
1084 fn run_state_record_tool_calls_accumulates() {
1085 let mut rs = RunState::default();
1086 rs.record_tool_calls(2);
1087 rs.record_tool_calls(3);
1088
1089 assert_eq!(rs.turn_tool_calls, 5);
1090 assert!(rs.run_has_tool_calls);
1091 }
1092
1093 #[test]
1094 fn run_state_record_reasoning_only() {
1095 let mut rs = RunState {
1096 empty_response_strikes: 2,
1097 ..RunState::default()
1098 };
1099
1100 let strikes = rs.record_reasoning_only();
1101
1102 assert_eq!(strikes, 1);
1103 assert_eq!(rs.reasoning_only_strikes, 1);
1104 assert_eq!(rs.empty_response_strikes, 0); }
1106
1107 #[test]
1108 fn run_state_record_reasoning_only_consecutive() {
1109 let mut rs = RunState::default();
1110
1111 assert_eq!(rs.record_reasoning_only(), 1);
1112 assert_eq!(rs.record_reasoning_only(), 2);
1113 assert_eq!(rs.record_reasoning_only(), 3);
1114 }
1115
1116 #[test]
1117 fn run_state_record_empty_response() {
1118 let mut rs = RunState {
1119 reasoning_only_strikes: 2,
1120 ..RunState::default()
1121 };
1122
1123 let strikes = rs.record_empty_response();
1124
1125 assert_eq!(strikes, 1);
1126 assert_eq!(rs.empty_response_strikes, 1);
1127 assert_eq!(rs.reasoning_only_strikes, 0); }
1129
1130 #[test]
1131 fn run_state_record_empty_response_consecutive() {
1132 let mut rs = RunState::default();
1133
1134 assert_eq!(rs.record_empty_response(), 1);
1135 assert_eq!(rs.record_empty_response(), 2);
1136 assert_eq!(rs.record_empty_response(), 3);
1137 }
1138
1139 #[test]
1140 fn run_state_branch_cross_reset() {
1141 let mut rs = RunState::default();
1143
1144 rs.record_reasoning_only();
1146 assert_eq!(rs.reasoning_only_strikes, 1);
1147
1148 rs.record_tool_calls(2);
1150 assert_eq!(rs.reasoning_only_strikes, 0);
1151 assert_eq!(rs.turn_tool_calls, 2);
1152
1153 let strikes = rs.record_reasoning_only();
1155 assert_eq!(strikes, 1);
1156 }
1157
1158 #[test]
1159 fn run_state_empty_to_reasoning_reset() {
1160 let mut rs = RunState::default();
1162
1163 rs.record_empty_response();
1164 rs.record_empty_response();
1165 assert_eq!(rs.empty_response_strikes, 2);
1166
1167 rs.record_reasoning_only();
1169 assert_eq!(rs.empty_response_strikes, 0);
1170 assert_eq!(rs.reasoning_only_strikes, 1);
1171 }
1172
1173 #[test]
1174 fn run_state_thinking_disabled_default() {
1175 let rs = RunState::default();
1176 assert!(!rs.thinking_disabled_for_rest_of_run);
1177 }
1178
1179 #[test]
1180 fn run_state_thinking_disabled_after_3_strikes() {
1181 let mut rs = RunState::default();
1182
1183 rs.record_reasoning_only();
1185 assert!(!rs.thinking_disabled_for_rest_of_run);
1186 assert_eq!(rs.reasoning_only_strikes, 1);
1187
1188 rs.record_reasoning_only();
1190 assert!(!rs.thinking_disabled_for_rest_of_run);
1191 assert_eq!(rs.reasoning_only_strikes, 2);
1192
1193 rs.record_reasoning_only();
1195 assert!(rs.thinking_disabled_for_rest_of_run);
1196 assert_eq!(rs.reasoning_only_strikes, 3);
1197 }
1198
1199 #[test]
1200 fn run_state_thinking_disabled_resets_on_new_run() {
1201 let mut rs = RunState::default();
1202
1203 rs.record_reasoning_only();
1205 rs.record_reasoning_only();
1206 rs.record_reasoning_only();
1207 assert!(rs.thinking_disabled_for_rest_of_run);
1208 assert_eq!(rs.reasoning_only_strikes, 3);
1209
1210 rs.reset_for_new_run();
1212 assert!(!rs.thinking_disabled_for_rest_of_run);
1213 assert_eq!(rs.reasoning_only_strikes, 0);
1214 }
1215
1216 #[test]
1217 fn run_state_thinking_disabled_stays_after_tool_calls() {
1218 let mut rs = RunState::default();
1219
1220 rs.record_reasoning_only();
1222 rs.record_reasoning_only();
1223 rs.record_reasoning_only();
1224 assert!(rs.thinking_disabled_for_rest_of_run);
1225
1226 rs.record_tool_calls(2);
1229 assert!(rs.thinking_disabled_for_rest_of_run);
1230 assert_eq!(rs.reasoning_only_strikes, 0); }
1232
1233 #[test]
1236 fn deserialize_legacy_flat_fields() {
1237 let json = r#"{
1239 "id": null,
1240 "chat_messages": [],
1241 "always_allowed_actions": [],
1242 "total_tool_calls": 5,
1243 "nudge_count": 3,
1244 "turn_tool_calls": 2,
1245 "reasoning_only_strikes": 1,
1246 "empty_response_strikes": 0
1247 }"#;
1248 let session: AgentSession = serde_json::from_str(json).unwrap();
1249 assert_eq!(session.run_state.nudge_count, 3);
1250 assert_eq!(session.run_state.turn_tool_calls, 2);
1251 assert_eq!(session.run_state.reasoning_only_strikes, 1);
1252 assert_eq!(session.run_state.empty_response_strikes, 0);
1253 assert!(!session.run_state.run_has_tool_calls); }
1255
1256 #[test]
1257 fn deserialize_new_run_state_format() {
1258 let json = r#"{
1260 "id": null,
1261 "chat_messages": [],
1262 "always_allowed_actions": [],
1263 "total_tool_calls": 5,
1264 "run_state": {
1265 "turn_tool_calls": 4,
1266 "run_has_tool_calls": true,
1267 "reasoning_only_strikes": 0,
1268 "empty_response_strikes": 1,
1269 "nudge_count": 2
1270 }
1271 }"#;
1272 let session: AgentSession = serde_json::from_str(json).unwrap();
1273 assert_eq!(session.run_state.turn_tool_calls, 4);
1274 assert!(session.run_state.run_has_tool_calls);
1275 assert_eq!(session.run_state.empty_response_strikes, 1);
1276 assert_eq!(session.run_state.nudge_count, 2);
1277 }
1278
1279 #[test]
1280 fn deserialize_run_state_takes_precedence_over_flat() {
1281 let json = r#"{
1283 "id": null,
1284 "chat_messages": [],
1285 "always_allowed_actions": [],
1286 "total_tool_calls": 0,
1287 "run_state": {
1288 "turn_tool_calls": 10,
1289 "run_has_tool_calls": true,
1290 "reasoning_only_strikes": 0,
1291 "empty_response_strikes": 0,
1292 "nudge_count": 0
1293 },
1294 "nudge_count": 99,
1295 "turn_tool_calls": 99
1296 }"#;
1297 let session: AgentSession = serde_json::from_str(json).unwrap();
1298 assert_eq!(session.run_state.turn_tool_calls, 10); assert_eq!(session.run_state.nudge_count, 0); }
1301
1302 #[test]
1303 fn deserialize_legacy_missing_optional_fields() {
1304 let json = r#"{
1306 "id": null,
1307 "chat_messages": [],
1308 "always_allowed_actions": [],
1309 "total_tool_calls": 0,
1310 "nudge_count": 1
1311 }"#;
1312 let session: AgentSession = serde_json::from_str(json).unwrap();
1313 assert_eq!(session.run_state.nudge_count, 1);
1314 assert_eq!(session.run_state.turn_tool_calls, 0); assert_eq!(session.run_state.reasoning_only_strikes, 0);
1316 assert_eq!(session.run_state.empty_response_strikes, 0);
1317 }
1318
1319 #[test]
1320 fn roundtrip_preserves_run_state() {
1321 let mut session = AgentSession::new(SessionId::new(1));
1322 session.run_state.nudge_count = 5;
1323 session.run_state.turn_tool_calls = 3;
1324 session.run_state.run_has_tool_calls = true;
1325 session.run_state.reasoning_only_strikes = 2;
1326
1327 let json = serde_json::to_string(&session).unwrap();
1328 let restored: AgentSession = serde_json::from_str(&json).unwrap();
1329 assert_eq!(restored.run_state.nudge_count, 5);
1330 assert_eq!(restored.run_state.turn_tool_calls, 3);
1331 assert!(restored.run_state.run_has_tool_calls);
1332 assert_eq!(restored.run_state.reasoning_only_strikes, 2);
1333 }
1334
1335 #[test]
1336 fn push_assistant_tool_calls_validates_json_args() {
1337 let mut s = make_session();
1338
1339 let valid_args = r#"{"path": "src/main.rs", "content": "fn main() {}"}"#;
1341 s.push_assistant_tool_calls(
1342 &[("id1".into(), "write_file".into(), valid_args.into())],
1343 None,
1344 None,
1345 );
1346 if let ChatMessage::Assistant {
1347 tool_calls: Some(ref tc),
1348 ..
1349 } = s.chat_messages[0]
1350 {
1351 assert_eq!(tc[0].arguments, valid_args);
1352 } else {
1353 panic!("expected Assistant message with tool_calls");
1354 }
1355
1356 let truncated_args = r#"{"path": "src/ui/markdown.rs", "content": "#;
1366 s.push_assistant_tool_calls(
1367 &[("id2".into(), "write_file".into(), truncated_args.into())],
1368 None,
1369 None,
1370 );
1371 if let ChatMessage::Assistant {
1372 tool_calls: Some(ref tc),
1373 ..
1374 } = s.chat_messages[1]
1375 {
1376 assert_eq!(tc[0].arguments, "{}");
1377 assert!(
1378 !tc[0].arguments.contains("tool_call_arguments_truncated"),
1379 "assistant arguments must not carry the poison wrapper object"
1380 );
1381 } else {
1382 panic!("expected Assistant message with tool_calls");
1383 }
1384 }
1385
1386 #[test]
1387 fn push_assistant_tool_calls_truncated_multibyte_no_panic() {
1388 let mut s = make_session();
1391
1392 let mut bad_args = "あ".repeat(70); bad_args.push_str("truncated"); s.push_assistant_tool_calls(&[("id1".into(), "tool".into(), bad_args)], None, None);
1400
1401 if let ChatMessage::Assistant {
1402 tool_calls: Some(ref tc),
1403 ..
1404 } = s.chat_messages[0]
1405 {
1406 assert_eq!(tc[0].arguments, "{}");
1407 } else {
1408 panic!("expected Assistant message with tool_calls");
1409 }
1410 }
1411
1412 #[test]
1413 fn push_assistant_tool_calls_then_tool_result_matches_anthropic_protocol() {
1414 let mut s = make_session();
1419
1420 let tool_calls = vec![
1422 (
1423 "call_00_VJtlnKha0ZZ2Yo8t5ysQ8883".to_string(),
1424 "write_file".to_string(),
1425 "{}".to_string(),
1426 ),
1427 (
1428 "call_01_abc123".to_string(),
1429 "bash".to_string(),
1430 r#"{"command": "ls"}"#.to_string(),
1431 ),
1432 ];
1433
1434 s.push_assistant_tool_calls(
1436 &tool_calls,
1437 Some("thinking...".to_string()),
1438 Some("I'll help you".to_string()),
1439 );
1440
1441 for (tc_id, _, _) in &tool_calls {
1443 s.push_tool_result(
1444 tc_id,
1445 "Tool call was not executed: the response hit the output token limit.",
1446 );
1447 }
1448
1449 if let ChatMessage::Assistant {
1452 tool_calls: Some(ref tc),
1453 ..
1454 } = s.chat_messages[0]
1455 {
1456 assert_eq!(tc.len(), 2);
1457 assert_eq!(tc[0].id, "call_00_VJtlnKha0ZZ2Yo8t5ysQ8883");
1458 assert_eq!(tc[0].name, "write_file");
1459 assert_eq!(tc[1].id, "call_01_abc123");
1460 assert_eq!(tc[1].name, "bash");
1461 } else {
1462 panic!("expected Assistant message with tool_calls");
1463 }
1464
1465 if let ChatMessage::Tool { tool_call_id, .. } = &s.chat_messages[1] {
1467 assert_eq!(tool_call_id, "call_00_VJtlnKha0ZZ2Yo8t5ysQ8883");
1468 } else {
1469 panic!("expected Tool message for first tool call");
1470 }
1471
1472 if let ChatMessage::Tool { tool_call_id, .. } = &s.chat_messages[2] {
1473 assert_eq!(tool_call_id, "call_01_abc123");
1474 } else {
1475 panic!("expected Tool message for second tool call");
1476 }
1477
1478 assert!(
1480 validate_message_sequence(&s.chat_messages).is_ok(),
1481 "message sequence should be valid with matching tool_use and tool_result"
1482 );
1483 }
1484}
1485
1486#[cfg(test)]
1487mod proptest_tests {
1488 use super::*;
1489 use proptest::prelude::*;
1490
1491 proptest! {
1492 #[test]
1495 fn reset_for_new_run_zeros_all_fields(
1496 turn_tool_calls in 0usize..1000,
1497 run_has_tool_calls in proptest::bool::ANY,
1498 reasoning_only_strikes in 0usize..100,
1499 empty_response_strikes in 0usize..100,
1500 nudge_count in 0usize..100,
1501 ) {
1502 let mut rs = RunState {
1503 turn_tool_calls,
1504 run_has_tool_calls,
1505 reasoning_only_strikes,
1506 empty_response_strikes,
1507 nudge_count,
1508 thinking_disabled_for_rest_of_run: true,
1509 original_thinking_enabled: true,
1510 truncation_strikes: 5,
1511 };
1512 rs.reset_for_new_run();
1513 assert_eq!(rs.turn_tool_calls, 0);
1514 assert!(!rs.run_has_tool_calls);
1515 assert_eq!(rs.reasoning_only_strikes, 0);
1516 assert_eq!(rs.empty_response_strikes, 0);
1517 assert_eq!(rs.nudge_count, 0);
1518 assert_eq!(rs.truncation_strikes, 0);
1519 assert!(!rs.thinking_disabled_for_rest_of_run);
1520 assert!(rs.original_thinking_enabled);
1522 }
1523
1524 #[test]
1525 fn record_tool_calls_accumulates(n in 0usize..100) {
1526 let mut rs = RunState::default();
1527 rs.record_tool_calls(n);
1528 assert_eq!(rs.turn_tool_calls, n);
1529 assert!(rs.run_has_tool_calls);
1530 assert_eq!(rs.reasoning_only_strikes, 0);
1531 assert_eq!(rs.empty_response_strikes, 0);
1532 }
1533
1534 #[test]
1535 fn record_reasoning_only_increments(count in 1usize..50) {
1536 let mut rs = RunState::default();
1537 for i in 1..=count {
1538 let strikes = rs.record_reasoning_only();
1539 assert_eq!(strikes, i);
1540 assert_eq!(rs.empty_response_strikes, 0);
1541 }
1542 }
1543
1544 #[test]
1545 fn record_empty_response_increments(count in 1usize..50) {
1546 let mut rs = RunState::default();
1547 for i in 1..=count {
1548 let strikes = rs.record_empty_response();
1549 assert_eq!(strikes, i);
1550 assert_eq!(rs.reasoning_only_strikes, 0);
1551 }
1552 }
1553
1554 #[test]
1557 fn push_assistant_tool_calls_valid_json_unchanged(args in r"\{[^{}]{0,200}\}") {
1558 if serde_json::from_str::<serde_json::Value>(&args).is_err() {
1560 return Ok(());
1561 }
1562 let mut s = make_session();
1563 s.push_assistant_tool_calls(
1564 &[("id".into(), "tool".into(), args.clone())],
1565 None,
1566 None,
1567 );
1568 if let ChatMessage::Assistant { tool_calls: Some(ref tc), .. } = s.chat_messages()[0] {
1569 assert_eq!(tc[0].arguments, args);
1570 } else {
1571 panic!("expected Assistant with tool_calls");
1572 }
1573 }
1574
1575 #[test]
1576 fn push_assistant_tool_calls_invalid_json_sanitized_to_empty(
1577 bad_args in "[a-z\u{4e00}-\u{9fff}]{0,300}"
1578 ) {
1579 if serde_json::from_str::<serde_json::Value>(&bad_args).is_ok() {
1581 return Ok(());
1582 }
1583 let mut s = make_session();
1584 s.push_assistant_tool_calls(
1585 &[("id".into(), "tool".into(), bad_args)],
1586 None,
1587 None,
1588 );
1589 if let ChatMessage::Assistant { tool_calls: Some(ref tc), .. } = s.chat_messages()[0] {
1590 serde_json::from_str::<serde_json::Value>(&tc[0].arguments)
1593 .expect("sanitized args must be valid JSON");
1594 assert_eq!(tc[0].arguments, "{}");
1595 } else {
1596 panic!("expected Assistant with tool_calls");
1597 }
1598 }
1599
1600 #[test]
1603 fn trim_oldest_turns_never_exceeds_max(turns in 1usize..20, max in 1usize..20) {
1604 let mut s = make_session();
1605 for i in 0..turns {
1606 s.push_message(MessageRole::User, format!("u{}", i));
1607 s.push_message(MessageRole::Assistant, format!("a{}", i));
1608 }
1609 s.trim_oldest_turns(max);
1610 assert!(s.turn_count() <= max || turns <= max);
1611 }
1612
1613 #[test]
1616 fn validate_simple_user_assistant_always_passes(count in 1usize..20) {
1617 let mut msgs = Vec::new();
1618 for i in 0..count {
1619 msgs.push(ChatMessage::user(format!("msg{}", i)));
1620 msgs.push(ChatMessage::assistant(format!("reply{}", i)));
1621 }
1622 assert!(validate_message_sequence(&msgs).is_ok());
1623 }
1624 }
1625}