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) --- Argcfile.sh | 33 ++++++++++++++ README.md | 4 +- config.example.yaml | 9 ++++ models.yaml | 22 +++++++++ src/client/mod.rs | 1 + src/client/reka.rs | 126 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 6 files changed, 194 insertions(+), 1 deletion(-) create mode 100644 src/client/reka.rs diff --git a/Argcfile.sh b/Argcfile.sh index 5f04b6f..96300b4 100755 --- a/Argcfile.sh +++ b/Argcfile.sh @@ -229,6 +229,27 @@ models-cohere() { } +# @cmd Chat with reka api +# @env REKA_API_KEY! +# @option -m --model=reka-flash $REKA_MODEL +# @flag -S --no-stream +# @arg text~ +chat-reka() { + _wrapper curl -i https://api.reka.ai/v1/chat \ +-X POST \ +-H 'Content-Type: application/json' \ +-H "X-Api-Key: $REKA_API_KEY" \ +-d "$(_build_body reka "$@")" +} + +# @cmd List reka models +# @env REKA_API_KEY! +models-reka() { + _wrapper curl https://api.reka.ai/v1/models \ +-H "X-Api-Key: $REKA_API_KEY" \ + +} + # @cmd Chat with ollama api # @option -m --model=codegemma $OLLAMA_MODEL # @flag -S --no-stream @@ -471,6 +492,18 @@ _build_body() { "model": "'$argc_model'", "message": "'"$*"'", "stream": '$stream' +}' + ;; + reka) + echo '{ + "model": "'$argc_model'", + "messages": [ + { + "role": "user", + "content": "'"$*"'" + } + ], + "stream": '$stream' }' ;; claude) diff --git a/README.md b/README.md index 0d0adcf..f7f6e23 100644 --- a/README.md +++ b/README.md @@ -31,6 +31,7 @@ AIChat is an all-in-one AI CLI tool featuring chat REPL, RAG, function calling, - Claude: Claude-3 (vision, paid, function-calling) - Mistral (paid, embedding, function-calling) - Cohere: Command-R/Command-R+ (paid, embedding, function-calling) +- Reka (paid, vision) - Perplexity: Llama-3/Mixtral (paid) - Groq: Llama-3/Mixtral/Gemma (free) - Ollama (free, local, embedding) @@ -43,8 +44,9 @@ AIChat is an all-in-one AI CLI tool featuring chat REPL, RAG, function calling, - Ernie (paid) - Qianwen (paid, vision, embedding) - Moonshot (paid) -- ZhipuAI: GLM-3.5/GLM-4 (paid, vision) - Deepseek (paid) +- ZhipuAI: GLM-3.5/GLM-4 (paid, vision) +- LingYiWanWu (paid) - Other openAI-compatible platforms ## Install diff --git a/config.example.yaml b/config.example.yaml index 9e32a91..6e4416c 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -137,6 +137,10 @@ clients: - type: cohere api_key: xxx # ENV: {client}_API_KEY + # See https://docs.reka.ai/quick-start + - type: reka + api_key: xxx # ENV: {client}_API_KEY + # See https://docs.perplexity.ai/docs/getting-started - type: openai-compatible name: perplexity @@ -238,6 +242,11 @@ clients: name: zhipuai api_key: xxx # ENV: {client}_API_KEY + # See https://platform.lingyiwanwu.com/docs + - type: openai-compatible + name: lingyiwanwu + api_key: xxx # ENV: {client}_API_KEY + # See https://docs.endpoints.anyscale.com/ - type: openai-compatible name: anyscale diff --git a/models.yaml b/models.yaml index 47eac21..b193200 100644 --- a/models.yaml +++ b/models.yaml @@ -199,6 +199,28 @@ default_chunk_size: 1000 max_concurrent_chunks: 96 +- platform: reka + docs: + # - https://www.reka.ai/ourmodels + # - https://www.reka.ai/reka-deploy + # - https://docs.reka.ai/api-reference/chat/create + models: + - name: reka-core + max_input_tokens: 128000 + input_price: 10 + output_price: 25 + supports_vision: true + - name: reka-flash + max_input_tokens: 128000 + input_price: 0.8 + output_price: 2 + supports_vision: true + - name: reka-edge + max_input_tokens: 128000 + input_price: 0.4 + output_price: 1 + supports_vision: true + - platform: perplexity # docs: # - https://docs.perplexity.ai/docs/model-cards 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