diff --git a/LocalModelIntegrator.slnx b/LocalModelIntegrator.slnx index d1ca95b..1651161 100644 --- a/LocalModelIntegrator.slnx +++ b/LocalModelIntegrator.slnx @@ -1,5 +1,8 @@ + + + diff --git a/src/LocalModelIntegrator/Services/ChatMessageNormalizer.cs b/src/LocalModelIntegrator/Services/ChatMessageNormalizer.cs new file mode 100644 index 0000000..41a93ce --- /dev/null +++ b/src/LocalModelIntegrator/Services/ChatMessageNormalizer.cs @@ -0,0 +1,45 @@ +using System.Collections.Generic; +using System.Linq; +using LocalModelIntegrator.Models; + +namespace LocalModelIntegrator.Services +{ + /// + /// Send-time payload normalization. Strict chat templates (Qwen among them) reject any + /// request whose system message is not the single first message ("System message must be + /// at the beginning"); permissive ones (Llama) don't care. Coalescing every system + /// message into one leading system message satisfies both, so it is applied to every + /// outgoing request regardless of backend. + /// + public static class ChatMessageNormalizer + { + /// + /// Returns a payload with all system messages merged (in order of appearance, joined + /// by a blank line) into a single system message at index 0. Non-system messages keep + /// their relative order and identity. The input list is never mutated - callers + /// persist it as conversation history. + /// + public static List CoalesceSystemMessages(List messages) + { + if (messages == null) + return null; + + List systemParts = messages + .Where(m => m.Role == "system" && !string.IsNullOrWhiteSpace(m.Content)) + .Select(m => m.Content) + .ToList(); + + // Already canonical (or nothing to do): single non-blank system at index 0, no others. + int systemCount = messages.Count(m => m.Role == "system"); + if (systemCount == 0 || + (systemCount == 1 && systemParts.Count == 1 && messages[0].Role == "system")) + return messages; + + var result = new List(messages.Count); + if (systemParts.Count > 0) + result.Add(new ChatMessage("system", string.Join("\n\n", systemParts))); + result.AddRange(messages.Where(m => m.Role != "system")); + return result; + } + } +} diff --git a/src/LocalModelIntegrator/Services/LLMService.cs b/src/LocalModelIntegrator/Services/LLMService.cs index a0a77f2..ac2d745 100644 --- a/src/LocalModelIntegrator/Services/LLMService.cs +++ b/src/LocalModelIntegrator/Services/LLMService.cs @@ -45,7 +45,7 @@ public async Task CallLLMAsync( string jsonRequest = dialect.BuildRequestJson(new DialectRequest { Model = options.ModelName, - Messages = messages, + Messages = ChatMessageNormalizer.CoalesceSystemMessages(messages), Temperature = options.Temperature, MaxTokens = options.MaxTokens, ContextWindowTokens = options.ContextWindowTokens @@ -122,7 +122,7 @@ public async Task CallLLMStreamingAsync( string jsonRequest = dialect.BuildRequestJson(new DialectRequest { Model = options.ModelName, - Messages = messages, + Messages = ChatMessageNormalizer.CoalesceSystemMessages(messages), Temperature = options.Temperature, MaxTokens = options.MaxTokens, Stream = true, // force streaming on regardless of the dialect's default @@ -262,7 +262,7 @@ public async Task CallLLMWithToolsAsync( string jsonRequest = dialect.BuildRequestJson(new DialectRequest { Model = options.ModelName, - Messages = messages, + Messages = ChatMessageNormalizer.CoalesceSystemMessages(messages), Temperature = options.Temperature, MaxTokens = options.MaxTokens, Stream = false, @@ -334,7 +334,7 @@ public async Task CallLLMWithToolsStreamingAsync( string jsonRequest = dialect.BuildRequestJson(new DialectRequest { Model = options.ModelName, - Messages = messages, + Messages = ChatMessageNormalizer.CoalesceSystemMessages(messages), Temperature = options.Temperature, MaxTokens = options.MaxTokens, Stream = true, @@ -498,7 +498,7 @@ public async Task CompleteCodeAsync( string jsonRequest = dialect.BuildRequestJson(new DialectRequest { Model = options.ModelName, - Messages = messages, + Messages = ChatMessageNormalizer.CoalesceSystemMessages(messages), Temperature = temperature, MaxTokens = maxTokens, Stream = false, diff --git a/tests/LocalModelIntegrator.Tests/ChatMessageNormalizerTests.cs b/tests/LocalModelIntegrator.Tests/ChatMessageNormalizerTests.cs new file mode 100644 index 0000000..1da1303 --- /dev/null +++ b/tests/LocalModelIntegrator.Tests/ChatMessageNormalizerTests.cs @@ -0,0 +1,148 @@ +using System.Collections.Generic; +using System.Linq; +using LocalModelIntegrator.Models; +using LocalModelIntegrator.Services; +using Newtonsoft.Json.Linq; +using Xunit; + +namespace LocalModelIntegrator.Tests +{ + // Strict chat templates (Qwen among them) reject any request whose system message is + // not the single first message ("System message must be at the beginning"). The + // normalizer coalesces every system message into one leading system message. + public class ChatMessageNormalizerTests + { + [Fact] + public void MidConversationSystemMessageIsMergedIntoTheLeadingOne() + { + var messages = new List + { + new ChatMessage("system", "Base prompt."), + new ChatMessage("user", "First question"), + new ChatMessage("assistant", "First answer"), + new ChatMessage("system", "Current Visual Studio context: Foo.cs"), + new ChatMessage("user", "Second question"), + }; + + List result = ChatMessageNormalizer.CoalesceSystemMessages(messages); + + Assert.Equal(4, result.Count); + Assert.Equal("system", result[0].Role); + Assert.Equal("Base prompt.\n\nCurrent Visual Studio context: Foo.cs", result[0].Content); + Assert.Equal(new[] { "user", "assistant", "user" }, result.Skip(1).Select(m => m.Role)); + Assert.Equal(new[] { "First question", "First answer", "Second question" }, + result.Skip(1).Select(m => m.Content)); + } + + [Fact] + public void TwoLeadingSystemMessagesBecomeOne() + { + var messages = new List + { + new ChatMessage("system", "Agent prompt."), + new ChatMessage("system", "Workspace orientation: bar"), + new ChatMessage("user", "Go"), + }; + + List result = ChatMessageNormalizer.CoalesceSystemMessages(messages); + + Assert.Equal(2, result.Count); + Assert.Equal("system", result[0].Role); + Assert.Equal("Agent prompt.\n\nWorkspace orientation: bar", result[0].Content); + Assert.Equal("user", result[1].Role); + } + + [Fact] + public void SingleLeadingSystemMessageIsLeftAlone() + { + var messages = new List + { + new ChatMessage("system", "Prompt"), + new ChatMessage("user", "Hi"), + new ChatMessage("assistant", "Hello"), + }; + + List result = ChatMessageNormalizer.CoalesceSystemMessages(messages); + + Assert.Equal(3, result.Count); + Assert.Equal(new[] { "system", "user", "assistant" }, result.Select(m => m.Role)); + Assert.Equal(new[] { "Prompt", "Hi", "Hello" }, result.Select(m => m.Content)); + } + + [Fact] + public void NoSystemMessagesMeansNoChange() + { + var messages = new List + { + new ChatMessage("user", "Hi"), + new ChatMessage("assistant", "Hello"), + }; + + List result = ChatMessageNormalizer.CoalesceSystemMessages(messages); + + Assert.Equal(new[] { "user", "assistant" }, result.Select(m => m.Role)); + } + + [Fact] + public void InputListAndItsMessagesAreNotMutated() + { + var trailingSystem = new ChatMessage("system", "Context note"); + var messages = new List + { + new ChatMessage("system", "Base"), + new ChatMessage("user", "Q"), + trailingSystem, + }; + + ChatMessageNormalizer.CoalesceSystemMessages(messages); + + // Callers persist this list as conversation history; normalization is send-time only. + Assert.Equal(3, messages.Count); + Assert.Equal("Base", messages[0].Content); + Assert.Equal("Context note", trailingSystem.Content); + Assert.Same(trailingSystem, messages[2]); + } + + [Fact] + public void ToolCallMetadataOnNonSystemMessagesIsPreserved() + { + var assistant = new ChatMessage("assistant", null) + { + ToolCalls = JArray.Parse("[{\"id\":\"call_1\"}]") + }; + var toolResult = new ChatMessage("tool", "result") { ToolCallId = "call_1" }; + var messages = new List + { + new ChatMessage("system", "Base"), + new ChatMessage("user", "Q"), + assistant, + toolResult, + new ChatMessage("system", "Active file: Foo.cs"), + new ChatMessage("user", "Next"), + }; + + List result = ChatMessageNormalizer.CoalesceSystemMessages(messages); + + Assert.Same(assistant, result[2]); + Assert.Same(toolResult, result[3]); + Assert.Equal("call_1", result[3].ToolCallId); + Assert.NotNull(result[2].ToolCalls); + } + + [Fact] + public void BlankSystemMessagesAreDroppedNotJoined() + { + var messages = new List + { + new ChatMessage("system", "Base"), + new ChatMessage("user", "Q"), + new ChatMessage("system", " "), + }; + + List result = ChatMessageNormalizer.CoalesceSystemMessages(messages); + + Assert.Equal(2, result.Count); + Assert.Equal("Base", result[0].Content); + } + } +} diff --git a/tests/LocalModelIntegrator.Tests/LocalModelIntegrator.Tests.csproj b/tests/LocalModelIntegrator.Tests/LocalModelIntegrator.Tests.csproj new file mode 100644 index 0000000..25912bc --- /dev/null +++ b/tests/LocalModelIntegrator.Tests/LocalModelIntegrator.Tests.csproj @@ -0,0 +1,24 @@ + + + net472 + 9.0 + false + LocalModelIntegrator.Tests + + + + + + + + + + + + + + all + + +