summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure.rs (renamed from src/client/azure_openai.rs)31
-rw-r--r--src/client/mod.rs8
2 files changed, 11 insertions, 28 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure.rs
index d1dc43b..bae7851 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure.rs
@@ -1,34 +1,33 @@
use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, Model};
+use super::{AzureClient, ExtraConfig, PromptType, SendData, Model};
use crate::utils::PromptKind;
-use anyhow::{anyhow, Result};
+use anyhow::Result;
use async_trait::async_trait;
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
-use std::env;
-
#[derive(Debug, Clone, Deserialize)]
-pub struct AzureOpenAIConfig {
+pub struct AzureConfig {
pub name: Option<String>,
pub api_base: Option<String>,
pub api_key: Option<String>,
- pub models: Vec<AzureOpenAIModel>,
+ pub models: Vec<AzureModel>,
pub extra: Option<ExtraConfig>,
}
#[derive(Debug, Clone, Deserialize)]
-pub struct AzureOpenAIModel {
+pub struct AzureModel {
name: String,
max_tokens: Option<usize>,
}
-openai_compatible_client!(AzureOpenAIClient);
+openai_compatible_client!(AzureClient);
-impl AzureOpenAIClient {
+impl AzureClient {
config_get_fn!(api_base, get_api_base);
+ config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 4] = [
("api_base", "API Base:", true, PromptKind::String),
@@ -42,7 +41,7 @@ impl AzureOpenAIClient {
),
];
- pub fn list_models(local_config: &AzureOpenAIConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &AzureConfig, client_index: usize) -> Vec<Model> {
let client_name = Self::name(local_config);
local_config
@@ -57,18 +56,8 @@ impl AzureOpenAIClient {
}
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
- let api_key = self.config.api_key.clone();
- let api_key = api_key
- .or_else(|| {
- let env_prefix = match &self.config.name {
- None => "AZURE".into(),
- Some(v) => v.to_uppercase(),
- };
- env::var(format!("{env_prefix}_OPENAI_KEY")).ok()
- })
- .ok_or_else(|| anyhow!("Miss api_key"))?;
-
let api_base = self.get_api_base()?;
+ let api_key = self.get_api_key()?;
let body = openai_build_body(data, self.model.llm_name.clone());
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 7ac9aa0..55f0ed0 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -10,11 +10,5 @@ pub use model::*;
register_client!(
(openai, "openai", OpenAI, OpenAIConfig, OpenAIClient),
(localai, "localai", LocalAI, LocalAIConfig, LocalAIClient),
- (
- azure_openai,
- "azure-openai",
- AzureOpenAI,
- AzureOpenAIConfig,
- AzureOpenAIClient
- ),
+ (azure, "azure", Azure, AzureConfig, AzureClient),
);