summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-02-28 08:22:15 +0800
committerGitHub <noreply@github.com>2024-02-28 08:22:15 +0800
commit75fe0b9205cc5ad5285e53ee5881b6d89e8a6000 (patch)
tree83e648a48bac49a7eaaf456cb531a1119f10279a /src/client
parentc538533014caa2bc84724392f3db2fbb0f443209 (diff)
downloadaichat-75fe0b9205cc5ad5285e53ee5881b6d89e8a6000.tar.gz
feat: support mistral (#324)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/mistral.rs68
-rw-r--r--src/client/mod.rs1
2 files changed, 69 insertions, 0 deletions
diff --git a/src/client/mistral.rs b/src/client/mistral.rs
new file mode 100644
index 0000000..1bb4889
--- /dev/null
+++ b/src/client/mistral.rs
@@ -0,0 +1,68 @@
+use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
+use super::{ExtraConfig, MistralClient, Model, PromptType, SendData};
+
+use crate::utils::PromptKind;
+
+use anyhow::Result;
+use async_trait::async_trait;
+use reqwest::{Client as ReqwestClient, RequestBuilder};
+use serde::Deserialize;
+
+const API_URL: &str = "https://api.mistral.ai/v1/chat/completions";
+
+const MODELS: [(&str, usize, &str); 5] = [
+ ("mistral-small-latest", 32000, "text"),
+ ("mistral-medium-latest", 32000, "text"),
+ ("mistral-larget-latest", 32000, "text"),
+ ("open-mistral-7b", 32000, "text"),
+ ("open-mixtral-8x7b", 32000, "text"),
+];
+
+
+#[derive(Debug, Clone, Deserialize)]
+pub struct MistralConfig {
+ pub name: Option<String>,
+ pub api_key: Option<String>,
+ pub extra: Option<ExtraConfig>,
+}
+
+openai_compatible_client!(MistralClient);
+
+impl MistralClient {
+ config_get_fn!(api_key, get_api_key);
+
+ pub const PROMPTS: [PromptType<'static>; 1] = [
+ ("api_key", "API Key:", false, PromptKind::String),
+ ];
+
+ pub fn list_models(local_config: &MistralConfig) -> Vec<Model> {
+ let client_name = Self::name(local_config);
+ MODELS
+ .into_iter()
+ .map(|(name, max_tokens, capabilities)| {
+ Model::new(client_name, name)
+ .set_capabilities(capabilities.into())
+ .set_max_tokens(Some(max_tokens))
+ .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
+ })
+ .collect()
+ }
+
+ fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key().ok();
+
+ let mut body = openai_build_body(data, self.model.name.clone());
+ self.model.merge_extra_fields(&mut body);
+
+ let url = API_URL;
+
+ debug!("Mistral Request: {url} {body}");
+
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_key) = api_key {
+ builder = builder.bearer_auth(api_key);
+ }
+
+ Ok(builder)
+ }
+}
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 1517b93..b1d4979 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -11,6 +11,7 @@ register_client!(
(openai, "openai", OpenAIConfig, OpenAIClient),
(gemini, "gemini", GeminiConfig, GeminiClient),
(claude, "claude", ClaudeConfig, ClaudeClient),
+ (mistral, "mistral", MistralConfig, MistralClient),
(localai, "localai", LocalAIConfig, LocalAIClient),
(ollama, "ollama", OllamaConfig, OllamaClient),
(