summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/client/azure_openai.rs2
-rw-r--r--src/client/ernie.rs2
-rw-r--r--src/client/localai.rs2
-rw-r--r--src/client/model.rs31
-rw-r--r--src/client/openai.rs2
-rw-r--r--src/client/palm.rs2
-rw-r--r--src/client/qianwen.rs2
-rw-r--r--src/config/mod.rs12
8 files changed, 45 insertions, 10 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index f4b7916..4700ad7 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -66,6 +66,8 @@ impl AzureOpenAIClient {
&api_base, self.model.name
);
+ debug!("AzureOpenAI Request: {url} {body}");
+
let builder = client.post(url).header("api-key", api_key).json(&body);
Ok(builder)
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 5e32a81..084933c 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -85,6 +85,8 @@ impl ErnieClient {
&ACCESS_TOKEN
});
+ debug!("Ernie Request: {url} {body}");
+
let builder = client.post(url).json(&body);
Ok(builder)
diff --git a/src/client/localai.rs b/src/client/localai.rs
index 93853dc..9325e0f 100644
--- a/src/client/localai.rs
+++ b/src/client/localai.rs
@@ -68,6 +68,8 @@ impl LocalAIClient {
let url = format!("{}{chat_endpoint}", self.config.api_base);
+ debug!("LocalAI Request: {url} {body}");
+
let mut builder = client.post(url).json(&body);
if let Some(api_key) = api_key {
builder = builder.bearer_auth(api_key);
diff --git a/src/client/model.rs b/src/client/model.rs
index 82d47c9..16fe087 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -30,6 +30,37 @@ impl Model {
}
}
+ pub fn find(models: &[Self], value: &str) -> Option<Self> {
+ let mut model = None;
+ let (client_name, model_name) = match value.split_once(':') {
+ Some((client_name, model_name)) => {
+ if model_name.is_empty() {
+ (client_name, None)
+ } else {
+ (client_name, Some(model_name))
+ }
+ }
+ None => (value, None),
+ };
+ match model_name {
+ Some(model_name) => {
+ if let Some(found) = models.iter().find(|v| v.id() == value) {
+ model = Some(found.clone());
+ } else if let Some(found) = models.iter().find(|v| v.client_name == client_name) {
+ let mut found = found.clone();
+ found.name = model_name.to_string();
+ model = Some(found)
+ }
+ }
+ None => {
+ if let Some(found) = models.iter().find(|v| v.client_name == client_name) {
+ model = Some(found.clone());
+ }
+ }
+ }
+ model
+ }
+
pub fn id(&self) -> String {
format!("{}:{}", self.client_name, self.name)
}
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 6b5edab..06a5316 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -68,6 +68,8 @@ impl OpenAIClient {
let url = format!("{api_base}/chat/completions");
+ debug!("OpenAI Request: {url} {body}");
+
let mut builder = client.post(url).bearer_auth(api_key).json(&body);
if let Some(organization_id) = &self.config.organization_id {
diff --git a/src/client/palm.rs b/src/client/palm.rs
index 45496fa..dd063da 100644
--- a/src/client/palm.rs
+++ b/src/client/palm.rs
@@ -70,6 +70,8 @@ impl PaLMClient {
let url = format!("{API_BASE}{}:generateMessage?key={}", model, api_key);
+ debug!("PaLM Request: {url} {body}");
+
let builder = client.post(url).json(&body);
Ok(builder)
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index b8e0d6f..2ed9945 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -68,6 +68,8 @@ impl QianwenClient {
let stream = data.stream;
let body = build_body(data, self.model.name.clone());
+ debug!("Qianwen Request: {API_URL} {body}");
+
let mut builder = client.post(API_URL).bearer_auth(api_key).json(&body);
if stream {
builder = builder.header("X-DashScope-SSE", "enable");
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 1720772..e6fef5f 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -305,17 +305,9 @@ impl Config {
pub fn set_model(&mut self, value: &str) -> Result<()> {
let models = list_models(self);
- let mut model = None;
- let value = value.trim_end_matches(':');
- if value.contains(':') {
- if let Some(found) = models.iter().find(|v| v.id() == value) {
- model = Some(found.clone());
- }
- } else if let Some(found) = models.iter().find(|v| v.client_name == value) {
- model = Some(found.clone());
- }
+ let model = Model::find(&models, value);
match model {
- None => bail!("Unknown model '{}'", value),
+ None => bail!("Invalid model '{}'", value),
Some(model) => {
if let Some(session) = self.session.as_mut() {
session.set_model(model.clone())?;