1use std::pin::Pin;
13use std::sync::Arc;
14
15use pi_agent::AgentMessage;
16use pi_ai::{AssistantContent, Message, StopReason};
17use tokio_util::sync::CancellationToken;
18
19use super::AgentSession;
20use super::events::AgentSessionEvent;
21use crate::core::compaction::{
22 GenerateBranchSummaryOptions, SummarizeStreamFn, collect_entries_for_branch_summary,
23 generate_branch_summary,
24};
25use crate::core::export_html::{
26 ExportError, ExportOptions, RenderedResult, RenderedToolHtml, SessionExportState,
27 ToolHtmlRenderer, export_session_to_html,
28};
29use crate::core::session_transfer::{SessionTransferError, export_branch_to_jsonl};
30use crate::core::sessions::{SessionEntry, SessionError};
31
32#[derive(Debug, thiserror::Error)]
38pub enum TreeError {
39 #[error("Entry {0} not found")]
41 EntryNotFound(String),
42 #[error("No model available for summarization")]
44 NoModel,
45 #[error("Branch summarization failed: {0}")]
47 Summarization(String),
48 #[error(transparent)]
50 Export(#[from] SessionTransferError),
51 #[error(transparent)]
53 HtmlExport(#[from] ExportError),
54 #[error(transparent)]
56 Session(#[from] SessionError),
57}
58
59#[derive(Clone, Debug, Default)]
61pub struct NavigateTreeOptions {
62 pub summarize: bool,
64 pub custom_instructions: Option<String>,
66 pub replace_instructions: bool,
68 pub label: Option<String>,
70}
71
72#[derive(Clone, Debug, Default)]
74pub struct NavigateTreeResult {
75 pub editor_text: Option<String>,
77 pub cancelled: bool,
79 pub aborted: bool,
81 pub summary_entry: Option<SessionEntry>,
83}
84
85#[derive(Clone, Debug, PartialEq, Eq)]
87pub struct ForkableUserMessage {
88 pub entry_id: String,
90 pub text: String,
92}
93
94#[derive(Clone, Debug, Default)]
99pub struct SummarizationAuth {
100 pub api_key: Option<String>,
102 pub headers: Option<std::collections::BTreeMap<String, Option<String>>>,
104 pub env: Option<std::collections::BTreeMap<String, String>>,
106}
107
108#[derive(Clone, Debug, Default)]
110pub struct TreePreparation {
111 pub target_id: String,
113 pub old_leaf_id: Option<String>,
115 pub common_ancestor_id: Option<String>,
117 pub entries_to_summarize: Vec<SessionEntry>,
119 pub user_wants_summary: bool,
121 pub custom_instructions: Option<String>,
123 pub replace_instructions: bool,
125 pub label: Option<String>,
127}
128
129pub type ToolHtmlPreRenderer = Arc<
137 dyn Fn(
138 String,
139 String,
140 serde_json::Value,
141 ) -> Pin<Box<dyn Future<Output = Option<RenderedToolHtml>> + Send>>
142 + Send
143 + Sync,
144>;
145
146struct MapToolHtmlRenderer {
151 calls: std::collections::HashMap<String, RenderedToolHtml>,
152}
153
154impl ToolHtmlRenderer for MapToolHtmlRenderer {
155 fn render_call(
156 &self,
157 tool_call_id: &str,
158 _tool_name: &str,
159 _arguments: &serde_json::Value,
160 ) -> Option<String> {
161 self.calls
162 .get(tool_call_id)
163 .and_then(|r| r.call_html.clone())
164 }
165
166 fn render_result(
167 &self,
168 tool_call_id: &str,
169 _tool_name: &str,
170 _result: &[pi_ai::ToolResultContent],
171 _details: Option<&serde_json::Value>,
172 _is_error: bool,
173 ) -> Option<RenderedResult> {
174 self.calls.get(tool_call_id).map(|r| RenderedResult {
175 collapsed: r.result_html_collapsed.clone(),
176 expanded: r.result_html_expanded.clone(),
177 })
178 }
179}
180
181impl AgentSession {
182 pub async fn navigate_tree(
200 self: &Arc<Self>,
201 target_id: &str,
202 options: NavigateTreeOptions,
203 auth: SummarizationAuth,
204 summarizer: Option<&SummarizeStreamFn>,
205 ) -> Result<NavigateTreeResult, TreeError> {
206 let old_leaf_id = {
207 let sm = self.session_manager.lock().await;
208 sm.get_leaf_id().map(str::to_owned)
209 };
210
211 if Some(target_id) == old_leaf_id.as_deref() {
212 if let Some(label) = options.label.as_deref() {
216 self.session_manager
217 .lock()
218 .await
219 .append_label_change(target_id, Some(label))?;
220 }
221 return Ok(NavigateTreeResult::default());
222 }
223
224 if options.summarize {
225 let model = self.model();
226 if model.id.is_empty() {
227 return Err(TreeError::NoModel);
228 }
229 }
230
231 let (target_entry, preparation) = {
233 let sm = self.session_manager.lock().await;
234 let target_entry = sm
235 .get_entry(target_id)
236 .cloned()
237 .ok_or_else(|| TreeError::EntryNotFound(target_id.to_owned()))?;
238
239 let collected =
240 collect_entries_for_branch_summary(&sm, old_leaf_id.as_deref(), target_id);
241 let prep = TreePreparation {
242 target_id: target_id.to_owned(),
243 old_leaf_id: old_leaf_id.clone(),
244 common_ancestor_id: collected.common_ancestor_id.clone(),
245 entries_to_summarize: collected.entries.clone(),
246 user_wants_summary: options.summarize,
247 custom_instructions: options.custom_instructions.clone(),
248 replace_instructions: options.replace_instructions,
249 label: options.label.clone(),
250 };
251 (target_entry, prep)
252 };
253
254 let token = self.begin_branch_summary_abort();
256
257 let _before_handlers = self.has_extension_handlers("session_before_tree");
259
260 let result = self
261 .navigate_tree_inner(target_entry, preparation, auth, summarizer, &token)
262 .await;
263 self.clear_branch_summary_abort();
264 result
265 }
266
267 async fn navigate_tree_inner(
268 self: &Arc<Self>,
269 target_entry: SessionEntry,
270 preparation: TreePreparation,
271 auth: SummarizationAuth,
272 summarizer: Option<&SummarizeStreamFn>,
273 token: &CancellationToken,
274 ) -> Result<NavigateTreeResult, TreeError> {
275 let mut summary_text: Option<String> = None;
277 let mut summary_details: Option<serde_json::Value> = None;
278 let mut from_extension = false;
279 if preparation.user_wants_summary
280 && !preparation.entries_to_summarize.is_empty()
281 && let Some(stream_fn) = summarizer
282 {
283 let model = self.model();
284 let reserve_tokens = self
285 .lock_settings()
286 .get_branch_summary_settings()
287 .reserve_tokens;
288 let opts = GenerateBranchSummaryOptions {
289 model: model.clone(),
290 api_key: auth.api_key.clone(),
291 headers: auth.headers.clone(),
292 env: auth.env.clone(),
293 signal: token.clone(),
294 custom_instructions: preparation.custom_instructions.clone(),
295 replace_instructions: preparation.replace_instructions,
296 reserve_tokens: Some(reserve_tokens),
297 stream_fn: Arc::clone(stream_fn),
298 };
299 let result = generate_branch_summary(&preparation.entries_to_summarize, opts)
300 .await
301 .map_err(|e| TreeError::Summarization(e.to_string()))?;
302 if result.aborted.unwrap_or(false) {
303 return Ok(NavigateTreeResult {
304 cancelled: true,
305 aborted: true,
306 ..Default::default()
307 });
308 }
309 if let Some(err) = result.error.clone() {
310 return Err(TreeError::Summarization(err));
311 }
312 summary_text = result.summary;
313 summary_details = Some(serde_json::json!({
314 "readFiles": result.read_files.unwrap_or_default(),
315 "modifiedFiles": result.modified_files.unwrap_or_default(),
316 }));
317 from_extension = false;
318 }
319 let _ = from_extension;
320
321 let (new_leaf_id, editor_text) = compute_new_leaf_and_editor_text(&target_entry);
323
324 let summary_entry = {
326 let mut sm = self.session_manager.lock().await;
327 if let Some(text) = summary_text.as_deref() {
328 let id = sm.branch_with_summary(
329 new_leaf_id.as_deref(),
330 text,
331 summary_details.clone(),
332 from_extension.then_some(true),
333 )?;
334 if let Some(l) = preparation.label.as_deref() {
335 let _ = sm.append_label_change(&id, Some(l));
336 }
337 sm.get_entry(&id).cloned()
338 } else if new_leaf_id.is_none() {
339 sm.reset_leaf();
340 if let Some(l) = preparation.label.as_deref() {
341 let _ = sm.append_label_change(&preparation.target_id, Some(l));
342 }
343 None
344 } else {
345 sm.branch(new_leaf_id.as_deref().unwrap_or(""))?;
346 if let Some(l) = preparation.label.as_deref() {
347 let _ = sm.append_label_change(&preparation.target_id, Some(l));
348 }
349 None
350 }
351 };
352
353 let session_context = {
355 let sm = self.session_manager.lock().await;
356 sm.build_session_context()
357 .map_err(|e| TreeError::Summarization(e.to_string()))?
358 };
359 self.agent.replace_messages(session_context.messages);
360
361 let _ = self.has_extension_handlers("session_tree");
363
364 Ok(NavigateTreeResult {
365 editor_text,
366 cancelled: false,
367 aborted: false,
368 summary_entry,
369 })
370 }
371
372 fn begin_branch_summary_abort(&self) -> CancellationToken {
374 let token = CancellationToken::new();
375 let mut inner = self.lock_inner();
376 if let Some(prev) = inner.branch_summary_abort.take() {
377 prev.cancel();
378 }
379 inner.branch_summary_abort = Some(token.clone());
380 token
381 }
382
383 fn clear_branch_summary_abort(&self) {
385 self.lock_inner().branch_summary_abort = None;
386 }
387
388 pub fn abort_branch_summary(&self) {
390 let mut inner = self.lock_inner();
391 if let Some(token) = inner.branch_summary_abort.take() {
392 token.cancel();
393 }
394 }
395
396 pub async fn get_user_messages_for_forking(&self) -> Vec<ForkableUserMessage> {
401 let entries: Vec<SessionEntry> = {
402 let sm = self.session_manager.lock().await;
403 sm.get_entries().into_iter().cloned().collect()
404 };
405 let mut out = Vec::new();
406 for entry in entries {
407 if let SessionEntry::Message(m) = &entry {
408 if m.message.role() != "user" {
409 continue;
410 }
411 let text = extract_user_message_text(&m.message);
412 if !text.is_empty() {
413 out.push(ForkableUserMessage {
414 entry_id: m.id.clone(),
415 text,
416 });
417 }
418 }
419 }
420 out
421 }
422
423 pub async fn export_to_html(
433 &self,
434 output_path: Option<&str>,
435 tool_pre_renderer: Option<ToolHtmlPreRenderer>,
436 ) -> Result<String, ExportError> {
437 let snapshot = self.agent.state();
439 let state = SessionExportState::from_agent_snapshot(&snapshot);
440
441 let theme_name = self.lock_settings().get_theme();
443
444 let map_renderer = if let Some(renderer) = tool_pre_renderer {
446 let entries = {
447 let sm = self.session_manager.lock().await;
448 sm.get_branch(None).into_iter().cloned().collect::<Vec<_>>()
449 };
450 let mut calls: std::collections::HashMap<String, RenderedToolHtml> =
451 std::collections::HashMap::new();
452 for entry in &entries {
453 if let SessionEntry::Message(m) = entry {
454 if let Some(Message::Assistant(assistant)) = m.message.as_llm() {
455 for block in &assistant.content {
456 if let AssistantContent::ToolCall(call) = block {
457 let args = serde_json::Value::Object((*call.arguments).clone());
458 if let Some(rendered) =
459 renderer(call.id.clone(), call.name.clone(), args).await
460 {
461 calls.insert(call.id.clone(), rendered);
462 }
463 }
464 }
465 }
466 if let Some(Message::ToolResult(result)) = m.message.as_llm() {
467 let payload = serde_json::to_value(&result.content).unwrap_or_default();
468 if let Some(rendered) = renderer(
469 result.tool_call_id.clone(),
470 result.tool_name.clone(),
471 payload,
472 )
473 .await
474 {
475 let entry = calls.entry(result.tool_call_id.clone()).or_default();
476 entry.result_html_collapsed = rendered.result_html_collapsed;
477 entry.result_html_expanded = rendered.result_html_expanded;
478 }
479 }
480 }
481 }
482 Some(MapToolHtmlRenderer { calls })
483 } else {
484 None
485 };
486
487 let sm = self.session_manager.lock().await;
488 let opts = ExportOptions {
489 output_path: output_path.map(std::path::PathBuf::from),
490 theme_name,
491 theme: None,
492 tool_renderer: map_renderer.as_ref().map(|r| r as &dyn ToolHtmlRenderer),
493 };
494 export_session_to_html(&sm, Some(&state), opts)
495 }
496
497 pub async fn export_to_jsonl(
503 &self,
504 output_path: Option<&str>,
505 ) -> Result<String, SessionTransferError> {
506 let sm = self.session_manager.lock().await;
507 export_branch_to_jsonl(&sm, output_path)
508 }
509
510 #[must_use]
514 pub fn get_last_assistant_text(&self) -> Option<String> {
515 let messages = self.agent.transcript();
516 for message in messages.into_iter().rev() {
517 if message.role() != "assistant" {
518 continue;
519 }
520 let Some(Message::Assistant(assistant)) = message.as_llm() else {
521 continue;
522 };
523 if matches!(assistant.stop_reason, StopReason::Aborted) && assistant.content.is_empty()
524 {
525 continue;
526 }
527 let text: String = assistant
528 .content
529 .iter()
530 .filter_map(|c| match c {
531 AssistantContent::Text(t) => Some(t.text.to_string()),
532 _ => None,
533 })
534 .collect();
535 let trimmed = text.trim();
536 if trimmed.is_empty() {
537 return None;
538 }
539 return Some(trimmed.to_owned());
540 }
541 None
542 }
543
544 pub async fn set_session_name(&self, name: &str) -> Result<(), SessionError> {
552 let resolved_name = {
553 let mut sm = self.session_manager.lock().await;
554 sm.append_session_info(name)?;
555 sm.get_session_name()
556 };
557 self.emit_public(AgentSessionEvent::SessionInfoChanged {
558 name: resolved_name.clone(),
559 });
560 if self.has_extension_handlers("session_info_changed") {
561 let runner = self.hooks.runner();
562 tokio::spawn(async move {
564 let _ = runner
565 .emit(AgentSessionEvent::SessionInfoChanged {
566 name: resolved_name,
567 })
568 .await;
569 });
570 }
571 Ok(())
572 }
573}
574
575fn compute_new_leaf_and_editor_text(
581 target_entry: &SessionEntry,
582) -> (Option<String>, Option<String>) {
583 match target_entry {
584 SessionEntry::Message(m) if m.message.role() == "user" => {
585 let text = extract_user_message_text(&m.message);
586 let text = if text.is_empty() { None } else { Some(text) };
587 (m.parent_id.clone(), text)
588 }
589 SessionEntry::CustomMessage(m) => {
590 let text = extract_custom_message_text(&m.content);
591 (m.parent_id.clone(), text)
592 }
593 _ => (target_entry.id().map(str::to_owned), None),
594 }
595}
596
597pub(super) fn extract_user_message_text(message: &AgentMessage) -> String {
599 let Some(Message::User(user)) = message.as_llm() else {
600 return String::new();
601 };
602 match &user.content {
603 pi_ai::UserMessageContent::Text(s) => s.clone(),
604 pi_ai::UserMessageContent::Blocks(blocks) => {
605 let mut out = String::new();
606 for block in blocks {
607 if let pi_ai::UserContent::Text(t) = block {
608 out.push_str(&t.text.read());
609 }
610 }
611 out
612 }
613 }
614}
615
616#[must_use]
618pub fn extract_user_message_text_pub(message: &AgentMessage) -> String {
619 extract_user_message_text(message)
620}
621
622fn extract_custom_message_text(
623 content: &crate::core::messages::CustomMessageContent,
624) -> Option<String> {
625 use crate::core::messages::CustomMessageContent;
626 let text = match content {
627 CustomMessageContent::Text(s) => s.clone(),
628 CustomMessageContent::Blocks(blocks) => {
629 let mut out = String::new();
630 for block in blocks {
631 if let pi_ai::UserContent::Text(t) = block {
632 out.push_str(&t.text.read());
633 }
634 }
635 out
636 }
637 };
638 if text.is_empty() { None } else { Some(text) }
639}
640
641#[cfg(test)]
646mod tests {
647 use super::*;
648 use crate::core::agent_session::{AgentSession, AgentSessionConfig};
649 use crate::core::sessions::SessionManager;
650 use crate::core::settings::SettingsManager;
651 use futures::stream::{self, BoxStream, StreamExt};
652 use pi_ai::{
653 AssistantMessage, AssistantMessageEvent, Context, Model, ModelCost, ModelInput, Provider,
654 ProviderError, StreamOptions, Usage,
655 };
656
657 type TestResult<T = ()> = Result<T, Box<dyn std::error::Error>>;
658
659 fn missing(context: &'static str) -> std::io::Error {
660 std::io::Error::other(context)
661 }
662
663 fn test_model() -> Model {
664 Model {
665 id: "m".to_owned(),
666 name: "m".to_owned(),
667 api: "test-api".to_owned(),
668 provider: "test-provider".to_owned(),
669 base_url: String::new(),
670 reasoning: false,
671 thinking_level_map: None,
672 input: vec![ModelInput::Text],
673 cost: ModelCost::default(),
674 context_window: 8_192,
675 max_tokens: 1_024,
676 headers: None,
677 compat: None,
678 extra: std::collections::BTreeMap::new(),
679 }
680 }
681
682 #[derive(Clone)]
683 struct StubProvider;
684
685 impl Provider for StubProvider {
686 fn stream(
687 &self,
688 _model: &Model,
689 _context: Context,
690 _options: StreamOptions,
691 ) -> BoxStream<'static, Result<AssistantMessageEvent, ProviderError>> {
692 stream::empty().boxed()
693 }
694 }
695
696 fn assistant_with_usage(text: &str, usage: Usage) -> AssistantMessage {
697 let mut message =
698 AssistantMessage::new("test-api", "test-provider", "m", pi_agent::now_millis());
699 message
700 .content
701 .push(pi_ai::AssistantContent::Text(pi_ai::TextContent::new(text)));
702 message.stop_reason = pi_ai::StopReason::Stop;
703 message.usage = usage;
704 message
705 }
706
707 fn make_session() -> TestResult<Arc<AgentSession>> {
708 let config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
709 AgentSession::new(config).map_err(Into::into)
710 }
711
712 #[tokio::test]
713 async fn get_user_messages_for_forking_returns_user_text() -> TestResult {
714 let session = make_session()?;
715 {
716 let mut sm = session.session_manager.lock().await;
717 sm.append_message(&AgentMessage::Llm(Box::new(Message::User(
718 pi_ai::UserMessage::new(pi_ai::UserMessageContent::Text("hello".into()), 0),
719 ))))?;
720 sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
721 assistant_with_usage("hi back", Usage::default()),
722 ))))?;
723 sm.append_message(&AgentMessage::Llm(Box::new(Message::User(
724 pi_ai::UserMessage::new(pi_ai::UserMessageContent::Text("second".into()), 1),
725 ))))?;
726 }
727 let messages = session.get_user_messages_for_forking().await;
728 assert_eq!(messages.len(), 2);
729 assert_eq!(messages[0].text, "hello");
730 assert_eq!(messages[1].text, "second");
731 Ok(())
732 }
733
734 #[tokio::test]
735 async fn get_last_assistant_text_skips_aborted_empty() -> TestResult {
736 let session = make_session()?;
737 let mut aborted = AssistantMessage::new("test-api", "test-provider", "m", 0);
738 aborted.stop_reason = pi_ai::StopReason::Aborted;
739 let mut good = AssistantMessage::new("test-api", "test-provider", "m", 1);
740 good.content
741 .push(pi_ai::AssistantContent::Text(pi_ai::TextContent::new(
742 "real text",
743 )));
744 good.stop_reason = pi_ai::StopReason::Stop;
745
746 session
747 .agent
748 .push_message(AgentMessage::Llm(Box::new(Message::Assistant(aborted))));
749 session
750 .agent
751 .push_message(AgentMessage::Llm(Box::new(Message::Assistant(good))));
752
753 let text = session.get_last_assistant_text();
754 assert_eq!(text.as_deref(), Some("real text"));
755 Ok(())
756 }
757
758 #[tokio::test]
759 async fn get_last_assistant_text_returns_none_when_only_aborted() -> TestResult {
760 let session = make_session()?;
761 let mut aborted = AssistantMessage::new("test-api", "test-provider", "m", 0);
762 aborted.stop_reason = pi_ai::StopReason::Aborted;
763 session
764 .agent
765 .push_message(AgentMessage::Llm(Box::new(Message::Assistant(aborted))));
766 let text = session.get_last_assistant_text();
767 assert_eq!(text, None);
768 Ok(())
769 }
770
771 #[tokio::test]
772 async fn set_session_name_emits_event_and_persists() -> TestResult {
773 let session = make_session()?;
774 let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<Option<String>>();
775 let _unsub = session.subscribe(move |event| {
776 if let AgentSessionEvent::SessionInfoChanged { name } = event {
777 let _ = tx.send(name.clone());
778 }
779 });
780 session.set_session_name("My Session").await?;
781 let name = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await?;
782 let name = name.ok_or_else(|| missing("session name receiver closed"))?;
783 assert_eq!(name.as_deref(), Some("My Session"));
784 let persisted = session.session_name().await;
785 assert_eq!(persisted.as_deref(), Some("My Session"));
786 Ok(())
787 }
788
789 #[tokio::test]
790 async fn set_session_name_collapses_newlines_to_single_spaces() -> TestResult {
791 let session = make_session()?;
792 session.set_session_name("line1\n\nline2").await?;
793 let persisted = session.session_name().await;
794 assert_eq!(persisted.as_deref(), Some("line1 line2"));
795 Ok(())
796 }
797
798 #[tokio::test]
799 async fn navigate_tree_noop_when_already_at_target() -> TestResult {
800 let session = make_session()?;
801 let id = {
802 let mut sm = session.session_manager.lock().await;
803 sm.append_message(&AgentMessage::Llm(Box::new(Message::User(
804 pi_ai::UserMessage::new(pi_ai::UserMessageContent::Text("x".into()), 0),
805 ))))?
806 };
807 let result = session
808 .navigate_tree(
809 &id,
810 NavigateTreeOptions::default(),
811 SummarizationAuth::default(),
812 None,
813 )
814 .await?;
815 assert!(!result.cancelled);
816 assert!(result.summary_entry.is_none());
817 Ok(())
818 }
819
820 #[tokio::test]
821 async fn navigate_tree_unknown_target_errors() -> TestResult {
822 let session = make_session()?;
823 let result = session
824 .navigate_tree(
825 "missing",
826 NavigateTreeOptions::default(),
827 SummarizationAuth::default(),
828 None,
829 )
830 .await;
831 assert!(matches!(result, Err(TreeError::EntryNotFound(_))));
832 Ok(())
833 }
834
835 #[tokio::test]
836 async fn navigate_tree_to_user_message_sets_editor_text_and_leaf_to_parent() -> TestResult {
837 let session = make_session()?;
838 let user_id = {
839 let mut sm = session.session_manager.lock().await;
840 sm.append_message(&AgentMessage::Llm(Box::new(Message::User(
841 pi_ai::UserMessage::new(pi_ai::UserMessageContent::Text("hello".into()), 0),
842 ))))?
843 };
844 {
845 let mut sm = session.session_manager.lock().await;
846 sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
847 assistant_with_usage("reply", Usage::default()),
848 ))))?;
849 }
850 let result = session
851 .navigate_tree(
852 &user_id,
853 NavigateTreeOptions::default(),
854 SummarizationAuth::default(),
855 None,
856 )
857 .await?;
858 assert_eq!(result.editor_text.as_deref(), Some("hello"));
859 let leaf = {
860 let sm = session.session_manager.lock().await;
861 sm.get_leaf_id().map(str::to_owned)
862 };
863 assert!(
864 leaf.is_none(),
865 "expected null leaf after navigating to root user, got {leaf:?}"
866 );
867 Ok(())
868 }
869
870 #[tokio::test]
871 async fn navigate_tree_to_assistant_sets_leaf_to_target() -> TestResult {
872 let session = make_session()?;
873 let (id1, _id2) = {
874 let mut sm = session.session_manager.lock().await;
875 let a = sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
876 assistant_with_usage("first", Usage::default()),
877 ))))?;
878 let b = sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
879 assistant_with_usage("second", Usage::default()),
880 ))))?;
881 (a, b)
882 };
883 let result = session
884 .navigate_tree(
885 &id1,
886 NavigateTreeOptions::default(),
887 SummarizationAuth::default(),
888 None,
889 )
890 .await?;
891 assert!(result.editor_text.is_none());
892 let leaf = {
893 let sm = session.session_manager.lock().await;
894 sm.get_leaf_id().map(str::to_owned)
895 };
896 assert_eq!(leaf.as_deref(), Some(id1.as_str()));
897 Ok(())
898 }
899
900 #[tokio::test]
901 async fn navigate_tree_attaches_label_to_target_when_no_summary() -> TestResult {
902 let session = make_session()?;
903 let id = {
904 let mut sm = session.session_manager.lock().await;
905 sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
906 assistant_with_usage("hi", Usage::default()),
907 ))))?
908 };
909 session
910 .navigate_tree(
911 &id,
912 NavigateTreeOptions {
913 label: Some("bookmark".into()),
914 ..Default::default()
915 },
916 SummarizationAuth::default(),
917 None,
918 )
919 .await?;
920 let label = {
921 let sm = session.session_manager.lock().await;
922 sm.get_label(&id).map(str::to_owned)
923 };
924 assert_eq!(label.as_deref(), Some("bookmark"));
925 Ok(())
926 }
927
928 #[tokio::test]
929 async fn navigate_tree_summarize_requires_model() -> TestResult {
930 let session = make_session()?;
931 let target = {
932 let mut sm = session.session_manager.lock().await;
933 sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
934 assistant_with_usage("first", Usage::default()),
935 ))))?
936 };
937 {
938 let mut sm = session.session_manager.lock().await;
939 sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
940 assistant_with_usage("second", Usage::default()),
941 ))))?;
942 }
943 let mut empty_model = test_model();
944 empty_model.id = String::new();
945 session.agent.set_model(empty_model);
946 let result = session
947 .navigate_tree(
948 &target,
949 NavigateTreeOptions {
950 summarize: true,
951 ..Default::default()
952 },
953 SummarizationAuth::default(),
954 None,
955 )
956 .await;
957 assert!(matches!(result, Err(TreeError::NoModel)), "got {result:?}");
958 Ok(())
959 }
960
961 #[tokio::test]
962 async fn session_abort_cancels_branch_summary() -> TestResult {
963 let session = make_session()?;
964 let target = {
965 let mut manager = session.session_manager.lock().await;
966 let target = manager.append_message(&AgentMessage::Llm(Box::new(
967 Message::Assistant(assistant_with_usage("first", Usage::default())),
968 )))?;
969 manager.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(
970 assistant_with_usage("second", Usage::default()),
971 ))))?;
972 target
973 };
974 let summarizer: SummarizeStreamFn = Arc::new(move |_model, _context, options| {
975 let signal = options.signal.clone();
976 Box::pin(async move {
977 if let Some(signal) = signal {
978 signal.cancelled().await;
979 }
980 let mut message = AssistantMessage::new("test-api", "test-provider", "m", 1);
981 message.stop_reason = pi_ai::StopReason::Stop;
982 let events = stream::iter(vec![Ok(AssistantMessageEvent::Done {
983 reason: pi_ai::DoneReason::Stop,
984 message,
985 })]);
986 Box::pin(events)
987 as std::pin::Pin<
988 Box<
989 dyn futures::Stream<Item = Result<AssistantMessageEvent, ProviderError>>
990 + Send,
991 >,
992 >
993 })
994 });
995 let navigation = tokio::spawn({
996 let session = Arc::clone(&session);
997 async move {
998 session
999 .navigate_tree(
1000 &target,
1001 NavigateTreeOptions {
1002 summarize: true,
1003 ..NavigateTreeOptions::default()
1004 },
1005 SummarizationAuth::default(),
1006 Some(&summarizer),
1007 )
1008 .await
1009 }
1010 });
1011 for _ in 0..100 {
1012 if session.is_summarizing() {
1013 break;
1014 }
1015 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
1016 }
1017 assert!(session.is_summarizing());
1018
1019 session.abort().await;
1020 let joined = tokio::time::timeout(std::time::Duration::from_secs(2), navigation).await?;
1021 let result = joined?;
1022 assert!(
1023 matches!(
1024 &result,
1025 Err(TreeError::Summarization(message))
1026 if message.eq_ignore_ascii_case("summarization cancelled")
1027 ),
1028 "expected cancelled summarization, got {result:?}"
1029 );
1030 assert!(!session.is_summarizing());
1031 Ok(())
1032 }
1033
1034 struct ExportTestTool;
1036
1037 impl pi_agent::AgentTool for ExportTestTool {
1038 fn name(&self) -> &'static str {
1039 "exportTestTool"
1040 }
1041 fn label(&self) -> &'static str {
1042 "Export Test Tool"
1043 }
1044 fn description(&self) -> &'static str {
1045 "A test tool for export assertions."
1046 }
1047 fn parameters(&self) -> &serde_json::Value {
1048 static EMPTY: std::sync::LazyLock<serde_json::Value> =
1049 std::sync::LazyLock::new(|| serde_json::Value::Object(serde_json::Map::new()));
1050 &EMPTY
1051 }
1052 fn validate_arguments(
1053 &self,
1054 args: &serde_json::Map<String, serde_json::Value>,
1055 ) -> std::result::Result<serde_json::Map<String, serde_json::Value>, pi_agent::ToolError>
1056 {
1057 Ok(args.clone())
1058 }
1059 fn execute(
1060 &self,
1061 _tool_call_id: &str,
1062 _args: serde_json::Map<String, serde_json::Value>,
1063 _cancel: tokio_util::sync::CancellationToken,
1064 _updates: pi_agent::ToolUpdates,
1065 ) -> Pin<
1066 Box<
1067 dyn Future<
1068 Output = std::result::Result<
1069 pi_agent::AgentToolResult,
1070 pi_agent::ToolError,
1071 >,
1072 > + Send,
1073 >,
1074 > {
1075 Box::pin(async { Ok(pi_agent::AgentToolResult::default()) })
1076 }
1077 }
1078
1079 fn export_test_session(cwd: &str) -> TestResult<Arc<AgentSession>> {
1080 use crate::core::settings::{Settings, SettingsManagerCreateOptions};
1081
1082 let session_manager = SessionManager::create(cwd, None, None)?;
1083 let mut settings_manager = SettingsManager::in_memory(
1084 &Settings::default(),
1085 SettingsManagerCreateOptions {
1086 project_trusted: true,
1087 },
1088 );
1089 settings_manager.set_theme("dark");
1090
1091 let mut config = AgentSessionConfig::test_config(Arc::new(StubProvider), test_model())?;
1092 config.session_manager = session_manager;
1093 config.settings_manager = settings_manager;
1094 config.system_prompt = "EXPORTED SYSTEM PROMPT".into();
1095 config.cwd = cwd.to_owned();
1096 let session = AgentSession::new(config)?;
1097 session.agent.set_tools(vec![
1098 Arc::new(ExportTestTool) as Arc<dyn pi_agent::AgentTool>
1099 ]);
1100 session
1101 .agent
1102 .set_system_prompt("EXPORTED SYSTEM PROMPT".into());
1103 Ok(session)
1104 }
1105
1106 async fn append_export_messages(session: &AgentSession) -> TestResult {
1107 let mut sm = session.session_manager.lock().await;
1108 sm.append_message(&AgentMessage::Llm(Box::new(Message::User(
1109 pi_ai::UserMessage::new(pi_ai::UserMessageContent::Text("hi".into()), 0),
1110 ))))?;
1111 let mut assistant =
1112 AssistantMessage::new("test-api", "test-provider", "m", pi_agent::now_millis());
1113 assistant
1114 .content
1115 .push(AssistantContent::ToolCall(pi_ai::ToolCall::new(
1116 "tc-1",
1117 "customTool",
1118 serde_json::Map::new(),
1119 )));
1120 assistant.stop_reason = pi_ai::StopReason::Stop;
1121 sm.append_message(&AgentMessage::Llm(Box::new(Message::Assistant(assistant))))?;
1122 Ok(())
1123 }
1124
1125 fn export_tool_renderer() -> ToolHtmlPreRenderer {
1126 use crate::core::export_html::RenderedToolHtml;
1127
1128 Arc::new(|_id: String, name: String, _args: serde_json::Value| {
1129 Box::pin(async move {
1130 if name == "customTool" {
1131 Some(RenderedToolHtml {
1132 call_html: Some("<div class='custom-tool'>RENDERED_HTML</div>".into()),
1133 result_html_collapsed: None,
1134 result_html_expanded: None,
1135 })
1136 } else {
1137 None
1138 }
1139 })
1140 })
1141 }
1142
1143 fn decode_export_data(html: &str) -> TestResult<serde_json::Value> {
1144 use base64::Engine as _;
1145
1146 let marker = "<script id=\"session-data\" type=\"application/json\">";
1147 let start = html
1148 .find(marker)
1149 .ok_or_else(|| missing("session-data script marker"))?
1150 + marker.len();
1151 let end = html[start..]
1152 .find("</script>")
1153 .ok_or_else(|| missing("session-data terminator"))?
1154 + start;
1155 let decoded = base64::engine::general_purpose::STANDARD.decode(html[start..end].trim())?;
1156 serde_json::from_slice(&decoded).map_err(Into::into)
1157 }
1158
1159 fn assert_export_data(data: &serde_json::Value) -> TestResult {
1160 assert_eq!(
1161 data["systemPrompt"].as_str(),
1162 Some("EXPORTED SYSTEM PROMPT"),
1163 "systemPrompt should be embedded from agent state"
1164 );
1165 let tools = data
1166 .get("tools")
1167 .and_then(serde_json::Value::as_array)
1168 .ok_or_else(|| missing("tools array should be present"))?;
1169 assert!(
1170 !tools.is_empty(),
1171 "tools should be non-empty from agent state"
1172 );
1173 assert!(
1174 tools.iter().any(|tool| tool["name"] == "exportTestTool"),
1175 "tools should contain exportTestTool"
1176 );
1177 let rendered = data
1178 .get("renderedTools")
1179 .and_then(serde_json::Value::as_object)
1180 .ok_or_else(|| missing("renderedTools should be present"))?;
1181 assert!(
1182 rendered.contains_key("tc-1"),
1183 "renderedTools should contain tc-1, got keys: {:?}",
1184 rendered.keys().collect::<Vec<_>>()
1185 );
1186 let call_html = rendered["tc-1"]["callHtml"]
1187 .as_str()
1188 .ok_or_else(|| missing("callHtml should be present"))?;
1189 assert!(
1190 call_html.contains("RENDERED_HTML"),
1191 "callHtml should contain rendered content"
1192 );
1193 Ok(())
1194 }
1195
1196 #[tokio::test]
1197 async fn export_to_html_embeds_state_tools_theme_rendered_tools() -> TestResult {
1198 let tmp = tempfile::tempdir()?;
1199 let cwd = tmp.path().to_string_lossy().into_owned();
1200 let session = export_test_session(&cwd)?;
1201 append_export_messages(&session).await?;
1202 let output_path = tmp.path().join("session.html");
1203 let output_path = output_path
1204 .to_str()
1205 .ok_or_else(|| missing("temporary export path should be UTF-8"))?;
1206
1207 let output = session
1208 .export_to_html(Some(output_path), Some(export_tool_renderer()))
1209 .await?;
1210 let html = std::fs::read_to_string(output)?;
1211 assert!(
1212 !html.contains("{{SESSION_DATA}}"),
1213 "template placeholder unfilled"
1214 );
1215 let data = decode_export_data(&html)?;
1216 assert_export_data(&data)?;
1217 Ok(())
1218 }
1219}