summaryrefslogtreecommitdiffstats
path: root/src/client/qianwen.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/qianwen.rs')
-rw-r--r--src/client/qianwen.rs37
1 files changed, 25 insertions, 12 deletions
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 7391a38..b33d13e 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,6 +1,6 @@
use super::{
- maybe_catch_error, message::*, sse_stream, Client, CompletionDetails, ExtraConfig, Model,
- ModelConfig, PromptAction, PromptKind, QianwenClient, SendData, SsMmessage, SseHandler,
+ maybe_catch_error, message::*, sse_stream, Client, CompletionOutput, ExtraConfig, Model,
+ ModelData, PromptAction, PromptKind, QianwenClient, SendData, SsMmessage, SseHandler,
};
use crate::utils::{base64_decode, sha256};
@@ -26,7 +26,7 @@ pub struct QianwenConfig {
pub name: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -62,7 +62,7 @@ impl QianwenClient {
}
fn is_vl(&self) -> bool {
- self.model.name.starts_with("qwen-vl")
+ self.model.name().starts_with("qwen-vl")
}
}
@@ -74,9 +74,9 @@ impl Client for QianwenClient {
&self,
client: &ReqwestClient,
mut data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
let api_key = self.get_api_key()?;
- patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
+ patch_messages(self.model.name(), &api_key, &mut data.messages).await?;
let builder = self.request_builder(client, data)?;
send_message(builder, self.is_vl()).await
}
@@ -88,16 +88,17 @@ impl Client for QianwenClient {
mut data: SendData,
) -> Result<()> {
let api_key = self.get_api_key()?;
- patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
+ patch_messages(self.model.name(), &api_key, &mut data.messages).await?;
let builder = self.request_builder(client, data)?;
send_message_streaming(builder, handler, self.is_vl()).await
}
}
-async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<(String, CompletionDetails)> {
+async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<CompletionOutput> {
let data: Value = builder.send().await?.json().await?;
maybe_catch_error(&data)?;
+ debug!("non-stream-data: {data}");
extract_completion_text(&data, is_vl)
}
@@ -109,6 +110,7 @@ async fn send_message_streaming(
let handle = |message: SsMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
maybe_catch_error(&data)?;
+ debug!("stream-data: {data}");
if is_vl {
if let Some(text) =
data["output"]["choices"][0]["message"]["content"][0]["text"].as_str()
@@ -129,10 +131,12 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
let mut has_upload = false;
+ let mut is_tool_call = false;
let input = if is_vl {
let messages: Vec<Value> = messages
.into_iter()
@@ -154,6 +158,10 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
}
})
.collect(),
+ MessageContent::ToolResults(_) => {
+ is_tool_call = true;
+ vec![]
+ }
};
json!({ "role": role, "content": content })
})
@@ -167,6 +175,9 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
"messages": messages,
})
};
+ if is_tool_call {
+ bail!("The client does not support function calling",);
+ }
let mut parameters = json!({});
if stream {
@@ -184,7 +195,7 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
}
let body = json!({
- "model": &model.name,
+ "model": &model.name(),
"input": input,
"parameters": parameters
});
@@ -192,7 +203,7 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
Ok((body, has_upload))
}
-fn extract_completion_text(data: &Value, is_vl: bool) -> Result<(String, CompletionDetails)> {
+fn extract_completion_text(data: &Value, is_vl: bool) -> Result<CompletionOutput> {
let err = || anyhow!("Invalid response data: {data}");
let text = if is_vl {
data["output"]["choices"][0]["message"]["content"][0]["text"]
@@ -201,13 +212,15 @@ fn extract_completion_text(data: &Value, is_vl: bool) -> Result<(String, Complet
} else {
data["output"]["text"].as_str().ok_or_else(err)?
};
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls: vec![],
id: data["request_id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["input_tokens"].as_u64(),
output_tokens: data["usage"]["output_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
/// Patch messages, upload embedded images to oss