diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-14 11:16:55 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-14 11:16:55 +0800 |
| commit | 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (patch) | |
| tree | fe60e7b0be7c4ea40051f40f88362e0487da6355 /src/config/input.rs | |
| parent | 154c1e0b4b7fad08094c601893b081581d2ee0c8 (diff) | |
| download | aichat-5284a18248bb8e48eaa4a1e6ddcf73d944d23783.tar.gz | |
refactor: config::Input (#503)
Diffstat (limited to 'src/config/input.rs')
| -rw-r--r-- | src/config/input.rs | 84 |
1 files changed, 77 insertions, 7 deletions
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<String>, data_urls: HashMap<String, String>, @@ -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<InputContext>) -> 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<String>, context: InputContext) -> Result<Self> { + pub fn new( + config: &GlobalConfig, + text: &str, + files: Vec<String>, + context: Option<InputContext>, + ) -> Result<Self> { 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<Box<dyn Client>> { + init_client(&self.config) + } + + pub fn prepare_send_data(&self, stream: bool) -> Result<SendData> { + 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<Vec<Message>> { + 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), |
