summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-12-04 21:03:59 +0800
committerGitHub <noreply@github.com>2024-12-04 21:03:59 +0800
commit3a3388375be05758d5f5574cce1d0f8eb7bdd604 (patch)
treecc8e19e218c20f28a04feca33850fce4bb2cbe98 /src/client/common.rs
parent7d42fe9429f75d195f865b07cef10d040d5397f2 (diff)
downloadaichat-3a3388375be05758d5f5574cce1d0f8eb7bdd604.tar.gz
refactor: improve retrieve model (#1036)
- check the model type while retrieve model - select chat/reranker model even if it is missed in client models - find predefined-models for openai-compatible client with startsWith - remove client::ApiType
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs38
1 files changed, 7 insertions, 31 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 9c516c8..e0ec861 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -19,7 +19,7 @@ use tokio::sync::mpsc::unbounded_channel;
const MODELS_YAML: &str = include_str!("../../models.yaml");
lazy_static::lazy_static! {
- pub static ref ALL_MODELS: Vec<BuiltinModels> = serde_yaml::from_str(MODELS_YAML).unwrap();
+ pub static ref ALL_PREDEFINED_MODELS: Vec<PredefinedModels> = serde_yaml::from_str(MODELS_YAML).unwrap();
static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap();
}
@@ -144,23 +144,23 @@ pub trait Client: Sync + Send {
&self,
client: &reqwest::Client,
mut request_data: RequestData,
- api_type: ApiType,
) -> RequestBuilder {
- self.patch_request_data(&mut request_data, api_type);
+ self.patch_request_data(&mut request_data);
request_data.into_builder(client)
}
- fn patch_request_data(&self, request_data: &mut RequestData, api_type: ApiType) {
+ fn patch_request_data(&self, request_data: &mut RequestData) {
+ let model_type = self.model().model_type();
let map = std::env::var(get_env_name(&format!(
"patch_{}_{}",
self.model().client_name(),
- api_type.name(),
+ model_type.api_name(),
)))
.ok()
.and_then(|v| serde_json::from_str(&v).ok())
.or_else(|| {
self.patch_config()
- .and_then(|v| api_type.extract_patch(v))
+ .and_then(|v| model_type.extract_patch(v))
.cloned()
});
let map = match map {
@@ -200,30 +200,6 @@ pub struct RequestPatch {
pub type ApiPatch = IndexMap<String, Value>;
-#[derive(Debug, Clone, Copy, PartialEq, Eq)]
-pub enum ApiType {
- ChatCompletions,
- Embeddings,
- Rerank,
-}
-
-impl ApiType {
- pub fn name(&self) -> &str {
- match self {
- ApiType::ChatCompletions => "chat_completions",
- ApiType::Embeddings => "embeddings",
- ApiType::Rerank => "rerank",
- }
- }
- pub fn extract_patch<'a>(&self, patch: &'a RequestPatch) -> Option<&'a ApiPatch> {
- match self {
- ApiType::ChatCompletions => patch.chat_completions.as_ref(),
- ApiType::Embeddings => patch.embeddings.as_ref(),
- ApiType::Rerank => patch.rerank.as_ref(),
- }
- }
-}
-
pub struct RequestData {
pub url: String,
pub headers: IndexMap<String, String>,
@@ -383,7 +359,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St
config["api_base"] = api_base.into();
}
prompts.push(("api_key", "API Key:", false, PromptKind::String));
- if !ALL_MODELS.iter().any(|v| v.platform == name) {
+ if !ALL_PREDEFINED_MODELS.iter().any(|v| v.platform == name) {
prompts.extend([
("models[].name", "Model Name:", true, PromptKind::String),
(