fix: better hardening of tool calls
This commit is contained in:
10
src/a2ui.rs
10
src/a2ui.rs
@@ -1179,7 +1179,10 @@ fn validate_enum(
|
|||||||
.as_str()
|
.as_str()
|
||||||
.ok_or_else(|| format!("component `{id}` {field} must be a string"))?;
|
.ok_or_else(|| format!("component `{id}` {field} must be a string"))?;
|
||||||
if !allowed.contains(&value) {
|
if !allowed.contains(&value) {
|
||||||
return Err(format!("component `{id}` has invalid {field} `{value}`"));
|
return Err(format!(
|
||||||
|
"component `{id}` has invalid {field} `{value}`; expected one of: {}",
|
||||||
|
allowed.join(", ")
|
||||||
|
));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
@@ -2245,7 +2248,10 @@ mod tests {
|
|||||||
json!({"version":VERSION,"createSurface":{"surfaceId":"s","catalogId":CATALOG_ID}}),
|
json!({"version":VERSION,"createSurface":{"surfaceId":"s","catalogId":CATALOG_ID}}),
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
assert!(apply(&mut store, json!({"version":VERSION,"updateComponents":{"surfaceId":"s","components":[{"id":"root","component":"Text","text":"x","variant":"h1"}]}})).is_err());
|
assert_eq!(
|
||||||
|
apply(&mut store, json!({"version":VERSION,"updateComponents":{"surfaceId":"s","components":[{"id":"root","component":"Text","text":"x","variant":"h1"}]}})).unwrap_err(),
|
||||||
|
"component `root` has invalid variant `h1`; expected one of: caption, body"
|
||||||
|
);
|
||||||
assert!(apply(&mut store, json!({"version":VERSION,"updateComponents":{"surfaceId":"s","components":[{"id":"root","component":"Text","text":{"call":"formatDate","args":{"value":"2026-01-01T00:00:00Z"}}}]}})).is_err());
|
assert!(apply(&mut store, json!({"version":VERSION,"updateComponents":{"surfaceId":"s","components":[{"id":"root","component":"Text","text":{"call":"formatDate","args":{"value":"2026-01-01T00:00:00Z"}}}]}})).is_err());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -66,6 +66,9 @@ pub(crate) struct ChatMessage {
|
|||||||
pub(super) a2ui_open_urls: Vec<String>,
|
pub(super) a2ui_open_urls: Vec<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const TOOL_PROTOCOL_CORRECTION: &str = "Tool protocol correction (no tool was executed): you emitted a complete tool call inside private reasoning. End reasoning, then emit the tool call again as assistant content using the required syntax; do not merely discuss it.";
|
||||||
|
const EMPTY_RESPONSE_CORRECTION: &str = "Protocol correction: your previous response was empty. Continue the task now with either a valid tool call or a final answer.";
|
||||||
|
|
||||||
impl ChatMessage {
|
impl ChatMessage {
|
||||||
pub(super) fn append(&mut self, reasoning: bool, chunk: &str) {
|
pub(super) fn append(&mut self, reasoning: bool, chunk: &str) {
|
||||||
if reasoning {
|
if reasoning {
|
||||||
@@ -90,6 +93,33 @@ impl ChatMessage {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn has_misplaced_tool_call(model: ModelChoice, message: &ChatMessage) -> bool {
|
||||||
|
message.content.trim().is_empty()
|
||||||
|
&& message.reasoning.as_deref().is_some_and(|reasoning| {
|
||||||
|
matches!(
|
||||||
|
crate::agent::parse_tool_calls(model, reasoning),
|
||||||
|
Ok((_, calls)) if !calls.is_empty()
|
||||||
|
)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_empty_response(message: &ChatMessage) -> bool {
|
||||||
|
message.content.trim().is_empty()
|
||||||
|
&& message
|
||||||
|
.reasoning
|
||||||
|
.as_deref()
|
||||||
|
.is_none_or(|reasoning| reasoning.trim().is_empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn correction_already_sent(conversation: &[ChatMessage], correction: &str) -> bool {
|
||||||
|
conversation
|
||||||
|
.iter()
|
||||||
|
.rev()
|
||||||
|
.skip(1)
|
||||||
|
.take_while(|message| !message.user)
|
||||||
|
.any(|message| message.tool && message.content == correction)
|
||||||
|
}
|
||||||
|
|
||||||
impl From<StoredMessage> for ChatMessage {
|
impl From<StoredMessage> for ChatMessage {
|
||||||
fn from(message: StoredMessage) -> Self {
|
fn from(message: StoredMessage) -> Self {
|
||||||
let mut message = Self {
|
let mut message = Self {
|
||||||
@@ -554,7 +584,7 @@ impl App {
|
|||||||
#[cfg(target_os = "macos")]
|
#[cfg(target_os = "macos")]
|
||||||
let mut start_queued = false;
|
let mut start_queued = false;
|
||||||
#[cfg(target_os = "macos")]
|
#[cfg(target_os = "macos")]
|
||||||
let mut a2ui_feedback = None;
|
let mut continuation_feedback = None;
|
||||||
#[cfg(target_os = "macos")]
|
#[cfg(target_os = "macos")]
|
||||||
let mut a2ui_changed = false;
|
let mut a2ui_changed = false;
|
||||||
#[cfg(target_os = "macos")]
|
#[cfg(target_os = "macos")]
|
||||||
@@ -625,6 +655,40 @@ impl App {
|
|||||||
!self.queued_inputs.is_empty() || self.manual_compaction_queued;
|
!self.queued_inputs.is_empty() || self.manual_compaction_queued;
|
||||||
}
|
}
|
||||||
Ok(_) => {
|
Ok(_) => {
|
||||||
|
let model = self.config.model;
|
||||||
|
let protocol_recovery = if self
|
||||||
|
.conversation
|
||||||
|
.last()
|
||||||
|
.is_some_and(|message| has_misplaced_tool_call(model, message))
|
||||||
|
{
|
||||||
|
Some((
|
||||||
|
TOOL_PROTOCOL_CORRECTION,
|
||||||
|
"Correcting tool call…",
|
||||||
|
"The model emitted a complete tool call inside private reasoning twice; no tool was executed.",
|
||||||
|
))
|
||||||
|
} else if self.conversation.last().is_some_and(is_empty_response) {
|
||||||
|
Some((
|
||||||
|
EMPTY_RESPONSE_CORRECTION,
|
||||||
|
"Retrying empty response…",
|
||||||
|
"The model returned an empty response twice; generation stopped.",
|
||||||
|
))
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
if let Some((correction, activity, repeated_error)) = protocol_recovery
|
||||||
|
{
|
||||||
|
self.generating = false;
|
||||||
|
self.tool_cards.clear();
|
||||||
|
self.active_generation = None;
|
||||||
|
if correction_already_sent(&self.conversation, correction) {
|
||||||
|
self.activity = Some("Failed".into());
|
||||||
|
self.error = Some(repeated_error.into());
|
||||||
|
} else {
|
||||||
|
self.activity = Some(activity.into());
|
||||||
|
continuation_feedback = Some(correction.to_owned());
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
}
|
||||||
let (validation_errors, replies, open_urls, error_surface_id) = self
|
let (validation_errors, replies, open_urls, error_surface_id) = self
|
||||||
.conversation
|
.conversation
|
||||||
.last_mut()
|
.last_mut()
|
||||||
@@ -673,7 +737,7 @@ impl App {
|
|||||||
"\nThere is no active A2UI surface. Correct this by emitting createSurface with a new surfaceId and a complete root component tree; do not retry updateComponents for an earlier surface.",
|
"\nThere is no active A2UI surface. Correct this by emitting createSurface with a new surfaceId and a complete root component tree; do not retry updateComponents for an earlier surface.",
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
a2ui_feedback = Some(feedback);
|
continuation_feedback = Some(feedback);
|
||||||
self.active_generation = None;
|
self.active_generation = None;
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
@@ -681,7 +745,7 @@ impl App {
|
|||||||
self.generating = false;
|
self.generating = false;
|
||||||
self.activity = Some("Continuing A2UI function call…".into());
|
self.activity = Some("Continuing A2UI function call…".into());
|
||||||
self.tool_cards.clear();
|
self.tool_cards.clear();
|
||||||
a2ui_feedback = Some(format!(
|
continuation_feedback = Some(format!(
|
||||||
"A2UI client response:\n{}",
|
"A2UI client response:\n{}",
|
||||||
replies
|
replies
|
||||||
.iter()
|
.iter()
|
||||||
@@ -799,7 +863,7 @@ impl App {
|
|||||||
self.start_next_queued();
|
self.start_next_queued();
|
||||||
}
|
}
|
||||||
#[cfg(target_os = "macos")]
|
#[cfg(target_os = "macos")]
|
||||||
if let Some(feedback) = a2ui_feedback
|
if let Some(feedback) = continuation_feedback
|
||||||
&& let Err(error) = self.continue_after_tool_result(&feedback)
|
&& let Err(error) = self.continue_after_tool_result(&feedback)
|
||||||
{
|
{
|
||||||
self.generating = false;
|
self.generating = false;
|
||||||
@@ -1510,8 +1574,58 @@ pub(super) fn session_title(reply: &str) -> Option<String> {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::{
|
use super::{
|
||||||
ChatMessage, compacted_context_start, has_chat_after_last_compaction, queued_prompt,
|
ChatMessage, TOOL_PROTOCOL_CORRECTION, compacted_context_start, correction_already_sent,
|
||||||
|
has_chat_after_last_compaction, has_misplaced_tool_call, is_empty_response, queued_prompt,
|
||||||
};
|
};
|
||||||
|
use crate::model::ModelChoice;
|
||||||
|
|
||||||
|
fn assistant(reasoning: Option<&str>, content: &str) -> ChatMessage {
|
||||||
|
ChatMessage {
|
||||||
|
id: 1,
|
||||||
|
user: false,
|
||||||
|
tool: false,
|
||||||
|
system: false,
|
||||||
|
compaction: false,
|
||||||
|
compaction_tail_start: None,
|
||||||
|
reasoning: reasoning.map(str::to_owned),
|
||||||
|
reasoning_complete: false,
|
||||||
|
reasoning_open: false,
|
||||||
|
content: content.to_owned(),
|
||||||
|
markdown: iced::widget::markdown::Content::new(),
|
||||||
|
a2ui_lines_processed: 0,
|
||||||
|
a2ui_errors: Vec::new(),
|
||||||
|
a2ui_replies: Vec::new(),
|
||||||
|
a2ui_open_urls: Vec::new(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn complete_tool_call_in_reasoning_gets_one_correction() {
|
||||||
|
let call = r#"<|DSML|tool_calls>
|
||||||
|
<|DSML|invoke name="bash">
|
||||||
|
<|DSML|parameter name="command" string="true">pwd</|DSML|parameter>
|
||||||
|
</|DSML|invoke>
|
||||||
|
</|DSML|tool_calls>"#;
|
||||||
|
let malformed = assistant(Some(call), "");
|
||||||
|
assert!(has_misplaced_tool_call(
|
||||||
|
ModelChoice::DeepSeekV4Flash,
|
||||||
|
&malformed
|
||||||
|
));
|
||||||
|
assert!(!has_misplaced_tool_call(
|
||||||
|
ModelChoice::DeepSeekV4Flash,
|
||||||
|
&assistant(Some(call), "I was only discussing this call.")
|
||||||
|
));
|
||||||
|
assert!(is_empty_response(&assistant(Some(""), "")));
|
||||||
|
assert!(!is_empty_response(&assistant(Some("still thinking"), "")));
|
||||||
|
|
||||||
|
let mut correction = assistant(None, TOOL_PROTOCOL_CORRECTION);
|
||||||
|
correction.tool = true;
|
||||||
|
let conversation = vec![assistant(Some(call), ""), correction, malformed];
|
||||||
|
assert!(correction_already_sent(
|
||||||
|
&conversation,
|
||||||
|
TOOL_PROTOCOL_CORRECTION
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn queued_guidance_is_one_reference_style_user_turn() {
|
fn queued_guidance_is_one_reference_style_user_turn() {
|
||||||
|
|||||||
Reference in New Issue
Block a user