Skip to main content

agent_base/engine/
session.rs

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/// Run-level state tracking for the react loop.
10///
11/// Manages counters and flags that track the current run's progress.
12/// All fields are reset at the start of each run (when a new user message arrives).
13///
14/// # Backward compatibility
15///
16/// Before this struct existed, `nudge_count` and `turn_tool_calls` were flat
17/// fields on `AgentSession`.  The [`RawAgentSession`] deserialization shim
18/// migrates them automatically.
19#[derive(Clone, Debug, Default, Serialize, Deserialize)]
20pub struct RunState {
21    /// Number of tool calls already executed in the current turn.
22    /// Reset to 0 at the start of each turn (when a new user message arrives).
23    /// Used by `TurnToolLimitMiddleware` to enforce per-turn tool call limits.
24    pub turn_tool_calls: usize,
25    /// Whether any tools were called in the current run.
26    /// Reset to false at the start of each run. Used by the completion judge
27    /// to determine if tools were used before a text-only response.
28    pub run_has_tool_calls: bool,
29    /// Number of consecutive LLM turns that produced only reasoning_content
30    /// (no text, no tool call). Reset to 0 when a normal response or tool call
31    /// is produced. Used by the react loop to fail instead of looping forever
32    /// on a reasoning-model runaway.
33    pub reasoning_only_strikes: usize,
34    /// Number of consecutive LLM turns that produced a completely empty response
35    /// (no text, no reasoning, no tool call). Reset to 0 when a normal response
36    /// or tool call is produced. Used by the react loop to retry a bounded
37    /// number of times, then fail instead of looping forever.
38    pub empty_response_strikes: usize,
39    /// Number of tool-enforcement nudges issued in the current turn.
40    /// Reset to 0 at the start of each turn (when a new user message arrives).
41    /// Used by `ToolEnforcementMiddleware` to cap nudge attempts per turn.
42    pub nudge_count: usize,
43    /// Thinking is disabled for the rest of the current run.
44    /// Set to true when reasoning_only_strikes reaches the maximum (3).
45    /// Reset to false at the start of each new run (when a new user message arrives).
46    /// This allows the model to continue without thinking after too many
47    /// reasoning-only responses.
48    #[serde(default)]
49    pub thinking_disabled_for_rest_of_run: bool,
50
51    // ─── New fields for thinking guard ─────────────────────────────
52    /// Original thinking configuration (for restoration)
53    ///
54    /// Records whether thinking was originally enabled when the session started.
55    /// Used to restore thinking to its original state after DisableThinking → RestoreThinking cycle.
56    #[serde(default)]
57    pub original_thinking_enabled: bool,
58
59    /// Number of consecutive LLM turns where the truncation guard fired
60    /// (all tool_calls had invalid/truncated arguments). Reset to 0 when
61    /// at least one valid tool call executes successfully. Used to break
62    /// the re-issue death spiral: after N consecutive truncations, stop
63    /// retrying and fail with an actionable message.
64    #[serde(default)]
65    pub truncation_strikes: usize,
66}
67
68impl RunState {
69    /// Reset all run-level state for a new run (when a new user message arrives).
70    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        // Note: original_thinking_enabled is NOT reset here
79        // It should be set once when the session starts
80    }
81
82    /// Record tool calls (branch 3: tool calls).
83    /// Resets reasoning_only_strikes, empty_response_strikes, and truncation_strikes.
84    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    /// Record reasoning-only response (branch 1).
93    /// Resets empty_response_strikes.
94    /// Returns the new strike count.
95    /// If strikes reach 3, disables thinking for the rest of the run.
96    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    /// Record empty response (branch 2).
110    /// Resets reasoning_only_strikes.
111    /// Returns the new strike count.
112    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    /// Record a truncation guard firing (all tool_calls had invalid arguments).
119    /// Returns the new strike count. Caller should check against threshold.
120    pub fn record_truncation(&mut self) -> usize {
121        self.truncation_strikes += 1;
122        self.truncation_strikes
123    }
124}
125
126/// Deserialization shim for backward compatibility.
127///
128/// Before `RunState` was introduced, `nudge_count`, `turn_tool_calls`,
129/// `reasoning_only_strikes`, and `empty_response_strikes` were flat fields
130/// on `AgentSession`.  This struct accepts **both** the old flat format and
131/// the new nested `run_state` format, migrating legacy data on the fly.
132#[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    // ── new format (preferred) ──
140    run_state: Option<RunState>,
141
142    // ── legacy flat fields (fallback) ──
143    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/// Stable identity + rolling state of a single chat thread.
169///
170/// `Default` returns a fresh session; load persisted state via serde.
171/// Backward-compatible with the old flat-field format (pre-`RunState` migration).
172#[derive(Clone, Debug, Default, Serialize, Deserialize)]
173#[serde(from = "RawAgentSession")]
174pub struct AgentSession {
175    id: Option<SessionId>,
176    /// LLM API format messages, sent directly to the provider.
177    /// This is the single source of truth for the conversation state.
178    chat_messages: Vec<ChatMessage>,
179    always_allowed_actions: HashSet<String>,
180    /// Total number of tool calls made in this session (across all turns).
181    /// Used by middleware for decisions like "first_turn_only" enforcement.
182    pub total_tool_calls: usize,
183    /// Run-level state tracking for the react loop.
184    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    /// Derive a simplified `Vec<Message>` view from the canonical `chat_messages`.
203    /// Assistant messages that contain only tool_calls (no text content) are
204    /// skipped, since they have no corresponding simplified representation.
205    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    /// 可变引用,仅用于需要直接操作消息的高级场景。
225    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    /// Push a message as ephemeral (user/system only): the LLM sees it this
249    /// turn, then `remove_ephemeral_messages` strips it at turn end — from
250    /// both memory and persistence.
251    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                // Ephemeral assistant/tool messages are not a thing; fall back
258                // to a regular push rather than silently mis-marking.
259                self.push_message(role, content);
260                return;
261            }
262        };
263        self.chat_messages.push(chat_msg);
264    }
265
266    /// Replace the system prompt — the first **non-ephemeral** System message
267    /// (wherever it sits, matching how `context` locates it); inserted at
268    /// index 0 if none exists. All other messages are untouched. Hosts that
269    /// own prompt composition (e.g. skill activation) re-bake through this.
270    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    /// Push an assistant message with reasoning/thinking content preserved.
288    /// This allows the LLM to see its own prior reasoning in subsequent turns,
289    /// preventing it from re-deriving the same conclusions every turn.
290    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                // Invalid JSON args (provider truncated the stream mid-generation)
331                // are sanitized to an empty object. We deliberately do NOT wrap
332                // them in a descriptive `{error, original_args_preview, message}`
333                // object: "{}" is valid JSON so the next request is never rejected
334                // with 400, and — critically — a rich error object here becomes an
335                // imitation vector: the model replays it verbatim as the next
336                // call's arguments and, because it parses, it slips past the react
337                // truncation guard. The truncation explanation lives in the
338                // paired tool_result, which is where the model reads feedback.
339                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    /// 移除所有临时消息(ephemeral=true)。
371    ///
372    /// 在 turn 结束时调用,确保注入的临时内容不残留到下一轮。
373    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    /// Count the number of conversation turns.
387    /// A turn starts with a User message and includes subsequent Assistant/Tool messages.
388    pub fn turn_count(&self) -> usize {
389        self.chat_messages
390            .iter()
391            .filter(|m| matches!(m, ChatMessage::User { .. }))
392            .count()
393    }
394
395    /// Remove the oldest turns from the front until turn count ≤ max_turns.
396    /// Preserves the System message at index 0 if present.
397    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        // Find User message positions (turn boundaries) in chat_messages
405        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        // Preserve system prefix: count leading System messages
423        let system_prefix = self
424            .chat_messages
425            .iter()
426            .take_while(|m| matches!(m, ChatMessage::System { .. }))
427            .count();
428
429        // Drain from system_prefix up to the start of the (turns_to_remove + 1)-th turn
430        let drain_end = user_positions[turns_to_remove];
431        if system_prefix >= drain_end {
432            return; // nothing to drain after system messages
433        }
434
435        self.chat_messages.drain(system_prefix..drain_end);
436    }
437
438    /// Remove the last message from `chat_messages`.
439    /// Used by the max_message_tokens safety valve to discard oversized messages.
440    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    /// Replace chat messages — only for persistence restore.
479    /// Validates message sequence before replacing.
480    ///
481    /// 仅供持久化恢复使用。调用方必须保证 messages 序列合法。
482    pub fn set_chat_messages(&mut self, messages: Vec<ChatMessage>) -> Result<(), String> {
483        validate_message_sequence(&messages)?;
484        // Recalculate total_tool_calls from the incoming messages so middleware
485        // decisions (e.g. first_turn_only enforcement) see the correct count.
486        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
501/// Validate that a chat message sequence is well-formed for LLM API consumption.
502///
503/// Checks:
504/// - At least one non-System/non-Custom message (System maps to the
505///   top-level `system` parameter, Custom is stripped — a sequence of only
506///   those leaves an empty `messages` array, which providers reject with
507///   HTTP 400). Guards compaction outputs against producing an unsendable
508///   window.
509/// - No Tool message without a preceding Assistant with matching tool_call
510/// - No duplicate Tool messages for the same tool_call_id
511/// - All tool_calls in an Assistant batch must be answered before the next Assistant batch
512/// - No unanswered tool calls at the end of the sequence
513pub 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                // Remove the ID on match — also detects duplicates (second remove returns false)
537                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                // Previous batch must be fully answered before a new batch starts
549                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    // All tool calls must be answered by the end of the sequence
562    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        // One user turn: User -> Assistant(tool_calls) -> Tool -> Assistant(text)
608        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        // Turn 1
626        s.push_message(MessageRole::User, "u1");
627        s.push_message(MessageRole::Assistant, "a1");
628        // Turn 2
629        s.push_message(MessageRole::User, "u2");
630        s.push_message(MessageRole::Assistant, "a2");
631        // Turn 3
632        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        // System message preserved
638        assert!(matches!(s.chat_messages()[0], ChatMessage::System { .. }));
639        // Oldest user message is u2
640        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        // Turn 1 with tool call
649        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        // Turn 2
654        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        // chat_messages should have lost 4 entries (User, Assistant(tool), Tool, Assistant(text))
662        assert_eq!(s.chat_messages().len(), chat_before - 4);
663        // simple_messages (derived from chat_messages, tool_calls-only filtered) loses 3 entries
664        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); // only User in simple_messages (tool_calls-only filtered)
685        s.pop_last_message();
686        assert_eq!(s.chat_messages().len(), 1);
687        assert_eq!(s.simple_messages().len(), 1); // simple_messages unchanged (still just User)
688    }
689
690    #[test]
691    fn test_pop_last_message_empty_session() {
692        let mut s = make_session();
693        s.pop_last_message(); // should not panic
694        assert_eq!(s.chat_messages().len(), 0);
695    }
696
697    // ── B5: remaining session lifecycle paths ──────────────────────────────
698
699    #[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        // 非临时 System 不存在 → 插到最前;ephemeral nudge 保留原。
849        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"); // only c1 answered
886        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        // A compaction that leaves only System/Custom messages would send an
983        // empty `messages` array (System maps to the `system` param) — the
984        // contract rejects it.
985        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    // ── RunState tests ────────────────────────────────────────────────────
1030
1031    #[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        // original_thinking_enabled is NOT reset
1064        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); // reset
1080        assert_eq!(rs.empty_response_strikes, 0); // reset
1081    }
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); // reset
1105    }
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); // reset
1128    }
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        // Simulate: reasoning only → tool calls → reasoning only
1142        let mut rs = RunState::default();
1143
1144        // Branch 1: reasoning only
1145        rs.record_reasoning_only();
1146        assert_eq!(rs.reasoning_only_strikes, 1);
1147
1148        // Branch 3: tool calls (should reset reasoning_only_strikes)
1149        rs.record_tool_calls(2);
1150        assert_eq!(rs.reasoning_only_strikes, 0);
1151        assert_eq!(rs.turn_tool_calls, 2);
1152
1153        // Branch 1 again: reasoning only (should start from 1, not 2)
1154        let strikes = rs.record_reasoning_only();
1155        assert_eq!(strikes, 1);
1156    }
1157
1158    #[test]
1159    fn run_state_empty_to_reasoning_reset() {
1160        // Simulate: empty → empty → reasoning only (should reset empty strikes)
1161        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        // Branch 1: reasoning only (should reset empty_response_strikes)
1168        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        // After 1st reasoning-only: not disabled
1184        rs.record_reasoning_only();
1185        assert!(!rs.thinking_disabled_for_rest_of_run);
1186        assert_eq!(rs.reasoning_only_strikes, 1);
1187
1188        // After 2nd reasoning-only: not disabled
1189        rs.record_reasoning_only();
1190        assert!(!rs.thinking_disabled_for_rest_of_run);
1191        assert_eq!(rs.reasoning_only_strikes, 2);
1192
1193        // After 3rd reasoning-only: disabled!
1194        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        // Simulate 3 reasoning-only responses
1204        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        // Reset for new run (new user message)
1211        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        // Simulate 3 reasoning-only responses
1221        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        // Tool calls should NOT reset thinking_disabled_for_rest_of_run
1227        // (it should stay disabled for the rest of the run)
1228        rs.record_tool_calls(2);
1229        assert!(rs.thinking_disabled_for_rest_of_run);
1230        assert_eq!(rs.reasoning_only_strikes, 0); // strikes reset, but thinking stays disabled
1231    }
1232
1233    // ── Backward-compatible deserialization ────────────────────────────────
1234
1235    #[test]
1236    fn deserialize_legacy_flat_fields() {
1237        // Old format: nudge_count, turn_tool_calls, etc. as flat fields
1238        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); // default
1254    }
1255
1256    #[test]
1257    fn deserialize_new_run_state_format() {
1258        // New format: nested run_state
1259        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        // When both are present, run_state wins
1282        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); // run_state wins
1299        assert_eq!(session.run_state.nudge_count, 0); // run_state wins
1300    }
1301
1302    #[test]
1303    fn deserialize_legacy_missing_optional_fields() {
1304        // Old format with some fields missing (defaults to 0)
1305        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); // missing → 0
1315        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        // Valid JSON args should pass through unchanged
1340        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        // Truncated (invalid JSON) args must be sanitized to an empty object,
1357        // NOT wrapped in a descriptive error object. Two reasons:
1358        //   1. "{}" is valid JSON, so the next request is never rejected 400.
1359        //   2. a rich `{error, original_args_preview, message}` object in
1360        //      assistant history is an imitation vector — the model replays it
1361        //      verbatim as the next call's arguments, and because it parses as
1362        //      valid JSON it slips past the react truncation guard and resurfaces
1363        //      as a downstream ToolArgsInvalid. The truncation explanation
1364        //      belongs in the tool_result, not the assistant arguments.
1365        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        // Bug-2: Invalid JSON args with multi-byte UTF-8 chars near byte 200
1389        // cause a panic at char boundary when slicing &args[..args.len().min(200)].
1390        let mut s = make_session();
1391
1392        // Build invalid JSON with CJK chars that straddle the 200-byte boundary.
1393        // "あ" = 3 bytes in UTF-8. Repeating ~70 times = ~210 bytes, then add invalid suffix.
1394        let mut bad_args = "あ".repeat(70); // 70 * 3 = 210 bytes
1395        bad_args.push_str("truncated"); // makes it invalid JSON
1396
1397        // This must NOT panic. Invalid args are sanitized to "{}" (a fixed
1398        // literal), so no slicing of the multibyte string happens at all.
1399        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        // Regression test for: when LLM response is truncated (finish_reason=max_tokens),
1415        // the code must push assistant message WITH tool_use blocks (not plain text)
1416        // so that subsequent tool_result messages can match the tool_use_id.
1417        // This is required by Anthropic protocol.
1418        let mut s = make_session();
1419
1420        // Simulate truncated tool call response
1421        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        // Push assistant message with tool_calls (as the fix does)
1435        s.push_assistant_tool_calls(
1436            &tool_calls,
1437            Some("thinking...".to_string()),
1438            Some("I'll help you".to_string()),
1439        );
1440
1441        // Push tool results for each tool call
1442        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        // Verify the message sequence is valid
1450        // 1. Assistant message should have tool_calls
1451        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        // 2. Tool result messages should have matching tool_use_id
1466        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        // 3. Validate message sequence (this would catch the bug)
1479        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        // ── RunState property tests ─────────────────────────────────────────
1493
1494        #[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            // original_thinking_enabled is NOT reset
1521            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        // ── push_assistant_tool_calls property tests ────────────────────────
1555
1556        #[test]
1557        fn push_assistant_tool_calls_valid_json_unchanged(args in r"\{[^{}]{0,200}\}") {
1558            // Only test strings that are actually valid JSON objects
1559            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            // Skip if it happens to be valid JSON
1580            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                // Must be valid JSON (no 400 on replay) AND carry no poison
1591                // wrapper the model could echo back as arguments.
1592                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        // ── trim_oldest_turns property tests ────────────────────────────────
1601
1602        #[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        // ── validate_message_sequence property tests ────────────────────────
1614
1615        #[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}