// Copyright (c) 2026 Probo Inc . // // Permission is hereby granted, free of charge, to any person obtaining a copy // of this software and associated documentation files (the "Software"), to deal // in the Software without restriction, including without limitation the rights // to use, copy, modify, merge, publish, distribute, sublicense, and/or sell // copies of the Software, and to permit persons to whom the Software is // furnished to do so, subject to the following conditions: // // The above copyright notice and this permission notice shall be included in // all copies or substantial portions of the Software. // // THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR // IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, // FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE // AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER // LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, // OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE // SOFTWARE. package agent_test import ( "context" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.probo.inc/probo/pkg/agent" "go.probo.inc/probo/pkg/llm" ) func TestNewMemorySession(t *testing.T) { t.Parallel() t.Run( "returns non-nil session", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() assert.NotNil(t, s) }, ) } func TestMemorySession_Load(t *testing.T) { t.Parallel() t.Run( "returns nil for unknown session ID", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() msgs, err := s.Load(context.Background(), "unknown") require.NoError(t, err) assert.Nil(t, msgs) }, ) t.Run( "returns saved messages", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() messages := []llm.Message{ { Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "hello"}}, }, { Role: llm.RoleAssistant, Parts: []llm.Part{llm.TextPart{Text: "hi there"}}, }, } err := s.Save(context.Background(), "sess-1", messages) require.NoError(t, err) loaded, err := s.Load(context.Background(), "sess-1") require.NoError(t, err) require.Len(t, loaded, 2) assert.Equal(t, llm.RoleUser, loaded[0].Role) assert.Equal(t, "hello", loaded[0].Text()) assert.Equal(t, llm.RoleAssistant, loaded[1].Role) assert.Equal(t, "hi there", loaded[1].Text()) }, ) t.Run( "returns a defensive copy of messages", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() messages := []llm.Message{ { Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "original"}}, }, } err := s.Save(context.Background(), "sess-copy", messages) require.NoError(t, err) loaded, err := s.Load(context.Background(), "sess-copy") require.NoError(t, err) loaded[0].Parts = []llm.Part{llm.TextPart{Text: "mutated"}} reloaded, err := s.Load(context.Background(), "sess-copy") require.NoError(t, err) assert.Equal(t, "original", reloaded[0].Text()) }, ) t.Run( "different session IDs are independent", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() err := s.Save( context.Background(), "sess-a", []llm.Message{ {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "alpha"}}}, }, ) require.NoError(t, err) err = s.Save( context.Background(), "sess-b", []llm.Message{ {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "beta"}}}, }, ) require.NoError(t, err) a, err := s.Load(context.Background(), "sess-a") require.NoError(t, err) require.Len(t, a, 1) assert.Equal(t, "alpha", a[0].Text()) b, err := s.Load(context.Background(), "sess-b") require.NoError(t, err) require.Len(t, b, 1) assert.Equal(t, "beta", b[0].Text()) }, ) } func TestMemorySession_Save(t *testing.T) { t.Parallel() t.Run( "stores a defensive copy of input messages", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() messages := []llm.Message{ { Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "before"}}, }, } err := s.Save(context.Background(), "sess-def", messages) require.NoError(t, err) messages[0].Parts = []llm.Part{llm.TextPart{Text: "after"}} loaded, err := s.Load(context.Background(), "sess-def") require.NoError(t, err) assert.Equal(t, "before", loaded[0].Text()) }, ) t.Run( "overwrites previous messages for same session ID", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() err := s.Save( context.Background(), "sess-ow", []llm.Message{ {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "first"}}}, }, ) require.NoError(t, err) err = s.Save( context.Background(), "sess-ow", []llm.Message{ {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "second"}}}, {Role: llm.RoleAssistant, Parts: []llm.Part{llm.TextPart{Text: "reply"}}}, }, ) require.NoError(t, err) loaded, err := s.Load(context.Background(), "sess-ow") require.NoError(t, err) require.Len(t, loaded, 2) assert.Equal(t, "second", loaded[0].Text()) assert.Equal(t, "reply", loaded[1].Text()) }, ) t.Run( "preserves tool calls and tool call ID", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() messages := []llm.Message{ { Role: llm.RoleAssistant, ToolCalls: []llm.ToolCall{ { ID: "call-1", Function: llm.FunctionCall{ Name: "get_weather", Arguments: `{"city":"Paris"}`, }, }, }, }, { Role: llm.RoleTool, ToolCallID: "call-1", Parts: []llm.Part{llm.TextPart{Text: "sunny"}}, }, } err := s.Save(context.Background(), "sess-tc", messages) require.NoError(t, err) loaded, err := s.Load(context.Background(), "sess-tc") require.NoError(t, err) require.Len(t, loaded, 2) require.Len(t, loaded[0].ToolCalls, 1) assert.Equal(t, "call-1", loaded[0].ToolCalls[0].ID) assert.Equal(t, "get_weather", loaded[0].ToolCalls[0].Function.Name) assert.Equal(t, `{"city":"Paris"}`, loaded[0].ToolCalls[0].Function.Arguments) assert.Equal(t, "call-1", loaded[1].ToolCallID) assert.Equal(t, "sunny", loaded[1].Text()) }, ) t.Run( "handles empty message slice", func(t *testing.T) { t.Parallel() s := agent.NewMemorySession() err := s.Save(context.Background(), "sess-empty", []llm.Message{}) require.NoError(t, err) loaded, err := s.Load(context.Background(), "sess-empty") require.NoError(t, err) assert.Empty(t, loaded) }, ) }