diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-28 06:04:36 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-28 06:04:36 +0800 |
| commit | 0e740d81e94505bd57036755abaaecb12c3b26e3 (patch) | |
| tree | 49000370fb12e4e5e1f1bd4d145104f2f502aa16 /src/client/cloudflare.rs | |
| parent | f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff) | |
| download | aichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz | |
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/cloudflare.rs')
| -rw-r--r-- | src/client/cloudflare.rs | 81 |
1 files changed, 46 insertions, 35 deletions
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 3ee0a91..a24a1c6 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -26,54 +26,64 @@ impl CloudflareClient { ("account_id", "Account ID:", true, PromptKind::String), ("api_key", "API Key:", true, PromptKind::String), ]; +} - fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> { - let account_id = self.get_account_id()?; - let api_key = self.get_api_key()?; +impl_client_trait!( + CloudflareClient, + ( + prepare_chat_completions, + chat_completions, + chat_completions_streaming + ), + (prepare_embeddings, embeddings), + (noop_prepare_rerank, noop_rerank), +); - let url = format!( - "{API_BASE}/accounts/{account_id}/ai/run/{}", - self.model.name() - ); +fn prepare_chat_completions( + self_: &CloudflareClient, + data: ChatCompletionsData, +) -> Result<RequestData> { + let account_id = self_.get_account_id()?; + let api_key = self_.get_api_key()?; - let body = build_chat_completions_body(data, &self.model)?; + let url = format!( + "{API_BASE}/accounts/{account_id}/ai/run/{}", + self_.model.name() + ); - let mut request_data = RequestData::new(url, body); + let body = build_chat_completions_body(data, &self_.model)?; - request_data.bearer_auth(api_key); + let mut request_data = RequestData::new(url, body); - Ok(request_data) - } + request_data.bearer_auth(api_key); + + Ok(request_data) +} - fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> { - let account_id = self.get_account_id()?; - let api_key = self.get_api_key()?; +fn prepare_embeddings(self_: &CloudflareClient, data: EmbeddingsData) -> Result<RequestData> { + let account_id = self_.get_account_id()?; + let api_key = self_.get_api_key()?; - let url = format!( - "{API_BASE}/accounts/{account_id}/ai/run/{}", - self.model.name() - ); + let url = format!( + "{API_BASE}/accounts/{account_id}/ai/run/{}", + self_.model.name() + ); - let body = json!({ - "text": data.texts, - }); + let body = json!({ + "text": data.texts, + }); - let mut request_data = RequestData::new(url, body); + let mut request_data = RequestData::new(url, body); - request_data.bearer_auth(api_key); + request_data.bearer_auth(api_key); - Ok(request_data) - } + Ok(request_data) } -impl_client_trait!( - CloudflareClient, - chat_completions, - chat_completions_streaming, - embeddings -); - -async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> { +async fn chat_completions( + builder: RequestBuilder, + _model: &Model, +) -> Result<ChatCompletionsOutput> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -88,6 +98,7 @@ async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutp async fn chat_completions_streaming( builder: RequestBuilder, handler: &mut SseHandler, + _model: &Model, ) -> Result<()> { let handle = |message: SseMmessage| -> Result<bool> { if message.data == "[DONE]" { @@ -103,7 +114,7 @@ async fn chat_completions_streaming( sse_stream(builder, handle).await } -async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> { +async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; |
