From 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 14 May 2024 11:16:55 +0800 Subject: refactor: config::Input (#503) --- src/config/input.rs | 84 ++++++++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 77 insertions(+), 7 deletions(-) (limited to 'src/config/input.rs') diff --git a/src/config/input.rs b/src/config/input.rs index 0527b49..20aa755 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,7 +1,9 @@ -use super::role::Role; -use super::session::Session; +use super::{role::Role, session::Session, GlobalConfig}; -use crate::client::{ImageUrl, MessageContent, MessageContentPart, ModelCapabilities}; +use crate::client::{ + init_client, Client, ImageUrl, Message, MessageContent, MessageContentPart, ModelCapabilities, + SendData, +}; use crate::utils::{base64_encode, sha256}; use anyhow::{bail, Context, Result}; @@ -24,6 +26,7 @@ lazy_static! { #[derive(Debug, Clone)] pub struct Input { + config: GlobalConfig, text: String, medias: Vec, data_urls: HashMap, @@ -31,16 +34,22 @@ pub struct Input { } impl Input { - pub fn from_str(text: &str, context: InputContext) -> Self { + pub fn from_str(config: &GlobalConfig, text: &str, context: Option) -> Self { Self { + config: config.clone(), text: text.to_string(), medias: Default::default(), data_urls: Default::default(), - context, + context: context.unwrap_or_else(|| InputContext::from_config(config)), } } - pub fn new(text: &str, files: Vec, context: InputContext) -> Result { + pub fn new( + config: &GlobalConfig, + text: &str, + files: Vec, + context: Option, + ) -> Result { let mut texts = vec![text.to_string()]; let mut medias = vec![]; let mut data_urls = HashMap::new(); @@ -78,10 +87,11 @@ impl Input { } Ok(Self { + config: config.clone(), text: texts.join("\n"), medias, data_urls, - context, + context: context.unwrap_or_else(|| InputContext::from_config(config)), }) } @@ -101,6 +111,61 @@ impl Input { self.text = text; } + pub fn create_client(&self) -> Result> { + init_client(&self.config) + } + + pub fn prepare_send_data(&self, stream: bool) -> Result { + let messages = self.build_messages()?; + self.config.read().model.max_input_tokens_limit(&messages)?; + let (temperature, top_p) = if let Some(session) = self.session(&self.config.read().session) + { + (session.temperature(), session.top_p()) + } else if let Some(role) = self.role() { + (role.temperature, role.top_p) + } else { + let config = self.config.read(); + (config.temperature, config.top_p) + }; + Ok(SendData { + messages, + temperature, + top_p, + stream, + }) + } + + pub fn maybe_print_input_tokens(&self) { + if self.config.read().dry_run { + if let Ok(messages) = self.build_messages() { + let tokens = self.config.read().model.total_tokens(&messages); + println!(">>> This message consumes {tokens} tokens. <<<"); + } + } + } + + pub fn build_messages(&self) -> Result> { + let messages = if let Some(session) = self.session(&self.config.read().session) { + session.build_messages(self) + } else if let Some(role) = self.role() { + role.build_messages(self) + } else { + let message = Message::new(self); + vec![message] + }; + Ok(messages) + } + + pub fn echo_messages(&self) -> String { + if let Some(session) = self.session(&self.config.read().session) { + session.echo_messages(self) + } else if let Some(role) = self.role() { + role.echo_messages(self) + } else { + self.render() + } + } + pub fn role(&self) -> Option<&Role> { self.context.role.as_ref() } @@ -207,6 +272,11 @@ impl InputContext { Self { role, session } } + pub fn from_config(config: &GlobalConfig) -> Self { + let config = config.read(); + InputContext::new(config.role.clone(), config.session.is_some()) + } + pub fn role(role: Role) -> Self { Self { role: Some(role), -- cgit v1.2.3