From 8aa16cb20abe6d36d977a25393c5b97a214fc0e2 Mon Sep 17 00:00:00 2001 From: Pekka Heikura Date: Sat, 30 May 2026 09:43:17 +0300 Subject: [PATCH 1/2] feat: pluggable tool-call adapter + OpenAI tool support (closes #95, #96, #97) Introduces IToolCallAdapter (Core) and ToolCallAdapterRegistry keyed on GGUF architecture. Ships adapters for Qwen2/3 ( wrapper), Qwen3-Coder (bare ), Llama-3 (<|python_tag|>...<|eom_id|>), and DeepSeek-R1 (<|tool_calls_begin|>...<|tool_calls_end|>). ChatTemplateRenderer resolves the adapter once at engine-load time. AnthropicEndpoints: replaces the hard-coded matcher in HandleStreaming + non-streaming ParseToolCalls with the adapter's open/close marker API. tool_result history-replay now goes through adapter.RenderToolResult so families like Llama-3 (role: ipython) replay correctly. OpenAiEndpoints: adds Tools / ToolChoice on ChatCompletionRequest, tool_calls on OaiAssistantMessage + ChunkDelta, OaiToolCall/OaiToolCallDelta records, history role=tool messages, finish_reason=tool_calls. Streaming buffering activates only when tools are in the request so the no-tools per-chunk cadence is preserved. Tests: 21 adapter unit tests (Qwen, Qwen3-Coder, Llama, DeepSeek + registry), 7 endpoint tests covering Qwen3-Coder /v1/messages non-stream + stream (closes #95), OpenAI /v1/chat/completions non-stream + stream + tool history + no-tools regression guard. Full suite green: 505/505. Co-Authored-By: Claude Opus 4.7 --- src/SharpInference.Core/ToolCallAdapter.cs | 569 ++++++++++++++++++ src/SharpInference.Server/ChatTemplate.cs | 10 + .../Endpoints/AnthropicEndpoints.cs | 149 +++-- .../Endpoints/OpenAiEndpoints.cs | 427 +++++++++++-- .../SharpInferenceJsonContext.cs | 8 + .../ToolCallAdapterTests.cs | 242 ++++++++ .../ToolCallEndpointTests.cs | 287 +++++++++ 7 files changed, 1567 insertions(+), 125 deletions(-) create mode 100644 src/SharpInference.Core/ToolCallAdapter.cs create mode 100644 tests/SharpInference.Tests.Core/ToolCallAdapterTests.cs create mode 100644 tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs diff --git a/src/SharpInference.Core/ToolCallAdapter.cs b/src/SharpInference.Core/ToolCallAdapter.cs new file mode 100644 index 0000000..d34bb10 --- /dev/null +++ b/src/SharpInference.Core/ToolCallAdapter.cs @@ -0,0 +1,569 @@ +using System.Text; +using System.Text.Json; + +namespace SharpInference.Core; + +/// +/// Translation layer between an LLM's free-text tool-call output and the +/// structured shape the API endpoints emit. One +/// implementation per model-family wire format (Qwen3, Qwen3-Coder, Llama-3, +/// DeepSeek-R1, ...). Resolved at engine-load time via +/// keyed on the GGUF +/// general.architecture field. +/// +/// +/// The interface covers both batched (non-streaming) parsing and an +/// open/close-marker model used by the streaming endpoint state machine. +/// +/// +public interface IToolCallAdapter +{ + /// Architecture key this adapter handles (e.g. "qwen3moe"). + string Architecture { get; } + + /// + /// Extracts structured tool calls from a complete model output. Returns the + /// plain text (with tool-call blocks stripped) plus the parsed calls in the + /// order they appeared. Used by the non-streaming endpoint code paths. + /// + (string PlainText, IReadOnlyList Calls) Parse(string rawOutput); + + /// + /// Length of the longest open marker this adapter recognises. The streaming + /// endpoint keeps up to MaxOpenTagLength - 1 trailing bytes buffered so + /// a marker that arrives split across chunks is still matched. + /// + int MaxOpenTagLength { get; } + + /// + /// Scans at or after + /// for a tool-call open marker. Returns the index of the marker's first byte + /// (the streamer flushes everything before it as plain text) and writes the + /// start of the buffered block region to . + /// For most adapters this is past the marker; adapters whose call name lives + /// inside the marker (e.g. Qwen3-Coder's <function=NAME>) return + /// the same index for both, so the marker stays in the block. Returns -1 when + /// no marker is present in the buffer. + /// + int FindOpenMarker(string buffer, int startSearch, out int contentStart); + + /// + /// Scans at or after + /// for the matching close marker. The return value is the EXCLUSIVE end index + /// of the buffered block region (i.e. buf[contentStart..returned] is + /// what gets handed to ); for most adapters that's + /// the marker's first byte, but adapters whose parser needs the close marker + /// itself (e.g. Qwen3-Coder) return one-past-end. + /// receives the index past the close marker (where streaming resumes scanning). + /// Returns -1 when the close marker is not yet present. + /// + int FindCloseMarker(string buffer, int startSearch, out int afterClose); + + /// + /// Parses one or more tool calls out of a buffered region. The slice handed in + /// matches the contract documented on and + /// . Multi-call wire formats (e.g. DeepSeek's + /// tool_calls_begin..end wrapper) may return more than one call. + /// + IReadOnlyList ParseBlock(string block); + + /// + /// Builds the role/content dictionary for replaying a prior tool result on + /// this model's chat template. Default Qwen/OpenAI shape is + /// {role:"tool", content:string}; some families wrap differently. + /// + Dictionary RenderToolResult(string toolUseId, string resultContent); +} + +/// +/// Process-wide registry of instances keyed by GGUF +/// architecture. Ships defaults for the families enumerated in issue #96; callers +/// can additional adapters at startup. Lookups for unknown +/// architectures fall through to the Qwen3-style wrapper adapter (no regression +/// for any model whose template already emits <tool_call> blocks). +/// +public static class ToolCallAdapterRegistry +{ + private static readonly Dictionary _adapters = + new(StringComparer.Ordinal) + { + // Wrapper-style: Qwen2, Qwen3, Qwen3.x, Qwen3-MoE variants — all use + // the ... envelope; supports both Qwen3.6 XML + // () and standard JSON ({"name":"...", ...}) + // payloads inside. + ["qwen2"] = new QwenToolCallAdapter("qwen2"), + ["qwen3"] = new QwenToolCallAdapter("qwen3"), + ["qwen3moe"] = new QwenToolCallAdapter("qwen3moe"), + ["qwen35moe"] = new QwenToolCallAdapter("qwen35moe"), + + // Qwen3-Coder: bare ... with no + // wrapper. Closes #95. + ["qwen3coder"] = new QwenCoderToolCallAdapter(), + + // Llama-3.x tool-calling format: <|python_tag|>{json}<|eom_id|> + ["llama"] = new LlamaToolCallAdapter(), + ["llama4"] = new LlamaToolCallAdapter(), + + // DeepSeek-R1 wrapped multi-call envelope. + ["deepseek2"] = new DeepSeekToolCallAdapter(), + }; + + /// + /// Fallback used when no adapter is registered for the requested architecture. + /// The Qwen3 wrapper adapter is the broadest compatible default — every model + /// whose chat template renders OpenAI-style tools tends to emit + /// <tool_call> blocks under it. + /// + public static IToolCallAdapter DefaultAdapter { get; } = new QwenToolCallAdapter("default"); + + /// Returns the adapter registered for , or the default. + public static IToolCallAdapter Get(string? architecture) + { + if (architecture is { Length: > 0 } && _adapters.TryGetValue(architecture, out var a)) + return a; + return DefaultAdapter; + } + + /// Registers (or replaces) the adapter for . + public static void Register(string architecture, IToolCallAdapter adapter) + { + ArgumentException.ThrowIfNullOrEmpty(architecture); + ArgumentNullException.ThrowIfNull(adapter); + _adapters[architecture] = adapter; + } +} + +// ── Adapter base helpers ────────────────────────────────────────────────────── + +internal static class ToolCallParseHelpers +{ + public static Dictionary DefaultRenderToolResult(string _, string resultContent) => + new(StringComparer.Ordinal) { ["role"] = "tool", ["content"] = resultContent }; + + /// + /// Parses a single <function=name><parameter=k>v</parameter>...</function> + /// block. Tolerant of a missing trailing </function> (Qwen3-Coder occasionally + /// stops generating before emitting it). + /// + public static ParsedToolCall? ParseXmlFunctionBlock(string block) + { + const string funcTag = "', funcStart); + if (nameEnd < 0) return null; + string funcName = block[(funcStart + funcTag.Length)..nameEnd].Trim(); + if (funcName.Length == 0) return null; + + int funcBodyEnd = block.IndexOf(funcClose, nameEnd, StringComparison.Ordinal); + string funcBody = funcBodyEnd >= 0 + ? block[(nameEnd + 1)..funcBodyEnd] + : block[(nameEnd + 1)..]; + + var args = new Dictionary(StringComparer.Ordinal); + int p = 0; + while (p < funcBody.Length) + { + int paramStart = funcBody.IndexOf(paramTag, p, StringComparison.Ordinal); + if (paramStart < 0) break; + + int paramNameEnd = funcBody.IndexOf('>', paramStart); + if (paramNameEnd < 0) break; + string paramName = funcBody[(paramStart + paramTag.Length)..paramNameEnd].Trim(); + + int valueStart = paramNameEnd + 1; + int paramEnd = funcBody.IndexOf(paramClose, valueStart, StringComparison.Ordinal); + if (paramEnd < 0) break; + + args[paramName] = funcBody[valueStart..paramEnd].Trim(); + p = paramEnd + paramClose.Length; + } + + return new ParsedToolCall(funcName, args); + } + + /// + /// Parses a {"name":"...","arguments":{...}} JSON object (Qwen3 default + /// payload and Llama-3 emit shape). arguments may itself be a JSON + /// string (some templates double-encode); both shapes are accepted. When + /// is absent the alternate key (arguments + /// vs parameters) is tried as a fallback — Llama-3 fine-tunes drift + /// between both spellings. + /// + public static ParsedToolCall? ParseJsonCallBlock(string block, string argumentsKey = "arguments") + { + try + { + using var doc = JsonDocument.Parse(block); + var root = doc.RootElement; + string name = root.TryGetProperty("name", out var n) ? n.GetString() ?? "" : ""; + if (name.Length == 0) return null; + + var args = new Dictionary(StringComparer.Ordinal); + string altKey = argumentsKey == "arguments" ? "parameters" : "arguments"; + if (root.TryGetProperty(argumentsKey, out var argsElem) + || root.TryGetProperty(altKey, out argsElem)) + { + object? parsed = argsElem.ValueKind == JsonValueKind.String + ? JsonElementToObject(JsonDocument.Parse(argsElem.GetString() ?? "{}").RootElement) + : JsonElementToObject(argsElem); + if (parsed is Dictionary d) args = d; + } + + return new ParsedToolCall(name, args); + } + catch (JsonException) + { + return null; + } + } + + public static object? JsonElementToObject(JsonElement el) => el.ValueKind switch + { + JsonValueKind.Object => el.EnumerateObject() + .ToDictionary(p => p.Name, + p => JsonElementToObject(p.Value), + StringComparer.Ordinal), + JsonValueKind.Array => el.EnumerateArray() + .Select(JsonElementToObject) + .ToList(), + JsonValueKind.String => el.GetString(), + JsonValueKind.Number => el.TryGetInt64(out long l) ? l + : el.TryGetDouble(out double d) ? d + : (object?)el.GetRawText(), + JsonValueKind.True => true, + JsonValueKind.False => false, + _ => null, + }; +} + +// ── Concrete adapters ───────────────────────────────────────────────────────── + +/// +/// Qwen2/3 family: model output wraps each call in +/// <tool_call>...</tool_call>. The payload inside is either +/// standard JSON ({"name":"...","arguments":{...}}) or — on the +/// Qwen3.6 line — the XML shape <function=name><parameter=...>...</parameter></function>. +/// +public sealed class QwenToolCallAdapter(string architecture) : IToolCallAdapter +{ + public const string OpenMarker = ""; + public const string CloseMarker = ""; + + public string Architecture { get; } = architecture; + public int MaxOpenTagLength => OpenMarker.Length; + + public (string PlainText, IReadOnlyList Calls) Parse(string rawOutput) + { + var calls = new List(); + var plain = new StringBuilder(); + int pos = 0; + + while (pos < rawOutput.Length) + { + int start = rawOutput.IndexOf(OpenMarker, pos, StringComparison.Ordinal); + if (start < 0) + { + plain.Append(rawOutput, pos, rawOutput.Length - pos); + break; + } + + plain.Append(rawOutput, pos, start - pos); + int contentStart = start + OpenMarker.Length; + int end = rawOutput.IndexOf(CloseMarker, contentStart, StringComparison.Ordinal); + if (end < 0) break; + + string block = rawOutput[contentStart..end].Trim(); + pos = end + CloseMarker.Length; + foreach (var c in ParseBlock(block)) calls.Add(c); + } + + return (plain.ToString(), calls); + } + + public int FindOpenMarker(string buffer, int startSearch, out int contentStart) + { + int idx = buffer.IndexOf(OpenMarker, startSearch, StringComparison.Ordinal); + contentStart = idx < 0 ? -1 : idx + OpenMarker.Length; + return idx; + } + + public int FindCloseMarker(string buffer, int startSearch, out int afterClose) + { + int idx = buffer.IndexOf(CloseMarker, startSearch, StringComparison.Ordinal); + afterClose = idx < 0 ? -1 : idx + CloseMarker.Length; + return idx; + } + + public IReadOnlyList ParseBlock(string block) + { + block = block.Trim(); + ParsedToolCall? call = block.StartsWith(" RenderToolResult(string toolUseId, string resultContent) => + ToolCallParseHelpers.DefaultRenderToolResult(toolUseId, resultContent); +} + +/// +/// Qwen3-Coder: emits bare <function=name>...</function> with +/// no surrounding <tool_call> envelope. Closes #95. +/// +public sealed class QwenCoderToolCallAdapter : IToolCallAdapter +{ + public const string OpenMarker = " "qwen3coder"; + public int MaxOpenTagLength => OpenMarker.Length; + + public (string PlainText, IReadOnlyList Calls) Parse(string rawOutput) + { + var calls = new List(); + var plain = new StringBuilder(); + int pos = 0; + + while (pos < rawOutput.Length) + { + int start = rawOutput.IndexOf(OpenMarker, pos, StringComparison.Ordinal); + if (start < 0) + { + plain.Append(rawOutput, pos, rawOutput.Length - pos); + break; + } + + plain.Append(rawOutput, pos, start - pos); + int end = rawOutput.IndexOf(CloseMarker, start, StringComparison.Ordinal); + // Block must include the prefix so the parser can read the name. + string block = end >= 0 + ? rawOutput[start..(end + CloseMarker.Length)] + : rawOutput[start..]; + pos = end >= 0 ? end + CloseMarker.Length : rawOutput.Length; + + var call = ToolCallParseHelpers.ParseXmlFunctionBlock(block); + if (call.HasValue) calls.Add(call.Value); + } + + return (plain.ToString(), calls); + } + + public int FindOpenMarker(string buffer, int startSearch, out int contentStart) + { + int idx = buffer.IndexOf(OpenMarker, startSearch, StringComparison.Ordinal); + // contentStart points AT the open marker (not past it) — Qwen3-Coder's call name + // lives inside the open marker (), so the block we hand to + // ParseBlock must include it. + contentStart = idx; + return idx; + } + + public int FindCloseMarker(string buffer, int startSearch, out int afterClose) + { + int idx = buffer.IndexOf(CloseMarker, startSearch, StringComparison.Ordinal); + if (idx < 0) { afterClose = -1; return -1; } + // Include the close marker in the buffered region: the block we hand to + // ParseBlock must include so ParseXmlFunctionBlock's name-end + // scan terminates correctly. + afterClose = idx + CloseMarker.Length; + return afterClose; + } + + public IReadOnlyList ParseBlock(string block) + { + var call = ToolCallParseHelpers.ParseXmlFunctionBlock(block.Trim()); + return call.HasValue ? [call.Value] : []; + } + + public Dictionary RenderToolResult(string toolUseId, string resultContent) => + ToolCallParseHelpers.DefaultRenderToolResult(toolUseId, resultContent); +} + +/// +/// Llama-3.x tool-calling: <|python_tag|>{"name":"...","parameters":{...}}<|eom_id|>. +/// The close marker may also be <|eot_id|> on shorter outputs — both are recognised. +/// +public sealed class LlamaToolCallAdapter : IToolCallAdapter +{ + public const string OpenMarker = "<|python_tag|>"; + public const string EomMarker = "<|eom_id|>"; + public const string EotMarker = "<|eot_id|>"; + + public string Architecture => "llama"; + public int MaxOpenTagLength => OpenMarker.Length; + + public (string PlainText, IReadOnlyList Calls) Parse(string rawOutput) + { + var calls = new List(); + var plain = new StringBuilder(); + int pos = 0; + + while (pos < rawOutput.Length) + { + int start = rawOutput.IndexOf(OpenMarker, pos, StringComparison.Ordinal); + if (start < 0) + { + plain.Append(rawOutput, pos, rawOutput.Length - pos); + break; + } + + plain.Append(rawOutput, pos, start - pos); + int contentStart = start + OpenMarker.Length; + int end = FindEarliest(rawOutput, contentStart, EomMarker, EotMarker, out int afterClose); + string block = end >= 0 ? rawOutput[contentStart..end] : rawOutput[contentStart..]; + pos = afterClose >= 0 ? afterClose : rawOutput.Length; + + foreach (var c in ParseBlock(block)) calls.Add(c); + } + + return (plain.ToString(), calls); + } + + public int FindOpenMarker(string buffer, int startSearch, out int contentStart) + { + int idx = buffer.IndexOf(OpenMarker, startSearch, StringComparison.Ordinal); + contentStart = idx < 0 ? -1 : idx + OpenMarker.Length; + return idx; + } + + public int FindCloseMarker(string buffer, int startSearch, out int afterClose) + { + int idx = FindEarliest(buffer, startSearch, EomMarker, EotMarker, out afterClose); + return idx; + } + + public IReadOnlyList ParseBlock(string block) + { + // Llama-3 emits {"name":"...", "parameters":{...}}; some fine-tunes still use + // "arguments". ParseJsonCallBlock handles both keys. + var call = ToolCallParseHelpers.ParseJsonCallBlock(block.Trim(), argumentsKey: "parameters"); + return call.HasValue ? [call.Value] : []; + } + + public Dictionary RenderToolResult(string toolUseId, string resultContent) => + new(StringComparer.Ordinal) + { + ["role"] = "ipython", // Llama-3 echoes tool results under role=ipython. + ["content"] = resultContent, + }; + + private static int FindEarliest(string buffer, int startSearch, string a, string b, out int afterClose) + { + int ia = buffer.IndexOf(a, startSearch, StringComparison.Ordinal); + int ib = buffer.IndexOf(b, startSearch, StringComparison.Ordinal); + int idx = (ia, ib) switch + { + (< 0, < 0) => -1, + (< 0, _) => ib, + (_, < 0) => ia, + _ => Math.Min(ia, ib), + }; + if (idx < 0) { afterClose = -1; return -1; } + afterClose = idx + (idx == ia ? a.Length : b.Length); + return idx; + } +} + +/// +/// DeepSeek-R1 family: <|tool_calls_begin|>...<|tool_calls_end|> +/// wraps one or more <|tool_call_begin|>name<|tool_sep|>{json}<|tool_call_end|> +/// inner blocks. +/// +public sealed class DeepSeekToolCallAdapter : IToolCallAdapter +{ + public const string OuterOpen = "<|tool_calls_begin|>"; + public const string OuterClose = "<|tool_calls_end|>"; + public const string InnerOpen = "<|tool_call_begin|>"; + public const string InnerClose = "<|tool_call_end|>"; + public const string Separator = "<|tool_sep|>"; + + public string Architecture => "deepseek2"; + public int MaxOpenTagLength => OuterOpen.Length; + + public (string PlainText, IReadOnlyList Calls) Parse(string rawOutput) + { + var calls = new List(); + var plain = new StringBuilder(); + int pos = 0; + + while (pos < rawOutput.Length) + { + int start = rawOutput.IndexOf(OuterOpen, pos, StringComparison.Ordinal); + if (start < 0) + { + plain.Append(rawOutput, pos, rawOutput.Length - pos); + break; + } + + plain.Append(rawOutput, pos, start - pos); + int contentStart = start + OuterOpen.Length; + int end = rawOutput.IndexOf(OuterClose, contentStart, StringComparison.Ordinal); + string block = end >= 0 ? rawOutput[contentStart..end] : rawOutput[contentStart..]; + pos = end >= 0 ? end + OuterClose.Length : rawOutput.Length; + + foreach (var c in ParseBlock(block)) calls.Add(c); + } + + return (plain.ToString(), calls); + } + + public int FindOpenMarker(string buffer, int startSearch, out int contentStart) + { + int idx = buffer.IndexOf(OuterOpen, startSearch, StringComparison.Ordinal); + contentStart = idx < 0 ? -1 : idx + OuterOpen.Length; + return idx; + } + + public int FindCloseMarker(string buffer, int startSearch, out int afterClose) + { + int idx = buffer.IndexOf(OuterClose, startSearch, StringComparison.Ordinal); + afterClose = idx < 0 ? -1 : idx + OuterClose.Length; + return idx; + } + + public IReadOnlyList ParseBlock(string block) + { + var results = new List(); + int p = 0; + while (p < block.Length) + { + int innerStart = block.IndexOf(InnerOpen, p, StringComparison.Ordinal); + if (innerStart < 0) break; + int contentStart = innerStart + InnerOpen.Length; + int innerEnd = block.IndexOf(InnerClose, contentStart, StringComparison.Ordinal); + if (innerEnd < 0) break; + + string inner = block[contentStart..innerEnd]; + p = innerEnd + InnerClose.Length; + + int sep = inner.IndexOf(Separator, StringComparison.Ordinal); + if (sep < 0) continue; + + string name = inner[..sep].Trim(); + string argsJson = inner[(sep + Separator.Length)..].Trim(); + if (name.Length == 0) continue; + + var argDict = new Dictionary(StringComparer.Ordinal); + try + { + using var doc = JsonDocument.Parse(argsJson); + if (ToolCallParseHelpers.JsonElementToObject(doc.RootElement) is Dictionary d) + argDict = d; + } + catch (JsonException) { /* keep empty arg dict for malformed JSON */ } + + results.Add(new ParsedToolCall(name, argDict)); + } + return results; + } + + public Dictionary RenderToolResult(string toolUseId, string resultContent) => + ToolCallParseHelpers.DefaultRenderToolResult(toolUseId, resultContent); +} diff --git a/src/SharpInference.Server/ChatTemplate.cs b/src/SharpInference.Server/ChatTemplate.cs index 3a1407a..e4e864e 100644 --- a/src/SharpInference.Server/ChatTemplate.cs +++ b/src/SharpInference.Server/ChatTemplate.cs @@ -59,6 +59,7 @@ public sealed class ChatTemplateRenderer { private JinjaChatTemplate? _template; private string _architecture; + private IToolCallAdapter _toolCallAdapter; /// Architecture string used to pick a fallback template when no Jinja is loaded. public string Architecture => _architecture; @@ -66,12 +67,20 @@ public sealed class ChatTemplateRenderer /// Compiled Jinja template, if the loaded model shipped one. public JinjaChatTemplate? JinjaTemplate => _template; + /// + /// Tool-call translation layer matched to the loaded model's architecture. Resolved + /// once at time from ; + /// endpoint code reads this without having to look up by string each request. + /// + public IToolCallAdapter ToolCallAdapter => _toolCallAdapter; + /// Default architecture (used both for fallback and exposed via ). /// Optional compiled Jinja template; null means "use the hardcoded fallback". public ChatTemplateRenderer(string architecture = "qwen2", JinjaChatTemplate? template = null) { _architecture = architecture; _template = template; + _toolCallAdapter = ToolCallAdapterRegistry.Get(architecture); } /// @@ -83,6 +92,7 @@ public void Configure(string architecture, JinjaChatTemplate? template) { _architecture = architecture; _template = template; + _toolCallAdapter = ToolCallAdapterRegistry.Get(architecture); } /// Messages in order (system, user, assistant, ...). diff --git a/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs b/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs index 977a6a7..987de61 100644 --- a/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs +++ b/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs @@ -54,10 +54,12 @@ await ctx.Response.WriteAsync( // once that many reasoning tokens have streamed. bool enableThinking = req.Thinking?.Type != "disabled"; + var adapter = chatTemplate.ToolCallAdapter; + string prompt; if (req.Tools is { Length: > 0 }) { - var (richMessages, tools) = BuildRichMessageList(req); + var (richMessages, tools) = BuildRichMessageList(req, adapter); prompt = chatTemplate.Format(richMessages, enableThinking, tools); } else @@ -78,17 +80,17 @@ await ctx.Response.WriteAsync( if (req.Stream == true) { - await HandleStreaming(ctx, engine, metrics, prompt, sp, msgId, modelId); + await HandleStreaming(ctx, engine, metrics, adapter, prompt, sp, msgId, modelId); } else { - await HandleNonStreaming(ctx, engine, metrics, prompt, sp, msgId, modelId); + await HandleNonStreaming(ctx, engine, metrics, adapter, prompt, sp, msgId, modelId); } } private static async Task HandleNonStreaming( - HttpContext ctx, IInferenceEngine engine, ServerMetrics metrics, string prompt, SamplingParams sp, - string msgId, string modelId) + HttpContext ctx, IInferenceEngine engine, ServerMetrics metrics, IToolCallAdapter adapter, + string prompt, SamplingParams sp, string msgId, string modelId) { var thinkingSb = new StringBuilder(); var textSb = new StringBuilder(); @@ -114,7 +116,7 @@ private static async Task HandleNonStreaming( } var rawText = textSb.ToString(); - var (plainText, toolCalls) = ParseToolCalls(rawText); + var (plainText, toolCalls) = ParseToolCalls(adapter, rawText); if (plainText.Length > 0) contentList.Add(new AContent("text", Text: plainText)); @@ -146,8 +148,8 @@ await ctx.Response.WriteAsync( } private static async Task HandleStreaming( - HttpContext ctx, IInferenceEngine engine, ServerMetrics metrics, string prompt, SamplingParams sp, - string msgId, string modelId) + HttpContext ctx, IInferenceEngine engine, ServerMetrics metrics, IToolCallAdapter adapter, + string prompt, SamplingParams sp, string msgId, string modelId) { ctx.Response.ContentType = "text/event-stream"; ctx.Response.Headers.CacheControl = "no-cache"; @@ -172,13 +174,14 @@ await WriteAnthropicEvent(ctx.Response, "message_start", int outputTokens = 0; bool hasToolCalls = false; - // Tool-call streaming state machine. - // Text output is buffered to detect "" tags before flushing as text_delta. - const string ToolCallOpenTag = ""; - const string ToolCallCloseTag = ""; + // Tool-call streaming state machine — adapter-driven so the open/close markers + // match the loaded model's wire format (Qwen3 , Qwen3-Coder , + // Llama-3 <|python_tag|>, DeepSeek-R1 <|tool_calls_begin|>, ...). + int maxOpenLen = adapter.MaxOpenTagLength; bool inToolCall = false; + int toolCallContentStart = -1; // index into toolCallBuf where call content begins var toolCallBuf = new StringBuilder(); - string pendingText = ""; // holds text while scanning for ToolCallOpenTag + string pendingText = ""; async Task FlushTextDelta(string text) { @@ -213,62 +216,64 @@ await WriteAnthropicEvent(ctx.Response, "content_block_delta", JsonSerializer.Serialize(delta, SharpInferenceJsonContext.Default.AContentBlockDeltaEvent)); } - async Task EmitToolCallBlock(string blockContent) + async Task EmitToolCallsFromBlock(string blockContent) { - var (_, calls) = JinjaChatTemplate.ParseToolCalls( - $"{blockContent}"); - if (calls.Count == 0) return; - - var tc = calls[0]; - string name = tc.Name; - string argsJson = JinjaChatTemplate.SerializeToJson(tc.Arguments); - - // Close any open text block. - if (textOpen) + var calls = adapter.ParseBlock(blockContent); + foreach (var tc in calls) { - await WriteAnthropicEvent(ctx.Response, "content_block_stop", - JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", textIndex), - SharpInferenceJsonContext.Default.AContentBlockStopEvent)); - textOpen = false; - nextBlockIndex = textIndex + 1; - } + string name = tc.Name; + string argsJson = JinjaChatTemplate.SerializeToJson(tc.Arguments); + + // Close any open text block at the first tool call. + if (textOpen) + { + await WriteAnthropicEvent(ctx.Response, "content_block_stop", + JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", textIndex), + SharpInferenceJsonContext.Default.AContentBlockStopEvent)); + textOpen = false; + nextBlockIndex = textIndex + 1; + } - var id = $"toolu_{Guid.NewGuid():N}"; - int toolIdx = nextBlockIndex++; - hasToolCalls = true; + var id = $"toolu_{Guid.NewGuid():N}"; + int toolIdx = nextBlockIndex++; + hasToolCalls = true; - var toolStart = new AContentBlockStartEvent("content_block_start", toolIdx, - new AContentBlock("tool_use", Id: id, Name: name, Input: EmptyJsonObject)); - await WriteAnthropicEvent(ctx.Response, "content_block_start", - JsonSerializer.Serialize(toolStart, SharpInferenceJsonContext.Default.AContentBlockStartEvent)); + var toolStart = new AContentBlockStartEvent("content_block_start", toolIdx, + new AContentBlock("tool_use", Id: id, Name: name, Input: EmptyJsonObject)); + await WriteAnthropicEvent(ctx.Response, "content_block_start", + JsonSerializer.Serialize(toolStart, SharpInferenceJsonContext.Default.AContentBlockStartEvent)); - var inputDelta = new AContentBlockDeltaEvent("content_block_delta", toolIdx, - new AContentDelta("input_json_delta", PartialJson: argsJson)); - await WriteAnthropicEvent(ctx.Response, "content_block_delta", - JsonSerializer.Serialize(inputDelta, SharpInferenceJsonContext.Default.AContentBlockDeltaEvent)); + var inputDelta = new AContentBlockDeltaEvent("content_block_delta", toolIdx, + new AContentDelta("input_json_delta", PartialJson: argsJson)); + await WriteAnthropicEvent(ctx.Response, "content_block_delta", + JsonSerializer.Serialize(inputDelta, SharpInferenceJsonContext.Default.AContentBlockDeltaEvent)); - await WriteAnthropicEvent(ctx.Response, "content_block_stop", - JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", toolIdx), - SharpInferenceJsonContext.Default.AContentBlockStopEvent)); + await WriteAnthropicEvent(ctx.Response, "content_block_stop", + JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", toolIdx), + SharpInferenceJsonContext.Default.AContentBlockStopEvent)); + } } // Process a text-stream chunk through the tool-call detection state machine. - // Calls FlushTextDelta/EmitToolCallBlock — must not be called after the main - // loop's finally block starts closing blocks. + // Adapter-driven so a different model family (e.g. Qwen3-Coder, Llama-3) plugs + // its own open/close marker scan in without changes here. async Task ProcessTextChunk(string chunk) { if (inToolCall) { toolCallBuf.Append(chunk); string buf = toolCallBuf.ToString(); - int closeIdx = buf.IndexOf(ToolCallCloseTag, StringComparison.Ordinal); + int closeIdx = adapter.FindCloseMarker(buf, toolCallContentStart, out int afterClose); if (closeIdx >= 0) { - string json = buf[..closeIdx]; - string remaining = buf[(closeIdx + ToolCallCloseTag.Length)..]; + // Block content is whatever lies between the open marker's contentStart + // and the close marker's start index. afterClose points past the marker. + string block = buf[toolCallContentStart..closeIdx]; + string remaining = buf[afterClose..]; toolCallBuf.Clear(); + toolCallContentStart = -1; inToolCall = false; - await EmitToolCallBlock(json); + await EmitToolCallsFromBlock(block); if (remaining.Length > 0) await ProcessTextChunk(remaining); } @@ -277,23 +282,31 @@ async Task ProcessTextChunk(string chunk) pendingText += chunk; - int openIdx = pendingText.IndexOf(ToolCallOpenTag, StringComparison.Ordinal); + int openIdx = adapter.FindOpenMarker(pendingText, 0, out int contentStart); if (openIdx >= 0) { if (openIdx > 0) await FlushTextDelta(pendingText[..openIdx]); - string afterTag = pendingText[(openIdx + ToolCallOpenTag.Length)..]; - pendingText = ""; + + // Capture whatever the adapter considers "block content" — for most adapters + // that's everything past the open marker, but Qwen3-Coder's name lives inside + // the open marker so its contentStart points AT the marker, not past it. inToolCall = true; toolCallBuf.Clear(); - if (afterTag.Length > 0) - await ProcessTextChunk(afterTag); + toolCallBuf.Append(pendingText, contentStart, pendingText.Length - contentStart); + toolCallContentStart = 0; + pendingText = ""; + + // The newly buffered region may itself already contain a complete close marker + // (e.g. when a single chunk delivered the whole call). Re-enter to check. + if (toolCallBuf.Length > 0) + await ProcessTextChunk(""); return; } - // No found: flush everything except the last (tag-length - 1) chars - // which might be the start of a partial tag match. - int safeLen = Math.Max(0, pendingText.Length - (ToolCallOpenTag.Length - 1)); + // No open marker found: flush everything except the last (maxOpenLen - 1) chars + // which might be the start of a partial marker match. + int safeLen = Math.Max(0, pendingText.Length - (maxOpenLen - 1)); if (safeLen > 0) { await FlushTextDelta(pendingText[..safeLen]); @@ -425,14 +438,14 @@ private static string MakeSignatureStub(string thinking) private static readonly JsonElement EmptyJsonObject = JsonDocument.Parse("{}").RootElement.Clone(); /// - /// Parses Qwen3-style <tool_call>...</tool_call> blocks from model output. - /// Supports both Qwen3.6 XML format and standard JSON format. - /// Returns the plain text (with tool_call tags removed) and a list of parsed tool calls. + /// Extracts tool calls from raw model output via the adapter for the loaded model. + /// Returns the plain text (with tool-call blocks stripped) and a list of parsed + /// calls tagged with Anthropic-style toolu_ identifiers. /// private static (string text, List<(string id, string name, string argsJson)> toolCalls) - ParseToolCalls(string output) + ParseToolCalls(IToolCallAdapter adapter, string output) { - var (plainText, calls) = JinjaChatTemplate.ParseToolCalls(output); + var (plainText, calls) = adapter.Parse(output); var toolCalls = calls .Select(c => ($"toolu_{Guid.NewGuid():N}", c.Name, JinjaChatTemplate.SerializeToJson(c.Arguments))) .ToList(); @@ -463,10 +476,11 @@ private static (string text, List<(string id, string name, string argsJson)> too /// Builds rich message dictionaries and converts Anthropic tool definitions to /// the OpenAI/Qwen3 function format expected by the Jinja chat template. /// Handles multi-content messages: tool_use content blocks become tool_calls - /// entries on assistant messages; tool_result blocks become role="tool" messages. + /// entries on assistant messages; tool_result blocks become the role/content shape + /// the adapter specifies (defaults to role="tool"). /// private static (List> messages, List? tools) - BuildRichMessageList(AnthropicMessageRequest req) + BuildRichMessageList(AnthropicMessageRequest req, IToolCallAdapter adapter) { var messages = new List>(); @@ -548,7 +562,10 @@ private static (List> messages, List? tools ? c.GetString() ?? "" : ExtractTextFromArray(c); } - messages.Add(new() { ["role"] = "tool", ["content"] = resultContent }); + string toolUseId = block.TryGetProperty("tool_use_id", out var tid) + ? tid.GetString() ?? "" + : ""; + messages.Add(adapter.RenderToolResult(toolUseId, resultContent)); } } diff --git a/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs b/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs index 7d00915..e8113aa 100644 --- a/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs +++ b/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs @@ -5,6 +5,7 @@ using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Routing; using Microsoft.Extensions.Options; +using SharpInference.Core; using SharpInference.Engine; namespace SharpInference.Server.Endpoints; @@ -50,8 +51,23 @@ await ctx.Response.WriteAsync( metrics.RecordRequest(); bool enableThinking = req.EnableThinking ?? true; - var messages = BuildMessageList(req.Messages, req.ResponseFormat?.Type); - var prompt = chatTemplate.Format(messages, enableThinking); + var adapter = chatTemplate.ToolCallAdapter; + + // Tool-aware rendering: if either tool definitions or a history-side + // tool_call / tool message is present, route through the rich-message path + // so the chat template can inject the tool schema and replay prior calls. + bool toolsActive = req.Tools is { Length: > 0 } || HasToolMessages(req.Messages); + string prompt; + if (toolsActive) + { + var (richMessages, tools) = BuildRichMessageList(req, adapter); + prompt = chatTemplate.Format(richMessages, enableThinking, tools); + } + else + { + var messages = BuildMessageList(req.Messages, req.ResponseFormat?.Type); + prompt = chatTemplate.Format(messages, enableThinking); + } // Parse logit_bias: OpenAI sends {"tokenId": biasValue} with string keys IReadOnlyDictionary? logitBias = null; @@ -75,81 +91,235 @@ await ctx.Response.WriteAsync( if (req.Stream == true) { - ctx.Response.ContentType = "text/event-stream"; - ctx.Response.Headers.CacheControl = "no-cache"; - ctx.Response.Headers.Connection = "keep-alive"; + await HandleStreaming(ctx, engine, metrics, adapter, toolsActive, prompt, sp, requestId, created); + } + else + { + await HandleNonStreaming(ctx, engine, metrics, adapter, toolsActive, prompt, sp, requestId, created); + } + } + + private static async Task HandleNonStreaming( + HttpContext ctx, IInferenceEngine engine, ServerMetrics metrics, IToolCallAdapter adapter, + bool toolsActive, string prompt, SamplingParams sp, string requestId, long created) + { + var textSb = new StringBuilder(); + var reasoningSb = new StringBuilder(); + int textTokens = 0; + int reasoningTokens = 0; + + await foreach (var c in engine.GenerateChunksAsync(prompt, sp, ctx.RequestAborted)) + { + if (c.Kind == GenerateChunkKind.Thinking) + { + reasoningSb.Append(c.Text); + reasoningTokens++; + } + else + { + textSb.Append(c.Text); + textTokens++; + } + } + + int completionTokens = textTokens + reasoningTokens; + metrics.RecordTokens(completionTokens); + + var rawText = textSb.ToString(); + IReadOnlyList parsedCalls; + string plainText; + if (toolsActive) + { + (plainText, parsedCalls) = adapter.Parse(rawText); + } + else + { + // No tools declared on this request → don't try to interpret the output as + // structured calls. Surfaces raw text identically to the pre-tools behaviour. + plainText = rawText; + parsedCalls = []; + } - // First chunk: role delta - var firstChunk = new ChatCompletionChunk( + OaiToolCall[]? toolCalls = null; + string finishReason = "stop"; + string? content = plainText; + if (parsedCalls.Count > 0) + { + toolCalls = parsedCalls + .Select(c => new OaiToolCall( + Id: $"call_{Guid.NewGuid():N}", + Type: "function", + Function: new OaiToolCallFunction(c.Name, JinjaChatTemplate.SerializeToJson(c.Arguments)))) + .ToArray(); + finishReason = "tool_calls"; + // OpenAI returns content: null when a tool_calls array is present and there + // was no accompanying text; an empty string would be a wire-shape mismatch. + if (plainText.Length == 0) content = null; + } + + var message = new OaiAssistantMessage( + "assistant", + content, + reasoningSb.Length > 0 ? reasoningSb.ToString() : null, + toolCalls); + var usage = new ChatUsage( + 0, completionTokens, completionTokens, + reasoningTokens > 0 ? new CompletionTokensDetails(reasoningTokens) : null); + + var response = new ChatCompletionResponse( + requestId, "chat.completion", created, engine.ModelId, + [new CompletionChoice(0, message, finishReason)], + usage); + + ctx.Response.ContentType = "application/json"; + await ctx.Response.WriteAsync( + JsonSerializer.Serialize(response, SharpInferenceJsonContext.Default.ChatCompletionResponse), + ctx.RequestAborted); + } + + private static async Task HandleStreaming( + HttpContext ctx, IInferenceEngine engine, ServerMetrics metrics, IToolCallAdapter adapter, + bool toolsActive, string prompt, SamplingParams sp, string requestId, long created) + { + ctx.Response.ContentType = "text/event-stream"; + ctx.Response.Headers.CacheControl = "no-cache"; + ctx.Response.Headers.Connection = "keep-alive"; + + // First chunk: role delta + var firstChunk = new ChatCompletionChunk( + requestId, "chat.completion.chunk", created, engine.ModelId, + [new ChunkChoice(0, new ChunkDelta("assistant", null), null)]); + await WriteEvent(ctx.Response, JsonSerializer.Serialize(firstChunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); + + int maxOpenLen = adapter.MaxOpenTagLength; + bool inToolCall = false; + int toolCallContentStart = -1; + var toolCallBuf = new StringBuilder(); + string pendingText = ""; + int toolCallIndex = 0; // monotonic index into delta.tool_calls + bool hasToolCalls = false; + long tokenCount = 0; + + async Task WriteContentDelta(string text) + { + if (text.Length == 0) return; + var chunk = new ChatCompletionChunk( requestId, "chat.completion.chunk", created, engine.ModelId, - [new ChunkChoice(0, new ChunkDelta("assistant", null), null)]); - await WriteEvent(ctx.Response, JsonSerializer.Serialize(firstChunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); + [new ChunkChoice(0, new ChunkDelta(null, text), null)]); + await WriteEvent(ctx.Response, JsonSerializer.Serialize(chunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); + } - long tokenCount = 0; - await foreach (var c in engine.GenerateChunksAsync(prompt, sp, ctx.RequestAborted)) + async Task EmitToolCallsFromBlock(string blockContent) + { + var calls = adapter.ParseBlock(blockContent); + foreach (var tc in calls) { - tokenCount++; - ChunkDelta delta = c.Kind == GenerateChunkKind.Thinking - ? new ChunkDelta(null, null, c.Text) - : new ChunkDelta(null, c.Text); + hasToolCalls = true; + var delta = new OaiToolCallDelta( + Index: toolCallIndex, + Id: $"call_{Guid.NewGuid():N}", + Type: "function", + Function: new OaiToolCallFunction(tc.Name, JinjaChatTemplate.SerializeToJson(tc.Arguments))); var chunk = new ChatCompletionChunk( requestId, "chat.completion.chunk", created, engine.ModelId, - [new ChunkChoice(0, delta, null)]); + [new ChunkChoice(0, new ChunkDelta(null, null, null, [delta]), null)]); await WriteEvent(ctx.Response, JsonSerializer.Serialize(chunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); + toolCallIndex++; } + } - // Final chunk with finish_reason - var finalChunk = new ChatCompletionChunk( - requestId, "chat.completion.chunk", created, engine.ModelId, - [new ChunkChoice(0, new ChunkDelta(null, null), "stop")]); - await WriteEvent(ctx.Response, JsonSerializer.Serialize(finalChunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); - await ctx.Response.WriteAsync("data: [DONE]\n\n", ctx.RequestAborted); - await ctx.Response.Body.FlushAsync(ctx.RequestAborted); + async Task ProcessTextChunk(string chunk) + { + if (inToolCall) + { + toolCallBuf.Append(chunk); + string buf = toolCallBuf.ToString(); + int closeIdx = adapter.FindCloseMarker(buf, toolCallContentStart, out int afterClose); + if (closeIdx >= 0) + { + string block = buf[toolCallContentStart..closeIdx]; + string remaining = buf[afterClose..]; + toolCallBuf.Clear(); + toolCallContentStart = -1; + inToolCall = false; + await EmitToolCallsFromBlock(block); + if (remaining.Length > 0) + await ProcessTextChunk(remaining); + } + return; + } + + pendingText += chunk; + int openIdx = adapter.FindOpenMarker(pendingText, 0, out int contentStart); + if (openIdx >= 0) + { + if (openIdx > 0) + await WriteContentDelta(pendingText[..openIdx]); + + inToolCall = true; + toolCallBuf.Clear(); + toolCallBuf.Append(pendingText, contentStart, pendingText.Length - contentStart); + toolCallContentStart = 0; + pendingText = ""; + + if (toolCallBuf.Length > 0) + await ProcessTextChunk(""); + return; + } - metrics.RecordTokens(tokenCount); + int safeLen = Math.Max(0, pendingText.Length - (maxOpenLen - 1)); + if (safeLen > 0) + { + await WriteContentDelta(pendingText[..safeLen]); + pendingText = pendingText[safeLen..]; + } } - else - { - var textSb = new StringBuilder(); - var reasoningSb = new StringBuilder(); - int textTokens = 0; - int reasoningTokens = 0; + try + { await foreach (var c in engine.GenerateChunksAsync(prompt, sp, ctx.RequestAborted)) { + tokenCount++; if (c.Kind == GenerateChunkKind.Thinking) { - reasoningSb.Append(c.Text); - reasoningTokens++; + var delta = new ChunkDelta(null, null, c.Text); + var chunk = new ChatCompletionChunk( + requestId, "chat.completion.chunk", created, engine.ModelId, + [new ChunkChoice(0, delta, null)]); + await WriteEvent(ctx.Response, JsonSerializer.Serialize(chunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); + } + else if (toolsActive) + { + await ProcessTextChunk(c.Text); } else { - textSb.Append(c.Text); - textTokens++; + // No tools declared → skip the buffering state machine and forward each + // chunk as a separate content delta (clients rely on the streaming cadence). + await WriteContentDelta(c.Text); } } + } + finally + { + // Flush any remaining buffered text (cannot contain a partial open marker now). + if (!inToolCall && pendingText.Length > 0) + { + try { await WriteContentDelta(pendingText); } catch { /* response aborted */ } + pendingText = ""; + } + } - int completionTokens = textTokens + reasoningTokens; - metrics.RecordTokens(completionTokens); - - var message = new OaiAssistantMessage( - "assistant", - textSb.ToString(), - reasoningSb.Length > 0 ? reasoningSb.ToString() : null); - var usage = new ChatUsage( - 0, completionTokens, completionTokens, - reasoningTokens > 0 ? new CompletionTokensDetails(reasoningTokens) : null); - - var response = new ChatCompletionResponse( - requestId, "chat.completion", created, engine.ModelId, - [new CompletionChoice(0, message, "stop")], - usage); + // Final chunk with finish_reason + var finishReason = hasToolCalls ? "tool_calls" : "stop"; + var finalChunk = new ChatCompletionChunk( + requestId, "chat.completion.chunk", created, engine.ModelId, + [new ChunkChoice(0, new ChunkDelta(null, null), finishReason)]); + await WriteEvent(ctx.Response, JsonSerializer.Serialize(finalChunk, SharpInferenceJsonContext.Default.ChatCompletionChunk)); + await ctx.Response.WriteAsync("data: [DONE]\n\n", ctx.RequestAborted); + await ctx.Response.Body.FlushAsync(ctx.RequestAborted); - ctx.Response.ContentType = "application/json"; - await ctx.Response.WriteAsync( - JsonSerializer.Serialize(response, SharpInferenceJsonContext.Default.ChatCompletionResponse), - ctx.RequestAborted); - } + metrics.RecordTokens(tokenCount); } private static Task HandleListModels(HttpContext ctx, IInferenceEngine engine) @@ -179,6 +349,94 @@ private static Task HandleListModels(HttpContext ctx, IInferenceEngine engine) return list; } + private static bool HasToolMessages(OaiMessage[]? messages) + { + if (messages is null) return false; + foreach (var m in messages) + if (m.Role == "tool" || m.ToolCalls is { Length: > 0 }) + return true; + return false; + } + + /// + /// Builds rich message dictionaries and converts OpenAI tool definitions to the + /// {type:"function", function:{...}} shape the Jinja chat template expects. + /// Mirrors ' BuildRichMessageList but for the + /// OpenAI wire shape: tool definitions live at top level, assistant tool calls + /// arrive as a tool_calls array, and tool results arrive as a separate + /// role:"tool" message with tool_call_id. + /// + private static (List> messages, List? tools) + BuildRichMessageList(ChatCompletionRequest req, IToolCallAdapter adapter) + { + var messages = new List>(); + + foreach (var m in req.Messages!) + { + var role = m.Role ?? "user"; + var content = m.Content ?? ""; + + if (role == "tool") + { + messages.Add(adapter.RenderToolResult(m.ToolCallId ?? "", content)); + continue; + } + + if (role == "assistant") + { + string textStr = ChatTemplate.ScrubAssistantThinking(content); + var msg = new Dictionary(StringComparer.Ordinal) + { + ["role"] = "assistant", + ["content"] = textStr, + }; + if (m.ToolCalls is { Length: > 0 }) + { + var toolCalls = new List(); + foreach (var tc in m.ToolCalls) + { + toolCalls.Add(new Dictionary(StringComparer.Ordinal) + { + ["id"] = tc.Id, + ["type"] = tc.Type, + ["function"] = new Dictionary(StringComparer.Ordinal) + { + ["name"] = tc.Function.Name, + // OpenAI stringifies arguments; pass through as-is. + ["arguments"] = tc.Function.Arguments, + }, + }); + } + msg["tool_calls"] = (object?)toolCalls; + } + messages.Add(msg); + continue; + } + + // system / user / other roles + messages.Add(new(StringComparer.Ordinal) { ["role"] = role, ["content"] = content }); + } + + List? toolsList = null; + if (req.Tools is { Length: > 0 }) + { + toolsList = req.Tools + .Select(t => (object?)new Dictionary(StringComparer.Ordinal) + { + ["type"] = t.Type, + ["function"] = new Dictionary(StringComparer.Ordinal) + { + ["name"] = t.Function.Name, + ["description"] = t.Function.Description, + ["parameters"] = t.Function.Parameters, + }, + }) + .ToList(); + } + + return (messages, toolsList); + } + private static async Task WriteEvent(HttpResponse response, string data) { await response.WriteAsync($"data: {data}\n\n", response.HttpContext.RequestAborted); @@ -198,9 +456,46 @@ public sealed record ChatCompletionRequest( Dictionary? LogitBias, ResponseFormat? ResponseFormat, [property: JsonPropertyName("enable_thinking")] bool? EnableThinking = null, - [property: JsonPropertyName("max_thinking_tokens")] int? MaxThinkingTokens = null); + [property: JsonPropertyName("max_thinking_tokens")] int? MaxThinkingTokens = null, + OaiTool[]? Tools = null, + [property: JsonPropertyName("tool_choice")] JsonElement? ToolChoice = null); + +/// +/// Message in an OpenAI /v1/chat/completions request. Both single-string +/// content and structured fields are supported; an assistant message echoing +/// a prior tool call uses in place of (or alongside) text, +/// and a role: "tool" message carries + the result +/// text in . +/// +public sealed record OaiMessage( + string? Role, + string? Content, + [property: JsonPropertyName("tool_call_id")] string? ToolCallId = null, + [property: JsonPropertyName("tool_calls")] OaiToolCall[]? ToolCalls = null, + string? Name = null); + +/// +/// OpenAI tool definition. function.parameters is the JSON Schema; we keep it +/// as a raw so it round-trips through the Jinja template unchanged. +/// +public sealed record OaiTool(string Type, OaiToolFunction Function); + +public sealed record OaiToolFunction( + string Name, + string? Description, + JsonElement? Parameters); + +/// +/// Single tool-call entry. Emitted on assistant messages (in responses) and accepted +/// on history-side assistant messages (in requests). OpenAI's spec uses a stringly +/// JSON-encoded arguments field. +/// +public sealed record OaiToolCall( + string Id, + string Type, + OaiToolCallFunction Function); -public sealed record OaiMessage(string? Role, string? Content); +public sealed record OaiToolCallFunction(string Name, string Arguments); public sealed record ChatCompletionResponse( string Id, @@ -213,8 +508,9 @@ public sealed record ChatCompletionResponse( public sealed record CompletionChoice(int Index, OaiAssistantMessage Message, string? FinishReason); public sealed record OaiAssistantMessage( string Role, - string Content, - [property: JsonPropertyName("reasoning_content")] string? ReasoningContent = null); + string? Content, + [property: JsonPropertyName("reasoning_content")] string? ReasoningContent = null, + [property: JsonPropertyName("tool_calls")] OaiToolCall[]? ToolCalls = null); public sealed record ChatUsage( int PromptTokens, int CompletionTokens, @@ -234,7 +530,20 @@ public sealed record ChunkChoice(int Index, ChunkDelta Delta, string? FinishReas public sealed record ChunkDelta( string? Role, string? Content, - [property: JsonPropertyName("reasoning_content")] string? ReasoningContent = null); + [property: JsonPropertyName("reasoning_content")] string? ReasoningContent = null, + [property: JsonPropertyName("tool_calls")] OaiToolCallDelta[]? ToolCalls = null); + +/// +/// Per-chunk tool-call delta. OpenAI streams partial tool calls in array index order +/// — clients reconstruct each call by concatenating the function.arguments JSON +/// fragments across deltas sharing the same . We emit the full call +/// in a single delta (matching how the engine surfaces a complete tool_use block). +/// +public sealed record OaiToolCallDelta( + int Index, + string? Id, + string? Type, + OaiToolCallFunction? Function); public sealed record ModelsResponse(string Object, ModelInfo[] Data); public sealed record ModelInfo(string Id, string Object, long Created, string OwnedBy); diff --git a/src/SharpInference.Server/SharpInferenceJsonContext.cs b/src/SharpInference.Server/SharpInferenceJsonContext.cs index 173c2c8..b6806d9 100644 --- a/src/SharpInference.Server/SharpInferenceJsonContext.cs +++ b/src/SharpInference.Server/SharpInferenceJsonContext.cs @@ -16,6 +16,14 @@ namespace SharpInference.Server; DefaultIgnoreCondition = JsonIgnoreCondition.WhenWritingNull)] [JsonSerializable(typeof(ChatCompletionRequest))] [JsonSerializable(typeof(OaiMessage[]))] +[JsonSerializable(typeof(OaiTool))] +[JsonSerializable(typeof(OaiTool[]))] +[JsonSerializable(typeof(OaiToolFunction))] +[JsonSerializable(typeof(OaiToolCall))] +[JsonSerializable(typeof(OaiToolCall[]))] +[JsonSerializable(typeof(OaiToolCallFunction))] +[JsonSerializable(typeof(OaiToolCallDelta))] +[JsonSerializable(typeof(OaiToolCallDelta[]))] [JsonSerializable(typeof(ChatCompletionResponse))] [JsonSerializable(typeof(CompletionChoice[]))] [JsonSerializable(typeof(OaiAssistantMessage))] diff --git a/tests/SharpInference.Tests.Core/ToolCallAdapterTests.cs b/tests/SharpInference.Tests.Core/ToolCallAdapterTests.cs new file mode 100644 index 0000000..e32b131 --- /dev/null +++ b/tests/SharpInference.Tests.Core/ToolCallAdapterTests.cs @@ -0,0 +1,242 @@ +using SharpInference.Core; + +namespace SharpInference.Tests.Core; + +/// +/// Per-adapter unit tests covering each family's wire format against representative +/// model outputs. Streaming tests use the open/close marker API directly so they +/// also lock in the contract the server's streaming state machine depends on. +/// +public sealed class ToolCallAdapterTests +{ + // ── Registry ────────────────────────────────────────────────────────────── + + [Fact] + public void Registry_ResolvesQwenArchitectures() + { + Assert.IsType(ToolCallAdapterRegistry.Get("qwen2")); + Assert.IsType(ToolCallAdapterRegistry.Get("qwen3")); + Assert.IsType(ToolCallAdapterRegistry.Get("qwen3moe")); + Assert.IsType(ToolCallAdapterRegistry.Get("qwen35moe")); + } + + [Fact] + public void Registry_ResolvesQwenCoder() => + Assert.IsType(ToolCallAdapterRegistry.Get("qwen3coder")); + + [Fact] + public void Registry_ResolvesLlamaFamilies() + { + Assert.IsType(ToolCallAdapterRegistry.Get("llama")); + Assert.IsType(ToolCallAdapterRegistry.Get("llama4")); + } + + [Fact] + public void Registry_ResolvesDeepSeek() => + Assert.IsType(ToolCallAdapterRegistry.Get("deepseek2")); + + [Fact] + public void Registry_UnknownArchFallsBackToDefault() + { + var a = ToolCallAdapterRegistry.Get("never-heard-of-it"); + Assert.Same(ToolCallAdapterRegistry.DefaultAdapter, a); + } + + [Fact] + public void Registry_NullOrEmptyArchFallsBackToDefault() + { + Assert.Same(ToolCallAdapterRegistry.DefaultAdapter, ToolCallAdapterRegistry.Get(null)); + Assert.Same(ToolCallAdapterRegistry.DefaultAdapter, ToolCallAdapterRegistry.Get("")); + } + + // ── Qwen (wrapper) adapter ──────────────────────────────────────────────── + + [Fact] + public void Qwen_Parse_JsonCall() + { + var a = new QwenToolCallAdapter("qwen3moe"); + var raw = "\n{\"name\":\"get_weather\",\"arguments\":{\"city\":\"Paris\"}}\n"; + var (plain, calls) = a.Parse(raw); + Assert.Equal("", plain); + Assert.Single(calls); + Assert.Equal("get_weather", calls[0].Name); + Assert.Equal("Paris", calls[0].Arguments["city"]); + } + + [Fact] + public void Qwen_Parse_XmlFunctionCall() + { + // Qwen3.6 alt payload: v + var a = new QwenToolCallAdapter("qwen3moe"); + var raw = "/etc/passwd"; + var (_, calls) = a.Parse(raw); + Assert.Single(calls); + Assert.Equal("read_file", calls[0].Name); + Assert.Equal("/etc/passwd", calls[0].Arguments["path"]); + } + + [Fact] + public void Qwen_Parse_TextBeforeAndAfterToolCall() + { + var a = new QwenToolCallAdapter("qwen3moe"); + var raw = "Let me check.{\"name\":\"x\",\"arguments\":{}} Done."; + var (plain, calls) = a.Parse(raw); + Assert.Equal("Let me check. Done.", plain); + Assert.Single(calls); + } + + [Fact] + public void Qwen_FindMarkers_RoundTripsBlock() + { + var a = new QwenToolCallAdapter("qwen3moe"); + var buf = "noise{\"name\":\"x\",\"arguments\":{\"k\":1}}tail"; + int open = a.FindOpenMarker(buf, 0, out int contentStart); + Assert.Equal(5, open); + Assert.Equal(5 + "".Length, contentStart); + int close = a.FindCloseMarker(buf, contentStart, out int afterClose); + Assert.True(close > contentStart); + Assert.Equal(close + "".Length, afterClose); + + var block = buf[contentStart..close]; + var calls = a.ParseBlock(block); + Assert.Single(calls); + Assert.Equal("x", calls[0].Name); + } + + // ── Qwen3-Coder adapter (closes #95) ────────────────────────────────────── + + [Fact] + public void QwenCoder_Parse_BareFunctionBlock() + { + var a = new QwenCoderToolCallAdapter(); + var raw = "Paris"; + var (plain, calls) = a.Parse(raw); + Assert.Equal("", plain); + Assert.Single(calls); + Assert.Equal("get_weather", calls[0].Name); + Assert.Equal("Paris", calls[0].Arguments["city"]); + } + + [Fact] + public void QwenCoder_Parse_TextBeforeFunctionBlock() + { + var a = new QwenCoderToolCallAdapter(); + var raw = "I'll check.\nx"; + var (plain, calls) = a.Parse(raw); + Assert.Equal("I'll check.\n", plain); + Assert.Single(calls); + } + + [Fact] + public void QwenCoder_Parse_MultipleFunctionBlocks() + { + var a = new QwenCoderToolCallAdapter(); + var raw = "1" + + "2"; + var (_, calls) = a.Parse(raw); + Assert.Equal(2, calls.Count); + Assert.Equal("a", calls[0].Name); + Assert.Equal("b", calls[1].Name); + } + + [Fact] + public void QwenCoder_Streaming_BlockIncludesOpenMarker() + { + // The name is inside the open marker, so the streaming block MUST contain it. + var a = new QwenCoderToolCallAdapter(); + var buf = "/"; + int open = a.FindOpenMarker(buf, 0, out int contentStart); + Assert.Equal(0, open); + Assert.Equal(0, contentStart); // contentStart == openIdx → block keeps the marker + int close = a.FindCloseMarker(buf, contentStart, out int afterClose); + var block = buf[contentStart..close]; + Assert.StartsWith(""; + var (plain, calls) = a.Parse(raw); + Assert.Equal("", plain); + Assert.Single(calls); + Assert.Equal("get_weather", calls[0].Name); + Assert.Equal("Paris", calls[0].Arguments["city"]); + } + + [Fact] + public void Llama_Parse_PythonTagBlock_WithEot() + { + // Some short tool outputs close with <|eot_id|> instead of <|eom_id|>. + var a = new LlamaToolCallAdapter(); + var raw = "<|python_tag|>{\"name\":\"x\",\"parameters\":{}}<|eot_id|>"; + var (_, calls) = a.Parse(raw); + Assert.Single(calls); + Assert.Equal("x", calls[0].Name); + } + + [Fact] + public void Llama_Parse_AcceptsArgumentsKey() + { + // Some fine-tunes use the OpenAI "arguments" key instead of "parameters". + var a = new LlamaToolCallAdapter(); + var raw = "<|python_tag|>{\"name\":\"x\",\"arguments\":{\"k\":1}}<|eom_id|>"; + var (_, calls) = a.Parse(raw); + Assert.Single(calls); + Assert.Equal(1L, calls[0].Arguments["k"]); + } + + [Fact] + public void Llama_RenderToolResult_UsesIpythonRole() + { + var a = new LlamaToolCallAdapter(); + var msg = a.RenderToolResult("call_1", "result text"); + Assert.Equal("ipython", msg["role"]); + Assert.Equal("result text", msg["content"]); + } + + // ── DeepSeek adapter ────────────────────────────────────────────────────── + + [Fact] + public void DeepSeek_Parse_SingleInnerCall() + { + var a = new DeepSeekToolCallAdapter(); + var raw = "<|tool_calls_begin|>" + + "<|tool_call_begin|>get_weather<|tool_sep|>{\"city\":\"Paris\"}<|tool_call_end|>" + + "<|tool_calls_end|>"; + var (_, calls) = a.Parse(raw); + Assert.Single(calls); + Assert.Equal("get_weather", calls[0].Name); + Assert.Equal("Paris", calls[0].Arguments["city"]); + } + + [Fact] + public void DeepSeek_Parse_MultipleInnerCalls() + { + var a = new DeepSeekToolCallAdapter(); + var raw = "<|tool_calls_begin|>" + + "<|tool_call_begin|>a<|tool_sep|>{\"k\":1}<|tool_call_end|>" + + "<|tool_call_begin|>b<|tool_sep|>{\"k\":2}<|tool_call_end|>" + + "<|tool_calls_end|>"; + var (_, calls) = a.Parse(raw); + Assert.Equal(2, calls.Count); + Assert.Equal("a", calls[0].Name); + Assert.Equal("b", calls[1].Name); + } + + [Fact] + public void DeepSeek_Parse_PlainTextStaysPlain() + { + var a = new DeepSeekToolCallAdapter(); + var (plain, calls) = a.Parse("just an answer with no tool call"); + Assert.Equal("just an answer with no tool call", plain); + Assert.Empty(calls); + } +} diff --git a/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs b/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs new file mode 100644 index 0000000..3b32892 --- /dev/null +++ b/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs @@ -0,0 +1,287 @@ +using System.Net; +using System.Net.Http.Json; +using System.Text.Json; +using Microsoft.AspNetCore.Mvc.Testing; +using Microsoft.Extensions.DependencyInjection; +using SharpInference.Engine; +using SharpInference.Server; + +namespace SharpInference.Tests.Server; + +/// +/// End-to-end coverage for the tool-call wire formats wired up by issues #95–#97: +/// Qwen3-Coder's bare <function=> shape on /v1/messages, and the +/// OpenAI /v1/chat/completions tool-call request + response parity. +/// +/// The fake engine emits canned script output regardless of architecture, so we +/// exercise the parser by swapping the configured architecture via +/// . +/// +public sealed class ToolCallEndpointTests +{ + private static HttpClient CreateClient( + FakeInferenceEngine fake, + string architecture = "qwen2") => + new WebApplicationFactory() + .WithWebHostBuilder(b => b.ConfigureServices(s => + { + s.Configure(o => o.Architecture = architecture); + s.AddSingleton(fake); + })) + .CreateClient(); + + // ── /v1/messages with Qwen3-Coder bare-function shape (#95) ──────────────── + + [Fact] + public async Task Anthropic_QwenCoder_NonStreaming_ParsesBareFunctionAsToolUse() + { + var fake = new FakeInferenceEngine("qwen3-coder", [ + (GenerateChunkKind.Text, ""), + (GenerateChunkKind.Text, "Paris"), + (GenerateChunkKind.Text, ""), + ]); + var client = CreateClient(fake, "qwen3coder"); + + var req = new + { + model = "qwen3-coder", + messages = new[] { new { role = "user", content = "Weather?" } }, + max_tokens = 50, + stream = false, + tools = new[] { new + { + name = "get_weather", + description = "Get weather", + input_schema = new { type = "object", properties = new { city = new { type = "string" } } } + } } + }; + var response = await client.PostAsJsonAsync("/v1/messages", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + + var json = await response.Content.ReadAsStringAsync(); + using var doc = JsonDocument.Parse(json); + Assert.Equal("tool_use", doc.RootElement.GetProperty("stop_reason").GetString()); + + var content = doc.RootElement.GetProperty("content"); + Assert.Equal(1, content.GetArrayLength()); + var block = content[0]; + Assert.Equal("tool_use", block.GetProperty("type").GetString()); + Assert.Equal("get_weather", block.GetProperty("name").GetString()); + Assert.Equal("Paris", block.GetProperty("input").GetProperty("city").GetString()); + } + + [Fact] + public async Task Anthropic_QwenCoder_Streaming_EmitsToolUseEvents() + { + var fake = new FakeInferenceEngine("qwen3-coder", [ + (GenerateChunkKind.Text, ""), + (GenerateChunkKind.Text, "ls"), + (GenerateChunkKind.Text, ""), + ]); + var client = CreateClient(fake, "qwen3coder"); + + var req = new + { + model = "qwen3-coder", + messages = new[] { new { role = "user", content = "ls" } }, + max_tokens = 50, + stream = true, + tools = new[] { new { name = "bash", description = "shell", input_schema = new { type = "object" } } } + }; + var response = await client.PostAsJsonAsync("/v1/messages", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + var body = await response.Content.ReadAsStringAsync(); + + Assert.Contains("\"type\":\"tool_use\"", body); + Assert.Contains("\"name\":\"bash\"", body); + Assert.Contains("input_json_delta", body); + Assert.Contains("event: message_stop", body); + // The stop_reason on the terminating message_delta must be tool_use. + Assert.Contains("tool_use", body); + } + + // ── /v1/chat/completions tool-call non-streaming (#97) ──────────────────── + + [Fact] + public async Task OpenAi_WithTools_NonStreaming_EmitsToolCallsArray() + { + var fake = new FakeInferenceEngine("test-model", [ + (GenerateChunkKind.Text, ""), + (GenerateChunkKind.Text, "{\"name\":\"get_weather\",\"arguments\":{\"city\":\"Paris\"}}"), + (GenerateChunkKind.Text, ""), + ]); + var client = CreateClient(fake); + + var req = new + { + model = "test-model", + messages = new[] { new { role = "user", content = "Weather?" } }, + max_tokens = 50, + stream = false, + tools = new[] { new + { + type = "function", + function = new + { + name = "get_weather", + description = "Get weather", + parameters = new { type = "object", properties = new { city = new { type = "string" } } } + } + } } + }; + var response = await client.PostAsJsonAsync("/v1/chat/completions", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + var json = await response.Content.ReadAsStringAsync(); + using var doc = JsonDocument.Parse(json); + + var choice = doc.RootElement.GetProperty("choices")[0]; + Assert.Equal("tool_calls", choice.GetProperty("finish_reason").GetString()); + + var message = choice.GetProperty("message"); + // Content must be null (or omitted) when only tool_calls were produced. + if (message.TryGetProperty("content", out var c)) + Assert.True(c.ValueKind == JsonValueKind.Null, "content must be null when only tool_calls produced"); + + var toolCalls = message.GetProperty("tool_calls"); + Assert.Equal(1, toolCalls.GetArrayLength()); + var call = toolCalls[0]; + Assert.Equal("function", call.GetProperty("type").GetString()); + Assert.True(call.TryGetProperty("id", out _), "tool_call must have id"); + Assert.Equal("get_weather", call.GetProperty("function").GetProperty("name").GetString()); + var argsStr = call.GetProperty("function").GetProperty("arguments").GetString(); + Assert.NotNull(argsStr); + using var argsDoc = JsonDocument.Parse(argsStr!); + Assert.Equal("Paris", argsDoc.RootElement.GetProperty("city").GetString()); + } + + [Fact] + public async Task OpenAi_WithTools_NonStreaming_TextBeforeCall_SurfacesBoth() + { + var fake = new FakeInferenceEngine("test-model", [ + (GenerateChunkKind.Text, "Looking it up. "), + (GenerateChunkKind.Text, "{\"name\":\"x\",\"arguments\":{}}"), + ]); + var client = CreateClient(fake); + + var req = new + { + model = "test-model", + messages = new[] { new { role = "user", content = "go" } }, + max_tokens = 50, + stream = false, + tools = new[] { new { type = "function", function = new { name = "x", description = "", parameters = new { type = "object" } } } } + }; + var response = await client.PostAsJsonAsync("/v1/chat/completions", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + + var json = await response.Content.ReadAsStringAsync(); + using var doc = JsonDocument.Parse(json); + var choice = doc.RootElement.GetProperty("choices")[0]; + Assert.Equal("tool_calls", choice.GetProperty("finish_reason").GetString()); + + var message = choice.GetProperty("message"); + Assert.Equal("Looking it up. ", message.GetProperty("content").GetString()); + Assert.Equal(1, message.GetProperty("tool_calls").GetArrayLength()); + } + + // ── /v1/chat/completions tool-call streaming (#97) ──────────────────────── + + [Fact] + public async Task OpenAi_WithTools_Streaming_EmitsToolCallDelta() + { + var fake = new FakeInferenceEngine("test-model", [ + (GenerateChunkKind.Text, ""), + (GenerateChunkKind.Text, "{\"name\":\"bash\",\"arguments\":{\"cmd\":\"ls\"}}"), + (GenerateChunkKind.Text, ""), + ]); + var client = CreateClient(fake); + + var req = new + { + model = "test-model", + messages = new[] { new { role = "user", content = "ls" } }, + max_tokens = 50, + stream = true, + tools = new[] { new { type = "function", function = new { name = "bash", description = "shell", parameters = new { type = "object" } } } } + }; + var response = await client.PostAsJsonAsync("/v1/chat/completions", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + var body = await response.Content.ReadAsStringAsync(); + + // tool_calls delta carrying name + arguments, finish_reason flips to tool_calls. + Assert.Contains("\"tool_calls\":[", body); + Assert.Contains("\"name\":\"bash\"", body); + Assert.Contains("\"finish_reason\":\"tool_calls\"", body); + Assert.Contains("[DONE]", body); + } + + // ── /v1/chat/completions tool history echo (#97) ────────────────────────── + + [Fact] + public async Task OpenAi_ToolMessageInHistory_IsAccepted() + { + var fake = new FakeInferenceEngine("test-model"); + var client = CreateClient(fake); + + var req = new + { + model = "test-model", + messages = new object[] + { + new { role = "user", content = "Weather?" }, + new + { + role = "assistant", + content = (string?)null, + tool_calls = new[] + { + new + { + id = "call_1", + type = "function", + function = new { name = "get_weather", arguments = "{\"city\":\"Paris\"}" } + } + } + }, + new { role = "tool", tool_call_id = "call_1", content = "Sunny, 22C" }, + }, + max_tokens = 20, + stream = false, + tools = new[] { new { type = "function", function = new { name = "get_weather", description = "weather", parameters = new { type = "object" } } } } + }; + var response = await client.PostAsJsonAsync("/v1/chat/completions", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + var json = await response.Content.ReadAsStringAsync(); + // FakeEngine has no clue about tools — it just echoes; we only want to confirm + // the rich-message + role:"tool" path doesn't bomb out on parse. + Assert.Contains("chat.completion", json); + } + + // ── /v1/chat/completions no-tools path stays unchanged ──────────────────── + + [Fact] + public async Task OpenAi_NoTools_StreamingPreservesPerChunkContentDeltas() + { + // Sanity: streaming without tools must NOT activate the buffering state machine, + // so per-chunk content_deltas continue to arrive separately. + var fake = new FakeInferenceEngine("test-model", [ + (GenerateChunkKind.Text, "Hi"), + (GenerateChunkKind.Text, "!"), + ]); + var client = CreateClient(fake); + + var req = new + { + model = "test-model", + messages = new[] { new { role = "user", content = "Hi" } }, + max_tokens = 10, + stream = true, + }; + var response = await client.PostAsJsonAsync("/v1/chat/completions", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + var body = await response.Content.ReadAsStringAsync(); + Assert.Contains("\"content\":\"Hi\"", body); + Assert.Contains("\"content\":\"!\"", body); + Assert.Contains("\"finish_reason\":\"stop\"", body); + } +} From 322fbb25aab0543f9ffe54b2e71519a9a71c3e14 Mon Sep 17 00:00:00 2001 From: Pekka Heikura Date: Sat, 30 May 2026 11:13:21 +0300 Subject: [PATCH 2/2] =?UTF-8?q?fix:=20PR=20#98=20review=20feedback=20?= =?UTF-8?q?=E2=80=94=20thread-safety,=20NRE=20guards,=20truncated=20tool-c?= =?UTF-8?q?all=20surfacing?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - ToolCallAdapterRegistry: ConcurrentDictionary instead of Dictionary (was racy between request-time Get and startup-time Register). - ParseJsonCallBlock: dispose the inner JsonDocument when arguments is a JSON string (leak under high QPS). - ParseXmlFunctionBlock: skip tags with empty names instead of inserting an empty-string dict key. - OaiToolCall.Function and OaiTool.Function are now nullable to match runtime reality of source-gen JSON; rich-message builder uses ?. so malformed history / tool defs no longer NRE the request. - Anthropic + OpenAI streaming: when the stream ends mid-tool-call, flush the buffered bytes as a text content delta and signal truncation via finish_reason=length (OpenAI) / stop_reason=max_tokens (Anthropic) instead of silently dropping the partial call. - Request-deserialization catches narrowed to JsonException + BadHttpRequestException; OperationCanceledException now propagates. Error body now includes the parse error message instead of being empty. - Streaming finally-block catches narrowed to OperationCanceledException + IOException; non-aborted errors propagate instead of being silently swallowed. - AnthropicEndpoints non-streaming tool input fallback narrowed to JsonException, and reuses the existing EmptyJsonObject instead of parsing "{}" per call. Tests: - New OpenAi_WithTools_Streaming_SplitOpenMarkerAcrossChunks regression guard for the maxOpenLen buffering — delivers across two engine chunks. - Anthropic streaming test now parses the message_delta event and asserts stop_reason==tool_use explicitly rather than substring-matching. Verified end-to-end against a running Qwen3-8B server: malformed JSON now returns a 400 with the parse error in the body, null function on tool defs no longer NREs, OpenAI tool-history echo path produces the expected reply, Anthropic SSE shape unchanged for the non-tool case. Co-Authored-By: Claude Opus 4.7 --- src/SharpInference.Core/ToolCallAdapter.cs | 19 +++++-- .../Endpoints/AnthropicEndpoints.cs | 49 +++++++++++++--- .../Endpoints/OpenAiEndpoints.cs | 54 ++++++++++++++---- .../ToolCallEndpointTests.cs | 57 ++++++++++++++++++- 4 files changed, 153 insertions(+), 26 deletions(-) diff --git a/src/SharpInference.Core/ToolCallAdapter.cs b/src/SharpInference.Core/ToolCallAdapter.cs index d34bb10..64a25f4 100644 --- a/src/SharpInference.Core/ToolCallAdapter.cs +++ b/src/SharpInference.Core/ToolCallAdapter.cs @@ -84,7 +84,7 @@ public interface IToolCallAdapter /// public static class ToolCallAdapterRegistry { - private static readonly Dictionary _adapters = + private static readonly System.Collections.Concurrent.ConcurrentDictionary _adapters = new(StringComparer.Ordinal) { // Wrapper-style: Qwen2, Qwen3, Qwen3.x, Qwen3-MoE variants — all use @@ -180,7 +180,9 @@ internal static class ToolCallParseHelpers int paramEnd = funcBody.IndexOf(paramClose, valueStart, StringComparison.Ordinal); if (paramEnd < 0) break; - args[paramName] = funcBody[valueStart..paramEnd].Trim(); + // Skip nameless tags rather than inserting an empty-string key. + if (paramName.Length > 0) + args[paramName] = funcBody[valueStart..paramEnd].Trim(); p = paramEnd + paramClose.Length; } @@ -209,9 +211,16 @@ internal static class ToolCallParseHelpers if (root.TryGetProperty(argumentsKey, out var argsElem) || root.TryGetProperty(altKey, out argsElem)) { - object? parsed = argsElem.ValueKind == JsonValueKind.String - ? JsonElementToObject(JsonDocument.Parse(argsElem.GetString() ?? "{}").RootElement) - : JsonElementToObject(argsElem); + object? parsed; + if (argsElem.ValueKind == JsonValueKind.String) + { + using var innerDoc = JsonDocument.Parse(argsElem.GetString() ?? "{}"); + parsed = JsonElementToObject(innerDoc.RootElement); + } + else + { + parsed = JsonElementToObject(argsElem); + } if (parsed is Dictionary d) args = d; } diff --git a/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs b/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs index 987de61..130e2ea 100644 --- a/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs +++ b/src/SharpInference.Server/Endpoints/AnthropicEndpoints.cs @@ -30,9 +30,22 @@ private static async Task HandleMessages( { req = await ctx.Request.ReadFromJsonAsync(SharpInferenceJsonContext.Default.AnthropicMessageRequest, ctx.RequestAborted); } - catch + catch (JsonException ex) { ctx.Response.StatusCode = 400; + ctx.Response.ContentType = "application/json"; + await ctx.Response.WriteAsync( + JsonSerializer.Serialize(new AErrorResponse("invalid_request_error", ex.Message), + SharpInferenceJsonContext.Default.AErrorResponse), ctx.RequestAborted); + return; + } + catch (Microsoft.AspNetCore.Http.BadHttpRequestException ex) + { + ctx.Response.StatusCode = 400; + ctx.Response.ContentType = "application/json"; + await ctx.Response.WriteAsync( + JsonSerializer.Serialize(new AErrorResponse("invalid_request_error", ex.Message), + SharpInferenceJsonContext.Default.AErrorResponse), ctx.RequestAborted); return; } @@ -124,8 +137,8 @@ private static async Task HandleNonStreaming( foreach (var (id, name, argsJson) in toolCalls) { JsonElement inputEl; - try { using var d = JsonDocument.Parse(argsJson); inputEl = d.RootElement.Clone(); } - catch { inputEl = JsonDocument.Parse("{}").RootElement.Clone(); } + try { using var d = JsonDocument.Parse(argsJson); inputEl = d.RootElement.Clone(); } + catch (JsonException) { inputEl = EmptyJsonObject; } contentList.Add(new AContent("tool_use", Id: id, Name: name, Input: inputEl)); } @@ -314,6 +327,7 @@ async Task ProcessTextChunk(string chunk) } } + bool truncatedToolCall = false; try { await foreach (var chunk in engine.GenerateChunksAsync(prompt, sp, ctx.RequestAborted)) @@ -349,13 +363,25 @@ await WriteAnthropicEvent(ctx.Response, "content_block_delta", finally { // Flush any remaining buffered text (cannot contain a partial tool_call open tag). - // Discard incomplete tool calls (inToolCall still true on stream end). if (!inToolCall && pendingText.Length > 0) { - try { await FlushTextDelta(pendingText); } catch { /* response already aborted */ } + try { await FlushTextDelta(pendingText); } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } pendingText = ""; } + // Stream ended mid-tool-call: surface the buffered bytes so the client sees the + // truncated output (rather than an empty assistant turn) and signal max_tokens. + if (inToolCall && toolCallBuf.Length > 0) + { + truncatedToolCall = true; + try { await FlushTextDelta(toolCallBuf.ToString()); } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } + toolCallBuf.Clear(); + } + // Close thinking block if it was opened but never got a following text block. if (thinkingOpen && !thinkingClosed) { @@ -369,7 +395,8 @@ await WriteAnthropicEvent(ctx.Response, "content_block_stop", JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", 0), SharpInferenceJsonContext.Default.AContentBlockStopEvent)); } - catch { /* response already aborted */ } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } } if (textOpen) { @@ -379,7 +406,8 @@ await WriteAnthropicEvent(ctx.Response, "content_block_stop", JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", textIndex), SharpInferenceJsonContext.Default.AContentBlockStopEvent)); } - catch { /* response already aborted */ } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } } } @@ -399,11 +427,14 @@ await WriteAnthropicEvent(ctx.Response, "content_block_stop", JsonSerializer.Serialize(new AContentBlockStopEvent("content_block_stop", idx), SharpInferenceJsonContext.Default.AContentBlockStopEvent)); } - catch { /* response already aborted */ } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } } // message_delta - var stopReason = hasToolCalls ? "tool_use" : "end_turn"; + var stopReason = hasToolCalls + ? "tool_use" + : truncatedToolCall ? "max_tokens" : "end_turn"; var msgDelta = new AMessageDeltaEvent("message_delta", new AMessageDelta(stopReason, null), new AUsage(0, outputTokens)); await WriteAnthropicEvent(ctx.Response, "message_delta", diff --git a/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs b/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs index e8113aa..398948c 100644 --- a/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs +++ b/src/SharpInference.Server/Endpoints/OpenAiEndpoints.cs @@ -32,9 +32,22 @@ private static async Task HandleChatCompletion( { req = await ctx.Request.ReadFromJsonAsync(SharpInferenceJsonContext.Default.ChatCompletionRequest, ctx.RequestAborted); } - catch + catch (JsonException ex) { ctx.Response.StatusCode = 400; + ctx.Response.ContentType = "application/json"; + await ctx.Response.WriteAsync( + JsonSerializer.Serialize(new ErrorResponse("invalid_request_error", ex.Message), + SharpInferenceJsonContext.Default.ErrorResponse), ctx.RequestAborted); + return; + } + catch (Microsoft.AspNetCore.Http.BadHttpRequestException ex) + { + ctx.Response.StatusCode = 400; + ctx.Response.ContentType = "application/json"; + await ctx.Response.WriteAsync( + JsonSerializer.Serialize(new ErrorResponse("invalid_request_error", ex.Message), + SharpInferenceJsonContext.Default.ErrorResponse), ctx.RequestAborted); return; } @@ -275,6 +288,7 @@ async Task ProcessTextChunk(string chunk) } } + bool truncatedToolCall = false; try { await foreach (var c in engine.GenerateChunksAsync(prompt, sp, ctx.RequestAborted)) @@ -305,13 +319,29 @@ async Task ProcessTextChunk(string chunk) // Flush any remaining buffered text (cannot contain a partial open marker now). if (!inToolCall && pendingText.Length > 0) { - try { await WriteContentDelta(pendingText); } catch { /* response aborted */ } + try { await WriteContentDelta(pendingText); } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } pendingText = ""; } + + // Stream ended mid-tool-call (model hit max_tokens / EOS before emitting the close + // marker). Surface the buffered bytes so the client sees the truncated output rather + // than getting an empty response, and signal length-based truncation via finish_reason. + if (inToolCall && toolCallBuf.Length > 0) + { + truncatedToolCall = true; + try { await WriteContentDelta(toolCallBuf.ToString()); } + catch (OperationCanceledException) { /* client disconnected */ } + catch (IOException) { /* response stream closed */ } + toolCallBuf.Clear(); + } } // Final chunk with finish_reason - var finishReason = hasToolCalls ? "tool_calls" : "stop"; + var finishReason = hasToolCalls + ? "tool_calls" + : truncatedToolCall ? "length" : "stop"; var finalChunk = new ChatCompletionChunk( requestId, "chat.completion.chunk", created, engine.ModelId, [new ChunkChoice(0, new ChunkDelta(null, null), finishReason)]); @@ -395,15 +425,17 @@ private static (List> messages, List? tools var toolCalls = new List(); foreach (var tc in m.ToolCalls) { + // Source-gen JSON sets non-nullable record properties to null when the + // payload omits them — guard so malformed history doesn't NRE the request. toolCalls.Add(new Dictionary(StringComparer.Ordinal) { ["id"] = tc.Id, ["type"] = tc.Type, ["function"] = new Dictionary(StringComparer.Ordinal) { - ["name"] = tc.Function.Name, + ["name"] = tc.Function?.Name, // OpenAI stringifies arguments; pass through as-is. - ["arguments"] = tc.Function.Arguments, + ["arguments"] = tc.Function?.Arguments, }, }); } @@ -426,9 +458,9 @@ private static (List> messages, List? tools ["type"] = t.Type, ["function"] = new Dictionary(StringComparer.Ordinal) { - ["name"] = t.Function.Name, - ["description"] = t.Function.Description, - ["parameters"] = t.Function.Parameters, + ["name"] = t.Function?.Name, + ["description"] = t.Function?.Description, + ["parameters"] = t.Function?.Parameters, }, }) .ToList(); @@ -477,8 +509,10 @@ public sealed record OaiMessage( /// /// OpenAI tool definition. function.parameters is the JSON Schema; we keep it /// as a raw so it round-trips through the Jinja template unchanged. +/// is nullable because source-gen deserialization sets a missing +/// JSON object to null even on a non-nullable property; the endpoint guards on it. /// -public sealed record OaiTool(string Type, OaiToolFunction Function); +public sealed record OaiTool(string Type, OaiToolFunction? Function); public sealed record OaiToolFunction( string Name, @@ -493,7 +527,7 @@ public sealed record OaiToolFunction( public sealed record OaiToolCall( string Id, string Type, - OaiToolCallFunction Function); + OaiToolCallFunction? Function); public sealed record OaiToolCallFunction(string Name, string Arguments); diff --git a/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs b/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs index 3b32892..4bab708 100644 --- a/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs +++ b/tests/SharpInference.Tests.Server/ToolCallEndpointTests.cs @@ -96,8 +96,61 @@ public async Task Anthropic_QwenCoder_Streaming_EmitsToolUseEvents() Assert.Contains("\"name\":\"bash\"", body); Assert.Contains("input_json_delta", body); Assert.Contains("event: message_stop", body); - // The stop_reason on the terminating message_delta must be tool_use. - Assert.Contains("tool_use", body); + + // The stop_reason on the terminating message_delta MUST be tool_use — parse it out of + // the SSE stream rather than substring-matching, otherwise an earlier tool_use type + // string would mask a regression that emits end_turn at the end. + var stopReason = ExtractMessageDeltaStopReason(body); + Assert.Equal("tool_use", stopReason); + } + + private static string? ExtractMessageDeltaStopReason(string sseBody) + { + // Locate the message_delta event, then parse its data line. + const string evt = "event: message_delta\n"; + int e = sseBody.IndexOf(evt, StringComparison.Ordinal); + if (e < 0) return null; + int dStart = sseBody.IndexOf("data: ", e, StringComparison.Ordinal); + if (dStart < 0) return null; + dStart += "data: ".Length; + int dEnd = sseBody.IndexOf('\n', dStart); + if (dEnd < 0) dEnd = sseBody.Length; + using var doc = JsonDocument.Parse(sseBody[dStart..dEnd]); + return doc.RootElement.GetProperty("delta").GetProperty("stop_reason").GetString(); + } + + [Fact] + public async Task OpenAi_WithTools_Streaming_SplitOpenMarkerAcrossChunks() + { + // Regression guard for the maxOpenLen / pendingText buffering: deliver `` so the open marker arrives split across two engine chunks. Without buffering, + // the first chunk would leak through as content and the tool call would not be detected. + var fake = new FakeInferenceEngine("test-model", [ + (GenerateChunkKind.Text, ""), + (GenerateChunkKind.Text, "{\"name\":\"bash\",\"arguments\":{\"cmd\":\"ls\"}}"), + (GenerateChunkKind.Text, ""), + ]); + var client = CreateClient(fake); + + var req = new + { + model = "test-model", + messages = new[] { new { role = "user", content = "ls" } }, + max_tokens = 50, + stream = true, + tools = new[] { new { type = "function", function = new { name = "bash", description = "shell", parameters = new { type = "object" } } } } + }; + var response = await client.PostAsJsonAsync("/v1/chat/completions", req); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + var body = await response.Content.ReadAsStringAsync(); + + // The first chunk's bytes must NEVER appear as a content delta — they are part of + // the open marker the buffer has to retain. + Assert.DoesNotContain("\"content\":\"