Merge pull request #44 from dj-oyu/perf/hot-path-optimization

perf: optimize LLM iteration hot path
This commit is contained in:
dj-oyu 2026-03-15 19:51:55 +09:00 committed by GitHub
commit dbbb8a580d
5 changed files with 218 additions and 84 deletions

View file

@ -718,30 +718,27 @@ func sanitizeHistoryForProvider(history []providers.Message) []providers.Message
return history
}
sanitized := make([]providers.Message, 0, len(history))
for _, msg := range history {
// Single-pass sanitization: filters orphaned messages and validates
// tool-call completeness using forward-looking checks.
result := make([]providers.Message, 0, len(history))
for i := 0; i < len(history); i++ {
msg := history[i]
switch msg.Role {
case "system":
// Drop system messages from history. BuildMessages always
// constructs its own single system message (static + dynamic +
// summary); extra system messages would break providers that
// only accept one (Anthropic, Codex).
logger.DebugCF("agent", "Dropping system message from history", map[string]any{})
continue
case "tool":
if len(sanitized) == 0 {
if len(result) == 0 {
logger.DebugCF("agent", "Dropping orphaned leading tool message", map[string]any{})
continue
}
// Walk backwards to find the nearest assistant message,
// skipping over any preceding tool messages (multi-tool-call case).
foundAssistant := false
for i := len(sanitized) - 1; i >= 0; i-- {
if sanitized[i].Role == "tool" {
for j := len(result) - 1; j >= 0; j-- {
if result[j].Role == "tool" {
continue
}
if sanitized[i].Role == "assistant" && len(sanitized[i].ToolCalls) > 0 {
if result[j].Role == "assistant" && len(result[j].ToolCalls) > 0 {
foundAssistant = true
}
break
@ -750,15 +747,15 @@ func sanitizeHistoryForProvider(history []providers.Message) []providers.Message
logger.DebugCF("agent", "Dropping orphaned tool message", map[string]any{})
continue
}
sanitized = append(sanitized, msg)
result = append(result, msg)
case "assistant":
if len(msg.ToolCalls) > 0 {
if len(sanitized) == 0 {
if len(result) == 0 {
logger.DebugCF("agent", "Dropping assistant tool-call turn at history start", map[string]any{})
continue
}
prev := sanitized[len(sanitized)-1]
prev := result[len(result)-1]
if prev.Role != "user" && prev.Role != "tool" {
logger.DebugCF(
"agent",
@ -767,68 +764,57 @@ func sanitizeHistoryForProvider(history []providers.Message) []providers.Message
)
continue
}
// Forward-looking completeness check: verify all tool_call IDs
// have matching tool result messages immediately following.
expected := make(map[string]bool, len(msg.ToolCalls))
for _, tc := range msg.ToolCalls {
expected[tc.ID] = false
}
toolMsgCount := 0
for j := i + 1; j < len(history); j++ {
next := history[j]
if next.Role == "system" {
continue // system messages will be dropped; skip over them
}
if next.Role != "tool" {
break
}
toolMsgCount++
if _, exists := expected[next.ToolCallID]; exists {
expected[next.ToolCallID] = true
}
}
allFound := true
for toolCallID, found := range expected {
if !found {
allFound = false
logger.DebugCF(
"agent",
"Dropping assistant message with incomplete tool results",
map[string]any{
"missing_tool_call_id": toolCallID,
"expected_count": len(expected),
"found_count": toolMsgCount,
},
)
break
}
}
if !allFound {
i += toolMsgCount
continue
}
}
sanitized = append(sanitized, msg)
result = append(result, msg)
default:
sanitized = append(sanitized, msg)
result = append(result, msg)
}
}
// Second pass: ensure every assistant message with tool_calls has matching
// tool result messages following it. This is required by strict providers
// like DeepSeek that enforce: "An assistant message with 'tool_calls' must
// be followed by tool messages responding to each 'tool_call_id'."
final := make([]providers.Message, 0, len(sanitized))
for i := 0; i < len(sanitized); i++ {
msg := sanitized[i]
if msg.Role == "assistant" && len(msg.ToolCalls) > 0 {
// Collect expected tool_call IDs
expected := make(map[string]bool, len(msg.ToolCalls))
for _, tc := range msg.ToolCalls {
expected[tc.ID] = false
}
// Check following messages for matching tool results
toolMsgCount := 0
for j := i + 1; j < len(sanitized); j++ {
if sanitized[j].Role != "tool" {
break
}
toolMsgCount++
if _, exists := expected[sanitized[j].ToolCallID]; exists {
expected[sanitized[j].ToolCallID] = true
}
}
// If any tool_call_id is missing, drop this assistant message and its partial tool messages
allFound := true
for toolCallID, found := range expected {
if !found {
allFound = false
logger.DebugCF(
"agent",
"Dropping assistant message with incomplete tool results",
map[string]any{
"missing_tool_call_id": toolCallID,
"expected_count": len(expected),
"found_count": toolMsgCount,
},
)
break
}
}
if !allFound {
// Skip this assistant message and its tool messages
i += toolMsgCount
continue
}
}
final = append(final, msg)
}
return final
return result
}
func (cb *ContextBuilder) AddToolResult(

View file

@ -176,7 +176,9 @@ func (la *LegacyAdapter) AddFullMessage(sessionKey string, msg providers.Message
c.dirty = true
}
// GetHistory returns a defensive copy of the session messages.
// GetHistory returns the session messages directly (read-only contract).
// Callers must not mutate the returned slice. If mutation is needed,
// copy the slice first or use SetHistory.
func (la *LegacyAdapter) GetHistory(key string) []providers.Message {
la.mu.RLock()
@ -213,11 +215,7 @@ func (la *LegacyAdapter) GetHistory(key string) []providers.Message {
defer la.mu.RUnlock()
history := make([]providers.Message, len(c.messages))
copy(history, c.messages)
return history
return c.messages
}
// SetHistory replaces the session's message history entirely.

View file

@ -277,7 +277,7 @@ func TestBackend_TruncateHistory_LargerThanLen(t *testing.T) {
}
}
func TestBackend_GetHistory_DefensiveCopy(t *testing.T) {
func TestBackend_GetHistory_ReadOnlyContract(t *testing.T) {
for name, be := range backends(t) {
t.Run(name, func(t *testing.T) {
be.GetOrCreate("k1")
@ -285,13 +285,14 @@ func TestBackend_GetHistory_DefensiveCopy(t *testing.T) {
be.AddMessage("k1", "user", "hello")
h1 := be.GetHistory("k1")
h1[0].Content = "modified"
h2 := be.GetHistory("k1")
if h2[0].Content != "hello" {
t.Errorf("defensive copy failed: %q", h2[0].Content)
// Read-only contract: both calls return the same backing data.
if len(h1) != len(h2) {
t.Errorf("expected same length, got %d vs %d", len(h1), len(h2))
}
if h1[0].Content != "hello" || h2[0].Content != "hello" {
t.Errorf("expected content 'hello'")
}
})
}

View file

@ -24,6 +24,10 @@ type ToolRegistry struct {
tools map[string]*ToolEntry
mu sync.RWMutex
version atomic.Uint64 // incremented on Register/RegisterHidden for cache invalidation
// Cached provider definitions, invalidated when version changes.
cachedDefs []providers.ToolDefinition
cachedVersion uint64
}
func NewToolRegistry() *ToolRegistry {
@ -81,6 +85,9 @@ func (r *ToolRegistry) PromoteTools(names []string, ttl int) {
}
}
}
if promoted > 0 {
r.version.Add(1) // invalidate ToProviderDefs cache
}
logger.DebugCF(
"tools",
"PromoteTools completed",
@ -92,11 +99,16 @@ func (r *ToolRegistry) PromoteTools(names []string, ttl int) {
func (r *ToolRegistry) TickTTL() {
r.mu.Lock()
defer r.mu.Unlock()
changed := false
for _, entry := range r.tools {
if !entry.IsCore && entry.TTL > 0 {
entry.TTL--
changed = true
}
}
if changed {
r.version.Add(1) // invalidate ToProviderDefs cache
}
}
// Version returns the current registry version (atomically).
@ -260,9 +272,25 @@ func (r *ToolRegistry) GetDefinitions() []map[string]any {
// ToProviderDefs converts tool definitions to provider-compatible format.
// This is the format expected by LLM provider APIs.
// Results are cached and invalidated when the registry version changes
// (i.e. when tools are registered). Callers must not mutate the returned slice.
func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition {
r.mu.RLock()
defer r.mu.RUnlock()
v := r.version.Load()
if r.cachedVersion == v && r.cachedDefs != nil {
defs := r.cachedDefs
r.mu.RUnlock()
return defs
}
r.mu.RUnlock()
r.mu.Lock()
defer r.mu.Unlock()
// Double-check after upgrading to write lock.
v = r.version.Load()
if r.cachedVersion == v && r.cachedDefs != nil {
return r.cachedDefs
}
sorted := r.sortedToolNames()
definitions := make([]providers.ToolDefinition, 0, len(sorted))
@ -299,6 +327,8 @@ func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition {
},
})
}
r.cachedDefs = definitions
r.cachedVersion = v
return definitions
}

View file

@ -3,6 +3,7 @@ package tools
import (
"context"
"encoding/json"
"fmt"
"strings"
"sync"
"testing"
@ -360,3 +361,121 @@ func TestToolRegistry_ConcurrentAccess(t *testing.T) {
t.Error("expected tools to be registered after concurrent access")
}
}
func TestToolRegistry_ToProviderDefs_Cache(t *testing.T) {
r := NewToolRegistry()
params := map[string]any{"type": "object", "properties": map[string]any{}}
for i := range 10 {
r.Register(&mockRegistryTool{
name: string(rune('a' + i)),
desc: "tool",
params: params,
result: SilentResult("ok"),
})
}
defs1 := r.ToProviderDefs()
defs2 := r.ToProviderDefs()
// Should return the same backing slice (cached).
if &defs1[0] != &defs2[0] {
t.Error("expected cached result to return same slice")
}
// Registering a new tool should invalidate the cache.
r.Register(&mockRegistryTool{
name: "new_tool",
desc: "new",
params: params,
result: SilentResult("ok"),
})
defs3 := r.ToProviderDefs()
if len(defs3) != 11 {
t.Errorf("expected 11 defs after new registration, got %d", len(defs3))
}
if &defs1[0] == &defs3[0] {
t.Error("expected cache invalidation after Register")
}
}
func TestToolRegistry_ToProviderDefs_CacheInvalidatedByTTL(t *testing.T) {
r := NewToolRegistry()
r.RegisterHidden(&mockRegistryTool{
name: "hidden",
desc: "hidden tool",
params: map[string]any{"type": "object"},
result: SilentResult("ok"),
})
// Hidden tool with TTL=0 should not appear.
defs1 := r.ToProviderDefs()
if len(defs1) != 0 {
t.Fatalf("expected 0 defs for hidden tool with TTL=0, got %d", len(defs1))
}
// Promote the tool.
r.PromoteTools([]string{"hidden"}, 2)
defs2 := r.ToProviderDefs()
if len(defs2) != 1 {
t.Fatalf("expected 1 def after promote, got %d", len(defs2))
}
// Tick TTL twice to expire.
r.TickTTL()
r.TickTTL()
defs3 := r.ToProviderDefs()
if len(defs3) != 0 {
t.Errorf("expected 0 defs after TTL expiry, got %d", len(defs3))
}
}
func BenchmarkToProviderDefs(b *testing.B) {
r := NewToolRegistry()
params := map[string]any{
"type": "object",
"properties": map[string]any{
"path": map[string]any{"type": "string", "description": "file path"},
},
"required": []string{"path"},
}
for i := range 30 {
r.Register(&mockRegistryTool{
name: fmt.Sprintf("tool_%02d", i),
desc: fmt.Sprintf("Description for tool %d", i),
params: params,
result: SilentResult("ok"),
})
}
b.ResetTimer()
for range b.N {
r.ToProviderDefs()
}
}
func BenchmarkToProviderDefs_NoCache(b *testing.B) {
r := NewToolRegistry()
params := map[string]any{
"type": "object",
"properties": map[string]any{
"path": map[string]any{"type": "string", "description": "file path"},
},
"required": []string{"path"},
}
for i := range 30 {
r.Register(&mockRegistryTool{
name: fmt.Sprintf("tool_%02d", i),
desc: fmt.Sprintf("Description for tool %d", i),
params: params,
result: SilentResult("ok"),
})
}
b.ResetTimer()
for range b.N {
// Force cache miss by bumping version each time.
r.version.Add(1)
r.ToProviderDefs()
}
}