summaryrefslogtreecommitdiffstats
path: root/src/client/replicate.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/replicate.rs
parentf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff)
downloadaichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/replicate.rs')
-rw-r--r--src/client/replicate.rs53
1 files changed, 27 insertions, 26 deletions
diff --git a/src/client/replicate.rs b/src/client/replicate.rs
index d6ca401..f0db24a 100644
--- a/src/client/replicate.rs
+++ b/src/client/replicate.rs
@@ -2,7 +2,6 @@ use super::prompt_format::*;
use super::*;
use anyhow::{anyhow, Result};
-use async_trait::async_trait;
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -25,25 +24,9 @@ impl ReplicateClient {
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
-
- fn prepare_chat_completions(
- &self,
- data: ChatCompletionsData,
- api_key: &str,
- ) -> Result<RequestData> {
- let url = format!("{API_BASE}/models/{}/predictions", self.model.name());
-
- let body = build_chat_completions_body(data, &self.model)?;
-
- let mut request_data = RequestData::new(url, body);
-
- request_data.bearer_auth(api_key);
-
- Ok(request_data)
- }
}
-#[async_trait]
+#[async_trait::async_trait]
impl Client for ReplicateClient {
client_common_fns!();
@@ -52,10 +35,9 @@ impl Client for ReplicateClient {
client: &ReqwestClient,
data: ChatCompletionsData,
) -> Result<ChatCompletionsOutput> {
- let api_key = self.get_api_key()?;
- let request_data = self.prepare_chat_completions(data, &api_key)?;
+ let request_data = prepare_chat_completions(self, data)?;
let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
- chat_completions(client, builder, &api_key).await
+ chat_completions(builder, client, &self.get_api_key()?).await
}
async fn chat_completions_streaming_inner(
@@ -64,16 +46,32 @@ impl Client for ReplicateClient {
handler: &mut SseHandler,
data: ChatCompletionsData,
) -> Result<()> {
- let api_key = self.get_api_key()?;
- let request_data = self.prepare_chat_completions(data, &api_key)?;
+ let request_data = prepare_chat_completions(self, data)?;
let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
- chat_completions_streaming(client, builder, handler).await
+ chat_completions_streaming(builder, handler, client).await
}
}
+fn prepare_chat_completions(
+ self_: &ReplicateClient,
+ data: ChatCompletionsData,
+) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
+
+ let url = format!("{API_BASE}/models/{}/predictions", self_.model.name());
+
+ let body = build_chat_completions_body(data, &self_.model)?;
+
+ let mut request_data = RequestData::new(url, body);
+
+ request_data.bearer_auth(api_key);
+
+ Ok(request_data)
+}
+
async fn chat_completions(
- client: &ReqwestClient,
builder: RequestBuilder,
+ client: &ReqwestClient,
api_key: &str,
) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -106,9 +104,9 @@ async fn chat_completions(
}
async fn chat_completions_streaming(
- client: &ReqwestClient,
builder: RequestBuilder,
handler: &mut SseHandler,
+ client: &ReqwestClient,
) -> Result<()> {
let res = builder.send().await?;
let status = res.status();
@@ -126,6 +124,9 @@ async fn chat_completions_streaming(
if message.event == "done" {
return Ok(true);
}
+
+ debug!("stream-data: {}", message.data);
+
handler.text(&message.data)?;
Ok(false)
};