summaryrefslogtreecommitdiffstats
path: root/src/config/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-27 14:04:50 +0800
committerGitHub <noreply@github.com>2023-11-27 14:04:50 +0800
commit35c75506e2fd94e2285d8a3cb66208d518d5f992 (patch)
tree97fb643bb5e231022442c1cc617f2daeaa152a66 /src/config/mod.rs
parent5bfe95d31110e75e84626598f033805b0ae4326c (diff)
downloadaichat-35c75506e2fd94e2285d8a3cb66208d518d5f992.tar.gz
feat: support vision (#249)
* feat: support vision * clippy * implement vision * resolve data url to local file * add model openai:gpt-4-vision-preview * use newline to concate embeded text files * set max_tokens for gpt-4-vision-preview
Diffstat (limited to 'src/config/mod.rs')
-rw-r--r--src/config/mod.rs39
1 files changed, 21 insertions, 18 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 9c86e4c..08509be 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1,6 +1,8 @@
+mod input;
mod role;
mod session;
+pub use self::input::Input;
use self::role::Role;
use self::session::{Session, TEMP_SESSION_NAME};
@@ -78,7 +80,7 @@ pub struct Config {
#[serde(skip)]
pub model: Model,
#[serde(skip)]
- pub last_message: Option<(String, String)>,
+ pub last_message: Option<(Input, String)>,
#[serde(skip)]
pub temperature: Option<f64>,
}
@@ -200,15 +202,15 @@ impl Config {
Ok(path)
}
- pub fn save_message(&mut self, input: &str, output: &str) -> Result<()> {
- self.last_message = Some((input.to_string(), output.to_string()));
+ pub fn save_message(&mut self, input: Input, output: &str) -> Result<()> {
+ self.last_message = Some((input.clone(), output.to_string()));
if self.dry_run {
return Ok(());
}
if let Some(session) = self.session.as_mut() {
- session.add_message(input, output)?;
+ session.add_message(&input, output)?;
return Ok(());
}
@@ -220,13 +222,14 @@ impl Config {
return Ok(());
}
let timestamp = now();
+ let input_markdown = input.render();
let output = match self.role.as_ref() {
None => {
- format!("# CHAT:[{timestamp}]\n{input}\n--------\n{output}\n--------\n\n",)
+ format!("# CHAT:[{timestamp}]\n{input_markdown}\n--------\n{output}\n--------\n\n",)
}
Some(v) => {
format!(
- "# CHAT:[{timestamp}] ({})\n{input}\n--------\n{output}\n--------\n\n",
+ "# CHAT:[{timestamp}] ({})\n{input_markdown}\n--------\n{output}\n--------\n\n",
v.name,
)
}
@@ -292,23 +295,23 @@ impl Config {
Ok(())
}
- pub fn echo_messages(&self, content: &str) -> String {
+ pub fn echo_messages(&self, input: &Input) -> String {
if let Some(session) = self.session.as_ref() {
- session.echo_messages(content)
+ session.echo_messages(input)
} else if let Some(role) = self.role.as_ref() {
- role.echo_messages(content)
+ role.echo_messages(input)
} else {
- content.to_string()
+ input.render()
}
}
- pub fn build_messages(&self, content: &str) -> Result<Vec<Message>> {
+ pub fn build_messages(&self, input: &Input) -> Result<Vec<Message>> {
let messages = if let Some(session) = self.session.as_ref() {
- session.build_emssages(content)
+ session.build_emssages(input)
} else if let Some(role) = self.role.as_ref() {
- role.build_messages(content)
+ role.build_messages(input)
} else {
- let message = Message::new(content);
+ let message = Message::new(input);
vec![message]
};
Ok(messages)
@@ -586,7 +589,7 @@ impl Config {
Ok(dir) => dir,
Err(_) => return vec![],
};
- match read_dir(&sessions_dir) {
+ match read_dir(sessions_dir) {
Ok(rd) => {
let mut names = vec![];
for entry in rd.flatten() {
@@ -643,8 +646,8 @@ impl Config {
}
}
- pub fn prepare_send_data(&self, content: &str, stream: bool) -> Result<SendData> {
- let messages = self.build_messages(content)?;
+ pub fn prepare_send_data(&self, input: &Input, stream: bool) -> Result<SendData> {
+ let messages = self.build_messages(input)?;
self.model.max_tokens_limit(&messages)?;
Ok(SendData {
messages,
@@ -653,7 +656,7 @@ impl Config {
})
}
- pub fn maybe_print_send_tokens(&self, input: &str) {
+ 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);