Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 23 additions & 3 deletions internal/memory/extended/recall.go
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,26 @@ func (r *Recall) trackErr(err error) {
// blended ranking with a pure retention ordering.
func (r *Recall) queryAtomsWithPrediction(ctx context.Context, query string, recent []string, state UserState) ([]MemoryAtom, error) {
all := make(map[string]scoredAtomMeta)

// Run intent prediction concurrently with the literal query (which may
// include a paid LLM rerank). Sequentially, the rerank can consume most of
// the recall context deadline, starving the predictor of budget and making
// it fail with context-deadline-exceeded before its LLM call even starts.
predicting := r.predictor != nil && r.cfg.PredictiveIntents > 0 &&
r.cfg.FollowUpAnticipationEnabled != nil && *r.cfg.FollowUpAnticipationEnabled
type predictResult struct {
intents []PredictedIntent
err error
}
var predCh chan predictResult
if predicting {
predCh = make(chan predictResult, 1)
go func() {
intents, err := r.predictor.Predict(ctx, query, recent, state)
predCh <- predictResult{intents, err}
}()
}

literal, err := r.queryAtomsScored(ctx, query, false)
if err != nil {
r.trackErr(err)
Expand All @@ -121,9 +141,9 @@ func (r *Recall) queryAtomsWithPrediction(ctx context.Context, query string, rec
all[s.atom.ID] = s
}

if r.predictor != nil && r.cfg.PredictiveIntents > 0 &&
r.cfg.FollowUpAnticipationEnabled != nil && *r.cfg.FollowUpAnticipationEnabled {
intents, err := r.predictor.Predict(ctx, query, recent, state)
if predicting {
res := <-predCh
intents, err := res.intents, res.err
if err != nil {
r.trackErr(err)
log.Printf("extended memory: predicted-intent generation failed: %v", err)
Expand Down
75 changes: 75 additions & 0 deletions internal/memory/extended/recall_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"fmt"
"strings"
"sync"
"testing"
"time"

Expand Down Expand Up @@ -309,6 +310,80 @@ func TestQueryAtomsWithPredictionCompositeOrdering(t *testing.T) {
}
}

// rerankGateLLM blocks its SimpleCall until the predictor has started,
// simulating a slow rerank that would otherwise consume the recall deadline.
type rerankGateLLM struct {
wait <-chan struct{}
}

func (m *rerankGateLLM) SimpleCall(ctx context.Context, _, _ string) (string, error) {
select {
case <-m.wait:
return "", nil
case <-ctx.Done():
return "", ctx.Err()
}
}

// predictorSignalLLM closes the channel on its first call so the test can
// observe that prediction started while the rerank was still blocked.
type predictorSignalLLM struct {
ch chan struct{}
once sync.Once
resp string
}

func (m *predictorSignalLLM) SimpleCall(_ context.Context, _, _ string) (string, error) {
m.once.Do(func() { close(m.ch) })
return m.resp, nil
}

// TestPredictionRunsConcurrentWithRerank is a regression test for the
// "predictor: context deadline exceeded" failure: the predictor must run
// concurrently with the literal query (including its paid rerank) so a slow
// rerank cannot consume the shared recall deadline before prediction starts.
// Under the old sequential behavior the rerank gate never opens and the
// recall deadlocks, failing this test via the outer timeout.
func TestPredictionRunsConcurrentWithRerank(t *testing.T) {
dir := t.TempDir()
cfg := DefaultConfig()
cfg.Enabled = boolPtr(true)
cfg.SemanticSearchRerank = boolPtr(true)
cfg.SemanticSearchMinScore = 0.01
cfg.PredictiveIntents = 1
cfg.FollowUpAnticipationEnabled = boolPtr(true)
// The mock embedder rates these two English sentences as near-duplicates;
// disable semantic dedup so both atoms reach the rerank path.
cfg.SemanticDedupThreshold = floatPtr(0)

predictorStarted := make(chan struct{})
em := New(dir, &rerankGateLLM{wait: predictorStarted}, cfg)
em.index.newEmb = func() embedding.TextEmbedder { return newMockEmbedder(vectorDim) }
em.index.emb = newMockEmbedder(vectorDim)
em.recall.predictor = NewPredictor(&predictorSignalLLM{
ch: predictorStarted,
resp: `[{"text":"follow-up question","confidence":0.9}]`,
}, cfg)
defer em.Close()

_ = em.AddAtom(context.Background(), makeSearchableAtom("User prefers Go for backend services"))
_ = em.AddAtom(context.Background(), makeSearchableAtom("Run go test ./... to verify"))

done := make(chan error, 1)
go func() {
_, err := em.recall.queryAtomsWithPrediction(context.Background(), "Go backend", nil, UserState{})
done <- err
}()
select {
case err := <-done:
if err != nil {
t.Fatalf("queryAtomsWithPrediction failed: %v", err)
}
case <-time.After(10 * time.Second):
t.Fatal("recall deadlocked: predictor did not run concurrently with the rerank")
}
}

// newIntentWeightedRecallEM builds an enabled ExtendedMemory whose recall
// uses the mock embedder, no rerank, and prediction enabled. The corpus is
// two atoms with (almost) disjoint character sets so the literal query
Expand Down
Loading