diff options
| author | sigoden <sigoden@gmail.com> | 2024-03-05 06:52:47 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-03-05 06:52:47 +0800 |
| commit | be4e5e569a61c54d8ac8fb77144e9b1d01e3b81f (patch) | |
| tree | 9823481746dd81dca8e2fb80ef91166a8c3a9a54 /src/client/claude.rs | |
| parent | 3f693ea060d96b0adc397768c3b2dd47708a20ce (diff) | |
| download | aichat-be4e5e569a61c54d8ac8fb77144e9b1d01e3b81f.tar.gz | |
feat: support claude-3 (#336)
Diffstat (limited to 'src/client/claude.rs')
| -rw-r--r-- | src/client/claude.rs | 61 |
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<()> { |
