Skip to main content

codei_agent/
loop_.rs

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