Skip to main content

llimorse/
agent.rs

1//! Agent harness around an LLM client.
2
3use 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/// Agent harness around an LLM client.
25#[derive(Debug)]
26pub struct Agent {
27    /// LLM client
28    client: Client,
29
30    /// Current chat history
31    history: ChatHistory,
32
33    /// Pending tool calls the LLM is waiting for
34    pending_calls: Vec<ToolCall>,
35
36    /// Available tools
37    tools: HashMap<String, Box<dyn Tool>>,
38
39    /// Tool definition as passed to the LLM
40    tool_definitions: Vec<ToolDefinition>,
41}
42
43/// Wrapper for the chat history to ensure that everything that is added goes through the
44/// `listener` as well.
45#[derive(Debug)]
46struct ChatHistory {
47    /// The full current chat history
48    history: Vec<ChatMessage>,
49
50    /// Push every single new `ChatMessage` through here
51    listener: Box<dyn ChatListener>,
52}
53
54/// Trait for an object listening to chat updates
55pub trait ChatListener: fmt::Debug + Send {
56    /// The given chat message is added to the chat history; log it.
57    fn log_message(&mut self, message: &ChatMessage);
58
59    /// Log more than a single message at once.
60    ///
61    /// The default implementation calls [`Self::log_message()`] for each.
62    fn log_messages(&mut self, messages: &[ChatMessage]) {
63        for message in messages {
64            self.log_message(message)
65        }
66    }
67}
68
69/// The current operational stage of the agent
70#[derive(Clone, Copy, Debug, Default, Eq, NoUninit, PartialEq)]
71#[repr(u8)]
72pub enum AgentStage {
73    /// Waiting for a prompt
74    #[default]
75    Idle,
76
77    /// Processing a prompt or a tool call result
78    Prefill,
79
80    /// Currently producing reasoning content
81    Reasoning,
82
83    /// Currently producing the actual response output
84    ResponseGeneration,
85
86    /// Currently writing out tool call parameters
87    ToolCallGeneration,
88
89    /// Awaiting tool execution
90    ToolExecution,
91}
92
93impl ChatListener for () {
94    fn log_message(&mut self, _message: &ChatMessage) {}
95    fn log_messages(&mut self, _messages: &[ChatMessage]) {}
96}
97
98/// Tool available to an agent.
99///
100/// This trait is what is needed by [`Agent`] to use a tool.
101pub trait Tool: fmt::Debug + Send {
102    /// Tool name
103    fn name(&self) -> String;
104
105    /// Tool description, if any
106    fn description(&self) -> Option<String>;
107
108    /// Tool arguments JSON schema
109    fn schema(&self) -> Schema;
110
111    /// Execute this tool, arguments given in JSON format (unparsed)
112    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    /// Format the tool call arguments for `Display`
119    fn fmt_call_display(&self, f: &mut fmt::Formatter<'_>, arguments: &str) -> fmt::Result;
120
121    /// Format the tool call result for `Display`
122    fn fmt_call_result_display(&self, f: &mut fmt::Formatter<'_>, result: &str) -> fmt::Result;
123}
124
125/// Connects a tool to its parameter and result types.
126///
127/// This exists so [`CallableTool`] is forced to use the correct types (when using the
128/// [`tool!`](super::tool) macro).  It is separate from [`CallableTool`] because the macro
129/// implements this, and the user implements the latter.
130pub trait ToolState {
131    /// Associated parameter type.
132    type ParamType: for<'a> Deserialize<'a> + fmt::Display + JsonSchema;
133
134    /// Associated return type.
135    type ResultType: for<'a> Deserialize<'a> + Serialize + fmt::Display;
136}
137
138/// User-defined trait for a tool.
139#[allow(async_fn_in_trait)]
140pub trait CallableTool: ToolState {
141    /// Execute a tool call.
142    async fn execute(
143        &self,
144        agent: &Agent,
145        arguments: <Self as ToolState>::ParamType,
146    ) -> Result<<Self as ToolState>::ResultType>;
147}
148
149/// Request currently being executed by the LLM.
150#[pin_project(project = AgentRunningProjection)]
151pub struct AgentRunning<'a, S: Stream<Item = reqwest::Result<bytes::Bytes>>> {
152    /// The results as it is begin generated
153    #[pin]
154    streaming: StreamingResult<S>,
155
156    /// Reference to the [`Agent`] object to push back the results
157    agent: &'a mut Agent,
158
159    /// Whether the request is done
160    terminated: bool,
161}
162
163/// `Future` to await a full agent response without streaming
164#[pin_project]
165pub struct AgentResponse<'a, S: Stream<Item = reqwest::Result<bytes::Bytes>>> {
166    /// Don’t tell anyone, but we actually do still stream, secretly.
167    #[pin]
168    stream: AgentRunning<'a, S>,
169}
170
171impl Agent {
172    /// Create a new agent harness around the given [`Client`].
173    pub fn new(client: Client) -> Self {
174        Self::new_with_listener(client, ())
175    }
176
177    /// Return the entire chat history so far
178    pub fn history(&self) -> &[ChatMessage] {
179        self.history.history()
180    }
181
182    /// Return the last content returned by the LLM after the last input from our side.
183    ///
184    /// Input from our side are:
185    /// - User messages
186    /// - System messages
187    /// - Tool call results
188    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!() // checked in `take_while`
197                };
198                msg.content.as_ref()
199            })
200    }
201}
202
203impl Agent {
204    /// Create a new agent harness around the given [`Client`] with `listener` receiving new chat
205    /// messages.
206    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    /// Push the given message on top of the chat history.
217    pub fn push(&mut self, message: impl Into<ChatMessage>) {
218        self.history.push(message.into());
219    }
220
221    /// Push a full history, e.g. when resuming.
222    ///
223    /// Note that each entry still goes through the [`ChatListener`].
224    pub fn push_history(&mut self, messages: Vec<ChatMessage>) {
225        self.history.push_vec(messages);
226    }
227
228    /// Push the given system-level message on top of the chat history.
229    pub fn push_system(&mut self, message: impl Into<String>) {
230        self.push(SystemMessage::from(message.into()))
231    }
232
233    /// Push the given user message on top of the chat history.
234    pub fn push_user(&mut self, message: impl Into<String>) {
235        self.push(UserMessage::from(message.into()))
236    }
237
238    /// Push the given tool call on top of the chat history, and push it into the pending calls
239    /// list (to execute it via [`Self::execute_pending_calls()`]).
240    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    /// Add the given tool, with the given state, to the agent.
250    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    /// Submit the current chat history.
265    ///
266    /// This includes pending tool call results (from [`Agent::execute_pending_calls()`]) and user
267    /// messages.
268    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    /// Executes all pending tool calls requested by the LLM.
291    ///
292    /// Return whether any have been executed, in which case the results will need to be submitted
293    /// to the LLM.
294    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        // TODO: We need to create an object here that auto-pushes [CANCELED] tool call errors in
305        // case the future is dropped before we get to push real results. It would also need to
306        // have access to `results` and push the already-finished tool call results.
307
308        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    /// Performs the actual tool call, returning a `Result<_>`.
338    ///
339    /// To be usable by the LLM, this needs to be called by something that catches the errors and
340    /// properly formats them for the LLM.
341    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    /// Returns an object that implements [`fmt::Display`] to properly format the call.
357    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    /// Returns an object that implements [`fmt::Display`] to properly format the result.
365    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    /// Return a strong reference to the current client state object
378    pub fn client_state_arc(&self) -> Arc<RwLock<ClientState>> {
379        self.client.state_arc()
380    }
381
382    /// Return the current client state object
383    pub fn client_state(&self) -> impl Deref<Target = ClientState> {
384        self.client.state()
385    }
386
387    /// Return the immutable client information
388    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    /// Same as [`Agent::execute_pending_calls()`].
395    ///
396    /// The problem is that [`AgentRunning`] retains a reference to [`Agent`] while it lives, so
397    /// without dropping it, [`Agent::execute_pending_calls()`] cannot be run.  This function plugs
398    /// that gap, doing both (dropping and executing the calls).
399    ///
400    /// Must only be called after the request has run its course, with success.
401    ///
402    /// # Panics
403    ///
404    /// Panics if [`AgentRunning::is_terminated()`] is false.
405    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    /// Await the full response instead of a stream of parts.
420    pub fn full_response(self) -> AgentResponse<'a, S> {
421        AgentResponse { stream: self }
422    }
423
424    /// Return the agent for this in-progress request
425    pub fn agent(&self) -> &Agent {
426        self.agent
427    }
428
429    /// Abort the incoming transmission, and treat it as finished
430    pub fn force_finalize(&mut self) {
431        self.streaming.force_finalize();
432        self.terminated = true;
433
434        // `force_finalize()` *must* create this message
435        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    /// Mark the stream as terminated.
448    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                // I think the tool calls need to remain in the history, so we cannot just
479                // `.take()` them...?
480                if let Some(ref tool_calls) = message.tool_calls {
481                    this.agent.pending_calls.extend(tool_calls.iter().cloned());
482                }
483
484                // The reasoning, we might want to remove because most models won’t feed it back,
485                // but some do, so... keep it.
486
487                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    /// Same as [`Agent::execute_pending_calls()`].
503    ///
504    /// The problem is that [`AgentResponse`] retains a reference to [`Agent`] while it lives, so
505    /// without dropping it, [`Agent::execute_pending_calls()`] cannot be run.  This function plugs
506    /// that gap, doing both (dropping and executing the calls).
507    ///
508    /// Must only be called after the request has run its course, with success.
509    ///
510    /// # Panics
511    ///
512    /// Panics if `self` has not been awaited yet.
513    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
529/// Helper struct for properly formatting call parameters for display.
530pub struct DisplayCall<'a> {
531    /// Agent; required to parse the call parameters
532    agent: &'a Agent,
533
534    /// Raw call parameters
535    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
556/// Helper struct for properly formatting a call result for display.
557pub struct DisplayCallResult<'a> {
558    /// Agent; required to parse the call result
559    agent: &'a Agent,
560
561    /// Raw call parameters
562    call: &'a ToolCallParams,
563
564    /// Raw tool result
565    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    /// Create a new `ChatHistory` instance, notifying `listener` on additions.
586    fn with_listener(listener: impl ChatListener + 'static) -> Self {
587        ChatHistory {
588            history: Vec::new(),
589            listener: Box::new(listener),
590        }
591    }
592
593    /// Return all chat messages in the history.
594    fn history(&self) -> &[ChatMessage] {
595        &self.history
596    }
597
598    /// Append the given `message` to the history.
599    fn push(&mut self, message: ChatMessage) {
600        self.listener.log_message(&message);
601        self.history.push(message);
602    }
603
604    /// Append the given `messages` to the history.
605    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    /// Return whether the agent is currently processing something (generating or in prefill)
613    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}