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