diff options
Diffstat (limited to 'src/client/cloudflare.rs')
| -rw-r--r-- | src/client/cloudflare.rs | 31 |
1 files changed, 19 insertions, 12 deletions
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 15dd5bb..966cee4 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -1,5 +1,5 @@ use super::{ - catch_error, sse_stream, Client, CloudflareClient, CompletionData, CompletionOutput, + catch_error, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, CloudflareClient, ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, SseMmessage, }; @@ -30,15 +30,15 @@ impl CloudflareClient { ("api_key", "API Key:", true, PromptKind::String), ]; - fn request_builder( + fn chat_completions_builder( &self, client: &ReqwestClient, - data: CompletionData, + data: ChatCompletionsData, ) -> Result<RequestBuilder> { let account_id = self.get_account_id()?; let api_key = self.get_api_key()?; - let mut body = build_body(data, &self.model)?; + let mut body = build_chat_completions_body(data, &self.model)?; self.patch_request_body(&mut body); let url = format!( @@ -54,9 +54,13 @@ impl CloudflareClient { } } -impl_client_trait!(CloudflareClient, send_message, send_message_streaming); +impl_client_trait!( + CloudflareClient, + chat_completions, + chat_completions_streaming +); -async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> { +async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -65,10 +69,13 @@ async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> { } debug!("non-stream-data: {data}"); - extract_completion(&data) + extract_chat_completions(&data) } -async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> { +async fn chat_completions_streaming( + builder: RequestBuilder, + handler: &mut SseHandler, +) -> Result<()> { let handle = |message: SseMmessage| -> Result<bool> { if message.data == "[DONE]" { return Ok(true); @@ -83,8 +90,8 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle sse_stream(builder, handle).await } -fn build_body(data: CompletionData, model: &Model) -> Result<Value> { - let CompletionData { +fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> { + let ChatCompletionsData { messages, temperature, top_p, @@ -113,10 +120,10 @@ fn build_body(data: CompletionData, model: &Model) -> Result<Value> { Ok(body) } -fn extract_completion(data: &Value) -> Result<CompletionOutput> { +fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> { let text = data["result"]["response"] .as_str() .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - Ok(CompletionOutput::new(text)) + Ok(ChatCompletionsOutput::new(text)) } |
