use std::collections::BTreeSet;
use crate::message::{
self, AssistantContent, CallId, Message, ToolName, ToolResultContent, UserContent,
};
use crate::tool::ToolResult;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum TranscriptError {
#[error("consecutive assistant messages at index {index}")]
ConsecutiveAssistant {
index: usize,
},
#[error("tool call `{call_id}` at index {index} has no result in the following message")]
UnansweredToolCall {
index: usize,
call_id: CallId,
},
#[error(
"tool result `{call_id}` at index {index} answers no call from the preceding assistant message"
)]
OrphanToolResult {
index: usize,
call_id: CallId,
},
}
pub fn validate_canonical(messages: &[Message]) -> Result<(), TranscriptError> {
let mut prev_assistant_calls: Option<BTreeSet<CallId>> = None;
let mut prev_was_assistant = false;
for (index, message) in messages.iter().enumerate() {
match message {
Message::Assistant(turn) => {
if prev_was_assistant {
return Err(TranscriptError::ConsecutiveAssistant { index });
}
if let Some(call_id) = prev_assistant_calls
.take()
.and_then(|pending| pending.into_iter().next())
{
return Err(TranscriptError::UnansweredToolCall {
index: index - 1,
call_id,
});
}
let calls: BTreeSet<CallId> =
turn.tool_calls().map(|call| call.id.clone()).collect();
prev_assistant_calls = (!calls.is_empty()).then_some(calls);
prev_was_assistant = true;
}
Message::User { content } => {
let mut pending = prev_assistant_calls.take().unwrap_or_default();
for item in content.iter() {
if let UserContent::ToolResult(result) = item {
let id = result.call.clone();
if !pending.remove(&id) {
return Err(TranscriptError::OrphanToolResult { index, call_id: id });
}
}
}
if let Some(call_id) = pending.into_iter().next() {
return Err(TranscriptError::UnansweredToolCall {
index: index.saturating_sub(1),
call_id,
});
}
prev_was_assistant = false;
}
Message::System { .. } => {
prev_was_assistant = false;
}
}
}
if let Some(call_id) = prev_assistant_calls.and_then(|pending| pending.into_iter().next()) {
return Err(TranscriptError::UnansweredToolCall {
index: messages.len().saturating_sub(1),
call_id,
});
}
Ok(())
}
pub fn tool_result_output(call: CallId, name: ToolName, result: &ToolResult) -> UserContent {
UserContent::ToolResult(message::ToolResult {
call,
name,
content: result.output().clone().into_content(),
is_error: !result.is_success(),
})
}
pub fn tool_result_message(call: CallId, name: ToolName, message: String) -> UserContent {
UserContent::ToolResult(message::ToolResult {
call,
name,
content: vec![ToolResultContent::text(message)],
is_error: true,
})
}
pub fn invalid_arguments_feedback(tool: &str, raw: &str) -> String {
format!(
"The arguments for tool `{tool}` are not a JSON object: {raw}. \
Call the tool again with a JSON object as its arguments."
)
}
pub const TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER: &str =
"Tool not executed because another tool call in the same assistant turn was invalid.";
pub fn invalid_call_feedback(
content: &[AssistantContent],
invalid: &CallId,
feedback: &str,
) -> Vec<UserContent> {
content
.iter()
.filter_map(|part| match part {
AssistantContent::ToolCall(call) if &call.id == invalid => Some(tool_result_message(
call.id.clone(),
call.function.name.clone(),
feedback.to_owned(),
)),
AssistantContent::ToolCall(call) => Some(not_executed(call)),
_ => None,
})
.collect()
}
pub fn not_executed(call: &crate::message::ToolCall) -> UserContent {
tool_result_message(
call.id.clone(),
call.function.name.clone(),
TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER.to_owned(),
)
}
pub fn is_empty_assistant_turn(content: &[AssistantContent]) -> bool {
content.iter().all(AssistantContent::is_blank)
}
pub fn assistant_text_from_choice(content: &[AssistantContent]) -> String {
content
.iter()
.filter_map(|part| match part {
AssistantContent::Text(text) => Some(text.text.as_str()),
_ => None,
})
.collect()
}
#[cfg(test)]
mod validator_tests;