diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 16 | ||||
| -rw-r--r-- | src/config/mod.rs | 11 | ||||
| -rw-r--r-- | src/config/session.rs | 22 |
3 files changed, 39 insertions, 10 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index bc58f32..55a4d45 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -29,7 +29,7 @@ pub struct Input { regenerate: bool, medias: Vec<String>, data_urls: HashMap<String, String>, - tool_call: Option<ToolResults>, + tool_results: Option<ToolResults>, rag_name: Option<String>, role: Role, with_session: bool, @@ -48,7 +48,7 @@ impl Input { regenerate: false, medias: Default::default(), data_urls: Default::default(), - tool_call: None, + tool_results: None, rag_name: None, role, with_session, @@ -104,7 +104,7 @@ impl Input { regenerate: false, medias, data_urls, - tool_call: Default::default(), + tool_results: Default::default(), rag_name: None, role, with_session, @@ -120,6 +120,10 @@ impl Input { self.data_urls.clone() } + pub fn tool_results(&self) -> &Option<ToolResults> { + &self.tool_results + } + pub fn text(&self) -> String { match self.patched_text.clone() { Some(text) => text, @@ -184,11 +188,11 @@ impl Input { } pub fn merge_tool_call(mut self, output: String, tool_results: Vec<ToolResult>) -> Self { - match self.tool_call.as_mut() { + match self.tool_results.as_mut() { Some(exist_tool_results) => { exist_tool_results.extend(tool_results, output); } - None => self.tool_call = Some(ToolResults::new(tool_results, output)), + None => self.tool_results = Some(ToolResults::new(tool_results, output)), } self } @@ -228,7 +232,7 @@ impl Input { } else { self.role().build_messages(self) }; - if let Some(tool_results) = &self.tool_call { + if let Some(tool_results) = &self.tool_results { messages.push(Message::new( MessageRole::Assistant, MessageContent::ToolResults(tool_results.clone()), diff --git a/src/config/mod.rs b/src/config/mod.rs index e6dd5c8..9da8dae 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1050,7 +1050,16 @@ impl Config { if let Some(session) = &self.session { let render_options = self.render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; - session.render(&mut markdown_render) + let agent_info: Option<(String, Vec<String>)> = self.agent.as_ref().map(|agent| { + let functions = agent + .functions() + .declarations() + .iter() + .filter_map(|v| if v.agent { Some(v.name.clone()) } else { None }) + .collect(); + (agent.name().to_string(), functions) + }); + session.render(&mut markdown_render, &agent_info) } else { bail!("No session") } diff --git a/src/config/session.rs b/src/config/session.rs index c072d1c..34d6393 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -160,7 +160,11 @@ impl Session { Ok(output) } - pub fn render(&self, render: &mut MarkdownRender) -> Result<String> { + pub fn render( + &self, + render: &mut MarkdownRender, + agent_info: &Option<(String, Vec<String>)>, + ) -> Result<String> { let mut items = vec![]; if let Some(path) = &self.path { @@ -205,7 +209,10 @@ impl Session { for message in &self.messages { match message.role { MessageRole::System => { - lines.push(render.render(&message.content.render_input(resolve_url_fn))); + lines.push( + render + .render(&message.content.render_input(resolve_url_fn, agent_info)), + ); } MessageRole::Assistant => { if let MessageContent::Text(text) = &message.content { @@ -217,9 +224,12 @@ impl Session { lines.push(format!( "{}){}", self.name, - message.content.render_input(resolve_url_fn) + message.content.render_input(resolve_url_fn, agent_info) )); } + MessageRole::Tool => { + lines.push(message.content.render_input(resolve_url_fn, agent_info)); + } } } } @@ -409,6 +419,12 @@ impl Session { .push(Message::new(MessageRole::User, input.message_content())); } self.data_urls.extend(input.data_urls()); + if let Some(tool_results) = input.tool_results() { + self.messages.push(Message::new( + MessageRole::Tool, + MessageContent::ToolResults(tool_results.clone()), + )) + } self.messages.push(Message::new( MessageRole::Assistant, MessageContent::Text(output.to_string()), |
