diff options
| author | sigoden <sigoden@gmail.com> | 2024-01-13 19:52:07 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-01-13 19:52:07 +0800 |
| commit | fe35cfd9419302f01baf9672493c0b0a4b41d889 (patch) | |
| tree | 94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/gemini.rs | |
| parent | 4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff) | |
| download | aichat-fe35cfd9419302f01baf9672493c0b0a4b41d889.tar.gz | |
feat: supports model capabilities (#297)
1. automatically switch to the model that has the necessary capabilities.
2. throw an error if the client does not have a model with the necessary capabilities
Diffstat (limited to 'src/client/gemini.rs')
| -rw-r--r-- | src/client/gemini.rs | 17 |
1 files changed, 8 insertions, 9 deletions
diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 98fd60d..6ff3f95 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -3,7 +3,7 @@ use super::{ SendData, TokensCountFactors, }; -use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind}; +use crate::{render::ReplyHandler, utils::PromptKind}; use anyhow::{anyhow, bail, Result}; use async_trait::async_trait; @@ -14,10 +14,10 @@ use serde_json::{json, Value}; const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/"; -const MODELS: [(&str, usize); 3] = [ - ("gemini-pro", 32768), - ("gemini-pro-vision", 16384), - ("gemini-ultra", 32768), +const MODELS: [(&str, usize, &str); 3] = [ + ("gemini-pro", 32768, "text"), + ("gemini-pro-vision", 16384, "vision"), + ("gemini-ultra", 32768, "text"), ]; const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2); @@ -31,9 +31,7 @@ pub struct GeminiConfig { #[async_trait] impl Client for GeminiClient { - fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>) { - (&self.global_config, &self.config.extra) - } + client_common_fns!(); async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { let builder = self.request_builder(client, data)?; @@ -61,8 +59,9 @@ impl GeminiClient { let client_name = Self::name(local_config); MODELS .into_iter() - .map(|(name, max_tokens)| { + .map(|(name, max_tokens, capabilities)| { Model::new(client_name, name) + .set_capabilities(capabilities.into()) .set_max_tokens(Some(max_tokens)) .set_tokens_count_factors(TOKENS_COUNT_FACTORS) }) |
