summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/config/mod.rs11
-rw-r--r--src/main.rs22
2 files changed, 11 insertions, 22 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 81fa6f8..bd02ec3 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -211,12 +211,12 @@ impl Default for Config {
pub type GlobalConfig = Arc<RwLock<Config>>;
impl Config {
- pub fn init(working_mode: WorkingMode, model_id: Option<&str>) -> Result<Self> {
+ pub fn init(working_mode: WorkingMode) -> Result<Self> {
let config_path = Self::config_file()?;
let mut config = if !config_path.exists() {
- match model_id {
- Some(v) => Self::load_dynamic(v)?,
- None => {
+ match env::var(get_env_name("platform")) {
+ Ok(v) => Self::load_dynamic(&v)?,
+ Err(_) => {
if *IS_STDOUT_TERMINAL {
create_config_file(&config_path)?;
}
@@ -1563,6 +1563,9 @@ impl Config {
}
fn load_envs(&mut self) {
+ if let Ok(v) = env::var(get_env_name("model")) {
+ self.model_id = v;
+ }
if let Some(v) = read_env_value::<f64>("temperature") {
self.temperature = v;
}
diff --git a/src/main.rs b/src/main.rs
index eefdc41..3c8cf65 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -54,19 +54,10 @@ async fn main() -> Result<()> {
} else {
WorkingMode::Command
};
- let mut model_id = cli.model.clone();
- if model_id.is_none() {
- if let Ok(v) = env::var(get_env_name("model")) {
- model_id = Some(v);
- }
- }
setup_logger(working_mode.is_serve())?;
- let config = Arc::new(RwLock::new(Config::init(
- working_mode,
- model_id.as_deref(),
- )?));
+ let config = Arc::new(RwLock::new(Config::init(working_mode)?));
let highlight = config.read().highlight;
- if let Err(err) = run(config, cli, text, model_id).await {
+ if let Err(err) = run(config, cli, text).await {
let highlight = stderr().is_terminal() && highlight;
render_error(err, highlight);
std::process::exit(1);
@@ -74,12 +65,7 @@ async fn main() -> Result<()> {
Ok(())
}
-async fn run(
- config: GlobalConfig,
- cli: Cli,
- text: Option<String>,
- model_id: Option<String>,
-) -> Result<()> {
+async fn run(config: GlobalConfig, cli: Cli, text: Option<String>) -> Result<()> {
let abort_signal = create_abort_signal();
if let Some(addr) = cli.serve {
@@ -143,7 +129,7 @@ async fn run(
println!("{sessions}");
return Ok(());
}
- if let Some(model_id) = &model_id {
+ if let Some(model_id) = &cli.model {
config.write().set_model(model_id)?;
}
if cli.no_stream {