summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-26 21:06:11 +0800
committerGitHub <noreply@github.com>2025-02-26 21:06:11 +0800
commitcb54f8643a965fbc35b1f07eb3ea026692eb8091 (patch)
tree719617569d379e4c34badbff95a77d0771df7213
parent6b28b1c4fcf57edc4375703bca34eeafa597c905 (diff)
downloadaichat-cb54f8643a965fbc35b1f07eb3ea026692eb8091.tar.gz
refactor: several improvements (#1203)
-rw-r--r--models.yaml40
-rw-r--r--src/client/bedrock.rs17
-rw-r--r--src/client/claude.rs9
-rw-r--r--src/client/openai.rs6
-rw-r--r--src/client/vertexai.rs5
5 files changed, 58 insertions, 19 deletions
diff --git a/models.yaml b/models.yaml
index 97e26e0..272da39 100644
--- a/models.yaml
+++ b/models.yaml
@@ -124,7 +124,7 @@
output_price: 0
supports_vision: true
supports_function_calling: true
- - name: gemini-2.0-flash-lite-preview
+ - name: gemini-2.0-flash-lite
max_input_tokens: 1048576
max_output_tokens: 8192
input_price: 0
@@ -200,7 +200,7 @@
thinking:
type: enabled
budget_tokens: 16000
- - name: claude-3-5-sonnet-latest
+ - name: claude-3-5-sonnet-20241022
max_input_tokens: 200000
max_output_tokens: 8192
require_max_tokens: true
@@ -208,7 +208,7 @@
output_price: 15
supports_vision: true
supports_function_calling: true
- - name: claude-3-5-sonnet-20241022
+ - name: claude-3-5-sonnet-20240620
max_input_tokens: 200000
max_output_tokens: 8192
require_max_tokens: true
@@ -216,14 +216,6 @@
output_price: 15
supports_vision: true
supports_function_calling: true
- - name: claude-3-5-haiku-latest
- max_input_tokens: 200000
- max_output_tokens: 8192
- require_max_tokens: true
- input_price: 0.8
- output_price: 4
- supports_vision: true
- supports_function_calling: true
- name: claude-3-5-haiku-20241022
max_input_tokens: 200000
max_output_tokens: 8192
@@ -499,7 +491,7 @@
output_price: 0.6
supports_vision: true
supports_function_calling: true
- - name: gemini-2.0-flash-lite-preview-02-05
+ - name: gemini-2.0-flash-lite-001
max_input_tokens: 1048576
max_output_tokens: 8192
input_price: 0.075
@@ -1341,6 +1333,30 @@
output_price: 0.4
supports_vision: true
supports_function_calling: true
+ - name: google/gemini-2.0-flash-lite-001
+ max_input_tokens: 1048576
+ input_price: 0.075
+ output_price: 0.3
+ supports_vision: true
+ supports_function_calling: true
+ - name: anthropic/claude-3.7-sonnet
+ max_input_tokens: 200000
+ max_output_tokens: 8192
+ require_max_tokens: true
+ input_price: 3
+ output_price: 15
+ supports_vision: true
+ supports_function_calling: true
+ - name: anthropic/claude-3.7-sonnet:thinking
+ max_input_tokens: 200000
+ max_output_tokens: 24000
+ require_max_tokens: true
+ input_price: 3
+ output_price: 15
+ supports_vision: true
+ patch:
+ body:
+ include_reasoning: true
- name: anthropic/claude-3.5-sonnet
max_input_tokens: 200000
max_output_tokens: 8192
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();