use async_openai::config::OpenAIConfig;
use async_openai::types::chat::{
ChatCompletionMessageToolCalls, ChatCompletionRequestMessage,
ChatCompletionRequestToolMessage, ChatCompletionResponseStream,
ChatCompletionTools, CreateChatCompletionRequest, CreateChatCompletionResponse,
};
use crate::config::{resolve_profile, ResolvedModel, RobitConfig};
use crate::error::LlmError;
fn validate_and_filter_messages(mut messages: Vec<ChatCompletionRequestMessage>) -> Vec<ChatCompletionRequestMessage> {
let original_len = messages.len();
messages.retain(|msg| {
match msg {
ChatCompletionRequestMessage::Assistant(assistant_msg) => {
let has_content = assistant_msg.content.is_some();
let has_tool_calls = assistant_msg.tool_calls.is_some();
if !has_content && !has_tool_calls {
tracing::warn!("Filtering out invalid assistant message (has neither content nor tool_calls)");
false
} else {
true
}
}
_ => true
}
});
let filtered_len = messages.len();
if filtered_len < original_len {
tracing::info!("Filtered {} invalid messages from history", original_len - filtered_len);
}
messages
}
fn repair_tool_pairing(
messages: Vec<ChatCompletionRequestMessage>,
) -> Vec<ChatCompletionRequestMessage> {
use std::collections::HashSet;
let mut open_ids: HashSet<String> = HashSet::new();
let mut keep: Vec<bool> = Vec::with_capacity(messages.len());
let mut dropped = 0usize;
for msg in &messages {
match msg {
ChatCompletionRequestMessage::Assistant(a) => {
if let Some(tool_calls) = &a.tool_calls {
for tc in tool_calls {
if let ChatCompletionMessageToolCalls::Function(f) = tc {
open_ids.insert(f.id.clone());
}
}
}
keep.push(true);
}
ChatCompletionRequestMessage::Tool(t) => {
if open_ids.remove(&t.tool_call_id) {
keep.push(true);
} else {
tracing::trace!(
"repair_tool_pairing: dropping orphaned tool message \
(tool_call_id='{}' not declared by any preceding assistant tool_calls)",
t.tool_call_id
);
keep.push(false);
dropped += 1;
}
}
_ => keep.push(true),
}
}
let mut missing = open_ids;
if dropped == 0 && missing.is_empty() {
return messages; }
if dropped > 0 {
tracing::warn!(
"repair_tool_pairing: dropped {} orphaned tool message(s) not declared by any assistant tool_calls",
dropped
);
}
let mut synthesized = 0usize;
let mut result: Vec<ChatCompletionRequestMessage> = Vec::with_capacity(messages.len());
for (msg, keep) in messages.into_iter().zip(keep.into_iter()) {
if !keep {
continue;
}
let missing_here: Vec<String> = match &msg {
ChatCompletionRequestMessage::Assistant(a) => a
.tool_calls
.as_ref()
.map(|tool_calls| {
tool_calls
.iter()
.filter_map(|tc| {
if let ChatCompletionMessageToolCalls::Function(f) = tc {
if missing.remove(&f.id) {
Some(f.id.clone())
} else {
None
}
} else {
None
}
})
.collect()
})
.unwrap_or_default(),
_ => Vec::new(),
};
result.push(msg);
for id in missing_here {
tracing::trace!(
"repair_tool_pairing: synthesizing missing tool response for tool_call_id='{}'",
id
);
synthesized += 1;
result.push(ChatCompletionRequestMessage::Tool(
ChatCompletionRequestToolMessage {
content: "[Tool result unavailable — session history was restored without this result]"
.to_string()
.into(),
tool_call_id: id,
}
.into(),
));
}
}
if synthesized > 0 {
tracing::warn!(
"repair_tool_pairing: synthesized {} missing tool response(s) for declared tool_calls",
synthesized
);
}
result
}
fn repair_interleaved_tool_responses(
messages: Vec<ChatCompletionRequestMessage>,
) -> Vec<ChatCompletionRequestMessage> {
use std::collections::HashSet;
{
let mut pending: HashSet<&String> = HashSet::new();
let mut interleaved = false;
'scan: for msg in &messages {
match msg {
ChatCompletionRequestMessage::Assistant(a) => {
if let Some(tool_calls) = &a.tool_calls {
pending = tool_calls
.iter()
.filter_map(|tc| {
if let ChatCompletionMessageToolCalls::Function(f) = tc {
Some(&f.id)
} else {
None
}
})
.collect();
}
}
ChatCompletionRequestMessage::Tool(t) => {
pending.remove(&t.tool_call_id);
}
_ => {
if !pending.is_empty() {
interleaved = true;
break 'scan;
}
}
}
}
if !interleaved {
return messages;
}
}
tracing::warn!(
"repair_interleaved_tool_responses: moving non-tool message(s) out of a \
tool_calls → tool-responses window (providers reject interleaved messages \
with a 400 error)"
);
let mut result: Vec<ChatCompletionRequestMessage> = Vec::with_capacity(messages.len());
let mut deferred: Vec<ChatCompletionRequestMessage> = Vec::new();
let mut pending: HashSet<String> = HashSet::new();
for msg in messages {
match &msg {
ChatCompletionRequestMessage::Assistant(a) => {
if let Some(tool_calls) = &a.tool_calls {
if !deferred.is_empty() {
result.append(&mut deferred);
}
pending = tool_calls
.iter()
.filter_map(|tc| {
if let ChatCompletionMessageToolCalls::Function(f) = tc {
Some(f.id.clone())
} else {
None
}
})
.collect();
result.push(msg);
} else if pending.is_empty() {
result.push(msg);
} else {
deferred.push(msg);
}
}
ChatCompletionRequestMessage::Tool(t) => {
let responded = pending.remove(&t.tool_call_id);
result.push(msg);
if responded && pending.is_empty() && !deferred.is_empty() {
result.append(&mut deferred);
}
}
_ => {
if pending.is_empty() {
result.push(msg);
} else {
deferred.push(msg);
}
}
}
}
result.append(&mut deferred);
result
}
pub struct LlmClient {
client: async_openai::Client<OpenAIConfig>,
model: String,
resolved: ResolvedModel,
}
impl LlmClient {
pub fn from_config(
config: &RobitConfig,
profile_name: Option<&str>,
) -> Result<Self, LlmError> {
let resolved = resolve_profile(config, profile_name)?;
let oc = OpenAIConfig::new()
.with_api_base(&resolved.base_url)
.with_api_key(&resolved.api_key);
let client = async_openai::Client::with_config(oc);
Ok(Self {
client,
model: resolved.model_id.clone(),
resolved,
})
}
pub async fn chat_stream(
&self,
messages: Vec<ChatCompletionRequestMessage>,
tools: Option<Vec<ChatCompletionTools>>,
) -> Result<ChatCompletionResponseStream, LlmError> {
let messages = validate_and_filter_messages(messages);
let messages = repair_interleaved_tool_responses(messages);
let messages = repair_tool_pairing(messages);
let msg_count = messages.len();
tracing::trace!("Creating chat stream for model={}, messages={}", self.model, msg_count);
let request = CreateChatCompletionRequest {
model: self.model.clone(),
messages,
tools,
stream: Some(true),
stream_options: Some(async_openai::types::chat::ChatCompletionStreamOptions {
include_usage: Some(true),
include_obfuscation: None,
}),
max_completion_tokens: self.resolved.max_tokens,
temperature: self.resolved.temperature,
..Default::default()
};
let stream = self.client.chat().create_stream(request).await;
if let Err(e) = &stream {
tracing::error!("Chat stream creation failed: {:?}", e);
}
let stream = stream?;
Ok(stream)
}
pub async fn chat(
&self,
messages: Vec<ChatCompletionRequestMessage>,
tools: Option<Vec<ChatCompletionTools>>,
) -> Result<CreateChatCompletionResponse, LlmError> {
let messages = validate_and_filter_messages(messages);
let messages = repair_interleaved_tool_responses(messages);
let messages = repair_tool_pairing(messages);
let request = CreateChatCompletionRequest {
model: self.model.clone(),
messages,
tools,
max_completion_tokens: self.resolved.max_tokens,
temperature: self.resolved.temperature,
..Default::default()
};
let response = self.client.chat().create(request).await?;
Ok(response)
}
pub fn model(&self) -> &str {
&self.model
}
pub fn profile(&self) -> &str {
&self.resolved.profile_name
}
pub fn resolved(&self) -> &ResolvedModel {
&self.resolved
}
pub fn supports_images(&self) -> bool {
self.resolved.supports_images
}
pub fn supports_tools(&self) -> bool {
self.resolved.supports_tools
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_openai::types::chat::{
ChatCompletionMessageToolCall, ChatCompletionRequestAssistantMessage,
ChatCompletionRequestUserMessage, FunctionCall,
};
fn user_msg(text: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::User(ChatCompletionRequestUserMessage {
content: text.to_string().into(),
name: None,
})
}
fn assistant_text_msg(text: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::Assistant(ChatCompletionRequestAssistantMessage {
content: Some(text.to_string().into()),
name: None,
tool_calls: None,
refusal: None,
audio: None,
#[allow(deprecated)]
function_call: None,
})
}
fn assistant_tool_call_msg(id: &str, name: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::Assistant(ChatCompletionRequestAssistantMessage {
content: None,
name: None,
tool_calls: Some(vec![ChatCompletionMessageToolCalls::Function(
ChatCompletionMessageToolCall {
id: id.to_string(),
function: FunctionCall {
name: name.to_string(),
arguments: "{}".to_string(),
},
},
)]),
refusal: None,
audio: None,
#[allow(deprecated)]
function_call: None,
})
}
fn tool_msg(id: &str, text: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::Tool(ChatCompletionRequestToolMessage {
content: text.to_string().into(),
tool_call_id: id.to_string(),
})
}
#[test]
fn repair_drops_orphaned_tool_messages() {
let messages = vec![
user_msg("hello"),
assistant_text_msg("hi"),
user_msg("do something"),
tool_msg("call_1", "orphaned result"),
assistant_text_msg("done"),
];
let repaired = repair_tool_pairing(messages);
assert_eq!(repaired.len(), 4, "orphaned tool message should be dropped");
assert!(
!repaired
.iter()
.any(|m| matches!(m, ChatCompletionRequestMessage::Tool(_))),
"no tool messages should remain"
);
}
#[test]
fn repair_synthesizes_missing_tool_responses() {
let messages = vec![
user_msg("do something"),
assistant_tool_call_msg("call_1", "bash"),
user_msg("next question"),
];
let repaired = repair_tool_pairing(messages);
assert_eq!(repaired.len(), 4, "placeholder tool response should be added");
assert!(matches!(repaired[2], ChatCompletionRequestMessage::Tool(_)));
if let ChatCompletionRequestMessage::Tool(t) = &repaired[2] {
assert_eq!(t.tool_call_id, "call_1");
}
}
#[test]
fn repair_keeps_valid_pairing_untouched() {
let messages = vec![
user_msg("do something"),
assistant_tool_call_msg("call_1", "bash"),
tool_msg("call_1", "ok"),
assistant_text_msg("done"),
];
let repaired = repair_tool_pairing(messages.clone());
assert_eq!(repaired.len(), messages.len(), "valid history must not change");
}
#[test]
fn repair_handles_mixed_valid_and_orphaned() {
let messages = vec![
user_msg("a"),
assistant_tool_call_msg("call_1", "read"),
tool_msg("call_1", "result 1"), tool_msg("call_ghost", "ghost result"), assistant_text_msg("done"),
];
let repaired = repair_tool_pairing(messages);
assert_eq!(repaired.len(), 4);
let tool_ids: Vec<&str> = repaired
.iter()
.filter_map(|m| {
if let ChatCompletionRequestMessage::Tool(t) = m {
Some(t.tool_call_id.as_str())
} else {
None
}
})
.collect();
assert_eq!(tool_ids, vec!["call_1"]);
}
fn assistant_multi_tool_call_msg(
calls: &[(&str, &str)],
) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::Assistant(ChatCompletionRequestAssistantMessage {
content: None,
name: None,
tool_calls: Some(
calls
.iter()
.map(|(id, name)| {
ChatCompletionMessageToolCalls::Function(ChatCompletionMessageToolCall {
id: id.to_string(),
function: FunctionCall {
name: name.to_string(),
arguments: "{}".to_string(),
},
})
})
.collect(),
),
refusal: None,
audio: None,
#[allow(deprecated)]
function_call: None,
})
}
fn image_user_msg(label: &str) -> ChatCompletionRequestMessage {
user_msg(&format!("[工具返回的图片] {}", label))
}
#[test]
fn interleave_repair_moves_user_messages_after_tool_batch() {
let messages = vec![
user_msg("generate images"),
assistant_multi_tool_call_msg(&[
("call_0", "read"),
("call_1", "read"),
("call_2", "read"),
]),
tool_msg("call_0", "Image file: a.png"),
image_user_msg("a.png"),
tool_msg("call_1", "Image file: b.png"),
image_user_msg("b.png"),
tool_msg("call_2", "Image file: c.png"),
image_user_msg("c.png"),
];
let repaired = repair_interleaved_tool_responses(messages);
assert_eq!(repaired.len(), 8, "no message may be dropped");
let role_kinds: Vec<&str> = repaired
.iter()
.map(|m| match m {
ChatCompletionRequestMessage::User(_) => "user",
ChatCompletionRequestMessage::Assistant(a) => {
if a.tool_calls.is_some() {
"assistant+tool_calls"
} else {
"assistant"
}
}
ChatCompletionRequestMessage::Tool(_) => "tool",
_ => "other",
})
.collect();
assert_eq!(
role_kinds,
vec![
"user",
"assistant+tool_calls",
"tool",
"tool",
"tool",
"user",
"user",
"user",
]
);
let user_texts: Vec<String> = repaired
.iter()
.filter_map(|m| {
if let ChatCompletionRequestMessage::User(u) = m {
if let async_openai::types::chat::ChatCompletionRequestUserMessageContent::Text(t) = &u.content {
Some(t.clone())
} else {
None
}
} else {
None
}
})
.collect();
assert_eq!(
user_texts.last().map(|t| t.contains("c.png")),
Some(true),
"deferred messages keep their original order (c.png last)"
);
}
#[test]
fn interleave_repair_keeps_valid_history_untouched() {
let messages = vec![
user_msg("look at this"),
assistant_tool_call_msg("call_1", "read"),
tool_msg("call_1", "Image file: a.png"),
image_user_msg("a.png"),
assistant_text_msg("looks great"),
];
let repaired = repair_interleaved_tool_responses(messages.clone());
assert_eq!(
repaired.len(),
messages.len(),
"valid history must not change length"
);
let summarize = |m: &ChatCompletionRequestMessage| match m {
ChatCompletionRequestMessage::User(u) => format!("user:{:?}", u.content),
ChatCompletionRequestMessage::Assistant(a) => format!(
"assistant:{:?}:{:?}",
a.content, a.tool_calls.as_ref().map(|tcs| tcs.len())
),
ChatCompletionRequestMessage::Tool(t) => {
format!("tool:{}:{:?}", t.tool_call_id, t.content)
}
_ => "other".to_string(),
};
let before: Vec<String> = messages.iter().map(summarize).collect();
let after: Vec<String> = repaired.iter().map(summarize).collect();
assert_eq!(before, after, "valid history must not be reordered");
}
#[test]
fn interleave_repair_truncated_batch_defers_to_end() {
let messages = vec![
user_msg("do something"),
assistant_multi_tool_call_msg(&[("call_0", "read"), ("call_1", "read")]),
tool_msg("call_0", "result 0"),
image_user_msg("a.png"),
];
let repaired = repair_interleaved_tool_responses(messages);
assert_eq!(repaired.len(), 4);
assert!(matches!(repaired[1], ChatCompletionRequestMessage::Assistant(_)));
assert!(matches!(repaired[2], ChatCompletionRequestMessage::Tool(_)));
assert!(matches!(repaired[3], ChatCompletionRequestMessage::User(_)));
}
#[test]
fn interleave_repair_flushes_deferred_before_next_assistant_batch() {
let messages = vec![
user_msg("start"),
assistant_multi_tool_call_msg(&[("call_0", "read"), ("call_1", "read")]),
tool_msg("call_0", "result 0"),
image_user_msg("a.png"),
assistant_tool_call_msg("call_2", "bash"),
tool_msg("call_2", "ok"),
];
let repaired = repair_interleaved_tool_responses(messages);
let role_kinds: Vec<&str> = repaired
.iter()
.map(|m| match m {
ChatCompletionRequestMessage::User(_) => "user",
ChatCompletionRequestMessage::Assistant(_) => "assistant",
ChatCompletionRequestMessage::Tool(_) => "tool",
_ => "other",
})
.collect();
assert_eq!(
role_kinds,
vec!["user", "assistant", "tool", "user", "assistant", "tool"]
);
}
}