1use super::client::{Client, ClientInfo, ClientState};
4use super::line_format::{
5 AssistantMessage, ChatMessage, FunctionDefinition, SystemMessage, ToolCall, ToolCallParams,
6 ToolChoiceMode, ToolDefinition, ToolResult, UserMessage,
7};
8use super::streaming_result::{StreamingChunk, StreamingResult};
9use anyhow::{Result, anyhow, bail};
10use bytemuck::NoUninit;
11use futures::stream::{FusedStream, FuturesUnordered};
12use futures::{Stream, StreamExt};
13use parking_lot::RwLock;
14use pin_project::pin_project;
15use schemars::{JsonSchema, Schema};
16use serde::{Deserialize, Serialize};
17use std::collections::HashMap;
18use std::ops::Deref;
19use std::pin::Pin;
20use std::sync::Arc;
21use std::task::{Context, Poll};
22use std::{fmt, mem};
23
24#[derive(Debug)]
26pub struct Agent {
27 client: Client,
29
30 history: ChatHistory,
32
33 pending_calls: Vec<ToolCall>,
35
36 tools: HashMap<String, Box<dyn Tool>>,
38
39 tool_definitions: Vec<ToolDefinition>,
41}
42
43#[derive(Debug)]
46struct ChatHistory {
47 history: Vec<ChatMessage>,
49
50 listener: Box<dyn ChatListener>,
52}
53
54pub trait ChatListener: fmt::Debug + Send {
56 fn log_message(&mut self, message: &ChatMessage);
58
59 fn log_messages(&mut self, messages: &[ChatMessage]) {
63 for message in messages {
64 self.log_message(message)
65 }
66 }
67}
68
69#[derive(Clone, Copy, Debug, Default, Eq, NoUninit, PartialEq)]
71#[repr(u8)]
72pub enum AgentStage {
73 #[default]
75 Idle,
76
77 Prefill,
79
80 Reasoning,
82
83 ResponseGeneration,
85
86 ToolCallGeneration,
88
89 ToolExecution,
91}
92
93impl ChatListener for () {
94 fn log_message(&mut self, _message: &ChatMessage) {}
95 fn log_messages(&mut self, _messages: &[ChatMessage]) {}
96}
97
98pub trait Tool: fmt::Debug + Send {
102 fn name(&self) -> String;
104
105 fn description(&self) -> Option<String>;
107
108 fn schema(&self) -> Schema;
110
111 fn execute_unparsed<'a>(
113 &'a self,
114 agent: &'a Agent,
115 arguments: &'a str,
116 ) -> Pin<Box<dyn Future<Output = Result<String>> + 'a>>;
117
118 fn fmt_call_display(&self, f: &mut fmt::Formatter<'_>, arguments: &str) -> fmt::Result;
120
121 fn fmt_call_result_display(&self, f: &mut fmt::Formatter<'_>, result: &str) -> fmt::Result;
123}
124
125pub trait ToolState {
131 type ParamType: for<'a> Deserialize<'a> + fmt::Display + JsonSchema;
133
134 type ResultType: for<'a> Deserialize<'a> + Serialize + fmt::Display;
136}
137
138#[allow(async_fn_in_trait)]
140pub trait CallableTool: ToolState {
141 async fn execute(
143 &self,
144 agent: &Agent,
145 arguments: <Self as ToolState>::ParamType,
146 ) -> Result<<Self as ToolState>::ResultType>;
147}
148
149#[pin_project(project = AgentRunningProjection)]
151pub struct AgentRunning<'a, S: Stream<Item = reqwest::Result<bytes::Bytes>>> {
152 #[pin]
154 streaming: StreamingResult<S>,
155
156 agent: &'a mut Agent,
158
159 terminated: bool,
161}
162
163#[pin_project]
165pub struct AgentResponse<'a, S: Stream<Item = reqwest::Result<bytes::Bytes>>> {
166 #[pin]
168 stream: AgentRunning<'a, S>,
169}
170
171impl Agent {
172 pub fn new(client: Client) -> Self {
174 Self::new_with_listener(client, ())
175 }
176
177 pub fn history(&self) -> &[ChatMessage] {
179 self.history.history()
180 }
181
182 pub fn last_result(&self) -> Option<&String> {
189 self.history
190 .history()
191 .iter()
192 .rev()
193 .take_while(|msg| matches!(msg, ChatMessage::Assistant(_)))
194 .find_map(|msg| {
195 let ChatMessage::Assistant(msg) = msg else {
196 unreachable!() };
198 msg.content.as_ref()
199 })
200 }
201}
202
203impl Agent {
204 pub fn new_with_listener(client: Client, listener: impl ChatListener + 'static) -> Self {
207 Agent {
208 client,
209 history: ChatHistory::with_listener(listener),
210 pending_calls: Vec::new(),
211 tools: HashMap::new(),
212 tool_definitions: Vec::new(),
213 }
214 }
215
216 pub fn push(&mut self, message: impl Into<ChatMessage>) {
218 self.history.push(message.into());
219 }
220
221 pub fn push_history(&mut self, messages: Vec<ChatMessage>) {
225 self.history.push_vec(messages);
226 }
227
228 pub fn push_system(&mut self, message: impl Into<String>) {
230 self.push(SystemMessage::from(message.into()))
231 }
232
233 pub fn push_user(&mut self, message: impl Into<String>) {
235 self.push(UserMessage::from(message.into()))
236 }
237
238 pub fn push_tool_call(&mut self, call: ToolCall) {
241 self.pending_calls.push(call.clone());
242 self.push(AssistantMessage {
243 content: None,
244 reasoning_content: None,
245 tool_calls: Some(vec![call]),
246 });
247 }
248
249 pub fn add_tool<T: Tool + 'static>(&mut self, state: T) {
251 let definition = ToolDefinition::Function {
252 function: FunctionDefinition {
253 name: state.name(),
254 description: state.description(),
255 parameters: Some(state.schema()),
256 strict: true,
257 },
258 };
259
260 self.tools.insert(state.name(), Box::new(state));
261 self.tool_definitions.push(definition);
262 }
263
264 pub async fn submit(
269 &mut self,
270 ) -> Result<AgentRunning<'_, impl Stream<Item = reqwest::Result<bytes::Bytes>>>> {
271 let streaming = {
272 self.client
273 .chat_stream(
274 self.history.history(),
275 &self.tool_definitions,
276 ToolChoiceMode::Auto,
277 )
278 .await?
279 };
280
281 self.client.state_mut().operation_stage = AgentStage::Prefill;
282
283 Ok(AgentRunning {
284 streaming,
285 agent: self,
286 terminated: false,
287 })
288 }
289
290 pub async fn execute_pending_calls<
295 F1: FnMut(&Agent, &ToolCall) -> Result<()>,
296 F2: FnMut(&Agent, &ToolCall, &Result<String>) -> Result<()>,
297 >(
298 &mut self,
299 mut tool_guard: F1,
300 mut tool_result_guard: F2,
301 ) -> bool {
302 self.client.state_mut().operation_stage = AgentStage::ToolExecution;
303
304 let mut results = Vec::<ChatMessage>::with_capacity(self.pending_calls.len());
309 let mut futs = FuturesUnordered::new();
310 for call in mem::take(&mut self.pending_calls) {
311 if let Err(err) = tool_guard(self, &call) {
312 results.push(ToolResult::rejected(call.id, err).into());
313 } else {
314 futs.push(async { (self.execute_call(&call.call).await, call) })
315 }
316 }
317
318 while let Some((result, call)) = futs.next().await {
319 if let Err(err) = tool_result_guard(self, &call, &result) {
320 results.push(ToolResult::rejected(call.id, err).into());
321 } else {
322 results.push(ToolResult::new(call.id, result).into());
323 }
324 }
325 drop(futs);
326
327 self.client.state_mut().operation_stage = AgentStage::Idle;
328
329 if results.is_empty() {
330 false
331 } else {
332 self.history.push_vec(results);
333 true
334 }
335 }
336
337 async fn execute_call(&self, tool_call: &ToolCallParams) -> Result<String> {
342 match tool_call {
343 ToolCallParams::Function { function } => {
344 let state = self
345 .tools
346 .get(function.name.as_str())
347 .ok_or_else(|| anyhow!("No such function: {}", function.name))?;
348
349 state.execute_unparsed(self, &function.arguments).await
350 }
351
352 ToolCallParams::Custom { custom } => bail!("No such tool: {}", custom.name),
353 }
354 }
355
356 pub fn display_call<'a>(&'a self, tool_call: &'a ToolCallParams) -> impl fmt::Display + 'a {
358 DisplayCall {
359 agent: self,
360 call: tool_call,
361 }
362 }
363
364 pub fn display_call_result<'a>(
366 &'a self,
367 tool_call: &'a ToolCallParams,
368 result: &'a String,
369 ) -> impl fmt::Display + 'a {
370 DisplayCallResult {
371 agent: self,
372 call: tool_call,
373 result,
374 }
375 }
376
377 pub fn client_state_arc(&self) -> Arc<RwLock<ClientState>> {
379 self.client.state_arc()
380 }
381
382 pub fn client_state(&self) -> impl Deref<Target = ClientState> {
384 self.client.state()
385 }
386
387 pub fn client_info(&self) -> &ClientInfo {
389 self.client.info()
390 }
391}
392
393impl<'a, S: Stream<Item = reqwest::Result<bytes::Bytes>>> AgentRunning<'a, S> {
394 pub async fn execute_pending_calls<
406 F1: FnMut(&Agent, &ToolCall) -> Result<()>,
407 F2: FnMut(&Agent, &ToolCall, &Result<String>) -> Result<()>,
408 >(
409 self,
410 tool_guard: F1,
411 tool_result_guard: F2,
412 ) -> bool {
413 assert!(self.terminated);
414 self.agent
415 .execute_pending_calls(tool_guard, tool_result_guard)
416 .await
417 }
418
419 pub fn full_response(self) -> AgentResponse<'a, S> {
421 AgentResponse { stream: self }
422 }
423
424 pub fn agent(&self) -> &Agent {
426 self.agent
427 }
428
429 pub fn force_finalize(&mut self) {
431 self.streaming.force_finalize();
432 self.terminated = true;
433
434 let message = self
436 .streaming
437 .full_message()
438 .expect("Failed to generate any message");
439 if let Some(ref tool_calls) = message.tool_calls {
440 self.agent.pending_calls.extend(tool_calls.iter().cloned());
441 }
442 self.agent.push(message);
443 }
444}
445
446impl<S: Stream<Item = reqwest::Result<bytes::Bytes>>> AgentRunningProjection<'_, '_, S> {
447 fn terminate(&mut self) {
449 *self.terminated = true;
450 }
451}
452
453impl<S: Stream<Item = reqwest::Result<bytes::Bytes>>> Stream for AgentRunning<'_, S> {
454 type Item = Result<StreamingChunk>;
455
456 fn poll_next(
457 self: Pin<&mut Self>,
458 ctx: &mut Context<'_>,
459 ) -> Poll<Option<Result<StreamingChunk>>> {
460 let mut this = self.project();
461
462 match this.streaming.as_mut().poll_next(ctx) {
463 Poll::Pending => Poll::Pending,
464 Poll::Ready(Some(Err(err))) => {
465 this.terminate();
466 Poll::Ready(Some(Err(err)))
467 }
468 Poll::Ready(Some(Ok(chunk))) => Poll::Ready(Some(Ok(chunk))),
469 Poll::Ready(None) => {
470 this.terminate();
471
472 let Some(message) = this.streaming.as_mut().full_message_pinned() else {
473 return Poll::Ready(Some(Err(anyhow!(
474 "Assistant did not generate a complete message"
475 ))));
476 };
477
478 if let Some(ref tool_calls) = message.tool_calls {
481 this.agent.pending_calls.extend(tool_calls.iter().cloned());
482 }
483
484 this.agent.push(message);
488
489 Poll::Ready(None)
490 }
491 }
492 }
493}
494
495impl<S: Stream<Item = reqwest::Result<bytes::Bytes>>> FusedStream for AgentRunning<'_, S> {
496 fn is_terminated(&self) -> bool {
497 self.terminated
498 }
499}
500
501impl<'a, S: Stream<Item = reqwest::Result<bytes::Bytes>>> AgentResponse<'a, S> {
502 pub async fn execute_pending_calls<
514 F1: FnMut(&Agent, &ToolCall) -> Result<()>,
515 F2: FnMut(&Agent, &ToolCall, &Result<String>) -> Result<()>,
516 >(
517 self,
518 tool_guard: F1,
519 tool_result_guard: F2,
520 ) -> bool {
521 assert!(self.stream.is_terminated());
522 self.stream
523 .agent
524 .execute_pending_calls(tool_guard, tool_result_guard)
525 .await
526 }
527}
528
529pub struct DisplayCall<'a> {
531 agent: &'a Agent,
533
534 call: &'a ToolCallParams,
536}
537
538impl fmt::Display for DisplayCall<'_> {
539 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
540 match self.call {
541 ToolCallParams::Function { function } => {
542 let Some(state) = self.agent.tools.get(function.name.as_str()) else {
543 return write!(f, "[unknown function {}]", function.name);
544 };
545
546 write!(f, "[function call] {}(", function.name)?;
547 state.fmt_call_display(f, &function.arguments)?;
548 write!(f, ")")
549 }
550
551 ToolCallParams::Custom { custom } => write!(f, "[unknown tool {}]", custom.name),
552 }
553 }
554}
555
556pub struct DisplayCallResult<'a> {
558 agent: &'a Agent,
560
561 call: &'a ToolCallParams,
563
564 result: &'a String,
566}
567
568impl fmt::Display for DisplayCallResult<'_> {
569 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
570 match self.call {
571 ToolCallParams::Function { function } => {
572 let Some(state) = self.agent.tools.get(function.name.as_str()) else {
573 return write!(f, "[unknown function {}]", function.name);
574 };
575
576 state.fmt_call_result_display(f, self.result)
577 }
578
579 ToolCallParams::Custom { custom } => write!(f, "[unknown tool {}]", custom.name),
580 }
581 }
582}
583
584impl ChatHistory {
585 fn with_listener(listener: impl ChatListener + 'static) -> Self {
587 ChatHistory {
588 history: Vec::new(),
589 listener: Box::new(listener),
590 }
591 }
592
593 fn history(&self) -> &[ChatMessage] {
595 &self.history
596 }
597
598 fn push(&mut self, message: ChatMessage) {
600 self.listener.log_message(&message);
601 self.history.push(message);
602 }
603
604 fn push_vec(&mut self, mut messages: Vec<ChatMessage>) {
606 self.listener.log_messages(&messages);
607 self.history.append(&mut messages);
608 }
609}
610
611impl AgentStage {
612 pub fn is_processing(&self) -> bool {
614 match self {
615 AgentStage::Idle | AgentStage::ToolExecution => false,
616 AgentStage::Prefill
617 | AgentStage::Reasoning
618 | AgentStage::ResponseGeneration
619 | AgentStage::ToolCallGeneration => true,
620 }
621 }
622}