Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions rust/src/tool-parser/benches/qwen3_coder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ use utils::{feed_external_parser, feed_parser, openai_tools};

const CHUNK_CHARS: usize = 7;
const LONG_NORMAL_TEXT_REPEATS: usize = 2048;
const LONG_TOOL_BODY_REPEATS: usize = 8192;

fn mixed_fixture() -> String {
concat!(
Expand Down Expand Up @@ -39,6 +40,17 @@ fn long_normal_text_fixture() -> String {
line.repeat(LONG_NORMAL_TEXT_REPEATS)
}

fn long_tool_call_fixture() -> String {
let location = "x".repeat(LONG_TOOL_BODY_REPEATS);
format!(
"<tool_call>\n\
<function=get_weather>\n\
<parameter=location>{location}</parameter>\n\
</function>\n\
</tool_call>"
)
}

fn native_parser(tools: &[Tool]) -> Box<dyn ToolParser> {
Qwen3CoderToolParser::create(tools).expect("Qwen Coder parser should initialize")
}
Expand Down Expand Up @@ -112,6 +124,7 @@ fn bench_qwen3_coder(c: &mut Criterion) {
let tools = test_tools();
let mixed_text = mixed_fixture();
let long_normal_text = long_normal_text_fixture();
let long_tool_call = long_tool_call_fixture();

run_stream_group(
c,
Expand All @@ -132,6 +145,16 @@ fn bench_qwen3_coder(c: &mut Criterion) {
&long_normal_text,
0,
);

run_stream_group(
c,
"qwen3_coder/long_tool_call_body",
&tools,
&long_tool_call,
CHUNK_CHARS,
"",
1,
);
}

criterion_group!(benches, bench_qwen3_coder);
Expand Down
42 changes: 28 additions & 14 deletions rust/src/tool-parser/src/deepseek_dsml/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use winnow::stream::Partial;
use winnow::token::{literal, rest, take_until};

use super::parameters::ToolSchemas;
use super::utils::{parse_buffered_event, safe_text_len};
use super::utils::{MarkerScanState, parse_buffered_event, safe_text_len, take_until_marker};
use super::{Result, ToolCallDelta, ToolParserOutput};
use crate::Tool;

Expand Down Expand Up @@ -39,10 +39,10 @@ impl DsmlTokens {
};
}

#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[derive(Debug, Clone, PartialEq, Eq)]
enum DsmlMode {
Text,
ToolBlock,
ToolBlock { invoke_end_scan: MarkerScanState },
Done,
}

Expand Down Expand Up @@ -94,7 +94,11 @@ impl DeepSeekDsmlToolParser {
DsmlEvent::Text { len: consumed_len } => {
output.normal_text.push_str(&self.buffer[..consumed_len]);
}
DsmlEvent::ToolCallsStart => self.mode = DsmlMode::ToolBlock,
DsmlEvent::ToolCallsStart => {
self.mode = DsmlMode::ToolBlock {
invoke_end_scan: MarkerScanState::default(),
};
}
DsmlEvent::Invoke { name, raw_params } => {
let mut arguments = serde_json::Map::with_capacity(raw_params.len());
for param in raw_params {
Expand Down Expand Up @@ -140,7 +144,7 @@ impl DeepSeekDsmlToolParser {
self.buffer.push_str(chunk);

while let Some((event, consumed_len)) = parse_buffered_event(&self.buffer, |input| {
parse_next_dsml_event(input, self.mode, self.tokens)
parse_next_dsml_event(input, &mut self.mode, self.tokens)
})? {
self.apply_event(event, output)?;
self.buffer.drain(..consumed_len);
Expand All @@ -154,7 +158,7 @@ impl DeepSeekDsmlToolParser {
match self.mode {
DsmlMode::Text => output.normal_text.push_str(&self.buffer),
DsmlMode::Done => {}
DsmlMode::ToolBlock => {
DsmlMode::ToolBlock { .. } => {
return Err(parsing_failed!("incomplete DeepSeek DSML tool call"));
}
}
Expand All @@ -166,12 +170,14 @@ impl DeepSeekDsmlToolParser {
/// Parse a DSML event for the current parser mode.
fn parse_next_dsml_event(
input: &mut DsmlInput<'_>,
mode: DsmlMode,
mode: &mut DsmlMode,
tokens: DsmlTokens,
) -> ModalResult<DsmlEvent> {
match mode {
DsmlMode::Text => parse_text_event(input, tokens),
DsmlMode::ToolBlock => parse_tool_block_event(input, tokens),
DsmlMode::ToolBlock { invoke_end_scan } => {
parse_tool_block_event(input, tokens, invoke_end_scan)
}
DsmlMode::Done => ignored_rest_event(input),
}
}
Expand All @@ -186,11 +192,16 @@ fn parse_text_event(input: &mut DsmlInput<'_>, tokens: DsmlTokens) -> ModalResul
}

/// Parse a tool-block DSML event.
fn parse_tool_block_event(input: &mut DsmlInput<'_>, tokens: DsmlTokens) -> ModalResult<DsmlEvent> {
fn parse_tool_block_event(
input: &mut DsmlInput<'_>,
tokens: DsmlTokens,
invoke_end_scan: &mut MarkerScanState,
) -> ModalResult<DsmlEvent> {
ws0.void().parse_next(input)?;
alt((invoke_event, |input: &mut DsmlInput<'_>| {
tool_calls_end_event(input, tokens)
}))
alt((
|input: &mut DsmlInput<'_>| invoke_event(input, invoke_end_scan),
|input: &mut DsmlInput<'_>| tool_calls_end_event(input, tokens),
))
.parse_next(input)
}

Expand All @@ -217,14 +228,17 @@ fn safe_text_event(input: &mut DsmlInput<'_>, tokens: DsmlTokens) -> ModalResult
}

/// Parse a DSML invoke block.
fn invoke_event(input: &mut DsmlInput<'_>) -> ModalResult<DsmlEvent> {
fn invoke_event(
input: &mut DsmlInput<'_>,
invoke_end_scan: &mut MarkerScanState,
) -> ModalResult<DsmlEvent> {
let (name, body) = seq!(
_: literal(INVOKE_START),
_: ws1,
dsml_name_attr,
_: ws0,
_: ">",
take_until(0.., INVOKE_END),
take_until_marker(INVOKE_END, invoke_end_scan),
_: literal(INVOKE_END),
)
.parse_next(input)?;
Expand Down
32 changes: 22 additions & 10 deletions rust/src/tool-parser/src/glm_xml/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use winnow::stream::Partial;
use winnow::token::{literal, rest, take_until, take_while};

use super::parameters::ToolSchemas;
use super::utils::{parse_buffered_event, safe_text_len};
use super::utils::{MarkerScanState, parse_buffered_event, safe_text_len, take_until_marker};
use super::{Result, ToolCallDelta, ToolParserOutput};
use crate::Tool;

Expand All @@ -24,10 +24,10 @@ const ARG_VALUE_END: &str = "</arg_value>";

type GlmInput<'i> = Partial<&'i str>;

#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[derive(Debug, Clone, PartialEq, Eq)]
enum GlmMode {
Text,
ToolCall,
ToolCall { tool_call_end_scan: MarkerScanState },
AfterToolCall,
}

Expand Down Expand Up @@ -81,7 +81,11 @@ impl GlmXmlToolParser {
GlmEvent::Text { len: consumed_len } => {
output.normal_text.push_str(&self.buffer[..consumed_len]);
}
GlmEvent::ToolCallStart => self.mode = GlmMode::ToolCall,
GlmEvent::ToolCallStart => {
self.mode = GlmMode::ToolCall {
tool_call_end_scan: MarkerScanState::default(),
};
}
GlmEvent::ToolCall { name, raw_params } => {
self.mode = GlmMode::AfterToolCall;
let arguments = self.tool_parameters.convert_params_with_schema(&name, raw_params);
Expand Down Expand Up @@ -110,7 +114,7 @@ impl GlmXmlToolParser {
self.buffer.push_str(chunk);

while let Some((event, consumed_len)) = parse_buffered_event(&self.buffer, |input| {
parse_next_glm_event(input, self.mode, self.separator)
parse_next_glm_event(input, &mut self.mode, self.separator)
})? {
self.apply_event(event, output)?;
self.buffer.drain(..consumed_len);
Expand All @@ -124,7 +128,9 @@ impl GlmXmlToolParser {
if !self.buffer.is_empty() {
match self.mode {
GlmMode::Text => output.normal_text.push_str(&self.buffer),
GlmMode::ToolCall => return Err(parsing_failed!("incomplete GLM MoE tool call")),
GlmMode::ToolCall { .. } => {
return Err(parsing_failed!("incomplete GLM MoE tool call"));
}
GlmMode::AfterToolCall => {}
}
}
Expand All @@ -136,12 +142,14 @@ impl GlmXmlToolParser {
/// Parse a GLM event for the current parser mode.
fn parse_next_glm_event(
input: &mut GlmInput<'_>,
mode: GlmMode,
mode: &mut GlmMode,
separator: Separator,
) -> ModalResult<GlmEvent> {
match mode {
GlmMode::Text => parse_text_event(input),
GlmMode::ToolCall => tool_call_event(input, separator),
GlmMode::ToolCall { tool_call_end_scan } => {
tool_call_event(input, separator, tool_call_end_scan)
}
GlmMode::AfterToolCall => after_tool_call_event(input),
}
}
Expand Down Expand Up @@ -173,9 +181,13 @@ fn ignored_rest_event(input: &mut GlmInput<'_>) -> ModalResult<GlmEvent> {
}

/// Parse a complete GLM tool call.
fn tool_call_event(input: &mut GlmInput<'_>, separator: Separator) -> ModalResult<GlmEvent> {
fn tool_call_event(
input: &mut GlmInput<'_>,
separator: Separator,
tool_call_end_scan: &mut MarkerScanState,
) -> ModalResult<GlmEvent> {
let (body,) = seq!(
take_until(0.., TOOL_CALL_END),
take_until_marker(TOOL_CALL_END, tool_call_end_scan),
_: literal(TOOL_CALL_END),
)
.parse_next(input)?;
Expand Down
42 changes: 30 additions & 12 deletions rust/src/tool-parser/src/hy_v3.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use winnow::stream::Partial;
use winnow::token::{literal, rest, take_until};

use super::parameters::ToolSchemas;
use super::utils::{parse_buffered_event, safe_text_len};
use super::utils::{MarkerScanState, parse_buffered_event, safe_text_len, take_until_marker};
use super::{Result, ToolCallDelta, ToolParser, ToolParserOutput};
use crate::Tool;

Expand All @@ -21,10 +21,10 @@ const ARG_VALUE_END: &str = "</arg_value>";

type HyV3Input<'i> = Partial<&'i str>;

#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[derive(Debug, Clone, PartialEq, Eq)]
enum HyV3Mode {
Text,
ToolBlock,
ToolBlock { tool_call_end_scan: MarkerScanState },
Done,
}

Expand Down Expand Up @@ -81,7 +81,11 @@ impl HyV3ToolParser {
HyV3Event::Text { len: consumed_len } => {
output.normal_text.push_str(&self.buffer[..consumed_len]);
}
HyV3Event::ToolBlockStart => self.mode = HyV3Mode::ToolBlock,
HyV3Event::ToolBlockStart => {
self.mode = HyV3Mode::ToolBlock {
tool_call_end_scan: MarkerScanState::default(),
};
}
HyV3Event::ToolCall { name, raw_params } => {
let arguments = self.tool_parameters.convert_params_with_schema(&name, raw_params);
let arguments = serde_json::to_string(&arguments)
Expand Down Expand Up @@ -113,7 +117,7 @@ impl ToolParser for HyV3ToolParser {
self.buffer.push_str(chunk);

while let Some((event, consumed_len)) = parse_buffered_event(&self.buffer, |input| {
parse_next_hy_v3_event(input, self.mode)
parse_next_hy_v3_event(input, &mut self.mode)
})? {
self.apply_event(event, output)?;
self.buffer.drain(..consumed_len);
Expand All @@ -126,7 +130,7 @@ impl ToolParser for HyV3ToolParser {
let mut output = ToolParserOutput::default();
match self.mode {
HyV3Mode::Text => output.normal_text.push_str(&self.buffer),
HyV3Mode::ToolBlock => return Err(parsing_failed!("incomplete HY3 tool call")),
HyV3Mode::ToolBlock { .. } => return Err(parsing_failed!("incomplete HY3 tool call")),
HyV3Mode::Done => {}
}
let _ = self.reset();
Expand All @@ -141,10 +145,15 @@ impl ToolParser for HyV3ToolParser {
}

/// Parse a HY3 event for the current parser mode.
fn parse_next_hy_v3_event(input: &mut HyV3Input<'_>, mode: HyV3Mode) -> ModalResult<HyV3Event> {
fn parse_next_hy_v3_event(
input: &mut HyV3Input<'_>,
mode: &mut HyV3Mode,
) -> ModalResult<HyV3Event> {
match mode {
HyV3Mode::Text => parse_text_event(input),
HyV3Mode::ToolBlock => parse_tool_block_event(input),
HyV3Mode::ToolBlock { tool_call_end_scan } => {
parse_tool_block_event(input, tool_call_end_scan)
}
HyV3Mode::Done => ignored_rest_event(input),
}
}
Expand All @@ -165,8 +174,14 @@ fn safe_text_event(input: &mut HyV3Input<'_>) -> ModalResult<HyV3Event> {
}

/// Parse one event inside a HY3 tool block.
fn parse_tool_block_event(input: &mut HyV3Input<'_>) -> ModalResult<HyV3Event> {
alt((tool_block_end_event, tool_call_event)).parse_next(input)
fn parse_tool_block_event(
input: &mut HyV3Input<'_>,
tool_call_end_scan: &mut MarkerScanState,
) -> ModalResult<HyV3Event> {
alt((tool_block_end_event, |input: &mut HyV3Input<'_>| {
tool_call_event(input, tool_call_end_scan)
}))
.parse_next(input)
}

/// Parse a HY3 tool-block end marker.
Expand All @@ -175,13 +190,16 @@ fn tool_block_end_event(input: &mut HyV3Input<'_>) -> ModalResult<HyV3Event> {
}

/// Parse a complete HY3 tool-call block.
fn tool_call_event(input: &mut HyV3Input<'_>) -> ModalResult<HyV3Event> {
fn tool_call_event(
input: &mut HyV3Input<'_>,
tool_call_end_scan: &mut MarkerScanState,
) -> ModalResult<HyV3Event> {
let (name, body) = seq!(
_: ws0,
_: literal(TOOL_CALL_START),
take_until(0.., TOOL_SEP),
_: literal(TOOL_SEP),
take_until(0.., TOOL_CALL_END),
take_until_marker(TOOL_CALL_END, tool_call_end_scan),
_: literal(TOOL_CALL_END),
)
.parse_next(input)?;
Expand Down
Loading
Loading