diff --git a/chains/chains_test.go b/chains/chains_test.go index c88fa426..9bb2b353 100644 --- a/chains/chains_test.go +++ b/chains/chains_test.go @@ -8,7 +8,6 @@ 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 57541ee2..cd36960e 100644 --- a/chains/conversation_test.go +++ b/chains/conversation_test.go @@ -7,7 +7,6 @@ import ( "testing" "github.com/stretchr/testify/require" - "github.com/tmc/langchaingo/llms/openai" "github.com/tmc/langchaingo/memory" ) @@ -18,10 +17,10 @@ func TestConversation(t *testing.T) { if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" { t.Skip("OPENAI_API_KEY not set") } - model, err := openai.New() + llm, err := openai.New() require.NoError(t, err) - c := NewConversation(model, memory.NewBuffer()) + c := NewConversation(llm, memory.NewBuffer()) _, err = Run(context.Background(), c, "Hi! I'm Jim") require.NoError(t, err) @@ -31,6 +30,8 @@ func TestConversation(t *testing.T) { } func TestConversationMemoryPrune(t *testing.T) { + t.Parallel() + if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" { t.Skip("OPENAI_API_KEY not set") } @@ -38,12 +39,16 @@ func TestConversationMemoryPrune(t *testing.T) { llm, err := openai.New() require.NoError(t, err) - c := NewConversation(llm, memory.NewTokenBuffer(llm, 100, memory.WithReturnMessages(true))) + c := NewConversation(llm, memory.NewTokenBuffer(llm, 50)) _, 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'`) + require.True(t, strings.Contains(res, "Jim"), `result does contain the keyword 'Jim'`) + // this message will hit the maxTokenLimit and will initiate the prune of the messages to fit the context + res, err = Run(context.Background(), c, "Are you sure that my name is Jim?") + require.NoError(t, err) + require.True(t, strings.Contains(res, "Jim"), `result does contain the keyword 'Jim'`) } diff --git a/memory/buffer_test.go b/memory/buffer_test.go index a3b641ac..ce1a43e2 100644 --- a/memory/buffer_test.go +++ b/memory/buffer_test.go @@ -5,7 +5,6 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/tmc/langchaingo/schema" ) @@ -87,7 +86,7 @@ func (t testChatMessageHistory) AddMessage(_ schema.ChatMessage) { func (t testChatMessageHistory) Clear() { } -func (t testChatMessageHistory) SetMessages(messages []schema.ChatMessage) { +func (t testChatMessageHistory) SetMessages(_ []schema.ChatMessage) { } func (t testChatMessageHistory) Messages() []schema.ChatMessage { diff --git a/memory/token_buffer.go b/memory/token_buffer.go index d3a558ad..fcd9b1b2 100644 --- a/memory/token_buffer.go +++ b/memory/token_buffer.go @@ -36,7 +36,7 @@ func (tb *TokenBuffer) LoadMemoryVariables(inputs map[string]any) (map[string]an return tb.Buffer.LoadMemoryVariables(inputs) } -// SaveContext uses Buffer method for saving context and prunes memory buffer. +// SaveContext uses Buffer method for saving context and prunes memory buffer if needed. func (tb *TokenBuffer) SaveContext(inputValues map[string]any, outputValues map[string]any) error { err := tb.Buffer.SaveContext(inputValues, outputValues) if err != nil { @@ -49,8 +49,9 @@ func (tb *TokenBuffer) SaveContext(inputValues map[string]any, outputValues map[ if currBufferLength > tb.MaxTokenLimit { // while currBufferLength is greater than MaxTokenLimit we keep removing messages from the memory + // from the oldest for currBufferLength > tb.MaxTokenLimit { - tb.chatHistory.SetMessages(append(tb.chatHistory.Messages()[:0], tb.chatHistory.Messages()[1:]...)) + tb.ChatHistory.SetMessages(append(tb.ChatHistory.Messages()[:0], tb.ChatHistory.Messages()[1:]...)) currBufferLength, err = tb.getNumTokensFromMessages() if err != nil { return err @@ -68,7 +69,7 @@ func (tb *TokenBuffer) Clear() error { func (tb *TokenBuffer) getNumTokensFromMessages() (int, error) { sum := 0 - for _, message := range tb.chatHistory.Messages() { + 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 diff --git a/memory/token_buffer_test.go b/memory/token_buffer_test.go index 9a5e94b3..aed00791 100644 --- a/memory/token_buffer_test.go +++ b/memory/token_buffer_test.go @@ -5,7 +5,6 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/tmc/langchaingo/llms/openai" "github.com/tmc/langchaingo/schema" ) @@ -79,37 +78,3 @@ func TestTokenBufferMemoryWithPreLoadedHistory(t *testing.T) { 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) - -}