summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-14 11:16:55 +0800
committerGitHub <noreply@github.com>2024-05-14 11:16:55 +0800
commit5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (patch)
treefe60e7b0be7c4ea40051f40f88362e0487da6355 /src/config
parent154c1e0b4b7fad08094c601893b081581d2ee0c8 (diff)
downloadaichat-5284a18248bb8e48eaa4a1e6ddcf73d944d23783.tar.gz
refactor: config::Input (#503)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs84
-rw-r--r--src/config/mod.rs68
-rw-r--r--src/config/session.rs9
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) {