diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-07 11:54:56 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-07 11:54:56 +0800 |
| commit | d40f104f667073f320a01d9c1a91aa88225ccaeb (patch) | |
| tree | df915b72efb7051a53a95439c2d8cc4627921c8e /src | |
| parent | 9a8b302432a3f9bfa1e467dde027fc92dacce3e2 (diff) | |
| download | aichat-d40f104f667073f320a01d9c1a91aa88225ccaeb.tar.gz | |
feat: allow the use of an unlisted model (#219)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/azure_openai.rs | 2 | ||||
| -rw-r--r-- | src/client/ernie.rs | 2 | ||||
| -rw-r--r-- | src/client/localai.rs | 2 | ||||
| -rw-r--r-- | src/client/model.rs | 31 | ||||
| -rw-r--r-- | src/client/openai.rs | 2 | ||||
| -rw-r--r-- | src/client/palm.rs | 2 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 12 |
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())?; |
