From 98ac7e2b573fb5d367fb6ce399605b1fe599f69e Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 18 Jun 2024 12:37:25 +0800 Subject: feat: support reka client (#614) --- src/client/mod.rs | 1 + src/client/reka.rs | 126 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 127 insertions(+) create mode 100644 src/client/reka.rs (limited to 'src') diff --git a/src/client/mod.rs b/src/client/mod.rs index 579c2a3..6710e20 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -24,6 +24,7 @@ register_client!( (gemini, "gemini", GeminiConfig, GeminiClient), (claude, "claude", ClaudeConfig, ClaudeClient), (cohere, "cohere", CohereConfig, CohereClient), + (reka, "reka", RekaConfig, RekaClient), (ollama, "ollama", OllamaConfig, OllamaClient), ( azure_openai, diff --git a/src/client/reka.rs b/src/client/reka.rs new file mode 100644 index 0000000..46f9118 --- /dev/null +++ b/src/client/reka.rs @@ -0,0 +1,126 @@ +use super::*; + +use anyhow::{anyhow, Result}; +use reqwest::{Client as ReqwestClient, RequestBuilder}; +use serde::Deserialize; +use serde_json::{json, Value}; + +const CHAT_COMPLETIONS_API_URL: &str = "https://api.reka.ai/v1/chat"; + +#[derive(Debug, Clone, Deserialize, Default)] +pub struct RekaConfig { + pub name: Option, + pub api_key: Option, + #[serde(default)] + pub models: Vec, + pub patches: Option, + pub extra: Option, +} + +impl RekaClient { + config_get_fn!(api_key, get_api_key); + + pub const PROMPTS: [PromptAction<'static>; 1] = + [("api_key", "API Key:", true, PromptKind::String)]; + + fn chat_completions_builder( + &self, + client: &ReqwestClient, + data: ChatCompletionsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let mut body = build_chat_completions_body(data, &self.model); + self.patch_chat_completions_body(&mut body); + + let url = CHAT_COMPLETIONS_API_URL; + + debug!("Reka Chat Completions Request: {url} {body}"); + + let builder = client.post(url).header("x-api-key", api_key).json(&body); + + Ok(builder) + } +} + +impl_client_trait!( + RekaClient, + chat_completions, + chat_completions_streaming +); + +async fn chat_completions(builder: RequestBuilder) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + + debug!("non-stream-data: {data}"); + extract_chat_completions(&data) +} + +async fn chat_completions_streaming( + builder: RequestBuilder, + handler: &mut SseHandler, +) -> Result<()> { + let mut prev_text = String::new(); + let handle = |message: SseMmessage| -> Result { + let data: Value = serde_json::from_str(&message.data)?; + debug!("stream-data: {data}"); + if let Some(text) = data["responses"][0]["chunk"]["content"].as_str() { + let delta_text = &text[prev_text.len()..]; + prev_text = text.to_string(); + handler.text(delta_text)?; + } + Ok(false) + }; + + sse_stream(builder, handle).await +} + +fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value { + let ChatCompletionsData { + mut messages, + temperature, + top_p, + functions: _, + stream, + } = data; + + patch_system_message(&mut messages); + + let mut body = json!({ + "model": &model.name(), + "messages": messages, + }); + + if let Some(v) = model.max_tokens_param() { + body["max_tokens"] = v.into(); + } + if let Some(v) = temperature { + body["temperature"] = v.into(); + } + if let Some(v) = top_p { + body["top_p"] = v.into(); + } + if stream { + body["stream"] = true.into(); + } + + body +} + +fn extract_chat_completions(data: &Value) -> Result { + let text = data["responses"][0]["message"]["content"].as_str().ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + + let output = ChatCompletionsOutput { + text: text.to_string(), + tool_calls: vec![], + id: data["id"].as_str().map(|v| v.to_string()), + input_tokens: data["usage"]["input_tokens"].as_u64(), + output_tokens: data["usage"]["output_tokens"].as_u64(), + }; + Ok(output) +} -- cgit v1.2.3