From 4db9b309803796bc5f996d0b3713344eb44207ec Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 25 Apr 2024 10:15:54 +0800 Subject: refactor: rewrite list models of all clients (#436) --- src/client/claude.rs | 38 +++++++++++++++++++------------------- 1 file changed, 19 insertions(+), 19 deletions(-) (limited to 'src/client/claude.rs') diff --git a/src/client/claude.rs b/src/client/claude.rs index 054731c..68ab509 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -15,13 +15,6 @@ use serde_json::{json, Value}; const API_BASE: &str = "https://api.anthropic.com/v1/messages"; -const MODELS: [(&str, usize, &str); 3] = [ - // https://docs.anthropic.com/claude/docs/models-overview - ("claude-3-opus-20240229", 200000, "text,vision"), - ("claude-3-sonnet-20240229", 200000, "text,vision"), - ("claude-3-haiku-20240307", 200000, "text,vision"), -]; - #[derive(Debug, Clone, Deserialize)] pub struct ClaudeConfig { pub name: Option, @@ -52,7 +45,16 @@ impl Client for ClaudeClient { } impl ClaudeClient { - list_models_fn!(ClaudeConfig, &MODELS); + list_models_fn!( + ClaudeConfig, + [ + // https://docs.anthropic.com/claude/docs/models-overview + ("claude-3-opus-20240229", "text,vision", 200000, 4096), + ("claude-3-sonnet-20240229", "text,vision", 200000, 4096), + ("claude-3-haiku-20240307", "text,vision", 200000, 4096), + ] + ); + config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = @@ -61,7 +63,7 @@ impl ClaudeClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let api_key = self.get_api_key().ok(); - let body = build_body(data, &self.model)?; + let body = claude_build_body(data, &self.model)?; let url = API_BASE; @@ -136,7 +138,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand Ok(()) } -fn build_body(data: SendData, model: &Model) -> Result { +pub fn claude_build_body(data: SendData, model: &Model) -> Result { let SendData { mut messages, temperature, @@ -191,18 +193,16 @@ fn build_body(data: SendData, model: &Model) -> Result { ); } - let max_tokens = model.max_output_tokens.unwrap_or(4096); - let mut body = json!({ "model": &model.name, - "max_tokens": max_tokens, "messages": messages, }); - - if let Some(system) = system_message { - body["system"] = system.into(); + if let Some(v) = system_message { + body["system"] = v.into(); + } + if let Some(v) = model.max_output_tokens { + body["max_tokens"] = v.into(); } - if let Some(v) = temperature { body["temperature"] = v.into(); } @@ -218,8 +218,8 @@ fn build_body(data: SendData, model: &Model) -> Result { fn catch_error(data: &Value, status: u16) -> Result<()> { debug!("Invalid response, status: {status}, data: {data}"); if let Some(error) = data["error"].as_object() { - if let (Some(type_), Some(message)) = (error["type"].as_str(), error["message"].as_str()) { - bail!("{message} (type: {type_})"); + if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) { + bail!("{message} (type: {typ})"); } } bail!("Invalid response, status: {status}, data: {data}"); -- cgit v1.2.3