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}