diff options
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 38 |
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), ( |
