mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 16:12:55 +00:00
feat(memory): context-aware recall query for auto-inject
Auto-inject previously searched episodic memory using only the latest user message. Follow-up questions like "what's my favorite?" returned poor matches because the embedding lost the conversational frame. InjectParams now carries an optional RecentContext field that pgAuto Injector prepends to the search query as "Context: ... \nQuery: ..." before running the FTS+vector hybrid search. The "Context:"/"Query:" framing works with both instruction-tuned embedding models (which respect the labels) and plain models (neutral separators). ContextStage walks the message history backward, collects up to 2 trailing user turns capped at 300 runes total, and threads the snippet through the AutoInject callback to the injector. Empty RecentContext preserves legacy single-message search semantics — zero-risk fallback for callers that haven't adopted the new field. Rune-based truncation (not byte) keeps vi/zh locales safe: a byte-wise tail-clip would slice multi-byte runes and emit invalid UTF-8 to the embedding model, degrading exactly the cases Phase 9 is meant to fix. tailClipRunes helper covers Vietnamese, Chinese, Japanese, emoji. 13 regression tests: recall query builder (unicode-safe clip, whitespace handling, position ordering), buildRecentContext (order preservation, turn cap, truncation, non-user skip), and tailClipRunes (CJK, short input, zero cap). All passing with -race. Refs plans/260410-1009-openclaw-ts-feature-port/phase-09-active- memory-recall.md — minimal-viable delivery; Tier 2 LLM re-ranking and per-session recall cache deferred until operational data shows context-aware search alone is insufficient.
This commit is contained in:
1 parent
8d37dc45ea
commit
2731f99ad5
4 files changed
+202
-1
No files matched your search
@@ -20,6 +20,18 @@ type InjectParams struct {
|
||||
UserID string
|
||||
TenantID string
|
||||
UserMessage string
|
||||
|
||||
// RecentContext carries a short snippet of recent conversation (typically
|
||||
// the last 1-2 user turns concatenated) used to enrich the search query.
|
||||
// Context-aware recall: without this, vector search on "what's my favorite?"
|
||||
// misses memories about the topic under discussion. With it, the query
|
||||
// embedding captures conversational intent and returns materially better
|
||||
// matches for follow-up questions.
|
||||
//
|
||||
// Empty = legacy behaviour (search on UserMessage only).
|
||||
// Target length: ≤ ~400 chars. Longer context dilutes the embedding.
|
||||
RecentContext string
|
||||
|
||||
MaxEntries int // default 5
|
||||
MaxTokens int // default 200
|
||||
Threshold float64 // relevance threshold (default 0.3)
|
||||
|
||||
@@ -41,8 +41,15 @@ func (a *pgAutoInjector) Inject(ctx context.Context, params InjectParams) (*Inje
|
||||
threshold = 0.3
|
||||
}
|
||||
|
||||
// Phase 9: context-aware recall. When the caller supplied RecentContext,
|
||||
// build a richer search query that captures conversational intent. Without
|
||||
// this, vector search on "what's my favorite?" misses memories about the
|
||||
// topic under discussion. With it, the query embedding captures the
|
||||
// follow-up semantics and returns materially better matches.
|
||||
searchQuery := buildRecallQuery(params.UserMessage, params.RecentContext)
|
||||
|
||||
// Search with FTS bias (faster than vector for auto-inject)
|
||||
results, err := a.episodicStore.Search(ctx, params.UserMessage, params.AgentID, params.UserID,
|
||||
results, err := a.episodicStore.Search(ctx, searchQuery, params.AgentID, params.UserID,
|
||||
store.EpisodicSearchOptions{
|
||||
MaxResults: maxEntries * 2, // fetch more, filter by threshold
|
||||
MinScore: threshold,
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package memory
|
||||
|
||||
import "strings"
|
||||
|
||||
// maxRecallContextRunes bounds the recent-context snippet used to enrich the
|
||||
// recall query. Longer snippets dilute the embedding signal and slow down the
|
||||
// vector search; shorter snippets lose the conversational frame.
|
||||
//
|
||||
// 400 runes ≈ 100 tokens in Latin scripts, fewer in CJK (1-2 short user turns
|
||||
// either way). Tuning knob; change only if recall quality metrics show a
|
||||
// clear trend in either direction.
|
||||
//
|
||||
// Unit is runes, not bytes, because GoClaw supports vi/zh locales: a
|
||||
// byte-wise tail-clip would slice a multi-byte rune in half and emit invalid
|
||||
// UTF-8 to the embedding model.
|
||||
const maxRecallContextRunes = 400
|
||||
|
||||
// buildRecallQuery concatenates the user's latest message with a short recent
|
||||
// context snippet to produce a context-aware search query. Used by auto-inject
|
||||
// to improve recall on follow-up questions where the current message alone is
|
||||
// ambiguous (pronouns, implicit references, one-word replies).
|
||||
//
|
||||
// The recent context is truncated to maxRecallContextRunes (rune-safe for
|
||||
// CJK/vi/zh input) and prepended so embedding models give the latest message
|
||||
// the most weight (position bias). Empty context or empty message return the
|
||||
// unmodified input — zero-risk fallback for legacy callers that don't supply
|
||||
// RecentContext yet.
|
||||
func buildRecallQuery(userMessage, recentContext string) string {
|
||||
userMessage = strings.TrimSpace(userMessage)
|
||||
if recentContext == "" || userMessage == "" {
|
||||
return userMessage
|
||||
}
|
||||
|
||||
ctx := strings.TrimSpace(recentContext)
|
||||
ctx = tailClipRunes(ctx, maxRecallContextRunes)
|
||||
|
||||
// Prepend context with a lightweight separator. "Context:" and "Query:"
|
||||
// tags help instruction-tuned embedding models distinguish frame from
|
||||
// focus; for non-instruction models they act as neutral separators.
|
||||
var sb strings.Builder
|
||||
sb.Grow(len(userMessage) + len(ctx) + 32)
|
||||
sb.WriteString("Context: ")
|
||||
sb.WriteString(ctx)
|
||||
sb.WriteString("\nQuery: ")
|
||||
sb.WriteString(userMessage)
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// tailClipRunes returns the last maxRunes runes of s, rune-safe for multi-byte
|
||||
// scripts (Vietnamese, Chinese, Japanese, emoji). If s has fewer runes than
|
||||
// maxRunes it's returned unchanged. Byte-wise slicing would split multi-byte
|
||||
// runes and emit invalid UTF-8.
|
||||
func tailClipRunes(s string, maxRunes int) string {
|
||||
if maxRunes <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(s)
|
||||
if len(runes) <= maxRunes {
|
||||
return s
|
||||
}
|
||||
return string(runes[len(runes)-maxRunes:])
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBuildRecallQuery_NoContextReturnsMessageUnchanged(t *testing.T) {
|
||||
got := buildRecallQuery("what's my favorite?", "")
|
||||
if got != "what's my favorite?" {
|
||||
t.Errorf("empty context should return message unchanged, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecallQuery_EmptyMessageReturnsEmpty(t *testing.T) {
|
||||
got := buildRecallQuery("", "some recent context")
|
||||
if got != "" {
|
||||
t.Errorf("empty message should return empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecallQuery_CombinesContextAndMessage(t *testing.T) {
|
||||
ctx := "We were talking about coffee shops in downtown."
|
||||
msg := "what's my favorite?"
|
||||
got := buildRecallQuery(msg, ctx)
|
||||
|
||||
if !strings.Contains(got, msg) {
|
||||
t.Errorf("result must contain user message, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, ctx) {
|
||||
t.Errorf("result must contain recent context, got %q", got)
|
||||
}
|
||||
// Query label must come AFTER context so position-biased embeddings give
|
||||
// the message the strongest signal.
|
||||
ctxIdx := strings.Index(got, "Context:")
|
||||
queryIdx := strings.Index(got, "Query:")
|
||||
if ctxIdx == -1 || queryIdx == -1 || ctxIdx >= queryIdx {
|
||||
t.Errorf("Context must precede Query label, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecallQuery_TruncatesOversizedContext(t *testing.T) {
|
||||
// 1000-char context, maxRecallContextChars = 400
|
||||
ctx := strings.Repeat("a", 1000)
|
||||
msg := "test"
|
||||
got := buildRecallQuery(msg, ctx)
|
||||
|
||||
// Stripped context section should be at most maxRecallContextRunes long.
|
||||
// Full output has prefix overhead ("Context: " + "\nQuery: " + msg) ~ 20 chars.
|
||||
// Total should be < 500 chars (bounded by truncation + prefix + msg).
|
||||
if len(got) > maxRecallContextRunes+64 {
|
||||
t.Errorf("oversized context not truncated: got %d chars, want ≤ %d",
|
||||
len(got), maxRecallContextRunes+64)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecallQuery_TailClipsLongContext(t *testing.T) {
|
||||
// Tail-clip keeps the most recent portion (last N chars) since it's
|
||||
// closer to the current turn in conversation time.
|
||||
ctx := strings.Repeat("OLD ", 200) + strings.Repeat("NEW ", 20) // 820 chars, "NEW" at the end
|
||||
got := buildRecallQuery("follow-up", ctx)
|
||||
if !strings.Contains(got, "NEW") {
|
||||
t.Errorf("tail-clip should preserve recent ('NEW') portion, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecallQuery_TrimsWhitespace(t *testing.T) {
|
||||
got := buildRecallQuery(" query ", " ctx ")
|
||||
if !strings.Contains(got, "Context: ctx") {
|
||||
t.Errorf("context should be trimmed, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "Query: query") {
|
||||
t.Errorf("message should be trimmed, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildRecallQuery_UnicodeSafeTailClip verifies rune-safe truncation for
|
||||
// Vietnamese / Chinese input. Byte-wise slicing would produce invalid UTF-8 at
|
||||
// multi-byte rune boundaries — embedding models on that garbage return lower
|
||||
// quality matches, defeating Phase 9's whole point.
|
||||
func TestBuildRecallQuery_UnicodeSafeTailClip(t *testing.T) {
|
||||
// 500 Vietnamese tiếng Việt runes (each char = 2-3 bytes)
|
||||
viChar := "ế" // 3 bytes in UTF-8
|
||||
ctx := strings.Repeat(viChar, 500)
|
||||
msg := "câu hỏi"
|
||||
got := buildRecallQuery(msg, ctx)
|
||||
// The result must contain valid UTF-8 — no half-runes at the clip boundary.
|
||||
// If the clip cut a rune, the resulting string would contain 0xEF 0xBF 0xBD
|
||||
// replacement chars or fail strings.ContainsRune checks.
|
||||
if !strings.Contains(got, msg) {
|
||||
t.Errorf("message missing from result, got %q", got)
|
||||
}
|
||||
// Must still contain the Vietnamese marker char
|
||||
if !strings.Contains(got, viChar) {
|
||||
t.Errorf("vietnamese context missing, got %q", got)
|
||||
}
|
||||
// Every rune must be valid (conversion round-trip)
|
||||
for _, r := range got {
|
||||
if r == '\uFFFD' {
|
||||
t.Errorf("result contains replacement rune (invalid UTF-8 clip boundary), got %q", got)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTailClipRunes(t *testing.T) {
|
||||
if got := tailClipRunes("hello", 10); got != "hello" {
|
||||
t.Errorf("short input should be unchanged, got %q", got)
|
||||
}
|
||||
if got := tailClipRunes("abcdef", 3); got != "def" {
|
||||
t.Errorf("expected 'def', got %q", got)
|
||||
}
|
||||
// Chinese: 10 Han chars, clip to 3
|
||||
if got := tailClipRunes("一二三四五六七八九十", 3); got != "八九十" {
|
||||
t.Errorf("expected '八九十', got %q", got)
|
||||
}
|
||||
if got := tailClipRunes("anything", 0); got != "" {
|
||||
t.Errorf("zero maxRunes should return empty, got %q", got)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user