summaryrefslogtreecommitdiffstats
path: root/src/config.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-08 19:43:11 +0800
committerGitHub <noreply@github.com>2023-03-08 19:43:11 +0800
commit56483e04f0b8c74e9f7f809e1a0035ca876e35f6 (patch)
tree0bb0e5cc584abc2a3e9c31c78a1ccaea24fbb8cd /src/config.rs
parent2539e24fe9eadf69cb20b01d2303a99fa2e8a363 (diff)
downloadaichat-56483e04f0b8c74e9f7f809e1a0035ca876e35f6.tar.gz
feat: add role-specific config (#42)
Now we can set role temperature. For example: ``` - name: shell prompt: > I want you to act as a linux shell expert. I want you to answer only with bash code. Do not write explanations. temperature: 0.4 ```
Diffstat (limited to 'src/config.rs')
-rw-r--r--src/config.rs30
1 files changed, 26 insertions, 4 deletions
diff --git a/src/config.rs b/src/config.rs
index ec897b4..1542f9e 100644
--- a/src/config.rs
+++ b/src/config.rs
@@ -10,9 +10,9 @@ use std::{
use anyhow::{anyhow, Context, Result};
use inquire::{Confirm, Text};
-use serde::Deserialize;
+use serde::{Deserialize, Serialize};
-use crate::utils::now;
+use crate::utils::{emphasis, now};
const CONFIG_FILE_NAME: &str = "config.yaml";
const ROLES_FILE_NAME: &str = "roles.yaml";
@@ -153,7 +153,19 @@ impl Config {
pub fn change_role(&mut self, name: &str) -> String {
match self.find_role(name) {
Some(role) => {
- let output = format!("{}>> {}", role.name, role.prompt.trim());
+ let temperature = match role.temperature {
+ Some(v) => format!("{v}"),
+ None => "null".into(),
+ };
+ let output = format!(
+ "{}: {}\n{}: {}\n{}: {}",
+ emphasis("name"),
+ role.name,
+ emphasis("prompt"),
+ role.prompt.trim(),
+ emphasis("temperature"),
+ temperature
+ );
self.role = Some(role);
output
}
@@ -165,6 +177,7 @@ impl Config {
self.role = Some(Role {
name: TEMP_ROLE_NAME.into(),
prompt: prompt.into(),
+ temperature: self.temperature,
});
}
@@ -178,6 +191,13 @@ impl Config {
})
}
+ pub fn get_temperature(&self) -> Option<f64> {
+ self.role
+ .as_ref()
+ .and_then(|v| v.temperature)
+ .or(self.temperature)
+ }
+
pub fn merge_prompt(&self, content: &str) -> String {
match self.get_prompt() {
Some(prompt) => format!("{}\n{content}", prompt.trim()),
@@ -305,12 +325,14 @@ impl Config {
}
}
-#[derive(Debug, Clone, Deserialize)]
+#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Role {
/// Role name
pub name: String,
/// Prompt text send to ai for setting up a role
pub prompt: String,
+ /// What sampling temperature to use, between 0 and 2
+ pub temperature: Option<f64>,
}
fn create_config_file(config_path: &Path) -> Result<()> {