summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-07 11:54:56 +0800
committerGitHub <noreply@github.com>2023-11-07 11:54:56 +0800
commitd40f104f667073f320a01d9c1a91aa88225ccaeb (patch)
treedf915b72efb7051a53a95439c2d8cc4627921c8e /src/client
parent9a8b302432a3f9bfa1e467dde027fc92dacce3e2 (diff)
downloadaichat-d40f104f667073f320a01d9c1a91aa88225ccaeb.tar.gz
feat: allow the use of an unlisted model (#219)
Diffstat (limited to 'src/client')
-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
7 files changed, 43 insertions, 0 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");