summaryrefslogtreecommitdiffstats
path: root/src/client/gemini.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-28 06:04:36 +0800
committerGitHub <noreply@github.com>2024-07-28 06:04:36 +0800
commit0e740d81e94505bd57036755abaaecb12c3b26e3 (patch)
tree49000370fb12e4e5e1f1bd4d145104f2f502aa16 /src/client/gemini.rs
parentf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff)
downloadaichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/gemini.rs')
-rw-r--r--src/client/gemini.rs83
1 files changed, 45 insertions, 38 deletions
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index aa1a5b1..2616218 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -23,57 +23,64 @@ impl GeminiClient {
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
+}
+
+impl_client_trait!(
+ GeminiClient,
+ (
+ prepare_chat_completions,
+ gemini_chat_completions,
+ gemini_chat_completions_streaming
+ ),
+ (prepare_embeddings, gemini_embeddings),
+ (noop_prepare_rerank, noop_rerank),
+);
- fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
+fn prepare_chat_completions(
+ self_: &GeminiClient,
+ data: ChatCompletionsData,
+) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
- let func = match data.stream {
- true => "streamGenerateContent",
- false => "generateContent",
- };
+ let func = match data.stream {
+ true => "streamGenerateContent",
+ false => "generateContent",
+ };
- let url = format!("{API_BASE}{}:{}?key={}", &self.model.name(), func, api_key);
+ let url = format!("{API_BASE}{}:{}?key={}", self_.model.name(), func, api_key);
- let body = gemini_build_chat_completions_body(data, &self.model)?;
+ let body = gemini_build_chat_completions_body(data, &self_.model)?;
- let request_data = RequestData::new(url, body);
+ let request_data = RequestData::new(url, body);
- Ok(request_data)
- }
+ Ok(request_data)
+}
- fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
+fn prepare_embeddings(self_: &GeminiClient, data: EmbeddingsData) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
- let url = format!(
- "{API_BASE}{}:embedContent?key={}",
- &self.model.name(),
- api_key
- );
+ let url = format!(
+ "{API_BASE}{}:embedContent?key={}",
+ self_.model.name(),
+ api_key
+ );
- let body = json!({
- "content": {
- "parts": [
- {
- "text": data.texts[0],
- }
- ]
- }
- });
+ let body = json!({
+ "content": {
+ "parts": [
+ {
+ "text": data.texts[0],
+ }
+ ]
+ }
+ });
- let request_data = RequestData::new(url, body);
+ let request_data = RequestData::new(url, body);
- Ok(request_data)
- }
+ Ok(request_data)
}
-impl_client_trait!(
- GeminiClient,
- gemini_chat_completions,
- gemini_chat_completions_streaming,
- gemini_embeddings
-);
-
-async fn gemini_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
+async fn gemini_embeddings(builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;