summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/ernie.rs6
-rw-r--r--src/client/message.rs30
-rw-r--r--src/client/vertexai.rs11
3 files changed, 41 insertions, 6 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 21362c1..428d263 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -238,7 +238,7 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu
stream,
} = data;
- patch_system_message(&mut messages);
+ let system_message = extract_system_message(&mut messages);
let messages: Vec<Value> = messages
.into_iter()
@@ -269,6 +269,10 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu
"messages": messages,
});
+ if let Some(v) = system_message {
+ body["system"] = v.into();
+ }
+
if let Some(v) = model.max_tokens_param() {
body["max_output_tokens"] = v.into();
}
diff --git a/src/client/message.rs b/src/client/message.rs
index d7ba698..77adf47 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -21,6 +21,30 @@ impl Message {
pub fn new(role: MessageRole, content: MessageContent) -> Self {
Self { role, content }
}
+
+ pub fn merge_system(&mut self, system: &str) {
+ match &mut self.content {
+ MessageContent::Text(text) => {
+ self.content = MessageContent::Array(vec![
+ MessageContentPart::Text {
+ text: system.to_string(),
+ },
+ MessageContentPart::Text {
+ text: text.to_string(),
+ },
+ ]);
+ }
+ MessageContent::Array(list) => {
+ list.insert(
+ 0,
+ MessageContentPart::Text {
+ text: system.to_string(),
+ },
+ );
+ }
+ _ => {}
+ }
+ }
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
@@ -124,12 +148,10 @@ pub struct ImageUrl {
pub fn patch_system_message(messages: &mut Vec<Message>) {
if messages[0].role.is_system() {
let system_message = messages.remove(0);
- if let (Some(message), MessageContent::Text(system_text)) =
+ if let (Some(message), MessageContent::Text(system)) =
(messages.get_mut(0), system_message.content)
{
- if let MessageContent::Text(text) = message.content.clone() {
- message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
- }
+ message.merge_system(&system);
}
}
}
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 16910b6..cc7e464 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -262,7 +262,12 @@ pub fn gemini_build_chat_completions_body(
stream: _,
} = data;
- patch_system_message(&mut messages);
+ let system_message = if model.name().starts_with("gemini-1.5-") {
+ extract_system_message(&mut messages)
+ } else {
+ patch_system_message(&mut messages);
+ None
+ };
let mut network_image_urls = vec![];
let contents: Vec<Value> = messages
@@ -333,6 +338,10 @@ pub fn gemini_build_chat_completions_body(
let mut body = json!({ "contents": contents, "generationConfig": {} });
+ if let Some(v) = system_message {
+ body["systemInstruction"] = json!({ "parts": [{"text": v }] });
+ }
+
if let Some(v) = model.max_tokens_param() {
body["generationConfig"]["maxOutputTokens"] = v.into();
}