1use async_openai::types::chat::{
4 ChatCompletionMessageToolCall, ChatCompletionMessageToolCalls,
5 ChatCompletionRequestAssistantMessage, ChatCompletionRequestMessage,
6 ChatCompletionRequestSystemMessage, ChatCompletionRequestToolMessage,
7 ChatCompletionRequestUserMessage, ChatCompletionRequestUserMessageContent,
8 ChatCompletionRequestUserMessageContentPart,
9 ChatCompletionRequestMessageContentPartText,
10 ChatCompletionRequestMessageContentPartImage,
11 FunctionCall,
12};
13
14use async_openai::types::chat::ImageUrl;
16use futures_util::StreamExt;
17use robit_ai::config::ContextConfig;
18use robit_ai::LlmClient;
19use std::any::Any;
20use std::collections::HashMap;
21use std::path::PathBuf;
22use std::sync::Arc;
23use tokio::sync::mpsc;
24
25use crate::context::{ContextManager, TruncationAction, TruncationResult};
26use crate::error::{AgentError, Result};
27use crate::event::{new_session_id, AgentEvent, FrontendMessage, MediaAttachment, SessionId};
28use crate::frontend::Frontend;
29use crate::media;
30use crate::prompt::PromptBuilder;
31use crate::skill::SkillRegistry;
32use crate::tool::async_runner::{AsyncTaskDone, AsyncTaskRunner};
33use crate::tool::task_registry::{AsyncTaskRecord, AsyncTaskStatus, TaskRegistry};
34use crate::tool::{ToolCallInfo, ToolContext, ToolImage, ToolRegistry, ToolResult};
35use tokio_util::sync::CancellationToken;
36
37pub struct AgentSession {
43 pub session_id: SessionId,
44 pub history: Vec<ChatCompletionRequestMessage>,
45 pub working_dir: PathBuf,
46 pub last_known_prompt_tokens: Option<u32>,
52 pub snapshot_message_count: usize,
56}
57
58impl AgentSession {
59 fn new(session_id: SessionId, working_dir: PathBuf, system_prompt: String) -> Self {
60 let system_msg = ChatCompletionRequestMessage::System(
61 ChatCompletionRequestSystemMessage {
62 content: system_prompt.into(),
63 name: None,
64 }
65 .into(),
66 );
67
68 Self {
69 session_id,
70 history: vec![system_msg],
71 working_dir,
72 last_known_prompt_tokens: None,
73 snapshot_message_count: 0,
74 }
75 }
76
77 pub fn with_history(
79 session_id: SessionId,
80 working_dir: PathBuf,
81 system_prompt: String,
82 history: Vec<ChatCompletionRequestMessage>,
83 ) -> Self {
84 let system_msg = ChatCompletionRequestMessage::System(
86 ChatCompletionRequestSystemMessage {
87 content: system_prompt.into(),
88 name: None,
89 }
90 .into(),
91 );
92
93 let mut full_history = vec![system_msg];
95 full_history.extend(history);
96
97 Self {
98 session_id,
99 history: full_history,
100 working_dir,
101 last_known_prompt_tokens: None,
102 snapshot_message_count: 0,
103 }
104 }
105}
106
107pub struct Agent {
113 llm_client: Arc<LlmClient>,
114 tools: Arc<ToolRegistry>,
115 skills: Arc<SkillRegistry>,
116 sessions: HashMap<SessionId, AgentSession>,
117 default_session_id: SessionId,
118 context_manager: ContextManager,
119 frontend: Arc<dyn Frontend>,
120 auto_approve: bool,
121 extensions: HashMap<String, Arc<dyn Any + Send + Sync>>,
123 pending_truncation: Option<(SessionId, crate::context::TruncationResult)>,
125 async_runner: AsyncTaskRunner,
128 done_rx: Option<mpsc::Receiver<AsyncTaskDone>>,
131 pending_tasks: HashMap<String, PendingTask>,
133 task_registry: TaskRegistry,
136}
137
138struct PendingTask {
140 cancel: CancellationToken,
141 tool_name: String,
142}
143
144impl Agent {
145 pub fn new(
147 llm_client: Arc<LlmClient>,
148 tools: Arc<ToolRegistry>,
149 skills: Arc<SkillRegistry>,
150 frontend: Arc<dyn Frontend>,
151 context_config: Option<&ContextConfig>,
152 context_window: Option<u64>,
153 working_dir: PathBuf,
154 auto_approve: bool,
155 extensions: HashMap<String, Arc<dyn Any + Send + Sync>>,
156 ) -> Self {
157 let prompt_builder = PromptBuilder::with_working_dir(Some(&working_dir));
158 let context_manager = ContextManager::new(context_window, context_config);
159
160 let skill_descs = skills.skill_descriptions();
163 let system_prompt = prompt_builder.build_system_prompt(&skill_descs, &working_dir);
164
165 let session_id = new_session_id();
167 let session = AgentSession::new(session_id.clone(), working_dir, system_prompt);
168
169 let mut sessions = HashMap::new();
170 sessions.insert(session_id.clone(), session);
171
172 let (done_tx, done_rx) = mpsc::channel::<AsyncTaskDone>(32);
173 let async_runner = AsyncTaskRunner::new(done_tx);
174 let task_registry = TaskRegistry::new();
175
176 Self {
177 llm_client,
178 tools,
179 skills,
180 sessions,
181 default_session_id: session_id,
182 context_manager,
183 frontend,
184 auto_approve,
185 extensions,
186 pending_truncation: None,
187 async_runner,
188 done_rx: Some(done_rx),
189 pending_tasks: HashMap::new(),
190 task_registry,
191 }
192 }
193
194 pub fn with_history(
196 llm_client: Arc<LlmClient>,
197 tools: Arc<ToolRegistry>,
198 skills: Arc<SkillRegistry>,
199 frontend: Arc<dyn Frontend>,
200 context_config: Option<&ContextConfig>,
201 context_window: Option<u64>,
202 working_dir: PathBuf,
203 auto_approve: bool,
204 extensions: HashMap<String, Arc<dyn Any + Send + Sync>>,
205 session_id: SessionId,
206 history: Vec<ChatCompletionRequestMessage>,
207 ) -> Self {
208 tracing::info!(
209 "Agent::with_history: session_id={}, received {} history messages",
210 session_id,
211 history.len()
212 );
213 let prompt_builder = PromptBuilder::with_working_dir(Some(&working_dir));
214 let context_manager = ContextManager::new(context_window, context_config);
215
216 let skill_descs = skills.skill_descriptions();
219 let system_prompt = prompt_builder.build_system_prompt(&skill_descs, &working_dir);
220
221 let mut session = AgentSession::with_history(
223 session_id.clone(),
224 working_dir,
225 system_prompt,
226 history,
227 );
228
229 tracing::debug!(
230 "Agent::with_history: after adding system prompt, session history length = {}",
231 session.history.len()
232 );
233 let supports_images = llm_client.supports_images();
236 sanitize_history_for_model(&mut session.history, supports_images);
237 let truncation_result = context_manager.maybe_truncate(
239 &mut session.history,
240 session.last_known_prompt_tokens,
241 session.snapshot_message_count,
242 );
243 if truncation_result.rounds_removed > 0 {
244 tracing::info!(
245 "Agent::with_history: truncated {} rounds ({} messages), needs_compression={}",
246 truncation_result.rounds_removed,
247 truncation_result.messages_removed,
248 truncation_result.needs_compression
249 );
250 }
251 tracing::debug!(
252 "Agent::with_history: after truncation, session history length = {}",
253 session.history.len()
254 );
255
256 let pending_truncation = if truncation_result.needs_compression {
257 Some((session_id.clone(), truncation_result))
258 } else {
259 None
260 };
261
262 let mut sessions = HashMap::new();
263 sessions.insert(session_id.clone(), session);
264
265 let (done_tx, done_rx) = mpsc::channel::<AsyncTaskDone>(32);
266 let async_runner = AsyncTaskRunner::new(done_tx);
267 let task_registry = TaskRegistry::new();
268
269 Self {
270 llm_client,
271 tools,
272 skills,
273 sessions,
274 default_session_id: session_id,
275 context_manager,
276 frontend,
277 auto_approve,
278 extensions,
279 pending_truncation,
280 async_runner,
281 done_rx: Some(done_rx),
282 pending_tasks: HashMap::new(),
283 task_registry,
284 }
285 }
286
287 pub async fn run(mut self, mut message_rx: mpsc::Receiver<FrontendMessage>) {
290 tracing::info!("Agent started, session: {}", self.default_session_id);
291
292 if self.pending_truncation.is_some() {
295 tracing::info!("=== Starting pending compression processing ===");
296 let session_id = self.default_session_id.clone();
297 let mut iterations = 0;
298 const MAX_COMPRESSION_ITERATIONS: usize = 20;
299
300 loop {
301 let pending = self.pending_truncation.take();
303 let result = match pending {
304 Some((_, r)) => r,
305 None => break,
306 };
307
308 iterations += 1;
309 if iterations > MAX_COMPRESSION_ITERATIONS {
310 tracing::warn!("Reached max compression iterations ({}), stopping", MAX_COMPRESSION_ITERATIONS);
311 break;
312 }
313
314 tracing::info!("Compression iteration {}: action={:?}, removed_rounds={}, removed_msgs={}",
315 iterations, result.action, result.rounds_removed, result.messages_removed);
316
317 if let Some(session) = self.sessions.get_mut(&session_id) {
319 apply_compression_result(&self.llm_client, &mut session.history, &result).await;
320 session.last_known_prompt_tokens = None;
322 session.snapshot_message_count = 0;
323 }
324
325 let needs_more = if let Some(session) = self.sessions.get(&session_id) {
327 let estimated = self.context_manager.estimate_context_tokens(
328 &session.history,
329 session.last_known_prompt_tokens,
330 session.snapshot_message_count,
331 );
332 estimated > self.context_manager.truncation_threshold()
333 } else {
334 false
335 };
336
337 if !needs_more {
338 tracing::info!("Context below threshold after {} compression iterations", iterations);
339 break;
340 }
341
342 if let Some(session) = self.sessions.get_mut(&session_id) {
344 let next_result = self.context_manager.maybe_truncate(
345 &mut session.history,
346 session.last_known_prompt_tokens,
347 session.snapshot_message_count,
348 );
349 if next_result.needs_compression {
350 self.pending_truncation = Some((session_id.clone(), next_result));
351 } else if next_result.messages_removed > 0 {
352 tracing::info!("Truncation without compression: {} messages removed", next_result.messages_removed);
354 self.pending_truncation = Some((session_id.clone(), next_result));
356 } else {
357 break;
358 }
359 }
360 }
361
362 tracing::info!("=== Compression processing finished ({} iterations) ===", iterations);
363 } else {
364 tracing::debug!("No pending compression needed");
365 }
366
367 let mut done_rx = self
371 .done_rx
372 .take()
373 .expect("done_rx is consumed exactly once in run()");
374
375 loop {
376 tokio::select! {
377 msg = message_rx.recv() => {
378 let Some(msg) = msg else { break; };
379 match msg {
380 FrontendMessage::UserInput { text, attachments } => {
381 if text == "/exit" || text == "/quit" {
382 break;
383 }
384 if text == "/clear" {
385 self.clear_session();
386 let _ = self
387 .frontend
388 .on_event(AgentEvent::TextDelta(
389 "\n[Conversation history cleared]\n".to_string(),
390 ))
391 .await;
392 let _ = self.frontend.on_event(AgentEvent::TurnComplete).await;
393 continue;
394 }
395
396 if let Some((skill, args)) = self.skills.match_trigger(&text) {
398 let skill = skill.clone();
399 self.run_skill_turn(&skill, &args).await;
400 continue;
401 }
402
403 self.run_turn(&text, attachments).await;
404 }
405 FrontendMessage::Cancel => {
406 self.handle_cancel_all().await;
408 }
409 FrontendMessage::CancelTask { task_id } => {
410 self.handle_cancel_task(&task_id).await;
411 }
412 FrontendMessage::ConfirmationResponse { .. } => {
413 tracing::warn!("Unexpected ConfirmationResponse outside tool confirmation");
416 }
417 }
418 }
419 done = done_rx.recv() => {
420 let Some(done) = done else { break; };
421 self.handle_async_done(done).await;
422 }
423 }
424 }
425
426 if !self.pending_tasks.is_empty() {
432 let remaining = self.pending_tasks.len();
433 tracing::warn!(
434 "[async] Agent exiting with {} pending task(s), cancelling and draining...",
435 remaining
436 );
437 for (_, pending) in self.pending_tasks.drain() {
439 pending.cancel.cancel();
440 }
441 let drain_deadline = tokio::time::Instant::now()
444 + tokio::time::Duration::from_secs(5);
445 while tokio::time::Instant::now() < drain_deadline {
446 match tokio::time::timeout(
447 tokio::time::Duration::from_millis(500),
448 done_rx.recv(),
449 )
450 .await
451 {
452 Ok(Some(done)) => {
453 tracing::info!(
454 "[async] drained result after shutdown: task_id={}, tool={}, cancelled={}",
455 done.task_id, done.tool_name, done.cancelled
456 );
457 self.handle_async_done(done).await;
458 }
459 Ok(None) => {
460 tracing::debug!("[async] done_tx closed during drain");
462 break;
463 }
464 Err(_) => {
465 }
467 }
468 }
469 tracing::info!("[async] drain phase complete");
470 }
471
472 tracing::info!("Agent stopped");
473 }
474
475 async fn run_turn(&mut self, user_input: &str, attachments: Vec<MediaAttachment>) {
477 let session_id = self.default_session_id.clone();
478
479 let user_message = self.build_user_message(user_input, &attachments).await;
481
482 if let Some(session) = self.sessions.get_mut(&session_id) {
484 session.history.push(user_message);
485 }
486
487 self.run_agent_loop(&session_id).await;
489 }
490
491 async fn run_agent_loop(&mut self, session_id: &SessionId) {
495 let max_tool_calls = self.context_manager.max_tool_calls_per_turn;
496 let max_iterations = 20;
497 let mut total_tool_calls = 0usize;
498 for iteration in 0..max_iterations {
499 match self.run_one_step(session_id).await {
500 Ok(0) => {
501 let _ = self.frontend.on_event(AgentEvent::TurnComplete).await;
502 return;
503 }
504 Ok(tool_call_count) => {
505 total_tool_calls += tool_call_count;
506
507 if total_tool_calls >= max_tool_calls {
509 tracing::warn!(
510 "Tool call limit reached: {} >= {} (max_tool_calls_per_turn), forcing turn completion",
511 total_tool_calls,
512 max_tool_calls
513 );
514 let _ = self
515 .frontend
516 .on_event(AgentEvent::TextDelta(
517 format!(
518 "\n\n[Tool call limit reached ({} calls). Please summarize progress and continue in the next message.]\n",
519 total_tool_calls
520 ),
521 ))
522 .await;
523 let _ = self.frontend.on_event(AgentEvent::TurnComplete).await;
524 return;
525 }
526
527 tracing::debug!(
528 "Iteration {}: {} tool calls executed (total: {}/{}), continuing loop",
529 iteration,
530 tool_call_count,
531 total_tool_calls,
532 max_tool_calls
533 );
534 }
535 Err(e) => {
536 let _ = self.frontend.on_event(AgentEvent::Error(e)).await;
537 let _ = self.frontend.on_event(AgentEvent::TurnComplete).await;
538 return;
539 }
540 }
541 }
542
543 let _ = self
545 .frontend
546 .on_event(AgentEvent::Error(AgentError::InternalError(
547 format!("Max iterations reached ({})", max_iterations),
548 )))
549 .await;
550 let _ = self.frontend.on_event(AgentEvent::TurnComplete).await;
551 }
552
553 async fn run_one_step(&mut self, session_id: &SessionId) -> Result<usize> {
556 let session = self
557 .sessions
558 .get_mut(session_id)
559 .ok_or_else(|| AgentError::InternalError("Session not found".to_string()))?;
560
561 let truncation_result = self.context_manager.maybe_truncate(
563 &mut session.history,
564 session.last_known_prompt_tokens,
565 session.snapshot_message_count,
566 );
567
568 if truncation_result.needs_compression {
570 apply_compression_result(&self.llm_client, &mut session.history, &truncation_result).await;
571 session.last_known_prompt_tokens = None;
573 session.snapshot_message_count = 0;
574
575 tracing::info!(
576 "Compression completed: action={:?}, removed_rounds={}",
577 truncation_result.action, truncation_result.rounds_removed
578 );
579 } else if truncation_result.messages_removed > 0 {
580 session.last_known_prompt_tokens = None;
582 session.snapshot_message_count = 0;
583
584 tracing::info!(
585 "Context truncated without compression: {} messages removed",
586 truncation_result.messages_removed
587 );
588 }
589
590 let tools_param = if self.llm_client.supports_tools() {
595 let tool_schemas = self.tools.tool_schemas();
596 if tool_schemas.is_empty() {
597 None
598 } else {
599 Some(tool_schemas)
600 }
601 } else {
602 None
603 };
604
605 let estimated_prompt = self.context_manager.estimate_context_tokens(
607 &session.history,
608 session.last_known_prompt_tokens,
609 session.snapshot_message_count,
610 );
611 let calibration_tag = if session.last_known_prompt_tokens.is_some() {
612 "calibrated"
613 } else {
614 "heuristic"
615 };
616 tracing::info!(
617 "LLM call: ~{} prompt tokens ({}), {} messages",
618 estimated_prompt,
619 calibration_tag,
620 session.history.len(),
621 );
622
623 if !self.llm_client.supports_images() {
627 sanitize_history_for_model(&mut session.history, false);
628 }
629
630 let mut stream = match self
632 .llm_client
633 .chat_stream(session.history.clone(), tools_param)
634 .await
635 {
636 Ok(s) => s,
637 Err(e) => {
638 tracing::error!("LLM chat_stream failed: {:?}", e);
639 return Err(e.into());
640 }
641 };
642 tracing::trace!("LLM stream obtained, starting to collect response");
643
644 let mut full_text = String::new();
646 let mut tool_call_chunks: HashMap<usize, ToolCallAccumulator> = HashMap::new();
647 let mut api_usage: Option<async_openai::types::chat::CompletionUsage> = None;
648
649 let mut chunk_count = 0;
650 while let Some(chunk_result) = stream.next().await {
651 let chunk = match chunk_result {
652 Ok(c) => c,
653 Err(e) => {
654 let llm_error = robit_ai::LlmError::from_openai_error(e);
659 tracing::error!("Stream chunk error: {}", llm_error);
660 return Err(AgentError::LlmError(llm_error));
661 }
662 };
663 chunk_count += 1;
664
665 if let Some(ref usage) = chunk.usage {
667 api_usage = Some(usage.clone());
668 }
669
670 if let Some(choice) = chunk.choices.first() {
671 if let Some(content) = &choice.delta.content {
673 full_text.push_str(content);
674 let _ = self
675 .frontend
676 .on_event(AgentEvent::TextDelta(content.clone()))
677 .await;
678 }
679
680 if let Some(tool_calls) = &choice.delta.tool_calls {
682 for tc in tool_calls {
683 let acc = tool_call_chunks
684 .entry(tc.index as usize)
685 .or_insert_with(ToolCallAccumulator::new);
686
687 if let Some(id) = &tc.id {
688 if !id.is_empty() {
690 acc.id = Some(id.clone());
691 }
692 }
693 if let Some(function) = &tc.function {
694 if let Some(name) = &function.name {
695 if !name.is_empty() {
697 acc.name = Some(name.clone());
698 }
699 }
700 if let Some(args) = &function.arguments {
701 acc.arguments.push_str(args);
702 }
703 }
704 }
705 }
706 }
707 }
708
709 tracing::debug!("Stream collection complete: {} chunks, {} chars of text", chunk_count, full_text.len());
710
711 let assembled_tool_calls: Vec<ChatCompletionMessageToolCall> = {
713 let mut indices: Vec<usize> = tool_call_chunks.keys().cloned().collect();
714 indices.sort();
715 indices
716 .into_iter()
717 .filter_map(|idx| tool_call_chunks.remove(&idx)?.into_tool_call())
718 .collect()
719 };
720
721 let estimated_response = crate::context::estimate_tokens(&full_text);
723 if let Some(ref usage) = api_usage {
724 tracing::info!(
725 "LLM response: API usage = {} prompt + {} completion = {} total tokens. Estimated: ~{} prompt + ~{} response = ~{} total",
726 usage.prompt_tokens,
727 usage.completion_tokens,
728 usage.total_tokens,
729 estimated_prompt,
730 estimated_response,
731 estimated_prompt + estimated_response,
732 );
733 } else {
734 tracing::info!(
735 "LLM response: {} chars, ~{} estimated tokens ({} tool calls). API usage not available from streaming.",
736 full_text.len(),
737 estimated_response,
738 assembled_tool_calls.len(),
739 );
740 }
741
742 if let Some(ref usage) = api_usage {
747 session.last_known_prompt_tokens = Some(usage.prompt_tokens);
748 session.snapshot_message_count = session.history.len();
749 tracing::trace!(
750 "Token calibration updated: prompt_tokens={} at {} messages",
751 usage.prompt_tokens, session.history.len()
752 );
753 }
754
755 downgrade_history_images(&mut session.history);
762
763 let content = if full_text.is_empty() {
765 None
766 } else {
767 Some(full_text.clone().into())
768 };
769 let tool_calls = if assembled_tool_calls.is_empty() {
770 None
771 } else {
772 Some(
773 assembled_tool_calls
774 .clone()
775 .into_iter()
776 .map(ChatCompletionMessageToolCalls::Function)
777 .collect(),
778 )
779 };
780
781 if content.is_some() || tool_calls.is_some() {
783 let assistant_msg = ChatCompletionRequestMessage::Assistant(
784 ChatCompletionRequestAssistantMessage {
785 content,
786 name: None,
787 tool_calls,
788 refusal: None,
789 audio: None,
790 #[allow(deprecated)]
791 function_call: None,
792 }
793 .into(),
794 );
795
796 session.history.push(assistant_msg);
797 } else {
798 tracing::warn!("Not adding empty assistant message to history (no content and no tool calls)");
799 }
800
801 if assembled_tool_calls.is_empty() {
803 return Ok(0);
804 }
805
806 self.execute_tool_calls(session_id, &assembled_tool_calls).await
807 }
808
809 async fn execute_tool_calls(
812 &mut self,
813 session_id: &SessionId,
814 assembled_tool_calls: &[ChatCompletionMessageToolCall],
815 ) -> Result<usize> {
816 let working_dir = {
818 let session = self
819 .sessions
820 .get(session_id)
821 .ok_or_else(|| AgentError::InternalError("Session not found".to_string()))?;
822 session.working_dir.clone()
823 };
824
825 let mut batch_images: Vec<ToolImage> = Vec::new();
828
829 for (tc_idx, tc) in assembled_tool_calls.iter().enumerate() {
831 tracing::info!(
832 "Executing tool [{}/{}]: name='{}', id='{}', args={}",
833 tc_idx + 1,
834 assembled_tool_calls.len(),
835 tc.function.name,
836 tc.id,
837 truncate_for_log(&tc.function.arguments, 80)
838 );
839
840 let tc_info = ToolCallInfo {
841 id: tc.id.clone(),
842 name: tc.function.name.clone(),
843 arguments: tc.function.arguments.clone(),
844 };
845
846 if let Err(e) = self
851 .frontend
852 .on_event(AgentEvent::ToolCallRequested {
853 tool_call_id: tc_info.id.clone(),
854 name: tc_info.name.clone(),
855 arguments: tc_info.arguments.clone(),
856 })
857 .await
858 {
859 tracing::warn!(
860 "[tool] ToolCallRequested delivery FAILED (user feedback may be lost): tool_call_id='{}', name='{}', error={}",
861 tc_info.id,
862 tc_info.name,
863 e
864 );
865 }
866
867 let requires_confirm = self.tools.requires_confirmation(&tc.function.name);
869 let approved = if requires_confirm && !self.auto_approve {
870 tracing::trace!(
871 "[tool] requesting user confirmation: tool_call_id='{}', name='{}'",
872 tc_info.id,
873 tc_info.name
874 );
875 match self.frontend.request_tool_confirmation(&tc_info).await {
876 Ok(approved) => {
877 tracing::trace!(
878 "[tool] confirmation response: tool_call_id='{}', name='{}', approved={}",
879 tc_info.id,
880 tc_info.name,
881 approved
882 );
883 approved
884 }
885 Err(e) => {
886 tracing::warn!(
887 "[tool] confirmation request failed: tool_call_id='{}', name='{}', error={}",
888 tc_info.id,
889 tc_info.name,
890 e
891 );
892 return Err(e);
893 }
894 }
895 } else {
896 tracing::trace!(
897 "[tool] skipping confirmation (requires_confirm={}, auto_approve={})",
898 requires_confirm,
899 self.auto_approve
900 );
901 true
902 };
903
904 let result = if approved {
906 let args: serde_json::Value = serde_json::from_str(&tc.function.arguments)
907 .unwrap_or(serde_json::Value::Null);
908
909 let cancel_token = CancellationToken::new();
913
914 let ctx = ToolContext {
915 working_dir: working_dir.clone(),
916 session_id: session_id.clone(),
917 tool_call_id: tc.id.clone(),
918 frontend: self.frontend.clone(),
919 extensions: self.extensions.clone(),
920 supports_images: self.llm_client.supports_images(),
921 async_runner: self.async_runner.clone(),
922 cancel_token: cancel_token.clone(),
923 task_registry: self.task_registry.clone(),
924 };
925
926 let result = self.tools.execute(&tc.function.name, args, &ctx).await;
927 tracing::trace!(
928 "[tool] execution returned: tool_call_id='{}', name='{}', is_pending={}, is_error={}, content_len={}",
929 tc_info.id,
930 tc_info.name,
931 result.is_pending,
932 result.is_error,
933 result.content.len()
934 );
935
936 if result.is_pending {
941 if let Some(tid) = &result.pending_task_id {
942 tracing::info!(
943 "[async] task submitted: task_id={}, tool={}, tool_call_id={}",
944 tid,
945 tc.function.name,
946 tc.id
947 );
948 self.pending_tasks.insert(
949 tid.clone(),
950 PendingTask {
951 cancel: cancel_token,
952 tool_name: tc.function.name.clone(),
953 },
954 );
955 self.task_registry.register(AsyncTaskRecord {
956 task_id: tid.clone(),
957 tool_name: tc.function.name.clone(),
958 tool_call_id: tc.id.clone(),
959 session_id: session_id.clone(),
960 status: AsyncTaskStatus::Pending,
961 started_at: std::time::Instant::now(),
962 result_summary: None,
963 });
964 } else {
965 tracing::warn!(
966 "[async] tool {} returned is_pending without pending_task_id",
967 tc.function.name
968 );
969 }
970 }
971
972 result
973 } else {
974 tracing::trace!(
975 "[tool] tool call rejected by user: tool_call_id='{}', name='{}'",
976 tc_info.id,
977 tc_info.name
978 );
979 ToolResult::error("User rejected this tool call")
980 };
981
982 let raw_len = result.content.len();
984 let truncated_result = ToolResult {
985 content: self.context_manager.truncate_tool_output(&result.content),
986 is_error: result.is_error,
987 images: result.images.clone(),
988 is_pending: result.is_pending,
989 pending_task_id: result.pending_task_id.clone(),
990 };
991 if truncated_result.content.len() != raw_len {
992 tracing::trace!(
993 "[tool] output truncated: tool_call_id='{}', name='{}', raw_len={}, truncated_len={}",
994 tc_info.id,
995 tc_info.name,
996 raw_len,
997 truncated_result.content.len()
998 );
999 }
1000
1001 if let Err(e) = self
1004 .frontend
1005 .on_event(AgentEvent::ToolCallResult {
1006 tool_call_id: tc.id.clone(),
1007 result: truncated_result.clone(),
1008 })
1009 .await
1010 {
1011 tracing::warn!(
1012 "[tool] ToolCallResult delivery FAILED (user feedback may be lost): tool_call_id='{}', name='{}', error={}",
1013 tc_info.id,
1014 tc_info.name,
1015 e
1016 );
1017 }
1018
1019 let tool_msg = ChatCompletionRequestMessage::Tool(
1021 ChatCompletionRequestToolMessage {
1022 content: truncated_result.content.into(),
1023 tool_call_id: tc.id.clone(),
1024 }
1025 .into(),
1026 );
1027
1028 let session = self
1029 .sessions
1030 .get_mut(session_id)
1031 .ok_or_else(|| AgentError::InternalError("Session not found".to_string()))?;
1032 session.history.push(tool_msg);
1033
1034 batch_images.extend(truncated_result.images);
1037 }
1038
1039 if self.llm_client.supports_images() {
1046 if let Some(image_msg) = build_image_user_message(&batch_images) {
1047 let session = self
1048 .sessions
1049 .get_mut(session_id)
1050 .ok_or_else(|| AgentError::InternalError("Session not found".to_string()))?;
1051 session.history.push(image_msg);
1052 }
1053 }
1054
1055 Ok(assembled_tool_calls.len())
1056 }
1057
1058 fn clear_session(&mut self) {
1060 if let Some(session) = self.sessions.get_mut(&self.default_session_id) {
1061 session.history.truncate(1);
1062 }
1063 }
1064
1065 async fn build_user_message(
1067 &self,
1068 text: &str,
1069 attachments: &[MediaAttachment],
1070 ) -> ChatCompletionRequestMessage {
1071 if self.llm_client.supports_images()
1073 && !attachments.is_empty()
1074 && attachments.iter().any(|a| a.is_image())
1075 {
1076 self.build_multimodal_message(text, attachments)
1077 .await
1078 } else {
1079 let mut full_text = text.to_string();
1081 for attachment in attachments {
1082 full_text = format!("{}\n{}", full_text, attachment.describe());
1083 }
1084 ChatCompletionRequestMessage::User(ChatCompletionRequestUserMessage {
1085 content: full_text.into(),
1086 name: None,
1087 })
1088 }
1089 }
1090
1091 async fn build_multimodal_message(
1093 &self,
1094 text: &str,
1095 attachments: &[MediaAttachment],
1096 ) -> ChatCompletionRequestMessage {
1097 let mut parts = vec![ChatCompletionRequestUserMessageContentPart::Text(
1098 ChatCompletionRequestMessageContentPartText {
1099 text: text.to_string(),
1100 prompt_cache_breakpoint: None,
1101 },
1102 )];
1103
1104 for attachment in attachments {
1106 if attachment.is_image() {
1107 match media::download_and_encode_base64(
1109 &attachment.url,
1110 &attachment.content_type,
1111 self.context_manager.max_image_dimension,
1112 )
1113 .await
1114 {
1115 Ok(encoded) => {
1116 parts.push(ChatCompletionRequestUserMessageContentPart::ImageUrl(
1117 ChatCompletionRequestMessageContentPartImage {
1118 image_url: ImageUrl {
1119 url: encoded.data_url,
1120 detail: None,
1121 },
1122 prompt_cache_breakpoint: None,
1123 },
1124 ));
1125 }
1126 Err(e) => {
1127 tracing::warn!("Failed to encode image: {}", e);
1128 let desc = attachment.describe();
1130 let current_text = match &mut parts[0] {
1131 ChatCompletionRequestUserMessageContentPart::Text(t) => &mut t.text,
1132 _ => unreachable!(),
1133 };
1134 *current_text = format!("{}\n{}", current_text, desc);
1135 }
1136 }
1137 } else {
1138 let desc = attachment.describe();
1140 let current_text = match &mut parts[0] {
1141 ChatCompletionRequestUserMessageContentPart::Text(t) => &mut t.text,
1142 _ => unreachable!(),
1143 };
1144 *current_text = format!("{}\n{}", current_text, desc);
1145 }
1146 }
1147
1148 ChatCompletionRequestMessage::User(ChatCompletionRequestUserMessage {
1149 content: ChatCompletionRequestUserMessageContent::Array(parts),
1150 name: None,
1151 })
1152 }
1153
1154 async fn run_skill_turn(&mut self, skill: &crate::skill::Skill, args: &str) {
1159 let _ = self
1161 .frontend
1162 .on_event(AgentEvent::SkillTriggered {
1163 name: skill.frontmatter.name.clone(),
1164 description: skill.frontmatter.description.clone(),
1165 })
1166 .await;
1167
1168 let session_id = self.default_session_id.clone();
1169
1170 let skill_message = format!(
1172 "## Skill: {}\n\n{}\n\n{}",
1173 skill.frontmatter.name,
1174 skill.frontmatter.description,
1175 skill.content
1176 );
1177
1178 let skill_msg = ChatCompletionRequestMessage::System(
1179 ChatCompletionRequestSystemMessage {
1180 content: skill_message.into(),
1181 name: Some(skill.frontmatter.name.clone()),
1182 }
1183 .into(),
1184 );
1185
1186 if let Some(session) = self.sessions.get_mut(&session_id) {
1187 session.history.push(skill_msg);
1188 }
1189
1190 let user_content = if args.is_empty() {
1192 "(User triggered skill, no additional arguments)".to_string()
1193 } else {
1194 args.to_string()
1195 };
1196
1197 if let Some(session) = self.sessions.get_mut(&session_id) {
1198 session.history.push(ChatCompletionRequestMessage::User(
1199 ChatCompletionRequestUserMessage {
1200 content: user_content.into(),
1201 name: None,
1202 }
1203 .into(),
1204 ));
1205 }
1206
1207 let max_iterations = 20;
1209 let mut completed = false;
1210 for iteration in 0..max_iterations {
1211 match self.run_one_step(&session_id).await {
1212 Ok(tool_call_count) => {
1213 if tool_call_count == 0 {
1214 completed = true;
1215 break;
1216 }
1217 tracing::debug!(
1218 "Skill iteration {}: tool calls executed",
1219 iteration
1220 );
1221 }
1222 Err(e) => {
1223 let _ = self.frontend.on_event(AgentEvent::Error(e)).await;
1224 break;
1225 }
1226 }
1227 }
1228
1229 if !completed {
1230 let _ = self
1231 .frontend
1232 .on_event(AgentEvent::Error(AgentError::InternalError(
1233 format!("Max iterations reached ({})", max_iterations),
1234 )))
1235 .await;
1236 }
1237
1238 let _ = self.frontend.on_event(AgentEvent::TurnComplete).await;
1239
1240 if let Some(session) = self.sessions.get_mut(&session_id) {
1242 let skill_name = skill.frontmatter.name.clone();
1243 session.history.retain(|msg| {
1244 !matches!(
1245 msg,
1246 ChatCompletionRequestMessage::System(s)
1247 if s.name.as_deref() == Some(&skill_name)
1248 )
1249 });
1250 }
1251 }
1252
1253 async fn handle_async_done(&mut self, done: AsyncTaskDone) {
1256 tracing::info!(
1257 "[async] task done: task_id={}, tool={}, session={}, cancelled={}, is_error={}",
1258 done.task_id,
1259 done.tool_name,
1260 done.session_id,
1261 done.cancelled,
1262 done.result.is_error
1263 );
1264
1265 self.pending_tasks.remove(&done.task_id);
1267
1268 let status = if done.cancelled {
1270 AsyncTaskStatus::Cancelled
1271 } else if done.result.is_error {
1272 AsyncTaskStatus::Failed
1273 } else {
1274 AsyncTaskStatus::Completed
1275 };
1276 let summary = summarize_result(&done.result.content);
1277 self.task_registry
1278 .update(&done.task_id, status, Some(summary));
1279
1280 let _ = self
1283 .frontend
1284 .on_event(AgentEvent::AsyncToolCompleted {
1285 task_id: done.task_id.clone(),
1286 tool_call_id: done.tool_call_id.clone(),
1287 result: done.result.clone(),
1288 })
1289 .await;
1290
1291 let session_id = done.session_id.clone();
1295 if !self.sessions.contains_key(&session_id) {
1296 tracing::error!(
1297 "[async] task {} (tool={}) finished but session {} not found; dropping result. \
1298 This means the Agent exited or the session was cleaned up before the task completed. \
1299 Result: {} chars, is_error={}, cancelled={}",
1300 done.task_id, done.tool_name, session_id,
1301 done.result.content.len(), done.result.is_error, done.cancelled
1302 );
1303 return;
1304 }
1305
1306 let notice = format!(
1310 "[后台任务完成通知] task_id={} (工具: {})\n{}",
1311 done.task_id, done.tool_name, done.result.content
1312 );
1313 if let Some(session) = self.sessions.get_mut(&session_id) {
1314 session.history.push(ChatCompletionRequestMessage::User(
1315 ChatCompletionRequestUserMessage {
1316 content: notice.into(),
1317 name: None,
1318 },
1319 ));
1320
1321 if self.llm_client.supports_images() {
1324 if let Some(image_msg) = build_image_user_message(&done.result.images) {
1325 session.history.push(image_msg);
1326 }
1327 }
1328 }
1329
1330 self.run_agent_loop(&session_id).await;
1332 }
1333
1334 async fn handle_cancel_task(&mut self, task_id: &str) {
1337 match self.pending_tasks.remove(task_id) {
1338 Some(pending) => {
1339 tracing::info!(
1340 "[async] cancelling task {} (tool={})",
1341 task_id,
1342 pending.tool_name
1343 );
1344 pending.cancel.cancel();
1345 }
1346 None => {
1347 tracing::warn!("[async] cancel request for unknown task {}", task_id);
1348 }
1349 }
1350 }
1351
1352 async fn handle_cancel_all(&mut self) {
1354 let count = self.pending_tasks.len();
1355 if count == 0 {
1356 tracing::info!("[async] Cancel requested, no pending tasks");
1357 return;
1358 }
1359 tracing::info!("[async] cancelling all {} pending task(s)", count);
1360 for (_, pending) in self.pending_tasks.drain() {
1361 pending.cancel.cancel();
1362 }
1363 }
1364}
1365
1366impl Drop for Agent {
1367 fn drop(&mut self) {
1368 let count = self.pending_tasks.len();
1371 if count > 0 {
1372 tracing::info!(
1373 "[async] Agent dropped, cancelling {} pending task(s)",
1374 count
1375 );
1376 for (_, pending) in self.pending_tasks.drain() {
1377 pending.cancel.cancel();
1378 }
1379 }
1380 }
1381}
1382
1383async fn apply_compression_result(
1395 llm_client: &LlmClient,
1396 history: &mut [ChatCompletionRequestMessage],
1397 result: &TruncationResult,
1398) {
1399 if !result.needs_compression {
1400 return;
1401 }
1402
1403 let pos = result.insert_position;
1404 if pos >= history.len() {
1405 tracing::warn!("Insert position {} out of bounds (history len: {})", pos, history.len());
1406 return;
1407 }
1408
1409 let (content, name) = match &result.action {
1410 TruncationAction::NewSegment => {
1411 let summary = generate_summary(llm_client, &result.removed_messages).await;
1412 (
1413 format!("[Summary: {}]", summary),
1414 "summary_segment".to_string(),
1415 )
1416 }
1417 TruncationAction::MergeSegments { summaries, .. } => {
1418 let merged = merge_summaries(llm_client, summaries).await;
1419 let current_level = crate::context::get_merge_level(&history[pos]);
1421 let name = if current_level == 0 {
1422 "summary_segment".to_string()
1423 } else {
1424 format!("summary_segment_m{}", current_level)
1425 };
1426 (format!("[Summary: {}]", merged), name)
1427 }
1428 TruncationAction::TruncateOnly => return,
1429 };
1430
1431 tracing::info!("Compression applied at position {}: {}", pos, name);
1432
1433 history[pos] = ChatCompletionRequestMessage::User(
1434 ChatCompletionRequestUserMessage {
1435 content: content.into(),
1436 name: Some(name),
1437 }
1438 );
1439}
1440
1441async fn generate_summary(
1443 llm_client: &LlmClient,
1444 removed_messages: &[ChatCompletionRequestMessage],
1445) -> String {
1446 tracing::debug!("Generating summary: removed_messages count = {}", removed_messages.len());
1447 let transcript = crate::context::format_removed_messages_as_transcript(removed_messages);
1448 tracing::debug!("Formatted transcript length: {} characters", transcript.len());
1449
1450 let system_prompt = "Summarize the following conversation transcript in 1-2 concise sentences. Focus on: what the user asked for, what actions were taken, and the outcomes. Be brief and factual.";
1451
1452 let messages = vec![
1453 ChatCompletionRequestMessage::System(
1454 ChatCompletionRequestSystemMessage {
1455 content: system_prompt.into(),
1456 name: None,
1457 }
1458 ),
1459 ChatCompletionRequestMessage::User(
1460 ChatCompletionRequestUserMessage {
1461 content: format!("Conversation transcript:\n\n{}", transcript).into(),
1462 name: None,
1463 }
1464 ),
1465 ];
1466
1467 tracing::info!("Calling LLM to generate summary...");
1468 match llm_client.chat(messages, None).await {
1469 Ok(response) => {
1470 tracing::info!("LLM responded successfully for summary generation");
1471 tracing::debug!("Number of choices in response: {}", response.choices.len());
1472 if let Some(choice) = response.choices.first() {
1473 tracing::debug!("Choice index: 0, has content: {}", choice.message.content.is_some());
1474 if let Some(content) = &choice.message.content {
1475 let summary = content.trim().to_string();
1476 if !summary.is_empty() {
1477 tracing::info!("Successfully generated summary (length: {})", summary.len());
1478 return summary;
1479 }
1480 }
1481 }
1482 tracing::warn!("Summary generation returned empty response, using fallback");
1483 "Conversation history compressed.".to_string()
1484 }
1485 Err(e) => {
1486 tracing::error!("Summary generation failed with error: {}, using fallback", e);
1487 "Conversation history compressed.".to_string()
1488 }
1489 }
1490}
1491
1492async fn merge_summaries(
1494 llm_client: &LlmClient,
1495 summaries: &[String],
1496) -> String {
1497 tracing::info!("Merging {} summary segments...", summaries.len());
1498
1499 let numbered: Vec<String> = summaries
1500 .iter()
1501 .enumerate()
1502 .map(|(i, s)| format!("[{}] {}", i + 1, s))
1503 .collect();
1504 let joined = numbered.join("\n\n");
1505
1506 let system_prompt = "You are given multiple conversation summaries from different time periods, ordered from oldest to newest. Merge them into a single concise summary (2-3 sentences) that preserves all key information.
1507
1508Key points to preserve:
1509- User goals and requests
1510- Important decisions made
1511- Technical context (file paths, APIs, architectures)
1512- Major outcomes and conclusions
1513
1514Do not simply concatenate — synthesize into a coherent narrative.";
1515
1516 let messages = vec![
1517 ChatCompletionRequestMessage::System(
1518 ChatCompletionRequestSystemMessage {
1519 content: system_prompt.into(),
1520 name: None,
1521 }
1522 ),
1523 ChatCompletionRequestMessage::User(
1524 ChatCompletionRequestUserMessage {
1525 content: format!("Summaries to merge:\n\n{}", joined).into(),
1526 name: None,
1527 }
1528 ),
1529 ];
1530
1531 match llm_client.chat(messages, None).await {
1532 Ok(response) => {
1533 if let Some(choice) = response.choices.first() {
1534 if let Some(content) = &choice.message.content {
1535 let summary = content.trim().to_string();
1536 if !summary.is_empty() {
1537 tracing::info!("Successfully merged {} summaries (length: {})", summaries.len(), summary.len());
1538 return summary;
1539 }
1540 }
1541 }
1542 tracing::warn!("Summary merge returned empty response, using fallback");
1543 "Multiple earlier conversation segments merged.".to_string()
1544 }
1545 Err(e) => {
1546 tracing::error!("Summary merge failed with error: {}, using fallback", e);
1547 "Multiple earlier conversation segments merged.".to_string()
1548 }
1549 }
1550}
1551
1552#[derive(Debug)]
1558struct ToolCallAccumulator {
1559 id: Option<String>,
1560 name: Option<String>,
1561 arguments: String,
1562}
1563
1564impl ToolCallAccumulator {
1565 fn new() -> Self {
1566 Self {
1567 id: None,
1568 name: None,
1569 arguments: String::new(),
1570 }
1571 }
1572
1573 fn into_tool_call(self) -> Option<ChatCompletionMessageToolCall> {
1575 let id = self.id?;
1576 let name = self.name?;
1577
1578 tracing::trace!(
1579 "Tool call assembled: id='{}', name='{}', args={}",
1580 id,
1581 name,
1582 truncate_for_log(&self.arguments, 80)
1583 );
1584
1585 Some(ChatCompletionMessageToolCall {
1586 id,
1587 function: FunctionCall {
1588 name,
1589 arguments: self.arguments,
1590 },
1591 })
1592 }
1593}
1594
1595fn truncate_for_log(s: &str, max_chars: usize) -> String {
1599 let char_count = s.chars().count();
1600 if char_count <= max_chars {
1601 s.to_string()
1602 } else {
1603 let preview: String = s.chars().take(max_chars).collect();
1604 format!("{}... ({} chars total)", preview, char_count)
1605 }
1606}
1607
1608fn build_image_user_message(images: &[ToolImage]) -> Option<ChatCompletionRequestMessage> {
1612 if images.is_empty() {
1613 return None;
1614 }
1615 let mut parts = vec![ChatCompletionRequestUserMessageContentPart::Text(
1616 ChatCompletionRequestMessageContentPartText {
1617 text: format!(
1618 "[工具返回的图片] {}",
1619 images
1620 .iter()
1621 .map(|i| i.label.as_str())
1622 .collect::<Vec<_>>()
1623 .join(", ")
1624 ),
1625 prompt_cache_breakpoint: None,
1626 },
1627 )];
1628 for img in images {
1629 parts.push(ChatCompletionRequestUserMessageContentPart::ImageUrl(
1630 ChatCompletionRequestMessageContentPartImage {
1631 image_url: ImageUrl {
1632 url: img.data_url.clone(),
1633 detail: None,
1634 },
1635 prompt_cache_breakpoint: None,
1636 },
1637 ));
1638 }
1639 Some(ChatCompletionRequestMessage::User(ChatCompletionRequestUserMessage {
1640 content: ChatCompletionRequestUserMessageContent::Array(parts),
1641 name: None,
1642 }))
1643}
1644
1645fn sanitize_history_for_model(
1653 history: &mut Vec<ChatCompletionRequestMessage>,
1654 supports_images: bool,
1655) {
1656 if supports_images {
1657 return;
1658 }
1659
1660 let mut sanitized_count = 0usize;
1661 for msg in history.iter_mut() {
1662 if let ChatCompletionRequestMessage::User(user_msg) = msg {
1663 if let ChatCompletionRequestUserMessageContent::Array(parts) = &user_msg.content {
1664 let has_image = parts
1666 .iter()
1667 .any(|p| matches!(p, ChatCompletionRequestUserMessageContentPart::ImageUrl(_)));
1668 if has_image {
1669 let text: String = parts
1671 .iter()
1672 .filter_map(|p| {
1673 if let ChatCompletionRequestUserMessageContentPart::Text(t) = p {
1674 Some(t.text.as_str())
1675 } else {
1676 None
1677 }
1678 })
1679 .collect::<Vec<_>>()
1680 .join("\n");
1681
1682 user_msg.content = ChatCompletionRequestUserMessageContent::Text(text);
1683 sanitized_count += 1;
1684 }
1685 }
1686 }
1687 }
1688
1689 if sanitized_count > 0 {
1690 tracing::info!(
1691 "sanitize_history_for_model: downgraded {} message(s) with image_url to text \
1692 (model does not support images)",
1693 sanitized_count
1694 );
1695 }
1696}
1697
1698fn downgrade_history_images(history: &mut Vec<ChatCompletionRequestMessage>) -> usize {
1711 let mut downgraded = 0usize;
1712 for msg in history.iter_mut() {
1713 if let ChatCompletionRequestMessage::User(user_msg) = msg {
1714 if let ChatCompletionRequestUserMessageContent::Array(parts) = &user_msg.content {
1715 let image_count = parts
1716 .iter()
1717 .filter(|p| {
1718 matches!(
1719 p,
1720 ChatCompletionRequestUserMessageContentPart::ImageUrl(_)
1721 )
1722 })
1723 .count();
1724 if image_count > 0 {
1725 let mut text: String = parts
1726 .iter()
1727 .filter_map(|p| {
1728 if let ChatCompletionRequestUserMessageContentPart::Text(t) = p {
1729 Some(t.text.as_str())
1730 } else {
1731 None
1732 }
1733 })
1734 .collect::<Vec<_>>()
1735 .join("\n");
1736 if !text.is_empty() {
1737 text.push('\n');
1738 }
1739 text.push_str(&format!(
1740 "[历史图片已省略:{image_count} 张图片已发送给模型,为控制请求体积不再重复发送]"
1741 ));
1742 user_msg.content = ChatCompletionRequestUserMessageContent::Text(text);
1743 downgraded += 1;
1744 }
1745 }
1746 }
1747 }
1748
1749 if downgraded > 0 {
1750 tracing::info!(
1751 "downgrade_history_images: replaced {} image message(s) with text placeholders \
1752 (images are sent once, then dropped to keep request bodies small)",
1753 downgraded
1754 );
1755 }
1756 downgraded
1757}
1758
1759fn summarize_result(content: &str) -> String {
1761 const MAX: usize = 500;
1762 let char_count = content.chars().count();
1763 if char_count <= MAX {
1764 content.to_string()
1765 } else {
1766 let truncated: String = content.chars().take(MAX).collect();
1767 format!("{}... (truncated, {} chars total)", truncated, char_count)
1768 }
1769}
1770
1771#[cfg(test)]
1772mod tests {
1773 use super::*;
1774 use crate::event::AgentEvent;
1775 use crate::frontend::Frontend;
1776 use crate::skill::SkillRegistry;
1777 use crate::tool::{Tool, ToolContext};
1778 use async_trait::async_trait;
1779 use robit_ai::config::{ModelConfig, ProviderConfig, RobitConfig};
1780 use serde_json::Value;
1781
1782 struct NoopFrontend;
1784
1785 #[async_trait]
1786 impl Frontend for NoopFrontend {
1787 async fn on_event(&self, _event: AgentEvent) -> Result<()> {
1788 Ok(())
1789 }
1790
1791 async fn request_tool_confirmation(&self, _info: &ToolCallInfo) -> Result<bool> {
1792 Ok(true)
1793 }
1794 }
1795
1796 struct ImageTool;
1799
1800 #[async_trait]
1801 impl Tool for ImageTool {
1802 fn name(&self) -> &str {
1803 "fake_image_tool"
1804 }
1805
1806 fn description(&self) -> &str {
1807 "Returns an image"
1808 }
1809
1810 fn parameters_schema(&self) -> Value {
1811 serde_json::json!({"type": "object", "properties": {}})
1812 }
1813
1814 fn requires_confirmation(&self) -> bool {
1815 false
1816 }
1817
1818 async fn execute(&self, _args: Value, _ctx: &ToolContext) -> Result<ToolResult> {
1819 Ok(ToolResult {
1820 content: "Image file: x.png".to_string(),
1821 is_error: false,
1822 images: vec![ToolImage {
1823 data_url: "data:image/png;base64,Zm9v".to_string(),
1824 label: "x.png".to_string(),
1825 }],
1826 is_pending: false,
1827 pending_task_id: None,
1828 })
1829 }
1830 }
1831
1832 fn vision_llm_client() -> Arc<LlmClient> {
1835 let config = RobitConfig {
1836 default_model: Some("test/vision".to_string()),
1837 providers: HashMap::from([(
1838 "test".to_string(),
1839 ProviderConfig {
1840 name: Some("Test".to_string()),
1841 base_url: "http://127.0.0.1:1".to_string(),
1842 api_key: "sk-test".to_string(),
1843 models: vec![ModelConfig {
1844 id: "vision".to_string(),
1845 name: Some("Vision".to_string()),
1846 context_window: None,
1847 max_output_tokens: None,
1848 temperature: None,
1849 max_tokens: None,
1850 supports_images: Some(true),
1851 supports_tools: Some(true),
1852 }],
1853 },
1854 )]),
1855 app: None,
1856 channels: None,
1857 default_image_model: None,
1858 image_providers: HashMap::new(),
1859 };
1860 Arc::new(LlmClient::from_config(&config, None).unwrap())
1861 }
1862
1863 fn tool_call(id: &str) -> ChatCompletionMessageToolCall {
1864 ChatCompletionMessageToolCall {
1865 id: id.to_string(),
1866 function: FunctionCall {
1867 name: "fake_image_tool".to_string(),
1868 arguments: "{}".to_string(),
1869 },
1870 }
1871 }
1872
1873 #[tokio::test]
1878 async fn parallel_image_tool_results_keep_tool_messages_contiguous() {
1879 let mut tools = ToolRegistry::new();
1880 tools.register(ImageTool);
1881 let mut agent = Agent::new(
1882 vision_llm_client(),
1883 Arc::new(tools),
1884 Arc::new(SkillRegistry::new(vec![], &[])),
1885 Arc::new(NoopFrontend),
1886 None,
1887 None,
1888 PathBuf::from("."),
1889 true,
1890 HashMap::new(),
1891 );
1892
1893 let session_id = agent.default_session_id.clone();
1894 let calls = vec![tool_call("call_0"), tool_call("call_1"), tool_call("call_2")];
1895 let executed = agent.execute_tool_calls(&session_id, &calls).await.unwrap();
1896 assert_eq!(executed, 3);
1897
1898 let session = agent.sessions.get(&session_id).unwrap();
1899 assert_eq!(session.history.len(), 5, "3 tool messages + 1 image user message");
1902 let kinds: Vec<&str> = session
1903 .history
1904 .iter()
1905 .skip(1)
1906 .map(|m| match m {
1907 ChatCompletionRequestMessage::Tool(_) => "tool",
1908 ChatCompletionRequestMessage::User(_) => "user",
1909 ChatCompletionRequestMessage::Assistant(_) => "assistant",
1910 _ => "other",
1911 })
1912 .collect();
1913 assert_eq!(
1914 kinds,
1915 vec!["tool", "tool", "tool", "user"],
1916 "tool responses must be contiguous after the assistant tool_calls \
1917 message; image user message(s) go after the batch"
1918 );
1919 }
1920
1921 fn image_user_message(text: &str, image_count: usize) -> ChatCompletionRequestMessage {
1924 let mut parts = vec![ChatCompletionRequestUserMessageContentPart::Text(
1925 ChatCompletionRequestMessageContentPartText {
1926 text: text.to_string(),
1927 prompt_cache_breakpoint: None,
1928 },
1929 )];
1930 for _ in 0..image_count {
1931 parts.push(ChatCompletionRequestUserMessageContentPart::ImageUrl(
1932 ChatCompletionRequestMessageContentPartImage {
1933 image_url: ImageUrl {
1934 url: "data:image/png;base64,Zm9v".to_string(),
1935 detail: None,
1936 },
1937 prompt_cache_breakpoint: None,
1938 },
1939 ));
1940 }
1941 ChatCompletionRequestMessage::User(ChatCompletionRequestUserMessage {
1942 content: ChatCompletionRequestUserMessageContent::Array(parts),
1943 name: None,
1944 })
1945 }
1946
1947 fn plain_user_message() -> ChatCompletionRequestMessage {
1948 ChatCompletionRequestMessage::User(ChatCompletionRequestUserMessage {
1949 content: ChatCompletionRequestUserMessageContent::Text("普通消息".to_string()),
1950 name: None,
1951 })
1952 }
1953
1954 fn history_user_text(msg: &ChatCompletionRequestMessage) -> String {
1955 match msg {
1956 ChatCompletionRequestMessage::User(u) => match &u.content {
1957 ChatCompletionRequestUserMessageContent::Text(t) => t.clone(),
1958 _ => panic!("expected Text content after downgrade"),
1959 },
1960 _ => panic!("expected User message"),
1961 }
1962 }
1963
1964 #[test]
1965 fn downgrade_history_images_replaces_base64_with_placeholder() {
1966 let mut history = vec![
1967 image_user_message("[工具返回的图片] a.png, b.png", 2),
1968 plain_user_message(),
1969 ];
1970 let n = downgrade_history_images(&mut history);
1971 assert_eq!(n, 1, "exactly one image-bearing message downgraded");
1972 let text = history_user_text(&history[0]);
1973 assert!(
1974 text.contains("[工具返回的图片] a.png, b.png"),
1975 "original text part must survive, got: {text}"
1976 );
1977 assert!(text.contains("2"), "placeholder should mention image count");
1978 assert!(
1979 !text.contains("base64"),
1980 "no base64 payload may remain in history, got: {text}"
1981 );
1982 }
1983
1984 #[test]
1985 fn downgrade_image_only_message_produces_placeholder_text() {
1986 let mut history = vec![image_user_message("", 1)];
1989 let n = downgrade_history_images(&mut history);
1990 assert_eq!(n, 1);
1991 let text = history_user_text(&history[0]);
1992 assert!(!text.trim().is_empty(), "placeholder text must be non-empty");
1993 }
1994
1995 #[test]
1996 fn downgrade_history_images_keeps_plain_messages_untouched() {
1997 let mut history = vec![plain_user_message(), plain_user_message()];
1998 let n = downgrade_history_images(&mut history);
1999 assert_eq!(n, 0);
2000 assert_eq!(history_user_text(&history[0]), "普通消息");
2001 }
2002
2003 #[test]
2004 fn downgrade_history_images_is_idempotent() {
2005 let mut history = vec![image_user_message("[工具返回的图片] a.png", 1)];
2006 assert_eq!(downgrade_history_images(&mut history), 1);
2007 assert_eq!(
2008 downgrade_history_images(&mut history),
2009 0,
2010 "second pass must find nothing and append no extra note"
2011 );
2012 }
2013}