summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-06 07:09:14 +0800
committerGitHub <noreply@github.com>2024-09-06 07:09:14 +0800
commit791b6150afd626baabc2d71acc1dc3dcb70ad4d0 (patch)
tree08eb8fe26088f0b0cdb2bf6a021307f3168fc293 /src
parentb16913fec38f604ab045d6f16ddba62481b09a6c (diff)
downloadaichat-791b6150afd626baabc2d71acc1dc3dcb70ad4d0.tar.gz
feat: add config `serve_addr` & env $SERVE_ADDR for specifying serve addr (#839)
Diffstat (limited to 'src')
-rw-r--r--src/config/mod.rs14
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/serve.rs3
3 files changed, 16 insertions, 3 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 3ef68c6..36a3349 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -57,6 +57,8 @@ pub const TEMP_SESSION_NAME: &str = "temp";
const CLIENTS_FIELD: &str = "clients";
+const SERVE_ADDR: &str = "127.0.0.1:8000";
+
const SUMMARIZE_PROMPT: &str =
"Summarize the discussion briefly in 200 words or less to use as a prompt for future context.";
const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: ";
@@ -126,6 +128,8 @@ pub struct Config {
pub left_prompt: Option<String>,
pub right_prompt: Option<String>,
+ pub serve_addr: Option<String>,
+
pub clients: Vec<ClientConfig>,
#[serde(skip)]
@@ -191,6 +195,8 @@ impl Default for Config {
left_prompt: None,
right_prompt: None,
+ serve_addr: None,
+
clients: vec![],
role: None,
@@ -392,6 +398,10 @@ impl Config {
flags
}
+ pub fn serve_addr(&self) -> String {
+ self.serve_addr.clone().unwrap_or_else(|| SERVE_ADDR.into())
+ }
+
pub fn log(is_serve: bool) -> Result<(LevelFilter, Option<PathBuf>)> {
let log_level = env::var(get_env_name("log_level"))
.ok()
@@ -1848,6 +1858,10 @@ impl Config {
if let Some(v) = read_env_value::<String>("right_prompt") {
self.right_prompt = v;
}
+
+ if let Some(v) = read_env_value::<String>("serve_addr") {
+ self.serve_addr = v;
+ }
}
fn load_functions(&mut self) -> Result<()> {
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 0d9353b..eec43e8 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -262,7 +262,7 @@ impl Repl {
None => println!("Usage: .prompt <text>..."),
},
".role" => match args {
- Some(args) => match args.split_once(|c| c == '\n' || c == ' ') {
+ Some(args) => match args.split_once(['\n', ' ']) {
Some((name, text)) => {
let role = self.config.read().retrieve_role(name.trim())?;
let input = Input::from_str(&self.config, text.trim(), Some(role));
diff --git a/src/serve.rs b/src/serve.rs
index dc16d84..59507a4 100644
--- a/src/serve.rs
+++ b/src/serve.rs
@@ -32,7 +32,6 @@ use tokio::{
use tokio_graceful::Shutdown;
use tokio_stream::wrappers::UnboundedReceiverStream;
-const DEFAULT_ADDRESS: &str = "127.0.0.1:8000";
const DEFAULT_MODEL_NAME: &str = "default";
const PLAYGROUND_HTML: &[u8] = include_bytes!("../assets/playground.html");
const ARENA_HTML: &[u8] = include_bytes!("../assets/arena.html");
@@ -50,7 +49,7 @@ pub async fn run(config: GlobalConfig, addr: Option<String>) -> Result<()> {
addr
}
}
- None => DEFAULT_ADDRESS.to_string(),
+ None => config.read().serve_addr(),
};
let server = Arc::new(Server::new(&config));
let listener = TcpListener::bind(&addr).await?;