chore: merge user message
This commit is contained in:
parent
1d4fa2727c
commit
c241764b7d
2 changed files with 85 additions and 9 deletions
|
|
@ -12,7 +12,6 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
"github.com/sipeed/picoclaw/pkg/skills"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/utils"
|
||||
)
|
||||
|
||||
type ContextBuilder struct {
|
||||
|
|
@ -201,12 +200,23 @@ func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary str
|
|||
// Diegox-17
|
||||
// --- FIN DEL FIX ---
|
||||
|
||||
var coalescedUser *providers.Message
|
||||
if len(history) > 0 && history[len(history)-1].Role == "user" {
|
||||
logger.WarnCF("agent", "Removing trailing user message from history to prevent consecutive user messages",
|
||||
map[string]interface{}{
|
||||
"content_preview": utils.Truncate(history[len(history)-1].Content, 50),
|
||||
})
|
||||
last := history[len(history)-1]
|
||||
history = history[:len(history)-1]
|
||||
merged := last.Content
|
||||
if merged != "" && currentMessage != "" {
|
||||
merged += "\n\n" + currentMessage
|
||||
} else if currentMessage != "" {
|
||||
merged = currentMessage
|
||||
}
|
||||
coalescedUser = &providers.Message{Role: "user", Content: merged}
|
||||
logger.InfoCF("agent", "Coalesced consecutive user messages",
|
||||
map[string]interface{}{
|
||||
"prev_len": len(last.Content),
|
||||
"new_len": len(currentMessage),
|
||||
"merged_len": len(merged),
|
||||
})
|
||||
}
|
||||
|
||||
messages = append(messages, providers.Message{
|
||||
|
|
@ -216,10 +226,11 @@ func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary str
|
|||
|
||||
messages = append(messages, history...)
|
||||
|
||||
messages = append(messages, providers.Message{
|
||||
Role: "user",
|
||||
Content: currentMessage,
|
||||
})
|
||||
if coalescedUser != nil {
|
||||
messages = append(messages, *coalescedUser)
|
||||
} else {
|
||||
messages = append(messages, providers.Message{Role: "user", Content: currentMessage})
|
||||
}
|
||||
|
||||
return messages
|
||||
}
|
||||
|
|
|
|||
65
pkg/agent/context_test.go
Normal file
65
pkg/agent/context_test.go
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
package agent
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// helper to create a temporary workspace for ContextBuilder
|
||||
func withTempWorkspace(t *testing.T, fn func(string)) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
// ensure required subdirs exist if any future logic reads them
|
||||
_ = os.MkdirAll(filepath.Join(dir, "skills"), 0755)
|
||||
fn(dir)
|
||||
}
|
||||
|
||||
func TestBuildMessages_CoalescesConsecutiveUserMessages(t *testing.T) {
|
||||
withTempWorkspace(t, func(ws string) {
|
||||
cb := NewContextBuilder(ws)
|
||||
|
||||
history := []providers.Message{
|
||||
{Role: "system", Content: "sys"},
|
||||
{Role: "assistant", Content: "hi"},
|
||||
{Role: "user", Content: "first"},
|
||||
}
|
||||
|
||||
msgs := cb.BuildMessages(history, "", "second", nil, "cli", "chat1")
|
||||
if len(msgs) < 2 {
|
||||
t.Fatalf("expected at least 2 messages, got %d", len(msgs))
|
||||
}
|
||||
if msgs[0].Role != "system" {
|
||||
t.Fatalf("expected first message to be system, got %s", msgs[0].Role)
|
||||
}
|
||||
// ensure no consecutive user messages and last is a single coalesced user
|
||||
last := msgs[len(msgs)-1]
|
||||
if last.Role != "user" {
|
||||
t.Fatalf("expected last message role=user, got %s", last.Role)
|
||||
}
|
||||
if last.Content != "first\n\nsecond" {
|
||||
t.Fatalf("expected coalesced content 'first\\n\\nsecond', got %q", last.Content)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestBuildMessages_AppendsUserIfNoConsecutive(t *testing.T) {
|
||||
withTempWorkspace(t, func(ws string) {
|
||||
cb := NewContextBuilder(ws)
|
||||
history := []providers.Message{
|
||||
{Role: "assistant", Content: "hi"},
|
||||
}
|
||||
msgs := cb.BuildMessages(history, "", "second", nil, "cli", "chat1")
|
||||
if len(msgs) < 3 {
|
||||
t.Fatalf("expected at least 3 messages (system + history + user), got %d", len(msgs))
|
||||
}
|
||||
if msgs[len(msgs)-1].Role != "user" {
|
||||
t.Fatalf("expected last message role=user, got %s", msgs[len(msgs)-1].Role)
|
||||
}
|
||||
if msgs[len(msgs)-2].Role == "user" {
|
||||
t.Fatalf("did not expect consecutive user messages")
|
||||
}
|
||||
})
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue