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    /// 转义并截取参数前 `max_chars` 个字符,供截断 WARN 直接展示原始内容。
322    /// Rust debug 转义让未闭合 JSON、控制字符与 CJK 边界清晰可辨,线上
323    /// 排查 provider 截断时不必再间接还原线内容。
324    fn preview_args(args: &str, max_chars: usize) -> String {
325        format!("{:?}", args.chars().take(max_chars).collect::<String>())
326    }
327
328    pub fn push_assistant_tool_calls(
329        &mut self,
330        tool_calls: &[(String, String, String)],
331        reasoning: Option<String>,
332        content: Option<String>,
333    ) {
334        let calls: Vec<ToolCallMessage> = tool_calls
335            .iter()
336            .map(|(id, name, args)| {
337                // Invalid JSON args (provider truncated the stream mid-generation)
338                // are sanitized to an empty object. We deliberately do NOT wrap
339                // them in a descriptive `{error, original_args_preview, message}`
340                // object: "{}" is valid JSON so the next request is never rejected
341                // with 400, and — critically — a rich error object here becomes an
342                // imitation vector: the model replays it verbatim as the next
343                // call's arguments and, because it parses, it slips past the react
344                // truncation guard. The truncation explanation lives in the
345                // paired tool_result, which is where the model reads feedback.
346                let valid_args = if serde_json::from_str::<serde_json::Value>(args).is_ok() {
347                    args.clone()
348                } else {
349                    tracing::warn!(
350                        tool_name = %name,
351                        args_len = args.len(),
352                        args_preview = %Self::preview_args(args, 80),
353                        "tool call arguments are not valid JSON (provider truncated them), \
354                         sanitizing to empty object; re-issue instruction goes to the tool_result"
355                    );
356                    "{}".to_string()
357                };
358                ToolCallMessage {
359                    id: id.clone(),
360                    name: name.clone(),
361                    arguments: valid_args,
362                }
363            })
364            .collect();
365        self.chat_messages.push(ChatMessage::Assistant {
366            content,
367            reasoning_content: reasoning,
368            tool_calls: Some(calls),
369            thinking_signature: None,
370        });
371    }
372
373    pub fn push_tool_result(&mut self, tool_call_id: &str, content: impl Into<String>) {
374        self.chat_messages
375            .push(ChatMessage::tool(tool_call_id, content));
376    }
377
378    /// 移除所有临时消息(ephemeral=true)。
379    ///
380    /// 在 turn 结束时调用,确保注入的临时内容不残留到下一轮。
381    pub fn remove_ephemeral_messages(&mut self) {
382        let before = self.chat_messages.len();
383        self.chat_messages.retain(|m| !m.is_ephemeral());
384        let removed = before - self.chat_messages.len();
385        if removed > 0 {
386            tracing::debug!(
387                removed,
388                remaining = self.chat_messages.len(),
389                "ephemeral messages cleaned up"
390            );
391        }
392    }
393
394    /// Count the number of conversation turns.
395    /// A turn starts with a User message and includes subsequent Assistant/Tool messages.
396    pub fn turn_count(&self) -> usize {
397        self.chat_messages
398            .iter()
399            .filter(|m| matches!(m, ChatMessage::User { .. }))
400            .count()
401    }
402
403    /// Remove the oldest turns from the front until turn count ≤ max_turns.
404    /// Preserves the System message at index 0 if present.
405    pub fn trim_oldest_turns(&mut self, max_turns: usize) {
406        let current_turns = self.turn_count();
407        if current_turns <= max_turns {
408            return;
409        }
410        let turns_to_remove = current_turns - max_turns;
411
412        // Find User message positions (turn boundaries) in chat_messages
413        let user_positions: Vec<usize> = self
414            .chat_messages
415            .iter()
416            .enumerate()
417            .filter_map(|(i, m)| {
418                if matches!(m, ChatMessage::User { .. }) {
419                    Some(i)
420                } else {
421                    None
422                }
423            })
424            .collect();
425
426        if user_positions.len() <= turns_to_remove {
427            return;
428        }
429
430        // Preserve system prefix: count leading System messages
431        let system_prefix = self
432            .chat_messages
433            .iter()
434            .take_while(|m| matches!(m, ChatMessage::System { .. }))
435            .count();
436
437        // Drain from system_prefix up to the start of the (turns_to_remove + 1)-th turn
438        let drain_end = user_positions[turns_to_remove];
439        if system_prefix >= drain_end {
440            return; // nothing to drain after system messages
441        }
442
443        self.chat_messages.drain(system_prefix..drain_end);
444    }
445
446    /// Remove the last message from `chat_messages`.
447    /// Used by the max_message_tokens safety valve to discard oversized messages.
448    pub fn pop_last_message(&mut self) {
449        self.chat_messages.pop();
450    }
451
452    pub fn close_dangling_tool_calls(&mut self, error_summary: &str) {
453        let assistant_idx = self.chat_messages.iter().rposition(
454            |m| matches!(m, ChatMessage::Assistant { tool_calls: Some(tc), .. } if !tc.is_empty()),
455        );
456
457        let Some(assistant_idx) = assistant_idx else {
458            return;
459        };
460
461        let ChatMessage::Assistant {
462            tool_calls: Some(tc),
463            ..
464        } = &self.chat_messages[assistant_idx]
465        else {
466            return;
467        };
468
469        let all_ids: Vec<String> = tc.iter().map(|t| t.id.clone()).collect();
470
471        let answered_ids: Vec<String> = self.chat_messages[assistant_idx + 1..]
472            .iter()
473            .filter_map(|m| match m {
474                ChatMessage::Tool { tool_call_id, .. } => Some(tool_call_id.clone()),
475                _ => None,
476            })
477            .collect();
478
479        for id in &all_ids {
480            if !answered_ids.iter().any(|a| a == id) {
481                self.push_tool_result(id, error_summary);
482            }
483        }
484    }
485
486    /// Replace chat messages — only for persistence restore.
487    /// Validates message sequence before replacing.
488    ///
489    /// 仅供持久化恢复使用。调用方必须保证 messages 序列合法。
490    pub fn set_chat_messages(&mut self, messages: Vec<ChatMessage>) -> Result<(), String> {
491        validate_message_sequence(&messages)?;
492        // Recalculate total_tool_calls from the incoming messages so middleware
493        // decisions (e.g. first_turn_only enforcement) see the correct count.
494        self.total_tool_calls = messages
495            .iter()
496            .filter_map(|m| match m {
497                ChatMessage::Assistant {
498                    tool_calls: Some(tc),
499                    ..
500                } => Some(tc.len()),
501                _ => None,
502            })
503            .sum();
504        self.chat_messages = messages;
505        Ok(())
506    }
507}
508
509/// Validate that a chat message sequence is well-formed for LLM API consumption.
510///
511/// Checks:
512/// - At least one non-System/non-Custom message (System maps to the
513///   top-level `system` parameter, Custom is stripped — a sequence of only
514///   those leaves an empty `messages` array, which providers reject with
515///   HTTP 400). Guards compaction outputs against producing an unsendable
516///   window.
517/// - No Tool message without a preceding Assistant with matching tool_call
518/// - No duplicate Tool messages for the same tool_call_id
519/// - All tool_calls in an Assistant batch must be answered before the next Assistant batch
520/// - No unanswered tool calls at the end of the sequence
521pub fn validate_message_sequence(messages: &[ChatMessage]) -> Result<(), String> {
522    if !messages
523        .iter()
524        .any(|m| !matches!(m, ChatMessage::System { .. } | ChatMessage::Custom { .. }))
525    {
526        return Err(
527            "sequence contains no sendable message: System/Custom alone leave the \
528             provider `messages` array empty"
529                .to_string(),
530        );
531    }
532
533    let mut pending_tool_call_ids: HashSet<String> = HashSet::new();
534
535    for (i, msg) in messages.iter().enumerate() {
536        match msg {
537            ChatMessage::Tool { tool_call_id, .. } => {
538                if pending_tool_call_ids.is_empty() {
539                    return Err(format!(
540                        "message[{}]: Tool message with call_id '{}' has no preceding tool_call",
541                        i, tool_call_id
542                    ));
543                }
544                // Remove the ID on match — also detects duplicates (second remove returns false)
545                if !pending_tool_call_ids.remove(tool_call_id) {
546                    return Err(format!(
547                        "message[{}]: Tool message with call_id '{}' does not match any pending tool_call (already answered or unknown)",
548                        i, tool_call_id
549                    ));
550                }
551            }
552            ChatMessage::Assistant {
553                tool_calls: Some(tc),
554                ..
555            } => {
556                // Previous batch must be fully answered before a new batch starts
557                if !pending_tool_call_ids.is_empty() {
558                    return Err(format!(
559                        "message[{}]: Assistant message with new tool_calls appears before pending calls were answered: {:?}",
560                        i, pending_tool_call_ids
561                    ));
562                }
563                pending_tool_call_ids = tc.iter().map(|t| t.id.clone()).collect();
564            }
565            _ => {}
566        }
567    }
568
569    // All tool calls must be answered by the end of the sequence
570    if !pending_tool_call_ids.is_empty() {
571        return Err(format!(
572            "message sequence ends with unanswered tool calls: {:?}",
573            pending_tool_call_ids
574        ));
575    }
576
577    Ok(())
578}
579
580#[cfg(test)]
581fn make_session() -> AgentSession {
582    AgentSession::new(SessionId::new(1))
583}
584
585#[cfg(test)]
586mod tests {
587    use super::*;
588
589    #[test]
590    fn test_turn_count_empty() {
591        let s = make_session();
592        assert_eq!(s.turn_count(), 0);
593    }
594
595    #[test]
596    fn test_turn_count_with_system_and_user() {
597        let mut s = make_session();
598        s.push_message(MessageRole::System, "system");
599        assert_eq!(s.turn_count(), 0);
600        s.push_message(MessageRole::User, "hello");
601        assert_eq!(s.turn_count(), 1);
602        s.push_message(MessageRole::Assistant, "hi");
603        assert_eq!(s.turn_count(), 1);
604        s.push_message(MessageRole::User, "bye");
605        assert_eq!(s.turn_count(), 2);
606    }
607
608    #[test]
609    fn test_turn_count_with_tool_calls() {
610        let mut s = make_session();
611        s.push_message(MessageRole::User, "do something");
612        s.push_assistant_tool_calls(&[("id1".into(), "tool".into(), "{}".into())], None, None);
613        s.push_tool_result("id1", "result");
614        s.push_message(MessageRole::Assistant, "done");
615        // One user turn: User -> Assistant(tool_calls) -> Tool -> Assistant(text)
616        assert_eq!(s.turn_count(), 1);
617    }
618
619    #[test]
620    fn test_trim_oldest_turns_noop() {
621        let mut s = make_session();
622        s.push_message(MessageRole::User, "hello");
623        s.push_message(MessageRole::Assistant, "hi");
624        s.trim_oldest_turns(5);
625        assert_eq!(s.turn_count(), 1);
626        assert_eq!(s.chat_messages().len(), 2);
627    }
628
629    #[test]
630    fn test_trim_oldest_turns_removes_old() {
631        let mut s = make_session();
632        s.push_message(MessageRole::System, "sys");
633        // Turn 1
634        s.push_message(MessageRole::User, "u1");
635        s.push_message(MessageRole::Assistant, "a1");
636        // Turn 2
637        s.push_message(MessageRole::User, "u2");
638        s.push_message(MessageRole::Assistant, "a2");
639        // Turn 3
640        s.push_message(MessageRole::User, "u3");
641        s.push_message(MessageRole::Assistant, "a3");
642
643        s.trim_oldest_turns(2);
644        assert_eq!(s.turn_count(), 2);
645        // System message preserved
646        assert!(matches!(s.chat_messages()[0], ChatMessage::System { .. }));
647        // Oldest user message is u2
648        assert!(
649            matches!(s.chat_messages()[1], ChatMessage::User { ref content, .. } if content == "u2")
650        );
651    }
652
653    #[test]
654    fn test_trim_oldest_turns_with_tool_calls() {
655        let mut s = make_session();
656        // Turn 1 with tool call
657        s.push_message(MessageRole::User, "u1");
658        s.push_assistant_tool_calls(&[("id1".into(), "t".into(), "{}".into())], None, None);
659        s.push_tool_result("id1", "r1");
660        s.push_message(MessageRole::Assistant, "a1");
661        // Turn 2
662        s.push_message(MessageRole::User, "u2");
663        s.push_message(MessageRole::Assistant, "a2");
664
665        let msg_before = s.simple_messages().len();
666        let chat_before = s.chat_messages().len();
667        s.trim_oldest_turns(1);
668        assert_eq!(s.turn_count(), 1);
669        // chat_messages should have lost 4 entries (User, Assistant(tool), Tool, Assistant(text))
670        assert_eq!(s.chat_messages().len(), chat_before - 4);
671        // simple_messages (derived from chat_messages, tool_calls-only filtered) loses 3 entries
672        assert_eq!(s.simple_messages().len(), msg_before - 3);
673    }
674
675    #[test]
676    fn test_pop_last_message_text() {
677        let mut s = make_session();
678        s.push_message(MessageRole::User, "hello");
679        s.push_message(MessageRole::Assistant, "hi");
680        assert_eq!(s.chat_messages().len(), 2);
681        s.pop_last_message();
682        assert_eq!(s.chat_messages().len(), 1);
683        assert_eq!(s.simple_messages().len(), 1);
684    }
685
686    #[test]
687    fn test_pop_last_message_tool_calls_only() {
688        let mut s = make_session();
689        s.push_message(MessageRole::User, "do it");
690        s.push_assistant_tool_calls(&[("id1".into(), "t".into(), "{}".into())], None, None);
691        assert_eq!(s.chat_messages().len(), 2);
692        assert_eq!(s.simple_messages().len(), 1); // only User in simple_messages (tool_calls-only filtered)
693        s.pop_last_message();
694        assert_eq!(s.chat_messages().len(), 1);
695        assert_eq!(s.simple_messages().len(), 1); // simple_messages unchanged (still just User)
696    }
697
698    #[test]
699    fn test_pop_last_message_empty_session() {
700        let mut s = make_session();
701        s.pop_last_message(); // should not panic
702        assert_eq!(s.chat_messages().len(), 0);
703    }
704
705    // ── B5: remaining session lifecycle paths ──────────────────────────────
706
707    #[test]
708    fn test_id_and_action_allowlist() {
709        let mut s = make_session();
710        assert_eq!(s.id(), Some(SessionId::new(1)));
711        assert!(!s.is_action_allowed("approve:rm"));
712        s.allow_action("approve:rm");
713        assert!(s.is_action_allowed("approve:rm"));
714        assert!(!s.is_action_allowed("approve:shell"));
715    }
716
717    #[test]
718    fn test_chat_messages_mut() {
719        let mut s = make_session();
720        s.chat_messages_mut().push(ChatMessage::user("direct"));
721        assert_eq!(s.chat_messages().len(), 1);
722    }
723
724    #[test]
725    fn test_push_message_tool_role() {
726        let mut s = make_session();
727        s.push_message(MessageRole::Tool, "result");
728        assert!(matches!(s.chat_messages()[0], ChatMessage::Tool { .. }));
729    }
730
731    #[test]
732    fn test_push_assistant_with_reasoning() {
733        let mut s = make_session();
734        s.push_assistant_with_reasoning("answer", "thinking");
735        match &s.chat_messages()[0] {
736            ChatMessage::Assistant {
737                content,
738                reasoning_content,
739                ..
740            } => {
741                assert_eq!(content.as_deref(), Some("answer"));
742                assert_eq!(reasoning_content.as_deref(), Some("thinking"));
743            }
744            other => panic!("unexpected message: {other:?}"),
745        }
746    }
747
748    #[test]
749    fn test_push_user_message_with_images() {
750        let mut s = make_session();
751        s.push_user_message_with_images(
752            "look",
753            vec![ImageAttachment::Url {
754                url: "http://x".into(),
755                detail: None,
756            }],
757        );
758        match &s.chat_messages()[0] {
759            ChatMessage::User { images, .. } => assert_eq!(images.len(), 1),
760            other => panic!("unexpected message: {other:?}"),
761        }
762    }
763
764    #[test]
765    fn test_push_assistant_tool_call_singular() {
766        let mut s = make_session();
767        s.push_assistant_tool_call("call_1", "bash", "{}");
768        match &s.chat_messages()[0] {
769            ChatMessage::Assistant {
770                tool_calls: Some(tc),
771                ..
772            } => {
773                assert_eq!(tc.len(), 1);
774                assert_eq!(tc[0].id, "call_1");
775                assert_eq!(tc[0].name, "bash");
776            }
777            other => panic!("unexpected message: {other:?}"),
778        }
779    }
780
781    #[test]
782    fn test_simple_messages_filters_empty_content_tool_calls() {
783        let mut s = make_session();
784        s.chat_messages_mut().push(ChatMessage::Assistant {
785            content: Some(String::new()),
786            reasoning_content: None,
787            tool_calls: Some(vec![ToolCallMessage {
788                id: "c".into(),
789                name: "t".into(),
790                arguments: "{}".into(),
791            }]),
792            thinking_signature: None,
793        });
794        assert!(s.simple_messages().is_empty());
795    }
796
797    #[test]
798    fn test_remove_ephemeral_messages() {
799        let mut s = make_session();
800        s.push_message(MessageRole::System, "keep");
801        s.chat_messages_mut()
802            .push(ChatMessage::user_ephemeral("temp"));
803        s.chat_messages_mut()
804            .push(ChatMessage::system_ephemeral("temp2"));
805        s.push_message(MessageRole::User, "keep2");
806        assert_eq!(s.chat_messages().len(), 4);
807        s.remove_ephemeral_messages();
808        assert_eq!(s.chat_messages().len(), 2);
809        assert!(s.chat_messages().iter().all(|m| !m.is_ephemeral()));
810    }
811
812    #[test]
813    fn test_set_system_prompt_replaces_first_non_ephemeral_system() {
814        let mut s = make_session();
815        s.push_message(MessageRole::System, "old prompt");
816        s.push_message(MessageRole::User, "hi");
817        s.push_message(MessageRole::Assistant, "hello");
818
819        s.set_system_prompt("new prompt");
820
821        let msgs = s.chat_messages();
822        assert_eq!(msgs.len(), 3, "history length unchanged");
823        assert!(
824            matches!(&msgs[0], ChatMessage::System { content, ephemeral: false } if content == "new prompt")
825        );
826        assert!(matches!(&msgs[1], ChatMessage::User { content, .. } if content == "hi"));
827        assert!(
828            matches!(&msgs[2], ChatMessage::Assistant { content: Some(c), .. } if c == "hello")
829        );
830    }
831
832    #[test]
833    fn test_set_system_prompt_inserts_when_absent() {
834        let mut s = make_session();
835        s.push_message(MessageRole::User, "hi");
836
837        s.set_system_prompt("fresh prompt");
838
839        let msgs = s.chat_messages();
840        assert_eq!(msgs.len(), 2);
841        assert!(
842            matches!(&msgs[0], ChatMessage::System { content, ephemeral: false } if content == "fresh prompt")
843        );
844        assert!(matches!(&msgs[1], ChatMessage::User { .. }));
845    }
846
847    #[test]
848    fn test_set_system_prompt_skips_ephemeral_system_and_inserts() {
849        let mut s = make_session();
850        s.chat_messages_mut()
851            .push(ChatMessage::system_ephemeral("ephemeral nudge"));
852        s.push_message(MessageRole::User, "hi");
853
854        s.set_system_prompt("real prompt");
855
856        // 非临时 System 不存在 → 插到最前;ephemeral nudge 保留原。
857        let msgs = s.chat_messages();
858        assert_eq!(msgs.len(), 3);
859        assert!(
860            matches!(&msgs[0], ChatMessage::System { content, ephemeral: false } if content == "real prompt")
861        );
862        assert!(matches!(
863            &msgs[1],
864            ChatMessage::System {
865                ephemeral: true,
866                ..
867            }
868        ));
869        assert!(matches!(&msgs[2], ChatMessage::User { .. }));
870    }
871
872    #[test]
873    fn test_close_dangling_tool_calls_noop_without_tool_call() {
874        let mut s = make_session();
875        s.push_message(MessageRole::User, "hi");
876        s.push_message(MessageRole::Assistant, "hi");
877        s.close_dangling_tool_calls("failed");
878        assert_eq!(s.chat_messages().len(), 2);
879    }
880
881    #[test]
882    fn test_close_dangling_tool_calls_adds_missing_results() {
883        let mut s = make_session();
884        s.push_message(MessageRole::User, "do");
885        s.push_assistant_tool_calls(
886            &[
887                ("c1".into(), "t".into(), "{}".into()),
888                ("c2".into(), "t".into(), "{}".into()),
889            ],
890            None,
891            None,
892        );
893        s.push_tool_result("c1", "ok"); // only c1 answered
894        s.close_dangling_tool_calls("failed");
895
896        let tool_results: Vec<(String, String)> = s
897            .chat_messages()
898            .iter()
899            .filter_map(|m| match m {
900                ChatMessage::Tool {
901                    tool_call_id,
902                    name: _,
903                    content,
904                } => Some((tool_call_id.clone(), content.clone())),
905                _ => None,
906            })
907            .collect();
908        assert_eq!(tool_results.len(), 2);
909        assert!(
910            tool_results
911                .iter()
912                .any(|(id, c)| id == "c2" && c == "failed")
913        );
914    }
915
916    #[test]
917    fn test_set_chat_messages_recalculates_total_tool_calls() {
918        let mut s = make_session();
919        let msgs = vec![
920            ChatMessage::user("do"),
921            ChatMessage::assistant_tool_call("c1", "t", "{}"),
922            ChatMessage::tool("c1", "result"),
923        ];
924        s.set_chat_messages(msgs).unwrap();
925        assert_eq!(s.total_tool_calls, 1);
926    }
927}
928
929#[cfg(test)]
930mod validate_tests {
931    use super::*;
932
933    #[test]
934    fn test_valid_simple_sequence() {
935        let msgs = vec![ChatMessage::user("hello"), ChatMessage::assistant("hi")];
936        assert!(validate_message_sequence(&msgs).is_ok());
937    }
938
939    #[test]
940    fn test_valid_tool_call_sequence() {
941        let msgs = vec![
942            ChatMessage::user("run command"),
943            ChatMessage::assistant_tool_call("call_1", "bash", r#"{"cmd":"ls"}"#),
944            ChatMessage::tool("call_1", "file1 file2"),
945            ChatMessage::assistant("done"),
946        ];
947        assert!(validate_message_sequence(&msgs).is_ok());
948    }
949
950    #[test]
951    fn test_valid_multi_tool_call_sequence() {
952        let msgs = vec![
953            ChatMessage::user("run commands"),
954            ChatMessage::Assistant {
955                content: None,
956                reasoning_content: None,
957                tool_calls: Some(vec![
958                    crate::types::ToolCallMessage {
959                        id: "call_1".into(),
960                        name: "bash".into(),
961                        arguments: "{}".into(),
962                    },
963                    crate::types::ToolCallMessage {
964                        id: "call_2".into(),
965                        name: "read".into(),
966                        arguments: "{}".into(),
967                    },
968                ]),
969                thinking_signature: None,
970            },
971            ChatMessage::tool("call_1", "result1"),
972            ChatMessage::tool("call_2", "result2"),
973            ChatMessage::assistant("done"),
974        ];
975        assert!(validate_message_sequence(&msgs).is_ok());
976    }
977
978    #[test]
979    fn test_orphaned_tool_result() {
980        let msgs = vec![
981            ChatMessage::user("hello"),
982            ChatMessage::tool("call_1", "orphaned result"),
983        ];
984        let err = validate_message_sequence(&msgs).unwrap_err();
985        assert!(err.contains("no preceding tool_call"));
986    }
987
988    #[test]
989    fn test_system_only_sequence_rejected() {
990        // A compaction that leaves only System/Custom messages would send an
991        // empty `messages` array (System maps to the `system` param) — the
992        // contract rejects it.
993        let msgs = vec![
994            ChatMessage::system("prompt"),
995            ChatMessage::system_ephemeral("reminder"),
996            ChatMessage::Custom {
997                role: "artifact".into(),
998                data: serde_json::json!({"id": "x"}),
999            },
1000        ];
1001        let err = validate_message_sequence(&msgs).unwrap_err();
1002        assert!(err.contains("no sendable message"));
1003    }
1004
1005    #[test]
1006    fn test_system_plus_user_ok() {
1007        let msgs = vec![ChatMessage::system("prompt"), ChatMessage::user("hi")];
1008        assert!(validate_message_sequence(&msgs).is_ok());
1009    }
1010
1011    #[test]
1012    fn test_mismatched_tool_call_id() {
1013        let msgs = vec![
1014            ChatMessage::user("run"),
1015            ChatMessage::assistant_tool_call("call_1", "bash", "{}"),
1016            ChatMessage::tool("call_2", "wrong id"),
1017        ];
1018        let err = validate_message_sequence(&msgs).unwrap_err();
1019        assert!(err.contains("does not match"));
1020    }
1021
1022    #[test]
1023    fn test_set_chat_messages_valid() {
1024        let mut s = make_session();
1025        let msgs = vec![ChatMessage::user("hello"), ChatMessage::assistant("hi")];
1026        assert!(s.set_chat_messages(msgs.clone()).is_ok());
1027        assert_eq!(s.chat_messages().len(), 2);
1028    }
1029
1030    #[test]
1031    fn test_set_chat_messages_invalid() {
1032        let mut s = make_session();
1033        let msgs = vec![ChatMessage::tool("call_1", "orphaned")];
1034        assert!(s.set_chat_messages(msgs).is_err());
1035    }
1036
1037    // ── RunState tests ────────────────────────────────────────────────────
1038
1039    #[test]
1040    fn run_state_default() {
1041        let rs = RunState::default();
1042        assert_eq!(rs.turn_tool_calls, 0);
1043        assert!(!rs.run_has_tool_calls);
1044        assert_eq!(rs.reasoning_only_strikes, 0);
1045        assert_eq!(rs.empty_response_strikes, 0);
1046        assert_eq!(rs.nudge_count, 0);
1047    }
1048
1049    #[test]
1050    fn run_state_reset_for_new_run() {
1051        let mut rs = RunState {
1052            turn_tool_calls: 5,
1053            run_has_tool_calls: true,
1054            reasoning_only_strikes: 2,
1055            empty_response_strikes: 1,
1056            nudge_count: 3,
1057            thinking_disabled_for_rest_of_run: true,
1058            original_thinking_enabled: true,
1059            truncation_strikes: 4,
1060        };
1061
1062        rs.reset_for_new_run();
1063
1064        assert_eq!(rs.turn_tool_calls, 0);
1065        assert!(!rs.run_has_tool_calls);
1066        assert_eq!(rs.reasoning_only_strikes, 0);
1067        assert_eq!(rs.empty_response_strikes, 0);
1068        assert_eq!(rs.nudge_count, 0);
1069        assert_eq!(rs.truncation_strikes, 0);
1070        assert!(!rs.thinking_disabled_for_rest_of_run);
1071        // original_thinking_enabled is NOT reset
1072        assert!(rs.original_thinking_enabled);
1073    }
1074
1075    #[test]
1076    fn run_state_record_tool_calls() {
1077        let mut rs = RunState {
1078            reasoning_only_strikes: 2,
1079            empty_response_strikes: 1,
1080            ..RunState::default()
1081        };
1082
1083        rs.record_tool_calls(3);
1084
1085        assert_eq!(rs.turn_tool_calls, 3);
1086        assert!(rs.run_has_tool_calls);
1087        assert_eq!(rs.reasoning_only_strikes, 0); // reset
1088        assert_eq!(rs.empty_response_strikes, 0); // reset
1089    }
1090
1091    #[test]
1092    fn run_state_record_tool_calls_accumulates() {
1093        let mut rs = RunState::default();
1094        rs.record_tool_calls(2);
1095        rs.record_tool_calls(3);
1096
1097        assert_eq!(rs.turn_tool_calls, 5);
1098        assert!(rs.run_has_tool_calls);
1099    }
1100
1101    #[test]
1102    fn run_state_record_reasoning_only() {
1103        let mut rs = RunState {
1104            empty_response_strikes: 2,
1105            ..RunState::default()
1106        };
1107
1108        let strikes = rs.record_reasoning_only();
1109
1110        assert_eq!(strikes, 1);
1111        assert_eq!(rs.reasoning_only_strikes, 1);
1112        assert_eq!(rs.empty_response_strikes, 0); // reset
1113    }
1114
1115    #[test]
1116    fn run_state_record_reasoning_only_consecutive() {
1117        let mut rs = RunState::default();
1118
1119        assert_eq!(rs.record_reasoning_only(), 1);
1120        assert_eq!(rs.record_reasoning_only(), 2);
1121        assert_eq!(rs.record_reasoning_only(), 3);
1122    }
1123
1124    #[test]
1125    fn run_state_record_empty_response() {
1126        let mut rs = RunState {
1127            reasoning_only_strikes: 2,
1128            ..RunState::default()
1129        };
1130
1131        let strikes = rs.record_empty_response();
1132
1133        assert_eq!(strikes, 1);
1134        assert_eq!(rs.empty_response_strikes, 1);
1135        assert_eq!(rs.reasoning_only_strikes, 0); // reset
1136    }
1137
1138    #[test]
1139    fn run_state_record_empty_response_consecutive() {
1140        let mut rs = RunState::default();
1141
1142        assert_eq!(rs.record_empty_response(), 1);
1143        assert_eq!(rs.record_empty_response(), 2);
1144        assert_eq!(rs.record_empty_response(), 3);
1145    }
1146
1147    #[test]
1148    fn run_state_branch_cross_reset() {
1149        // Simulate: reasoning only → tool calls → reasoning only
1150        let mut rs = RunState::default();
1151
1152        // Branch 1: reasoning only
1153        rs.record_reasoning_only();
1154        assert_eq!(rs.reasoning_only_strikes, 1);
1155
1156        // Branch 3: tool calls (should reset reasoning_only_strikes)
1157        rs.record_tool_calls(2);
1158        assert_eq!(rs.reasoning_only_strikes, 0);
1159        assert_eq!(rs.turn_tool_calls, 2);
1160
1161        // Branch 1 again: reasoning only (should start from 1, not 2)
1162        let strikes = rs.record_reasoning_only();
1163        assert_eq!(strikes, 1);
1164    }
1165
1166    #[test]
1167    fn run_state_empty_to_reasoning_reset() {
1168        // Simulate: empty → empty → reasoning only (should reset empty strikes)
1169        let mut rs = RunState::default();
1170
1171        rs.record_empty_response();
1172        rs.record_empty_response();
1173        assert_eq!(rs.empty_response_strikes, 2);
1174
1175        // Branch 1: reasoning only (should reset empty_response_strikes)
1176        rs.record_reasoning_only();
1177        assert_eq!(rs.empty_response_strikes, 0);
1178        assert_eq!(rs.reasoning_only_strikes, 1);
1179    }
1180
1181    #[test]
1182    fn run_state_thinking_disabled_default() {
1183        let rs = RunState::default();
1184        assert!(!rs.thinking_disabled_for_rest_of_run);
1185    }
1186
1187    #[test]
1188    fn run_state_thinking_disabled_after_3_strikes() {
1189        let mut rs = RunState::default();
1190
1191        // After 1st reasoning-only: not disabled
1192        rs.record_reasoning_only();
1193        assert!(!rs.thinking_disabled_for_rest_of_run);
1194        assert_eq!(rs.reasoning_only_strikes, 1);
1195
1196        // After 2nd reasoning-only: not disabled
1197        rs.record_reasoning_only();
1198        assert!(!rs.thinking_disabled_for_rest_of_run);
1199        assert_eq!(rs.reasoning_only_strikes, 2);
1200
1201        // After 3rd reasoning-only: disabled!
1202        rs.record_reasoning_only();
1203        assert!(rs.thinking_disabled_for_rest_of_run);
1204        assert_eq!(rs.reasoning_only_strikes, 3);
1205    }
1206
1207    #[test]
1208    fn run_state_thinking_disabled_resets_on_new_run() {
1209        let mut rs = RunState::default();
1210
1211        // Simulate 3 reasoning-only responses
1212        rs.record_reasoning_only();
1213        rs.record_reasoning_only();
1214        rs.record_reasoning_only();
1215        assert!(rs.thinking_disabled_for_rest_of_run);
1216        assert_eq!(rs.reasoning_only_strikes, 3);
1217
1218        // Reset for new run (new user message)
1219        rs.reset_for_new_run();
1220        assert!(!rs.thinking_disabled_for_rest_of_run);
1221        assert_eq!(rs.reasoning_only_strikes, 0);
1222    }
1223
1224    #[test]
1225    fn run_state_thinking_disabled_stays_after_tool_calls() {
1226        let mut rs = RunState::default();
1227
1228        // Simulate 3 reasoning-only responses
1229        rs.record_reasoning_only();
1230        rs.record_reasoning_only();
1231        rs.record_reasoning_only();
1232        assert!(rs.thinking_disabled_for_rest_of_run);
1233
1234        // Tool calls should NOT reset thinking_disabled_for_rest_of_run
1235        // (it should stay disabled for the rest of the run)
1236        rs.record_tool_calls(2);
1237        assert!(rs.thinking_disabled_for_rest_of_run);
1238        assert_eq!(rs.reasoning_only_strikes, 0); // strikes reset, but thinking stays disabled
1239    }
1240
1241    // ── Backward-compatible deserialization ────────────────────────────────
1242
1243    #[test]
1244    fn deserialize_legacy_flat_fields() {
1245        // Old format: nudge_count, turn_tool_calls, etc. as flat fields
1246        let json = r#"{
1247            "id": null,
1248            "chat_messages": [],
1249            "always_allowed_actions": [],
1250            "total_tool_calls": 5,
1251            "nudge_count": 3,
1252            "turn_tool_calls": 2,
1253            "reasoning_only_strikes": 1,
1254            "empty_response_strikes": 0
1255        }"#;
1256        let session: AgentSession = serde_json::from_str(json).unwrap();
1257        assert_eq!(session.run_state.nudge_count, 3);
1258        assert_eq!(session.run_state.turn_tool_calls, 2);
1259        assert_eq!(session.run_state.reasoning_only_strikes, 1);
1260        assert_eq!(session.run_state.empty_response_strikes, 0);
1261        assert!(!session.run_state.run_has_tool_calls); // default
1262    }
1263
1264    #[test]
1265    fn deserialize_new_run_state_format() {
1266        // New format: nested run_state
1267        let json = r#"{
1268            "id": null,
1269            "chat_messages": [],
1270            "always_allowed_actions": [],
1271            "total_tool_calls": 5,
1272            "run_state": {
1273                "turn_tool_calls": 4,
1274                "run_has_tool_calls": true,
1275                "reasoning_only_strikes": 0,
1276                "empty_response_strikes": 1,
1277                "nudge_count": 2
1278            }
1279        }"#;
1280        let session: AgentSession = serde_json::from_str(json).unwrap();
1281        assert_eq!(session.run_state.turn_tool_calls, 4);
1282        assert!(session.run_state.run_has_tool_calls);
1283        assert_eq!(session.run_state.empty_response_strikes, 1);
1284        assert_eq!(session.run_state.nudge_count, 2);
1285    }
1286
1287    #[test]
1288    fn deserialize_run_state_takes_precedence_over_flat() {
1289        // When both are present, run_state wins
1290        let json = r#"{
1291            "id": null,
1292            "chat_messages": [],
1293            "always_allowed_actions": [],
1294            "total_tool_calls": 0,
1295            "run_state": {
1296                "turn_tool_calls": 10,
1297                "run_has_tool_calls": true,
1298                "reasoning_only_strikes": 0,
1299                "empty_response_strikes": 0,
1300                "nudge_count": 0
1301            },
1302            "nudge_count": 99,
1303            "turn_tool_calls": 99
1304        }"#;
1305        let session: AgentSession = serde_json::from_str(json).unwrap();
1306        assert_eq!(session.run_state.turn_tool_calls, 10); // run_state wins
1307        assert_eq!(session.run_state.nudge_count, 0); // run_state wins
1308    }
1309
1310    #[test]
1311    fn deserialize_legacy_missing_optional_fields() {
1312        // Old format with some fields missing (defaults to 0)
1313        let json = r#"{
1314            "id": null,
1315            "chat_messages": [],
1316            "always_allowed_actions": [],
1317            "total_tool_calls": 0,
1318            "nudge_count": 1
1319        }"#;
1320        let session: AgentSession = serde_json::from_str(json).unwrap();
1321        assert_eq!(session.run_state.nudge_count, 1);
1322        assert_eq!(session.run_state.turn_tool_calls, 0); // missing → 0
1323        assert_eq!(session.run_state.reasoning_only_strikes, 0);
1324        assert_eq!(session.run_state.empty_response_strikes, 0);
1325    }
1326
1327    #[test]
1328    fn roundtrip_preserves_run_state() {
1329        let mut session = AgentSession::new(SessionId::new(1));
1330        session.run_state.nudge_count = 5;
1331        session.run_state.turn_tool_calls = 3;
1332        session.run_state.run_has_tool_calls = true;
1333        session.run_state.reasoning_only_strikes = 2;
1334
1335        let json = serde_json::to_string(&session).unwrap();
1336        let restored: AgentSession = serde_json::from_str(&json).unwrap();
1337        assert_eq!(restored.run_state.nudge_count, 5);
1338        assert_eq!(restored.run_state.turn_tool_calls, 3);
1339        assert!(restored.run_state.run_has_tool_calls);
1340        assert_eq!(restored.run_state.reasoning_only_strikes, 2);
1341    }
1342
1343    #[test]
1344    fn push_assistant_tool_calls_validates_json_args() {
1345        let mut s = make_session();
1346
1347        // Valid JSON args should pass through unchanged
1348        let valid_args = r#"{"path": "src/main.rs", "content": "fn main() {}"}"#;
1349        s.push_assistant_tool_calls(
1350            &[("id1".into(), "write_file".into(), valid_args.into())],
1351            None,
1352            None,
1353        );
1354        if let ChatMessage::Assistant {
1355            tool_calls: Some(ref tc),
1356            ..
1357        } = s.chat_messages[0]
1358        {
1359            assert_eq!(tc[0].arguments, valid_args);
1360        } else {
1361            panic!("expected Assistant message with tool_calls");
1362        }
1363
1364        // Truncated (invalid JSON) args must be sanitized to an empty object,
1365        // NOT wrapped in a descriptive error object. Two reasons:
1366        //   1. "{}" is valid JSON, so the next request is never rejected 400.
1367        //   2. a rich `{error, original_args_preview, message}` object in
1368        //      assistant history is an imitation vector — the model replays it
1369        //      verbatim as the next call's arguments, and because it parses as
1370        //      valid JSON it slips past the react truncation guard and resurfaces
1371        //      as a downstream ToolArgsInvalid. The truncation explanation
1372        //      belongs in the tool_result, not the assistant arguments.
1373        let truncated_args = r#"{"path": "src/ui/markdown.rs", "content": "#;
1374        s.push_assistant_tool_calls(
1375            &[("id2".into(), "write_file".into(), truncated_args.into())],
1376            None,
1377            None,
1378        );
1379        if let ChatMessage::Assistant {
1380            tool_calls: Some(ref tc),
1381            ..
1382        } = s.chat_messages[1]
1383        {
1384            assert_eq!(tc[0].arguments, "{}");
1385            assert!(
1386                !tc[0].arguments.contains("tool_call_arguments_truncated"),
1387                "assistant arguments must not carry the poison wrapper object"
1388            );
1389        } else {
1390            panic!("expected Assistant message with tool_calls");
1391        }
1392    }
1393
1394    #[test]
1395    fn push_assistant_tool_calls_truncated_multibyte_no_panic() {
1396        // Bug-2: Invalid JSON args with multi-byte UTF-8 chars near byte 200
1397        // cause a panic at char boundary when slicing &args[..args.len().min(200)].
1398        let mut s = make_session();
1399
1400        // Build invalid JSON with CJK chars that straddle the 200-byte boundary.
1401        // "あ" = 3 bytes in UTF-8. Repeating ~70 times = ~210 bytes, then add invalid suffix.
1402        let mut bad_args = "あ".repeat(70); // 70 * 3 = 210 bytes
1403        bad_args.push_str("truncated"); // makes it invalid JSON
1404
1405        // This must NOT panic. Invalid args are sanitized to "{}" (a fixed
1406        // literal), so no slicing of the multibyte string happens at all.
1407        s.push_assistant_tool_calls(&[("id1".into(), "tool".into(), bad_args)], None, None);
1408
1409        if let ChatMessage::Assistant {
1410            tool_calls: Some(ref tc),
1411            ..
1412        } = s.chat_messages[0]
1413        {
1414            assert_eq!(tc[0].arguments, "{}");
1415        } else {
1416            panic!("expected Assistant message with tool_calls");
1417        }
1418    }
1419
1420    #[test]
1421    fn push_assistant_tool_calls_then_tool_result_matches_anthropic_protocol() {
1422        // Regression test for: when LLM response is truncated (finish_reason=max_tokens),
1423        // the code must push assistant message WITH tool_use blocks (not plain text)
1424        // so that subsequent tool_result messages can match the tool_use_id.
1425        // This is required by Anthropic protocol.
1426        let mut s = make_session();
1427
1428        // Simulate truncated tool call response
1429        let tool_calls = vec![
1430            (
1431                "call_00_VJtlnKha0ZZ2Yo8t5ysQ8883".to_string(),
1432                "write_file".to_string(),
1433                "{}".to_string(),
1434            ),
1435            (
1436                "call_01_abc123".to_string(),
1437                "bash".to_string(),
1438                r#"{"command": "ls"}"#.to_string(),
1439            ),
1440        ];
1441
1442        // Push assistant message with tool_calls (as the fix does)
1443        s.push_assistant_tool_calls(
1444            &tool_calls,
1445            Some("thinking...".to_string()),
1446            Some("I'll help you".to_string()),
1447        );
1448
1449        // Push tool results for each tool call
1450        for (tc_id, _, _) in &tool_calls {
1451            s.push_tool_result(
1452                tc_id,
1453                "Tool call was not executed: the response hit the output token limit.",
1454            );
1455        }
1456
1457        // Verify the message sequence is valid
1458        // 1. Assistant message should have tool_calls
1459        if let ChatMessage::Assistant {
1460            tool_calls: Some(ref tc),
1461            ..
1462        } = s.chat_messages[0]
1463        {
1464            assert_eq!(tc.len(), 2);
1465            assert_eq!(tc[0].id, "call_00_VJtlnKha0ZZ2Yo8t5ysQ8883");
1466            assert_eq!(tc[0].name, "write_file");
1467            assert_eq!(tc[1].id, "call_01_abc123");
1468            assert_eq!(tc[1].name, "bash");
1469        } else {
1470            panic!("expected Assistant message with tool_calls");
1471        }
1472
1473        // 2. Tool result messages should have matching tool_use_id
1474        if let ChatMessage::Tool { tool_call_id, .. } = &s.chat_messages[1] {
1475            assert_eq!(tool_call_id, "call_00_VJtlnKha0ZZ2Yo8t5ysQ8883");
1476        } else {
1477            panic!("expected Tool message for first tool call");
1478        }
1479
1480        if let ChatMessage::Tool { tool_call_id, .. } = &s.chat_messages[2] {
1481            assert_eq!(tool_call_id, "call_01_abc123");
1482        } else {
1483            panic!("expected Tool message for second tool call");
1484        }
1485
1486        // 3. Validate message sequence (this would catch the bug)
1487        assert!(
1488            validate_message_sequence(&s.chat_messages).is_ok(),
1489            "message sequence should be valid with matching tool_use and tool_result"
1490        );
1491    }
1492}
1493
1494#[cfg(test)]
1495mod proptest_tests {
1496    use super::*;
1497    use proptest::prelude::*;
1498
1499    proptest! {
1500        // ── RunState property tests ─────────────────────────────────────────
1501
1502        #[test]
1503        fn reset_for_new_run_zeros_all_fields(
1504            turn_tool_calls in 0usize..1000,
1505            run_has_tool_calls in proptest::bool::ANY,
1506            reasoning_only_strikes in 0usize..100,
1507            empty_response_strikes in 0usize..100,
1508            nudge_count in 0usize..100,
1509        ) {
1510            let mut rs = RunState {
1511                turn_tool_calls,
1512                run_has_tool_calls,
1513                reasoning_only_strikes,
1514                empty_response_strikes,
1515                nudge_count,
1516                thinking_disabled_for_rest_of_run: true,
1517                original_thinking_enabled: true,
1518                truncation_strikes: 5,
1519            };
1520            rs.reset_for_new_run();
1521            assert_eq!(rs.turn_tool_calls, 0);
1522            assert!(!rs.run_has_tool_calls);
1523            assert_eq!(rs.reasoning_only_strikes, 0);
1524            assert_eq!(rs.empty_response_strikes, 0);
1525            assert_eq!(rs.nudge_count, 0);
1526            assert_eq!(rs.truncation_strikes, 0);
1527            assert!(!rs.thinking_disabled_for_rest_of_run);
1528            // original_thinking_enabled is NOT reset
1529            assert!(rs.original_thinking_enabled);
1530        }
1531
1532        #[test]
1533        fn record_tool_calls_accumulates(n in 0usize..100) {
1534            let mut rs = RunState::default();
1535            rs.record_tool_calls(n);
1536            assert_eq!(rs.turn_tool_calls, n);
1537            assert!(rs.run_has_tool_calls);
1538            assert_eq!(rs.reasoning_only_strikes, 0);
1539            assert_eq!(rs.empty_response_strikes, 0);
1540        }
1541
1542        #[test]
1543        fn record_reasoning_only_increments(count in 1usize..50) {
1544            let mut rs = RunState::default();
1545            for i in 1..=count {
1546                let strikes = rs.record_reasoning_only();
1547                assert_eq!(strikes, i);
1548                assert_eq!(rs.empty_response_strikes, 0);
1549            }
1550        }
1551
1552        #[test]
1553        fn record_empty_response_increments(count in 1usize..50) {
1554            let mut rs = RunState::default();
1555            for i in 1..=count {
1556                let strikes = rs.record_empty_response();
1557                assert_eq!(strikes, i);
1558                assert_eq!(rs.reasoning_only_strikes, 0);
1559            }
1560        }
1561
1562        // ── push_assistant_tool_calls property tests ────────────────────────
1563
1564        #[test]
1565        fn push_assistant_tool_calls_valid_json_unchanged(args in r"\{[^{}]{0,200}\}") {
1566            // Only test strings that are actually valid JSON objects
1567            if serde_json::from_str::<serde_json::Value>(&args).is_err() {
1568                return Ok(());
1569            }
1570            let mut s = make_session();
1571            s.push_assistant_tool_calls(
1572                &[("id".into(), "tool".into(), args.clone())],
1573                None,
1574                None,
1575            );
1576            if let ChatMessage::Assistant { tool_calls: Some(ref tc), .. } = s.chat_messages()[0] {
1577                assert_eq!(tc[0].arguments, args);
1578            } else {
1579                panic!("expected Assistant with tool_calls");
1580            }
1581        }
1582
1583        #[test]
1584        fn push_assistant_tool_calls_invalid_json_sanitized_to_empty(
1585            bad_args in "[a-z\u{4e00}-\u{9fff}]{0,300}"
1586        ) {
1587            // Skip if it happens to be valid JSON
1588            if serde_json::from_str::<serde_json::Value>(&bad_args).is_ok() {
1589                return Ok(());
1590            }
1591            let mut s = make_session();
1592            s.push_assistant_tool_calls(
1593                &[("id".into(), "tool".into(), bad_args)],
1594                None,
1595                None,
1596            );
1597            if let ChatMessage::Assistant { tool_calls: Some(ref tc), .. } = s.chat_messages()[0] {
1598                // Must be valid JSON (no 400 on replay) AND carry no poison
1599                // wrapper the model could echo back as arguments.
1600                serde_json::from_str::<serde_json::Value>(&tc[0].arguments)
1601                    .expect("sanitized args must be valid JSON");
1602                assert_eq!(tc[0].arguments, "{}");
1603            } else {
1604                panic!("expected Assistant with tool_calls");
1605            }
1606        }
1607
1608        #[test]
1609        fn preview_args_escapes_and_caps_at_max_chars(payload in "[a-z\u{4e00}-\u{9fff}\n\t]{0,400}") {
1610            // 截断 WARN 的 args_preview:按字符截取(不是字节),控制字符以
1611            // debug 形式转义(每个输入字符最多膨胀为 2 个字符),长度封顶在
1612            // min(80, len) * 2 + 成对引号。
1613            let preview = AgentSession::preview_args(&payload, 80);
1614            let expected_cap = 2 * payload.chars().count().min(80) + 2;
1615            assert!(
1616                preview.chars().count() <= expected_cap,
1617                "preview must cap at escaped length plus quotes: {preview}"
1618            );
1619            if payload.contains('\n') || payload.contains('\t') {
1620                assert!(
1621                    preview.contains("\\n") || preview.contains("\\t"),
1622                    "control chars must be escaped: {preview}"
1623                );
1624            }
1625            if payload.chars().count() <= 80 && !payload.contains('\n') && !payload.contains('\t') {
1626                assert_eq!(preview, format!("{payload:?}"));
1627            }
1628        }
1629
1630        // ── trim_oldest_turns property tests ────────────────────────────────
1631
1632        #[test]
1633        fn trim_oldest_turns_never_exceeds_max(turns in 1usize..20, max in 1usize..20) {
1634            let mut s = make_session();
1635            for i in 0..turns {
1636                s.push_message(MessageRole::User, format!("u{}", i));
1637                s.push_message(MessageRole::Assistant, format!("a{}", i));
1638            }
1639            s.trim_oldest_turns(max);
1640            assert!(s.turn_count() <= max || turns <= max);
1641        }
1642
1643        // ── validate_message_sequence property tests ────────────────────────
1644
1645        #[test]
1646        fn validate_simple_user_assistant_always_passes(count in 1usize..20) {
1647            let mut msgs = Vec::new();
1648            for i in 0..count {
1649                msgs.push(ChatMessage::user(format!("msg{}", i)));
1650                msgs.push(ChatMessage::assistant(format!("reply{}", i)));
1651            }
1652            assert!(validate_message_sequence(&msgs).is_ok());
1653        }
1654    }
1655}