Skip to main content

zeph_core/agent/
utils.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4use zeph_llm::provider::{LlmProvider, Message, MessagePart, Role};
5
6use super::{Agent, CODE_CONTEXT_PREFIX};
7use crate::channel::Channel;
8use crate::metrics::{MetricsSnapshot, SECURITY_EVENT_CAP, SecurityEvent};
9use zeph_common::SecurityEventCategory;
10use zeph_tools::FilterStats;
11
12/// Fetch entity/edge/community counts from `store`, defaulting each to `0` on a per-metric error.
13///
14/// Shared by [`Agent::sync_graph_counts`] and the background graph-count-sync tasks in
15/// `persistence::extract` (post-extraction refresh and the periodic count-sync task) so the
16/// three call sites cannot drift apart.
17pub(super) async fn fetch_graph_counts(store: &zeph_memory::graph::GraphStore) -> (u64, u64, u64) {
18    let (entities, edges, communities) = tokio::join!(
19        store.entity_count(),
20        store.active_edge_count(),
21        store.community_count()
22    );
23    (
24        entities.unwrap_or(0).cast_unsigned(),
25        edges.unwrap_or(0).cast_unsigned(),
26        communities.unwrap_or(0).cast_unsigned(),
27    )
28}
29
30impl<C: Channel> Agent<C> {
31    /// Read the community-detection failure counter from `SemanticMemory` and update metrics.
32    pub fn sync_community_detection_failures(&self) {
33        if let Some(memory) = self.services.memory.persistence.memory.as_ref() {
34            let failures = memory.community_detection_failures();
35            self.update_metrics(|m| {
36                m.graph_community_detection_failures = failures;
37            });
38        }
39    }
40
41    /// Sync all graph counters (extraction count/failures) from `SemanticMemory` to metrics.
42    pub fn sync_graph_extraction_metrics(&self) {
43        if let Some(memory) = self.services.memory.persistence.memory.as_ref() {
44            let count = memory.graph_extraction_count();
45            let failures = memory.graph_extraction_failures();
46            self.update_metrics(|m| {
47                m.graph_extraction_count = count;
48                m.graph_extraction_failures = failures;
49            });
50        }
51    }
52
53    /// Fetch entity/edge/community counts from the graph store and write to metrics.
54    pub async fn sync_graph_counts(&self) {
55        let Some(memory) = self.services.memory.persistence.memory.as_ref() else {
56            return;
57        };
58        let Some(store) = memory.graph_store.as_ref() else {
59            return;
60        };
61        let (entities, edges, communities) = fetch_graph_counts(store).await;
62        self.update_metrics(|m| {
63            m.graph_entities_total = entities;
64            m.graph_edges_total = edges;
65            m.graph_communities_total = communities;
66        });
67    }
68
69    /// Perform a real health check on the vector store and update metrics.
70    pub async fn check_vector_store_health(&self, backend_name: &str) {
71        let connected = match self.services.memory.persistence.memory.as_ref() {
72            Some(m) => m.is_vector_store_connected().await,
73            None => false,
74        };
75        let name = backend_name.to_owned();
76        self.update_metrics(|m| {
77            m.qdrant_available = connected;
78            m.vector_backend = name;
79        });
80    }
81
82    /// Fetch compression-guidelines metadata from `SQLite` and write to metrics.
83    ///
84    /// Only fetches version and `created_at`; does not load the full guidelines text.
85    /// Feature-gated: compiled only when `compression-guidelines` is enabled.
86    pub async fn sync_guidelines_status(&self) {
87        let Some(memory) = self.services.memory.persistence.memory.as_ref() else {
88            return;
89        };
90        let cid = self.services.memory.persistence.conversation_id;
91        match memory.sqlite().load_compression_guidelines_meta(cid).await {
92            Ok((version, created_at)) => {
93                #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
94                let version_u32 = u32::try_from(version).unwrap_or(0);
95                self.update_metrics(|m| {
96                    m.guidelines_version = version_u32;
97                    m.guidelines_updated_at = created_at;
98                });
99            }
100            Err(e) => {
101                tracing::warn!("failed to sync guidelines status: {e:#}");
102            }
103        }
104    }
105
106    pub(super) fn record_filter_metrics(&mut self, fs: &FilterStats) {
107        let saved = fs.estimated_tokens_saved() as u64;
108        let raw = (fs.raw_chars / 4) as u64;
109        let confidence = fs.confidence;
110        let was_filtered = fs.filtered_chars < fs.raw_chars;
111        self.update_metrics(|m| {
112            m.filter_raw_tokens += raw;
113            m.filter_saved_tokens += saved;
114            m.filter_applications += 1;
115            m.filter_total_commands += 1;
116            if was_filtered {
117                m.filter_filtered_commands += 1;
118            }
119            if let Some(c) = confidence {
120                match c {
121                    zeph_tools::FilterConfidence::Full => {
122                        m.filter_confidence_full += 1;
123                    }
124                    zeph_tools::FilterConfidence::Partial => {
125                        m.filter_confidence_partial += 1;
126                    }
127                    zeph_tools::FilterConfidence::Fallback => {
128                        m.filter_confidence_fallback += 1;
129                    }
130                    _ => {}
131                }
132            }
133        });
134    }
135
136    pub(super) fn update_metrics(&self, f: impl FnOnce(&mut MetricsSnapshot)) {
137        if let Some(ref tx) = self.runtime.metrics.metrics_tx {
138            let elapsed = self.runtime.lifecycle.start_time.elapsed().as_secs();
139            tx.send_modify(|m| {
140                m.uptime_seconds = elapsed;
141                f(m);
142            });
143        }
144    }
145
146    /// Publish the effective context window limit from the active provider's budget into
147    /// [`MetricsSnapshot::context_max_tokens`].
148    ///
149    /// Call after the provider pool is constructed (builder) and on every successful `/provider`
150    /// switch so the TUI context gauge always reflects the active provider's window.
151    /// When no budget is configured the field is set to `0`, which the gauge renders as `"—"`.
152    pub(crate) fn publish_context_budget(&self) {
153        let max_tokens = self
154            .context_manager
155            .budget
156            .as_ref()
157            .map_or(0, |b| b.max_tokens() as u64);
158        self.update_metrics(|m| m.context_max_tokens = max_tokens);
159    }
160
161    /// Flush `metrics.pending_timings` into the rolling window and publish to the metrics snapshot.
162    ///
163    /// Call once per turn after all four phases have written to `pending_timings`.
164    /// Resets `pending_timings` to default after flushing.
165    pub(super) fn flush_turn_timings(&mut self) {
166        let timings = std::mem::take(&mut self.runtime.metrics.pending_timings);
167        tracing::debug!(
168            prepare_context_ms = timings.prepare_context_ms,
169            llm_chat_ms = timings.llm_chat_ms,
170            tool_exec_ms = timings.tool_exec_ms,
171            persist_message_ms = timings.persist_message_ms,
172            "turn timings"
173        );
174
175        if self.runtime.metrics.timing_window.len() >= 10 {
176            self.runtime.metrics.timing_window.pop_front();
177        }
178        self.runtime
179            .metrics
180            .timing_window
181            .push_back(timings.clone());
182
183        let count = self.runtime.metrics.timing_window.len();
184        let mut avg = crate::metrics::TurnTimings::default();
185        let mut max = crate::metrics::TurnTimings::default();
186        for t in &self.runtime.metrics.timing_window {
187            avg.prepare_context_ms = avg.prepare_context_ms.saturating_add(t.prepare_context_ms);
188            avg.llm_chat_ms = avg.llm_chat_ms.saturating_add(t.llm_chat_ms);
189            avg.tool_exec_ms = avg.tool_exec_ms.saturating_add(t.tool_exec_ms);
190            avg.persist_message_ms = avg.persist_message_ms.saturating_add(t.persist_message_ms);
191
192            max.prepare_context_ms = max.prepare_context_ms.max(t.prepare_context_ms);
193            max.llm_chat_ms = max.llm_chat_ms.max(t.llm_chat_ms);
194            max.tool_exec_ms = max.tool_exec_ms.max(t.tool_exec_ms);
195            max.persist_message_ms = max.persist_message_ms.max(t.persist_message_ms);
196        }
197        let n = count as u64;
198        avg.prepare_context_ms /= n;
199        avg.llm_chat_ms /= n;
200        avg.tool_exec_ms /= n;
201        avg.persist_message_ms /= n;
202
203        let total_ms = timings
204            .prepare_context_ms
205            .saturating_add(timings.llm_chat_ms)
206            .saturating_add(timings.tool_exec_ms)
207            .saturating_add(timings.persist_message_ms);
208
209        self.update_metrics(|m| {
210            m.last_turn_timings = timings;
211            m.avg_turn_timings = avg;
212            m.max_turn_timings = max;
213            m.timing_sample_count = n;
214        });
215
216        if let Some(ref recorder) = self.runtime.metrics.histogram_recorder {
217            recorder.observe_turn_duration(std::time::Duration::from_millis(total_ms));
218        }
219    }
220
221    /// Push the current classifier metrics snapshot into `MetricsSnapshot`.
222    ///
223    /// Call this after any classifier invocation (injection, PII, feedback) so the TUI panel
224    /// reflects the latest p50/p95 values. No-op when classifier metrics are not configured.
225    pub(super) fn push_classifier_metrics(&self) {
226        if let Some(ref m) = self.runtime.metrics.classifier_metrics {
227            let snapshot = m.snapshot();
228            self.update_metrics(|ms| ms.classifier = snapshot);
229        }
230    }
231
232    pub(super) fn push_security_event(
233        &self,
234        category: SecurityEventCategory,
235        source: &str,
236        detail: impl Into<String>,
237    ) {
238        if let Some(ref tx) = self.runtime.metrics.metrics_tx {
239            let event = SecurityEvent::new(category, source, detail);
240            let elapsed = self.runtime.lifecycle.start_time.elapsed().as_secs();
241            tx.send_modify(|m| {
242                m.uptime_seconds = elapsed;
243                if m.security_events.len() >= SECURITY_EVENT_CAP {
244                    m.security_events.pop_front();
245                }
246                m.security_events.push_back(event);
247            });
248        }
249    }
250
251    pub(super) fn recompute_prompt_tokens(&mut self) {
252        self.runtime.providers.cached_prompt_tokens = self
253            .msg
254            .messages
255            .iter()
256            .map(|m| self.runtime.metrics.token_counter.count_message_tokens(m) as u64)
257            .sum();
258    }
259
260    pub(super) fn push_message(&mut self, msg: Message) {
261        self.runtime.providers.cached_prompt_tokens +=
262            self.runtime
263                .metrics
264                .token_counter
265                .count_message_tokens(&msg) as u64;
266        if msg.role == zeph_llm::provider::Role::Assistant {
267            self.services.session.last_assistant_at = Some(std::time::Instant::now());
268        }
269        self.msg.messages.push(msg);
270        // Detect MagicDoc headers in tool output after pushing the message.
271        self.detect_magic_docs_in_messages();
272    }
273
274    /// Like [`Self::push_message`], but splices `msg` at `index` instead of appending it at the
275    /// true end — for repairing an out-of-order shutdown-flush tombstone (see
276    /// `shutdown::flush_orphaned_tool_use_on_shutdown`) where a later turn's message may already
277    /// have been appended after the orphaned assistant message this tombstone must immediately
278    /// follow. Token accounting and `MagicDoc` detection are position-independent, so both are
279    /// shared with `push_message`.
280    pub(super) fn insert_message(&mut self, index: usize, msg: Message) {
281        self.runtime.providers.cached_prompt_tokens +=
282            self.runtime
283                .metrics
284                .token_counter
285                .count_message_tokens(&msg) as u64;
286        if msg.role == zeph_llm::provider::Role::Assistant {
287            self.services.session.last_assistant_at = Some(std::time::Instant::now());
288        }
289        self.msg.messages.insert(index, msg);
290        self.detect_magic_docs_in_messages();
291    }
292
293    pub(crate) fn record_cost_and_cache(&self, input_tokens: u64, output_tokens: u64) {
294        let (cache_write, cache_read) = self.provider.last_cache_usage().unwrap_or((0, 0));
295
296        if let Some(ref tracker) = self.runtime.metrics.cost_tracker {
297            let provider_name = if self.runtime.config.active_provider_name.is_empty() {
298                self.provider.name()
299            } else {
300                self.runtime.config.active_provider_name.as_str()
301            };
302            tracker.record_usage(
303                provider_name,
304                self.provider.provider_kind_str(),
305                &self.runtime.config.model_name,
306                input_tokens,
307                cache_read,
308                cache_write,
309                output_tokens,
310            );
311            let breakdown = tracker.provider_breakdown();
312            self.update_metrics(|m| {
313                m.cost_spent_cents = tracker.current_spend();
314                m.cache_creation_tokens += cache_write;
315                m.cache_read_tokens += cache_read;
316                m.provider_cost_breakdown = breakdown;
317            });
318        } else if cache_write > 0 || cache_read > 0 {
319            self.update_metrics(|m| {
320                m.cache_creation_tokens += cache_write;
321                m.cache_read_tokens += cache_read;
322            });
323        }
324    }
325
326    pub(crate) fn record_successful_task(&self) {
327        if let Some(ref tracker) = self.runtime.metrics.cost_tracker {
328            tracker.record_successful_task();
329            self.update_metrics(|m| {
330                m.cost_cps_cents = tracker.cps();
331                m.cost_successful_tasks = tracker.successful_tasks();
332            });
333        }
334    }
335
336    /// Extract a redacted preview of the last assistant message.
337    ///
338    /// Walks `self.msg.messages` in reverse to find the most recent `Role::Assistant`
339    /// message, takes up to `max_chars` Unicode scalar values from `message.content`,
340    /// and applies [`crate::redact::scrub_content`] to redact any secrets.
341    ///
342    /// Returns an empty string when no assistant message exists in the current turn.
343    pub(super) fn last_assistant_preview(&self, max_chars: usize) -> String {
344        let raw = self
345            .msg
346            .messages
347            .iter()
348            .rev()
349            .find(|m| m.role == Role::Assistant)
350            .map_or("", |m| m.content.as_str());
351
352        if raw.is_empty() {
353            return String::new();
354        }
355
356        // Truncate to max_chars before redaction to bound redaction work.
357        let truncated: &str = if raw.chars().count() > max_chars {
358            let end = raw
359                .char_indices()
360                .nth(max_chars)
361                .map_or(raw.len(), |(i, _)| i);
362            &raw[..end]
363        } else {
364            raw
365        };
366
367        crate::redact::scrub_content(truncated).into_owned()
368    }
369
370    /// Inject pre-formatted code context into the message list.
371    /// The caller is responsible for retrieving and formatting the text.
372    pub fn inject_code_context(&mut self, text: &str) {
373        self.remove_code_context_messages();
374        if text.is_empty() || self.msg.messages.len() <= 1 {
375            return;
376        }
377        let content = format!("{CODE_CONTEXT_PREFIX}{text}");
378        self.msg.messages.insert(
379            1,
380            Message::from_parts(
381                Role::System,
382                vec![MessagePart::CodeContext { text: content }],
383            ),
384        );
385    }
386
387    #[must_use]
388    pub fn context_messages(&self) -> &[Message] {
389        &self.msg.messages
390    }
391
392    /// Truncate stale tool result content in old messages to bound in-memory growth.
393    ///
394    /// After the LLM has seen and responded to tool output, the full content is no longer
395    /// needed in the hot message list (it is already persisted to `SQLite`). Truncating keeps
396    /// the in-process message vec small across long sessions.
397    ///
398    /// Skips the last 2 messages so the LLM retains full context for the next turn.
399    ///
400    /// Truncated variants: `MessagePart::ToolResult` (content) and `MessagePart::ToolOutput` (body).
401    pub(super) fn truncate_old_tool_results(&mut self) {
402        const LIMIT: usize = 2048;
403        const SUFFIX: &str = "…[truncated]";
404
405        let len = self.msg.messages.len();
406        if len <= 2 {
407            return;
408        }
409        for msg in &mut self.msg.messages[..len - 2] {
410            for part in &mut msg.parts {
411                match part {
412                    MessagePart::ToolResult { content, .. } if content.len() > LIMIT => {
413                        content.truncate(content.floor_char_boundary(LIMIT));
414                        content.push_str(SUFFIX);
415                    }
416                    MessagePart::ToolOutput { body, .. } if body.len() > LIMIT => {
417                        body.truncate(body.floor_char_boundary(LIMIT));
418                        body.push_str(SUFFIX);
419                    }
420                    _ => {}
421                }
422            }
423        }
424    }
425}
426
427#[cfg(test)]
428mod tests {
429    use super::super::agent_tests::{
430        MockChannel, MockToolExecutor, create_test_registry, mock_provider,
431    };
432    use super::*;
433    use zeph_llm::provider::{MessageMetadata, MessagePart};
434    use zeph_memory::graph::GraphStore;
435    use zeph_memory::graph::types::EntityType;
436    use zeph_memory::store::SqliteStore;
437
438    async fn setup_graph_store() -> GraphStore {
439        let sqlite = SqliteStore::new(":memory:").await.unwrap();
440        GraphStore::new(sqlite.pool().clone())
441    }
442
443    #[tokio::test]
444    async fn fetch_graph_counts_empty_store_returns_zeros() {
445        let store = setup_graph_store().await;
446        assert_eq!(fetch_graph_counts(&store).await, (0, 0, 0));
447    }
448
449    #[tokio::test]
450    async fn fetch_graph_counts_reflects_actual_counts() {
451        let store = setup_graph_store().await;
452        let a = store
453            .upsert_entity("Alice", "Alice", EntityType::Person, None, None)
454            .await
455            .unwrap()
456            .0;
457        let b = store
458            .upsert_entity("Bob", "Bob", EntityType::Person, None, None)
459            .await
460            .unwrap()
461            .0;
462        store
463            .insert_edge(a, b, "knows", "Alice knows Bob", 1.0, None, None)
464            .await
465            .unwrap();
466        store
467            .upsert_community("cluster", "summary", &[a, b], None)
468            .await
469            .unwrap();
470
471        assert_eq!(fetch_graph_counts(&store).await, (2, 1, 1));
472    }
473
474    /// #5677 follow-up: each of the 3 metrics must fall back to `0` independently on its own
475    /// query error, not abort the other two — verified by breaking only `graph_communities`
476    /// while `graph_entities`/`graph_edges` stay intact and populated.
477    #[tokio::test]
478    async fn fetch_graph_counts_falls_back_to_zero_per_field_on_error() {
479        let sqlite = SqliteStore::new(":memory:").await.unwrap();
480        let pool = sqlite.pool().clone();
481        let store = GraphStore::new(pool.clone());
482        let a = store
483            .upsert_entity("Alice", "Alice", EntityType::Person, None, None)
484            .await
485            .unwrap()
486            .0;
487        let b = store
488            .upsert_entity("Bob", "Bob", EntityType::Person, None, None)
489            .await
490            .unwrap()
491            .0;
492        store
493            .insert_edge(a, b, "knows", "Alice knows Bob", 1.0, None, None)
494            .await
495            .unwrap();
496
497        sqlx::query("DROP TABLE graph_communities")
498            .execute(&pool)
499            .await
500            .unwrap();
501
502        assert_eq!(fetch_graph_counts(&store).await, (2, 1, 0));
503    }
504
505    #[test]
506    fn push_message_increments_cached_tokens() {
507        let provider = mock_provider(vec![]);
508        let channel = MockChannel::new(vec![]);
509        let registry = create_test_registry();
510        let executor = MockToolExecutor::no_tools();
511        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
512
513        let before = agent.runtime.providers.cached_prompt_tokens;
514        let msg = Message {
515            role: Role::User,
516            content: "hello world!!".to_string(),
517            parts: vec![],
518            metadata: MessageMetadata::default(),
519        };
520        let expected_delta = agent
521            .runtime
522            .metrics
523            .token_counter
524            .count_message_tokens(&msg) as u64;
525        agent.push_message(msg);
526        assert_eq!(
527            agent.runtime.providers.cached_prompt_tokens,
528            before + expected_delta
529        );
530    }
531
532    /// #5646: `insert_message` must splice at the given index (not append at the end) while
533    /// still tracking token accounting identically to `push_message` — direct coverage of the
534    /// method itself, complementing its indirect exercise via
535    /// `flush_orphaned_tests::flush_orphaned_inserts_tombstone_immediately_after_orphan_not_at_end`
536    /// and `focus_tests::persist_cancelled_tool_results_some_index_inserts_at_that_position`.
537    #[test]
538    fn insert_message_splices_at_index_and_tracks_tokens() {
539        let provider = mock_provider(vec![]);
540        let channel = MockChannel::new(vec![]);
541        let registry = create_test_registry();
542        let executor = MockToolExecutor::no_tools();
543        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
544
545        agent.msg.messages.push(Message {
546            role: Role::User,
547            content: "first".to_string(),
548            parts: vec![],
549            metadata: MessageMetadata::default(),
550        });
551        agent.msg.messages.push(Message {
552            role: Role::User,
553            content: "third".to_string(),
554            parts: vec![],
555            metadata: MessageMetadata::default(),
556        });
557        let insert_idx = agent.msg.messages.len() - 1;
558        let before_tokens = agent.runtime.providers.cached_prompt_tokens;
559
560        let msg = Message {
561            role: Role::User,
562            content: "second".to_string(),
563            parts: vec![],
564            metadata: MessageMetadata::default(),
565        };
566        let expected_delta = agent
567            .runtime
568            .metrics
569            .token_counter
570            .count_message_tokens(&msg) as u64;
571        agent.insert_message(insert_idx, msg);
572
573        assert_eq!(
574            agent.msg.messages[insert_idx].content, "second",
575            "message must be spliced at the given index"
576        );
577        assert_eq!(
578            agent.msg.messages[insert_idx + 1].content,
579            "third",
580            "the message previously at insert_idx must be pushed one slot forward"
581        );
582        assert_eq!(
583            agent.runtime.providers.cached_prompt_tokens,
584            before_tokens + expected_delta,
585            "insert_message must track token accounting identically to push_message"
586        );
587    }
588
589    #[test]
590    fn recompute_prompt_tokens_matches_sum() {
591        let provider = mock_provider(vec![]);
592        let channel = MockChannel::new(vec![]);
593        let registry = create_test_registry();
594        let executor = MockToolExecutor::no_tools();
595        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
596
597        agent.msg.messages.push(Message {
598            role: Role::User,
599            content: "1234".to_string(),
600            parts: vec![],
601            metadata: MessageMetadata::default(),
602        });
603        agent.msg.messages.push(Message {
604            role: Role::Assistant,
605            content: "5678".to_string(),
606            parts: vec![],
607            metadata: MessageMetadata::default(),
608        });
609
610        agent.recompute_prompt_tokens();
611
612        let expected: u64 = agent
613            .msg
614            .messages
615            .iter()
616            .map(|m| agent.runtime.metrics.token_counter.count_message_tokens(m) as u64)
617            .sum();
618        assert_eq!(agent.runtime.providers.cached_prompt_tokens, expected);
619    }
620
621    #[test]
622    fn inject_code_context_into_messages_with_existing_content() {
623        let provider = mock_provider(vec![]);
624        let channel = MockChannel::new(vec![]);
625        let registry = create_test_registry();
626        let executor = MockToolExecutor::no_tools();
627        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
628
629        // Add a user message so we have more than 1 message
630        agent.push_message(Message {
631            role: Role::User,
632            content: "question".to_string(),
633            parts: vec![],
634            metadata: MessageMetadata::default(),
635        });
636
637        agent.inject_code_context("some code here");
638
639        let found = agent.msg.messages.iter().any(|m| {
640            m.parts.iter().any(|p| {
641                matches!(p, MessagePart::CodeContext { text } if text.contains("some code here"))
642            })
643        });
644        assert!(found, "code context should be injected into messages");
645    }
646
647    #[test]
648    fn inject_code_context_empty_text_is_noop() {
649        let provider = mock_provider(vec![]);
650        let channel = MockChannel::new(vec![]);
651        let registry = create_test_registry();
652        let executor = MockToolExecutor::no_tools();
653        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
654
655        agent.push_message(Message {
656            role: Role::User,
657            content: "question".to_string(),
658            parts: vec![],
659            metadata: MessageMetadata::default(),
660        });
661        let count_before = agent.msg.messages.len();
662
663        agent.inject_code_context("");
664
665        // No code context message inserted for empty text
666        assert_eq!(agent.msg.messages.len(), count_before);
667    }
668
669    #[test]
670    fn inject_code_context_with_single_message_is_noop() {
671        let provider = mock_provider(vec![]);
672        let channel = MockChannel::new(vec![]);
673        let registry = create_test_registry();
674        let executor = MockToolExecutor::no_tools();
675        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
676        // Only system prompt → len == 1 → inject should be noop
677        let count_before = agent.msg.messages.len();
678
679        agent.inject_code_context("some code");
680
681        assert_eq!(agent.msg.messages.len(), count_before);
682    }
683
684    #[test]
685    fn context_messages_returns_all_messages() {
686        let provider = mock_provider(vec![]);
687        let channel = MockChannel::new(vec![]);
688        let registry = create_test_registry();
689        let executor = MockToolExecutor::no_tools();
690        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
691
692        agent.push_message(Message {
693            role: Role::User,
694            content: "test".to_string(),
695            parts: vec![],
696            metadata: MessageMetadata::default(),
697        });
698
699        assert_eq!(agent.context_messages().len(), agent.msg.messages.len());
700    }
701
702    #[test]
703    fn truncate_old_tool_results_truncates_stale_content() {
704        let provider = mock_provider(vec![]);
705        let channel = MockChannel::new(vec![]);
706        let registry = create_test_registry();
707        let executor = MockToolExecutor::no_tools();
708        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
709
710        let big_content = "x".repeat(4096);
711
712        // Message 0 (old) — should be truncated.
713        agent.msg.messages.push(Message {
714            role: Role::User,
715            content: String::new(),
716            parts: vec![MessagePart::ToolResult {
717                tool_use_id: "id1".to_string(),
718                content: big_content.clone(),
719                is_error: false,
720            }],
721            metadata: MessageMetadata::default(),
722        });
723        // Message 1 (old) — ToolOutput should also be truncated.
724        agent.msg.messages.push(Message {
725            role: Role::User,
726            content: String::new(),
727            parts: vec![MessagePart::ToolOutput {
728                tool_name: "shell".into(),
729                body: big_content.clone(),
730                compacted_at: None,
731            }],
732            metadata: MessageMetadata::default(),
733        });
734        // Message 2 (recent) — must NOT be truncated.
735        agent.msg.messages.push(Message {
736            role: Role::Assistant,
737            content: "reply".to_string(),
738            parts: vec![MessagePart::ToolResult {
739                tool_use_id: "id3".to_string(),
740                content: big_content.clone(),
741                is_error: false,
742            }],
743            metadata: MessageMetadata::default(),
744        });
745        // Message 3 (most recent) — must NOT be truncated.
746        agent.msg.messages.push(Message {
747            role: Role::User,
748            content: "last".to_string(),
749            parts: vec![MessagePart::ToolResult {
750                tool_use_id: "id4".to_string(),
751                content: big_content.clone(),
752                is_error: false,
753            }],
754            metadata: MessageMetadata::default(),
755        });
756
757        // Agent::new inserts a system prompt at index 0, so our messages are at 1..=4.
758        let base = agent.msg.messages.len() - 4;
759
760        agent.truncate_old_tool_results();
761
762        // Old ToolResult truncated.
763        if let MessagePart::ToolResult { content, .. } = &agent.msg.messages[base].parts[0] {
764            assert!(
765                content.ends_with("…[truncated]"),
766                "msg[base] should be truncated"
767            );
768            assert!(content.len() <= 2048 + 16);
769        } else {
770            panic!("expected ToolResult at msg[base]");
771        }
772
773        // Old ToolOutput truncated.
774        if let MessagePart::ToolOutput { body, .. } = &agent.msg.messages[base + 1].parts[0] {
775            assert!(
776                body.ends_with("…[truncated]"),
777                "msg[base+1] should be truncated"
778            );
779        } else {
780            panic!("expected ToolOutput at msg[base+1]");
781        }
782
783        // Recent messages untouched.
784        if let MessagePart::ToolResult { content, .. } = &agent.msg.messages[base + 2].parts[0] {
785            assert_eq!(content.len(), 4096, "msg[base+2] should NOT be truncated");
786        } else {
787            panic!("expected ToolResult at msg[base+2]");
788        }
789        if let MessagePart::ToolResult { content, .. } = &agent.msg.messages[base + 3].parts[0] {
790            assert_eq!(content.len(), 4096, "msg[base+3] should NOT be truncated");
791        } else {
792            panic!("expected ToolResult at msg[base+3]");
793        }
794    }
795
796    #[test]
797    fn truncate_old_tool_results_noop_when_few_messages() {
798        let provider = mock_provider(vec![]);
799        let channel = MockChannel::new(vec![]);
800        let registry = create_test_registry();
801        let executor = MockToolExecutor::no_tools();
802        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
803
804        let big = "y".repeat(4096);
805        agent.msg.messages.push(Message {
806            role: Role::User,
807            content: String::new(),
808            parts: vec![MessagePart::ToolResult {
809                tool_use_id: "id".to_string(),
810                content: big.clone(),
811                is_error: false,
812            }],
813            metadata: MessageMetadata::default(),
814        });
815        agent.msg.messages.push(Message {
816            role: Role::Assistant,
817            content: "ok".to_string(),
818            parts: vec![MessagePart::ToolResult {
819                tool_use_id: "id2".to_string(),
820                content: big.clone(),
821                is_error: false,
822            }],
823            metadata: MessageMetadata::default(),
824        });
825
826        // Agent::new inserts a system prompt at index 0; our messages are at 1 and 2.
827        let len_before = agent.msg.messages.len();
828        agent.truncate_old_tool_results();
829
830        // Neither message truncated — both fall in the last-2 window (len=3, skip last 2).
831        assert_eq!(agent.msg.messages.len(), len_before);
832        if let MessagePart::ToolResult { content, .. } =
833            &agent.msg.messages[len_before - 2].parts[0]
834        {
835            assert_eq!(
836                content.len(),
837                4096,
838                "second-to-last should not be truncated"
839            );
840        } else {
841            panic!("expected ToolResult");
842        }
843        if let MessagePart::ToolResult { content, .. } =
844            &agent.msg.messages[len_before - 1].parts[0]
845        {
846            assert_eq!(content.len(), 4096, "last should not be truncated");
847        } else {
848            panic!("expected ToolResult");
849        }
850    }
851
852    fn make_timings(ctx: u64, llm: u64, tool: u64, persist: u64) -> crate::metrics::TurnTimings {
853        crate::metrics::TurnTimings {
854            prepare_context_ms: ctx,
855            llm_chat_ms: llm,
856            tool_exec_ms: tool,
857            persist_message_ms: persist,
858        }
859    }
860
861    fn agent_with_metrics_watch() -> (
862        Agent<MockChannel>,
863        tokio::sync::watch::Receiver<crate::metrics::MetricsSnapshot>,
864    ) {
865        let provider = mock_provider(vec![]);
866        let channel = MockChannel::new(vec![]);
867        let registry = create_test_registry();
868        let executor = MockToolExecutor::no_tools();
869        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
870
871        let (tx, rx) = tokio::sync::watch::channel(crate::metrics::MetricsSnapshot::default());
872        agent.runtime.metrics.metrics_tx = Some(tx);
873        (agent, rx)
874    }
875
876    // T1-a: single flush — last_turn_timings equals the flushed value, count == 1.
877    #[test]
878    fn flush_turn_timings_single_flush() {
879        let (mut agent, rx) = agent_with_metrics_watch();
880
881        agent.runtime.metrics.pending_timings = make_timings(10, 200, 50, 5);
882        agent.flush_turn_timings();
883
884        let snap = rx.borrow();
885        assert_eq!(snap.last_turn_timings.prepare_context_ms, 10);
886        assert_eq!(snap.last_turn_timings.llm_chat_ms, 200);
887        assert_eq!(snap.last_turn_timings.tool_exec_ms, 50);
888        assert_eq!(snap.last_turn_timings.persist_message_ms, 5);
889        assert_eq!(snap.timing_sample_count, 1);
890        // avg == last when sample_count == 1
891        assert_eq!(snap.avg_turn_timings.llm_chat_ms, 200);
892    }
893
894    // T1-b: pending_timings reset to default after flush.
895    #[test]
896    fn flush_turn_timings_resets_pending() {
897        let provider = mock_provider(vec![]);
898        let channel = MockChannel::new(vec![]);
899        let registry = create_test_registry();
900        let executor = MockToolExecutor::no_tools();
901        let mut agent = Agent::new(provider, channel, registry, None, 5, executor);
902
903        agent.runtime.metrics.pending_timings = make_timings(10, 200, 50, 5);
904        agent.flush_turn_timings();
905
906        let p = &agent.runtime.metrics.pending_timings;
907        assert_eq!(p.prepare_context_ms, 0);
908        assert_eq!(p.llm_chat_ms, 0);
909        assert_eq!(p.tool_exec_ms, 0);
910        assert_eq!(p.persist_message_ms, 0);
911    }
912
913    // T1-c: window capped at 10; avg and max computed correctly.
914    #[test]
915    fn flush_turn_timings_window_capped_at_10() {
916        let (mut agent, rx) = agent_with_metrics_watch();
917
918        // Push 12 turns: llm_chat_ms = i * 10 for i in 1..=12.
919        for i in 1_u64..=12 {
920            agent.runtime.metrics.pending_timings = make_timings(0, i * 10, 0, 0);
921            agent.flush_turn_timings();
922        }
923
924        let snap = rx.borrow();
925        // Window holds last 10: turns 3..=12, llm values 30..=120.
926        assert_eq!(snap.timing_sample_count, 10);
927        // max = 120
928        assert_eq!(snap.max_turn_timings.llm_chat_ms, 120);
929        // avg of 30,40,...,120 = (30+120)*10/2/10 = 75
930        assert_eq!(snap.avg_turn_timings.llm_chat_ms, 75);
931    }
932}