summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--README.md1
-rw-r--r--config.example.yaml4
-rw-r--r--src/client/claude.rs171
-rw-r--r--src/client/mod.rs1
4 files changed, 177 insertions, 0 deletions
diff --git a/README.md b/README.md
index c8ee3a1..a8f3837 100644
--- a/README.md
+++ b/README.md
@@ -48,6 +48,7 @@ Download it from [GitHub Releases](https://github.com/sigoden/aichat/releases),
- LocalAI: opensource LLMs and other openai-compatible LLMs
- Ollama: opensource LLMs
- VertexAI: gemini-pro/gemini-pro-vision/gemini-ultra/gemini-ultra-vision
+- Claude: claude-instant-1.2/claude-2.0/claude-2.1
- Azure-OpenAI: user deployed gpt-3.5/gpt-4
- Ernie: ernie-bot-turbo/ernie-bot/ernie-bot-8k/ernie-bot-4
- Qianwen: qwen-turbo/qwen-plus/qwen-max/qwen-max-longcontext/qwen-vl-plus
diff --git a/config.example.yaml b/config.example.yaml
index 9f02b4c..08bf432 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -71,6 +71,10 @@ clients:
- type: qianwen
api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx
+ # See https://docs.anthropic.com/claude/reference/getting-started-with-the-api
+ - type: claude
+ api_key: xxx
+
# See https://cloud.google.com/vertex-ai
- type: vertexai
api_base: https://{REGION}-aiplatform.googleapis.com/v1/projects/{PROJECT_ID}/locations/{REGION}/publishers/google/models
diff --git a/src/client/claude.rs b/src/client/claude.rs
new file mode 100644
index 0000000..920a695
--- /dev/null
+++ b/src/client/claude.rs
@@ -0,0 +1,171 @@
+use super::{
+ patch_system_message, ClaudeClient, Client, ExtraConfig, Model, PromptType, SendData,
+ TokensCountFactors,
+};
+
+use crate::{render::ReplyHandler, utils::PromptKind};
+
+use anyhow::{anyhow, bail, Result};
+use async_trait::async_trait;
+use futures_util::StreamExt;
+use reqwest::{Client as ReqwestClient, RequestBuilder};
+use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt};
+use serde::Deserialize;
+use serde_json::{json, Value};
+
+const API_BASE: &str = "https://api.anthropic.com/v1/messages";
+
+const MODELS: [(&str, usize, &str); 3] = [
+ ("claude-2.1", 204096, "text"),
+ ("claude-2.0", 104096, "text"),
+ ("claude-instant-1.2", 104096, "text"),
+];
+
+const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2);
+
+#[derive(Debug, Clone, Deserialize)]
+pub struct ClaudeConfig {
+ pub name: Option<String>,
+ pub api_key: Option<String>,
+ pub extra: Option<ExtraConfig>,
+}
+
+#[async_trait]
+impl Client for ClaudeClient {
+ client_common_fns!();
+
+ async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> {
+ let builder = self.request_builder(client, data)?;
+ send_message(builder).await
+ }
+
+ async fn send_message_streaming_inner(
+ &self,
+ client: &ReqwestClient,
+ handler: &mut ReplyHandler,
+ data: SendData,
+ ) -> Result<()> {
+ let builder = self.request_builder(client, data)?;
+ send_message_streaming(builder, handler).await
+ }
+}
+
+impl ClaudeClient {
+ config_get_fn!(api_key, get_api_key);
+
+ pub const PROMPTS: [PromptType<'static>; 1] =
+ [("api_key", "API Key:", false, PromptKind::String)];
+
+ pub fn list_models(local_config: &ClaudeConfig) -> Vec<Model> {
+ let client_name = Self::name(local_config);
+ MODELS
+ .into_iter()
+ .map(|(name, max_tokens, capabilities)| {
+ Model::new(client_name, name)
+ .set_capabilities(capabilities.into())
+ .set_max_tokens(Some(max_tokens))
+ .set_tokens_count_factors(TOKENS_COUNT_FACTORS)
+ })
+ .collect()
+ }
+
+ fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key().ok();
+
+ let body = build_body(data, self.model.name.clone());
+
+ let url = API_BASE;
+
+ debug!("Claude Request: {url} {body}");
+
+ let mut builder = client.post(url).json(&body);
+ builder = builder.header("anthropic-version", "2023-06-01");
+ if let Some(api_key) = api_key {
+ builder = builder.header("x-api-key", api_key)
+ }
+
+ Ok(builder)
+ }
+}
+
+async fn send_message(builder: RequestBuilder) -> Result<String> {
+ let data: Value = builder.send().await?.json().await?;
+ check_error(&data)?;
+
+ let output = data["content"][0]["text"]
+ .as_str()
+ .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+
+ Ok(output.to_string())
+}
+
+async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> {
+ let mut es = builder.eventsource()?;
+ while let Some(event) = es.next().await {
+ match event {
+ Ok(Event::Open) => {}
+ Ok(Event::Message(message)) => {
+ let data: Value = serde_json::from_str(&message.data)?;
+ check_error(&data)?;
+ if let Some(typ) = data["type"].as_str() {
+ if typ == "content_block_delta" {
+ if let Some(text) = data["delta"]["text"].as_str() {
+ handler.text(text)?;
+ }
+ }
+ }
+ }
+ Err(err) => {
+ match err {
+ EventSourceError::StreamEnded => {}
+ EventSourceError::InvalidStatusCode(code, res) => {
+ let data: Value = res.json().await?;
+ check_error(&data)?;
+ bail!("Invalid status code: {code}");
+ }
+ _ => {
+ bail!("{}", err);
+ }
+ }
+ es.close();
+ }
+ }
+ }
+
+ Ok(())
+}
+
+fn build_body(data: SendData, model: String) -> Value {
+ let SendData {
+ mut messages,
+ temperature,
+ stream,
+ } = data;
+
+ patch_system_message(&mut messages);
+
+ let mut body = json!({
+ "model": model,
+ "max_tokens": 4096,
+ "messages": messages,
+ });
+
+ if let Some(v) = temperature {
+ body["temperature"] = (v / 2.0).into();
+ }
+ if stream {
+ body["stream"] = true.into();
+ }
+ body
+}
+
+fn check_error(data: &Value) -> Result<()> {
+ if let Some(error) = data["error"].as_object() {
+ if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
+ bail!("{typ}: {message}");
+ } else {
+ bail!("{}", Value::Object(error.clone()))
+ }
+ }
+ Ok(())
+}
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 320ddc4..1517b93 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -10,6 +10,7 @@ pub use model::*;
register_client!(
(openai, "openai", OpenAIConfig, OpenAIClient),
(gemini, "gemini", GeminiConfig, GeminiClient),
+ (claude, "claude", ClaudeConfig, ClaudeClient),
(localai, "localai", LocalAIConfig, LocalAIClient),
(ollama, "ollama", OllamaConfig, OllamaClient),
(