summaryrefslogtreecommitdiffstats
path: root/src/client/claude.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-05 06:52:47 +0800
committerGitHub <noreply@github.com>2024-03-05 06:52:47 +0800
commitbe4e5e569a61c54d8ac8fb77144e9b1d01e3b81f (patch)
tree9823481746dd81dca8e2fb80ef91166a8c3a9a54 /src/client/claude.rs
parent3f693ea060d96b0adc397768c3b2dd47708a20ce (diff)
downloadaichat-be4e5e569a61c54d8ac8fb77144e9b1d01e3b81f.tar.gz
feat: support claude-3 (#336)
Diffstat (limited to 'src/client/claude.rs')
-rw-r--r--src/client/claude.rs61
1 files changed, 56 insertions, 5 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 920a695..325bce6 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -3,7 +3,11 @@ use super::{
TokensCountFactors,
};
-use crate::{render::ReplyHandler, utils::PromptKind};
+use crate::{
+ client::{ImageUrl, MessageContent, MessageContentPart},
+ render::ReplyHandler,
+ utils::PromptKind,
+};
use anyhow::{anyhow, bail, Result};
use async_trait::async_trait;
@@ -15,7 +19,9 @@ use serde_json::{json, Value};
const API_BASE: &str = "https://api.anthropic.com/v1/messages";
-const MODELS: [(&str, usize, &str); 3] = [
+const MODELS: [(&str, usize, &str); 5] = [
+ ("claude-3-opus-20240229", 204096, "text,vision"),
+ ("claude-3-sonnet-20240229", 204096, "text,vision"),
("claude-2.1", 204096, "text"),
("claude-2.0", 104096, "text"),
("claude-instant-1.2", 104096, "text"),
@@ -72,7 +78,7 @@ impl ClaudeClient {
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 body = build_body(data, self.model.name.clone())?;
let url = API_BASE;
@@ -135,7 +141,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand
Ok(())
}
-fn build_body(data: SendData, model: String) -> Value {
+fn build_body(data: SendData, model: String) -> Result<Value> {
let SendData {
mut messages,
temperature,
@@ -144,6 +150,51 @@ fn build_body(data: SendData, model: String) -> Value {
patch_system_message(&mut messages);
+ let mut network_image_urls = vec![];
+ let messages: Vec<Value> = messages
+ .into_iter()
+ .map(|message| {
+ let role = message.role;
+ let content = match message.content {
+ MessageContent::Text(text) => vec![json!({"type": "text", "text": text})],
+ MessageContent::Array(list) => list
+ .into_iter()
+ .map(|item| match item {
+ MessageContentPart::Text { text } => json!({"type": "text", "text": text}),
+ MessageContentPart::ImageUrl {
+ image_url: ImageUrl { url },
+ } => {
+ if let Some((mime_type, data)) = url
+ .strip_prefix("data:")
+ .and_then(|v| v.split_once(";base64,"))
+ {
+ json!({
+ "type": "image",
+ "source": {
+ "type": "base64",
+ "media_type": mime_type,
+ "data": data,
+ }
+ })
+ } else {
+ network_image_urls.push(url.clone());
+ json!({ "url": url })
+ }
+ }
+ })
+ .collect(),
+ };
+ json!({ "role": role, "content": content })
+ })
+ .collect();
+
+ if !network_image_urls.is_empty() {
+ bail!(
+ "The model does not support network images: {:?}",
+ network_image_urls
+ );
+ }
+
let mut body = json!({
"model": model,
"max_tokens": 4096,
@@ -156,7 +207,7 @@ fn build_body(data: SendData, model: String) -> Value {
if stream {
body["stream"] = true.into();
}
- body
+ Ok(body)
}
fn check_error(data: &Value) -> Result<()> {