diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-26 21:06:11 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-26 21:06:11 +0800 |
| commit | cb54f8643a965fbc35b1f07eb3ea026692eb8091 (patch) | |
| tree | 719617569d379e4c34badbff95a77d0771df7213 /src | |
| parent | 6b28b1c4fcf57edc4375703bca34eeafa597c905 (diff) | |
| download | aichat-cb54f8643a965fbc35b1f07eb3ea026692eb8091.tar.gz | |
refactor: several improvements (#1203)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/bedrock.rs | 17 | ||||
| -rw-r--r-- | src/client/claude.rs | 9 | ||||
| -rw-r--r-- | src/client/openai.rs | 6 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 5 |
4 files changed, 30 insertions, 7 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 089f1c2..0dc84ca 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,6 +1,6 @@ use super::*; -use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256}; +use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256, strip_think_tag}; use anyhow::{bail, Context, Result}; use aws_smithy_eventstream::frame::{DecodedFrame, MessageFrameDecoder}; @@ -241,7 +241,9 @@ async fn chat_completions_streaming( "contentBlockDelta" => { if let Some(text) = data["delta"]["text"].as_str() { handler.text(text)?; - } else if let Some(text) = data["delta"]["reasoningContent"]["text"].as_str() { + } else if let Some(text) = + data["delta"]["reasoningContent"]["text"].as_str() + { if reasoning_state == 0 { handler.text("<think>\n")?; reasoning_state = 1; @@ -317,11 +319,16 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu let mut network_image_urls = vec![]; + let messages_len = messages.len(); let messages: Vec<Value> = messages .into_iter() - .flat_map(|message| { + .enumerate() + .flat_map(|(i, message)| { let Message { role, content } = message; match content { + MessageContent::Text(text) if role.is_assistant() && i != messages_len - 1 => { + vec![json!({ "role": role, "content": [ { "text": strip_think_tag(&text) } ] })] + } MessageContent::Text(text) => vec![json!({ "role": role, "content": [ @@ -469,7 +476,9 @@ fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> { text.push_str("\n\n"); } text.push_str(v); - } else if let Some(reasoning_text) = item["reasoningContent"]["reasoningText"].as_object() { + } else if let Some(reasoning_text) = + item["reasoningContent"]["reasoningText"].as_object() + { if let Some(text) = json_str_from_map(reasoning_text, "text") { reasoning = Some(text.to_string()); } diff --git a/src/client/claude.rs b/src/client/claude.rs index 202d2f7..4b77870 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,5 +1,7 @@ use super::*; +use crate::utils::strip_think_tag; + use anyhow::{bail, Context, Result}; use reqwest::RequestBuilder; use serde::Deserialize; @@ -169,11 +171,16 @@ pub fn claude_build_chat_completions_body( let mut network_image_urls = vec![]; + let messages_len = messages.len(); let messages: Vec<Value> = messages .into_iter() - .flat_map(|message| { + .enumerate() + .flat_map(|(i, message)| { let Message { role, content } = message; match content { + MessageContent::Text(text) if role.is_assistant() && i != messages_len - 1 => { + vec![json!({ "role": role, "content": strip_think_tag(&text) })] + } MessageContent::Text(text) => vec![json!({ "role": role, "content": text, diff --git a/src/client/openai.rs b/src/client/openai.rs index c002f52..e5f4d23 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -300,7 +300,11 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod }); if let Some(v) = model.max_tokens_param() { - if model.patch().and_then(|v| v.get("body").and_then(|v| v.get("max_tokens"))) == Some(&Value::Null) { + if model + .patch() + .and_then(|v| v.get("body").and_then(|v| v.get("max_tokens"))) + == Some(&Value::Null) + { body["max_completion_tokens"] = v.into(); } else { body["max_tokens"] = v.into(); diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 7cf5a1f..730c5d8 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -154,7 +154,10 @@ fn prepare_embeddings(self_: &VertexAIClient, data: &EmbeddingsData) -> Result<R let access_token = get_access_token(self_.name())?; let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers"); - let url = format!("{base_url}/google/models/{}:predict", self_.model.real_name()); + let url = format!( + "{base_url}/google/models/{}:predict", + self_.model.real_name() + ); let instances: Vec<_> = data.texts.iter().map(|v| json!({"content": v})).collect(); |
