Skip to main content

codei_agent/
loop_.rs

1use std::sync::{Arc, RwLock};
2
3use codei_config::{discover_skills, load_plugins, run_hooks, HookEvent, ResolvedConfig};
4use codei_llm::{create_provider_by_name, ChatRequest, LlmProvider, StreamEvent, ToolCall, Usage};
5use codei_mcp::McpManager;
6use codei_session::{ContextBuilder, Session, SessionStore, ToolCallRecord};
7use codei_tools::{
8    default_registry, register_mcp_tools, tool_definitions, ToolContext, ToolRegistry,
9};
10use futures_util::StreamExt;
11use tokio::sync::mpsc::UnboundedSender;
12use tracing::{debug, warn};
13
14use crate::error::AgentError;
15use crate::event::AgentEvent;
16use crate::prompt::{build_system_prompt, load_project_instructions};
17use crate::task_tool::{TaskDeps, TaskTool};
18use crate::tool_args::repair_tool_args;
19
20#[derive(Debug, Clone, Default)]
21pub struct TurnOutcome {
22    pub usage: Option<Usage>,
23}
24
25pub struct AgentLoop {
26    config: Arc<ResolvedConfig>,
27    model: Arc<RwLock<String>>,
28    provider_name: Arc<RwLock<String>>,
29    provider: Arc<RwLock<Arc<dyn LlmProvider>>>,
30    tools: ToolRegistry,
31    tool_ctx: ToolContext,
32    system_prompt: String,
33    max_tool_rounds: u32,
34    events: Option<UnboundedSender<AgentEvent>>,
35}
36
37impl AgentLoop {
38    pub fn new(
39        config: Arc<ResolvedConfig>,
40        model: Arc<RwLock<String>>,
41        provider: Arc<dyn LlmProvider>,
42        provider_name: String,
43        tool_ctx: ToolContext,
44        mcp: Option<Arc<McpManager>>,
45        events: Option<UnboundedSender<AgentEvent>>,
46    ) -> Self {
47        let project = load_project_instructions(&config);
48        let skills = discover_skills(&config);
49        let system_prompt = build_system_prompt(&config, &project, &skills);
50        let max_tool_rounds = config.config.agent.max_tool_rounds_per_turn;
51        let max_sub_rounds = (max_tool_rounds / 2).clamp(3, 12);
52
53        let mut tools = default_registry(&config);
54        if let Some(ref manager) = mcp {
55            register_mcp_tools(&mut tools, manager);
56        }
57
58        let deps = Arc::new(TaskDeps {
59            config: Arc::clone(&config),
60            model: Arc::clone(&model),
61            provider: Arc::new(RwLock::new(provider.clone())),
62            provider_name: Arc::new(RwLock::new(provider_name.clone())),
63            tool_ctx: tool_ctx.clone(),
64            mcp: mcp.clone(),
65            max_sub_rounds,
66            system_prompt: system_prompt.clone(),
67        });
68        tools.register(Box::new(TaskTool::new(deps)));
69
70        Self {
71            config,
72            model,
73            provider_name: Arc::new(RwLock::new(provider_name)),
74            provider: Arc::new(RwLock::new(provider)),
75            tools,
76            tool_ctx,
77            system_prompt,
78            max_tool_rounds,
79            events,
80        }
81    }
82
83    pub(crate) fn with_tools(parts: AgentParts) -> Self {
84        Self {
85            config: parts.config,
86            model: parts.model,
87            provider_name: Arc::new(RwLock::new(parts.provider_name)),
88            provider: Arc::new(RwLock::new(parts.provider)),
89            tools: parts.tools,
90            tool_ctx: parts.tool_ctx,
91            system_prompt: parts.system_prompt,
92            max_tool_rounds: parts.max_tool_rounds,
93            events: parts.events,
94        }
95    }
96
97    pub fn provider_name(&self) -> Arc<RwLock<String>> {
98        Arc::clone(&self.provider_name)
99    }
100
101    pub fn config(&self) -> &ResolvedConfig {
102        &self.config
103    }
104
105    pub fn model(&self) -> Arc<RwLock<String>> {
106        Arc::clone(&self.model)
107    }
108
109    pub fn provider(&self) -> Arc<RwLock<Arc<dyn LlmProvider>>> {
110        Arc::clone(&self.provider)
111    }
112
113    pub(crate) fn system_prompt(&self) -> &str {
114        &self.system_prompt
115    }
116
117    pub fn set_provider(&self, name: &str) -> Result<(), AgentError> {
118        let provider = create_provider_by_name(&self.config, name)?;
119        *self
120            .provider_name
121            .write()
122            .map_err(|_| AgentError::Stopped("provider lock poisoned".into()))? = name.to_string();
123        *self
124            .provider
125            .write()
126            .map_err(|_| AgentError::Stopped("provider lock poisoned".into()))? = provider;
127        Ok(())
128    }
129
130    pub async fn run_turn(
131        &self,
132        session: &mut Session,
133        user_input: &str,
134        store: &SessionStore,
135    ) -> Result<TurnOutcome, AgentError> {
136        if let Some(root) = &self.config.project_root {
137            let plugins = load_plugins(root);
138            run_hooks(
139                &plugins,
140                HookEvent::BeforeTurn,
141                &self.config.cwd,
142                &[("CODEI_PROMPT", user_input.to_string())],
143            )
144            .map_err(AgentError::Config)?;
145        }
146
147        session.push_user(user_input);
148        store.save(session)?;
149
150        let mut usage: Option<Usage> = None;
151        let mut rounds = 0u32;
152
153        loop {
154            if rounds >= self.max_tool_rounds {
155                return Err(AgentError::MaxToolRounds);
156            }
157            rounds += 1;
158
159            if self.compact_session_if_needed(session, store).await? {
160                debug!(
161                    keep = self.config.config.agent.compaction_keep_messages,
162                    "session auto-compacted with LLM summary"
163                );
164            }
165
166            let model = self.model.read().expect("model lock poisoned").clone();
167            let provider = self
168                .provider
169                .read()
170                .expect("provider lock poisoned")
171                .clone();
172            let request = ChatRequest {
173                model: model.clone(),
174                messages: ContextBuilder::build_with_config(
175                    session,
176                    &self.system_prompt,
177                    Some(&self.config.config.agent),
178                ),
179                tools: Some(tool_definitions(&self.tools)),
180                temperature: Some(self.config.config.defaults.temperature),
181                max_tokens: Some(self.config.config.defaults.max_tokens),
182            };
183
184            debug!(
185                round = rounds,
186                model = %model,
187                provider = %self.provider_name.read().expect("provider lock poisoned"),
188                message_count = request.messages.len(),
189                "agent llm round start"
190            );
191            for (index, msg) in request.messages.iter().enumerate() {
192                debug!(
193                    index,
194                    role = ?msg.role,
195                    tool_calls = msg.tool_calls.as_ref().map(|c| c.len()).unwrap_or(0),
196                    content = %truncate_opt(msg.content.as_deref(), 300),
197                    tool_call_id = ?msg.tool_call_id,
198                    "agent request message"
199                );
200                if let Some(calls) = &msg.tool_calls {
201                    for call in calls {
202                        debug!(
203                            id = %call.id,
204                            name = %call.name,
205                            arguments = %call.arguments,
206                            "agent request tool_call"
207                        );
208                    }
209                }
210            }
211
212            let stream = provider.chat(request).await?;
213            let response = self.collect_stream(stream).await?;
214
215            debug!(
216                round = rounds,
217                content_len = response.content.len(),
218                tool_count = response.tool_calls.len(),
219                "agent stream collected"
220            );
221            if response.tool_calls.is_empty() {
222                debug!(
223                    round = rounds,
224                    content_preview = %truncate(&response.content, 500),
225                    "agent text-only response (no tool calls)"
226                );
227            }
228            for call in &response.tool_calls {
229                debug!(
230                    id = %call.id,
231                    name = %call.name,
232                    arguments = %call.arguments,
233                    "agent tool_call final"
234                );
235            }
236            if response
237                .tool_calls
238                .iter()
239                .any(|c| c.arguments.trim().is_empty() || c.arguments.trim() == "{}")
240            {
241                warn!(
242                    round = rounds,
243                    "agent received tool_call with empty or {{}} arguments"
244                );
245            }
246
247            if let Some(u) = response.usage {
248                match &mut usage {
249                    Some(acc) => acc.add_assign(u),
250                    None => usage = Some(u),
251                }
252            }
253
254            if response.tool_calls.is_empty() {
255                session.push_assistant(response.content, None);
256                store.save(session)?;
257                self.emit(AgentEvent::TurnComplete { usage });
258                self.run_after_turn_hooks(user_input)?;
259                return Ok(TurnOutcome { usage });
260            }
261
262            let records: Vec<ToolCallRecord> = response
263                .tool_calls
264                .iter()
265                .map(|tc| ToolCallRecord {
266                    id: tc.id.clone(),
267                    name: tc.name.clone(),
268                    arguments: tc.arguments.clone(),
269                })
270                .collect();
271            let assistant_content = response.content.clone();
272            session.push_assistant(response.content, Some(records));
273            store.save(session)?;
274
275            for call in &response.tool_calls {
276                let args: serde_json::Value = serde_json::from_str(&call.arguments)
277                    .unwrap_or_else(|_| serde_json::json!({ "raw": call.arguments }));
278                let args = repair_tool_args(&call.name, &assistant_content, args);
279                debug!(name = %call.name, args = %args, "agent tool execute");
280                self.emit(AgentEvent::ToolStarted {
281                    name: call.name.clone(),
282                    args: args.clone(),
283                });
284
285                let result = match self.tools.execute(&self.tool_ctx, &call.name, args).await {
286                    Ok(result) => result,
287                    Err(err) => codei_tools::ToolResult {
288                        content: err.to_string(),
289                        is_error: true,
290                    },
291                };
292                debug!(
293                    name = %call.name,
294                    is_error = result.is_error,
295                    content = %truncate(&result.content, 800),
296                    "agent tool result"
297                );
298                self.emit(AgentEvent::ToolFinished {
299                    name: call.name.clone(),
300                    result: result.clone(),
301                });
302                session.push_tool(&call.id, result.content);
303                store.save(session)?;
304            }
305        }
306    }
307
308    fn run_after_turn_hooks(&self, user_input: &str) -> Result<(), AgentError> {
309        if let Some(root) = &self.config.project_root {
310            let plugins = load_plugins(root);
311            run_hooks(
312                &plugins,
313                HookEvent::AfterTurn,
314                &self.config.cwd,
315                &[("CODEI_PROMPT", user_input.to_string())],
316            )
317            .map_err(AgentError::Config)?;
318        }
319        Ok(())
320    }
321
322    async fn collect_stream(
323        &self,
324        mut stream: codei_llm::ChatStream,
325    ) -> Result<StreamedResponse, AgentError> {
326        let mut content = String::new();
327        let mut usage = None;
328        let mut pending_tools: std::collections::BTreeMap<
329            u32,
330            (Option<String>, Option<String>, String),
331        > = std::collections::BTreeMap::new();
332
333        while let Some(event) = stream.next().await {
334            match event? {
335                StreamEvent::TextDelta(text) => {
336                    self.emit(AgentEvent::AssistantDelta { text: text.clone() });
337                    content.push_str(&text);
338                }
339                StreamEvent::ToolCallDelta {
340                    index,
341                    id,
342                    name,
343                    arguments,
344                } => {
345                    debug!(
346                        index,
347                        id = ?id,
348                        name = ?name,
349                        arguments = ?arguments,
350                        "agent tool_call delta"
351                    );
352                    let entry = pending_tools.entry(index).or_default();
353                    if let Some(id) = id {
354                        entry.0 = Some(id);
355                    }
356                    if let Some(name) = name {
357                        entry.1 = Some(name);
358                    }
359                    if let Some(args) = arguments {
360                        entry.2.push_str(&args);
361                    }
362                }
363                StreamEvent::Usage(u) => usage = Some(u),
364                StreamEvent::Done => {}
365            }
366        }
367
368        let mut tool_calls = Vec::new();
369        for (_, (id, name, arguments)) in pending_tools {
370            if let Some(name) = name {
371                let id = id.unwrap_or_else(|| {
372                    warn!(
373                        name = %name,
374                        "tool call missing id; using synthetic id (function calling mode)"
375                    );
376                    format!("call_{name}")
377                });
378                tool_calls.push(ToolCall {
379                    id,
380                    name,
381                    arguments,
382                });
383            }
384        }
385
386        Ok(StreamedResponse {
387            content,
388            tool_calls,
389            usage,
390        })
391    }
392
393    fn emit(&self, event: AgentEvent) {
394        if let Some(tx) = &self.events {
395            let _ = tx.send(event);
396        }
397    }
398}
399
400struct StreamedResponse {
401    content: String,
402    tool_calls: Vec<ToolCall>,
403    usage: Option<Usage>,
404}
405
406pub(crate) struct AgentParts {
407    pub config: Arc<ResolvedConfig>,
408    pub model: Arc<RwLock<String>>,
409    pub provider: Arc<dyn LlmProvider>,
410    pub provider_name: String,
411    pub tool_ctx: ToolContext,
412    pub tools: ToolRegistry,
413    pub max_tool_rounds: u32,
414    pub system_prompt: String,
415    pub events: Option<UnboundedSender<AgentEvent>>,
416}
417
418fn truncate(value: &str, max: usize) -> String {
419    if value.len() <= max {
420        return value.to_string();
421    }
422    format!(
423        "{}… [truncated, total {} bytes]",
424        &value[..max],
425        value.len()
426    )
427}
428
429fn truncate_opt(value: Option<&str>, max: usize) -> String {
430    match value {
431        Some(text) => truncate(text, max),
432        None => String::from("<none>"),
433    }
434}