diff --git a/chains/chains.go b/chains/chains.go index 56c14f16..15649840 100644 --- a/chains/chains.go +++ b/chains/chains.go @@ -19,9 +19,9 @@ type Chain interface { Call(ctx context.Context, inputs map[string]any, options ...ChainCallOption) (map[string]any, error) // GetMemory gets the memory of the chain. GetMemory() schema.Memory - // InputKeys returns the input keys the chain expects. + // GetInputKeys returns the input keys the chain expects. GetInputKeys() []string - // OutputKeys returns the output keys the chain returns. + // GetOutputKeys returns the output keys the chain returns. GetOutputKeys() []string } diff --git a/chains/chains_test.go b/chains/chains_test.go index 9bb2b353..c88fa426 100644 --- a/chains/chains_test.go +++ b/chains/chains_test.go @@ -8,6 +8,7 @@ import ( "time" "github.com/stretchr/testify/require" + "github.com/tmc/langchaingo/llms" "github.com/tmc/langchaingo/prompts" "github.com/tmc/langchaingo/schema" diff --git a/chains/conversation_test.go b/chains/conversation_test.go index 87409fbc..57541ee2 100644 --- a/chains/conversation_test.go +++ b/chains/conversation_test.go @@ -7,12 +7,14 @@ import ( "testing" "github.com/stretchr/testify/require" + "github.com/tmc/langchaingo/llms/openai" "github.com/tmc/langchaingo/memory" ) func TestConversation(t *testing.T) { t.Parallel() + if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" { t.Skip("OPENAI_API_KEY not set") } @@ -27,3 +29,21 @@ func TestConversation(t *testing.T) { require.NoError(t, err) require.True(t, strings.Contains(res, "Jim"), `result does not contain the keyword 'Jim'`) } + +func TestConversationMemoryPrune(t *testing.T) { + if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" { + t.Skip("OPENAI_API_KEY not set") + } + + llm, err := openai.New() + require.NoError(t, err) + + c := NewConversation(llm, memory.NewTokenBuffer(llm, 100, memory.WithReturnMessages(true))) + _, err = Run(context.Background(), c, "Hi! I'm Jim") + require.NoError(t, err) + + res, err := Run(context.Background(), c, "What is my name?") + require.NoError(t, err) + require.True(t, strings.Contains(res, "Jim"), `result does not contain the keyword 'Jim'`) + +} diff --git a/memory/buffer_test.go b/memory/buffer_test.go index 23d3c701..a3b641ac 100644 --- a/memory/buffer_test.go +++ b/memory/buffer_test.go @@ -5,6 +5,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/tmc/langchaingo/schema" ) @@ -86,6 +87,9 @@ func (t testChatMessageHistory) AddMessage(_ schema.ChatMessage) { func (t testChatMessageHistory) Clear() { } +func (t testChatMessageHistory) SetMessages(messages []schema.ChatMessage) { +} + func (t testChatMessageHistory) Messages() []schema.ChatMessage { return []schema.ChatMessage{ schema.HumanChatMessage{Text: "user message test"}, diff --git a/memory/chat.go b/memory/chat.go index b116e89b..69970d9b 100644 --- a/memory/chat.go +++ b/memory/chat.go @@ -37,3 +37,7 @@ func (h *ChatMessageHistory) Clear() { func (h *ChatMessageHistory) AddMessage(message schema.ChatMessage) { h.messages = append(h.messages, message) } + +func (h *ChatMessageHistory) SetMessages(messages []schema.ChatMessage) { + h.messages = messages +} diff --git a/memory/token_buffer.go b/memory/token_buffer.go new file mode 100644 index 00000000..d3a558ad --- /dev/null +++ b/memory/token_buffer.go @@ -0,0 +1,81 @@ +package memory + +import ( + "github.com/tmc/langchaingo/llms" + "github.com/tmc/langchaingo/schema" +) + +// TokenBuffer for storing conversation memory. +type TokenBuffer struct { + Buffer + LLM llms.LanguageModel + MaxTokenLimit int +} + +// Statically assert that TokenBuffer implement the memory interface. +var _ schema.Memory = &TokenBuffer{} + +// NewTokenBuffer is a function for crating a new token buffer memory. +func NewTokenBuffer(llm llms.LanguageModel, maxTokenLimit int, options ...BufferOption) *TokenBuffer { + tb := &TokenBuffer{ + LLM: llm, + MaxTokenLimit: maxTokenLimit, + Buffer: *applyBufferOptions(options...), + } + + return tb +} + +// MemoryVariables uses Buffer method for memory variables. +func (tb *TokenBuffer) MemoryVariables() []string { + return tb.Buffer.MemoryVariables() +} + +// LoadMemoryVariables uses Buffer method for loading memory variables. +func (tb *TokenBuffer) LoadMemoryVariables(inputs map[string]any) (map[string]any, error) { + return tb.Buffer.LoadMemoryVariables(inputs) +} + +// SaveContext uses Buffer method for saving context and prunes memory buffer. +func (tb *TokenBuffer) SaveContext(inputValues map[string]any, outputValues map[string]any) error { + err := tb.Buffer.SaveContext(inputValues, outputValues) + if err != nil { + return err + } + currBufferLength, err := tb.getNumTokensFromMessages() + if err != nil { + return err + } + + if currBufferLength > tb.MaxTokenLimit { + // while currBufferLength is greater than MaxTokenLimit we keep removing messages from the memory + for currBufferLength > tb.MaxTokenLimit { + tb.chatHistory.SetMessages(append(tb.chatHistory.Messages()[:0], tb.chatHistory.Messages()[1:]...)) + currBufferLength, err = tb.getNumTokensFromMessages() + if err != nil { + return err + } + } + } + + return nil +} + +// Clear uses Buffer method for clearing buffer memory. +func (tb *TokenBuffer) Clear() error { + return tb.Buffer.Clear() +} + +func (tb *TokenBuffer) getNumTokensFromMessages() (int, error) { + sum := 0 + for _, message := range tb.chatHistory.Messages() { + bufferString, err := schema.GetBufferString([]schema.ChatMessage{message}, tb.Buffer.HumanPrefix, tb.Buffer.AIPrefix) + if err != nil { + return 0, err + } + + sum += tb.LLM.GetNumTokens(bufferString) + } + + return sum, nil +} diff --git a/memory/token_buffer_test.go b/memory/token_buffer_test.go new file mode 100644 index 00000000..9a5e94b3 --- /dev/null +++ b/memory/token_buffer_test.go @@ -0,0 +1,115 @@ +package memory + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/tmc/langchaingo/llms/openai" + "github.com/tmc/langchaingo/schema" +) + +func TestTokenBufferMemory(t *testing.T) { + t.Parallel() + + llm, err := openai.New() + require.NoError(t, err) + m := NewTokenBuffer(llm, 2000) + + result1, err := m.LoadMemoryVariables(map[string]any{}) + require.NoError(t, err) + expected1 := map[string]any{"history": ""} + assert.Equal(t, expected1, result1) + + err = m.SaveContext(map[string]any{"foo": "bar"}, map[string]any{"bar": "foo"}) + require.NoError(t, err) + + result2, err := m.LoadMemoryVariables(map[string]any{}) + require.NoError(t, err) + + expected2 := map[string]any{"history": "Human: bar\nAI: foo"} + assert.Equal(t, expected2, result2) +} + +func TestTokenBufferMemoryReturnMessage(t *testing.T) { + t.Parallel() + + llm, err := openai.New() + require.NoError(t, err) + m := NewTokenBuffer(llm, 2000, WithReturnMessages(true)) + + expected1 := map[string]any{"history": []schema.ChatMessage{}} + result1, err := m.LoadMemoryVariables(map[string]any{}) + require.NoError(t, err) + assert.Equal(t, expected1, result1) + + err = m.SaveContext(map[string]any{"foo": "bar"}, map[string]any{"bar": "foo"}) + require.NoError(t, err) + + result2, err := m.LoadMemoryVariables(map[string]any{}) + require.NoError(t, err) + + expectedChatHistory := NewChatMessageHistory( + WithPreviousMessages([]schema.ChatMessage{ + schema.HumanChatMessage{Text: "bar"}, + schema.AIChatMessage{Text: "foo"}, + }), + ) + + expected2 := map[string]any{"history": expectedChatHistory.Messages()} + assert.Equal(t, expected2, result2) +} + +func TestTokenBufferMemoryWithPreLoadedHistory(t *testing.T) { + t.Parallel() + + llm, err := openai.New() + require.NoError(t, err) + + m := NewTokenBuffer(llm, 2000, WithChatHistory(NewChatMessageHistory( + WithPreviousMessages([]schema.ChatMessage{ + schema.HumanChatMessage{Text: "bar"}, + schema.AIChatMessage{Text: "foo"}, + }), + ))) + + result, err := m.LoadMemoryVariables(map[string]any{}) + require.NoError(t, err) + expected := map[string]any{"history": "Human: bar\nAI: foo"} + assert.Equal(t, expected, result) +} + +func TestTokenBufferMemoryPrune(t *testing.T) { + t.Parallel() + + llm, err := openai.New() + require.NoError(t, err) + + m := NewTokenBuffer(llm, 20, WithChatHistory(NewChatMessageHistory( + WithPreviousMessages([]schema.ChatMessage{ + schema.HumanChatMessage{Text: "human message test for max token"}, + schema.AIChatMessage{Text: "ai message test for max token"}, + }), + ))) + + buffStringMsg1, err := schema.GetBufferString([]schema.ChatMessage{ + schema.HumanChatMessage{Text: "human message test for max token"}, + }, "Human", "AI") + require.NoError(t, err) + tokenNumMsg1 := m.LLM.GetNumTokens(buffStringMsg1) + assert.Equal(t, 9, tokenNumMsg1) + + buffStringMsg2, err := schema.GetBufferString([]schema.ChatMessage{ + schema.AIChatMessage{Text: "ai message test for max token"}, + }, "Human", "AI") + require.NoError(t, err) + tokenNumMsg2 := m.LLM.GetNumTokens(buffStringMsg2) + assert.Equal(t, 8, tokenNumMsg2) + + assert.Equal(t, tokenNumMsg1+tokenNumMsg2, 17) + + _, err = llm.Call() + require.NoError(t, err) + +} diff --git a/schema/chat_message_history.go b/schema/chat_message_history.go index 1fdde7a0..2507ee42 100644 --- a/schema/chat_message_history.go +++ b/schema/chat_message_history.go @@ -16,4 +16,7 @@ type ChatMessageHistory interface { // Messages get all messages from the store Messages() []ChatMessage + + // SetMessages replaces existing messages in the store + SetMessages(messages []ChatMessage) } diff --git a/schema/memory.go b/schema/memory.go index ebc00983..a6c75e2b 100644 --- a/schema/memory.go +++ b/schema/memory.go @@ -2,12 +2,12 @@ package schema // Memory is the interface for memory in chains. type Memory interface { - // Input keys this memory class will load dynamically. + // MemoryVariables Input keys this memory class will load dynamically. MemoryVariables() []string - // Return key-value pairs given the text input to the chain. + // LoadMemoryVariables Return key-value pairs given the text input to the chain. // If None, return all memories LoadMemoryVariables(inputs map[string]any) (map[string]any, error) - // Save the context of this model run to memory. + // SaveContext Save the context of this model run to memory. SaveContext(inputs map[string]any, outputs map[string]any) error // Clear memory contents. Clear() error