feat(memory): add DAG-based context budget compression
Introduce a directed acyclic graph for managing context window budget allocation. Nodes represent message segments with token costs; the compressor selects which nodes to retain or summarize based on importance scores and available token budget.
This commit is contained in:
parent
aaaafb723d
commit
b5ab7a6108
4 changed files with 720 additions and 0 deletions
109
pkg/memory/dag/budget.go
Normal file
109
pkg/memory/dag/budget.go
Normal file
|
|
@ -0,0 +1,109 @@
|
||||||
|
package dag
|
||||||
|
|
||||||
|
// BudgetConfig defines the percentage allocation for each context section.
|
||||||
|
// All percentages should sum to 100.
|
||||||
|
type BudgetConfig struct {
|
||||||
|
SystemPromptPct int // % for system prompt (identity, rules, skills)
|
||||||
|
ObservationsPct int // % for observation block
|
||||||
|
KnowledgePct int // % for knowledge block (Focus completions)
|
||||||
|
DAGSummariesPct int // % for DAG compressed history
|
||||||
|
RawTailPct int // % for raw recent messages (uncompressed tail)
|
||||||
|
ToolResultsPct int // % for tool call results
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultBudgetConfig returns a balanced allocation.
|
||||||
|
func DefaultBudgetConfig() BudgetConfig {
|
||||||
|
return BudgetConfig{
|
||||||
|
SystemPromptPct: 20,
|
||||||
|
ObservationsPct: 10,
|
||||||
|
KnowledgePct: 5,
|
||||||
|
DAGSummariesPct: 25,
|
||||||
|
RawTailPct: 30,
|
||||||
|
ToolResultsPct: 10,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Budget represents concrete token allocations computed from config and context window.
|
||||||
|
type Budget struct {
|
||||||
|
Total int
|
||||||
|
SystemPrompt int
|
||||||
|
Observations int
|
||||||
|
Knowledge int
|
||||||
|
DAGSummaries int
|
||||||
|
RawTail int
|
||||||
|
ToolResults int
|
||||||
|
}
|
||||||
|
|
||||||
|
// ComputeBudget calculates token allocations from a context window size and config.
|
||||||
|
func ComputeBudget(contextWindow int, cfg BudgetConfig) Budget {
|
||||||
|
return Budget{
|
||||||
|
Total: contextWindow,
|
||||||
|
SystemPrompt: contextWindow * cfg.SystemPromptPct / 100,
|
||||||
|
Observations: contextWindow * cfg.ObservationsPct / 100,
|
||||||
|
Knowledge: contextWindow * cfg.KnowledgePct / 100,
|
||||||
|
DAGSummaries: contextWindow * cfg.DAGSummariesPct / 100,
|
||||||
|
RawTail: contextWindow * cfg.RawTailPct / 100,
|
||||||
|
ToolResults: contextWindow * cfg.ToolResultsPct / 100,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remaining returns tokens available after accounting for used amounts.
|
||||||
|
func (b Budget) Remaining(usedSystem, usedObs, usedKnowledge, usedDAG, usedTail, usedTools int) int {
|
||||||
|
used := usedSystem + usedObs + usedKnowledge + usedDAG + usedTail + usedTools
|
||||||
|
remaining := b.Total - used
|
||||||
|
if remaining < 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return remaining
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelectDAGLevel determines which DAG compression level to use given
|
||||||
|
// the available token budget for DAG summaries.
|
||||||
|
func SelectDAGLevel(d *DAG, budgetTokens int) Level {
|
||||||
|
if d == nil || len(d.Nodes) == 0 {
|
||||||
|
return LevelRaw
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionTokens := d.TotalTokens(LevelSession)
|
||||||
|
if sessionTokens > 0 && sessionTokens <= budgetTokens {
|
||||||
|
sectionTokens := d.TotalTokens(LevelSection)
|
||||||
|
if sectionTokens > 0 && sectionTokens <= budgetTokens {
|
||||||
|
chunkTokens := d.TotalTokens(LevelChunk)
|
||||||
|
if chunkTokens <= budgetTokens {
|
||||||
|
return LevelChunk
|
||||||
|
}
|
||||||
|
return LevelSection
|
||||||
|
}
|
||||||
|
return LevelSession
|
||||||
|
}
|
||||||
|
|
||||||
|
return LevelSession
|
||||||
|
}
|
||||||
|
|
||||||
|
// RenderDAGForBudget renders DAG nodes at the most detailed level
|
||||||
|
// that fits within the given token budget.
|
||||||
|
func RenderDAGForBudget(d *DAG, budgetTokens int) string {
|
||||||
|
if d == nil || len(d.Nodes) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
level := SelectDAGLevel(d, budgetTokens)
|
||||||
|
return d.FormatLevel(level)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TailMessageCount estimates how many raw messages fit in the tail budget.
|
||||||
|
// Uses a rough average of ~50 tokens per message.
|
||||||
|
func TailMessageCount(tailBudget int) int {
|
||||||
|
const (
|
||||||
|
avgTokensPerMessage = 50
|
||||||
|
minTail = 4
|
||||||
|
)
|
||||||
|
if tailBudget <= 0 {
|
||||||
|
return minTail
|
||||||
|
}
|
||||||
|
count := tailBudget / avgTokensPerMessage
|
||||||
|
if count < minTail {
|
||||||
|
return minTail
|
||||||
|
}
|
||||||
|
return count
|
||||||
|
}
|
||||||
232
pkg/memory/dag/compress.go
Normal file
232
pkg/memory/dag/compress.go
Normal file
|
|
@ -0,0 +1,232 @@
|
||||||
|
package dag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"unicode/utf8"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Message is a minimal message representation for DAG compression.
|
||||||
|
type Message struct {
|
||||||
|
Role string
|
||||||
|
Content string
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompressorConfig controls the deterministic compression behavior.
|
||||||
|
type CompressorConfig struct {
|
||||||
|
ChunkSize int // Messages per chunk node (default 8)
|
||||||
|
SectionSize int // Chunks per section node (default 4)
|
||||||
|
MaxSentences int // Max sentences to extract per message (default 2)
|
||||||
|
TargetRatio float64 // Target compression ratio (default 0.25 = 4:1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultCompressorConfig returns sensible defaults.
|
||||||
|
func DefaultCompressorConfig() CompressorConfig {
|
||||||
|
return CompressorConfig{
|
||||||
|
ChunkSize: 8,
|
||||||
|
SectionSize: 4,
|
||||||
|
MaxSentences: 2,
|
||||||
|
TargetRatio: 0.25,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Compressor builds a DAG from raw messages using deterministic
|
||||||
|
// extractive summarization. No LLM calls — fully reproducible.
|
||||||
|
type Compressor struct {
|
||||||
|
cfg CompressorConfig
|
||||||
|
counter int
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCompressor creates a new deterministic compressor.
|
||||||
|
func NewCompressor(cfg CompressorConfig) *Compressor {
|
||||||
|
return &Compressor{cfg: cfg}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Compress builds a hierarchical DAG from the given messages.
|
||||||
|
// Messages are grouped into chunks, chunks into sections, and
|
||||||
|
// sections into a session summary.
|
||||||
|
func (c *Compressor) Compress(msgs []Message) *DAG {
|
||||||
|
d := NewDAG()
|
||||||
|
if len(msgs) == 0 {
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// Level 1: Chunk summaries
|
||||||
|
chunkNodes := c.buildChunks(msgs, d)
|
||||||
|
if len(chunkNodes) == 0 {
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// Level 2: Section summaries (groups of chunks)
|
||||||
|
sectionNodes := c.buildSections(chunkNodes, d)
|
||||||
|
|
||||||
|
// Level 3: Session summary (if multiple sections)
|
||||||
|
if len(sectionNodes) > 1 {
|
||||||
|
sessionNode := c.buildSessionSummary(sectionNodes, d)
|
||||||
|
d.SetRoots([]string{sessionNode.ID})
|
||||||
|
} else if len(sectionNodes) == 1 {
|
||||||
|
d.SetRoots([]string{sectionNodes[0].ID})
|
||||||
|
} else {
|
||||||
|
ids := make([]string, len(chunkNodes))
|
||||||
|
for i, n := range chunkNodes {
|
||||||
|
ids[i] = n.ID
|
||||||
|
}
|
||||||
|
d.SetRoots(ids)
|
||||||
|
}
|
||||||
|
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Compressor) nextID(prefix string) string {
|
||||||
|
c.counter++
|
||||||
|
return fmt.Sprintf("%s-%d", prefix, c.counter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Compressor) buildChunks(msgs []Message, d *DAG) []*Node {
|
||||||
|
var chunks []*Node
|
||||||
|
for i := 0; i < len(msgs); i += c.cfg.ChunkSize {
|
||||||
|
end := i + c.cfg.ChunkSize
|
||||||
|
if end > len(msgs) {
|
||||||
|
end = len(msgs)
|
||||||
|
}
|
||||||
|
|
||||||
|
chunk := msgs[i:end]
|
||||||
|
summary := c.extractChunkSummary(chunk)
|
||||||
|
node := &Node{
|
||||||
|
ID: c.nextID("chunk"),
|
||||||
|
Level: LevelChunk,
|
||||||
|
Summary: summary,
|
||||||
|
Tokens: estimateTokens(summary),
|
||||||
|
StartIdx: i,
|
||||||
|
EndIdx: end,
|
||||||
|
}
|
||||||
|
d.Add(node)
|
||||||
|
chunks = append(chunks, node)
|
||||||
|
}
|
||||||
|
return chunks
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Compressor) buildSections(chunks []*Node, d *DAG) []*Node {
|
||||||
|
if len(chunks) <= c.cfg.SectionSize {
|
||||||
|
return chunks
|
||||||
|
}
|
||||||
|
|
||||||
|
var sections []*Node
|
||||||
|
for i := 0; i < len(chunks); i += c.cfg.SectionSize {
|
||||||
|
end := i + c.cfg.SectionSize
|
||||||
|
if end > len(chunks) {
|
||||||
|
end = len(chunks)
|
||||||
|
}
|
||||||
|
|
||||||
|
group := chunks[i:end]
|
||||||
|
childIDs := make([]string, len(group))
|
||||||
|
var summaryParts []string
|
||||||
|
for j, ch := range group {
|
||||||
|
childIDs[j] = ch.ID
|
||||||
|
summaryParts = append(summaryParts, ch.Summary)
|
||||||
|
}
|
||||||
|
|
||||||
|
combined := strings.Join(summaryParts, " ")
|
||||||
|
summary := extractSentences(combined, c.cfg.MaxSentences)
|
||||||
|
node := &Node{
|
||||||
|
ID: c.nextID("section"),
|
||||||
|
Level: LevelSection,
|
||||||
|
Summary: summary,
|
||||||
|
Tokens: estimateTokens(summary),
|
||||||
|
StartIdx: group[0].StartIdx,
|
||||||
|
EndIdx: group[len(group)-1].EndIdx,
|
||||||
|
Children: childIDs,
|
||||||
|
}
|
||||||
|
d.Add(node)
|
||||||
|
sections = append(sections, node)
|
||||||
|
}
|
||||||
|
return sections
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Compressor) buildSessionSummary(sections []*Node, d *DAG) *Node {
|
||||||
|
childIDs := make([]string, len(sections))
|
||||||
|
var summaryParts []string
|
||||||
|
for i, s := range sections {
|
||||||
|
childIDs[i] = s.ID
|
||||||
|
summaryParts = append(summaryParts, s.Summary)
|
||||||
|
}
|
||||||
|
|
||||||
|
combined := strings.Join(summaryParts, " ")
|
||||||
|
summary := extractSentences(combined, c.cfg.MaxSentences)
|
||||||
|
node := &Node{
|
||||||
|
ID: c.nextID("session"),
|
||||||
|
Level: LevelSession,
|
||||||
|
Summary: summary,
|
||||||
|
Tokens: estimateTokens(summary),
|
||||||
|
StartIdx: sections[0].StartIdx,
|
||||||
|
EndIdx: sections[len(sections)-1].EndIdx,
|
||||||
|
Children: childIDs,
|
||||||
|
}
|
||||||
|
d.Add(node)
|
||||||
|
return node
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractChunkSummary produces a deterministic summary of a message chunk.
|
||||||
|
// Strategy: for each message, take the first sentence (or first N chars if short).
|
||||||
|
// Prefix with role to preserve conversational structure.
|
||||||
|
func (c *Compressor) extractChunkSummary(msgs []Message) string {
|
||||||
|
var parts []string
|
||||||
|
for _, m := range msgs {
|
||||||
|
sentence := extractSentences(m.Content, 1)
|
||||||
|
if sentence == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
parts = append(parts, fmt.Sprintf("%s: %s", m.Role, sentence))
|
||||||
|
}
|
||||||
|
return strings.Join(parts, " | ")
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractSentences pulls the first N sentences from text.
|
||||||
|
// A sentence ends at '.', '!', '?', or '\n\n'.
|
||||||
|
func extractSentences(text string, n int) string {
|
||||||
|
if n <= 0 || text == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
text = strings.TrimSpace(text)
|
||||||
|
var result []string
|
||||||
|
remaining := text
|
||||||
|
|
||||||
|
for i := 0; i < n && remaining != ""; i++ {
|
||||||
|
idx := findSentenceEnd(remaining)
|
||||||
|
if idx < 0 {
|
||||||
|
result = append(result, strings.TrimSpace(remaining))
|
||||||
|
break
|
||||||
|
}
|
||||||
|
sentence := strings.TrimSpace(remaining[:idx+1])
|
||||||
|
if sentence != "" {
|
||||||
|
result = append(result, sentence)
|
||||||
|
}
|
||||||
|
remaining = strings.TrimSpace(remaining[idx+1:])
|
||||||
|
}
|
||||||
|
|
||||||
|
joined := strings.Join(result, " ")
|
||||||
|
const maxLen = 200
|
||||||
|
if utf8.RuneCountInString(joined) > maxLen {
|
||||||
|
runes := []rune(joined)
|
||||||
|
return string(runes[:maxLen]) + "..."
|
||||||
|
}
|
||||||
|
return joined
|
||||||
|
}
|
||||||
|
|
||||||
|
func findSentenceEnd(s string) int {
|
||||||
|
for i, r := range s {
|
||||||
|
if r == '.' || r == '!' || r == '?' {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
if r == '\n' && i+1 < len(s) && s[i+1] == '\n' {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
|
// estimateTokens provides a rough token count (chars * 2/5 heuristic).
|
||||||
|
func estimateTokens(s string) int {
|
||||||
|
return utf8.RuneCountInString(s) * 2 / 5
|
||||||
|
}
|
||||||
245
pkg/memory/dag/dag_test.go
Normal file
245
pkg/memory/dag/dag_test.go
Normal file
|
|
@ -0,0 +1,245 @@
|
||||||
|
package dag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func makeMessages(n int) []Message {
|
||||||
|
msgs := make([]Message, n)
|
||||||
|
for i := range msgs {
|
||||||
|
role := "user"
|
||||||
|
if i%2 == 1 {
|
||||||
|
role = "assistant"
|
||||||
|
}
|
||||||
|
msgs[i] = Message{
|
||||||
|
Role: role,
|
||||||
|
Content: strings.Repeat("word ", 20) + ".",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return msgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompressor_EmptyInput(t *testing.T) {
|
||||||
|
c := NewCompressor(DefaultCompressorConfig())
|
||||||
|
d := c.Compress(nil)
|
||||||
|
assert.Empty(t, d.Nodes)
|
||||||
|
assert.Empty(t, d.Roots)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompressor_SmallInput(t *testing.T) {
|
||||||
|
c := NewCompressor(DefaultCompressorConfig())
|
||||||
|
msgs := []Message{
|
||||||
|
{Role: "user", Content: "Hello, how are you?"},
|
||||||
|
{Role: "assistant", Content: "I'm doing well. Thanks for asking!"},
|
||||||
|
}
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
require.Len(t, d.Nodes, 1)
|
||||||
|
chunk := d.NodesAtLevel(LevelChunk)
|
||||||
|
require.Len(t, chunk, 1)
|
||||||
|
assert.Equal(t, 0, chunk[0].StartIdx)
|
||||||
|
assert.Equal(t, 2, chunk[0].EndIdx)
|
||||||
|
assert.Contains(t, chunk[0].Summary, "user:")
|
||||||
|
assert.Contains(t, chunk[0].Summary, "assistant:")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompressor_ChunkSplitting(t *testing.T) {
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
msgs := makeMessages(12)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
chunks := d.NodesAtLevel(LevelChunk)
|
||||||
|
require.Len(t, chunks, 3)
|
||||||
|
|
||||||
|
assert.Equal(t, 0, chunks[0].StartIdx)
|
||||||
|
assert.Equal(t, 4, chunks[0].EndIdx)
|
||||||
|
assert.Equal(t, 4, chunks[1].StartIdx)
|
||||||
|
assert.Equal(t, 8, chunks[1].EndIdx)
|
||||||
|
assert.Equal(t, 8, chunks[2].StartIdx)
|
||||||
|
assert.Equal(t, 12, chunks[2].EndIdx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompressor_SectionBuilding(t *testing.T) {
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
cfg.SectionSize = 2
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
// 20 messages = 5 chunks, section_size=2 => 3 sections
|
||||||
|
msgs := makeMessages(20)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
chunks := d.NodesAtLevel(LevelChunk)
|
||||||
|
assert.Len(t, chunks, 5)
|
||||||
|
|
||||||
|
sections := d.NodesAtLevel(LevelSection)
|
||||||
|
assert.Len(t, sections, 3)
|
||||||
|
|
||||||
|
// First section covers chunks 0-1 (msgs 0-7)
|
||||||
|
assert.Equal(t, 0, sections[0].StartIdx)
|
||||||
|
assert.Equal(t, 8, sections[0].EndIdx)
|
||||||
|
assert.Len(t, sections[0].Children, 2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompressor_SessionSummary(t *testing.T) {
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
cfg.SectionSize = 2
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
// 24 messages => 6 chunks => 3 sections => 1 session
|
||||||
|
msgs := makeMessages(24)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
sessions := d.NodesAtLevel(LevelSession)
|
||||||
|
require.Len(t, sessions, 1)
|
||||||
|
assert.Equal(t, 0, sessions[0].StartIdx)
|
||||||
|
assert.Equal(t, 24, sessions[0].EndIdx)
|
||||||
|
assert.Len(t, sessions[0].Children, 3)
|
||||||
|
|
||||||
|
require.Len(t, d.Roots, 1)
|
||||||
|
assert.Equal(t, sessions[0].ID, d.Roots[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractSentences(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
text string
|
||||||
|
n int
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"single sentence", "Hello world.", 1, "Hello world."},
|
||||||
|
{"two sentences", "First sentence. Second sentence.", 2, "First sentence. Second sentence."},
|
||||||
|
{"extract one from many", "A. B. C. D.", 1, "A."},
|
||||||
|
{"empty", "", 1, ""},
|
||||||
|
{"no period", "Hello world", 1, "Hello world"},
|
||||||
|
{"question mark", "What? Who knows.", 1, "What?"},
|
||||||
|
{"exclamation", "Wow! Amazing.", 2, "Wow! Amazing."},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := extractSentences(tt.text, tt.n)
|
||||||
|
assert.Equal(t, tt.want, got)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractSentences_Truncation(t *testing.T) {
|
||||||
|
long := strings.Repeat("This is a very long sentence with many words. ", 20)
|
||||||
|
result := extractSentences(long, 5)
|
||||||
|
assert.LessOrEqual(t, len([]rune(result)), 210)
|
||||||
|
assert.True(t, strings.HasSuffix(result, "..."))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNode_FormatForPrompt(t *testing.T) {
|
||||||
|
n := &Node{
|
||||||
|
ID: "chunk-1",
|
||||||
|
Level: LevelChunk,
|
||||||
|
Summary: "User asked about auth. Assistant explained JWT flow.",
|
||||||
|
StartIdx: 0,
|
||||||
|
EndIdx: 8,
|
||||||
|
}
|
||||||
|
formatted := n.FormatForPrompt()
|
||||||
|
assert.Contains(t, formatted, "[chunk msgs 0-7]")
|
||||||
|
assert.Contains(t, formatted, "User asked about auth")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDAG_FormatLevel(t *testing.T) {
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
msgs := makeMessages(8)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
output := d.FormatLevel(LevelChunk)
|
||||||
|
assert.NotEmpty(t, output)
|
||||||
|
lines := strings.Split(strings.TrimSpace(output), "\n")
|
||||||
|
assert.Len(t, lines, 2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDAG_TotalTokens(t *testing.T) {
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
msgs := makeMessages(12)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
chunkTokens := d.TotalTokens(LevelChunk)
|
||||||
|
assert.Greater(t, chunkTokens, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComputeBudget(t *testing.T) {
|
||||||
|
cfg := DefaultBudgetConfig()
|
||||||
|
b := ComputeBudget(100000, cfg)
|
||||||
|
|
||||||
|
assert.Equal(t, 100000, b.Total)
|
||||||
|
assert.Equal(t, 20000, b.SystemPrompt)
|
||||||
|
assert.Equal(t, 10000, b.Observations)
|
||||||
|
assert.Equal(t, 5000, b.Knowledge)
|
||||||
|
assert.Equal(t, 25000, b.DAGSummaries)
|
||||||
|
assert.Equal(t, 30000, b.RawTail)
|
||||||
|
assert.Equal(t, 10000, b.ToolResults)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBudget_Remaining(t *testing.T) {
|
||||||
|
b := Budget{Total: 10000}
|
||||||
|
assert.Equal(t, 7000, b.Remaining(1000, 500, 500, 500, 500, 0))
|
||||||
|
assert.Equal(t, 0, b.Remaining(5000, 3000, 1000, 1000, 1000, 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSelectDAGLevel(t *testing.T) {
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
cfg.SectionSize = 2
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
msgs := makeMessages(24)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
chunkTokens := d.TotalTokens(LevelChunk)
|
||||||
|
|
||||||
|
// Large budget -> most detailed (chunk)
|
||||||
|
assert.Equal(t, LevelChunk, SelectDAGLevel(d, chunkTokens+1000))
|
||||||
|
|
||||||
|
// Very small budget -> session level
|
||||||
|
assert.Equal(t, LevelSession, SelectDAGLevel(d, 10))
|
||||||
|
|
||||||
|
// Nil DAG
|
||||||
|
assert.Equal(t, LevelRaw, SelectDAGLevel(nil, 1000))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTailMessageCount(t *testing.T) {
|
||||||
|
assert.Equal(t, 4, TailMessageCount(100)) // Minimum
|
||||||
|
assert.Equal(t, 20, TailMessageCount(1000)) // 1000/50
|
||||||
|
assert.Equal(t, 4, TailMessageCount(0)) // Zero budget
|
||||||
|
assert.Equal(t, 4, TailMessageCount(-1)) // Negative
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRenderDAGForBudget(t *testing.T) {
|
||||||
|
assert.Empty(t, RenderDAGForBudget(nil, 1000))
|
||||||
|
|
||||||
|
cfg := DefaultCompressorConfig()
|
||||||
|
cfg.ChunkSize = 4
|
||||||
|
cfg.SectionSize = 2
|
||||||
|
c := NewCompressor(cfg)
|
||||||
|
|
||||||
|
// 16 messages => 4 chunks => 2 sections => 1 session
|
||||||
|
msgs := makeMessages(16)
|
||||||
|
d := c.Compress(msgs)
|
||||||
|
|
||||||
|
// Large budget should get chunk-level detail
|
||||||
|
chunkTokens := d.TotalTokens(LevelChunk)
|
||||||
|
output := RenderDAGForBudget(d, chunkTokens+1000)
|
||||||
|
assert.NotEmpty(t, output)
|
||||||
|
assert.Contains(t, output, "[chunk")
|
||||||
|
}
|
||||||
134
pkg/memory/dag/node.go
Normal file
134
pkg/memory/dag/node.go
Normal file
|
|
@ -0,0 +1,134 @@
|
||||||
|
package dag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Level represents the compression tier of a DAG node.
|
||||||
|
type Level int
|
||||||
|
|
||||||
|
const (
|
||||||
|
LevelRaw Level = 0 // Original messages (not stored as nodes)
|
||||||
|
LevelChunk Level = 1 // Chunk summary (~8-16 messages)
|
||||||
|
LevelSection Level = 2 // Section summary (group of chunks)
|
||||||
|
LevelSession Level = 3 // Session summary (top-level)
|
||||||
|
)
|
||||||
|
|
||||||
|
func (l Level) String() string {
|
||||||
|
switch l {
|
||||||
|
case LevelRaw:
|
||||||
|
return "raw"
|
||||||
|
case LevelChunk:
|
||||||
|
return "chunk"
|
||||||
|
case LevelSection:
|
||||||
|
return "section"
|
||||||
|
case LevelSession:
|
||||||
|
return "session"
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("level-%d", l)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Node is a single node in the compression DAG. Each node stores an
|
||||||
|
// extractive summary and retains lossless pointers back to the
|
||||||
|
// original message range it covers.
|
||||||
|
type Node struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Level Level `json:"level"`
|
||||||
|
Summary string `json:"summary"`
|
||||||
|
Tokens int `json:"tokens"`
|
||||||
|
StartIdx int `json:"start_idx"` // Inclusive index into original message slice
|
||||||
|
EndIdx int `json:"end_idx"` // Exclusive index into original message slice
|
||||||
|
Children []string `json:"children,omitempty"` // Child node IDs (lower level)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MessageRange returns the [start, end) range of original messages this node covers.
|
||||||
|
func (n *Node) MessageRange() (int, int) {
|
||||||
|
return n.StartIdx, n.EndIdx
|
||||||
|
}
|
||||||
|
|
||||||
|
// Span returns how many original messages this node covers.
|
||||||
|
func (n *Node) Span() int {
|
||||||
|
return n.EndIdx - n.StartIdx
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatForPrompt renders the node as a compact block for context injection.
|
||||||
|
func (n *Node) FormatForPrompt() string {
|
||||||
|
return fmt.Sprintf("[%s msgs %d-%d] %s", n.Level, n.StartIdx, n.EndIdx-1, n.Summary)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DAG is the hierarchical compression tree. Nodes at higher levels
|
||||||
|
// summarize groups of lower-level nodes. The root level covers the
|
||||||
|
// entire session.
|
||||||
|
type DAG struct {
|
||||||
|
Nodes map[string]*Node `json:"nodes"`
|
||||||
|
Roots []string `json:"roots"` // Top-level node IDs (highest compression)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDAG creates an empty DAG.
|
||||||
|
func NewDAG() *DAG {
|
||||||
|
return &DAG{
|
||||||
|
Nodes: make(map[string]*Node),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add inserts a node into the DAG.
|
||||||
|
func (d *DAG) Add(node *Node) {
|
||||||
|
d.Nodes[node.ID] = node
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetRoots sets the top-level node IDs.
|
||||||
|
func (d *DAG) SetRoots(ids []string) {
|
||||||
|
d.Roots = ids
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get returns a node by ID, or nil if not found.
|
||||||
|
func (d *DAG) Get(id string) *Node {
|
||||||
|
return d.Nodes[id]
|
||||||
|
}
|
||||||
|
|
||||||
|
// NodesAtLevel returns all nodes at the given compression level, ordered by StartIdx.
|
||||||
|
func (d *DAG) NodesAtLevel(level Level) []*Node {
|
||||||
|
var result []*Node
|
||||||
|
for _, n := range d.Nodes {
|
||||||
|
if n.Level == level {
|
||||||
|
result = append(result, n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sortByStart(result)
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// TotalTokens returns the sum of tokens across all nodes at the given level.
|
||||||
|
func (d *DAG) TotalTokens(level Level) int {
|
||||||
|
total := 0
|
||||||
|
for _, n := range d.Nodes {
|
||||||
|
if n.Level == level {
|
||||||
|
total += n.Tokens
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return total
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatLevel renders all nodes at a given level as a prompt-ready string.
|
||||||
|
func (d *DAG) FormatLevel(level Level) string {
|
||||||
|
nodes := d.NodesAtLevel(level)
|
||||||
|
if len(nodes) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
var sb strings.Builder
|
||||||
|
for _, n := range nodes {
|
||||||
|
sb.WriteString(n.FormatForPrompt())
|
||||||
|
sb.WriteByte('\n')
|
||||||
|
}
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortByStart(nodes []*Node) {
|
||||||
|
for i := 1; i < len(nodes); i++ {
|
||||||
|
for j := i; j > 0 && nodes[j].StartIdx < nodes[j-1].StartIdx; j-- {
|
||||||
|
nodes[j], nodes[j-1] = nodes[j-1], nodes[j]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue