summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
Diffstat (limited to 'src/config')
-rw-r--r--src/config/mod.rs51
-rw-r--r--src/config/session.rs31
2 files changed, 40 insertions, 42 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 8589a8b..00a0b67 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -6,7 +6,7 @@ use self::session::{Session, TEMP_SESSION_NAME};
use crate::client::{
create_client_config, list_client_types, list_models, ClientConfig, ExtraConfig, Message,
- ModelInfo, OpenAIClient, SendData,
+ Model, OpenAIClient, SendData,
};
use crate::render::{MarkdownRender, RenderOptions};
use crate::utils::{get_env_name, light_theme_from_colorfgbg, now, prompt_op_err};
@@ -41,7 +41,8 @@ const CLIENTS_FIELD: &str = "clients";
#[serde(default)]
pub struct Config {
/// LLM model
- pub model: Option<String>,
+ #[serde(rename(serialize = "model", deserialize = "model"))]
+ pub model_id: Option<String>,
/// GPT temperature, between 0 and 2
#[serde(rename(serialize = "temperature", deserialize = "temperature"))]
pub default_temperature: Option<f64>,
@@ -73,7 +74,7 @@ pub struct Config {
#[serde(skip)]
pub session: Option<Session>,
#[serde(skip)]
- pub model_info: ModelInfo,
+ pub model: Model,
#[serde(skip)]
pub last_message: Option<(String, String)>,
#[serde(skip)]
@@ -83,7 +84,7 @@ pub struct Config {
impl Default for Config {
fn default() -> Self {
Self {
- model: None,
+ model_id: None,
default_temperature: None,
save: true,
highlight: true,
@@ -97,7 +98,7 @@ impl Default for Config {
roles: vec![],
role: None,
session: None,
- model_info: Default::default(),
+ model: Default::default(),
last_message: None,
temperature: None,
}
@@ -135,7 +136,7 @@ impl Config {
config.load_roles()?;
- config.setup_model_info()?;
+ config.setup_model()?;
config.setup_highlight();
config.setup_light_theme()?;
@@ -304,22 +305,22 @@ impl Config {
pub fn set_model(&mut self, value: &str) -> Result<()> {
let models = list_models(self);
- let mut model_info = None;
+ let mut model = None;
let value = value.trim_end_matches(':');
if value.contains(':') {
- if let Some(model) = models.iter().find(|v| v.id() == value) {
- model_info = Some(model.clone());
+ if let Some(found) = models.iter().find(|v| v.id() == value) {
+ model = Some(found.clone());
}
- } else if let Some(model) = models.iter().find(|v| v.client == value) {
- model_info = Some(model.clone());
+ } else if let Some(found) = models.iter().find(|v| v.client_name == value) {
+ model = Some(found.clone());
}
- match model_info {
+ match model {
None => bail!("Unknown model '{}'", value),
- Some(model_info) => {
+ Some(model) => {
if let Some(session) = self.session.as_mut() {
- session.set_model(model_info.clone())?;
+ session.set_model(model.clone())?;
}
- self.model_info = model_info;
+ self.model = model;
Ok(())
}
}
@@ -338,7 +339,7 @@ impl Config {
.clone()
.map_or_else(|| String::from("no"), |v| v.to_string());
let items = vec![
- ("model", self.model_info.id()),
+ ("model", self.model.id()),
("temperature", temperature),
("dry_run", self.dry_run.to_string()),
("save", self.save.to_string()),
@@ -471,18 +472,14 @@ impl Config {
}
self.session = Some(Session::new(
TEMP_SESSION_NAME,
- self.model_info.clone(),
+ self.model.clone(),
self.role.clone(),
));
}
Some(name) => {
let session_path = Self::session_file(name)?;
if !session_path.exists() {
- self.session = Some(Session::new(
- name,
- self.model_info.clone(),
- self.role.clone(),
- ));
+ self.session = Some(Session::new(name, self.model.clone(), self.role.clone()));
} else {
let session = Session::load(name, &session_path)?;
let model = session.model().to_string();
@@ -608,7 +605,7 @@ impl Config {
pub fn prepare_send_data(&self, content: &str, stream: bool) -> Result<SendData> {
let messages = self.build_messages(content)?;
- self.model_info.max_tokens_limit(&messages)?;
+ self.model.max_tokens_limit(&messages)?;
Ok(SendData {
messages,
temperature: self.get_temperature(),
@@ -619,7 +616,7 @@ impl Config {
pub fn maybe_print_send_tokens(&self, input: &str) {
if self.dry_run {
if let Ok(messages) = self.build_messages(input) {
- let tokens = self.model_info.total_tokens(&messages);
+ let tokens = self.model.total_tokens(&messages);
println!(">>> This message consumes {tokens} tokens. <<<");
}
}
@@ -666,8 +663,8 @@ impl Config {
Ok(())
}
- fn setup_model_info(&mut self) -> Result<()> {
- let model = match &self.model {
+ fn setup_model(&mut self) -> Result<()> {
+ let model = match &self.model_id {
Some(v) => v.clone(),
None => {
let models = list_models(self);
@@ -716,7 +713,7 @@ impl Config {
if let Some(model_name) = value.get("model").and_then(|v| v.as_str()) {
if model_name.starts_with("gpt") {
- self.model = Some(format!("{}:{}", OpenAIClient::NAME, model_name));
+ self.model_id = Some(format!("{}:{}", OpenAIClient::NAME, model_name));
}
}
diff --git a/src/config/session.rs b/src/config/session.rs
index 92e8c2a..1aebd64 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -1,5 +1,5 @@
use super::role::Role;
-use super::ModelInfo;
+use super::Model;
use crate::client::{Message, MessageRole};
use crate::render::MarkdownRender;
@@ -14,7 +14,8 @@ pub const TEMP_SESSION_NAME: &str = "temp";
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Session {
- model: String,
+ #[serde(rename(serialize = "model", deserialize = "model"))]
+ model_id: String,
temperature: Option<f64>,
messages: Vec<Message>,
#[serde(skip)]
@@ -26,21 +27,21 @@ pub struct Session {
#[serde(skip)]
pub role: Option<Role>,
#[serde(skip)]
- pub model_info: ModelInfo,
+ pub model: Model,
}
impl Session {
- pub fn new(name: &str, model_info: ModelInfo, role: Option<Role>) -> Self {
+ pub fn new(name: &str, model: Model, role: Option<Role>) -> Self {
let temperature = role.as_ref().and_then(|v| v.temperature);
Self {
- model: model_info.id(),
+ model_id: model.id(),
temperature,
messages: vec![],
name: name.to_string(),
path: None,
dirty: false,
role,
- model_info,
+ model,
}
}
@@ -61,7 +62,7 @@ impl Session {
}
pub fn model(&self) -> &str {
- &self.model
+ &self.model_id
}
pub fn temperature(&self) -> Option<f64> {
@@ -69,7 +70,7 @@ impl Session {
}
pub fn tokens(&self) -> usize {
- self.model_info.total_tokens(&self.messages)
+ self.model.total_tokens(&self.messages)
}
pub fn export(&self) -> Result<String> {
@@ -83,7 +84,7 @@ impl Session {
data["temperature"] = temperature.into();
}
data["total_tokens"] = tokens.into();
- if let Some(max_tokens) = self.model_info.max_tokens {
+ if let Some(max_tokens) = self.model.max_tokens {
data["max_tokens"] = max_tokens.into();
}
if percent != 0.0 {
@@ -103,13 +104,13 @@ impl Session {
items.push(("path", path.to_string()));
}
- items.push(("model", self.model_info.id()));
+ items.push(("model", self.model.id()));
if let Some(temperature) = self.temperature() {
items.push(("temperature", temperature.to_string()));
}
- if let Some(max_tokens) = self.model_info.max_tokens {
+ if let Some(max_tokens) = self.model.max_tokens {
items.push(("max_tokens", max_tokens.to_string()));
}
@@ -143,7 +144,7 @@ impl Session {
pub fn tokens_and_percent(&self) -> (usize, f32) {
let tokens = self.tokens();
- let max_tokens = self.model_info.max_tokens.unwrap_or_default();
+ let max_tokens = self.model.max_tokens.unwrap_or_default();
let percent = if max_tokens == 0 {
0.0
} else {
@@ -164,9 +165,9 @@ impl Session {
self.temperature = value;
}
- pub fn set_model(&mut self, model_info: ModelInfo) -> Result<()> {
- self.model = model_info.id();
- self.model_info = model_info;
+ pub fn set_model(&mut self, model: Model) -> Result<()> {
+ self.model_id = model.id();
+ self.model = model;
Ok(())
}