summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/model.rs6
-rw-r--r--src/config/mod.rs34
-rw-r--r--src/main.rs48
3 files changed, 53 insertions, 35 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
index acf04a3..17646c9 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -45,21 +45,21 @@ impl Model {
pub fn retrieve_chat(config: &Config, model_id: &str) -> Result<Self> {
match Self::find(&list_chat_models(config), model_id) {
Some(v) => Ok(v),
- None => bail!("Invalid chat model '{model_id}'"),
+ None => bail!("Unknown chat model '{model_id}'"),
}
}
pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result<Self> {
match Self::find(&list_embedding_models(config), model_id) {
Some(v) => Ok(v),
- None => bail!("Invalid embedding model '{model_id}'"),
+ None => bail!("Unknown embedding model '{model_id}'"),
}
}
pub fn retrieve_reranker(config: &Config, model_id: &str) -> Result<Self> {
match Self::find(&list_reranker_models(config), model_id) {
Some(v) => Ok(v),
- None => bail!("Invalid reranker model '{model_id}'"),
+ None => bail!("Unknown reranker model '{model_id}'"),
}
}
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 07d7b33..6d196d6 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -210,17 +210,20 @@ impl Default for Config {
pub type GlobalConfig = Arc<RwLock<Config>>;
impl Config {
- pub fn init(working_mode: WorkingMode) -> Result<Self> {
+ pub fn init(working_mode: WorkingMode, model_id: Option<&str>) -> Result<Self> {
let config_path = Self::config_file()?;
-
- let platform = env::var(get_env_name("platform")).ok();
- if *IS_STDOUT_TERMINAL && platform.is_none() && !config_path.exists() {
- create_config_file(&config_path)?;
- }
- let mut config = if platform.is_some() {
- Self::load_dynamic_config(&platform.unwrap())?
+ let mut config = if !config_path.exists() {
+ match model_id {
+ Some(v) => Self::load_dynamic(v)?,
+ None => {
+ if *IS_STDOUT_TERMINAL {
+ create_config_file(&config_path)?;
+ }
+ Self::load_from_file(&config_path)?
+ }
+ }
} else {
- Self::load_config_file(&config_path)?
+ Self::load_from_file(&config_path)?
};
config.working_mode = working_mode;
@@ -1498,7 +1501,7 @@ impl Config {
.with_context(|| format!("Failed to create/append {}", path.display()))
}
- fn load_config_file(config_path: &Path) -> Result<Self> {
+ fn load_from_file(config_path: &Path) -> Result<Self> {
let content = read_to_string(config_path)
.with_context(|| format!("Failed to load config at {}", config_path.display()))?;
let config: Self = serde_yaml::from_str(&content).map_err(|err| {
@@ -1520,7 +1523,11 @@ impl Config {
Ok(config)
}
- fn load_dynamic_config(platform: &str) -> Result<Self> {
+ fn load_dynamic(model_id: &str) -> Result<Self> {
+ let platform = match model_id.split_once(':') {
+ Some((v, _)) => v,
+ _ => model_id,
+ };
let is_openai_compatible = OPENAI_COMPATIBLE_PLATFORMS
.into_iter()
.any(|(name, _)| platform == name);
@@ -1530,7 +1537,7 @@ impl Config {
json!({ "type": platform })
};
let config = json!({
- "model": platform.to_string(),
+ "model": model_id.to_string(),
"save": false,
"clients": vec![client],
});
@@ -1540,9 +1547,6 @@ 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 50394fb..609b51b 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -46,18 +46,39 @@ async fn main() -> Result<()> {
let cli = Cli::parse();
let text = cli.text();
let text = aggregate_text(text)?;
- let file = &cli.file;
- let no_input = text.is_none() && file.is_empty();
let working_mode = if cli.serve.is_some() {
WorkingMode::Serve
- } else if no_input {
+ } else if text.is_none() && cli.file.is_empty() {
WorkingMode::Repl
} 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)?));
+ let config = Arc::new(RwLock::new(Config::init(
+ working_mode,
+ model_id.as_deref(),
+ )?));
+ let highlight = config.read().highlight;
+ if let Err(err) = run(config, cli, text, model_id).await {
+ let highlight = stderr().is_terminal() && highlight;
+ render_error(err, highlight);
+ std::process::exit(1);
+ }
+ Ok(())
+}
+async fn run(
+ config: GlobalConfig,
+ cli: Cli,
+ text: Option<String>,
+ model_id: Option<String>,
+) -> Result<()> {
let abort_signal = create_abort_signal();
if let Some(addr) = cli.serve {
@@ -127,7 +148,7 @@ async fn main() -> Result<()> {
println!("{sessions}");
return Ok(());
}
- if let Some(model_id) = &cli.model {
+ if let Some(model_id) = &model_id {
config.write().set_model(model_id)?;
}
if cli.save_session {
@@ -141,29 +162,22 @@ async fn main() -> Result<()> {
println!("{}", info);
return Ok(());
}
- if cli.execute {
- if no_input {
- bail!("No input");
- }
- let input = create_input(&config, text, file).await?;
+ let is_repl = config.read().working_mode.is_repl();
+ if cli.execute && !is_repl {
+ let input = create_input(&config, text, &cli.file).await?;
let shell = detect_shell();
shell_execute(&config, &shell, input).await?;
return Ok(());
}
config.write().apply_prelude()?;
- if let Err(err) = match no_input {
+ match is_repl {
false => {
- let mut input = create_input(&config, text, file).await?;
+ let mut input = create_input(&config, text, &cli.file).await?;
input.use_embeddings(abort_signal.clone()).await?;
start_directive(&config, input, cli.no_stream, cli.code, abort_signal).await
}
true => start_interactive(&config).await,
- } {
- let highlight = stderr().is_terminal() && config.read().highlight;
- render_error(err, highlight);
- std::process::exit(1);
}
- Ok(())
}
#[async_recursion]