summaryrefslogtreecommitdiffstats
path: root/src/client/gemini.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-01-13 19:52:07 +0800
committerGitHub <noreply@github.com>2024-01-13 19:52:07 +0800
commitfe35cfd9419302f01baf9672493c0b0a4b41d889 (patch)
tree94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/gemini.rs
parent4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff)
downloadaichat-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.rs17
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)
})