diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 84 | ||||
| -rw-r--r-- | src/config/mod.rs | 68 | ||||
| -rw-r--r-- | src/config/session.rs | 9 |
3 files changed, 85 insertions, 76 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), diff --git a/src/config/mod.rs b/src/config/mod.rs index 4a1867f..cd74e31 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -7,7 +7,7 @@ pub use self::role::{Role, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ - create_client_config, list_client_types, list_models, ClientConfig, Message, Model, SendData, + create_client_config, list_client_types, list_models, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::render::{MarkdownRender, RenderOptions}; @@ -305,7 +305,7 @@ impl Config { Ok(()) } - pub fn get_state(&self) -> State { + pub fn state(&self) -> State { if let Some(session) = &self.session { if session.is_empty() { if self.role.is_some() { @@ -359,28 +359,6 @@ impl Config { } } - pub fn echo_messages(&self, input: &Input) -> String { - if let Some(session) = input.session(&self.session) { - session.echo_messages(input) - } else if let Some(role) = input.role() { - role.echo_messages(input) - } else { - input.render() - } - } - - pub fn build_messages(&self, input: &Input) -> Result<Vec<Message>> { - let messages = if let Some(session) = input.session(&self.session) { - session.build_messages(input) - } else if let Some(role) = input.role() { - role.build_messages(input) - } else { - let message = Message::new(input); - vec![message] - }; - Ok(messages) - } - pub fn set_wrap(&mut self, value: &str) -> Result<()> { if value == "no" { self.wrap = None; @@ -402,7 +380,7 @@ impl Config { None => bail!("No model '{}'", value), Some(model) => { if let Some(session) = self.session.as_mut() { - session.set_model(model.clone())?; + session.set_model(&model); } self.model = model; Ok(()) @@ -625,7 +603,7 @@ impl Config { self.session = Some(Session::new(self, name)); } else { let session = Session::load(name, &session_path)?; - let model_id = session.model().to_string(); + let model_id = session.model_id().to_string(); self.session = Some(session); self.set_model(&model_id)?; } @@ -787,44 +765,6 @@ impl Config { render_prompt(right_prompt, &variables) } - pub fn prepare_send_data(&self, input: &Input, stream: bool) -> Result<SendData> { - let messages = self.build_messages(input)?; - let temperature = if let Some(session) = input.session(&self.session) { - session.temperature() - } else if let Some(role) = input.role() { - role.temperature - } else { - self.temperature - }; - let top_p = if let Some(session) = input.session(&self.session) { - session.top_p() - } else if let Some(role) = input.role() { - role.top_p - } else { - self.top_p - }; - self.model.max_input_tokens_limit(&messages)?; - Ok(SendData { - messages, - temperature, - top_p, - stream, - }) - } - - pub fn input_context(&self) -> InputContext { - InputContext::new(self.role.clone(), self.session.is_some()) - } - - pub fn maybe_print_send_tokens(&self, input: &Input) { - if self.dry_run { - if let Ok(messages) = self.build_messages(input) { - let tokens = self.model.total_tokens(&messages); - println!(">>> This message consumes {tokens} tokens. <<<"); - } - } - } - fn generate_prompt_context(&self) -> HashMap<&str, String> { let mut output = HashMap::new(); output.insert("model", self.model.id()); diff --git a/src/config/session.rs b/src/config/session.rs index 8e11d3f..a458d4e 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -74,7 +74,7 @@ impl Session { &self.name } - pub fn model(&self) -> &str { + pub fn model_id(&self) -> &str { &self.model_id } @@ -112,7 +112,7 @@ impl Session { let (tokens, percent) = self.tokens_and_percent(); let mut data = json!({ "path": self.path, - "model": self.model(), + "model": self.model_id(), }); if let Some(temperature) = self.temperature() { data["temperature"] = temperature.into(); @@ -240,14 +240,13 @@ impl Session { } } - pub fn set_model(&mut self, model: Model) -> Result<()> { + pub fn set_model(&mut self, model: &Model) { let model_id = model.id(); if self.model_id != model_id { self.model_id = model_id; self.dirty = true; } - self.model = model; - Ok(()) + self.model = model.clone(); } pub fn compress(&mut self, prompt: String) { |
