diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-01 06:01:10 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-01 06:01:10 +0800 |
| commit | 5eae392dbd4fc789829ac2f47a24614dbc5db5e9 (patch) | |
| tree | 331aded9b1badd5b66eb6b7afbbbe8f8227b87b8 /src/client/bedrock.rs | |
| parent | 8aa18b59f5be467fcadc9aa679bce58d66e62793 (diff) | |
| download | aichat-5eae392dbd4fc789829ac2f47a24614dbc5db5e9.tar.gz | |
refactore: add models for openai-compatible platforms (#471)
Diffstat (limited to 'src/client/bedrock.rs')
| -rw-r--r-- | src/client/bedrock.rs | 12 |
1 files changed, 4 insertions, 8 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index fb4dad8..d6ab696 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -2,7 +2,7 @@ use super::claude::{claude_build_body, claude_extract_completion}; use super::{ catch_error, generate_prompt, BedrockClient, Client, CompletionDetails, ExtraConfig, Model, ModelConfig, PromptAction, PromptFormat, PromptKind, SendData, SseHandler, - LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT, + LLAMA3_PROMPT_FORMAT, MISTRAL_PROMPT_FORMAT, }; use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256}; @@ -140,7 +140,7 @@ async fn send_message( match model_category { ModelCategory::Anthropic => claude_extract_completion(&data), - ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => llama_extract_completion(&data), + ModelCategory::MetaLlama3 => llama_extract_completion(&data), ModelCategory::Mistral => mistral_extrat_completion(&data), } } @@ -183,7 +183,7 @@ async fn send_message_streaming( } } } - ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => { + ModelCategory::MetaLlama3 => { if let Some(text) = data["generation"].as_str() { handler.text(text)?; } @@ -220,7 +220,6 @@ fn build_body(data: SendData, model: &Model, model_category: &ModelCategory) -> body["anthropic_version"] = "bedrock-2023-05-31".into(); Ok(body) } - ModelCategory::MetaLlama2 => meta_llama_build_body(data, model, LLAMA2_PROMPT_FORMAT), ModelCategory::MetaLlama3 => meta_llama_build_body(data, model, LLAMA3_PROMPT_FORMAT), ModelCategory::Mistral => mistral_build_body(data, model), } @@ -256,7 +255,7 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> { top_p, stream: _, } = data; - let prompt = generate_prompt(&messages, LLAMA2_PROMPT_FORMAT)?; + let prompt = generate_prompt(&messages, MISTRAL_PROMPT_FORMAT)?; let mut body = json!({ "prompt": prompt }); if let Some(v) = model.max_output_tokens { @@ -294,7 +293,6 @@ fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionDetails) #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ModelCategory { Anthropic, - MetaLlama2, MetaLlama3, Mistral, } @@ -305,8 +303,6 @@ impl FromStr for ModelCategory { fn from_str(s: &str) -> std::result::Result<Self, Self::Err> { if s.starts_with("anthropic.") { Ok(ModelCategory::Anthropic) - } else if s.starts_with("meta.llama2") { - Ok(ModelCategory::MetaLlama2) } else if s.starts_with("meta.llama3") { Ok(ModelCategory::MetaLlama3) } else if s.starts_with("mistral") { |
