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::{cap_output_tokens, 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 messages = ContextBuilder::build_with_config(
173                session,
174                &self.system_prompt,
175                Some(&self.config.config.agent),
176            );
177            let tools = Some(tool_definitions(&self.tools));
178            let configured_max = self.config.config.defaults.max_tokens;
179            let context_window = self.config.config.agent.context_window_tokens;
180            let max_tokens =
181                cap_output_tokens(&messages, tools.as_deref(), configured_max, context_window);
182            if max_tokens < configured_max {
183                debug!(
184                    configured_max,
185                    max_tokens, context_window, "max_tokens capped to fit context window"
186                );
187            }
188            let request = ChatRequest {
189                model: model.clone(),
190                messages,
191                tools,
192                temperature: Some(self.config.config.defaults.temperature),
193                max_tokens: Some(max_tokens),
194            };
195
196            debug!(
197                round = rounds,
198                model = %model,
199                provider = %self.provider_name.read().expect("provider lock poisoned"),
200                message_count = request.messages.len(),
201                "agent llm round start"
202            );
203            for (index, msg) in request.messages.iter().enumerate() {
204                debug!(
205                    index,
206                    role = ?msg.role,
207                    tool_calls = msg.tool_calls.as_ref().map(|c| c.len()).unwrap_or(0),
208                    content = %truncate_opt(msg.content.as_deref(), 300),
209                    tool_call_id = ?msg.tool_call_id,
210                    "agent request message"
211                );
212                if let Some(calls) = &msg.tool_calls {
213                    for call in calls {
214                        debug!(
215                            id = %call.id,
216                            name = %call.name,
217                            arguments = %call.arguments,
218                            "agent request tool_call"
219                        );
220                    }
221                }
222            }
223
224            let stream = provider.chat(request).await?;
225            let response = self.collect_stream(stream).await?;
226
227            debug!(
228                round = rounds,
229                content_len = response.content.len(),
230                tool_count = response.tool_calls.len(),
231                "agent stream collected"
232            );
233            if response.tool_calls.is_empty() {
234                debug!(
235                    round = rounds,
236                    content_preview = %truncate(&response.content, 500),
237                    "agent text-only response (no tool calls)"
238                );
239            }
240            for call in &response.tool_calls {
241                debug!(
242                    id = %call.id,
243                    name = %call.name,
244                    arguments = %call.arguments,
245                    "agent tool_call final"
246                );
247            }
248            if response
249                .tool_calls
250                .iter()
251                .any(|c| c.arguments.trim().is_empty() || c.arguments.trim() == "{}")
252            {
253                warn!(
254                    round = rounds,
255                    "agent received tool_call with empty or {{}} arguments"
256                );
257            }
258
259            if let Some(u) = response.usage {
260                match &mut usage {
261                    Some(acc) => acc.add_assign(u),
262                    None => usage = Some(u),
263                }
264            }
265
266            if response.tool_calls.is_empty() {
267                session.push_assistant(response.content, None);
268                store.save(session)?;
269                self.emit(AgentEvent::TurnComplete { usage });
270                self.run_after_turn_hooks(user_input)?;
271                return Ok(TurnOutcome { usage });
272            }
273
274            let records: Vec<ToolCallRecord> = response
275                .tool_calls
276                .iter()
277                .map(|tc| ToolCallRecord {
278                    id: tc.id.clone(),
279                    name: tc.name.clone(),
280                    arguments: tc.arguments.clone(),
281                })
282                .collect();
283            let assistant_content = response.content.clone();
284            session.push_assistant(response.content, Some(records));
285            store.save(session)?;
286
287            for call in &response.tool_calls {
288                let args: serde_json::Value = serde_json::from_str(&call.arguments)
289                    .unwrap_or_else(|_| serde_json::json!({ "raw": call.arguments }));
290                let args = repair_tool_args(&call.name, &assistant_content, args);
291                debug!(name = %call.name, args = %args, "agent tool execute");
292                self.emit(AgentEvent::ToolStarted {
293                    name: call.name.clone(),
294                    args: args.clone(),
295                });
296
297                let result = match self.tools.execute(&self.tool_ctx, &call.name, args).await {
298                    Ok(result) => result,
299                    Err(err) => codei_tools::ToolResult {
300                        content: err.to_string(),
301                        is_error: true,
302                    },
303                };
304                debug!(
305                    name = %call.name,
306                    is_error = result.is_error,
307                    content = %truncate(&result.content, 800),
308                    "agent tool result"
309                );
310                self.emit(AgentEvent::ToolFinished {
311                    name: call.name.clone(),
312                    result: result.clone(),
313                });
314                session.push_tool(&call.id, result.content);
315                store.save(session)?;
316            }
317        }
318    }
319
320    fn run_after_turn_hooks(&self, user_input: &str) -> Result<(), AgentError> {
321        if let Some(root) = &self.config.project_root {
322            let plugins = load_plugins(root);
323            run_hooks(
324                &plugins,
325                HookEvent::AfterTurn,
326                &self.config.cwd,
327                &[("CODEI_PROMPT", user_input.to_string())],
328            )
329            .map_err(AgentError::Config)?;
330        }
331        Ok(())
332    }
333
334    async fn collect_stream(
335        &self,
336        mut stream: codei_llm::ChatStream,
337    ) -> Result<StreamedResponse, AgentError> {
338        let mut content = String::new();
339        let mut usage = None;
340        let mut pending_tools: std::collections::BTreeMap<
341            u32,
342            (Option<String>, Option<String>, String),
343        > = std::collections::BTreeMap::new();
344
345        while let Some(event) = stream.next().await {
346            match event? {
347                StreamEvent::TextDelta(text) => {
348                    self.emit(AgentEvent::AssistantDelta { text: text.clone() });
349                    content.push_str(&text);
350                }
351                StreamEvent::ToolCallDelta {
352                    index,
353                    id,
354                    name,
355                    arguments,
356                } => {
357                    debug!(
358                        index,
359                        id = ?id,
360                        name = ?name,
361                        arguments = ?arguments,
362                        "agent tool_call delta"
363                    );
364                    let entry = pending_tools.entry(index).or_default();
365                    if let Some(id) = id {
366                        entry.0 = Some(id);
367                    }
368                    if let Some(name) = name {
369                        entry.1 = Some(name);
370                    }
371                    if let Some(args) = arguments {
372                        entry.2.push_str(&args);
373                    }
374                }
375                StreamEvent::Usage(u) => usage = Some(u),
376                StreamEvent::Done => {}
377            }
378        }
379
380        let mut tool_calls = Vec::new();
381        for (_, (id, name, arguments)) in pending_tools {
382            if let Some(name) = name {
383                let id = id.unwrap_or_else(|| {
384                    warn!(
385                        name = %name,
386                        "tool call missing id; using synthetic id (function calling mode)"
387                    );
388                    format!("call_{name}")
389                });
390                tool_calls.push(ToolCall {
391                    id,
392                    name,
393                    arguments,
394                });
395            }
396        }
397
398        Ok(StreamedResponse {
399            content,
400            tool_calls,
401            usage,
402        })
403    }
404
405    fn emit(&self, event: AgentEvent) {
406        if let Some(tx) = &self.events {
407            let _ = tx.send(event);
408        }
409    }
410}
411
412struct StreamedResponse {
413    content: String,
414    tool_calls: Vec<ToolCall>,
415    usage: Option<Usage>,
416}
417
418pub(crate) struct AgentParts {
419    pub config: Arc<ResolvedConfig>,
420    pub model: Arc<RwLock<String>>,
421    pub provider: Arc<dyn LlmProvider>,
422    pub provider_name: String,
423    pub tool_ctx: ToolContext,
424    pub tools: ToolRegistry,
425    pub max_tool_rounds: u32,
426    pub system_prompt: String,
427    pub events: Option<UnboundedSender<AgentEvent>>,
428}
429
430fn truncate(value: &str, max: usize) -> String {
431    if value.len() <= max {
432        return value.to_string();
433    }
434    format!(
435        "{}… [truncated, total {} bytes]",
436        &value[..max],
437        value.len()
438    )
439}
440
441fn truncate_opt(value: Option<&str>, max: usize) -> String {
442    match value {
443        Some(text) => truncate(text, max),
444        None => String::from("<none>"),
445    }
446}