From 9c6c9f10a27d0993636b453f39d8934c95c5c2b2 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 23 Apr 2024 18:14:47 +0800 Subject: feat: builtin models can be overwrited by models config (#429) --- src/client/claude.rs | 28 ++++++++-------------------- 1 file changed, 8 insertions(+), 20 deletions(-) (limited to 'src/client/claude.rs') diff --git a/src/client/claude.rs b/src/client/claude.rs index d87e64a..24bffbf 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,6 +1,6 @@ use super::{ patch_system_message, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent, - MessageContentPart, Model, PromptType, ReplyHandler, SendData, + MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -15,17 +15,19 @@ use serde_json::{json, Value}; const API_BASE: &str = "https://api.anthropic.com/v1/messages"; -const MODELS: [(&str, usize, isize, &str); 3] = [ +const MODELS: [(&str, usize, &str); 3] = [ // https://docs.anthropic.com/claude/docs/models-overview - ("claude-3-opus-20240229", 200000, 4096, "text,vision"), - ("claude-3-sonnet-20240229", 200000, 4096, "text,vision"), - ("claude-3-haiku-20240307", 200000, 4096, "text,vision"), + ("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, pub api_key: Option, + #[serde(default)] + pub models: Vec, pub extra: Option, } @@ -50,26 +52,12 @@ impl Client for ClaudeClient { } impl ClaudeClient { + list_models_fn!(ClaudeConfig, &MODELS); 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 { - let client_name = Self::name(local_config); - MODELS - .into_iter() - .map( - |(name, max_input_tokens, max_output_tokens, capabilities)| { - Model::new(client_name, name) - .set_capabilities(capabilities.into()) - .set_max_input_tokens(Some(max_input_tokens)) - .set_max_output_tokens(Some(max_output_tokens)) - }, - ) - .collect() - } - fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let api_key = self.get_api_key().ok(); -- cgit v1.2.3