mirror of
https://github.com/MindWorkAI/AI-Studio.git
synced 2026-10-10 16:13:47 +00:00
Fixed answers breaking off after an invalid tool call (#1036)
This commit is contained in:
1 parent
b895cf674b
commit
88f71d2405
6 files changed
+202
-32
No files matched your search
@@ -3,6 +3,7 @@ using System.Text.Json;
|
||||
|
||||
using AIStudio.Provider;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
|
||||
using Microsoft.Extensions.Logging.Abstractions;
|
||||
|
||||
@@ -19,6 +20,11 @@ namespace AIStudio.Tests.Provider.ToolCalling;
|
||||
///
|
||||
/// What a round asks for is one tool call at a time, wherever the provider lets it ask: a provider
|
||||
/// which rejects the question fails the whole request, so it is not asked at all.
|
||||
///
|
||||
/// What a round sends back of the model's calls is what the provider accepts. A broken call is
|
||||
/// answered all the same, but with a name and arguments the provider takes, since vLLM rejects the
|
||||
/// whole request over a single call in its history whose name breaks the rule for function names
|
||||
/// or whose arguments are not JSON.
|
||||
/// </remarks>
|
||||
[TestFixture]
|
||||
public sealed class ChatCompletionToolCallingAdapterTests
|
||||
@@ -99,6 +105,63 @@ public sealed class ChatCompletionToolCallingAdapterTests
|
||||
Assert.That(await SentRequest(mayAskForSequentialToolCalls: true, includeTools: false), Does.Not.Contain("parallel_tool_calls"));
|
||||
}
|
||||
|
||||
[Test]
|
||||
public async Task AMalformedFunctionNameIsAnsweredButNeverSentBack()
|
||||
{
|
||||
//
|
||||
// What a model behind vLLM once returned. With this name in the history, vLLM rejected the
|
||||
// next request, because it breaks the rule for function names:
|
||||
//
|
||||
const string MALFORMED_NAME = "1,2,3,4,5,6,7,8,9,10,11,12,13,14,15";
|
||||
|
||||
var (call, nextRequest) = await AnswerOneCall(MALFORMED_NAME, "{}", "The tool call was invalid.");
|
||||
|
||||
Assert.Multiple(() =>
|
||||
{
|
||||
Assert.That(call.IsValid, Is.False, "The call is answered as invalid instead of being run.");
|
||||
Assert.That(nextRequest, Does.Not.Contain(MALFORMED_NAME), "The provider would reject the whole request over this name.");
|
||||
Assert.That(nextRequest, Does.Contain("\"name\":\"invalid_tool_call\""), "The call itself still goes back, or its result would answer a call the provider does not know.");
|
||||
});
|
||||
}
|
||||
|
||||
[Test]
|
||||
public async Task AWellFormedUnknownNameGoesBackUnchanged()
|
||||
{
|
||||
//
|
||||
// A name which merely misses a tool is a name the provider accepts. Such a call stays valid,
|
||||
// so that the executor answers it with which tool is not available: that tells the model
|
||||
// more than an invalid call does.
|
||||
//
|
||||
var (call, nextRequest) = await AnswerOneCall("web_lookup", "{}", "Tool 'web_lookup' is not available.");
|
||||
|
||||
Assert.Multiple(() =>
|
||||
{
|
||||
Assert.That(call.IsValid, Is.True);
|
||||
Assert.That(call.ToolName, Is.EqualTo("web_lookup"));
|
||||
Assert.That(nextRequest, Does.Contain("\"name\":\"web_lookup\""));
|
||||
});
|
||||
}
|
||||
|
||||
[Test]
|
||||
public async Task BrokenArgumentsAreAnsweredButNeverSentBack()
|
||||
{
|
||||
//
|
||||
// Up to v0.28, vLLM parses the arguments of every call in the history, and a single one
|
||||
// which is not JSON fails the whole request:
|
||||
//
|
||||
const string BROKEN_ARGUMENTS = """{"query": "weather in Berl""";
|
||||
|
||||
var (call, nextRequest) = await AnswerOneCall("web_search", BROKEN_ARGUMENTS, "The tool call was invalid.");
|
||||
|
||||
Assert.Multiple(() =>
|
||||
{
|
||||
Assert.That(call.IsValid, Is.False, "The call is answered as invalid instead of being run.");
|
||||
Assert.That(nextRequest, Does.Not.Contain("weather in Berl"), "The provider would reject the whole request over these arguments.");
|
||||
Assert.That(nextRequest, Does.Contain("\"arguments\":\"{}\""), "The call itself still goes back, with arguments the provider can read.");
|
||||
Assert.That(nextRequest, Does.Contain("\"name\":\"web_search\""), "A valid name goes back as it came, even when the arguments next to it are broken.");
|
||||
});
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Runs one round and returns the request it sent, as it goes over the wire.
|
||||
/// </summary>
|
||||
@@ -126,6 +189,50 @@ public sealed class ChatCompletionToolCallingAdapterTests
|
||||
return usages;
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Runs a round in which the model calls one tool, answers the call, and runs the next round.
|
||||
/// </summary>
|
||||
/// <param name="functionName">The name the model writes into its call.</param>
|
||||
/// <param name="arguments">The arguments the model writes into its call.</param>
|
||||
/// <param name="result">The result the call is answered with.</param>
|
||||
/// <returns>The call as the loop sees it, and the request of the next round as it goes over the wire.</returns>
|
||||
private static async Task<(ToolCallingRequestedCall Call, string NextRequest)> AnswerOneCall(string functionName, string arguments, string result)
|
||||
{
|
||||
var toolCallLine = JsonSerializer.Serialize(new
|
||||
{
|
||||
Choices = new[] { new { Index = 0, Delta = new { ToolCalls = new[] { new { Index = 0, Id = "call_1", Type = "function", Function = new { Name = functionName, Arguments = arguments } } } } } },
|
||||
}, ProviderJsonOptions.OPTIONS);
|
||||
|
||||
var sent = new List<ChatCompletionAPIRequest>();
|
||||
var adapter = Adapter(true, sent.Add,
|
||||
[
|
||||
toolCallLine,
|
||||
"[DONE]",
|
||||
],
|
||||
[
|
||||
"""{"choices":[{"index":0,"delta":{"content":"Here is the answer."}}]}""",
|
||||
"[DONE]",
|
||||
]);
|
||||
|
||||
ToolCallingRequestedCall? call = null;
|
||||
await foreach (var streamEvent in adapter.ExecuteRoundAsync(null, true))
|
||||
if (streamEvent.Round is { } round)
|
||||
call = round.Calls.Single();
|
||||
|
||||
Assert.That(call, Is.Not.Null, "The first round has to end with the call.");
|
||||
var answeredCall = call!;
|
||||
|
||||
// What the loop does between two rounds:
|
||||
adapter.RecordAssistantTurn();
|
||||
adapter.RecordToolResult(answeredCall.CallId, result, isError: true);
|
||||
|
||||
await foreach (var _ in adapter.ExecuteRoundAsync(null, true))
|
||||
{
|
||||
}
|
||||
|
||||
return (answeredCall, JsonSerializer.Serialize(sent[^1], ProviderJsonOptions.OPTIONS));
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Builds an adapter whose requests are answered by the given rounds, one after another.
|
||||
/// </summary>
|
||||
|
||||
@@ -155,6 +155,26 @@ public sealed class ToolExecutorTests : ToolRegistryTestBase
|
||||
Assert.That(thread.IsWebAddressGivenToTheModel(new Uri("https://example.org/newsletter/2026-10")), Is.True, "The user may ask in the next message to open the link.");
|
||||
}
|
||||
|
||||
[TestCase("web_search", ExpectedResult = true)]
|
||||
[TestCase("read-web-page_2", ExpectedResult = true)]
|
||||
[TestCase(null, ExpectedResult = false)]
|
||||
[TestCase("", ExpectedResult = false)]
|
||||
[TestCase(" web_search", ExpectedResult = false)]
|
||||
[TestCase("web search", ExpectedResult = false)]
|
||||
[TestCase("1,2,3,4,5,6,7,8,9,10,11,12,13,14,15", ExpectedResult = false)]
|
||||
[TestCase("größe", ExpectedResult = false)]
|
||||
public bool AFunctionNameFollowsTheRuleOfTheProviders(string? functionName) => ToolExecutor.IsValidFunctionName(functionName);
|
||||
|
||||
[Test]
|
||||
public void AFunctionNameHasAtMost64Characters()
|
||||
{
|
||||
Assert.Multiple(() =>
|
||||
{
|
||||
Assert.That(ToolExecutor.IsValidFunctionName(new string('a', 64)), Is.True);
|
||||
Assert.That(ToolExecutor.IsValidFunctionName(new string('a', 65)), Is.False);
|
||||
});
|
||||
}
|
||||
|
||||
private static void AssertDemandsNothing(ToolCallOutcome outcome)
|
||||
{
|
||||
const string REASON = "Nothing reached the model, so there is nothing the chat has to keep.";
|
||||
|
||||
Reference in new issue
Block a user