mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 14:13:21 +00:00
feat(pipeline): per-provider context window, cache/tenant/pipeline hardening
Ship Group A (bug fixes) + Group B (pipeline enhancements) from plans/260410-1009-openclaw-ts-feature-port — five changes that share cmd/ wiring and pipeline plumbing so they commit as a unit. cache: InMemoryCache now supports periodic sweep + max-size cap via variadic options. PermissionCache wires 60s sweep + 10k entry cap + Close() hook in gateway shutdown so long-running gateways don't leak per-user permission entries. Backward compatible for zero-arg callers. store: ContactCollector.seen cache key now includes tenant_id + channel_ instance so the same sender in different tenants (or different bot instances in the same tenant) no longer silently skip upserts against each other. Zero-tenant (Desktop) behaviour preserved. pipeline: EffectiveContextWindow is resolved once per run in ContextStage via a ResolveContextWindow callback (backed by providers.ModelRegistry) so PruneStage bills history against the actual model window instead of a stale static config. Nil resolver / unknown model fall back to Config.ContextWindow for backward compatibility. Locked to the model observed at context build time to prevent mid-run budget drift. pipeline: PipelineConfig.ReserveTokens carves out an optional safety buffer subtracted from the history budget so compaction fires slightly before the hard limit — protects against provider over-delivery and token counter drift on streaming responses. Zero (default) preserves legacy budget math. agent: ModelRegistry flows gateway → ResolverDeps → LoopConfig → Loop → pipeline adapter so resolver can look up per-model capabilities at run time without re-touching gateway internals. 20 regression tests across cache, contact collector, and pipeline cover the critical paths: cross-tenant isolation, cache sweep + eviction + Close idempotency, per-model window override + fallback, and reserve token buffer behaviour. All passing with -race on both go build ./... and go build -tags sqliteonly ./....
This commit is contained in:
1 parent
77a80680ff
commit
8d37dc45ea
17 files changed
+785
-26
No files matched your search
+4
-2
@@ -284,7 +284,7 @@ func runGateway() {
|
||||
var mcpPool *mcpbridge.Pool
|
||||
var mediaStore *media.Store
|
||||
var postTurn tools.PostTurnProcessor
|
||||
contextFileInterceptor, mcpPool, mediaStore, postTurn = wireExtras(pgStores, agentRouter, providerRegistry, msgBus, pgStores.Sessions, toolsReg, toolPE, skillsLoader, hasMemory, traceCollector, workspace, cfg.Gateway.InjectionAction, cfg, sandboxMgr, redisClient, domainBus)
|
||||
contextFileInterceptor, mcpPool, mediaStore, postTurn = wireExtras(pgStores, agentRouter, providerRegistry, modelReg, msgBus, pgStores.Sessions, toolsReg, toolPE, skillsLoader, hasMemory, traceCollector, workspace, cfg.Gateway.InjectionAction, cfg, sandboxMgr, redisClient, domainBus)
|
||||
if mcpPool != nil {
|
||||
defer mcpPool.Stop()
|
||||
}
|
||||
@@ -511,8 +511,10 @@ func runGateway() {
|
||||
methods.NewTenantsMethods(pgStores.Tenants, msgBus, workspace).Register(server.Router())
|
||||
server.SetTenantsHandler(httpapi.NewTenantsHandler(pgStores.Tenants, msgBus, workspace))
|
||||
server.Router().SetTenantStore(pgStores.Tenants)
|
||||
// Permission cache for tenant membership checks
|
||||
// Permission cache for tenant membership checks. Store on deps so
|
||||
// lifecycle shutdown can call Close() to stop the sweep goroutines.
|
||||
permCache := cache.NewPermissionCache()
|
||||
deps.permCache = permCache
|
||||
msgBus.Subscribe("permission-cache", func(e bus.Event) {
|
||||
if p, ok := e.Payload.(bus.CacheInvalidatePayload); ok {
|
||||
permCache.HandleInvalidation(p)
|
||||
|
||||
@@ -3,6 +3,7 @@ package cmd
|
||||
import (
|
||||
"github.com/nextlevelbuilder/goclaw/internal/agent"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/bus"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/cache"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/channels"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/config"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/gateway"
|
||||
@@ -24,6 +25,7 @@ type gatewayDeps struct {
|
||||
agentRouter *agent.Router
|
||||
toolsReg *tools.Registry
|
||||
skillsLoader *skills.Loader // optional: enables skill creation in evolution approval
|
||||
permCache *cache.PermissionCache // nil if no tenant store; closed on shutdown to stop sweep goroutines
|
||||
workspace string
|
||||
dataDir string
|
||||
}
|
||||
@@ -159,6 +159,11 @@ func (d *gatewayDeps) runLifecycle(
|
||||
// Close provider resources (e.g. Claude CLI temp files)
|
||||
d.providerRegistry.Close()
|
||||
|
||||
// Stop permission cache sweep goroutines so they don't leak past shutdown.
|
||||
if d.permCache != nil {
|
||||
d.permCache.Close()
|
||||
}
|
||||
|
||||
// Stop sandbox pruning + release containers
|
||||
if deps.sandboxMgr != nil {
|
||||
deps.sandboxMgr.Stop()
|
||||
|
||||
@@ -42,6 +42,7 @@ func wireExtras(
|
||||
stores *store.Stores,
|
||||
agentRouter *agent.Router,
|
||||
providerReg *providers.Registry,
|
||||
modelReg providers.ModelRegistry,
|
||||
msgBus *bus.MessageBus,
|
||||
sessStore store.SessionStore,
|
||||
toolsReg *tools.Registry,
|
||||
@@ -144,6 +145,7 @@ func wireExtras(
|
||||
AgentStore: stores.Agents,
|
||||
ProviderStore: stores.Providers,
|
||||
ProviderReg: providerReg,
|
||||
ModelRegistry: modelReg,
|
||||
Bus: msgBus,
|
||||
Sessions: sessStore,
|
||||
Tools: toolsReg,
|
||||
|
||||
@@ -59,6 +59,19 @@ func (l *Loop) buildPipelineDeps(req *RunRequest, bridgeRS *runState) pipeline.P
|
||||
Compaction: l.compactionCfg,
|
||||
// V3 memory/retrieval flags removed — always true at runtime.
|
||||
},
|
||||
// Resolve per-model context window once per run. Falls back to
|
||||
// Config.ContextWindow when registry/model is unknown (existing
|
||||
// behaviour unchanged for tests and lite edition).
|
||||
ResolveContextWindow: func(provider, model string) int {
|
||||
if l.modelRegistry == nil || model == "" {
|
||||
return 0
|
||||
}
|
||||
spec := l.modelRegistry.Resolve(provider, model)
|
||||
if spec == nil {
|
||||
return 0
|
||||
}
|
||||
return spec.ContextWindow
|
||||
},
|
||||
EmitEvent: func(event any) {
|
||||
if ae, ok := event.(AgentEvent); ok {
|
||||
l.emit(ae)
|
||||
@@ -233,16 +246,19 @@ func convertRunResult(pr *pipeline.RunResult) *RunResult {
|
||||
|
||||
// makeAutoInjectCallback creates the AutoInject callback that captures agent/tenant context.
|
||||
// Returns nil if autoInjector is not configured (v3 retrieval disabled or no episodic store).
|
||||
func (l *Loop) makeAutoInjectCallback(req *RunRequest) func(ctx context.Context, userMessage, userID string) (string, error) {
|
||||
// Phase 9: plumbs recentContext through to enrich vector search queries for
|
||||
// context-aware recall.
|
||||
func (l *Loop) makeAutoInjectCallback(req *RunRequest) func(ctx context.Context, userMessage, userID, recentContext string) (string, error) {
|
||||
if l.autoInjector == nil {
|
||||
return nil
|
||||
}
|
||||
return func(ctx context.Context, userMessage, userID string) (string, error) {
|
||||
return func(ctx context.Context, userMessage, userID, recentContext string) (string, error) {
|
||||
result, err := l.autoInjector.Inject(ctx, memory.InjectParams{
|
||||
AgentID: l.agentUUID.String(),
|
||||
UserID: userID,
|
||||
TenantID: store.TenantIDFromContext(ctx).String(),
|
||||
UserMessage: userMessage,
|
||||
AgentID: l.agentUUID.String(),
|
||||
UserID: userID,
|
||||
TenantID: store.TenantIDFromContext(ctx).String(),
|
||||
UserMessage: userMessage,
|
||||
RecentContext: recentContext,
|
||||
})
|
||||
if err != nil || result == nil {
|
||||
return "", err
|
||||
|
||||
@@ -76,6 +76,7 @@ type Loop struct {
|
||||
defaultTimezone string // system default timezone for bootstrap pre-fill
|
||||
provider providers.Provider
|
||||
model string
|
||||
modelRegistry providers.ModelRegistry // resolves per-model context window at run time (nil = use static contextWindow)
|
||||
contextWindow int
|
||||
maxTokens int // max output tokens per LLM call (0 = default 8192)
|
||||
maxIterations int
|
||||
@@ -258,6 +259,10 @@ type LoopConfig struct {
|
||||
MemoryCfg *config.MemoryConfig
|
||||
SandboxCfg *sandbox.Config
|
||||
|
||||
// ModelRegistry resolves provider/model → ModelSpec for per-run context
|
||||
// window lookup. Nil = fall back to static LoopConfig.ContextWindow.
|
||||
ModelRegistry providers.ModelRegistry
|
||||
|
||||
Bus bus.EventPublisher
|
||||
DomainBus eventbus.DomainEventBus // V3 domain event bus for consolidation pipeline
|
||||
Sessions store.SessionStore
|
||||
@@ -412,6 +417,7 @@ func NewLoop(cfg LoopConfig) *Loop {
|
||||
agentType: cfg.AgentType,
|
||||
provider: cfg.Provider,
|
||||
model: cfg.Model,
|
||||
modelRegistry: cfg.ModelRegistry,
|
||||
contextWindow: cfg.ContextWindow,
|
||||
maxTokens: cfg.MaxTokens,
|
||||
maxIterations: cfg.MaxIterations,
|
||||
|
||||
@@ -30,6 +30,7 @@ type ResolverDeps struct {
|
||||
AgentStore store.AgentStore
|
||||
ProviderStore store.ProviderStore
|
||||
ProviderReg *providers.Registry
|
||||
ModelRegistry providers.ModelRegistry // per-model context window + capabilities lookup
|
||||
Bus bus.EventPublisher
|
||||
Sessions store.SessionStore
|
||||
Tools *tools.Registry
|
||||
@@ -412,6 +413,7 @@ func NewManagedResolver(deps ResolverDeps) ResolverFunc {
|
||||
AutoInjector: deps.AutoInjector,
|
||||
Provider: provider,
|
||||
Model: ag.Model,
|
||||
ModelRegistry: deps.ModelRegistry,
|
||||
ContextWindow: contextWindow,
|
||||
MaxTokens: ag.ParseMaxTokens(),
|
||||
MaxIterations: maxIter,
|
||||
|
||||
Vendored
+127
-7
@@ -2,29 +2,63 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// entry wraps a cached value with expiration metadata.
|
||||
// entry wraps a cached value with expiration and creation metadata.
|
||||
// createdAt is used for oldest-first eviction when maxSize is exceeded.
|
||||
type entry[V any] struct {
|
||||
value V
|
||||
expiresAt time.Time // zero means no expiry
|
||||
createdAt time.Time // set on Set(), used for eviction ordering
|
||||
}
|
||||
|
||||
func (e entry[V]) expired() bool {
|
||||
return !e.expiresAt.IsZero() && time.Now().After(e.expiresAt)
|
||||
}
|
||||
|
||||
// InMemoryCache is a thread-safe in-memory Cache implementation with TTL support.
|
||||
// InMemoryCache is a thread-safe in-memory Cache implementation with TTL support,
|
||||
// periodic sweep goroutine for expired entries, and optional max-size cap with
|
||||
// oldest-first eviction.
|
||||
type InMemoryCache[V any] struct {
|
||||
data sync.Map
|
||||
data sync.Map
|
||||
maxSize int // 0 = unlimited
|
||||
sweepInterval time.Duration // 0 = no periodic sweep (lazy eviction only)
|
||||
cancel context.CancelFunc
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
// NewInMemoryCache creates a new in-memory cache.
|
||||
func NewInMemoryCache[V any]() *InMemoryCache[V] {
|
||||
return &InMemoryCache[V]{}
|
||||
// CacheOption configures an InMemoryCache during construction.
|
||||
type CacheOption[V any] func(*InMemoryCache[V])
|
||||
|
||||
// WithMaxSize sets a maximum entry count. When exceeded during sweep, the
|
||||
// oldest 20% of entries are evicted. Zero = unlimited.
|
||||
func WithMaxSize[V any](n int) CacheOption[V] {
|
||||
return func(c *InMemoryCache[V]) { c.maxSize = n }
|
||||
}
|
||||
|
||||
// WithSweepInterval sets the periodic sweep interval for expired entries.
|
||||
// Zero disables the sweep goroutine (lazy eviction on Get only).
|
||||
func WithSweepInterval[V any](d time.Duration) CacheOption[V] {
|
||||
return func(c *InMemoryCache[V]) { c.sweepInterval = d }
|
||||
}
|
||||
|
||||
// NewInMemoryCache creates a new in-memory cache. Without options it behaves
|
||||
// exactly as before (lazy eviction, no size cap, no sweep goroutine).
|
||||
func NewInMemoryCache[V any](opts ...CacheOption[V]) *InMemoryCache[V] {
|
||||
c := &InMemoryCache[V]{}
|
||||
for _, opt := range opts {
|
||||
opt(c)
|
||||
}
|
||||
if c.sweepInterval > 0 {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
c.cancel = cancel
|
||||
go c.sweepLoop(ctx)
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *InMemoryCache[V]) Get(_ context.Context, key string) (V, bool) {
|
||||
@@ -43,7 +77,7 @@ func (c *InMemoryCache[V]) Get(_ context.Context, key string) (V, bool) {
|
||||
}
|
||||
|
||||
func (c *InMemoryCache[V]) Set(_ context.Context, key string, value V, ttl time.Duration) {
|
||||
e := entry[V]{value: value}
|
||||
e := entry[V]{value: value, createdAt: time.Now()}
|
||||
if ttl > 0 {
|
||||
e.expiresAt = time.Now().Add(ttl)
|
||||
}
|
||||
@@ -69,3 +103,89 @@ func (c *InMemoryCache[V]) Clear(_ context.Context) {
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
// Close stops the background sweep goroutine. Safe to call multiple times.
|
||||
// After Close, the cache remains readable/writable but without periodic sweep.
|
||||
func (c *InMemoryCache[V]) Close() {
|
||||
c.closeOnce.Do(func() {
|
||||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// sweepLoop runs the periodic expiry + size-cap cleanup in a background goroutine.
|
||||
func (c *InMemoryCache[V]) sweepLoop(ctx context.Context) {
|
||||
ticker := time.NewTicker(c.sweepInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
c.sweepOnce()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// sweepOnce deletes expired entries and, if maxSize is set and exceeded,
|
||||
// evicts the oldest 20% of remaining entries (by createdAt).
|
||||
func (c *InMemoryCache[V]) sweepOnce() {
|
||||
// Phase 1: collect expired keys (don't delete during Range iteration).
|
||||
// sync.Map's Range is safe with concurrent Store/Delete, but deleting
|
||||
// during iteration can cause the same key to be visited twice in
|
||||
// pathological cases. Collecting first is cleaner.
|
||||
var expiredKeys []string
|
||||
type keyAge struct {
|
||||
key string
|
||||
createdAt time.Time
|
||||
}
|
||||
var allAlive []keyAge
|
||||
|
||||
c.data.Range(func(k, v any) bool {
|
||||
key, ok := k.(string)
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
e, ok := v.(entry[V])
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
if e.expired() {
|
||||
expiredKeys = append(expiredKeys, key)
|
||||
} else {
|
||||
allAlive = append(allAlive, keyAge{key: key, createdAt: e.createdAt})
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
for _, k := range expiredKeys {
|
||||
c.data.Delete(k)
|
||||
}
|
||||
|
||||
// Phase 2: size-cap eviction. Sort alive entries by createdAt ascending,
|
||||
// evict oldest 20% if over maxSize.
|
||||
if c.maxSize > 0 && len(allAlive) > c.maxSize {
|
||||
sort.Slice(allAlive, func(i, j int) bool {
|
||||
return allAlive[i].createdAt.Before(allAlive[j].createdAt)
|
||||
})
|
||||
toEvict := len(allAlive) - c.maxSize + (c.maxSize / 5) // bring below cap + 20% headroom
|
||||
if toEvict > len(allAlive) {
|
||||
toEvict = len(allAlive)
|
||||
}
|
||||
for i := 0; i < toEvict; i++ {
|
||||
c.data.Delete(allAlive[i].key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// sizeLocked returns the current entry count (for tests and metrics).
|
||||
// Named sizeLocked for historical reasons; sync.Map needs no lock.
|
||||
func (c *InMemoryCache[V]) sizeLocked() int {
|
||||
n := 0
|
||||
c.data.Range(func(_, _ any) bool {
|
||||
n++
|
||||
return true
|
||||
})
|
||||
return n
|
||||
}
|
||||
Vendored
+104
@@ -90,3 +90,107 @@ func TestInMemoryCache_Clear(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestInMemoryCache_PeriodicSweep verifies expired entries are removed by
|
||||
// the background sweep goroutine (not just lazy on Get).
|
||||
func TestInMemoryCache_PeriodicSweep(t *testing.T) {
|
||||
c := NewInMemoryCache[string](
|
||||
WithSweepInterval[string](20*time.Millisecond),
|
||||
)
|
||||
defer c.Close()
|
||||
ctx := context.Background()
|
||||
|
||||
c.Set(ctx, "short", "v1", 10*time.Millisecond)
|
||||
c.Set(ctx, "long", "v2", 0) // no expiry
|
||||
|
||||
// Wait for sweep to run at least once after short entry expires
|
||||
time.Sleep(60 * time.Millisecond)
|
||||
|
||||
// Short entry should be removed by sweep even without Get
|
||||
if count := c.sizeLocked(); count != 1 {
|
||||
t.Errorf("expected 1 entry after sweep, got %d", count)
|
||||
}
|
||||
if _, ok := c.Get(ctx, "long"); !ok {
|
||||
t.Error("long-lived entry should still exist")
|
||||
}
|
||||
}
|
||||
|
||||
// TestInMemoryCache_MaxSizeEviction verifies oldest entries are evicted when
|
||||
// max size cap is reached.
|
||||
func TestInMemoryCache_MaxSizeEviction(t *testing.T) {
|
||||
c := NewInMemoryCache[int](
|
||||
WithMaxSize[int](5),
|
||||
WithSweepInterval[int](10*time.Millisecond),
|
||||
)
|
||||
defer c.Close()
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert 10 entries with distinct creation times to ensure oldest-first ordering
|
||||
for i := 0; i < 10; i++ {
|
||||
c.Set(ctx, string(rune('a'+i)), i, 0)
|
||||
time.Sleep(2 * time.Millisecond)
|
||||
}
|
||||
|
||||
// Trigger sweep by waiting for interval
|
||||
time.Sleep(30 * time.Millisecond)
|
||||
|
||||
if count := c.sizeLocked(); count > 5 {
|
||||
t.Errorf("expected size ≤ 5 after max-size eviction, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
// TestInMemoryCache_Close verifies Close stops the sweep goroutine (no leak).
|
||||
func TestInMemoryCache_Close(t *testing.T) {
|
||||
c := NewInMemoryCache[string](
|
||||
WithSweepInterval[string](10*time.Millisecond),
|
||||
)
|
||||
ctx := context.Background()
|
||||
c.Set(ctx, "k", "v", 0)
|
||||
|
||||
c.Close()
|
||||
// Close should be idempotent
|
||||
c.Close()
|
||||
|
||||
// After Close, cache is still readable for lazy access but sweep is stopped
|
||||
if _, ok := c.Get(ctx, "k"); !ok {
|
||||
t.Error("Get should still work after Close")
|
||||
}
|
||||
}
|
||||
|
||||
// TestInMemoryCache_ConcurrentSweepAndSet verifies no race between sweep and Set.
|
||||
func TestInMemoryCache_ConcurrentSweepAndSet(t *testing.T) {
|
||||
c := NewInMemoryCache[int](
|
||||
WithSweepInterval[int](1*time.Millisecond),
|
||||
WithMaxSize[int](100),
|
||||
)
|
||||
defer c.Close()
|
||||
ctx := context.Background()
|
||||
|
||||
done := make(chan bool)
|
||||
go func() {
|
||||
for i := 0; i < 500; i++ {
|
||||
c.Set(ctx, string(rune('a'+(i%26))), i, 5*time.Millisecond)
|
||||
}
|
||||
done <- true
|
||||
}()
|
||||
go func() {
|
||||
for i := 0; i < 500; i++ {
|
||||
_, _ = c.Get(ctx, string(rune('a'+(i%26))))
|
||||
}
|
||||
done <- true
|
||||
}()
|
||||
<-done
|
||||
<-done
|
||||
}
|
||||
|
||||
// TestInMemoryCache_BackwardCompatZeroArg verifies existing call sites with
|
||||
// zero-arg constructor still work (variadic options).
|
||||
func TestInMemoryCache_BackwardCompatZeroArg(t *testing.T) {
|
||||
c := NewInMemoryCache[string]()
|
||||
defer c.Close()
|
||||
ctx := context.Background()
|
||||
c.Set(ctx, "k", "v", 0)
|
||||
if v, ok := c.Get(ctx, "k"); !ok || v != "v" {
|
||||
t.Errorf("backward compat broken: got %v, %v", v, ok)
|
||||
}
|
||||
}
|
||||
Vendored
+32
-4
@@ -23,15 +23,43 @@ type PermissionCache struct {
|
||||
teamAccess *InMemoryCache[bool]
|
||||
}
|
||||
|
||||
// NewPermissionCache creates a new permission cache.
|
||||
// permissionCacheSweepInterval and permissionCacheMaxSize bound background
|
||||
// growth of per-user cache entries. Without these, long-running gateways with
|
||||
// many distinct users would accumulate unbounded entries (tenant_roles, agent
|
||||
// access, team access) even with a 30s TTL — lazy eviction only fires on Get,
|
||||
// so entries for disconnected users never get reclaimed.
|
||||
const (
|
||||
permissionCacheSweepInterval = 60 * time.Second
|
||||
permissionCacheMaxSize = 10_000
|
||||
)
|
||||
|
||||
// NewPermissionCache creates a new permission cache with periodic sweep
|
||||
// goroutines for all three inner caches. Call Close() on gateway shutdown to
|
||||
// stop the sweep goroutines.
|
||||
func NewPermissionCache() *PermissionCache {
|
||||
return &PermissionCache{
|
||||
tenantRole: NewInMemoryCache[string](),
|
||||
agentAccess: NewInMemoryCache[agentAccessEntry](),
|
||||
teamAccess: NewInMemoryCache[bool](),
|
||||
tenantRole: NewInMemoryCache[string](
|
||||
WithSweepInterval[string](permissionCacheSweepInterval),
|
||||
WithMaxSize[string](permissionCacheMaxSize),
|
||||
),
|
||||
agentAccess: NewInMemoryCache[agentAccessEntry](
|
||||
WithSweepInterval[agentAccessEntry](permissionCacheSweepInterval),
|
||||
WithMaxSize[agentAccessEntry](permissionCacheMaxSize),
|
||||
),
|
||||
teamAccess: NewInMemoryCache[bool](
|
||||
WithSweepInterval[bool](permissionCacheSweepInterval),
|
||||
WithMaxSize[bool](permissionCacheMaxSize),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
// Close stops all background sweep goroutines. Safe to call multiple times.
|
||||
func (pc *PermissionCache) Close() {
|
||||
pc.tenantRole.Close()
|
||||
pc.agentAccess.Close()
|
||||
pc.teamAccess.Close()
|
||||
}
|
||||
|
||||
const (
|
||||
tenantRoleTTL = 30 * time.Second
|
||||
agentAccessTTL = 30 * time.Second
|
||||
|
||||
@@ -3,6 +3,7 @@ package pipeline
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/bootstrap"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
@@ -34,6 +35,21 @@ func (s *ContextStage) Execute(ctx context.Context, state *RunState) error {
|
||||
state.Ctx = ctx
|
||||
}
|
||||
|
||||
// 0.5. Resolve the effective context window for this run's provider/model.
|
||||
// Done once here so PruneStage reads a stable value on every iteration and
|
||||
// the budget can't drift if the model somehow changes mid-run. A zero
|
||||
// result from the resolver (unknown model, no registry) leaves the field
|
||||
// zero — PruneStage then falls back to Config.ContextWindow.
|
||||
if s.deps.ResolveContextWindow != nil && state.Model != "" {
|
||||
providerID := ""
|
||||
if state.Provider != nil {
|
||||
providerID = state.Provider.Name()
|
||||
}
|
||||
if cw := s.deps.ResolveContextWindow(providerID, state.Model); cw > 0 {
|
||||
state.Context.EffectiveContextWindow = cw
|
||||
}
|
||||
}
|
||||
|
||||
// 1. Resolve workspace
|
||||
if s.deps.ResolveWorkspace != nil {
|
||||
ws, err := s.deps.ResolveWorkspace(ctx, state.Input)
|
||||
@@ -96,8 +112,11 @@ func (s *ContextStage) Execute(ctx context.Context, state *RunState) error {
|
||||
|
||||
// 8. Auto-inject L0 memory context into system prompt.
|
||||
// V3RetrievalEnabled check removed — auto-inject runs whenever AutoInject is available.
|
||||
// Phase 9: pass recent conversation context so vector search can resolve
|
||||
// pronouns and implicit references in follow-up questions.
|
||||
if s.deps.AutoInject != nil && state.Input.Message != "" {
|
||||
section, err := s.deps.AutoInject(ctx, state.Input.Message, state.Input.UserID)
|
||||
recentCtx := buildRecentContext(state.Messages.History())
|
||||
section, err := s.deps.AutoInject(ctx, state.Input.Message, state.Input.UserID, recentCtx)
|
||||
if err == nil && section != "" {
|
||||
state.Context.MemorySection = section
|
||||
sys := state.Messages.System()
|
||||
@@ -109,6 +128,50 @@ func (s *ContextStage) Execute(ctx context.Context, state *RunState) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildRecentContext concatenates the trailing user turns from the history
|
||||
// buffer into a short snippet suitable for enriching a recall query. Walks
|
||||
// backward from the end so we keep the most recent turns regardless of
|
||||
// earlier messages being pruned. Returns "" when there's no usable context.
|
||||
//
|
||||
// Budget: up to 2 user turns, max ~300 runes total. Rune (not byte) cap keeps
|
||||
// vi/zh locales safe — a byte-wise clip would slice multi-byte characters
|
||||
// and emit invalid UTF-8 to the embedding model. Tuning knob is intentional
|
||||
// here rather than config-driven — Phase 9 adds it only if operational data
|
||||
// shows variance across agent types.
|
||||
func buildRecentContext(history []providers.Message) string {
|
||||
const maxTurns = 2
|
||||
const maxRunes = 300
|
||||
if len(history) == 0 {
|
||||
return ""
|
||||
}
|
||||
turns := make([]string, 0, maxTurns)
|
||||
for i := len(history) - 1; i >= 0 && len(turns) < maxTurns; i-- {
|
||||
m := history[i]
|
||||
if m.Role != "user" || m.Content == "" {
|
||||
continue
|
||||
}
|
||||
turns = append([]string{m.Content}, turns...) // prepend to preserve order
|
||||
}
|
||||
if len(turns) == 0 {
|
||||
return ""
|
||||
}
|
||||
var sb strings.Builder
|
||||
for i, t := range turns {
|
||||
if i > 0 {
|
||||
sb.WriteString(" | ")
|
||||
}
|
||||
sb.WriteString(t)
|
||||
}
|
||||
joined := sb.String()
|
||||
// Rune-safe tail clip: keep the most recent portion of the conversation
|
||||
// (closest in time to the current turn) rather than the oldest.
|
||||
runes := []rune(joined)
|
||||
if len(runes) > maxRunes {
|
||||
joined = string(runes[len(runes)-maxRunes:])
|
||||
}
|
||||
return joined
|
||||
}
|
||||
|
||||
// toAnySlice converts []bootstrap.ContextFile to []any for ContextState.ContextFiles.
|
||||
// Phase 8 will remove this when ContextState uses typed field.
|
||||
func toAnySlice(files []bootstrap.ContextFile) []any {
|
||||
|
||||
@@ -18,12 +18,20 @@ type PipelineDeps struct {
|
||||
EventBus eventbus.DomainEventBus
|
||||
Config PipelineConfig
|
||||
|
||||
// ResolveContextWindow returns the effective context window (in tokens) for
|
||||
// a given provider/model pair. Nil = always use Config.ContextWindow.
|
||||
// Invoked ONCE per run by ContextStage and stored in RunState.Context.EffectiveContextWindow.
|
||||
ResolveContextWindow func(provider, model string) int
|
||||
|
||||
// Callbacks from agent.Loop — Phase 8 adapter wires these.
|
||||
EmitEvent func(event any)
|
||||
|
||||
// Auto-inject memory context (ContextStage, L0 tier).
|
||||
// Callback captures agent/tenant context via closure.
|
||||
AutoInject func(ctx context.Context, userMessage, userID string) (string, error)
|
||||
// Callback captures agent/tenant context via closure. recentContext carries
|
||||
// a short snippet of recent conversation (last 1-2 user turns) so the
|
||||
// downstream recall query can resolve pronouns and implicit references.
|
||||
// Empty recentContext = legacy single-message search semantics.
|
||||
AutoInject func(ctx context.Context, userMessage, userID, recentContext string) (string, error)
|
||||
|
||||
// InjectContext sets up agent/tenant/user/workspace/tool context values.
|
||||
// Wraps injectContext() for v3 pipeline. Called once at ContextStage start.
|
||||
@@ -92,6 +100,14 @@ type PipelineConfig struct {
|
||||
MaxTokens int
|
||||
Compaction *config.CompactionConfig
|
||||
|
||||
// ReserveTokens is a safety buffer subtracted from the history budget so
|
||||
// PruneStage compacts slightly before the hard limit. Prevents edge cases
|
||||
// where a provider returns more than MaxTokens output or where the token
|
||||
// counter's estimate drifts upward during streaming.
|
||||
// Zero (default) preserves legacy behavior: budget = contextWindow - overhead - MaxTokens.
|
||||
// Recommended: 5-10% of contextWindow for reasoning-heavy models.
|
||||
ReserveTokens int
|
||||
|
||||
// V3 memory/retrieval flags removed — always true at runtime.
|
||||
// Memory flush runs if callback != nil; auto-inject runs if AutoInject != nil.
|
||||
}
|
||||
@@ -31,8 +31,19 @@ func (s *PruneStage) Result() StageResult { return s.result }
|
||||
func (s *PruneStage) Execute(ctx context.Context, state *RunState) error {
|
||||
s.result = Continue
|
||||
|
||||
// Compute budget: context window minus overhead (system prompt + context files) minus output reserve
|
||||
budget := s.deps.Config.ContextWindow - state.Context.OverheadTokens - s.deps.Config.MaxTokens
|
||||
// Compute budget using the effective context window for this run's model.
|
||||
// ContextStage resolves EffectiveContextWindow once per run via ModelRegistry;
|
||||
// if zero (unknown model, registry not wired) fall back to the pipeline-wide
|
||||
// Config.ContextWindow for backward compatibility.
|
||||
//
|
||||
// ReserveTokens (optional, default 0) carves out a safety buffer so compaction
|
||||
// fires slightly before the hard limit — protects against provider over-delivery
|
||||
// and token-counter drift on streaming responses.
|
||||
contextWindow := state.Context.EffectiveContextWindow
|
||||
if contextWindow == 0 {
|
||||
contextWindow = s.deps.Config.ContextWindow
|
||||
}
|
||||
budget := contextWindow - state.Context.OverheadTokens - s.deps.Config.MaxTokens - s.deps.Config.ReserveTokens
|
||||
if budget <= 0 {
|
||||
return nil // no history budget, nothing to prune
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/bootstrap"
|
||||
@@ -500,6 +501,193 @@ func TestPruneStage_ZeroBudget_NoOp(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestPruneStage_EffectiveContextWindow_OverridesConfig verifies that when
|
||||
// ContextStage has resolved a per-model context window (e.g. gpt-4o=128k
|
||||
// vs Config.ContextWindow=10k), PruneStage uses the resolved value.
|
||||
//
|
||||
// Without this, swapping models at session time would leave the pipeline
|
||||
// pruning against the wrong budget. Phase 4 regression gate.
|
||||
func TestPruneStage_EffectiveContextWindow_OverridesConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
pruneCallCount := 0
|
||||
// Config says 10k window, but model-specific resolution says 128k.
|
||||
// History is 60k tokens — would be over budget at 10k (fires prune),
|
||||
// under budget at 128k (no prune should happen).
|
||||
deps := &PipelineDeps{
|
||||
Config: PipelineConfig{
|
||||
ContextWindow: 10_000, // stale/default — should be ignored
|
||||
MaxTokens: 1_000,
|
||||
},
|
||||
TokenCounter: &mockTokenCounter{countPerMessage: 600},
|
||||
PruneMessages: func(msgs []providers.Message, _ int) []providers.Message {
|
||||
pruneCallCount++
|
||||
return msgs
|
||||
},
|
||||
}
|
||||
stage := NewPruneStage(deps, nil)
|
||||
state := defaultState()
|
||||
state.Context.EffectiveContextWindow = 128_000 // resolved by ContextStage
|
||||
|
||||
history := make([]providers.Message, 100) // 100 * 600 = 60k < budget 127k
|
||||
for i := range history {
|
||||
history[i] = providers.Message{Role: "user", Content: "msg"}
|
||||
}
|
||||
state.Messages.SetHistory(history)
|
||||
|
||||
if err := stage.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("Execute() error: %v", err)
|
||||
}
|
||||
if pruneCallCount != 0 {
|
||||
t.Errorf("PruneMessages should not fire with 128k window, called %d times", pruneCallCount)
|
||||
}
|
||||
}
|
||||
|
||||
// --- buildRecentContext tests (Phase 9) ---
|
||||
|
||||
func TestBuildRecentContext_EmptyHistory(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := buildRecentContext(nil); got != "" {
|
||||
t.Errorf("empty history should return empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecentContext_SkipsNonUserMessages(t *testing.T) {
|
||||
t.Parallel()
|
||||
hist := []providers.Message{
|
||||
{Role: "user", Content: "hello"},
|
||||
{Role: "assistant", Content: "hi back"},
|
||||
{Role: "tool", Content: "tool output"},
|
||||
}
|
||||
got := buildRecentContext(hist)
|
||||
if got != "hello" {
|
||||
t.Errorf("should only include user messages, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecentContext_CapsAtTwoTurns(t *testing.T) {
|
||||
t.Parallel()
|
||||
hist := []providers.Message{
|
||||
{Role: "user", Content: "first"},
|
||||
{Role: "assistant", Content: "reply"},
|
||||
{Role: "user", Content: "second"},
|
||||
{Role: "assistant", Content: "reply"},
|
||||
{Role: "user", Content: "third"},
|
||||
}
|
||||
got := buildRecentContext(hist)
|
||||
// Should contain last two user turns ("second" and "third"), not "first"
|
||||
if !strings.Contains(got, "second") {
|
||||
t.Errorf("missing second turn, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "third") {
|
||||
t.Errorf("missing third turn, got %q", got)
|
||||
}
|
||||
if strings.Contains(got, "first") {
|
||||
t.Errorf("should exclude old turns, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecentContext_PreservesTurnOrder(t *testing.T) {
|
||||
t.Parallel()
|
||||
hist := []providers.Message{
|
||||
{Role: "user", Content: "earlier"},
|
||||
{Role: "user", Content: "later"},
|
||||
}
|
||||
got := buildRecentContext(hist)
|
||||
// "earlier" should come before "later" in the output
|
||||
earlierIdx := strings.Index(got, "earlier")
|
||||
laterIdx := strings.Index(got, "later")
|
||||
if earlierIdx == -1 || laterIdx == -1 || earlierIdx >= laterIdx {
|
||||
t.Errorf("order broken: got %q (earlier=%d, later=%d)", got, earlierIdx, laterIdx)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRecentContext_TruncatesLongMessages(t *testing.T) {
|
||||
t.Parallel()
|
||||
long := strings.Repeat("x", 500)
|
||||
hist := []providers.Message{
|
||||
{Role: "user", Content: long},
|
||||
}
|
||||
got := buildRecentContext(hist)
|
||||
if len(got) > 300 {
|
||||
t.Errorf("result should be capped at 300 chars, got %d", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
// --- Prune Stage tests continued ---
|
||||
|
||||
// TestPruneStage_ReserveTokens_BuffersBudget verifies that ReserveTokens is
|
||||
// subtracted from the usable budget so compaction fires earlier than the hard
|
||||
// limit. Phase 5 regression gate for provider over-delivery protection.
|
||||
func TestPruneStage_ReserveTokens_BuffersBudget(t *testing.T) {
|
||||
t.Parallel()
|
||||
pruneCallCount := 0
|
||||
// Without reserve: budget = 10000 - 0 - 1000 = 9000, softThreshold 6300
|
||||
// With reserve=2000: budget = 10000 - 0 - 1000 - 2000 = 7000, softThreshold 4900
|
||||
// 50 msgs * 100 = 5000 tokens — crosses the 4900 soft threshold only when reserve is set.
|
||||
deps := &PipelineDeps{
|
||||
Config: PipelineConfig{
|
||||
ContextWindow: 10_000,
|
||||
MaxTokens: 1_000,
|
||||
ReserveTokens: 2_000,
|
||||
},
|
||||
TokenCounter: &mockTokenCounter{countPerMessage: 100},
|
||||
PruneMessages: func(msgs []providers.Message, _ int) []providers.Message {
|
||||
pruneCallCount++
|
||||
return msgs[:1]
|
||||
},
|
||||
}
|
||||
stage := NewPruneStage(deps, nil)
|
||||
state := defaultState()
|
||||
history := make([]providers.Message, 50)
|
||||
for i := range history {
|
||||
history[i] = providers.Message{Role: "user", Content: "msg"}
|
||||
}
|
||||
state.Messages.SetHistory(history)
|
||||
|
||||
if err := stage.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("Execute() error: %v", err)
|
||||
}
|
||||
if pruneCallCount != 1 {
|
||||
t.Errorf("ReserveTokens should have triggered early prune, called %d times", pruneCallCount)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPruneStage_EffectiveContextWindow_ZeroFallsBackToConfig verifies
|
||||
// backward compatibility: when ContextStage hasn't resolved a model-specific
|
||||
// window (unknown model, nil resolver, no registry), PruneStage falls back
|
||||
// to Config.ContextWindow — matching pre-Phase-4 behavior.
|
||||
func TestPruneStage_EffectiveContextWindow_ZeroFallsBackToConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
pruneCallCount := 0
|
||||
deps := &PipelineDeps{
|
||||
Config: PipelineConfig{
|
||||
ContextWindow: 10_000, // used because EffectiveContextWindow=0
|
||||
MaxTokens: 1_000,
|
||||
},
|
||||
TokenCounter: &mockTokenCounter{countPerMessage: 600},
|
||||
PruneMessages: func(msgs []providers.Message, _ int) []providers.Message {
|
||||
pruneCallCount++
|
||||
return msgs[:1]
|
||||
},
|
||||
}
|
||||
stage := NewPruneStage(deps, nil)
|
||||
state := defaultState()
|
||||
// state.Context.EffectiveContextWindow left as zero
|
||||
|
||||
history := make([]providers.Message, 20) // 20 * 600 = 12k > 10k → should prune
|
||||
for i := range history {
|
||||
history[i] = providers.Message{Role: "user", Content: "msg"}
|
||||
}
|
||||
state.Messages.SetHistory(history)
|
||||
|
||||
if err := stage.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("Execute() error: %v", err)
|
||||
}
|
||||
if pruneCallCount == 0 {
|
||||
t.Error("PruneMessages should fire with 10k fallback window")
|
||||
}
|
||||
}
|
||||
|
||||
// --- ToolStage tests ---
|
||||
|
||||
func TestToolStage_NoToolCalls_NoOp(t *testing.T) {
|
||||
|
||||
@@ -15,6 +15,16 @@ type ContextState struct {
|
||||
Summary string // session summary for context continuity
|
||||
HadBootstrap bool
|
||||
OverheadTokens int // system prompt + context files (accurate via TokenCounter)
|
||||
|
||||
// EffectiveContextWindow is the context window size (in tokens) resolved
|
||||
// per-run from the provider/model pair via ModelRegistry. Resolved ONCE in
|
||||
// ContextStage and read by PruneStage on every iteration. Zero means "no
|
||||
// model-specific data available" and PruneStage falls back to
|
||||
// PipelineConfig.ContextWindow.
|
||||
//
|
||||
// Resolved once per run (not per iteration) to avoid budget skew — if the
|
||||
// model somehow changes mid-run a mismatch causes silent truncation loops.
|
||||
EffectiveContextWindow int
|
||||
}
|
||||
|
||||
// ThinkState: owned by ThinkStage.
|
||||
|
||||
@@ -26,7 +26,16 @@ func NewContactCollector(s ContactStore, c cache.Cache[bool]) *ContactCollector
|
||||
// contactType: "user" (individual sender), "group" (group chat entity), or "topic" (forum topic).
|
||||
// Pass empty threadID/threadType for base contacts (DM, group root).
|
||||
func (c *ContactCollector) EnsureContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType, threadID, threadType string) {
|
||||
key := channelType + ":" + senderID + ":" + threadID
|
||||
// Cache key must include every dimension the underlying DB unique constraint
|
||||
// uses, otherwise dedup skips legitimate upserts:
|
||||
// - tenantID: fixes cross-tenant leak (same sender in tenant A vs B)
|
||||
// - channelInstance: fixes collision when two bots in the same tenant share
|
||||
// overlapping sender ID spaces (e.g. two Telegram bot tokens with users
|
||||
// who happen to have the same Telegram user_id)
|
||||
// - threadID: different threads/topics track separate contacts
|
||||
// Zero UUID (Desktop / single-tenant) keeps legacy dedup semantics intact.
|
||||
tid := TenantIDFromContext(ctx)
|
||||
key := tid.String() + ":" + channelType + ":" + channelInstance + ":" + senderID + ":" + threadID
|
||||
if _, ok := c.seen.Get(ctx, key); ok {
|
||||
return
|
||||
}
|
||||
@@ -34,7 +43,13 @@ func (c *ContactCollector) EnsureContact(ctx context.Context, channelType, chann
|
||||
contactType = "user"
|
||||
}
|
||||
if err := c.store.UpsertContact(ctx, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType, threadID, threadType); err != nil {
|
||||
slog.Warn("contact_collector.upsert_failed", "error", err, "channel", channelType, "sender", senderID)
|
||||
slog.Warn("contact_collector.upsert_failed",
|
||||
"error", err,
|
||||
"tenant_id", tid,
|
||||
"channel", channelType,
|
||||
"instance", channelInstance,
|
||||
"sender", senderID,
|
||||
)
|
||||
return
|
||||
}
|
||||
c.seen.Set(ctx, key, true, contactSeenTTL)
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/cache"
|
||||
)
|
||||
|
||||
// mockContactStore records every UpsertContact call for assertion.
|
||||
// Only implements the methods used by ContactCollector (store.ContactStore).
|
||||
type mockContactStore struct {
|
||||
mu sync.Mutex
|
||||
upserts []mockUpsertCall
|
||||
}
|
||||
|
||||
type mockUpsertCall struct {
|
||||
tenantID uuid.UUID
|
||||
channelType string
|
||||
channelInstance string
|
||||
senderID string
|
||||
userID string
|
||||
displayName string
|
||||
username string
|
||||
peerKind string
|
||||
contactType string
|
||||
threadID string
|
||||
threadType string
|
||||
}
|
||||
|
||||
func (m *mockContactStore) UpsertContact(ctx context.Context, channelType, channelInstance, senderID, userID, displayName, username, peerKind, contactType, threadID, threadType string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.upserts = append(m.upserts, mockUpsertCall{
|
||||
tenantID: TenantIDFromContext(ctx),
|
||||
channelType: channelType,
|
||||
channelInstance: channelInstance,
|
||||
senderID: senderID,
|
||||
userID: userID,
|
||||
displayName: displayName,
|
||||
username: username,
|
||||
peerKind: peerKind,
|
||||
contactType: contactType,
|
||||
threadID: threadID,
|
||||
threadType: threadType,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockContactStore) ResolveTenantUserID(_ context.Context, _, _ string) (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// Stub methods to satisfy ContactStore interface (not used in these tests).
|
||||
func (m *mockContactStore) ListContacts(_ context.Context, _ ContactListOpts) ([]ChannelContact, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (m *mockContactStore) CountContacts(_ context.Context, _ ContactListOpts) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
func (m *mockContactStore) GetContactsBySenderIDs(_ context.Context, _ []string) (map[string]ChannelContact, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (m *mockContactStore) GetContactByID(_ context.Context, _ uuid.UUID) (*ChannelContact, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (m *mockContactStore) GetSenderIDsByContactIDs(_ context.Context, _ []uuid.UUID) ([]string, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (m *mockContactStore) MergeContacts(_ context.Context, _ []uuid.UUID, _ uuid.UUID) error {
|
||||
return nil
|
||||
}
|
||||
func (m *mockContactStore) UnmergeContacts(_ context.Context, _ []uuid.UUID) error { return nil }
|
||||
func (m *mockContactStore) GetContactsByMergedID(_ context.Context, _ uuid.UUID) ([]ChannelContact, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *mockContactStore) upsertCount() int {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return len(m.upserts)
|
||||
}
|
||||
|
||||
// TestContactCollector_SameTenantDedup verifies that repeated calls for the
|
||||
// same (tenant, channel, sender, thread) only hit the store once.
|
||||
func TestContactCollector_SameTenantDedup(t *testing.T) {
|
||||
mock := &mockContactStore{}
|
||||
c := NewContactCollector(mock, cache.NewInMemoryCache[bool]())
|
||||
|
||||
tenant := uuid.New()
|
||||
ctx := WithTenantID(context.Background(), tenant)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
c.EnsureContact(ctx, "telegram", "tg-main", "user-123", "uid-1", "Alice", "alice", "user", "user", "", "")
|
||||
}
|
||||
|
||||
if got := mock.upsertCount(); got != 1 {
|
||||
t.Errorf("same-tenant dedup broken: got %d upserts, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestContactCollector_CrossTenantIsolation verifies the core bug fix: same
|
||||
// (channel, sender, thread) in DIFFERENT tenants must produce separate upserts.
|
||||
// Before the fix, the cache key was missing tenantID so the second tenant's
|
||||
// upsert was silently skipped, causing a cross-tenant contact leak.
|
||||
func TestContactCollector_CrossTenantIsolation(t *testing.T) {
|
||||
mock := &mockContactStore{}
|
||||
c := NewContactCollector(mock, cache.NewInMemoryCache[bool]())
|
||||
|
||||
tenantA := uuid.New()
|
||||
tenantB := uuid.New()
|
||||
ctxA := WithTenantID(context.Background(), tenantA)
|
||||
ctxB := WithTenantID(context.Background(), tenantB)
|
||||
|
||||
// Same sender ID "user-123" in both tenants
|
||||
c.EnsureContact(ctxA, "telegram", "tg-main", "user-123", "uid-1", "Alice", "alice", "user", "user", "", "")
|
||||
c.EnsureContact(ctxB, "telegram", "tg-main", "user-123", "uid-1", "Alice", "alice", "user", "user", "", "")
|
||||
|
||||
if got := mock.upsertCount(); got != 2 {
|
||||
t.Errorf("cross-tenant isolation broken: got %d upserts, want 2", got)
|
||||
}
|
||||
|
||||
// Verify each tenant got its own upsert
|
||||
mock.mu.Lock()
|
||||
defer mock.mu.Unlock()
|
||||
seen := map[uuid.UUID]bool{}
|
||||
for _, u := range mock.upserts {
|
||||
seen[u.tenantID] = true
|
||||
}
|
||||
if !seen[tenantA] || !seen[tenantB] {
|
||||
t.Errorf("expected upserts for both tenants, got %+v", seen)
|
||||
}
|
||||
}
|
||||
|
||||
// TestContactCollector_ZeroTenantID verifies Desktop edition (no tenant) still
|
||||
// works with a zero/nil UUID — single-tenant semantics preserved.
|
||||
func TestContactCollector_ZeroTenantID(t *testing.T) {
|
||||
mock := &mockContactStore{}
|
||||
c := NewContactCollector(mock, cache.NewInMemoryCache[bool]())
|
||||
|
||||
// No tenant in context — Desktop / single-tenant mode
|
||||
ctx := context.Background()
|
||||
|
||||
c.EnsureContact(ctx, "telegram", "tg", "user-1", "uid", "", "", "user", "user", "", "")
|
||||
c.EnsureContact(ctx, "telegram", "tg", "user-1", "uid", "", "", "user", "user", "", "") // dup
|
||||
|
||||
if got := mock.upsertCount(); got != 1 {
|
||||
t.Errorf("zero-tenant dedup broken: got %d upserts, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestContactCollector_DifferentThreads verifies same sender in different
|
||||
// threads (same tenant) produces separate upserts.
|
||||
func TestContactCollector_DifferentThreads(t *testing.T) {
|
||||
mock := &mockContactStore{}
|
||||
c := NewContactCollector(mock, cache.NewInMemoryCache[bool]())
|
||||
|
||||
tenant := uuid.New()
|
||||
ctx := WithTenantID(context.Background(), tenant)
|
||||
|
||||
c.EnsureContact(ctx, "slack", "ws-1", "user-1", "uid", "", "", "user", "user", "thread-A", "channel")
|
||||
c.EnsureContact(ctx, "slack", "ws-1", "user-1", "uid", "", "", "user", "user", "thread-B", "channel")
|
||||
|
||||
if got := mock.upsertCount(); got != 2 {
|
||||
t.Errorf("different-thread isolation broken: got %d upserts, want 2", got)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user