refactor(chat): unify generation flow and fix tool/streaming bugs
- Unified synchronous generate() core with callback; both streaming and non-streaming paths run it via spawn_blocking (context: std Mutex). - build_prompt now passes tool definitions to the template (was dead) and embeds assistant tool-call history as XML matching the parser format; fixes double <tool_response> wrap and template set-scoping bug. - Tokenize with AddBos::Never (template owns <s>) to remove double BOS. - Streaming: preserve inter-word spaces (per-chunk trim removed), add [DONE] + usage chunk, emit error events, single-shot tool_calls delta. - Strict model validation (400 on unknown model); health/UI/README aligned to minicpm5-1b-fable5-v2-thinking; auth returns JSON errors; n_ctx/ n_batch/n_threads env-configurable. - Added 18 unit tests; cargo check/clippy/fmt clean. - scripts/smoke-test.sh for post-deploy verification on the VPS. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
344bc195fa
commit
b636496497
@@ -1,3 +1,6 @@
|
||||
pub mod use_cases;
|
||||
|
||||
pub use use_cases::{build_prompt, build_sampler, clean_text, parse_tool_calls, SamplerParams};
|
||||
pub use use_cases::{
|
||||
build_prompt, build_sampler, clean_text, parse_tool_calls, split_stream_chunk, validate_model,
|
||||
SamplerParams,
|
||||
};
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
{{- "\n" }}
|
||||
{{- tool | tojson }}
|
||||
{%- endfor %}
|
||||
{{- '\n</tools>\n\nTool usage guidelines:\n- You may call zero or more functions. If no function calls are needed, just answer normally.\n- When calling a function, use: <function name="name"><param name="key">value</param></function>' }}
|
||||
{{- '\n</tools>\n\nTool usage guidelines:\n- You may call zero or more functions. If no function calls are needed, just answer normally.\n- When calling a function, wrap each call in <tool_call> tags:\n<tool_call>\n<function=name>\n<parameter=key>value</parameter>\n</function>\n</tool_call>' }}
|
||||
{%- endset %}
|
||||
|
||||
{{- '<|im_start|>system\n' }}
|
||||
@@ -26,17 +26,8 @@
|
||||
{%- if (message.role == "user") or (message.role == "system" and not loop.first) %}
|
||||
{{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
|
||||
{%- elif message.role == "assistant" %}
|
||||
{%- if message.tool_calls %}
|
||||
{%- set content = message.content %}
|
||||
{%- for tool_call in message.tool_calls %}
|
||||
{%- if tool_call.type == "function" %}
|
||||
{%- set content = content + "\n<tool_call>\n<function=" + tool_call.function.name + ">\n" + (tool_call.function.arguments | tojson) + "\n</function>\n</tool_call>" %}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{{- '<|im_start|>assistant\n' + content + '<|im_end|>\n' }}
|
||||
{%- else %}
|
||||
{{- '<|im_start|>assistant\n' + content + '<|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
{#- Tool-call XML from history is pre-embedded in content by build_prompt. #}
|
||||
{{- '<|im_start|>assistant\n' + content + '<|im_end|>\n' }}
|
||||
{%- elif message.role == "tool" %}
|
||||
{{- '<|im_start|>user\n<tool_response>\n' + content + '\n</tool_response><|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
|
||||
@@ -5,21 +5,21 @@
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use minijinja::{Environment, Value};
|
||||
use llama_cpp_2::sampling::LlamaSampler;
|
||||
use minijinja::{Environment, Value};
|
||||
|
||||
use crate::domain::entity::{ChatMessage, ChatRequest, ToolCall, ToolCallFunction};
|
||||
|
||||
/// Build a prompt from messages using the GGUF's Jinja chat template.
|
||||
/// Build a prompt from messages using the project's Jinja chat template.
|
||||
///
|
||||
/// Renders the model's baked-in template via minijinja, passing the message
|
||||
/// history, optional tool definitions, and generation-prompt switches.
|
||||
/// Renders the template (see `templates/chat_template.jinja`) via minijinja,
|
||||
/// passing the message history, optional tool definitions, and the
|
||||
/// generation-prompt switches. The template owns the `<s>` BOS token, so
|
||||
/// tokenization must NOT add another one (`AddBos::Never`).
|
||||
pub fn build_prompt(
|
||||
model: &llama_cpp_2::model::LlamaModel,
|
||||
messages: &[ChatMessage],
|
||||
_tools: &Option<Vec<crate::domain::entity::ToolDef>>,
|
||||
tools: &Option<Vec<crate::domain::entity::ToolDef>>,
|
||||
) -> Result<String, String> {
|
||||
// Load the GGUF's chat template (embedded at crate build time)
|
||||
let template_str = include_str!("templates/chat_template.jinja");
|
||||
|
||||
let mut env = Environment::new();
|
||||
@@ -31,50 +31,55 @@ pub fn build_prompt(
|
||||
serde_json::to_string(value).unwrap_or_default()
|
||||
});
|
||||
|
||||
let tmpl = env.get_template("chat")
|
||||
let tmpl = env
|
||||
.get_template("chat")
|
||||
.map_err(|e| format!("Template get error: {e}"))?;
|
||||
|
||||
// Build messages as serde_json::Value for minijinja
|
||||
// Build messages as serde_json::Value for minijinja.
|
||||
// Assistant tool calls from history are embedded directly into the content
|
||||
// as XML (in the exact format the model is told to emit) — the template
|
||||
// cannot accumulate `set` variables across a loop, so this is built here.
|
||||
let mut msgs_val: Vec<Value> = Vec::new();
|
||||
for msg in messages {
|
||||
let mut m: HashMap<String, Value> = HashMap::new();
|
||||
m.insert("role".into(), Value::from(msg.role.clone()));
|
||||
|
||||
let content = msg.content.clone().unwrap_or_default();
|
||||
let mut content = msg.content.clone().unwrap_or_default();
|
||||
|
||||
// For assistant messages, check if there are tool_calls
|
||||
if msg.role == "assistant" {
|
||||
if let Some(tcs) = &msg.tool_calls {
|
||||
// Serialise tool calls per the template's expected format
|
||||
let tcs_val: Vec<Value> = tcs.iter().map(|tc| {
|
||||
let args: serde_json::Value =
|
||||
serde_json::from_str(&tc.function.arguments).unwrap_or_default();
|
||||
Value::from_serialize(&serde_json::json!({
|
||||
"id": tc.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.function.name,
|
||||
"arguments": args,
|
||||
for tc in tcs {
|
||||
if tc.call_type == "function" {
|
||||
content
|
||||
.push_str(&format!("\n<tool_call>\n<function={}>\n", tc.function.name));
|
||||
let args: serde_json::Value =
|
||||
serde_json::from_str(&tc.function.arguments).unwrap_or_default();
|
||||
if let Some(obj) = args.as_object() {
|
||||
for (k, v) in obj {
|
||||
// Strings stay raw (no JSON quotes) so they round-trip
|
||||
// through parse_tool_calls unchanged.
|
||||
let rendered = match v {
|
||||
serde_json::Value::String(s) => s.clone(),
|
||||
other => other.to_string(),
|
||||
};
|
||||
content
|
||||
.push_str(&format!("<parameter={k}>{rendered}</parameter>\n"));
|
||||
}
|
||||
}
|
||||
}))
|
||||
}).collect();
|
||||
m.insert("tool_calls".into(), Value::from(tcs_val));
|
||||
content.push_str("</function>\n</tool_call>");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Handle tool role messages
|
||||
if msg.role == "tool" {
|
||||
// Wrap in tool_response as the template expects
|
||||
let wrapped = format!("<tool_response>\n{}\n</tool_response>", content);
|
||||
m.insert("content".into(), Value::from(wrapped));
|
||||
} else {
|
||||
m.insert("content".into(), Value::from(content));
|
||||
}
|
||||
// NOTE: the template wraps tool-role content in <tool_response>; do not
|
||||
// wrap here or it would be double-wrapped.
|
||||
m.insert("content".into(), Value::from(content));
|
||||
|
||||
msgs_val.push(Value::from(m));
|
||||
}
|
||||
|
||||
// BOS token for sentencepiece / unigram models
|
||||
// BOS token for sentencepiece / unigram models (template-owned).
|
||||
let bos_token: &str = "<s>";
|
||||
|
||||
// Build context
|
||||
@@ -84,6 +89,14 @@ pub fn build_prompt(
|
||||
ctx.insert("add_generation_prompt".into(), Value::from(true));
|
||||
ctx.insert("enable_thinking".into(), Value::from(true));
|
||||
|
||||
// Tool definitions (optional) — previously dead, now actually rendered.
|
||||
if let Some(tools) = tools {
|
||||
if !tools.is_empty() {
|
||||
let tools_val: Vec<Value> = tools.iter().map(Value::from_serialize).collect();
|
||||
ctx.insert("tools".into(), Value::from(tools_val));
|
||||
}
|
||||
}
|
||||
|
||||
// Render
|
||||
let result = tmpl
|
||||
.render(&ctx)
|
||||
@@ -92,6 +105,17 @@ pub fn build_prompt(
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Validate that the requested model matches the single served model.
|
||||
pub fn validate_model(model: &str) -> Result<(), String> {
|
||||
if model != crate::config::MODEL_ID {
|
||||
return Err(format!(
|
||||
"Unknown model '{model}'. Available: {}",
|
||||
crate::config::MODEL_ID
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Parameters for building a [`LlamaSampler`] chain.
|
||||
pub struct SamplerParams {
|
||||
pub temperature: Option<f32>,
|
||||
@@ -174,13 +198,13 @@ pub fn build_sampler(params: &SamplerParams) -> LlamaSampler {
|
||||
///
|
||||
/// For thinking models, returns (reasoning, cleaned_answer).
|
||||
pub fn clean_text(text: &str) -> (String, String) {
|
||||
let text = text.replace("<|im_end|>", "")
|
||||
.replace("<|im_start|>", "");
|
||||
let text = text.replace("<|im_end|>", "").replace("<|im_start|>", "");
|
||||
|
||||
// Separate reasoning (between <think>/</think>) from answer
|
||||
let text = text.trim();
|
||||
let (reasoning, answer) = if let Some(close_idx) = text.find("</think>") {
|
||||
let reasoning = text[..close_idx].trim()
|
||||
let reasoning = text[..close_idx]
|
||||
.trim()
|
||||
.trim_start_matches("<think>")
|
||||
.trim()
|
||||
.to_string();
|
||||
@@ -250,9 +274,9 @@ pub fn parse_tool_calls(text: &str) -> (String, Vec<ToolCall>) {
|
||||
|
||||
for line in lines {
|
||||
let line = line.trim();
|
||||
if let Some(param) =
|
||||
line.strip_prefix("<parameter=")
|
||||
.and_then(|s| s.strip_suffix('>'))
|
||||
if let Some(param) = line
|
||||
.strip_prefix("<parameter=")
|
||||
.and_then(|s| s.strip_suffix('>'))
|
||||
{
|
||||
if let Some(p) = current_param.take() {
|
||||
args_map.insert(
|
||||
@@ -275,7 +299,10 @@ pub fn parse_tool_calls(text: &str) -> (String, Vec<ToolCall>) {
|
||||
}
|
||||
}
|
||||
if let Some(p) = current_param.take() {
|
||||
args_map.insert(p, serde_json::Value::String(current_value.trim().to_string()));
|
||||
args_map.insert(
|
||||
p,
|
||||
serde_json::Value::String(current_value.trim().to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
let args_json = serde_json::Value::Object(args_map).to_string();
|
||||
@@ -298,3 +325,325 @@ pub fn parse_tool_calls(text: &str) -> (String, Vec<ToolCall>) {
|
||||
|
||||
(cleaned, tool_calls)
|
||||
}
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════
|
||||
// STREAMING CHUNK SPLITTING
|
||||
// ═══════════════════════════════════════════════════════════════
|
||||
|
||||
/// Remove special tokens and tool-call XML markup from a text fragment.
|
||||
///
|
||||
/// Handles both fixed tags (`<|im_end|>`, `<think>`, `<tool_call>`, …) and the
|
||||
/// attribute-bearing openers used by this model's tool format (`<function=…>`,
|
||||
/// `<parameter=…>`), even when a tag straddles a token boundary.
|
||||
fn strip_markup(text: &str) -> String {
|
||||
let mut out = String::with_capacity(text.len());
|
||||
let mut rest = text;
|
||||
|
||||
while !rest.is_empty() {
|
||||
let Some(idx) = rest.find('<') else {
|
||||
out.push_str(rest);
|
||||
break;
|
||||
};
|
||||
out.push_str(&rest[..idx]);
|
||||
let tail = &rest[idx..];
|
||||
|
||||
// Fixed tags (no attribute content).
|
||||
let fixed = [
|
||||
"<|im_start|>",
|
||||
"<|im_end|>",
|
||||
"<think>",
|
||||
"</think>",
|
||||
"<tool_call>",
|
||||
"</tool_call>",
|
||||
"</function>",
|
||||
"</parameter>",
|
||||
];
|
||||
if let Some(tag) = fixed.iter().find(|t| tail.starts_with(**t)) {
|
||||
rest = &tail[tag.len()..];
|
||||
continue;
|
||||
}
|
||||
|
||||
// Attribute-bearing openers: <function=…> / <parameter=…>.
|
||||
if let Some(attr) = tail
|
||||
.strip_prefix("<function=")
|
||||
.or_else(|| tail.strip_prefix("<parameter="))
|
||||
{
|
||||
if let Some(end) = attr.find('>') {
|
||||
rest = &attr[end + 1..];
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// Not a known tag — keep this char and advance one UTF-8 char.
|
||||
let ch = tail.chars().next().expect("non-empty tail");
|
||||
out.push(ch);
|
||||
rest = &tail[ch.len_utf8()..];
|
||||
}
|
||||
|
||||
out
|
||||
}
|
||||
|
||||
/// Split an incremental streamed text fragment into `(reasoning, content)`
|
||||
/// deltas for SSE.
|
||||
///
|
||||
/// * `think_done` means the `</think>` boundary was already crossed **before**
|
||||
/// this fragment (i.e. it is not the chunk containing the first `</think>`).
|
||||
/// * Whitespace **inside** a fragment is preserved — only the whitespace
|
||||
/// sitting immediately around the `</think>` boundary is trimmed, so
|
||||
/// reasoning does not end with a dangling newline and content does not begin
|
||||
/// with one. (Trimming every fragment corrupted inter-word spaces.)
|
||||
/// * Before the first `</think>`, everything is emitted as `reasoning_content`;
|
||||
/// after it, as `content`. A stray second `</think>` is stripped, not split.
|
||||
pub fn split_stream_chunk(new_text: &str, think_done: bool) -> (Option<String>, String) {
|
||||
if !think_done {
|
||||
if let Some(pos) = new_text.find("</think>") {
|
||||
let mut before = strip_markup(&new_text[..pos]);
|
||||
let mut after = strip_markup(&new_text[pos + 8..]);
|
||||
|
||||
while before.ends_with(['\n', ' ', '\t']) {
|
||||
before.pop();
|
||||
}
|
||||
while after.starts_with(['\n', ' ', '\t']) {
|
||||
after.remove(0);
|
||||
}
|
||||
|
||||
let reasoning = if before.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(before)
|
||||
};
|
||||
(reasoning, after)
|
||||
} else {
|
||||
// Still thinking — everything is reasoning.
|
||||
let cleaned = strip_markup(new_text);
|
||||
let reasoning = if cleaned.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(cleaned)
|
||||
};
|
||||
(reasoning, String::new())
|
||||
}
|
||||
} else {
|
||||
// Think phase already ended — everything is content; strip stray markup.
|
||||
(None, strip_markup(new_text))
|
||||
}
|
||||
}
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════
|
||||
// TESTS
|
||||
// ═══════════════════════════════════════════════════════════════
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::domain::entity::{
|
||||
ChatMessage, ToolCallFunction, ToolCallResponse, ToolDef, ToolFunction,
|
||||
};
|
||||
|
||||
fn msg(role: &str, content: &str) -> ChatMessage {
|
||||
ChatMessage {
|
||||
role: role.into(),
|
||||
content: Some(content.into()),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
name: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_def() -> ToolDef {
|
||||
ToolDef {
|
||||
tool_type: "function".into(),
|
||||
function: ToolFunction {
|
||||
name: "get_weather".into(),
|
||||
description: "Get weather for a city".into(),
|
||||
parameters: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": { "type": "string" }
|
||||
},
|
||||
"required": ["city"]
|
||||
}),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// ── split_stream_chunk ──
|
||||
|
||||
#[test]
|
||||
fn split_chunk_before_think_is_reasoning() {
|
||||
let (reasoning, content) = split_stream_chunk("Hello ", false);
|
||||
assert_eq!(reasoning.as_deref(), Some("Hello "));
|
||||
assert_eq!(content, "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_chunk_preserves_internal_spaces() {
|
||||
// Regression: trimming every fragment used to eat inter-word spaces.
|
||||
let (r1, _) = split_stream_chunk("Hello", false);
|
||||
let (r2, _) = split_stream_chunk(" world", false);
|
||||
assert_eq!(format!("{}{}", r1.unwrap(), r2.unwrap()), "Hello world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_chunk_boundary_trims_only_edges() {
|
||||
let (reasoning, content) = split_stream_chunk("question\n</think>\n\nAnswer ", false);
|
||||
assert_eq!(reasoning.as_deref(), Some("question"));
|
||||
assert_eq!(content, "Answer ");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_chunk_after_think_is_content() {
|
||||
let (reasoning, content) = split_stream_chunk(" answer", true);
|
||||
assert_eq!(reasoning, None);
|
||||
assert_eq!(content, " answer");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_chunk_stray_think_tag_is_stripped_not_split() {
|
||||
// A second </think> (already past the boundary) must not restart
|
||||
// reasoning classification.
|
||||
let (reasoning, content) = split_stream_chunk("...</think>more", true);
|
||||
assert_eq!(reasoning, None);
|
||||
assert_eq!(content, "...more");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_chunk_strips_special_and_markup() {
|
||||
let (reasoning, content) = split_stream_chunk(
|
||||
"<|im_end|><think>Hello</think>\n<tool_call><function=get_weather>",
|
||||
false,
|
||||
);
|
||||
assert_eq!(reasoning.as_deref(), Some("Hello"));
|
||||
assert_eq!(content, "");
|
||||
}
|
||||
|
||||
// ── clean_text ──
|
||||
|
||||
#[test]
|
||||
fn clean_text_splits_reasoning_and_answer() {
|
||||
let (reasoning, answer) =
|
||||
clean_text("<think>Let me think\nabout it</think>\nThe answer is 42.");
|
||||
assert_eq!(reasoning, "Let me think\nabout it");
|
||||
assert_eq!(answer, "The answer is 42.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn clean_text_no_think_returns_answer() {
|
||||
let (reasoning, answer) = clean_text("Just an answer");
|
||||
assert_eq!(reasoning, "");
|
||||
assert_eq!(answer, "Just an answer");
|
||||
}
|
||||
|
||||
// ── parse_tool_calls ──
|
||||
|
||||
#[test]
|
||||
fn parse_single_tool_call() {
|
||||
let text = "I'll look that up.\n<tool_call>\n<function=get_weather>\n<parameter=city>Jakarta</parameter>\n</function>\n</tool_call>";
|
||||
let (cleaned, calls) = parse_tool_calls(text);
|
||||
assert_eq!(calls.len(), 1);
|
||||
assert_eq!(calls[0].function.name, "get_weather");
|
||||
assert!(calls[0].function.arguments.contains("Jakarta"));
|
||||
assert!(!cleaned.contains("<tool_call>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_multiple_tool_calls() {
|
||||
let text = "<tool_call>\n<function=a>\n<parameter=x>1</parameter>\n</function>\n</tool_call>\n<tool_call>\n<function=b>\n<parameter=y>2</parameter>\n</function>\n</tool_call>";
|
||||
let (_cleaned, calls) = parse_tool_calls(text);
|
||||
assert_eq!(calls.len(), 2);
|
||||
assert_eq!(calls[0].function.name, "a");
|
||||
assert_eq!(calls[1].function.name, "b");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_malformed_tool_call_returns_empty() {
|
||||
let (cleaned, calls) = parse_tool_calls("no calls here");
|
||||
assert!(calls.is_empty());
|
||||
assert_eq!(cleaned, "no calls here");
|
||||
}
|
||||
|
||||
// ── validate_model ──
|
||||
|
||||
#[test]
|
||||
fn validate_model_accepts_served_id() {
|
||||
assert!(validate_model(crate::config::MODEL_ID).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_model_rejects_unknown_id() {
|
||||
assert!(validate_model("minicpm-v-4.6").is_err());
|
||||
}
|
||||
|
||||
// ── build_sampler (smoke — no model required) ──
|
||||
|
||||
#[test]
|
||||
fn build_sampler_constructs_for_common_params() {
|
||||
let params = SamplerParams {
|
||||
temperature: Some(0.8),
|
||||
top_p: Some(0.9),
|
||||
top_k: Some(40),
|
||||
min_p: Some(0.05),
|
||||
repeat_penalty: Some(1.1),
|
||||
frequency_penalty: Some(0.0),
|
||||
presence_penalty: Some(0.0),
|
||||
seed: Some(42),
|
||||
};
|
||||
let _ = build_sampler(¶ms);
|
||||
|
||||
let greedy = SamplerParams {
|
||||
temperature: Some(0.0),
|
||||
..params
|
||||
};
|
||||
let _ = build_sampler(&greedy);
|
||||
}
|
||||
|
||||
// ── build_prompt ──
|
||||
|
||||
#[test]
|
||||
fn build_prompt_renders_messages() {
|
||||
let messages = vec![msg("system", "You are helpful."), msg("user", "Hi!")];
|
||||
let prompt = build_prompt(&messages, &None).unwrap();
|
||||
assert!(prompt.starts_with("<s>"));
|
||||
assert!(prompt.contains("<|im_start|>system\nYou are helpful."));
|
||||
assert!(prompt.contains("<|im_start|>user\nHi!"));
|
||||
assert!(prompt.contains("<|im_start|>assistant\n<think>\n"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_prompt_includes_tool_definitions() {
|
||||
let messages = vec![msg("user", "What's the weather?")];
|
||||
let prompt = build_prompt(&messages, &Some(vec![tool_def()])).unwrap();
|
||||
assert!(prompt.contains("<tools>"));
|
||||
assert!(prompt.contains("get_weather"));
|
||||
assert!(prompt.contains("Tool usage guidelines"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_prompt_wraps_tool_response_once() {
|
||||
let messages = vec![msg("user", "Weather?"), msg("tool", "Sunny")];
|
||||
let prompt = build_prompt(&messages, &None).unwrap();
|
||||
assert_eq!(prompt.matches("<tool_response>").count(), 1);
|
||||
assert_eq!(prompt.matches("</tool_response>").count(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_prompt_renders_assistant_tool_calls_history() {
|
||||
let assistant = ChatMessage {
|
||||
role: "assistant".into(),
|
||||
content: Some("".into()),
|
||||
tool_calls: Some(vec![ToolCallResponse {
|
||||
id: "call_1".into(),
|
||||
call_type: "function".into(),
|
||||
function: ToolCallFunction {
|
||||
name: "get_weather".into(),
|
||||
arguments: r#"{"city":"Jakarta"}"#.into(),
|
||||
},
|
||||
}]),
|
||||
tool_call_id: None,
|
||||
name: None,
|
||||
};
|
||||
let prompt = build_prompt(&[assistant], &None).unwrap();
|
||||
assert!(prompt.contains("<function=get_weather>"));
|
||||
assert!(prompt.contains("<parameter=city>Jakarta</parameter>"));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user